Async GRPO with LoRA across Hugging Face Jobs

Hugging Face 已在 AsyncGRPOTrainer (TRL v1.14) 中實現了 LoRA 支援,允許訓練器僅將小型 LoRA adapter 同步到 vLLM,而非完整的模型權重。透過利用 Hugging Face Storage Buckets 作為共享檔案系統,並使用自定義代理伺服器進行請求路由,此架構使得訓練與推理可以在不同的機器(Hugging Face Jobs)上運行,而無需 NCCL 或共享的本地磁碟。

架構:透過 Storage Buckets 進行分散式同步

AsyncGRPOTrainer 現在支援僅針對 adapter 的同步路徑,這消除了直接將 tensor 送往 vLLM 的需求。相反地,訓練器會將 LoRA adapter 儲存到 Storage Bucket 中的特定目錄,執行原子重新命名,並透過 /v1/load_lora_adapter 端點通知 vLLM。

由於 Hugging Face Jobs 可以使用 hf-mount 將 Storage Buckets 掛載為 FUSE 檔案系統,因此訓練器與 vLLM 副本(replicas)可以在不同的 VM 上共享相同的絕對路徑。這消除了訓練器與推理伺服器需要共享物理節點或密集叢集網路的要求。

Job Layout

系統由三個主要組件組成:

  • Trainer Job: 運行 AsyncGRPOTrainer 並搭配 LoRA 與 FSDP。
  • vLLM Jobs: 多個副本(replicas)提供基礎模型並從 bucket 中載入最新的 adapters。
  • Proxy Server: 運行在 trainer Job 上的一個小型 asyncio-based 代理伺服器,負責處理身份驗證標頭並將請求路由至 vLLM 副本。

代理伺服器:KV-Prefix 路由與廣播

為了在多個 vLLM 副本之間實現效率最大化,使用了一個自定義代理伺服器來管理請求分配與狀態同步。

透過 KV-Prefix 進行路由

為了避免冗餘的 prefill 計算,代理伺服器會根據 KV cache prefix 進行請求路由。它將 prompt 分成 16-token blocks,並計算以 adapter name 為種子的鏈式雜湊(chained hashes)。

路由器會追蹤哪個副本(replica)提供了哪些 block hash。如果請求的 prompt 與特定副本上已快取(cached)的 prefix 匹配(且該副本未超載),則請求會被路由至該處(即「親和性命中」,affinity hit)。這能防止系統在多次 rollout 中對同一個 prompt 重新計算 prefill,這對於在單一 prompt 下生成多個 completions 的 GRPO 至關重要。

狀態廣播

由於每個 vLLM 副本都是一個獨立的 Job,代理伺服器透過廣播狀態變更請求(例如 adapter 載入、暫停與恢復)來確保一致性。這確保了特定的 policy version name 會在整個集群中指向相同的權重。

效能優化與瓶頸分析

使用 sail/Sanity-Test-R1D-1.5B 資料集與 Qwen/Qwen2.5-Math-1.5B 模型,Hugging Face 進行了五次實驗運行,以優化管線(pipeline)。結果顯示,Async RL 的瓶頸可能會在訓練與生成之間切換。

關鍵優化項目

  1. Token-Budget Batching: 從每個設備的 train batch size 為 1 轉向 token-budget batching(例如 token_budget=16384),透過將多個序列打包進每個 row,減少了 microbatches 的數量,使 MFU 從 3.9% 提升至 19%。
  2. 停用 Gradient Checkpointing: 對於較小的模型(1.5B),停用 gradient checkpointing 減少了 forward+backward 的時間,透過消除冗餘的 forward passes,將瓶頸從訓練器轉移至生成副本。
  3. 增加 In-Flight Requests: 提高 max_inflight_tasks(例如提高至 384)讓系統能充分利用多個 vLLM 副本,防止客戶端併發限制(concurrency limit)限制了吞吐量。

最終結果

透過結合這些優化,完成 500 個 steps 的總時間從 3 小時 27 分鐘減少至 53 分鐘(提升了 3.9 倍速度)。

Metric Run 1 (Baseline) Run 5 (Optimized)
Wall Clock Time 3 h 27 min 53 min
Median Step Time 22.9 s 4.8 s
Samples Trained 64,000 84,078
MFU (Fwd/Bwd) 3.9% 23.5%
Mean Staleness 1.5 versions 2.0 versions

技術實作細節

  • vLLM Version: 固定於 v0.27.1 以確保與 runtime LoRA 端點的相容性。
  • Adapter Slots: 為了支援 max_staleness=4,vLLM 配置了 --max-loras 6,以確保在切換期間,當前 policy 與之前的版本仍保持載入狀態。
  • Consistency: 使用版本化的 adapter names 以防止 KV cache 錯誤地匹配到由舊版 policy version 產生的 prefix。

Sources