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)。收到權重更新時:

  1. 從 bucket 下載 delta safetensors 文件。
  2. 對於 anchors:加載所有 tensor 並為未來的 deltas 進行快照。
  3. 對於 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_weights API 時,後者將被移除。
  • 固定的 anchor 週期: 目前每 N 步生成一個 anchor;採用自適應策略(當累積漂移超過閾值時生成 anchor)可能會降低成本。
  • 多節點 FSDP2 訓練器: BF16ChangeDetector 是為單進程優化器 hook 構建的;多節點 FSDP2 的支持尚未經測試。
  • 掛載到優化器: 由於複雜的交互作用,從 Adam 統計數據預測掩碼仍然具有挑戰性。
  • 與傳輸壓縮疊加: 稀疏 safetensors 和逐分塊 gzip 是正交的,但尚未結合。

嘗試使用

Sources