LoGRA、低ランク勾配スケッチでLLM強化学習のメモリ障壁を下げる
背景
大規模言語モデルの推論能力を高める手段として、強化学習によるポストトレーニングの重要性が増しています。一方で、強化学習は通常の教師あり微調整より多くのメモリを必要とします。モデルの重みだけでなく、勾配やオプティマイザー状態、ポリシー更新に関わる情報も扱う必要があるためです。モデルが大きくなるほど、密なAdamのような従来方式ではハードウェアの制約に達しやすくなります。LoGRAは、この問題を勾配情報の保持方法から見直します。
LoGRAの仕組み
LoGRAの中心的な考え方は、学習に有用な勾配を完全な密行列として保持するのではなく、低ランクの勾配スケッチとして表現することです。低ランク表現はデータ量が小さいため、更新処理に必要なメモリを抑えられます。さらに、このコンパクトな表現はモデル更新だけでなく、ポリシー同期にも使われます。
ただし、勾配を圧縮するだけでは学習が不安定になる可能性があります。近似された更新が大きすぎると、ポリシーが急激に変化し、それまでの学習を乱すおそれがあるためです。そこでLoGRAは、予測KLによるステップ制御を導入します。更新を適用する前にポリシーの変化を推定し、その結果に応じて更新幅を調整します。つまり、メモリ削減とポリシーの変化量の管理を一体化しています。
論文要約で示された主な結果は次の通りです。
- 推論タスクで、性能を落とさず平均学習メモリを最大45.7%削減した;
- 密なAdamではメモリ不足になる条件で、単一の8 GPUノード上で27Bパラメータモデルを1,100ステップ超、安定して学習した;
- 実装はMoltライブラリのLoGRA用サンプルスクリプトとして公開されている。
意義と残る論点
LoGRAの意義は、単なるメモリ節約にとどまりません。これまでハードウェア上の理由で実行しにくかった大規模モデルのRL実験を、限られたGPU構成でも検討できる可能性があります。また、勾配圧縮は圧縮率だけを見るのではなく、圧縮後の更新が学習を壊さないかまで考える必要があることを示しています。
一方、今回確認できる素材は主に要約レベルであり、低ランク設定ごとの比較、タスク別の内訳、通信コスト、他のRLアルゴリズムでの結果までは示されていません。したがって45.7%という数字は、あらゆる環境で得られる固定的な改善幅ではなく、論文が報告する最大値として理解するのが適切です。異なるモデル構造や長期学習での有効性は、本文と実装を通じた追加検証が必要です。
コメント
ログイン状態を確認中…
コメントを読み込み中…