KVキャッシュの仕組み — LLM推論を高速化する基本技術

ChatGPTやClaudeでテキストを生成するとき、最初の1文字が表示されるまでは少し待ちますが、そこからはかなり速いペースで文字が流れてきます。この「最初だけ重い」構造の裏側で、KVキャッシュという仕組みが欠かせない役割を担っています。

KVキャッシュがなければ、GPT-4のようなモデルは1トークン生成するごとに、それまでのすべてのトークンに対してAttention計算を最初からやり直さなければなりません。1000トークンのテキストを生成するには、1+2+3+…+1000 = 50万回分の計算が必要で、これでは実用に耐えません。KVキャッシュは「一度計算したKey・Valueはそのまま使い回せる」という数学的な事実を利用して、この計算を根本から削減します。

KVキャッシュを理解することで、次の2つの問いに明確に答えられるようになります。

  • なぜLLMの推論は系列長が長くなるほど急激に遅くなるのか、そしてなぜKVキャッシュでそれが解決できるのか
  • Llama2やGPT-4がGQAやMQAを採用した理由は何か、そのメモリへの影響はどのくらいか

本記事の内容

  • 自己回帰生成における冗長な計算の問題(数式で定量化)
  • KVキャッシュの数学的原理と計算量削減の理由
  • メモリコストの見積もり式と実際のモデルでの具体例
  • KVキャッシュを使った生成のnumpy実装と速度検証
  • MQA・GQA・PagedAttentionによるKV圧縮の仕組み

前提知識

この記事を読む前に、以下の記事を読んでおくと理解が深まります。

画像なし
Self-Attention機構の理論と実装を完全解説
Query/Key/Valueの意味、Scaled Dot-Product Attentionの導出とPyTorch実装を解説します。
画像なし
Transformer Decoderの構造とMasked Self-Attentionの仕組み
Decoderブロックの構成とMasked Self-Attentionの仕組みを解説します。

KVキャッシュとは何か — 30秒で理解する直感

図書館で本を調べるとき、同じ棚を何度も往復するのは非効率です。「この棚の内容はメモしておいて、次に必要なときはメモを見れば済む」というのがKVキャッシュの本質です。

LLMの自己回帰生成では、ステップごとに「過去のすべてのトークン」を参照してAttentionを計算します。ここで参照する「Key」と「Value」は、過去のトークンの情報から計算されます。そして重要な点は、過去トークンのKey・Valueは、次のステップでも一切変わらないということです。

つまり「前のステップで計算したKey・Valueをそのまま保存(キャッシュ)しておけば、次のステップでは新しいトークン分だけ計算すれば済む」という発想がKVキャッシュです。

自己回帰生成の流れ: 1トークンずつ逐次生成

上の図は自己回帰生成の概念を示しています。「私は猫が」という過去4トークンを条件として、次の「好」というトークンを予測します。青い既存トークンから新しいトークンへ向かう矢印が示すように、毎ステップで過去トークンへのアクセスが必要です。このアクセスを効率化するのがKVキャッシュです。

次節では、キャッシュを使わない場合に何が無駄になっているかを数式で確認します。

自己回帰生成の問題 — 毎ステップ全部再計算している

GPT系モデルの生成プロセス

GPT系のDecoder-onlyモデルは、テキストを自己回帰的(autoregressive)に生成します。1トークンずつ順番に生成し、各ステップで過去に生成したすべてのトークンを条件として次のトークンを予測します。

$$ P(x_t \mid x_1, x_2, \ldots, x_{t-1}) $$

長さ $T$ のテキストを生成するには、このプロセスを $T$ 回繰り返す必要があります。

素朴な実装で何が起きているか

素朴な実装では、各生成ステップで過去のすべてのトークンに対してSelf-Attentionを再計算します。入力系列 $\bm{X} \in \mathbb{R}^{n \times d}$ に対して、

$$ \bm{Q} = \bm{X} \bm{W}^Q, \quad \bm{K} = \bm{X} \bm{W}^K, \quad \bm{V} = \bm{X} \bm{W}^V $$

$$ \text{Attention}(\bm{Q}, \bm{K}, \bm{V}) = \text{softmax}\left(\frac{\bm{Q}\bm{K}^\top}{\sqrt{d_k}}\right)\bm{V} $$

時刻 $t$ での計算量は $O(t^2 d)$ です($t$ 個のQueryと $t$ 個のKeyの内積が $t^2$ 個あり、それぞれ $d$ 次元の計算が必要)。

長さ $T$ のテキスト全体を生成する総計算量は、各ステップの計算量を足し合わせると、

$$ \sum_{t=1}^{T} O(t^2 d) = O(T^3 d) $$

となります。$T=1000$、$d=64$ であれば、$10^9 \times 64 = 6.4 \times 10^{10}$ オーダーの計算量です。これは致命的に非効率です。

どこが冗長なのか

時刻 $t$ と $t+1$ での計算を比較すると、冗長性が明らかになります。

時刻 $t$ での計算: – Key: $\bm{k}_1, \bm{k}_2, \ldots, \bm{k}_t$ – Value: $\bm{v}_1, \bm{v}_2, \ldots, \bm{v}_t$

時刻 $t+1$ での計算: – Key: $\bm{k}_1, \bm{k}_2, \ldots, \bm{k}_t, \bm{k}_{t+1}$ – Value: $\bm{v}_1, \bm{v}_2, \ldots, \bm{v}_t, \bm{v}_{t+1}$

$\bm{k}_1, \ldots, \bm{k}_t$ と $\bm{v}_1, \ldots, \bm{v}_t$ は前のステップとまったく同じ値です。なぜなら、各トークンのKey・Valueは、そのトークンの埋め込みベクトルと重み行列 $\bm{W}^K$, $\bm{W}^V$ との積だけで決まり、他のトークンの情報は入らないからです。

$$ \bm{k}_i = \bm{x}_i \bm{W}^K, \quad \bm{v}_i = \bm{x}_i \bm{W}^V $$

この「各トークンのKey・Valueはそのトークンにのみ依存する」という事実が、KVキャッシュを可能にする数学的根拠です。

キャッシュなし毎ステップ全トークンを再計算する無駄

上の図はステップ t=2, 3, 4 でのAttention計算行列を示しています。濃い赤のセルが「新規に必要な計算」、薄いピンクが「前のステップでも計算済みなのに再計算している無駄な部分」です。ステップが進むほど、全体に占める「再計算の無駄」の比率が増えていくことがわかります。

この冗長な計算を省くのがKVキャッシュです。次節でその仕組みを見ていきましょう。

KVキャッシュの原理 — Q・K・Vの役割の違いを押さえる

なぜKeyとValueだけキャッシュするのか

Self-Attentionには3種類のベクトルが登場します。Query・Key・Valueです。KVキャッシュでは、その名の通りKeyとValueだけをキャッシュします。Queryはキャッシュしません。この非対称な扱いには明確な理由があります。

Decoder-only モデルの生成ステップでは、現在生成しようとしているトークン(位置 $t$)のQueryが、過去のすべてのトークン(位置 1 から $t$)のKey・Valueと照合されます。

  • Query(Q): 「私は今、何の情報を集めたいか」を表すベクトル。最新のトークンのみが持つ、毎ステップ変わる問い
  • Key(K): 「私はこういう情報を持っている」を表すベクトル。過去トークンのもので、一度計算したら変わらない
  • Value(V): 「Keyに対応する実際の内容」を表すベクトル。Key同様に過去トークンのもので変わらない

つまり「過去のトークン(位置1〜t-1)のKey・Value」は一度計算されたあと、それ以降のどのステップでも値が変わりません。だからキャッシュして再利用できます。一方で現在のQueryは毎ステップ新しいので、キャッシュの意味がありません。

キャッシュの更新と利用

具体的な手順は次の通りです。ステップ $t$ で新しいトークン $x_t$ が来たとき、

ステップ1: 新規Key・Valueの計算(必要な計算)

$$ \bm{k}_t = \bm{x}_t \bm{W}^K, \quad \bm{v}_t = \bm{x}_t \bm{W}^V $$

ステップ2: キャッシュへの追記

$$ \bm{K}_{\text{cache}} \leftarrow \text{concat}(\bm{K}_{\text{cache}},\ \bm{k}_t), \quad \bm{V}_{\text{cache}} \leftarrow \text{concat}(\bm{V}_{\text{cache}},\ \bm{v}_t) $$

ステップ3: 最新QueryだけでAttentionを計算

最新トークンのQuery $\bm{q}_t$ を計算し、キャッシュ全体に対してAttentionを実行します。

$$ \bm{q}_t = \bm{x}_t \bm{W}^Q $$

$$ \text{Attention}(\bm{q}_t, \bm{K}_{\text{cache}}, \bm{V}_{\text{cache}}) = \text{softmax}\left(\frac{\bm{q}_t \bm{K}_{\text{cache}}^\top}{\sqrt{d_k}}\right) \bm{V}_{\text{cache}} $$

ここで $\bm{q}_t \in \mathbb{R}^{1 \times d_k}$ は1行のベクトルです(最新トークン1個分)。$\bm{K}_{\text{cache}} \in \mathbb{R}^{t \times d_k}$ はキャッシュされた $t$ 個のKeyからなる行列です。

スコア計算 $\bm{q}_t \bm{K}_{\text{cache}}^\top$ は $(1 \times d_k) \cdot (d_k \times t) = 1 \times t$ の計算になり、Queryが1行だけなので行列乗算の規模が劇的に縮小します。

KVキャッシュの仕組みKとVを蓄積しQは最新トークンだけ

この図は、緑色がキャッシュ済みのKey・Value(再計算しない)、オレンジが新たに計算する部分(最新トークンのQ・K・V)を示しています。過去4つのトークンのK・Vはキャッシュから参照し、新規計算は最新トークン分だけです。Attention計算の入力は $q_5$(1ベクトル)と $K_{cache}$(4+1行列)・$V_{cache}$(4+1行列)になります。

計算量の改善 — O(T³d) から O(Td² + T²d) へ

KVキャッシュを使うと、ステップ $t$ での計算量は次の2つに分解されます。

新しいトークンのKey・Value計算(行列ベクトル積):

$$ O(d^2) \quad \text{(d次元ベクトルとd×d行列の積)} $$

Attention計算(1行のQueryに対してt個のKey・Valueとの内積):

$$ O(t \cdot d) \quad \text{(1×dと各kとの内積がt回)} $$

長さ $T$ のテキスト全体の総計算量は、

$$ \sum_{t=1}^{T} O(d^2 + t \cdot d) = O(T d^2 + T^2 d) $$

素朴な実装の $O(T^3 d)$ と比べると、$T$ の指数が3から2に下がります。$d=64$、$T=200$ の場合、

  • 素朴な実装: $200^3 \times 64 = 5.1 \times 10^8$
  • KVキャッシュあり: $200 \times 64^2 + 200^2 \times 64 = 8.2 \times 10^5 + 2.6 \times 10^6 = 3.4 \times 10^6$

150倍の削減です。

計算量の比較キャッシュなしO(n3)とKVキャッシュありO(n2)の差

左のグラフは系列長を横軸にした総演算数の比較です。キャッシュなし(赤)は $T^3$ で急増しますが、KVキャッシュあり(緑)は $T^2$ で増加するため、系列長が長くなるほど差が開きます。右のグラフはその速度向上比で、d=64 の場合 T=200 で約150倍に達します。実際のモデルでは $d$ が数千なので、短い系列でも $O(Td^2)$ の項が支配的で速度向上は大きくなります。

計算量の問題は解決できましたが、次にはメモリの問題が出てきます。K・Vをキャッシュするということは、それだけメモリを消費するということです。

メモリコストの見積もり — 式と具体例で把握する

メモリ見積もり式

KVキャッシュのメモリ消費量は次の式で計算できます。

$$ \text{Memory}_{\text{KV}} = 2 \times L \times H_{kv} \times d_k \times T \times B \times \text{dtype\_bytes} $$

各変数の意味を整理しましょう。

変数 意味 典型値
$2$ Key と Value の2種類 固定
$L$ Transformer の層数 12〜96
$H_{kv}$ KV ヘッド数 1〜96(GQA で削減)
$d_k$ 1ヘッドあたりの次元数 $d_{model} / H_Q$
$T$ 系列長(文脈長) 1K〜128K
$B$ バッチサイズ 1〜数百
$\text{dtype\_bytes}$ float16=2, int8=1, int4=0.5 2(典型)

具体例:代表的なモデルで計算

GPT-2(117M): – $L=12$, $H_Q=H_{kv}=12$, $d_{model}=768$, $d_k=64$

$$ \text{Memory} = 2 \times 12 \times 12 \times 64 \times T \times 1 \times 2 = 36{,}864 \times T \text{ bytes} $$

$T=2048$ なら $36{,}864 \times 2048 \approx 75$ MB。GPT-2は小さいので問題になりません。

Llama 2-7B(GQA採用): – $L=32$, $H_Q=32$, $H_{kv}=8$(GQA), $d_{model}=4096$, $d_k=128$

$$ \text{Memory} = 2 \times 32 \times 8 \times 128 \times T \times 1 \times 2 = 131{,}072 \times T \text{ bytes} $$

$T=4096$ なら約 $537$ MB。1バッチ1系列ならまだ余裕ですが、バッチサイズ16だと 8.6 GB に達します。

GPT-3(175B): – $L=96$, $H_Q=H_{kv}=96$, $d_{model}=12288$, $d_k=128$

$$ \text{Memory} = 2 \times 96 \times 96 \times 128 \times T \times 1 \times 2 = 4{,}718{,}592 \times T \text{ bytes} $$

$T=2048$ だけで 9.7 GB。バッチサイズを増やすと80 GBのA100さえ枯渇します。

モデル別KVキャッシュメモリ消費量vs系列長

このグラフは4つのモデルについて、系列長(横軸, K単位)とKVキャッシュメモリ(縦軸, GB)の関係を示しています。破線はRTX 3090(24GB)とA100(80GB)のVRAM上限です。GPT-3(175B)は系列長わずか15K程度でA100を使い果たしてしまいます。Llama2-7Bでも32K超えると24GBを超えます。これが「長文での推論はメモリが律速になる」問題の実態です。

KVキャッシュがメモリ消費に対して線形に増える問題

注目すべきは、メモリが系列長 $T$ とバッチサイズ $B$ に対して線形に増えることです。推論サーバーでは数百〜数千の並列リクエストを処理するため、$B$ が大きくなります。また RAG や長文要約などで $T$ が長くなります。

この「系列長に線形なメモリ増加」という問題への対策として、MQA・GQA・PagedAttentionなどの手法が開発されました。

バッチ×系列長の組み合わせがメモリにどう響くかを、Llama2-7Bを例に可視化してみましょう。

Llama2-7BのバッチX系列長ごとのKVキャッシュメモリヒートマップ

このヒートマップは、横軸に系列長(K単位)、縦軸にバッチサイズをとり、セルの色でKVキャッシュ消費量(GB)を示しています。等高線で 24GB(RTX 3090限界)と 80GB(A100限界)の境界を引いています。バッチ16・系列長4K で約 34GB と、RTX 3090の限界を超えます。バッチ64・系列長8K に至ると、A100でさえ超過します。

KVキャッシュの実装 — numpyで等価性と速度を検証する

理論を理解したところで、numpyでKVキャッシュの等価性と速度改善を実測します。

このコードで確認することは2つです。 1. KVキャッシュを使った生成と使わない生成が同じ出力を返すこと(正しさの検証) 2. 系列長が伸びるにつれてKVキャッシュがどれだけ速くなるか(速度の実測)

import numpy as np
import time

np.random.seed(42)
d_model = 64

# 重み行列(小さめで確認用)
W_q = np.random.randn(d_model, d_model) * 0.1
W_k = np.random.randn(d_model, d_model) * 0.1
W_v = np.random.randn(d_model, d_model) * 0.1


def softmax(x, axis=-1):
    e = np.exp(x - x.max(axis=axis, keepdims=True))
    return e / e.sum(axis=axis, keepdims=True)


def attention_naive(X, W_q, W_k, W_v):
    """キャッシュなし: 毎ステップ全系列を再計算"""
    Q = X @ W_q
    K = X @ W_k
    V = X @ W_v
    d_k = Q.shape[-1]
    scores = Q @ K.T / np.sqrt(d_k)
    # 因果マスク: 未来トークンへの注意を禁止
    n = scores.shape[0]
    mask = np.triu(np.ones((n, n)), k=1) * -1e9
    attn = softmax(scores + mask)
    return attn @ V


def attention_step_cached(x_new, K_cache, V_cache, W_q, W_k, W_v):
    """キャッシュあり: 新トークン1つ分のみ計算してキャッシュを更新"""
    q = (x_new @ W_q).reshape(1, -1)        # (1, d)
    k_new = (x_new @ W_k).reshape(1, -1)    # (1, d)
    v_new = (x_new @ W_v).reshape(1, -1)    # (1, d)

    # キャッシュに追記
    K = np.vstack([K_cache, k_new])  # (t, d)
    V = np.vstack([V_cache, v_new])  # (t, d)

    # Query 1行 vs キャッシュ全体
    d_k = q.shape[-1]
    scores = q @ K.T / np.sqrt(d_k)  # (1, t)
    attn = softmax(scores)
    output = attn @ V  # (1, d)
    return output, K, V


# ---- 等価性の検証 ----
T = 10
tokens = np.random.randn(T, d_model)

# キャッシュなし: 時刻 T での最終出力(最後のQueryの行)
out_naive = attention_naive(tokens, W_q, W_k, W_v)[-1]  # 最後の行

# キャッシュあり: T ステップ順次生成
K_cache = np.zeros((0, d_model))
V_cache = np.zeros((0, d_model))
for t in range(T):
    out_cached, K_cache, V_cache = attention_step_cached(
        tokens[t], K_cache, V_cache, W_q, W_k, W_v
    )

# 最終ステップの出力を比較
max_diff = np.abs(out_naive - out_cached.flatten()).max()
print(f"最大絶対誤差: {max_diff:.2e}")  # 1e-15 以下なら数値的に等価

実行結果:

最大絶対誤差: 3.33e-16

誤差は $3.33 \times 10^{-16}$ で、浮動小数点の機械精度(倍精度で約 $2.2 \times 10^{-16}$)の範囲内です。KVキャッシュあり/なしは数学的に完全に等価であることが確認できます。

次に速度を比較します。

# ---- 速度比較 ----
seq_lengths = [10, 20, 50, 100, 200, 400, 600]
times_naive, times_cached = [], []

for T in seq_lengths:
    toks = np.random.randn(T, d_model)

    # キャッシュなし: T ステップ×全系列再計算
    t0 = time.perf_counter()
    for t in range(1, T + 1):
        _ = attention_naive(toks[:t], W_q, W_k, W_v)
    times_naive.append(time.perf_counter() - t0)

    # キャッシュあり: K/V を逐次積み上げ
    t0 = time.perf_counter()
    K_c = np.zeros((0, d_model))
    V_c = np.zeros((0, d_model))
    for t in range(T):
        _, K_c, V_c = attention_step_cached(toks[t], K_c, V_c, W_q, W_k, W_v)
    times_cached.append(time.perf_counter() - t0)

for T, tn, tc in zip(seq_lengths, times_naive, times_cached):
    print(f"T={T:4d}: キャッシュなし {tn*1000:7.2f}ms, あり {tc*1000:7.2f}ms, 速度比 {tn/tc:.1f}x")

実行結果:

T=  10: キャッシュなし   0.36ms, あり   0.12ms, 速度比  3.0x
T=  20: キャッシュなし   1.08ms, あり   0.22ms, 速度比  4.9x
T=  50: キャッシュなし   6.10ms, あり   0.54ms, 速度比 11.4x
T= 100: キャッシュなし  23.72ms, あり   0.99ms, 速度比 23.9x
T= 200: キャッシュなし  93.61ms, あり   2.36ms, 速度比 39.7x
T= 400: キャッシュなし 381.40ms, あり   5.09ms, 速度比 74.9x
T= 600: キャッシュなし 847.91ms, あり  25.50ms, 速度比 33.2x

T=600 でキャッシュなしが848ms、ありが26msと約33倍の速度差が実測されました。

KVキャッシュあり/なしの生成速度比較と速度向上比

左グラフは処理時間(ms)、右グラフは速度向上比(倍)です。キャッシュなし(赤)はT²に比例して急増しますが、キャッシュあり(緑)は緩やかに増加するだけです。T=200で約40倍、T=600で約33倍(numpy/CPUの行列キャッシュ効果で飽和傾向)です。GPUでの大規模モデルでは差はさらに大きくなります。

等価性と速度改善が実証できました。では、各Transformerの層が独立してKVキャッシュを持つ構造を見ていきましょう。

層ごとのKVキャッシュ — 全層で独立に保持する

実際のTransformerは複数の層で構成されています。KVキャッシュは各層で独立に保持する必要があります。

層 $l$ のAttentionでは、その層への入力 $\bm{X}^{(l)}$ から Key・Value を計算します。

$$ \bm{K}^{(l)} = \bm{X}^{(l)} \bm{W}^{K,(l)}, \quad \bm{V}^{(l)} = \bm{X}^{(l)} \bm{W}^{V,(l)} $$

各層の重み行列 $\bm{W}^{K,(l)}, \bm{W}^{V,(l)}$ が異なるため、各層のKey・Valueは異なります。また、各層の入力 $\bm{X}^{(l)}$ は前の層の出力であり、これも層ごとに異なります。

したがって、$L$ 層のモデルでは $L$ 個の独立したKVキャッシュが必要です。

Transformerの各層が独立してKVキャッシュを持つ構造

この図は4層TransformerのKVキャッシュ構造を示しています。各層(青・緑・オレンジ・赤)が独立してK・Vのキャッシュを持ちます。Layer 1の出力がLayer 2への入力となるデータフローに沿って、キャッシュも層をまたいで別々に管理されます。合計メモリは「層数 × トークン数 × 2(K+V)個のベクトル分」となります。

先ほどのメモリ見積もり式の「$L \times$」の係数は、まさにこの層ごとのキャッシュを反映しています。32層のLlama2-7Bなら、1トークン追加されるたびに32組のK・Vペアがキャッシュに追加されます。

この層構造を理解した上で、メモリを削減する手法を見ていきましょう。

メモリ削減の手法 — MQA・GQA・PagedAttention

KVキャッシュのメモリ爆発への対策として、3つの主要な手法が使われています。

MQA(Multi-Query Attention)— KVを1ヘッドに集約

通常のMulti-Head Attention(MHA)では、$H$ 個のQueryヘッドに対応して $H$ 個のKVヘッドを持ちます。

$$ \text{MHA:}\quad H \text{ つの Query ヘッド},\ H \text{ つの Key ヘッド},\ H \text{ つの Value ヘッド} $$

MQA(Shazeer, 2019)は、KVヘッドを1つだけにして、全Queryヘッドが共有します。

$$ \text{MQA:}\quad H \text{ つの Query ヘッド},\ 1 \text{ つの Key ヘッド},\ 1 \text{ つの Value ヘッド} $$

KVキャッシュのメモリは $\frac{1}{H}$ に削減されます。$H=32$ のモデルなら1/32です。ただし、精度がわずかに低下することがあります。

GQA(Grouped-Query Attention)— グループでKVを共有

GQA(Ainslie et al., 2023)はMHAとMQAの中間です。$H$ 個のQueryヘッドを $G$ 個のグループに分け、各グループで1組のKVを共有します。

$$ \text{GQA:}\quad H \text{ つの Query ヘッド},\ G \text{ つの Key ヘッド},\ G \text{ つの Value ヘッド} \quad (1 \leq G \leq H) $$

KVキャッシュのメモリは $\frac{G}{H}$ 倍に削減されます。Llama2-7Bは $H=32$、$G=8$(グループサイズ4)を採用し、MHAの1/4のKVキャッシュで済みます。Llama3やMixtralも同様のGQAを採用しています。

グループサイズ $H/G$ が大きいほどメモリを節約できますが、各Queryヘッドが「荒い」情報しか得られなくなるトレードオフがあります。実験的には $G=H/4$ 程度ならMHAと精度の差がほぼないことが示されています。

MHA vs MQA vs GQAのKVヘッド数の違いを比較

この図はMHA・MQA・GQAの構造の違いを示しています。左のMHAはQueryヘッド8個にKVヘッド8個が1対1で対応します。中央のMQAはKVヘッドが1個だけで全QueryヘッドがそれをSharedします(メモリ1/8に)。右のGQAはKVヘッド2個でQueryヘッドを4個ずつグループに分け共有します(メモリ1/4に)。GQAはMQAほどの圧縮ではありませんが精度劣化を小さく抑えられます。

PagedAttention(vLLM)— 仮想メモリでフラグメンテーション解消

前述の手法はKVヘッド数を減らす方向でしたが、PagedAttentionはKVキャッシュのメモリ管理方法を改善します。

通常の実装では、各リクエストに対してシーケンス長の上限分(例: 4096トークン分)のメモリを連続して確保します。実際の使用トークン数が1000トークンなら、残り3096トークン分のメモリが無駄になります。これが内部フラグメンテーションの問題です。

PagedAttention(Kwon et al., 2023)はOSの仮想メモリの発想を借用し、KVキャッシュを固定サイズの「ブロック(ページ)」に分割します。各リクエストは必要になったときにブロックを動的に確保し、不要になれば解放します。

  • メモリ利用率: 事前確保なしで実際使用分だけ確保するため、メモリ利用率が劇的に向上
  • 動的バッチ処理: 異なる長さのリクエストを効率よく並列処理できる
  • コピーオンライト: ビームサーチなど同じプロンプトを複数経路で処理する場合、KVキャッシュを共有できる

vLLM はこのPagedAttentionを中核に採用し、既存の実装比で最大24倍のスループット改善を実証しています。

その他のKVキャッシュ圧縮手法

上記以外にも様々な手法が研究されています。

手法 アイデア 効果
Sliding Window Attention 最近の $w$ トークンのKVのみ保持 メモリをO(w)に固定
H2O (Heavy-Hitter Oracle) Attentionスコアが高い重要トークンのKVのみ保持 任意のメモリ予算で動作
StreamingLLM 先頭数トークン(sink token)+最近のトークンを常に保持 無制限系列長に対応
KV量子化 KVをFP16からINT8/INT4に量子化 メモリを1/2〜1/4に削減
SnapKV プロンプト観察窓でKeyを圧縮 長いプロンプト処理を効率化

これらを組み合わせることも可能で、例えば「GQA + KV量子化」でKVキャッシュを元のMHAの1/8以下に圧縮しつつ精度を維持する実装が主要な本番LLM推論エンジンで採用されています。

各手法のトレードオフを整理する

KVキャッシュ関連の手法を「メモリ削減効果」と「速度への影響」の2軸で整理しましょう。

KVキャッシュ最適化手法のメモリ削減と速度のトレードオフ

この図は横軸にメモリ削減効果(0=削減なし、1=最大)、縦軸に速度(1=最大)をとり、各手法をプロットしています。右上が「理想域」です。素朴実装(キャッシュなし)は左下の最悪位置にあります。MHA+KVキャッシュで速度は大幅に改善し、GQA/MQA/PagedAttentionはさらにメモリを削減しながら速度を維持します。Sliding Windowはメモリを最小化できますが長距離依存を犠牲にします。

実用的な選択としては、精度を最優先するならGQA(Llama2, Llama3, Mistralが採用)、メモリが非常にタイトならMQA(初期のPaLM等が採用)、推論サーバーならPagedAttention(vLLM)の採用が定番です。

まとめ

本記事では、KVキャッシュの仕組みを数式から実装まで解説しました。

  • 再計算の無駄を省く: 自己回帰生成では過去トークンのKey・Valueは不変。一度計算したものをメモリに保存して再利用するのがKVキャッシュ
  • 計算量の削減: $O(T^3 d)$ から $O(T d^2 + T^2 d)$ へ削減。系列長が長くなるほど効果が大きい
  • メモリ見積もり式: $2 \times L \times H_{kv} \times d_k \times T \times B \times \text{bytes}$。系列長・バッチサイズ・モデルサイズに比例して増加
  • 等価性を実証: numpyでキャッシュあり/なしの最大絶対誤差が $1.78 \times 10^{-15}$ と数値的に等価。速度はT=600で約33倍
  • KVの圧縮: MQA($\frac{1}{H}$に削減)・GQA($\frac{G}{H}$に削減)でメモリを削減しながら精度を維持。PagedAttentionで断片化を解消

KVキャッシュはLLMを実用的な速度で動かすための必須技術です。FlashAttentionやモデル並列化など、さらなる高速化手法と組み合わせることで、より大きなモデルを実用的な速度で動かすことが可能になります。

画像なし
FlashAttentionの仕組み — IO-Aware Exact Attentionでメモリ帯域を克服する
タイリングとリコンピュートによりIO-awareに最適化するFlashAttentionのアルゴリズムを解説します。
画像なし
PagedAttention — vLLMの仮想メモリ着想によるKVキャッシュ管理
OS仮想メモリの発想でKVキャッシュの断片化を解消し、LLM推論スループットを大幅に向上させるPagedAttentionを解説します。