TRL 中的 Delta Weight Sync 實現極低頻寬的兆級參數模型訓練
TRL 中的 Delta Weight Sync 實現極低頻寬的兆級參數模型訓練
一 TB 的問題
非同步 RL 訓練需要每一步都將整個模型從訓練器(trainer)發送到推理引擎(inference engine)以保持策略同步。對於一個 bf16 格式的 7B 模型,這意味著每一步需要 14 GB。對於一個前沿的 1T 參數模型,每一步大約需要 1 TB。這種傳輸位於關鍵路徑上,導致 GPU 在不生成 token 的情況下出現計算閒置。
為什麼 bf16 RL 權重幾乎總是稀疏的
在連續的 RL 優化器步驟之間,大約 99% 的 bf16 權重保持位元完全相同(在最壞情況下也不低於 98%)。這是因為 bf16 的精度有限:如果更新的量級低於權重周圍可表示值之間間距的一半,更新就會被捨入吸收。在典型的 RL 學習率下(例如 3×10⁻⁶),大多數權重的更新大小都小於這個閾值,因此 bf16 表示形式不會改變。這種稀疏性是由算術規律保證的,而非偶然的測量結果。
HF Buckets 與架構
什麼是 Bucket?
Bucket 是 Hugging Face Hub 上的一種用於高頻對象存儲的 repo 類型。它不需要 commit 儀式或 PR 工作流。文件通過兩個函數進行添加、列出或下載:用於上傳的 batch_bucket_files 和用於下載的 download_bucket_files。在底層,Buckets 使用 Xet,這是 Hub 的內容定義分塊(content-defined chunking)存儲層,它根據內容對分塊進行去重。
三個方塊
該架構由三個組件和一個共享基層組成:
- Trainer: 擁有模型權重,運行優化器,並發送稀疏的 deltas(可以位於任何地方:單個 GPU、多個 GPU 或筆記型電腦)。
- HF Bucket: 一個包含
anchors/(用於完整快照)和deltas/(用於稀疏補丁)的單一 repo;這是雙方達成共識的唯一媒介。 - vLLM rollout server: 從 bucket 拉取數據,應用 deltas,並提供 rollouts(不一定與 trainer 位於同一位置)。
- Environment: 通過 HTTP 或函數調用連接到 rollout server。
Trainer 和 rollout server 從不直接交換權重數據;它們只共享一個帶有 bucket 座標的微小 POST 請求。所有的數據傳輸都發生在每一方與 bucket 之間,且是並行的。
協議
使用 safetensors 作為傳輸格式
我們使用 safetensors 作為磁碟和傳輸格式。Bucket 中存在兩種文件類型:
- Anchors: 具有完整 bf16 權重的正常檢查點(每 N 步寫入一次,默認 N=10)。
- Deltas: 對於每個變化的參數,存儲一個包含元素索引的 int32 tensor 和一個包含這些索引對應值的 bf16 tensor。
元數據(Metadata)指示文件是稀疏的還是 anchor,從而使接收方能夠進行相應的分支處理。
Trainer 端:來自優化器 Hook 的布林掩碼 (Boolean Mask)
BF16ChangeDetector 在優化器上註冊 pre-step 和 post-step hooks,以便在步驟前後對 bf16 權重進行快照。通過比較這些快照來計算變更元素的布林掩碼。使用這種地面真值(ground-truth)方法是因為從 Adam 統計數據預測掩碼的召回率(recall)很低(約 30%)。
vLLM 端:一個 30 行的擴展
我們實現了一個 DeltaWeightTransferEngine,通過 --worker-extension-cls 標誌插入 vLLM(無需 fork)。收到權重更新時:
- 從 bucket 下載 delta safetensors 文件。
- 對於 anchors:加載所有 tensor 並為未來的 deltas 進行快照。
- 對於 deltas:對於每個變化的參數,檢索索引和值,將其應用於本地 bf16 快照,並將重建的完整 tensor 饋送給 vLLM 的
load_weights。
在 Spaces 上實際運行
我們進行了一次完全解耦的訓練,沒有共享網絡:
- Trainer: 一個 GPU 節點。
- vLLM rollout server: 安裝了我們擴展的 Hugging Face Space (Docker SDK, L4 GPU)。
- Wordle environment: 第二個 Hugging Face Space (CPU),具有 256 個並發會話能力。
- Hub bucket: 用於權重 deltas 和 anchors 的中央 repo。
設置僅涉及幾次 hf CLI 調用。vLLM Space 的 Dockerfile 從 delta-weight-sync 分支安裝 TRL 並設置 worker extension class。訓練可以從任何可以通過 HTTPS 訪問 Spaces 和 bucket 的地方啟動。
這究竟解鎖了什麼?
- 無需集群的非同步 RL 訓練: 單個 GPU trainer 可以使用 Spaces 作為 rollout 和環境,權重通過 bucket 傳輸。
- 免費的多副本推理: 多個 vLLM Spaces 從同一個 bucket 拉取數據;Xet 對存儲的分塊進行去重,且 Hub 的邊緣緩存可以廉價地提供重複下載。
- 可調試的傳輸格式: Deltas 是可以使用 Python 中的
safe_open檢查的 safetensors 文件。 - 通往前沿規模之路: 對於 Qwen3-0.6B 模型,每一步的負載從 1.2 GB 降至 20–35 MB。對於 Llama-3.1-405B 模型(bf16 下為 810 GB),根據簡單計算,每一步的 deltas 約為 6 GB(對比 810 GB 全量),將推理暫停時間從約 8 秒(使用 100 GB/s NCCL)減少到幾秒鐘。在 1 GB/s 頻寬的跨雲環境下,全量廣播需要 13 分鐘;而 delta 僅需 6 秒。
我們還需解決的問題
- 兩個 CPU bf16 快照: Trainer 保留一個用於變化檢測;rollout server 保留一個用於為 vLLM 的
load_weights重建完整 tensor。當 vLLM 獲得稀疏load_weightsAPI 時,後者將被移除。 - 固定的 anchor 週期: 目前每 N 步生成一個 anchor;採用自適應策略(當累積漂移超過閾值時生成 anchor)可能會降低成本。
- 多節點 FSDP2 訓練器:
BF16ChangeDetector是為單進程優化器 hook 構建的;多節點 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。