TRL의 Delta Weight Sync를 통한 최소 대역폭 기반 조 단위 파라미터 모델 학습
TRL의 Delta Weight Sync를 통한 최소 대역폭 기반 조 단위 파라미터 모델 학습
1 Terabyte 문제
비동기(Async) RL 학습에서는 정책(policy)을 동기화하기 위해 매 단계마다 트레이너에서 추론 엔진으로 전체 모델을 전송해야 합니다. bf16 기반의 7B 모델의 경우, 단계당 14 GB가 소요됩니다. 1T 파라미터 규모의 프런티어 모델의 경우, 단계당 약 1 TB가 소요됩니다. 이 전송 과정은 크리티컬 패스(critical path)에 위치하여, GPU가 토큰을 생성하지 못하고 유휴 상태로 대기하는 시간을 발생시킵니다.
bf16 RL 가중치가 거의 항상 희소(Sparse)한 이유
연속적인 RL 옵티마이저 단계 사이에서 bf16 가중치의 약 99%는 비트 단위로 동일하게 유지됩니다(최악의 경우에도 98% 이상). 이는 bf16의 정밀도가 제한적이기 때문입니다. 업데이트 크기가 가중치 주변의 표현 가능한 값 사이 간격의 절반보다 작으면 업데이트가 반올림되어 흡수됩니다. 일반적인 RL 학습률(예: 3×10⁻⁶)에서 대부분의 가중치에 대한 업데이트 크기는 이 임계값보다 작으므로, bf16 표현은 변하지 않습니다. 이러한 희소성은 운이 좋은 측정이 아니라 산술적으로 보장되는 특성입니다.
HF Buckets와 아키텍처
Bucket이란 무엇인가?
Bucket은 고빈도 객체 스토리지(high-frequency object storage)를 위한 Hugging Face Hub의 리포지토리 유형입니다. 별도의 커밋 절차나 PR 워크플로우가 필요하지 않습니다. 파일은 batch_bucket_files(업로드용) 및 download_bucket_files(다운로드용)라는 두 가지 함수를 통해 추가, 목록 확인 또는 다운로드할 수 있습니다. 내부적으로 Bucket은 콘텐츠 기반 청킹(content-defined chunking) 스토리지 레이어인 Xet를 사용하여 콘텐츠를 기반으로 청크를 중복 제거합니다.
세 가지 구성 요소
아키텍처는 세 가지 구성 요소와 하나의 공유 기질로 구성됩니다:
- Trainer: 모델 가중치를 소유하고, 옵티마이저를 실행하며, 희소 델타(sparse deltas)를 생성합니다(단일 GPU, 다중 GPU 또는 노트북 등 어디든 위치 가능).
- HF Bucket: 전체 스냅샷을 위한
anchors/와 희소 패치를 위한deltas/를 포함하는 단일 리포지토리이며, 양측이 공유하는 유일한 매개체입니다. - vLLM rollout server: 버킷에서 데이터를 가져와 델타를 적용하고 롤아웃을 서비스합니다(트레이너와 반드시 같은 위치에 있을 필요는 없음).
- Environment: HTTP 또는 함수 호출을 통해 롤아웃 서버에 연결됩니다.
트레이너와 롤아웃 서버는 가중치 데이터를 직접 교환하지 않으며, 버킷 좌표가 포함된 아주 작은 POST 요청만 공유합니다. 모든 바이트 전송은 각 측과 버킷 사이에서 병렬로 이루어집니다.
프로토콜
전송 포맷으로서의 Safetensors
저장 및 전송 포맷으로 safetensors를 사용합니다. 버킷에는 두 가지 파일 유형이 존재합니다:
- Anchors: 전체 bf16 가중치를 가진 일반 체크포인트(N 단계마다 기록, 기본 N=10).
- Deltas: 변경된 각 파라미터에 대해, 요소 인덱스의 int32 텐서와 해당 인덱스의 값에 대한 bf16 텐서를 저장합니다.
메타데이터는 파일이 희소(sparse)한지 또는 앵커(anchor)인지 나타내어 수신 측에서 적절하게 분기할 수 있게 합니다.
Trainer 측: 옵티마이저 훅을 통한 불리언 마스크
BF16ChangeDetector는 옵티마이저에 pre-step 및 post-step 훅을 등록하여 단계 전후의 가중치를 bf16로 스냅샷합니다. 변경된 요소의 불리언 마스크는 이 스냅샷들을 비교하여 계산됩니다. Adam 통계로부터 마스크를 예측하는 방식은 재현율(recall)이 낮았기 때문에(~30%), 이 ground-truth 방식을 사용합니다.
vLLM 측: 30줄의 확장 기능
--worker-extension-cls 플래그를 통해 vLLM에 플러그인되는 DeltaWeightTransferEngine을 구현했습니다(포크가 필요 없음). 가중치 업데이트를 수신하면:
- 버킷에서 델타 safetensors 파일을 다운로드합니다.
- 앵커의 경우: 모든 텐서를 로드하고 향후 델타를 위해 스냅샷을 찍습니다.
- 델타의 경우: 변경된 각 파라미터에 대해 인덱스와 값을 가져와 로컬 bf16 스냅샷에 적용한 후, 재구성된 전체 텐서를 vLLM의
load_weights에 전달합니다.
실제 Spaces에서의 구현
공유 네트워크가 없는 완전한 분리형(disaggregated) 학습을 실행했습니다:
- Trainer: 단일 GPU 박스.
- vLLM rollout server: 확장이 설치된 Hugging Face Space (Docker SDK, L4 GPU).
- Wordle environment: 256개의 동시 세션 용량을 가진 두 번째 Hugging Face Space (CPU).
- Hub bucket: 가중치 델타 및 앵커를 위한 중앙 리포지토리.
설정에는 몇 가지 hf CLI 호출이 포함되었습니다. vLLM Space Dockerfile은 delta-weight-sync 브랜치에서 TRL을 설치하고 워커 확장 클래스를 설정합니다. 학습은 Spaces 및 버킷에 HTTPS로 접근할 수 있는 곳이라면 어디서든 시작할 수 있습니다.
이것이 실제로 가능하게 하는 것은 무엇인가?
- 클러스터 없는 비동기 RL 학습: 단일 GPU 트레이너가 Spaces를 롤아웃 및 환경으로 사용할 수 있으며, 가중치는 버킷을 통해 이동합니다.
- 무료로 제공되는 다중 복제본 추론: 여러 vLLM Spaces가 동일한 버킷에서 데이터를 가져옵니다. Xet는 저장된 청크를 중복 제거하며, Hub의 엣지 캐시는 반복된 다운로드를 저렴하게 처리합니다.
- 디버깅 가능한 전송 포맷: 델타는 Python의
safe_open으로 검사할 수 있는 safetensors 파일입니다. - 프런티어 규모로 가는 경로: Qwen3-0.6B 모델의 경우 단계당 페이로드가 1.2 GB에서 20–35 MB로 급감합니다. Llama-3.1-405B 모델(bf16 기준 810 GB)의 경우, 단순 계산상 단계당 약 6 GB의 델타가 발생하여(810 GB 전체 대비), 추론 일시 중지 시간이 ~8초(100 GB/s NCCL 기준)에서 단 몇 초로 단축됩니다. 1 GB/s 대역폭의 클라우드 간 통신 시, 전체 브로드캐스트에는 13분이 걸리지만 델타는 6초면 충분합니다.
남은 과제
- 두 개의 CPU bf16 스냅샷: 트레이너는 변경 감지를 위해 하나를 유지하고, 롤아웃 서버는 vLLM의
load_weights를 위한 전체 텐서 재구성을 위해 하나를 유지합니다. 후자는 vLLM에 희소load_weightsAPI가 도입되면 제거될 예정입니다. - 고정된 앵커 주기: 현재는 N 단계마다 앵커를 생성합니다. 누적 드리프트가 임계값을 초과할 때 앵커를 생성하는 적응형 정책을 통해 비용을 줄일 수 있습니다.
- 다중 노드 FSDP2 트레이너:
BF16ChangeDetector는 단일 프로세스 옵티마이저 훅을 위해 설계되었습니다. 다중 노드 FSDP2 지원은 아직 측정되지 않았습니다. - 옵티마이저 훅킹: 복잡한 상호작용으로 인해 Adam 통계로부터 마스크를 예측하는 것은 여전히 어려운 과제입니다.
- 전송 시 압축 결합: 희소 safetensors와 청크별 gzip은 서로 독립적이며 아직 결합되지 않았습니다.
시도해보기
- PR: huggingface/trl#5417 (branch:
delta-weight-sync). - 전체 Wordle 예제:
examples/scripts/openenv/async_wordle.py. - Spaces Dockerfiles:
examples/scripts/openenv/vllm_space/및examples/scripts/openenv/wordle_space/. - 배경 지식: 우리의 async RL landscape post, Fireworks 1 TB post, Cursor Composer 2 report.