每生成一个词为何又重算全文?KV Cache 的增量解码
从自回归推理的重复前缀计算出发,追踪 Key/Value 缓存的张量生长、等价性、复杂度与工程边界。
上一篇组装了 Transformer Decoder 的三条信息流:因果自注意力读目标前缀,交叉注意力读源记忆,FFN 逐位置加工特征。透明的推理循环每次都把完整前缀重新送入模型,虽然答案正确,却在不断重做已经做过的投影和注意力。
Key–Value Cache(KV Cache)的核心不是改变模型数学,而是保存每层历史 token 已经算好的 Key 和 Value。新一步只投影新 token,用它的 Query 读取“历史缓存 + 当前 token”。本文只讲透这个增量数据流。
01 无缓存解码到底重复了什么?#
设 prompt 长度为 ,已生成 个 token。无缓存方式在第 步将长度 的整个序列重新前向:
第 1 步: [prompt] ─► 重算 prompt 所有 Q/K/V
第 2 步: [prompt, y1] ─► 又重算 prompt 与 y1
第 3 步: [prompt, y1, y2] ─► 又重算 prompt、y1、y2
...
第 t 步: [prompt, y1, ..., y(t-1)] ─► 重算全部前缀text对一层自注意力,每个历史 token 的 Key/Value 只由它进入该层时的隐状态与已固定权重决定。在因果模型中,后来的 token 不能反过来改写历史位置的隐状态,因此历史 K/V 可以复用。
若每一步都重算长度 的完整自注意力,仅分数矩阵工作量的累积就约为:
用 KV Cache 后,每步只有一个新 Query 与 个 Key 计算分数:
这是简化的单层注意力量级,不包括 FFN、投影、内核常数、prompt 预填充和内存带宽。它说明缓存消除了哪类重复,不代表真实延迟会按同一比例下降。
02 每层缓存里究竟放什么?#
对 个头、每头宽度 的自注意力,在第 步开始时:
| 张量 | 形状 | 含义 |
|---|---|---|
| 新 token 隐状态 | [N,1,D] | 当前层只处理一个新位置 |
| 新 Query | [N,H,1,d] | 询问历史与当前信息 |
| 新 Key/Value | [N,H,1,d] | 把当前 token 加入可被未来查询的记忆 |
| 更新后 Key Cache | [N,H,t,d] | 位置 1 到 的全部 Key |
| 更新后 Value Cache | [N,H,t,d] | 位置 1 到 的全部 Value |
| 当步注意力分数 | [N,H,1,t] | 一个新 Query 读取 个 Key |
| 当步输出 | [N,1,D] | 仅产生最新位置的表示 |
历史缓存
K_cache [N,H,t-1,d] ─┐
V_cache [N,H,t-1,d] ─┤
├─ append new K/V ─► K,V [N,H,t,d]
新 token x_t [N,1,D] │
└─ Q/K/V projection ─┘
│
└─ q_t [N,H,1,d] @ K^T [N,H,d,t]
▼
scores [N,H,1,t]
▼ softmax @ V
output [N,H,1,d] ─► [N,1,D]text缓存是每一层各自一份。第 7 层的 Key/Value 来自第 7 层的输入表示,不能与第 3 层共用。若有 层,标准多头注意力缓存的元素数约为:
开头的 2 分别对应 Key 和 Value。若用 fp16/bfloat16,每元素通常 2 字节;批大小、层数和上下文长度都会线性放大缓存。
03 为什么缓存 K/V,通常不缓存 Query?#
在因果生成的第 步,我们只需要计算最新位置的输出。它的 Query 会读所有 。上一步的 已经用于产生上一位置输出,新 token 不会回头重算该输出,因此历史 Query 没有再次被使用。
时间 t-1: q_(t-1) 读 [k_1 ... k_(t-1)] ─► 输出已完成
时间 t: q_t 读 [k_1 ... k_(t-1), k_t]
未来 t+1: q_(t+1) 读 [k_1 ... k_t, k_(t+1)]textKey/Value 是未来 Query 反复查询的记忆,Query 是当步一次性的读取请求。“不缓存 Q”指不为未来步保留历史 Query;当前内核在计算期间当然仍需要 。
04 用一维头手算两步缓存#
为了可手算,令单头宽度 ,缩放因子为 1。预填充后缓存为:
第 3 个 token 投影得到 。追加后:
分数为 [1,2,0],softmax 近似为 [0.245,0.665,0.090],所以:
第 4 步只需追加 并计算新 。已缓存的 [1,2,0] 和 [10,20,30] 不变。若无缓存,模型会从前缀隐状态再算一次这三个 Key/Value,最终 不会因此更正确,只会更费计算。
05 预填充与逐 token 解码是两个阶段#
KV Cache 推理常分成:
- 预填充(Prefill):一次输入整个 prompt,使用 causal attention 并行计算其表示,同时写入每层 prompt K/V。
- 解码(Decode):每次只输入一个新 token,追加该 token 的 K/V,再用新 Query 读全部缓存。
prefill:
prompt [N,P] ─► 并行 causal forward ─► 每层 K/V [N,H,P,d] + 下一 token logits
decode step 1:
y1 [N,1] ─► 新 K/V ─► cache length P+1 ─► y2 logits
decode step 2:
y2 [N,1] ─► 新 K/V ─► cache length P+2 ─► y3 logitstext预填充往往更偏计算密集,因为可以并行处理许多 prompt token;单 token decode 往往更受缓存读写和内存带宽限制。性能报告应分开首 token 延迟(Time to First Token)与后续 token 间隔,不要只给一个平均数。
06 用 PyTorch 2.13 SDPA 写透明的动态 KV Cache#
下面的实现只展示单层 causal self-attention。为便于理解,它用 torch.cat 追加缓存;生产实现应避免每步重新分配并复制已有缓存。
from typing import NamedTuple
import torch
from torch import nn
from torch.nn import functional as F
class KVCache(NamedTuple):
key: torch.Tensor # [N,H,S,d]
value: torch.Tensor # [N,H,S,d]
class CachedCausalSelfAttention(nn.Module):
def __init__(self, model_dim: int, num_heads: int, dropout: float = 0.0):
super().__init__()
assert model_dim % num_heads == 0
self.num_heads = num_heads
self.head_dim = model_dim // num_heads
self.dropout = dropout
self.qkv = nn.Linear(model_dim, 3 * model_dim)
self.out = nn.Linear(model_dim, model_dim)
def _split_heads(self, x: torch.Tensor) -> torch.Tensor:
n, length, _ = x.shape
x = x.view(n, length, self.num_heads, self.head_dim)
return x.transpose(1, 2) # [N,H,L,d]
def forward(
self,
x: torch.Tensor, # prefill: [N,P,D]; decode: [N,1,D]
cache: KVCache | None = None,
use_cache: bool = False,
) -> tuple[torch.Tensor, KVCache | None]:
n, length, model_dim = x.shape
projected = self.qkv(x)
q_raw, k_raw, v_raw = projected.chunk(3, dim=-1)
q = self._split_heads(q_raw)
new_k = self._split_heads(k_raw)
new_v = self._split_heads(v_raw)
if cache is None:
key, value = new_k, new_v
# prefill 长度 > 1 时必须阻止 prompt 内的未来泄漏。
causal = length > 1
else:
assert length == 1, "增量 decode 每次只应输入一个新 token"
assert cache.key.shape[:-2] == new_k.shape[:-2]
key = torch.cat([cache.key, new_k], dim=-2)
value = torch.cat([cache.value, new_v], dim=-2)
# Key 只包含过去+当前,根本没有未来位置可见。
causal = False
attended = F.scaled_dot_product_attention(
q, key, value,
is_causal=causal,
dropout_p=self.dropout if self.training else 0.0,
) # [N,H,L,d]
merged = attended.transpose(1, 2).contiguous().view(n, length, model_dim)
output = self.out(merged)
next_cache = KVCache(key, value) if use_cache else None
return output, next_cachepythonPyTorch 2.13 的 scaled_dot_product_attention ↗ 接收 Query [N,...,H_q,L,E]、Key [N,...,H,S,E]与 Value [N,...,H,S,E_v],输出 [N,...,H_q,L,E_v]。当前 API 仍标记为 Beta,并会根据输入和硬件选择可用后端。
函数式 SDPA 会按 dropout_p 无条件应用 dropout,不会自动读取 module.eval();因此代码必须在 eval 时显式传 0.0。
07 如何证明缓存没有改变结果?#
在 dropout 关闭、相同精度与后端下,完整因果前向的每个位置输出,应与“预填充首个 token + 逐 token 缓存”的对应输出接近:
torch.manual_seed(0)
attention = CachedCausalSelfAttention(16, 4, dropout=0.0).eval()
x = torch.randn(2, 6, 16)
with torch.inference_mode():
full, _ = attention(x, cache=None, use_cache=False)
pieces = []
cache = None
for position in range(x.size(1)):
step, cache = attention(
x[:, position:position + 1],
cache=cache,
use_cache=True,
)
pieces.append(step)
incremental = torch.cat(pieces, dim=1)
torch.testing.assert_close(incremental, full, rtol=1e-5, atol=1e-6)
assert cache is not None
assert cache.key.shape == (2, 4, 6, 4)
assert cache.value.shape == (2, 4, 6, 4)python融合内核、低精度累加和不同计算顺序可使末位出现微小数值差异,因此使用容差比较,不要用逐 bit 相等。但若误差很大或随长度急剧放大,应先查位置编码偏移、层缓存对应、mask 与 cache 追加轴。
08 位置索引为什么必须跟着缓存长度走?#
增量步只输入 [N,1] 的 token,但它不是“位置 0”。若缓存已有 个位置,新 token 的绝对位置应是 :
past_length = 0 if cache is None else cache.key.size(-2)
position_ids = torch.arange(
past_length,
past_length + input_ids.size(1),
device=input_ids.device,
) # prefill 可能是 [0..P-1],decode 通常只有 [P+t]python学习式绝对位置要查正确行;旋转位置编码(Rotary Position Embedding, RoPE)要用正确角度旋转新 Q/K;相对位置偏置要知道 Query 的全局位置。若每步都把新 token 当成位置 0,形状不会报错,缓存等价性却会立即失效。
09 Encoder–Decoder 的 cross-attention 还能缓存什么?#
对 Encoder–Decoder Transformer,编码器 memory 在整个目标生成期间不变。每个解码层的 cross-attention 可以将 memory 投影成本层的 一次,后续每步只从新目标表示计算 Query:
生成前一次:
encoder memory [N,S,D] ─► 每层 cross-attn K/V projection
─► K_mem,V_mem [N,H,S,d]
每个 decode step:
新目标表示 [N,1,D] ─► q_t [N,H,1,d]
─► 读固定 K_mem,V_mem
─► cross-attn output [N,1,D]text这份 cross-attention K/V 的长度固定为 ;目标自注意力 K/V 则会随生成长度增长。两者的来源和生命周期不同,工程上不要放进一个无类型的“cache”列表里靠顺序猜。
10 从透明 cat 到生产缓存#
torch.cat([old,new], dim=-2) 每步都要为更长张量分配存储并复制历史内容。它适合教学和等价性测试,不是高并发服务的缓存管理策略。
| 策略 | 写入方式 | 优点 | 主要代价 |
|---|---|---|---|
动态 cat | 每步生成新张量 | 代码最透明 | 重复分配和拷贝 |
| 预分配静态 cache | 写入预定位置 | 形状稳定、少分配 | 需要最大长度和安全边界 |
| 分页 cache | 用固定大小块映射逻辑序列 | 易共享、减少外部碎片 | 需要块表、调度和专用内核 |
| 滑动窗口 | 仅保留最近 个位置 | 显存上界固定 | 丢弃窗口外直接证据 |
静态 cache 常用 cache_position 或等价索引写入,但“预分配成功”不代表 mask 正确:未写入槽位必须对当前 Query 不可见。分页 cache 还要在逻辑 token 位置和物理块地址之间维护正确映射。
11 Batch、EOS 与 Beam Search 为什么让缓存更难?#
同一 batch 中的序列可能在不同时刻产生 <eos>。有三种常见处理:
- 保留已完成行并对它们生成 padding,调度简单但继续占用计算。
- 将已完成序列从活跃 batch 移除,需要同步重排所有层 K/V 和输出索引。
- 用连续批处理调度不同请求,需要更完整的块管理与隔离。
Beam Search 会让一条前缀分叉成多个候选。新 beam 在分叉前的 K/V 完全相同,可以逻辑共享;候选重排后,缓存的 batch/beam 维也必须按中选 parent beam 同步重排。只重排 token id 而忘了 K/V,形状仍然合法,每条 beam 却在读别人的历史。
12 一条可执行的调试与性能验证路径#
- 先做逐位置等价测试。 关闭 dropout,比较完整因果前向与逐 token 缓存的所有位置,不只比最后 token。
- 检查每层 cache 长度。 prefill 后应为 ,每步只增加 1;不同层必须一致。
- 检查追加轴。 序列轴是
-2,头宽轴是-1;拼错轴有时会因数值巧合而暂时不报错。 - 打印新 Query 的全局位置。 它应等于已有 cache 长度,而不是每步都回到 0。
- 做 cache 污染测试。 两个请求交替生成,验证它们的缓存存储不共享可变写入区。
- 做 beam 重排测试。 人工交换 parent beam 索引,检查所有层的 K/V 首维都按同一映射更新。
- 分开测 prefill 和 decode。 分别记录首 token 延迟、每 token 延迟、吞吐、峰值缓存显存,并固定 batch、prompt 长度和生成长度。
- 用 profiler 确认少算了,而不是只看 wall time。 检查 decode 步的 QKV 投影输入长度是 1,且没有隐式重建全前缀。
13 最常见的缓存错误#
- prefill 没有 causal mask。 prompt 内的早期位置在预填充阶段偷看了后面。
- 对
[L=1,S>1]的增量 SDPA 盲传is_causal=True。 非方形 mask 对齐与想象不同,新 Query 可能读不到全部历史。 - 只缓存最后一层。 前面每层仍重算整个前缀,计算没有真正增量化。
- 复用了不同模型版本的 cache。 权重更新后历史 K/V 已失效,请求内更不能热切模型。
- 位置索引每步从 0 开始。 学习式位置或 RoPE 与无缓存前向不等价。
- 将不同请求的缓存写进同一槽位。 这不只是质量问题,还可能成为跨请求数据泄漏。
- beam 排序后未重排缓存。 token 前缀与 K/V 历史不再对应。
- 动态
cat误当生产优化。 注意力少算了,缓存却在每步整体复制。 - 认为 KV Cache 减少训练显存。 它主要是自回归推理优化;标准训练还需要前向激活以供反向。
- 只报告 tokens/s,不报告测试形状。 batch、prompt、生成长度不同时,数字无法比较。
14 缓存会在哪些场景失去优势?#
- 显存容量先到上限。 长上下文、大 batch、多层多头会让 K/V 显存线性增长,可能迫使 batch 降低并伤害吞吐。
- 只生成极少 token。 短输出时,缓存管理与内核启动的收益可能不明显,prompt prefill 仍占主导。
- 模型要修改历史表示。 非因果、双向或对整段反复编辑的架构不满足“未来不改写过去”前提。
- 窗口被截断。 滑动窗口缓存把超出 的 K/V 丢弃后,数学上已不再等价于全上下文模型。
- 缓存量化误差。 低比特 K/V 可降显存与带宽,但误差会影响后续每一个 Query,需要按任务和长度验证。
- 优化层次不同。 FlashAttention 类内核改善注意力内存访问,KV Cache 避免跨解码步重算;两者可同时使用,不是二选一。
15 与相近技术怎样区分?#
| 技术 | 主要减少什么 | 是否改变注意力连接 | 主要代价 |
|---|---|---|---|
| KV Cache | 跨生成步的历史 K/V 重算 | 否 | 持久显存与管理复杂度 |
| FlashAttention | 精确注意力中间矩阵的 HBM 访问 | 否 | 内核与硬件约束 |
| Multi-Query Attention | 让所有 Query 头共享更少 K/V 头 | 改变参数化 | 表示容量与训练选择 |
| Grouped-Query Attention | 让一组 Query 头共享 K/V 头 | 改变参数化 | 头数约束与质量折中 |
| Sliding-Window Attention | 每个 Query 的可见 Key 数 | 是 | 无法直接读取窗口外 token |
| Speculative Decoding | 大模型串行验证步数 | 不必然 | 草稿模型、验证与调度 |
PyTorch 当前的 SDPA 还支持实验性 enable_gqa=True,要求 Query 头数能整除 Key/Value 头数,且 Key 头数等于 Value 头数;当前后端和 Nested Tensor 仍有限制。这是减少缓存宽度的模型结构选择,不是打开 KV Cache 所必需的开关。
16 今天真正需要记住什么?#
- 因果模型中,后来 token 不会改写历史位置,因此每层历史 Key/Value 可以在解码步之间复用。
- 增量步的新 Query 是
[N,H,1,d],缓存为[N,H,t,d],当步分数只是[N,H,1,t]。 - 历史 Key/Value 会被未来 Query 重复读取,历史 Query 不会;所以通常只持久缓存 K/V。
- prefill 是并行处理 prompt 并建立缓存,decode 是每步追加一个 K/V;两者的性能瓶颈不同。
- KV Cache 是用显存换重复计算,不会消除逐 token 依赖;长上下文、大 batch 和多 beam 时,缓存管理可变成主要问题。
17 思考题与小练习#
- 对 的标准多头注意力,计算 fp16 K/V Cache 的元素数与约多少 GiB。若 Key/Value 头数降为 4,理论缓存降为原来多少?
- 修改
CachedCausalSelfAttention的等价测试,先用前 4 个 token 做一次 prefill,再逐个输入后 2 个 token。验证与长度 6 的完整 causal forward 逐位置一致。 - 将增量分支的
is_causal=False改成True,用三个可区分的 Value 打印输出。解释非方形因果对齐为何让唯一 Query 没有读到想象中的全部过去。
相关工作#
- Dai et al. (2019), Transformer-XL ↗:复用跨片段隐状态并引入相对位置,展示历史状态复用对长依赖的价值。
- Shazeer (2019), Fast Transformer Decoding: One Write-Head is All You Need ↗:提出 Multi-Query Attention,通过共享 K/V 头减少自回归解码的缓存带宽。
- Ainslie et al. (2023), GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints ↗:用 Grouped-Query Attention 在多头容量与 MQA 缓存成本之间折中。
- Kwon et al. (2023), Efficient Memory Management for Large Language Model Serving with PagedAttention ↗:将分页内存思路用于 KV Cache,降低碎片并支持高吞吐服务。
- Dao et al. (2022), FlashAttention ↗:从 IO 复杂度优化精确注意力,与跨步复用 K/V 的优化层次不同。
18 下一篇预告#
KV Cache 假设每个 token 带着正确的位置进入注意力。下一篇将回到这个前提,系统比较正弦绝对位置、学习式位置与旋转位置编码,并追踪位置信息究竟是加到隐状态,还是直接改写 Query–Key 的相似度。