AI・機械学習
上級

RLHF PPO訓練(Proximal Policy Optimization)(アールエルエイチエフピーピーオークンレン)

RLHF PPO訓練は、報酬モデルのスコアを最大化しつつ参照ポリシーからの乖離を制御する強化学習フェーズである。Proximal Policy Optimizationアルゴリズムでクリッピングベースの方策更新を行い、KLダイバージェンスペナルティで出力の安定性を維持する。

0 回閲覧
0 いいね

RLHF PPO訓練はRLHFパイプラインの最終段階(Stage 3)であり、学習済み報酬モデルのスコアを報酬信号としてLLMの方策(ポリシー)を強化学習で最適化するプロセスである。Schulman et al.(2017)が提案したPPOアルゴリズムがデファクト標準となっている。

概要

PPO(Proximal Policy Optimization)は方策勾配法の一種で、方策更新の幅をクリッピングにより制限することで学習の安定性を確保する。LLMのRLHFにおいては、各タイムステップでトークンを生成するアクションに対し、応答全体の完成後に報酬モデルから得たスコアを報酬信号として方策を更新する。

PPO訓練の目的関数は:

J(θ) = E[r_RM(x, y) - β * KL(π_θ || π_ref)]

ここでr_RM は報酬モデルスコア、β はKLペナルティ係数(通常0.01〜0.2)、π_θ は現在の方策、π_ref はSFTモデル(参照ポリシー)である。

PPOアルゴリズムの詳細

  • クリッピング目的関数: 方策比率 r_t = π_θ(a_t|s_t) / π_old(a_t|s_t) を [1-ε, 1+ε] にクリップ。ε=0.2が標準。大きすぎる更新を防ぎ学習を安定化する
  • GAE(Generalized Advantage Estimation): λ=0.95でアドバンテージを推定。バイアスとバリアンスのトレードオフを制御
  • バリューヘッド: LLMに価値関数ヘッド(線形層)を追加し、状態価値V(s)を推定。アドバンテージ計算に使用
  • ミニバッチ更新: 1エポックあたり4〜8ミニバッチで複数回更新。InstructGPTでは各バッチ256プロンプト
  • 学習率: 通常1e-6〜5e-6。SFT時の1/10程度に設定。コサインスケジューラーまたはウォームアップ後リニア減衰

PPO訓練のハイパーパラメータ

パラメータ典型値説明感度
KL係数β0.01〜0.2参照ポリシーからの乖離ペナルティ非常に高い
クリップε0.2方策更新のクリッピング範囲中程度
GAE λ0.95アドバンテージ推定の減衰率低い
学習率1e-6〜5e-6方策ネットワークの学習率高い
バッチサイズ64〜512プロンプト数/バッチ中程度
PPOエポック2〜4ミニバッチの反復回数中程度
応答最大長512〜2048生成トークン上限低い

KLペナルティの役割と設計

KLダイバージェンスペナルティはRLHF PPOの最も重要な正則化メカニズムである:

  • 目的: モデルが報酬モデルのスコアを過度に最適化(報酬ハッキング)することを防ぎ、SFTモデルの言語能力を維持する
  • 固定係数方式: β を定数として設定。InstructGPTで使用。シンプルだが最適値の事前決定が困難
  • 適応係数方式: 実測KL値が目標KL値(target_kl=6.0 natなど)を超えたらβ を増加、下回ったらβ を減少させるPID制御的アプローチ。Anthropicが採用
  • KL報酬方式: KLペナルティを報酬の一部として各トークンに分配する方式。トークンレベルの制御が可能
  • 過小β: 報酬ハッキングが発生し、冗長・不自然な応答を生成。過大β: ほぼSFTモデルと同じ出力になり学習効果なし

PPO訓練の実装上の課題

  • メモリ消費: 方策モデル・参照モデル・報酬モデル・価値関数の4モデルを同時にGPUメモリに保持する必要がある。70Bモデルでは8×A100 80GBでもメモリ逼迫
  • 学習不安定性: 報酬スコアが突然崩壊する「reward collapse」が発生しやすい。チェックポイント保存と早期停止が重要
  • サンプル効率: 1回のPPO更新に数千のプロンプト-応答ペアを生成する必要があり、推論コストが支配的
  • 分散訓練: DeepSpeed-Chat(Microsoft)やOpenRLHF はZeRO Stage 3 + テンソル並列でメモリ問題を緩和。vLLM連携で推論を高速化する構成も普及

PPOの代替手法

  • REINFORCE: PPOより単純だがバリアンスが大きく学習が不安定。小規模実験向き
  • ReMax: Liらが2024年に提案。REINFORCEのベースラインを最大報酬で置換し、PPO同等性能をバリューヘッドなしで達成
  • GRPO(Group Relative Policy Optimization): DeepSeek-Math(2024年)で提案。複数応答のグループ内相対順位で報酬を正規化。バリューモデル不要でメモリ50%削減
  • Rejection Sampling: N個生成してRM最上位を選択。Llama 2で採用。PPOより安定だがN倍の推論コスト

よくある質問(FAQ)

Q1: PPO訓練にはどのくらいの計算資源が必要か? A: 7Bモデルで4×A100 80GB、1〜3日程度が目安。70Bモデルでは8×A100以上で1週間前後。参照モデルをLoRAで圧縮し、vLLMで推論を高速化する構成が2025年のベストプラクティス。DeepSpeed-ChatのHybrid Engine は学習と推論を統合しスループットを2〜3倍改善する。

Q2: KL係数βはどう決めるのか? A: 適応方式(target_kl=6.0 nat程度)から開始し、報酬スコアとKL値の推移を監視するのが実用的。報酬が上昇しKLが10 nat以下なら順調。KLが20 natを超えたら報酬ハッキングの兆候。固定方式ではβ=0.05〜0.1が安全な初期値。

Q3: PPOとDPOはどちらが性能が上か? A: 大規模データ(50万ペア以上)・大規模モデル(70B+)ではPPOがDPOを上回る事例が多い(Llama 2、GPT-4)。中小規模(7B〜13B、10万ペア以下)ではDPOが同等以上の性能をはるかに少ない計算コストで達成する。2025年時点ではGRPO・ReMaxなどPPOの軽量代替が研究の主流。

まとめ

  • PPOはRLHF Stage 3の標準アルゴリズムで、クリッピングにより安定した方策更新を実現
  • KLペナルティが報酬ハッキング抑制の鍵。適応係数方式が推奨
  • メモリ消費(4モデル同時保持)が最大のボトルネック。DeepSpeed-Chat/GRPO等で緩和
  • 7Bモデルなら4×A100で数日、70Bモデルは8×A100で1週間程度の計算コスト