观文听傑

返回

上一篇组装了 Transformer Decoder 的三条信息流:因果自注意力读目标前缀,交叉注意力读源记忆,FFN 逐位置加工特征。透明的推理循环每次都把完整前缀重新送入模型,虽然答案正确,却在不断重做已经做过的投影和注意力。

Key–Value Cache(KV Cache)的核心不是改变模型数学,而是保存每层历史 token 已经算好的 Key 和 Value。新一步只投影新 token,用它的 Query 读取“历史缓存 + 当前 token”。本文只讲透这个增量数据流。

01 无缓存解码到底重复了什么?#

设 prompt 长度为 PP,已生成 tt 个 token。无缓存方式在第 t+1t+1 步将长度 P+tP+t 的整个序列重新前向:

第 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 可以复用。

若每一步都重算长度 tt 的完整自注意力,仅分数矩阵工作量的累积就约为:

t=1TO(t2D)=O(T3D)\sum_{t=1}^{T}O(t^2D)=O(T^3D)

用 KV Cache 后,每步只有一个新 Query 与 tt 个 Key 计算分数:

t=1TO(tD)=O(T2D)\sum_{t=1}^{T}O(tD)=O(T^2D)

这是简化的单层注意力量级,不包括 FFN、投影、内核常数、prompt 预填充和内存带宽。它说明缓存消除了哪类重复,不代表真实延迟会按同一比例下降。

02 每层缓存里究竟放什么?#

HH 个头、每头宽度 dd 的自注意力,在第 tt 步开始时:

张量形状含义
新 token 隐状态[N,1,D]当前层只处理一个新位置
新 Query[N,H,1,d]询问历史与当前信息
新 Key/Value[N,H,1,d]把当前 token 加入可被未来查询的记忆
更新后 Key Cache[N,H,t,d]位置 1 到 tt 的全部 Key
更新后 Value Cache[N,H,t,d]位置 1 到 tt 的全部 Value
当步注意力分数[N,H,1,t]一个新 Query 读取 tt 个 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 层共用。若有 LlayersL_{layers} 层,标准多头注意力缓存的元素数约为:

2LlayersNHTd2\,L_{layers}\,N\,H\,T\,d

开头的 2 分别对应 Key 和 Value。若用 fp16/bfloat16,每元素通常 2 字节;批大小、层数和上下文长度都会线性放大缓存。

03 为什么缓存 K/V,通常不缓存 Query?#

在因果生成的第 tt 步,我们只需要计算最新位置的输出。它的 Query qtq_t 会读所有 kt,vtk_{\le t},v_{\le t}。上一步的 qt1q_{t-1} 已经用于产生上一位置输出,新 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)]
text

Key/Value 是未来 Query 反复查询的记忆,Query 是当步一次性的读取请求。“不缓存 Q”指不为未来步保留历史 Query;当前内核在计算期间当然仍需要 qtq_t

04 用一维头手算两步缓存#

为了可手算,令单头宽度 d=1d=1,缩放因子为 1。预填充后缓存为:

K(2)=[1,2],V(2)=[10,20]K^{(2)}=[1,2],\qquad V^{(2)}=[10,20]

第 3 个 token 投影得到 q3=1,k3=0,v3=30q_3=1,k_3=0,v_3=30。追加后:

K(3)=[1,2,0],V(3)=[10,20,30]K^{(3)}=[1,2,0],\qquad V^{(3)}=[10,20,30]

分数为 [1,2,0],softmax 近似为 [0.245,0.665,0.090],所以:

o30.245(10)+0.665(20)+0.090(30)=18.45o_3\approx0.245(10)+0.665(20)+0.090(30)=18.45

第 4 步只需追加 k4,v4k_4,v_4 并计算新 q4q_4。已缓存的 [1,2,0][10,20,30] 不变。若无缓存,模型会从前缀隐状态再算一次这三个 Key/Value,最终 o3o_3 不会因此更正确,只会更费计算。

05 预填充与逐 token 解码是两个阶段#

KV Cache 推理常分成:

  1. 预填充(Prefill):一次输入整个 prompt,使用 causal attention 并行计算其表示,同时写入每层 prompt K/V。
  2. 解码(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 logits
text

预填充往往更偏计算密集,因为可以并行处理许多 prompt token;单 token decode 往往更受缓存读写和内存带宽限制。性能报告应分开首 token 延迟(Time to First Token)与后续 token 间隔,不要只给一个平均数。

06 用 PyTorch 2.13 SDPA 写透明的动态 KV Cache#

下面的实现只展示单层 causal self-attention。为便于理解,它用 torch.cat 追加缓存;生产实现应避免每步重新分配并复制已有缓存。

PyTorch 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 缓存”的对应输出接近:

融合内核、低精度累加和不同计算顺序可使末位出现微小数值差异,因此使用容差比较,不要用逐 bit 相等。但若误差很大或随长度急剧放大,应先查位置编码偏移、层缓存对应、mask 与 cache 追加轴。

08 位置索引为什么必须跟着缓存长度走?#

增量步只输入 [N,1] 的 token,但它不是“位置 0”。若缓存已有 P+tP+t 个位置,新 token 的绝对位置应是 P+tP+t

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 M[N,S,D]M[N,S,D] 在整个目标生成期间不变。每个解码层的 cross-attention 可以将 memory 投影成本层的 Kmem,VmemK_{mem},V_{mem} 一次,后续每步只从新目标表示计算 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 的长度固定为 SS;目标自注意力 K/V 则会随生成长度增长。两者的来源和生命周期不同,工程上不要放进一个无类型的“cache”列表里靠顺序猜。

10 从透明 cat 到生产缓存#

torch.cat([old,new], dim=-2) 每步都要为更长张量分配存储并复制历史内容。它适合教学和等价性测试,不是高并发服务的缓存管理策略。

策略写入方式优点主要代价
动态 cat每步生成新张量代码最透明重复分配和拷贝
预分配静态 cache写入预定位置形状稳定、少分配需要最大长度和安全边界
分页 cache用固定大小块映射逻辑序列易共享、减少外部碎片需要块表、调度和专用内核
滑动窗口仅保留最近 WW 个位置显存上界固定丢弃窗口外直接证据

静态 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 一条可执行的调试与性能验证路径#

  1. 先做逐位置等价测试。 关闭 dropout,比较完整因果前向与逐 token 缓存的所有位置,不只比最后 token。
  2. 检查每层 cache 长度。 prefill 后应为 PP,每步只增加 1;不同层必须一致。
  3. 检查追加轴。 序列轴是 -2,头宽轴是 -1;拼错轴有时会因数值巧合而暂时不报错。
  4. 打印新 Query 的全局位置。 它应等于已有 cache 长度,而不是每步都回到 0。
  5. 做 cache 污染测试。 两个请求交替生成,验证它们的缓存存储不共享可变写入区。
  6. 做 beam 重排测试。 人工交换 parent beam 索引,检查所有层的 K/V 首维都按同一映射更新。
  7. 分开测 prefill 和 decode。 分别记录首 token 延迟、每 token 延迟、吞吐、峰值缓存显存,并固定 batch、prompt 长度和生成长度。
  8. 用 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 仍占主导。
  • 模型要修改历史表示。 非因果、双向或对整段反复编辑的架构不满足“未来不改写过去”前提。
  • 窗口被截断。 滑动窗口缓存把超出 WW 的 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 今天真正需要记住什么?#

  1. 因果模型中,后来 token 不会改写历史位置,因此每层历史 Key/Value 可以在解码步之间复用。
  2. 增量步的新 Query 是 [N,H,1,d],缓存为 [N,H,t,d],当步分数只是 [N,H,1,t]
  3. 历史 Key/Value 会被未来 Query 重复读取,历史 Query 不会;所以通常只持久缓存 K/V。
  4. prefill 是并行处理 prompt 并建立缓存,decode 是每步追加一个 K/V;两者的性能瓶颈不同。
  5. KV Cache 是用显存换重复计算,不会消除逐 token 依赖;长上下文、大 batch 和多 beam 时,缓存管理可变成主要问题。

17 思考题与小练习#

  1. Llayers=24,N=4,H=16,T=2048,d=64L_{layers}=24,N=4,H=16,T=2048,d=64 的标准多头注意力,计算 fp16 K/V Cache 的元素数与约多少 GiB。若 Key/Value 头数降为 4,理论缓存降为原来多少?
  2. 修改 CachedCausalSelfAttention 的等价测试,先用前 4 个 token 做一次 prefill,再逐个输入后 2 个 token。验证与长度 6 的完整 causal forward 逐位置一致。
  3. 将增量分支的 is_causal=False 改成 True,用三个可区分的 Value 打印输出。解释非方形因果对齐为何让唯一 Query 没有读到想象中的全部过去。

相关工作#

18 下一篇预告#

KV Cache 假设每个 token 带着正确的位置进入注意力。下一篇将回到这个前提,系统比较正弦绝对位置、学习式位置与旋转位置编码,并追踪位置信息究竟是加到隐状态,还是直接改写 Query–Key 的相似度。

每生成一个词为何又重算全文?KV Cache 的增量解码
https://zwjcode.cn/blog/transformer-kv-cache-incremental-decoding
作者
发布于 2026年9月8日
版权协议 CC BY-NC-SA 4.0
评论加载似乎遇到了问题,请尝试刷新页面。