GPT-OSS 代理式強化學習訓練:實務回顧
TL;DR
Hugging Face 與 LinkedIn 的研究人員成功地透過解決 PPO on-policy 完整性、attention‑sink 反向傳播以及 MoE 記憶體實體化等關鍵不穩定性,為 GPT-OSS 模型開啟了代理式強化學習(RL)。這些修正使 GPT-OSS 能夠作為與環境與工具互動的多步決策代理的穩定骨幹。
代理式強化學習與 GPT-OSS
代理式強化學習與傳統的單回合 RL 不同,它會優化整個決策過程。模型不僅產生靜態回應,而是學會規劃行動、呼叫工具,並在多步軌跡中調整行為。這需要一個閉環系統,讓代理收集 rollout 軌跡、計算獎勵,並使用 PPO 或 GRPO 等演算法迭代更新其策略。
雖然 GPT-OSS 的表現已與 OpenAI o3-mini 與 o4-mini 相當,但其在代理式 RL 中的適用性尚未驗證。研究人員使用 verl 訓練框架,並在包括 GSM8K、Retool(代理式程式編寫任務)以及可驗證指令遵循等任務上測試了 GPT-OSS-20B 模型。
解決 PPO On-Policy 不穩定性
最初的訓練執行出現 KL 散度與熵指數爆炸,且獎勵未提升。團隊發現這是由於混合專家(Mixture of Experts,MoE)架構導致的 Proximal Policy Optimization(PPO)on-policy 完整性失效。
MoE 對數機率不匹配
在純 on-policy PPO 中,重要性抽樣比率必須恰好為 1。然而,在類似 GPT-OSS 的 MoE 架構中,門控網路可能因浮點差異或隨機性,在產生 rollout 的前向傳播與訓練的前向傳播之間將輸入路由至不同的專家。這會導致當前對數機率與舊對數機率不匹配,錯誤觸發 PPO clip,違反 on-policy 假設。
修正方法: 團隊實作了對數機率替換,當環境已知為 on-policy(minibatch 大小等於全局 batch 大小)時覆寫計算,透過設定 old_log_prob = log_prob.detach() 將重要性比率強制回到 1。
透過 Attention Sinks 校正訓練與推論不匹配
即使在修正 PPO 完整性後,梯度範數仍持續爆炸。研究人員發現了一個根本性的訓練與推論不匹配問題:推論引擎(SGLang)與訓練堆疊(使用 FlashAttention‑v2 的 FSDP)產生了不同的 token 級別機率。
Attention Sinks 的角色
GPT-OSS 使用 attention sinks——可學習的標量參數,在 softmax 計算中充當「虛擬 token」。這些 sink 讓模型將注意力質量分配給學習到的參數,而不是強迫其落在內容 token 上,從而提升串流推論的穩定性。
在 FlashAttention v3 中的實作
團隊發現 verl 硬編碼了 FlashAttention v2,該版本不支援 attention sinks,且 v2 與 v3 都未支援 sink 梯度所需的反向傳播。為了解決此問題,他們:
- 利用 vLLM FlashAttention 分支的前向傳播。
- 實作反向傳播以計算 sink 梯度 $\frac{\partial L}{\partial S_{h}}$。
結果: 此修正使得在 GSM8K、VerifyIf 與 Retool 任務上收斂速度顯著提升,且獎勵穩定改善。
為長上下文擴展記憶體效能
代理式 RL 需要隨著環境回饋被附加至軌跡上而擴大上下文視窗。團隊實作了兩項主要的記憶體最佳化,以防止 Out-of-Memory(OOM)失敗。
減少 MoE 專家實體化
研究人員發現 Hugging Face Transformers 的推論前向路徑會為所有專家複製隱藏狀態,導致在 GPU 記憶體中實體化極大的張量(例如,嘗試為 20B 模型分配 180 GiB)。他們修補了實作,改用更節省記憶體的執行路徑,將專家順序處理。
使用 FlashAttention v3 的序列平行化
為了進一步降低每個 GPU 的激活記憶體,團隊實作了序列平行化(上下文平行化)。此方法將輸入序列分割至多個裝置,減少峰值激活佔用。
由於 attention 層需要序列的所有 token 同時位於同一 GPU,團隊實作了全對全(all-to-all)通信策略:
- 注意力前(Pre-attention): 收集序列元素,於 attention‑head 級別進行切分。
- 注意力後(Post-attention): 將輸出重新分配回原始的序列平行布局。
此設計具備 attention‑sink 感知,且相容於 FlashAttention v3,使模型能處理多步代理所需的長上下文視窗。