Hugging Face Jobs를 통한 LoRA 기반 Async GRPO

Hugging Face는 AsyncGRPOTrainer(TRL v1.14)에 LoRA 지원을 구현하여, 트레이너가 전체 모델 가중치가 아닌 작은 LoRA 어댑터만 vLLM에 동기화할 수 있도록 했습니다. Hugging Face Storage Buckets를 공유 파일 시스템으로 활용하고 요청 라우팅을 위한 커스텀 프록시를 사용하여, 이 아키텍처는 NCCL이나 공유 로컬 디스크 없이도 학습과 추론이 서로 다른 머신(Hugging Face Jobs)에서 실행되도록 합니다.

아키텍처: Storage Buckets를 통한 분산 동기화

AsyncGRPOTrainer는 이제 텐서를 vLLM으로 직접 전송할 필요가 없는 어댑터 전용 동기화 경로를 지원합니다. 대신, 트레이너는 LoRA 어댑터를 Storage Bucket의 특정 디렉터리에 저장하고, 원자적 이름 변경(atomic rename)을 수행한 후 /v1/load_lora_adapter 엔드포인트를 통해 vLLM에 알림을 보냅니다.

Hugging Face Jobs는 hf-mount를 사용하여 Storage Buckets를 FUSE 파일 시스템으로 마운트할 수 있으므로, 트레이너와 vLLM 리플리카는 서로 다른 VM에서 동일한 절대 경로를 공유할 수 있습니다. 이를 통해 트레이너와 추론 서버가 물리적 노드나 고밀도 클러스터 네트워크를 공유해야 하는 요구 사항이 제거됩니다.

Job 레이아웃

시스템은 세 가지 주요 구성 요소로 이루어져 있습니다:

  • Trainer Job: LoRA와 FSDP를 사용하여 AsyncGRPOTrainer를 실행합니다.
  • vLLM Jobs: 베이스 모델을 서빙하고 버킷에서 최신 어댑터를 로드하는 여러 리플리카입니다.
  • Proxy Server: 인증 헤더를 처리하고 요청을 vLLM 리플리카로 라우팅하는 트레이너 Job에서 실행되는 작은 asyncio 기반 프록시입니다.

프록시: KV-Prefix 라우팅 및 브로드캐스팅

여러 vLLM 리플리카 간 효율성을 극대화하기 위해, 요청 분배와 상태 동기화를 관리하는 커스텀 프록시가 사용됩니다.

KV-Prefix 기반 라우팅

중복된 prefill 계산을 피하기 위해, 프록시는 KV 캐시 접두사(prefix)를 기반으로 요청을 라우팅합니다. 프롬프트를 16-token 블록으로 나누고 어댑터 이름을 시드로 사용하여 체인 해시를 계산합니다.

라우터는 어떤 리플리카가 어떤 블록 해시를 서빙했는지 추적합니다. 요청의 프롬프트가 특정 리플리카에 이미 캐시된 접두사와 일치하고(그 리플리카가 과부하 상태가 아니면), 요청은 해당 리플리카로 라우팅됩니다("affinity hit"). 이는 GRPO에서 단일 프롬프트에 대해 여러 완성본이 생성되므로, 여러 롤아웃에 걸쳐 동일한 프롬프트의 prefill을 다시 계산하는 것을 방지하는 데 중요합니다.

상태 브로드캐스팅

각 vLLM 리플리카는 별도의 Job이므로, 프록시는 어댑터 로드, 일시 중지, 재개와 같은 상태 변경 요청을 모든 리플리카로 브로드캐스트하여 일관성을 유지합니다. 이를 통해 특정 정책 버전 이름이 전체 플릿(fleet)에서 동일한 가중치를 가리키도록 보장합니다.

성능 최적화 및 병목 분석

sail/Sanity-Test-R1D-1.5B 데이터셋과 Qwen/Qwen2.5-Math-1.5B 모델을 사용하여 Hugging Face는 파이프라인을 최적화하기 위해 5회의 실험을 수행했습니다. 그 결과, Async RL의 병목 현상이 학습과 생성 사이에서 이동할 수 있음을 보여줍니다.

주요 최적화 사항

  1. Token-Budget Batching: 디바이스당 학습 배치 크기를 1에서 token-budget batching(예: token_budget=16384)으로 전환하여 여러 시퀀스를 각 행에 패킹하고 마이크로 배치의 수를 줄임으로써 MFU를 3.9%에서 19%로 향상시켰습니다.
  2. Gradient Checkpointing 비활성화: 작은 모델(1.5B)의 경우, gradient checkpointing을 비활성화하면 중복된 forward 패스를 제거하여 forward+backward 시간을 줄이고, 병목 현상을 트레이너에서 생성 리플리카로 이동시킵니다.
  3. In-Flight Requests 증가: max_inflight_tasks를 높임(예: 384으로)으로써 시스템이 여러 vLLM 리플리카를 완전히 활용할 수 있게 하여, 클라이언트 측 동시성 제한이 처리량을 저하시키는 것을 방지했습니다.

최종 결과

이러한 최적화들을 결합하여 500단계 완성에 걸리는 총 시간은 3시간 27분에서 53분으로 줄었습니다(3.9배 속도 향상).

지표 Run 1 (베이스라인) Run 5 (최적화됨)
Wall Clock Time 3 h 27 min 53 min
Median Step Time 22.9 s 4.8 s
MFU (Fwd/Bwd) 3.9% 23.5%
Samples Trained 64,000 84,078
Mean Staleness 1.5 versions 2.0 versions

기술 구현 세부 사항

  • vLLM 버전: 런타임 LoRA 엔드포인트와의 호환성을 위해 v0.27.1로 고정되었습니다.
  • Adapter Slots: max_staleness=4를 지원하기 위해, 스왑 동안 현재 정책과 이전 버전이 로드된 상태로 유지되도록 vLLM이 --max-loras 6으로 구성됩니다.
  • 일관성: 버전화된 어댑터 이름이 사용되어 KV 캐시가 이전 정책 버전으로 생성된 접두사와 잘못 일치하는 것을 방지합니다.

Sources