Direct Preference Optimization (DPO) を使用すると、開発者は個別の報酬モデルのトレーニングや強化学習 (RL) ループを実行することなく、大規模言語モデルをファインチューニングできます。これにより、計算コストと、従来の RLHF(人間のフィードバックからの強化学習)パイプラインで頻繁に発生する不安定性の両方を大幅に削減できます。

なぜ RLHF は負荷が高いと感じるのか

標準的な RLHF のレシピには、3つのステージがあります。まず、精選されたデータセットでベースモデルをファインチューニングします。次に、報酬モデルが、出力のペア間の人間の好みを予測することを学習します。最後に、実務者は Proximal Policy Optimization (PPO) — 古典的な RL アルゴリズム — を実行し、元のモデルに近づきつつ、予測報酬を最大化するように方策を押し進めます。

この3つのモデル構成は、メモリに同時に常駐し、更新のたびに新しいテキストをサンプリングするというコストのかかる RL ループを強制します。チームはしばしば、方策が報酬を「ハック」することを学習してしまい、プロキシ(代理)指標では高いスコアを出すものの、意図した品質を欠いた出力を生成してしまう現象に直面します。その結果、パイプラインは高コストで脆弱、かつスケーリングが困難なものになります。

DPO の代数的なショートカット

DPO は報酬モデルを完全にスキップします。重要な観察結果は、RLHF の目的関数(参照モデルからの乖離をペナルティとしつつ、期待報酬を最大化する)には閉形式の式が存在するということです。数学的に整理することで、任意のトークンに対する報酬は、方策によって割り当てられた対数確率と、参照モデルによって割り当てられた対数確率の差になります。

実際には、これは「報酬」が方策そのものの中に存在することを意味します。トレーニングは、好みのペアに対する単一の分類損失へと簡略化されます。選択された回答と拒否された回答が与えられたとき、モデルは選択されたテキストに対してより高い確率を割り当てるように促されます。サンプリングも、PPO の更新も、保存するための追加モデルも必要ありません。

新しい損失関数の仕組み

損失関数は、現在のモデルの方策における好ましい回答の対数確率を、参照モデルにおける対数確率と比較し、温度パラメータのようなハイパーパラメータ $\beta$ でスケーリングします。$\beta$ が高いと、方策は参照モデルの近くに留まるよう強制され、流暢さと安全性が維持されます。$\beta$ が低いと、方策はより遠くまで逸脱することができ、選択された回答に対する好みがより鮮明になります。

開発者にとって重要なメリット

  • 報酬モデルが不要 – 個別の予測器のために追加の人間のフィードバックを収集する必要がなくなります。
  • トレーニング中のサンプリングが不要 – モデルは勾配を計算するために新しいテキストを生成する必要がなく、GPU 時間を劇的に削減できます。
  • 安定性 – 分散の大きい RL 勾配の代わりに、標準的なバイナリ・クロスエントロピー損失を使用するため、発散が起こりにくくなります。
  • 効率性 – 各好みのペアに対して1回の勾配ステップを行うだけで十分であり、PPO よりもはるかに少ないエポック数でトレーニングが収束します。

初期の実験では、DPO は計算予算のわずかな一部を使用しながら、ベンチマークの好みのデータセットにおいて PPO の性能に匹敵、あるいはそれを上回ることが示されています。このコスト面での優位性により、多くのオープンソースプロジェクトが、すでに DPO またはその類似の変形手法をデフォルトのアライメント手法として採用しています。

トレードオフ

DPO は固定された好みのペアのセットに対して機能します。トレーニング中に新しい補完テキストをサンプリングしないため、元のデータセットに存在しなかった回答空間を探索することはできません。対照的に、オンラインの PPO 実行では、モデルを絶えず探索することで、斬新でより高い報酬を得られる振る舞いを発見できる可能性があります。

$\beta$ が低すぎたり、トレーニングのステップ数が多すぎたりすると、方策が参照モデルから大きく逸脱し、流暢さが失われたり、望ましくないアーティファクト(不自然な生成物)が生じたりする可能性があります。

まとめ

Direct Preference Optimization は、3つのモデルと重い RL を必要とする RLHF スタックを、人間の好みのペアから直接学習する、単一で安定した損失関数に置き換えます。その結果、トレーニングデータが必要な振る舞いを捉えており、かつドリフトパラメータを適切に管理できていれば、言語モデルのアライメントに向けた、より安価で予測可能な道筋が得られます。