推論段階・KV cache・メモリ

このページの目次

同じモデルでも入力が長くなり、同時要求が増えると何が変わるでしょうか。prefill と decode を分け、重みと過去の K/V を別々に計算します。再計算できる容量式を使い、計算量・記憶容量・実測速度を区別します。

接頭部分全体から新しい 1 位置へ

prefill は既知の入力を処理します。因果注意力では位置 i は先頭から i 個の位置だけを使いますが、既知の入力位置は一度の forward 内でまとめて計算できます。decode では、生成した token を次の forward に入力し、その K/V を計算して履歴に追加します。

未来の token は過去位置の表現を変えないため、履歴の K/V を再利用できます。ただし新しい query は参照可能な key とのスコアと value の集約を必要とします。図は 1 層・1 ヘッドの利用可能な query-key 対を数えます。実際のカーネルはブロック単位で計算し、完全なスコア行列を実体化しない場合もあります。KV cache の原理

n 位置の提示には n(n+1)/2 対があり、続く m 個の新しい位置をキャッシュ付きで処理すると mn+m(m+1)/2 対が追加されます。ここでは生成 token を後続 forward に入力する時点で数え、提示末尾の logits から最初の出力を選ぶために提示をもう一度処理するとは数えません。FLOPs・時間・特定エンジンの速度予測ではなく、依存関係の数です。

キャッシュ後も計算する位置はどれか

位置を増やすたび、キャッシュなしでは伸びた接頭部全体を再処理し、ありでは query を 1 行だけ追加します。n=6、m=2 なら 21+7+8=36、全体の再計算なら 21+28+36=85 です。

図を準備しています
キャッシュ後も計算する位置はどれか

キャッシュは過去 query の重複処理を省き、新 query は履歴 K/V を参照します。

重みとKVキャッシュのメモリ使用量の計算

重みのストレージはまず総量で概算

合計35×10⁹個のパラメータがあり、すべて理想的な4ビットで保存されると仮定すると、裸の重みは 17.5×10⁹ バイト、約 16.30 GiB となります。FP16の裸の重みであれば 70×10⁹ バイト、約 65.19 GiB です。実際の量子化では、スケールやグループ化のメタデータが含まれ、特定のテンソルはより高い精度を保持することがあります。量子化フォーマット名に含まれるビット数は、そのまま最終ファイルの平均ビット数として扱えません。

すべての重みが同じデバイスに常駐する場合、このデバイスのメモリは総量で予算を組む必要があります。デバイス間での分割、CPUアンロード、階層型キャッシュを行うこともできるため、「すべてのエクスパートが同じGPUに常駐しなければならない」というわけではありません。その代償として、異なる位置でエクスパートを選択する際に、転送や待機のコストを支払う必要があります。ロードできることと、目標レイテンシでサービスを提供できることは、別の課題です。

KVキャッシュは総パラメータラベルではなく、構造とシーケンスで見る

各層で完全なK/Vを保存し、各シーケンスが等長で標準的なレイアウトを採用する簡略化されたモデルの場合:

KVバイト数 ≈ 2 × L × B × S × Hkv × D × b

記号意味
2KとVの2つのキャッシュ
Lこのキャッシュ構造を採用する層数
B、S同時にキャッシュされるシーケンス数、各シーケンスの長さ
Hkv、DKVヘッド数、ヘッドあたりの次元
b各キャッシュ要素のバイト数

例えば、L=32、B=1、S=8192、Hkv=8、D=128、b=2 の場合、結果は 1 GiB です。長さが65536に拡大すると 8 GiB になります。他のパラメータが変わらず、KVヘッド数が32に変更された場合、それぞれ 4 GiB と 32 GiB になります。これにより、「64Kコンテキスト」と言うだけではVRAMを推定できない理由が説明できます。

GQA(Grouped-Query Attention)では、複数のクエリグループがより少ないKVヘッドを共有するため、キャッシュ量はKVヘッド数に基づいて計算され、クエリヘッド数に基づいて計算されるわけではありません。GQA 論文 スライディングウィンドウ、圧縮キャッシュ、混合ループ層、または異なる層構造の場合、実際のレイアウトに基づいて層ごとに計算する必要があり、この積算式を機械的に適用してはいけません。

from fractions import Fraction

def kv_bytes(layers, sequences, length, kv_heads, head_dim, item_bytes):
    return 2 * layers * sequences * length * kv_heads * head_dim * item_bytes

one = kv_bytes(32, 1, 8192, 8, 128, 2)
long = kv_bytes(32, 1, 65536, 8, 128, 2)
assert one == 1024**3
assert long == 8 * 1024**3
weights = Fraction(35_000_000_000 * 4, 8)
print("理想的な4ビット重みのGiB:", round(float(weights / 1024**3), 2))
print("例のKV GiB:", one // 1024**3, long // 1024**3)
# 理想的な4ビット重みのGiB: 16.3
# 例のKV GiB: 1 8

最後に、活性化、演算子のワークスペース、アロケータのオーバーヘッド、その他のプロセスの占用分を加算し、目標の入力長と同時実行数においてピーク値を測定する必要があります。上記の計算例は、特定の実際のモデルを測定したものではありません。具体的なデプロイメントと量子化の記録については、ローカルLLMデプロイメント を参照してください。

KV を増やす因子はどれか

まず 1 因子だけ、次に長さと並行数を同時に変えます。KV heads と保持済み位置数を使い、GiB は 2³⁰ byte で換算します。容量から実測スループットを直接求めるモデルではありません。

図を準備しています
KV を増やす因子はどれか

KV は位置数・系列数・KV ヘッド数の積で増え、総パラメータ数だけでは求まりません。

計算量から実際のスループットへ

トークンあたりの活性化パラメータは、特定の行列計算を粗く見積もるのに役立ちますが、トークン/秒に直接換算することはできません。長いプロンプトの処理(prefill)、段階的な生成(decode)、高同時実行数のバッチ処理では、ボトルネックが異なる可能性があります。

観察される現象考えられる制限比較すべき測定値
最初のトークンに非常に時間がかかる長い入力の処理、キュー待ち、プレフィックスの再利用状況キュー待ち時間、prefill時間、キャッシュヒット率
単一リクエストの生成が遅い重みまたはKVの読み取り、演算子の効率ステップごとのレイテンシ、コンテキスト長、メモリ帯域幅
バッチサイズを増やしても収益が小さいエクスパートのホットスポット、通信、ワークスペースまたはスケジューリング総スループットと各リクエストの尾部レイテンシ
アンロード後は実行できるがカクつくエクスパートまたは層のデータ転送デバイス間の転送量と待機時間

MoEのエクスパート行列は比較的小さい場合もありますが、分散と結合にはオーバーヘッドがあります。マルチデバイスでのエクスパート並列処理では、トークン表現の交換が必要になることもあります。2つのモデルの活性化パラメータが近接していても、層数、KV構造、量子化カーネル、バッチサイズ、ルーティング分布などが異なることで、速度が顕著に異なる可能性があります。したがって、「35B-A3Bが3Bの稠密モデルのように速い」というのは、検証待ちの性能仮説に過ぎません。

次に読む:推論時計算・候補・検証。