一句话总结:KV Cache 在每一层保存历史 token 的 Key 和 Value。Decode 时仍要为新 token 计算新的 Q、K、V,但不再重算旧 token 的 K、V。
22.1 重复计算藏在哪里?
第 20 章的教学版生成器每次都把末尾窗口重新送进模型。结果当然没错,但旧 token 已经算过的隐藏状态和 K/V 也跟着一遍遍重算。
继续用原来的例子。为了看清步骤,下面暂时把每个汉字当成一个教学单位;真实 tokenizer 未必这样切分:
输入:中华人民
第 1 步:中华人民 → 共
第 2 步:中华人民共 → 和
第 3 步:中华人民共和 → 国
没有 Cache 时,第 3 步会重新跑“中、华、人、民、共、和”的每一层。可是因果模型里,旧位置看不到未来;只要模型参数、旧前缀和位置没有改变,那些旧位置在各层产生的 K/V 也不会改变。
这正是可以省掉的部分。
和第 20 章的边界:第 20 章的固定正弦位置会在裁出新窗口后从 0 重新编号。窗口一旦滑动,保留 token 的位置也变了,它们的 hidden state 和 K/V 不再是原来的值。要和那个教学实现严格等价,窗口滑动时必须重算整个窗口;生产级 sliding-window Cache 需要模型本身支持的位置与缓存策略。
22.2 Cache 里存什么?每一步又算什么?
对第 层的一个 Attention 模块:
Prefill 一次处理完整 prompt,并在每一层建立:
随后 Decode 到位置 时,该层仍要从新 hidden state 计算:
再把新 K/V 追加进去:
最后只算新 Query 的这一行:
为什么不缓存旧 Q?
旧位置的输出早已用完,未来也不会回头重算它们的 Attention row;新位置只需要一个新 Query。相反,每个新 Query 都要和历史 Key 配对,再从历史 Value 取信息,所以历史 K/V 必须留下。
边界:Encoder–Decoder 模型还可以缓存 Decoder cross-attention 使用的 Encoder K/V。这里先讲 Decoder-only causal self-attention,不把“KV Cache 只存在于 Decoder”说成所有架构的绝对规则。
22.3 计算量:先说清楚在数哪一部分
设当前上下文长度为 。只看某一层的 Attention score 和 Value 加权:
| 当前 Decode step | 无 Cache | 有 Cache |
|---|---|---|
| Query–Key 配对 | ||
| Attention 主项 | ||
| 旧 token 的 K/V projection | 全部重算 | 不重算 |
在还没有碰到固定上下文窗口的上限时,如果从很短的前缀一路生成到长度 ,只对这部分 Attention 求和,可以简写为:
但这不是“整个模型固定快 N 倍”。每一步仍要为新 token 跑所有层的 Q/K/V projection、输出 projection、FFN、Norm 和 LM head,还要从显存读取不断增长的 Cache。真实加速取决于 prompt/output 长度、batch、模型、dtype、kernel、硬件和计时范围;不能把旧版某次 11 秒 vs 56 秒写成普遍保证。
22.4 内存公式:关键是 n_kv_heads
对于常见的密集 Cache,理论张量字节数是:
- :Key 与 Value;
- :batch(Beam Search 的展开也会影响实际 Cache);
- :已缓存长度;
- :层数;
- :KV heads 数,不一定等于 Query heads;
- :每个 head 的维度;
- :每个元素的字节数。
手算 Llama 2 7B
Llama 2 7B 有 32 层、32 个 KV heads、head dimension 128。按 FP16/BF16 的 2 bytes、batch 1、4096 tokens:
每 token = 2 × 32 × 32 × 128 × 2
= 524,288 bytes = 512 KiB
4096 tokens = 2,147,483,648 bytes = 2 GiB
这组数也顺便纠正一个常见口误:7B 权重按 2 bytes/parameter 约 14 GB(十进制),那是 FP16/BF16 的量级;FP32 约 28 GB,还没算运行时开销。
NVIDIA A10 有 24 GB GDDR6。即使先粗略减去约 14 GB 权重,也不能宣布剩余部分全归 KV Cache:框架、allocator、临时 workspace、activations 和其他 buffer 都要占空间。用“10 GB ÷ 2 GiB”得出“正好 5 路 4K”还混用了 GB 与 GiB;即使先统一单位,也只是不含开销的纸面计算,不是部署承诺。
同理,对这一个 MHA 配置做纯内存外推,32K 是 16 GiB、128K 是 64 GiB;这不代表原始 4K 模型突然获得了 128K 的位置与质量能力。
22.5 Prefill 与 Decode 是两种工作负载
| 阶段 | Prefill | Decode |
|---|---|---|
| 输入 | 完整 prompt | 新增的一个或少量 token |
| Cache | 为每层写入整段 K/V | 每层追加新 K/V |
| 并行性 | 序列位置可并行 | token 之间必须串行 |
| 常见指标 | TTFT(首 token 延迟) | TPOT(每个输出 token 延迟) |
| 常见瓶颈 | 长 prompt 常偏计算密集 | 小 batch 常偏带宽/延迟受限 |
“常见”不等于定律。短 prompt、很大的 batch、不同量化和 kernel 都会改变瓶颈。FlashAttention 优化一次 Attention 内的 I/O;KV Cache 避免 Decode steps 之间重算历史 K/V,二者解决的是不同层次的问题。
22.6 多轮对话不是自动白送
跨轮复用需要满足几个条件:
- 上一轮 Cache 还在同一个服务端会话里;
- 新请求的 token 序列以旧序列为完全相同的前缀;
- attention mask、position IDs /
cache_position接着走; - 模型、adapter 和影响 Attention 的配置没有改变。
如果改了 system prompt、chat template、历史消息或中间 token,最早变化点之后的 Cache 都失效。很多 API 本来就是无状态的,会重新 Prefill;“界面看起来像同一段对话”并不能证明后端复用了 KV。
流式“打字机效果”也不只是 Cache 的功劳:自回归模型本来就逐 token 生成,服务端还要把 token 逐步传给客户端。KV Cache 只是让每一步少做大量重复工作。
22.7 用可执行代码证明等价性
下面是一层、无 bias、无 、无位置编码的教学实现。它只验证 Cache 的核心递推:比较完整 causal Attention 与逐 token Cache。它不是生产 kernel,但会真正执行等价性断言。
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。当前版本可通过 use_cache=False 禁用,也提供 Dynamic、Static、offloaded、quantized 等 Cache 策略;具体默认值和兼容矩阵应以安装版本的文档为准,不在书里写死一组虚构 benchmark。
22.8 优化方向与容易混淆的地方
- MQA / GQA:减少 ,直接缩小每个 token 的 Cache;第 23 章细讲。
- KV quantization:降低每个元素的 bytes,但要带上量化元数据、kernel 支持和质量误差一起测。
- Offloading:把部分层的 Cache 搬到 CPU,省 GPU 内存但增加传输成本。
- PagedAttention:主要减少分配碎片、按需分块并支持共享;它不会凭空把每个仍然存活的 K/V 元素变少。
- Sliding / eviction:只有架构本身支持窗口,或使用经过验证的淘汰方法时才能安全做。简单扔掉最旧 K/V 会改变模型看到的上下文;StreamingLLM 也正是因为普通窗口会失效,才保留 attention sinks。
- 标准训练:完整 teacher-forced 序列只前向一次,还需要梯度;通常不使用生成式 KV Cache。不要把“通常禁用”说成理论上任何训练都绝不可能有状态缓存。
22.9 本章总结
- Cache 的是每一层历史 K/V;新 token 的 Q/K/V 仍然要算;
- 对单个 Decode step 的 Attention,矩阵从 变成 ,不代表整模型固定快 倍;
- 内存公式必须用 ,并把 batch、长度、层数、head dimension 和 dtype 算全;
- Llama 2 7B 的 4K FP16/BF16 Cache 是每个 batch element 约 2 GiB;14 GB 权重不是 FP32;
- 多轮复用要求 exact token prefix 与正确位置状态;
- Prefill 决定 TTFT,Decode 决定 TPOT,FlashAttention 与 KV Cache 解决不同问题。
本章交付物
- 能解释为什么新 token 仍要计算 Q、K、V
- 能区分 Attention 局部复杂度与整模型延迟
- 能用
n_kv_heads手算 Cache 字节数 - 能说出跨轮 Cache 何时可以复用、何时必须失效
- 能运行教学代码并验证 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 换取更小的 Cache 与更高的 Decode 吞吐。