M51:最小 Cached Decode:复用历史 K/V 而不改变输出语义
Cached decode 在每个生成步只计算新 token 的 K/V,并把它们追加到历史 cache;在相同模型和数值条件下,它应与无 cache 自回归计算保持语义等价。
内容类型:预习教材(不代表已完成)
日期:2026-10-13
阶段:P1 · AI Model Engineering 90
周次:W8 · 推理、KV Cache、量化与 serving
节奏:周二最小实现
状态:教材已备;学习未完成
标签:kv-cache、cached-decode、causal-attention、implementation、correctness
一句话定义
Cached decode 在每个生成步只计算新 token 的 K/V,并把它们追加到历史 cache;在相同模型和数值条件下,它应与无 cache 自回归计算保持语义等价。
学习目标
- 能描述 cache 的 shape、追加维度和每步输入。
- 能实现或跟随 tiny LM 的最小 cached decode。
- 能用 logits/token 一致性先验证正确,再讨论速度和内存。
核心知识
无 cache 生成第一个 token 时,输入完整 prompt;生成下一 token 时,又把 prompt + token₁ 全部送入模型。Cached 路线在 prefill 后保存每层 K/V,下一步只输入 token₁,并把新 K/V 追加到 sequence 维。
典型 cache shape 可表示为 [batch, kv_heads, sequence, head_dim],但库可能交换维度。最重要的是明确:在哪一维追加、query 的序列长度为何通常为 1、position id 从历史长度继续、causal mask 是否仍正确。
正确性对照要固定解码策略。若使用采样,随机数消费顺序可能造成差异;学习阶段优先 greedy argmax,逐步比较 logits 或 token。浮点后端与运算顺序可能产生微小误差,应预先定义容差。
机制/推导
在第 t 步,无 cache 为所有位置重新算 K₁…K_t、V₁…V_t;有 cache 只算 K_t、V_t,再形成:
K_cache ← concat(K_cache, K_t)
V_cache ← concat(V_cache, V_t)
新 query Q_t 对整个 K_cache 做注意力并加权 V_cache。旧 token 的输出无需重新生成,因为 causal Transformer 中旧位置看不到未来 token。
位置编码是常见陷阱。绝对位置、RoPE 或其他方案都要求新 token 使用正确的偏移;若每个 decode step 都从位置 0 开始,shape 仍可能正确,输出却失真。Batch 中不同序列长度还需要独立的有效长度或分页 cache 管理。
最小练习或观察步骤
- 使用 tiny LM 和固定 token 序列;先写无 cache 的 greedy 生成作为 oracle。
- 修改 attention 接口,使其接收可选
past_kv,返回new_kv。 - prefill 完整 prompt,确认每层 cache 的 sequence 长度等于 prompt 长度。
- 每步只输入最新 token,使用正确 position offset,并追加 cache。
- 逐步比较 cached/no-cache 的 logits 最大差值与 argmax token;不预设必然一致。
- 只有正确性边界通过后,再记录时间和峰值内存。
- 增加一个长度 1 prompt 或超过一次生成步的边界例。
常见误区
- 只比较最后文本,忽略中间 logits 已经偏离。
- cache 写入了旧 K/V,却仍把完整前缀作为新输入,造成重复。
- position id 每步从零开始。
- cache 没有按层分开,或在 beam/batch 间错误共享。
- 看到运行更快就忽略输出不等价。
金融 / Web3 / 文档场景连接
在生成审计说明或长文档摘要时,cache 错位可能产生流畅但不可察觉的内容差异。上线前的正确性对照应包含数字、地址、证据 ID 等高敏感 token,而不能只用普通自然语言观察。
自检问题
- cached decode 每一步重新计算哪些张量?
- 为什么旧位置表示无需因新 token 重新计算?
- position offset 错误为何可能不触发 shape 异常?
- 为什么先用 greedy 而不是随机采样做等价性检查?
专业课程对齐
- 精读 vLLM PagedAttention 设计文档 的 cache 块寻址、QK 计算与 value 聚合,注意这是高性能 kernel,先将其还原为单 batch 的逻辑 cache。
- 精读 Stanford CS336 的 Transformer attention 与 inference 系统课题,对齐 RoPE/position offset、MHA/GQA 与 cache shape。
- 选读 PyTorch Tutorials 的 scaled dot-product attention、inference mode 与 numerical accuracy 主题,用来写 oracle 对照。
深入学习提示
先用无 cache greedy decode 建立 oracle,再实现 prefill 与单 token decode,最后才读 PagedAttention 的分页布局。代码上逐层打印 [batch,kv_heads,sequence,head_dim],确认追加维、position id、causal mask 和 batch 内有效长度。每步比较 cached/no-cache logits 的最大绝对差、argmax token 与 cache length,正确后再量时和记峰值内存。反例是只比较最终文本:错误 position offset 可能偶然生成同 token;另一反例是对 cache 每步 concat 导致重分配,却把该开销归因于 attention 本身。
学后填写区
- Cache shape 与追加维度:
- Position 处理方式:
- 正确性对照(未运行可留空):
- 时间/内存观察(未运行可留空):
- 一个失败边界: