AIの基礎

KVキャッシュはビットの問題ではなく、幾何学的問題を持っています。

mm
Unite.AI を Google の優先ソースに追加

同じ2ビットの精度で、量子化軸を選択する1つの決定は、ベンチマークスコアを2.88から63.53まで変化させる。キーと値は反対の処理が必要であり、その理由はハードウェアではなく、注意の式にある。

Llama-2-13Bを考えてみましょう。キーと値のキャッシュを32の量子化グループサイズでグループ化し、他のすべてをそのままにしています。同じモデル、同じビット予算、同じグループサイズ、同じベンチマークです。

実装の決定の1つの要素によって、CoQAの精度は2.88または63.53のどちらかになります。フル精度でのスコアは66.37です。

問題は、使用する総ビット数ではなく、スケールファクターを計算するときにどの軸をグループ化軸として選択するかです。チャネルをキー(キー)のグループ化次元として、トークンを値(値)のグループ化次元として使用すると、フル精度のパフォーマンスに近いところになります。如果これらの選択を反転すると、品質が低下します。如果両方の選択を反転すると、モデルは機能しません。

同じ2ビットを同じキャッシュに費やす4つの方法。Llama-2-13BでのKIVIの消去実験の結果。

量子化は通常、1つのダイヤルとして考えられます。8ビット、4ビット、2ビット、精度のコストが付随しています。しかし、KVキャッシュ内ではそうではありません。座標系を選択することであり、キーと値には異なるシステムが適用されます。この記事では、その理由を説明します。簡単に言うと、量子化エラーはグループ内の値の範囲に依存します。キーと値には非常に異なる構造があります。多くの人が値の分布から正しい軸を導き出すことができないため、間違えることがあります。エラーがどのように変化するかを注意して消費した後に確認する必要があります。那により、 intermediate アクティベーションを圧縮するための一般的な原則と、再構築エラーを品質の代理として疑う理由が得られます。

KVキャッシュが問題を抱える理由

生成フェーズ中に、トランスフォーマーは、以前に処理したトークンのキーと値の投影(KV)データをすべてキャッシュに保存して、データを再計算する必要がないようにします。そのキャッシュは、コンテキストの長さとバッチサイズとともに線形に増加します。最終的に、このキャッシュはモデル自体よりも大きくなることになります。

この増加は、モデルのさまざまな部分のメモリ使用量を調べることで簡単に特定できます。LLaMA-7BのKVQuant分析では、シーケンス長が512のとき、重みは約98パーセントのメモリを占め、活性化は2パーセントでした。一方、コンテキストが128Kのとき、重みとKVキャッシュの比率は約16パーセントと84パーセントに逆転しました。KIVIの著者が引用したOPT-175Bの分析では、バッチサイズが512で、512トークンのプロンプトのとき、KVキャッシュは1.2TBに達し、モデル重みの何倍にもなりました。

ただし、容量はここでの問題の半分だけです。GPUは、生成するたびに、デバイスメモリからKVキャッシュ全体を読み出す必要があります。つまり、GPUがKVキャッシュを読み出している間、コンピュートコアはアイドル状態になります。したがって、キャッシュの全体的なサイズを削減すると、使用可能な処理ヘッドルームが増加し、データ転送に費やされる時間が短縮されます。

量子化エラーの実際

一様整数量子化は数学的に単純です。数字のグループに対して、最小の数字をゼロ点として記録し、表現可能なレベルの数でグループの範囲を割ってステップサイズを取得します。次に、各要素を最も近いステップに丸めます。2つの即時結果が得られます。最初に、要素ごとのエラーは半分のステップサイズに制限されます。2番目に、ステップサイズはグループの範囲を2^B – 1で割ったものです。2ビットの場合、4つのレベルしかありません。したがって、グループ内の他の要素よりも100倍大きい要素は、ただ悪くはならない。ステップサイズを大きくし、他の要素も粗くします。グループは損傷の単位です。軸を選択することは、どの要素が一緒に苦しむかを決定することです。質問を違う角度から見ると、「どのビット数を犠牲にできるか?」ではなく、「極端な値はどこにあり、隔離できるか?」ということになります。

キー:固定チャネルに外れ値がある

大規模言語モデルには、他の活性化よりも非常に大きい活性化があります。Sunと同僚は、さまざまなモデルファミリーのこれらの非常に大きな活性化をカタログ化しました。Mixtral 8x7Bでは、最大の大きさは約7000で、特徴の trung bình大きさは約0.3でした。つまり、4桁以上離れています。これらは非常にまれで、入力が変化してもほとんど変化しない次元に固定されています。偶然ではありません。注目に焦点をあてるために暗黙の偏見として機能します。キー キャッシュでは、この構造は非常に明確です。特定のチャネルは、シーケンス内の各トークンにわたって一貫して非常に大きな大きさを運びます。トークンでグループ化すると、各グループにはこれらの外れ値チャネルが含まれるため、各グループのステップ サイズは外れ値によって設定され、通常のチャネルはそれに従います。チャネルでグループ化すると、外れ値チャネルは独自のグループを形成します。その内部範囲は大きいですが、自己完結しています。通常のチャネルは単独で残ります。結果は一致しています。Llama-2-13Bの層とヘッドを平均すると、KIVIはトークンごとのグループ化でのキー再構築エラーを13.67、チャネルごとのグループ化でのエラーを4.55と報告しています。さらに重要なのは、スコア エラーが47.00対9.60であるということです。トークンごとの量子化では、スコア エラーが約5倍になります。スコアは、キーに対して有意義なメトリックと一致しています。チャネル量子化は両方の面で優れています。

値:直感が壊れるところ

値キャッシュには、チャネル外れ値パターンは表示されません。かなり平坦に見えます。範囲の議論によっては、どちらの軸でも同等の品質の圧縮が得られる可能性があります。

しかし、実際にはそうではありません。キー管理の実装に関係なく(2.80と2.88の結果)、チャネルごとの値を圧縮するとモデルが崩壊します。

そして、ここに重要な点があります。各値の圧縮された元のテンソルに対する生の再構築エラーを使用してこの損失を測定すると、チャネルごとの値量子化は実際に少し良く、3.73対4.57となります。如果あなたが圧縮を最も明らかな方法で検証した場合、あなたはモデルを壊す構成を選択することになります。

Llama-2-13Bでの値キャッシュ量子化エラー、2つの方法で測定。保存されたテンソル メトリックと消費された出力メトリックは、1桁以上異なります。

解決策は、値キャッシュが直接読み出されないことです。行列積によって消費されます。注目出力は、ソフトマックス注目スコアを重みとして、トークン全体の値ベクトルの加重平均です。したがって、関連するエラーは、テンソル自体内で発生するのではなく、このプロセス中に発生するエラーです。注目出力に基づいて測定すると、順序は完全に逆転しました。KIVIが報告した、トークンごとの値ベクトル量子化による注目出力の相対エラーは3.55でした。一方、チャネルごとの量子化では49.89でした。後者は、値の圧縮の仕方によっては、前者の14倍以上大きくなります。

説明は、注目スパース性にある。彼らは84.3パーセントと測定しました。大部分の情報は、重要度の高いトークンの少数に帰属します。トークンごとの量子化では、各トークンのエラーがそのトークンに限定され、重要でないトークンのエラーは近似ゼロの注目重みによって乗算され、実質的に消えます。一方、チャネルごとの量子化では、各トークンのエラーが共有チャネルスケール全体に広がり、表現の悪いトークンが重要なトークンの表現を汚染します。注目が効率的であることを可能にするスパース性は、トークンごとの量子化が安全であることを示すのと同じ特性です。

転送可能な教訓は、KVキャッシュを超えています。テンソルが消費される場所で圧縮エラーを測定することです。再構築エラーによって暗黙的に仮定されているのは、テンソルの各コンポーネントが最終出力に等しい重みで貢献することです。注目は明示的にそうではありません。入力を重み付け、ゲート化、またはスパース化するダウンストリーム操作は、この仮定を破壊します。以前の記事で説明したように、検索システムの評価メトリックの盲点と同様の結果です。

回転埋め込みはキーを複雑にします

回転位置埋め込み(RoPE)を使用する際にいくつかの問題があります。RoPEは、各トークンの相対的な位置に基づいてチャネルのペアを回転させます。その混合は、チャネルごとのキー量子化が機能した固定チャネル構造を部分的に溶解します。外れ値チャネルが回転して隣接するチャネルに影響を及ぼし、隣接するチャネルが範囲を継承します。KVQuantの答えは、順序付けです。回転が適用される前にキーを量子化し、量子化の後でRoPEを適用します。チャネルごとのキー量子化、非一様データ型、および外れ値の小さな部分を分離することにより、0.1のパープレクシティ低下を3ビットで実現し、1百万トークンのコンテキストでLlama-7Bを1つのA100-80GBで提供できるようになります。

RoPEの影響のレベルを理解することも重要です。RotateKVという論文の著者は、RoPEを追加すると量子化エラーが145パーセント増加したと報告し、外れ値チャネルは注目ヘッドごとに異なり、すべての場所で共有の回転行列を適用することは不十分であり、ヘッド適応回転がより良いことを指摘しました。

システム課税、そしてそれが詳細ではない理由

トークンごとの量子化は、デコーディングに適しています。各トークンが到着し、量子化し、シーケンスに追加します。何も動きません。

しかし、チャネルごとの量子化は適していません。チャネルの統計は、まだ生成されていないトークンをまたいでいるため、トークンが到着したときにスケールファクターを計算することはできません。KIVIの回避策は、最近のトークンを最大128個をフル精度で残りのバッファに保持し、十分な数が蓄積したときにグループ化して量子化します。

実際、残りのバッファは重要な役割を果たすことになります。GSM8KでLlama-2-7Bを使用すると、フル精度のスコアは13.50です。2ビットに完全量子化すると、正しい軸でスコアは5.76です。同じ軸で、2ビットで、最近生成されたトークンのフル精度のスライディングウィンドウを追加すると、スコアは12.74になります。最近生成されたトークンのフル精度のスライディングウィンドウは、難しい多段問題で失われたものの多くを回復します。計算の連鎖によって注目されるトークンを考慮すると、これは妥当です。

これらのことを正しく行うことで大きな利点があります。KIVIは、Llama-2-7Bのピークメモリ使用量が2.6倍少なく、バッチサイズが最大4倍になるだけでなく、実際のサービス タスクでのスループットが2.35〜3.47倍になることを報告しています。

これについてどうするか

  1. 両方に同じ量子化器を使用しないでください。 キー(チャネルごと)と値(トークンごと)に異なる量子化器を使用します。単一の量子化器を「KVキャッシュ」に適用するパイプラインは、少数のビットを使用して値を表すときに、可能な限り多くの品質をすでに犠牲にしている可能性があります。
  2. RoPEを適用する前にキーを量子化します。 これは、好みの問題ではなく、正しさの問題です。
  3. 最近生成されたトークンのフル精度ウィンドウを保存します。 そのウィンドウを保存することでメモリをほとんど使用しませんが、難しいタスクの精度の多くを生成するのはこの領域です。
  4. 再構築エラーで検証しないでください。 常に注目出力またはエンドタスクのパフォーマンスに基づいて検証します。保存メトリックはただ汚いだけではなく、値の場合は方向が逆です。
  5. 短いコンテキストの多肢選択ベンチマークで検証しないでください。 KIVIの著者は、キャッシュが時間の経過とともに構築され、生成から読み出されることを実行しない限り、キャッシュの設計に内在する故障を観察することができないため、MMLUのような閉じた質問のタスクを避けます。デコーディングステップは出力ロジットを読み出すだけで、キャッシュをほとんど使用しません。

研究の行方

幾何学的な問題の性質についてはまだ行うことがありますが、多くの研究者が、外れ値チャネルがさまざまなトランスフォーマーヘッドにどのように分布しているか、ハードウェアの制限がどのグループ化を最も安価にするかを研究しています。InnerQは、チャネルごとのキー正規化を事前充填中にキーとクエリの重みに折り込むことで、実行時に追加のオーバーヘッドを回避します。また、InnerQは、最近生成されたトークンと注目シンク トークンの両方の高精度ウィンドウを保存します。そうすることで、シンク チャネルの外れ値が隣接するチャネルを汚染する可能性を排除します。

他の人は、キャッシュ全体を保存するのではなく、必要に応じてキーと/または値(s)をリマテリアライズできるだけの情報を保存することを提案しています。

最後に、量子化が精度だけに影響を与えるわけではないことを覚えておくことが重要です。KVキャッシュを量子化すると、FP8キャッシュとトレーニング不要の回復プロトコルを使用して、最大97パーセントの損失を回復できることを実証した、最近の研究は、KVキャッシュを量子化するとアラインメントの低下につながることを示しています。したがって、構成がベンチマーク結果を保持している場合でも、他のすべての関心があるパラメータを保持しているわけではありません。

一般原則

量子化の考え方は「精度予算」として捉えられます。いくつのビットを犠牲にできるか。KVキャッシュは、より有用な質問は構造的なものであることを示しています。精度はグループで割り当てられ、グループは損傷の単位であり、グループ化する軸はどの要素が運命を共有するかを決定します。正しい軸は、テンソルが消費される方法であり、メモリに保存されている方法ではありません。キーはドット積計算で使用されます。1つの汚れたチャネルはすべてのスコアを汚染します。値は、トークン全体の値ベクトルの加重平均として、ソフトマックス注目スコアを重みとして消費されます。したがって、1つの汚れたトークンは単に重み付けされます。

2つのテンソルが同じ次元で、2つの連続する層で生成されていても、別々に扱われます。どの活性化を圧縮するかを考える場合は、どの操作がそれを縮小するか、そしてグループ化がそれを尊重しているかを尋ねる価値があります。

ヒマンシュ・ゴエルは、バイオメディカル、金融、規制文書ワークフローなどのハイステークス・ドメイン向けのリトリーバル増強生成を専門とするAI/ML研究者です。