GPT-OSS Agentic RL 訓練: 実践的な振り返り
TL;DR
Hugging FaceとLinkedInの研究者は、PPOオンポリシー整合性、注意シンクのバックワードパス、およびMoEメモリの具現化における重大な不安定性を解決することにより、GPT-OSSモデルに対するエージェント強化学習(RL)のロックを成功裏に解除しました。これらの修正により、GPT-OSSは環境やツールと相互作用するマルチステップ意思決定エージェントの安定したバックボーンとして機能するようになります。
Agentic RL と GPT-OSS
エージェントRLは、意思決定プロセス全体を最適化することにより、従来のシングルターンRLとは異なります。静的な応答を生成する代わりに、モデルはマルチステップ軌道上で行動を計画し、ツールを呼び出し、行動を適応させるように学習します。これにより、エージェントがロールアウト軌跡を収集し、報酬を計算し、PPOやGRPOなどのアルゴリズムを使って方針を反復的に更新する閉ループシステムが必要となります。
GPT-OSSはOpenAIのo3-miniおよびo4-miniと同等のパフォーマンスを示してきましたが、エージェントRLへの適合性は以前まで検証されていませんでした。研究者はverlトレーニングフレームワークを使用し、GSM8K、Retool(エージェントコーディングタスク)、および検証可能な指示従従タスクでGPT-OSS-20Bモデルをテストしました。
PPOオンポリシー不安定性の解決
初期のトレーニングランでは、報酬が増加しないままKLダイバージェントとエントロピーが爆発的に増加する現象が見られました。チームは、Mixture of Experts(MoE)アーキテクチャによって引き起こされるProximal Policy Optimization(PPO)オンポリシー整合性の失敗を特定しました。
MoEログ確率の不一致
純粋なオンポリシーPPOでは、重要度サンプリング比は正確に1でなければなりません。しかし、GPT-OSSなどのMoEアーキテクチャでは、浮動小数点の違いや確率性により、ガットネットワークがロールアウト生成に使用されるフォワードパスとトレーニングに使用されるフォワードパスの間で入力を異なるエキスパートにルーティングする可能性があります。これにより、現在のログ確率と古いログ確率の間に不一致が生じ、PPOクリップが誤ってトリガーされ、オンポリシーの仮定が違反されます。
修正: チームは、環境がオンポリシーであることがわかっている場合(ミニバッチサイズがグローバルバッチサイズと等しい場合)に計算をオーバーライドするログ確率の置換を実装し、old_log_prob = log_prob.detach() を設定することで重要度比を1に強制しました。
注意シンクを介したトレーニング-推論の不一致の修正
PPOの整合性を修正した後でも、勾配ノームは爆発し続けました。研究者は、推論エンジン(SGLang)とトレーニングスタック(FSDPとFlashAttention-v2)が異なるトークンレベルの確率を生成するという、トレーニング-推論の根本的な不一致を特定しました。
注意シンクの役割
GPT-OSSは、ソフトマックス計算において「仮想トークン」として機能する学習可能なスカラーパラメータである注意シンクを利用します。これらのシンクは、コンテンツトークンに強制するのではなく、学習されたパラメータに注意質量を割り当てることを可能にし、ストリーミング推論の安定性を向上させます。
FlashAttention v3での実装
チームは、verlが注意シンクをサポートしないFlashAttention v2にハードコードされていることを発見し、さらにv2およびv3のいずれもシンク勾配に必要なバックワードパスをサポートしていないことを突き止めました。これを解決するために、彼らは以下を行いました:
- vLLM FlashAttentionフォークからフォワードパスを活用しました。
- シンク勾配$rac{\partial L}{\partial S_{h}}$を計算するバックワードパスを実装しました。
結果: この修正により、GSM8K、VerifyIf、およびRetoolタスク全体で収束が大幅に高速化され、報酬の改善が安定しました。
長いコンテキスト向けのメモリ効率のスケーリング
エージェントRLでは、環境からのフィードバックが軌跡に追加されるに従って、コンテキストウィンドウを拡張する必要があります。チームは、Out-of-Memory(OOM)障害を防ぐために、2つの主要なメモリ最適化を実装しました。
MoEエキスパートの具現化の緩和
研究者は、Hugging Face Transformersの推論フォワードパスがすべてのエキスパートに対して隠れ状態を複製し、GPUメモリに極めて大きなテンソルを具現化していることを発見しました(例えば、20Bモデルに対して180 GiBの割り当てを試みるなど)。彼らは、エキスパートを順次処理するよりメモリ効率の高い実行パスを使用するように実装をパッチしました。
FlashAttention v3を用いたシーケンスパラレル
GPUごとのアクティベーションメモリをさらに削減するため、チームはシーケンスパラレル(コンテキストパラレル)を実装しました。これにより、入力シーケンスがデバイス間で分割され、ピークアクティベーションフットプリントが削減されます。
注意レイヤーは、シーケンスのすべてのトークンが1つのGPU上に存在することを必要とするため、チームはall-to-all通信戦略を実装しました:
- 事前注意: シーケンス要素を収集し、注意ヘッドレベルで分割します。
- 後処理: 出力を元のシーケンスパラレルレイアウトに再配布します。
この設計は注意シンク対応であり、FlashAttention v3と互換性があり、マルチステップエージェントに必要な長いコンテキストウィンドウをモデルが扱えるようにします。