LLMのための効率的な知識蒸留:オフラインTop-K LogitsおよびFused Chunked KL Loss
Multiverse Computingは、VRAM使用量とトレーニングコストを大幅に削減し、単一のGPUでロングコンテキストの蒸留を可能にする新しい知識蒸留手法を導入しました。オフラインのlogitキャッシュと、融合されたchunked KL-divergence lossを組み合わせることで、研究者たちは、フル語彙の確率分布に通常伴う膨大なメモリオーバーヘッドなしに、大規模な教師モデルをより小さな生徒モデルへと蒸留できるようになりました。
標準的な知識蒸留の高コスト
Kullback-Leibler (KL) 離散損失を用いた従来の「オンライン」知識蒸留は、教師モデルと生徒モデルの両方を同時にVRAMにロードする必要があるため、計算コストが高くなります。トレーニングの各ステップにおいて、教師モデルはすべてのトークン位置に対して全語彙にわたる確率分布を生成するために、フルフォワードパスを実行する必要があります。
大規模モデルの場合、これは巨大なメモリのボトルネックとなります。例えば、語彙数201,088、シーケンス長32K、バッチサイズ4のgpt-oss-120bのようなモデルを使用する場合、教師の確率テンソルだけでbfloat16において約50GBのVRAMを必要とします。勾配、アクティベーション、オプティマイザの状態を含めると、1回のイテレーションで約250GBのVRAMのピークに達する可能性があり、これはH200やB200のようなハイエンドGPUの容量さえも超えています。
スケーラブルな蒸留のための2つのシステム変更
これらのメモリ制約を解決するために、Multiverse Computingは2つの主要な技術的変更を実装しました。
1. Top-K Logit Cachingによるオフライン蒸留
各ステップで教師の出力を再計算する代わりに、システムは教師の出力を一度だけ計算し、位置ごとの上位100個の最も可能性の高いトークンをキャッシュします。その後、生徒モデルはこのキャッシュに対してトレーニングされます。これにより、トレーニングプロセス中に教師モデルをメモリに保持しておく必要がなくなり、同じキャッシュを異なる実験的なアブレーションで再利用することが可能になります。
2. Fused Chunked KL Loss
標準的なKL損失の実装では、生徒と教師の不一致を計算するために、全語彙×シーケンスのグリッドを構築します。Multiverse Computingの「fused chunked KL loss」は、モデルの出力プロジェクションを損失計算に直接融合させることで、これを最適化します。
フルロジットグリッドを実体化するのではなく、プロセスは次のように動作します:
- シーケンスのチャンクを一度に1つずつ処理します。
- その特定のチャンクに対してのみ、隠れ状態をロジットに投影します。
- 結果を進行中の損失に統合し、直ちにチャンクを破棄します。
- バックワードパスでは、各チャンクをオンザフライで再計算します。
このアプローチにより、ピークメモリは全語彙サイズに基づいて急増するのではなく、シーケンス長に対して線形に増加することが保証されます。
パフォーマンスとメモリのベンチマーク
単一のH200 GPU上で、Llama 3.1 8B Instruct(教師)と3.2B Llamaモデル(生徒)を使用し、8Kトークンのコンテキストで異なる損失実装を比較したところ、研究者たちはすべての手法がほぼ同一のトレーニング損失に到達することを発見しました。これは、上位100個のキャッシュされたロジットのみを使用したオフライン蒸留が、オンライン蒸留に対して実質的にロスレスであることを裏付けています。
| 手法 (8Kコンテキスト, 単一H200) | ピークメモリ | イテレーション時間 | スループット |
|---|---|---|---|
| Online distillation | 102.8 GB | 25.9 s | 237 TFLOP/s |
| Offline, dense KL | 78.3 GB | 18.5 s | 331 TFLOP/s |
| Offline, forward-chunked KL | 61.8 GB | 18.4 s | 335 TFLOP/s |
| Offline, fused chunked KL | 58.3 GB | 20.2 s | 304 TFLOP/s |
ロングコンテキストへのスケーリング
融合されたchunked lossの利点は、コンテキスト長が長くなるにつれてより顕著になります。トイ・アウトプットプロジェクションネットワークを使用した分離ベンチマークでは:
- 32Kトークン時: ピークメモリは85.2 GiB (dense loss) から 5.45 GiB (fully chunked) に減少し、15.6倍の削減を達成しました。
- 64Kトークン時: dense lossの実装は完全に失敗しました。
- 256Kトークン時: fully chunked lossは11.6 GiBを使用し、次に優れたチャンク版の134.2 GiBと比較して、イテレーションあたりの速度が3.3倍速くなりました。
GPT-OSS 20Bモデルを32,768トークンのコンテキストで蒸留する実世界のシナリオでは、融合損失により、必要なハードウェアが4つのGPUノードから1つに削減されました。これにより、ステップ時間が57.0秒から12.23秒へと5倍高速化し、GPUあたりのスループットは74.2から345.7 TFLOP/sへと増加しました。
生徒モデルの能力
このパイプラインの効率性により、Llama 3.1 8B Instructから蒸留された3.2Bパラメータの生徒モデルを生成する大規模な蒸留キャンペーンが可能になりました。このコンパクトな生徒モデルは、BoolQおよびHellaSwagにおいて教師の精度の大部分を維持しており、パラメータ数が半分以下であるにもかかわらず、MMLUベンチマークでは教師のスコアから約9ポイント以内に収まっています。
Sources
関連
- Dispatch
- Dispatch
- Dispatch
- Dispatch
- Dispatch