記事
推測によるソナーの加速

推測デコーディングは、大規模言語モデル(LLM)の生成速度を、素早く小さなドラフトモデルを使用して候補を生産し、大きなターゲットモデルで検証することにより加速します。
この方法では、単一トークンを生成する高価なターゲットランの代わりに、複数のトークンが一度に生成されます。ここでは、Sonarモデルでの中間トークンレイテンシーを削減するために適用されたさまざまな種類の推測デコーディングの実装の詳細を示します。
推測デコーディング
推測デコーディングは自然言語の構造とトランスフォーマーの自己回帰特性を活用してトークンの生成を加速させます。Llama-70Bのような大規模モデルがLlama-1Bのような小規模モデルよりも多くの知識を持っているとしても、簡単なタスクでは同様に機能することがあります。この重複は、より複雑な問題を大規模モデルに任せて、特定のシーケンスがより安価なモデルによって生成される方が良いということを示しています。課題は、どの完了がより良いかを決定し、小規模モデルの生成が大規模モデルと同じ品質かどうかを判断することにあります。
幸いにも、LLMは自己回帰トランスフォーマーです:トークンのシーケンスが与えられた場合、次のトークンの確率分布を出力します。さらに、入力シーケンスのトークンに関連する中間機能から導き出されるロジットも、モデルがその正確なトークンを発行する可能性がどれだけ高いかを示しています。この特性により推測が可能になります:もしトークンのシーケンスが小規模モデルから入力プレフィックスで生成された場合、それはターゲットモデルとどれだけ一致しているか確認するために大規模モデルに通すことができます。候補の各プレフィックスは確率でスコアが付けられ、受諾しきい値を超えた一番長いものが選ばれます。ボーナスとして、ターゲットモデルは次のトークンを無料で提供します:もしドラフトモデルが n トークンを生成すると、n + 1 までのトークンが一度に生成されることができます。

推論時、推測サンプリングプロセスは概ね4つのステージに分割できます:
プレフィル:ターゲットモデルとドラフトモデルの両方が、KVキャッシュエントリを埋めるために入力シーケンスで実行されなければなりません。Medusaのようないくつかのスキームでは、予測に簡単な密な層を使用しますが、本ポストではKVキャッシュ自身を必要とするトランスフォーマー系ドラフトに焦点を当てています。
ドラフト生成:ドラフトモデルは一定数の固定トークンを生成するまで繰り返します。ドラフトのシーケンスは線形であるか、モデルは一定の深さまでのツリー構造を探求することができます(EAGLE, Medusa)。ここでは線形シーケンスに焦点を当てています。
受容:ターゲットモデルがドラフトシーケンス上で実行され、各ドラフトトークンに対応するロジットを構築します。最も長い許可可能シーケンスの長さが決定されます。
ターゲット生成:ターゲットが生成するロジットは、シーケンスの不一致位置または末端において、次のトークンに対応します。これらのロジットは、シーケンスを締めくくる堅牢なトークンを提供するためにサンプリングされることができます。
推測デコーディングを実装するためにはさまざまな方法があります。本ポストでは、社内の1Bモデルを使用してSonarモデルを加速するために使用したスキーム、およびDeepSeek規模のモデルを加速するために構築している予測メカニズムに焦点を当てます。
ターゲット-ドラフト
推測デコーディングは、既存の小規模LLMをドラフトモデルとしてターゲットモデルに結合し、候補シーケンスを生成することによって達成できます。実際には、同じデータセット上でファインチューニングされたLlama-1Bモデルを使用してSonarを加速しました。この方法はドラフトをゼロからトレーニングする必要がない一方で、小規模モデルは依然としてかなりのKVキャッシュ容量を使用しわずかなプレフィルオーバーヘッドを導入し、TTFTを増加させます。
このスキームでは、デコーダはデコード専用バッチでのみ推測し、プレフィルまたは混合プレフィル-デコードバッチ中に標準のサンプリングを通じてトークンを生成します。プレフィルステージでは、ターゲットのロジットはすぐにサンプリングされ、新たに生成されたトークンをドラフトのKVキャッシュにプレフィルするのにも使用されます。ドラフトはまだサンプリングは行われませんが、それによって生成されたロジットはデコードステージに持ち越されます。

デコードでは、ドラフトモデルが進行し、各ステージでトップトークンをサンプリングします。望ましいドラフトの長さに達した後、トークンはターゲットモデルで実行され、サンプラーが推定長さを特定します。受容は、ドラフトとターゲットの確率分布全体を比較することによって決定されます。ターゲットは常に受け入れられたドラフトシーケンスの後に一つのロジットを出力するので、それをサンプリングして追加の出力を生成します。ドラフトモデルはまだ受け入れられたトークンを見ていませんが、再実行され次のデコードステップの準備として対応するKVキャッシュエントリを埋めます。
EAGLE
EAGLEは推測デコーディングスキームで、多くのドラフトシーケンス、つまりドラフトトークンの可能性あるツリー状の探索から生成されます。固定(EAGLE)または動的にシェイプされた(EAGLE-2)ツリーが、各ノードでTop-K候補を考慮してドラフトトークンの連続実行を通じて探索されます。次にシーケンスは評価され、適切な最も長いものが選ばれ、ターゲットの追加トークンも付加されます。

より正確な予測を達成するために、EAGLEドラフトモデルはトークンだけでなく、ターゲットモデルのターゲット機能(最終層の隠れ状態)を使用して予測します。EAGLEの欠点は、低レイテンシーバジェット内で適切な候補を生成するのに十分に正確な小さなドラフトモデルをトレーニングする必要があることです。通常、ドラフトモデルは、オリジナルモデルのデコーダ層と同一の単一トランスフォーマー層で、埋め込みやlm_headの射影と繋げてターゲットに密結合しています。これによりKVキャッシュ容量が少なくなるため、EAGLEは低メモリーフットプリントを持っています。
ターゲットモデルでツリー状シーケンスを検証するためには、カスタム注意マスクが使用されなければなりません。しかし、シーケンス全体にカスタム注意マスクを使用することは、現実的な入力長で注意時間を大幅に(最大50%)遅くし、推測から得られる一部の速度向上を無効化します。この理由から、完全なツリー探索を本番環境に導入していません。代わりに、DeepSeek-V3テクニカルレポートで提示されたMTPのようなスキームによる単一トークン予測の特別なケースに焦点を当てています。
MTP
このスキームは、トークンと共に隠れ状態を使用して予測を行うドラフト-ターゲットデコーディングに似ています。通常のドラフト-ターゲット推測よりも、プレフィルとデコードステージで少し多くの作業を行わなければなりません。ドラフトモデルはトークンと隠れ状態の両方を使用します:トークンt_{i+1}はトークンt_iに対応するロジットL_iからサンプリングされ、これらは順に隠れ状態H_iから導き出されます。結果として、入力トークンバッファは、ターゲットが出力する隠れ状態ベクトルに対して1ステップ左にシフトされなければなりません。下の図は、トレーニング時の対応を示し、推論時のシフトも示しています。

デコーディングフローは、隠れ状態とロジットの両方を持ち越す点以外はドラフト-ターゲットデコーディングにかなり似ています。私たちの実装は、関連するサンプリングとロジット処理すべてを共有しており、モデルのフォワード呼び出しを専門化しています。複数のトークンが予測される場合、ドラフトモデルはドラフト隠れ状態を使用して予測し、自身の特徴に基づいてKVキャッシュエントリを埋めます。長期間ではこれが精度を低下させる可能性があります。その後、ターゲット予測のためのKVキャッシュエントリを埋めるためにドラフトモデルを実行する際、全シーケンスでより正確なターゲット隠れ状態を入力として使用します。ドラフトモデルが小さいので、追加トークンを処理するためのコストは無視できます。
MTPヘッドのトレーニング
MTPを活用するために、Perplexityのデータセットでファインチューニングされているモデルに接続されたMTPヘッドをトレーニングするために必要なインフラストラクチャを構築しました。8xH100デバイスの1ノードで、Llama-1BからLlama-70B、DeepSeek V2-Liteまでのモデルのヘッドを約1日で構築できます。大規模モデルでは、ファインチューニングプロセス中に構築されたMTPヘッドを利用します。
MTPトレーニングの目的は、草案の隠れ状態とターゲットの隠れ状態から外挿されたロジットを、ターゲットの次のトークンロジットと隠れ状態に一致させることです。隠れ状態の推論は高価なので、推論用に最適化されたターゲットモデルの実装を使用してこれらを事前計算し、トレーニング中に使用します。しかし、推論MTP実装の検証と量子化や最適化による数値の違いが結果を妨げないことを確認するため、検証損失と精度の推定において、ターゲットモデルとドラフトモデルの推論実装を完全に再利用しています。
オリジナルの論文で使用されたShareGPTデータセットから大規模サンプルにスケールアップする際、EAGLE論文で概説および実装されているMTPヘッド構造が70Bサイズのモデルでトレーニングできなかったことに気づきました。ShareGPTにはより短いシーケンスが多く含まれていたのに対し、我々はやや少ない数のかなり長いプロンプトでトレーニングしています。オリジナルのEAGLEヘッドは典型的なトランスフォーマーとは少し構造が異なっていたため、削除されたRMS正規化レイヤーを再導入しました。これにより、トレーニングを収束させるだけでなく、ヘッドの精度が数ポイント向上しました。

レイヤーノルムはトレーニングを促進するだけでなく、ノルムを再導入することは数学的にも直感的です。MTPヘッドはターゲットモデルの埋め込みとロジット射影を再利用します。これらはLlama 70Bで約2 GBのサイズになる可能性があるので、トレーニング中は固定され、MTPレイヤーは元のモデルの射影レイヤーがトレーニング中に学習したベクトル空間に予測を埋め込むことを学ぶことが期待されます。ノルムを除去すると、一つのMLPはノルムを伴うMLPと同じ関数を学習することが期待されるため、ドラフトとターゲットモデルの隠れ状態の一致が妨げられます。
推測デコーディングを用いた推論
推論エンジンでは、入力シーケンスのトークンを生成するために、それらを適切なサイズのバッチにグループ化し、KVキャッシュに次のトークン用のページを割り当てなければなりません。入力トークンとKVページ情報は、モデルを実行して次のトークンをサンプリングするロジットを生成するために必要なすべての並列ランクにブロードキャストされるバッファに詰め込まれます。最後に、メタデータはGPUメモリにコピーされ、モデルが実行され、結果としてロジットが生成され、そこから次のトークンがサンプリングされます。
ラップで間に操作をすることでドラフトとターゲット推論サーバーを緩く結合するいくつかの実装とは異なり、我々のドラフト-ターゲットペアは密に結合され、一緒に生成を進めます。バッチのスケジューリングとKVページの割り当てはすべての推測デコーディング形式のためにモデル間で共有されます:これはモデルを包括的な推論サーバーに接続するロジックを統一し、すべてが同じインターフェースを公開します。
Perplexityの推論ランタイムは、注意カーネルを設定しスケジュールを決定するために必要なメタデータを決定するFlashInferを基盤としています。バッチを構成するいくつかの入力シーケンスがある場合、前フィル、デコード、または検証のために、CPU側の作業で中間バッファを割り当て、注意に使用される一定のバッファを埋める必要があります。この作業は、バッチスケジューリングおよびKVページ割り当てのコストに加えて、これらは最大限のGPU使用を可能にするためには隠さなければならないレイテンシーをもたらします。
推測なしで推論のCPU側およびGPU側の作業を完全に並列化しましたが、推測デコーディングのCPU-GPUのバランスがより複雑であることがわかりました。主な課題は、受け入れられたトークンの数が次の実行のシーケンス長さを決定するため、避けにくいGPUからCPUへの同期ポイントが発生することです。CPU作業のレイテンシーを最も隠すために、さまざまなスケジューリングスキームを実験しました。
ドラフト-ターゲットスケジュール
ターゲットモデルより小さいにも関わらず、全LLMがドラフトとして使用されると、GPUに大きなレイテンシーを導入し、CPU側の高価な操作を隠す余地を提供します。小さなモデルはテンソル並列処理から恩恵を受けないため、ターゲットとドラフトが分散されるランク数には不一致があります。私たちの実装では、ドラフトモデルはTPグループのリーダーランクでのみ実行されます。

前述のように、一回のデコードステップでロジットが次回の実行に持ち越されます。これにより、ドラフトモデルの一回の実行をCPU側のバッチスケジューリング作業と重ねることが可能になります。バッチがまとめられた後、サンプラーとドラフトへの繰り返し呼び出しがドラフトトークンを生成します。同時に、ターゲットモデルのバッチを検証のために組み立て、並列ワーカーで同期されます。ターゲットロジットは検証され、受け入れられたシーケンスの長さを決定するためにサンプリングされます。この時点で、次のシーケンスの長さを決定するためにGPUからCPUへの同期が必要です。ドラフトモデルはリーダーノードでのみ実行されるため、そのバッチは順次設定され、その実行はKVキャッシュエントリを追加トークンで埋めるために開始されます。現在のランでドラフトによって生成されたロジットは、次回の実行で最初のドラフトトークンをサンプルするために使用されます。最も重要なのは、ドラフトが動作中の間、次のバッチがスケジュールできることです。
単一トークンのためのMTPスケジュール
ランタイムはEagleスタイルのドラフトツリー探索をまだ提供していませんが、モデルサイズの単一トランスフォーマーデコーダ層によって生成されるドラフトトークンの線形シーケンスを考慮した、このスキームの特別なケースを実装しました。このスキームは、DeepSeek R1のオープンソースウェイトを使用したドラフト予測に使用できます。大規模MTPレイヤーが十分に高い受け入れ率を達成し、オーバーヘッドを正当化するため、単一トークンを予測する特に興味深いサブケースです。
MTPスケジューリングは、ドラフトモデルがはるかに高速でCPUサイドのレイテンシーを隠す要姧が少ないため、やや複雑です。加えて、ドラフトはターゲットモデルと共に分割されているため、バッチ情報に対する共用メモリ転送が必要です。実行は、前述のスキームと同様、持ち越しロジットから最初のトークンをサンプリングするバッチ情報の転送から始まります。次に、ターゲットはトークンを検証し、2 * Dトークンを処理します。Dはデコードバッチサイズです。これはMixture-of-Experts(MoE)モデルのマイクロバッチ処理に理想的で、Infinibandのような遅いインターコネクトで等分に分割できます。ターゲットの隠れ状態は次のドラフトランに持ち越され、一方ロジットはサンプラーで検証されます。

GPUにわずかな追加作業を行うことで、ドラフトシーケンス受け入れ後のCPU-GPU同期を回避します。ターゲットの入力トークンがシフトされた後、カーネルが次のターゲットトークンを対応する位置に挿入します。ドラフトは同じバッチ情報で再実行され、KVキャッシュエントリを埋め、次のラン用のロジットと隠れ状態を構築し、受け入れられなかったトークンで若干の冗長な作業を行います。このような場合、ドラフトモデルが小さいため、未使用作業のレイテンシーはわずかしか測定されません。ドラフトランと並行して、シーケンスの長さがCPUで決定され、次のバッチのスケジュリングが開始され、GPU作業の終了を待つことはありません。
ドラフト層での追加作業は注意によってほとんど目立ちませんが、MLPレイヤーはより問題があります。行列積命令がトークン数の次元に沿って64の境界にパッドされるため、倍増が大幅に多くのブロックを必要としない場合、オーバーヘッドは隠されます。長いドラフトシーケンスではオーバーヘッドはより高価であり、通常のドラフトターゲットモデルで使用されるスキームがより効果的です。