一文要約:KV Cacheは各層で過去のtokenのKeyとValueを保存します。decode時も新しいtokenのQ、K、Vは計算しますが、過去のtokenのK/Vは再計算しません。
22.1 重複計算はどこに隠れているか
第20章の教材用の生成器は、stepごとに末尾のwindowをモデルへ再入力します。出力は正しいのですが、過去のtokenのhidden stateとK/Vまで何度も計算します。
蒔絵の工房を例にしましょう。次の角括弧は仕組みを見やすくするための単位で、実Tokenizerの分割を主張するものではありません。
入力:[蒔絵] [の] [粉筒]
Step 1:[蒔絵] [の] [粉筒] → [から]
Step 2:[蒔絵] [の] [粉筒] [から] → [金粉]
Step 3:[蒔絵] [の] [粉筒] [から] [金粉] → [を蒔く]
Cacheがなければ、step 3で古い単位をすべて各層へ通し直します。しかし因果モデルの過去位置は未来のtokenを見ません。モデルの重み、過去のprefix、位置が変わらなければ、その位置が各層で作ったK/Vも変わりません。
ここが再利用できる計算です。
第20章との境界:第20章の固定正弦位置は、新しいwindowへ切り出すたびに0から振り直されます。windowが滑ると、残ったtokenの位置も変わるため、hidden stateとK/Vは以前の値のままではありません。その教材実装と厳密に同じ結果を得るには、windowが滑った時点で全体を再計算する必要があります。本番のsliding-window Cacheには、モデルが対応する位置とキャッシュの方針が必要です。
22.2 何を保存し、毎Step何を計算するのか
第 層のAttention moduleを考えます。
PrefillはPrompt全体を一度処理し、各層に次を作ります。
Decode位置 では、その層は新しいHidden stateから今も次を計算します。
新しいK/Vを追加します。
そして新しいQueryの1行だけを計算します。
なぜ過去のQは保存しないのか
過去位置の出力はすでに使い終わり、将来のStepで過去のAttention rowを再計算する必要はありません。新しい位置に必要なのは新しいQuery一つです。一方、そのQueryは過去すべてのKeyと照合し、Valueから情報を取るため、過去K/Vは残します。
境界:Encoder–Decoder modelでは、DecoderのCross-attentionが使うEncoder K/VもCacheできます。本章はDecoder-onlyのCausal self-attentionを扱います。「KV CacheはDecoderにしか存在しない」を全Architectureの絶対則にはしません。
22.3 計算量:何を数えた値かを先に言う
現在のContext lengthを とします。1層のAttention scoreとValue weightingだけを見ると:
| 現在のDecode step | Cacheなし | Cacheあり |
|---|---|---|
| Query–Keyの組 | ||
| Attentionの主項 | ||
| 過去TokenのK/V projection | すべて再計算 | 再利用 |
固定context windowの上限に達する前に、短いprefixから長さ まで生成し、このAttention部分だけを合計すると:
これは「モデル全体が常にN倍速い」という意味ではありません。各stepで新しいtokenを全層のQ/K/V投影、出力投影、FFN、Norm、LM headへ通し、増え続けるCacheをメモリから読みます。実際の加速率はpromptとoutputの長さ、batch、モデル、dtype、kernel、ハードウェア、計測範囲で変わります。過去の一測定で11秒対56秒だったとしても、普遍的な保証にはできません。
22.4 メモリ:重要なのは n_kv_heads
一般的なDense Cacheの理論Tensor bytesは:
- :KeyとValue;
- :Batch size(Beam展開も実Cacheへ影響します);
- :Cache済みSequence length;
- :Layer数;
- :KV heads数。Query heads数と同じとは限らない;
- :Head dimension;
- :1要素のbytes。
Llama 2 7Bを手計算する
Llama 2 7Bは32層、32 KV heads、Head dimension 128です。FP16/BF16の2 bytes、Batch 1、4096 tokensなら:
1 token = 2 × 32 × 32 × 128 × 2
= 524,288 bytes = 512 KiB
4096 tokens = 2,147,483,648 bytes = 2 GiB
ここでよくある単位間違いも直せます。7B weightsを2 bytes/parameterで置いた約14 GBはFP16/BF16の規模です。FP32は約28 GBで、Runtime overheadはまだ含みません。
NVIDIA A10は24 GBのGDDR6を搭載します。概算で14 GBの重みを引いても、残りをすべてKVに使えるわけではありません。フレームワーク、allocator、一時workspace、activation、その他のbufferもメモリを使います。「10 GB ÷ 2 GiBで4K requestがちょうど5件」という計算はGBとGiBも混ぜています。単位を揃えても、それはオーバーヘッドを含まない紙上の計算であり、デプロイ時の保証ではありません。
この MHA構成だけをMemory計算で外挿すれば、32Kで16 GiB、128Kで64 GiBです。しかしその計算だけで、元の4K Modelが正しい128KのPosition behaviorや品質を得るわけではありません。
22.5 PrefillとDecodeは別のワークロード
| Phase | Prefill | Decode |
|---|---|---|
| Input | Prompt全体 | 新しい1個または少数Token |
| Cache | 各層へPromptのK/Vを書く | 各層へ新K/Vを追加 |
| Parallelism | Prompt位置を並列処理できる | Output tokensは直列 |
| よく使う指標 | TTFT(最初のTokenまで) | TPOT(Output tokenごとの時間) |
| よくある律速 | 長いPromptは計算寄り | 小Batchは帯域/Latency寄り |
「よくある」であって法則ではありません。短いPrompt、大きなBatch、Quantization、Kernelで境界は動きます。FlashAttentionは一回のAttention内部のI/Oを減らし、KV CacheはDecode stepsをまたぐ過去K/Vの再計算を避けます。解く階層が違います。
22.6 Multi-turnでの再利用は条件付き
Turnをまたぐ再利用には次が必要です。
- 前回のCacheが同じServer-side sessionに残っている;
- 新しいToken列が完全に同じToken prefixから始まる;
- Attention maskとPosition IDs /
cache_positionが正しく続く; - Model、Adapter、Attentionに影響する設定が変わっていない。
System prompt、Chat template、履歴、途中Tokenを変えれば、最初の変更点以降のCacheは無効です。多くのAPIはStatelessで、送信された履歴を再びPrefillします。画面上で一つの会話に見えても、Server-side KV reuseの証明にはなりません。
「タイプライター表示」もCacheだけが作るものではありません。Autoregressive modelは元からTokenを順番に生成し、ServerもClientへ順次送ります。KV Cacheは各Stepの重複作業を大きく減らす役目です。
22.7 実行できるコードで等価性を証明する
次は1層、biasなし、 なし、位置符号化なしの教材実装です。完全な因果Attentionとtokenごとのcached decodeを比べ、Cacheの中心的な漸化式だけを確認します。本番のkernelではありませんが、等価性のassertionは実際に動きます。
import math
import torch
def split_heads(x, n_heads):
batch, time, width = x.shape
if width % n_heads != 0:
raise ValueError("width must be divisible by n_heads")
head_dim = width // n_heads
return x.view(batch, time, n_heads, head_dim).transpose(1, 2)
def merge_heads(x):
batch, n_heads, time, head_dim = x.shape
return x.transpose(1, 2).contiguous().view(
batch, time, n_heads * head_dim
)
def project(x, w_q, w_k, w_v, n_heads):
return (
split_heads(x @ w_q, n_heads),
split_heads(x @ w_k, n_heads),
split_heads(x @ w_v, n_heads),
)
def full_causal_attention(x, w_q, w_k, w_v, n_heads):
q, k, v = project(x, w_q, w_k, w_v, n_heads)
scores = q @ k.transpose(-2, -1) / math.sqrt(q.size(-1))
time = x.size(1)
future = torch.triu(
torch.ones(time, time, dtype=torch.bool, device=x.device),
diagonal=1,
)
scores = scores.masked_fill(future, -torch.inf)
return merge_heads(torch.softmax(scores, dim=-1) @ v)
def cached_step(x_new, cache, w_q, w_k, w_v, n_heads):
if x_new.size(1) != 1:
raise ValueError("cached_step expects exactly one new token")
q, k_new, v_new = project(x_new, w_q, w_k, w_v, n_heads)
if cache is None:
k, v = k_new, v_new
else:
k_past, v_past = cache
k = torch.cat((k_past, k_new), dim=-2)
v = torch.cat((v_past, v_new), dim=-2)
# The query is the newest position, so every cached position is visible.
scores = q @ k.transpose(-2, -1) / math.sqrt(q.size(-1))
output = merge_heads(torch.softmax(scores, dim=-1) @ v)
return output, (k, v)
if __name__ == "__main__":
torch.manual_seed(22)
x = torch.randn(2, 7, 12, dtype=torch.float64)
weights = [torch.randn(12, 12, dtype=x.dtype) for _ in range(3)]
with torch.inference_mode():
expected = full_causal_attention(x, *weights, n_heads=3)
cache = None
pieces = []
for position in range(x.size(1)):
output, cache = cached_step(
x[:, position:position + 1],
cache,
*weights,
n_heads=3,
)
pieces.append(output)
actual = torch.cat(pieces, dim=1)
torch.testing.assert_close(actual, expected, atol=1e-10, rtol=1e-10)
assert cache[0].shape == (2, 3, 7, 4)
assert cache[1].shape == (2, 3, 7, 4)
print("cached decode matches full causal attention")
完全なTransformerではCacheは層ごとです。新Tokenが第1層を通り、その新しいHidden stateが第2層へ入り、順に上がります。Token embeddingから全層分のK/Vを最初に一度だけ投影することはできません。
Hugging Face generate()は通常、past_key_values、Mask、cache_positionを管理します。現在のReleaseはuse_cache=Falseで無効化でき、Dynamic、Static、Offloaded、Quantizedなどの方式を持ちます。互換性一覧はインストール済みVersionのDocumentationで確認し、出所のないBenchmarkを本へ固定しません。
22.8 最適化方向と混同しやすい点
- MQA / GQA: を減らし、1 TokenあたりのCacheを直接小さくします。第23章で扱います。
- KV quantization:1要素のbytesを減らしますが、Metadata、Kernel support、品質誤差まで測ります。
- Offloading:一部LayerのCacheをCPUへ移し、GPUメモリと転送コストを交換します。
- PagedAttention:主にメモリ割り当ての断片化を減らし、必要時のBlock allocationと共有を可能にします。生きているK/V要素そのものを消す技術ではありません。
- Sliding / eviction:アーキテクチャ自体がsliding windowに対応しているか、検証済みのeviction手法が新しい意味論を定める場合だけ安全です。最も古いK/Vを盲目的に捨てれば、モデルが見る文脈は変わります。StreamingLLMがattention sinksを残すのも、単純な「直近tokenだけのwindow」が失敗するためです。
- 標準Training:完全なTeacher-forced sequenceを一度処理しGradientが必要なので、Generation用KV Cacheは通常無効です。「通常無効」は「どんなTrainingにも状態Cacheは理論上あり得ない」より正確です。
22.9 本章のまとめ
- 各層の過去K/Vを保存し、新TokenのQ/K/Vは今も計算する。
- 1 Decode stepのAttention matrixは から になるが、モデル全体が自動で 倍速くなるわけではない。
- メモリ式は を使い、Batch、Length、Layers、Head dimension、dtypeを含める。
- Llama 2 7Bの4K FP16/BF16 CacheはBatch elementごとに約2 GiB。14 GB weightsはFP32ではない。
- Multi-turnでの再利用にはExact token prefixと正しいPosition stateが必要。
- PrefillはTTFT、DecodeはTPOTに効き、FlashAttentionとKV Cacheは別の問題を解く。
章末チェックリスト
- 新Tokenで新Q、K、Vが必要な理由を説明できる
- Attention complexityとEnd-to-end latencyを分けられる
-
n_kv_headsを使ってCache bytesを計算できる - 会話Cacheを再利用できる条件と無効条件を言える
- 教育Codeを実行し、Cached decodeと完全Causal Attentionを一致させられる
一次資料
- Hugging Face Transformers: Caching
- Hugging Face Transformers: Cache strategies
- Llama 2
- PagedAttention / vLLM
- StreamingLLM
- NVIDIA A10 specifications
次章予告
なぜ はQuery heads数と同じとは限らないのでしょうか。次章ではMHA、MQA、GQAを比較し、K/V共有がModel capacityと小さなCache、高いDecode throughputをどう交換するかを見ます。