返回 M01~M90 教材库
M51 · 预习教材教材已备 ≠ 学习已完成

M51:最小 Cached Decode:复用历史 K/V 而不改变输出语义

Cached decode 在每个生成步只计算新 token 的 K/V,并把它们追加到历史 cache;在相同模型和数值条件下,它应与无 cache 自回归计算保持语义等价。

2026-10-13kv-cache、cached-decode、causal-attention、implementation、correctness

内容类型:预习教材(不代表已完成)
日期: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 自回归计算保持语义等价。

学习目标

  1. 能描述 cache 的 shape、追加维度和每步输入。
  2. 能实现或跟随 tiny LM 的最小 cached decode。
  3. 能用 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 管理。

最小练习或观察步骤

  1. 使用 tiny LM 和固定 token 序列;先写无 cache 的 greedy 生成作为 oracle。
  2. 修改 attention 接口,使其接收可选 past_kv,返回 new_kv
  3. prefill 完整 prompt,确认每层 cache 的 sequence 长度等于 prompt 长度。
  4. 每步只输入最新 token,使用正确 position offset,并追加 cache。
  5. 逐步比较 cached/no-cache 的 logits 最大差值与 argmax token;不预设必然一致。
  6. 只有正确性边界通过后,再记录时间和峰值内存。
  7. 增加一个长度 1 prompt 或超过一次生成步的边界例。

常见误区

  • 只比较最后文本,忽略中间 logits 已经偏离。
  • cache 写入了旧 K/V,却仍把完整前缀作为新输入,造成重复。
  • position id 每步从零开始。
  • cache 没有按层分开,或在 beam/batch 间错误共享。
  • 看到运行更快就忽略输出不等价。

金融 / Web3 / 文档场景连接

在生成审计说明或长文档摘要时,cache 错位可能产生流畅但不可察觉的内容差异。上线前的正确性对照应包含数字、地址、证据 ID 等高敏感 token,而不能只用普通自然语言观察。

自检问题

  1. cached decode 每一步重新计算哪些张量?
  2. 为什么旧位置表示无需因新 token 重新计算?
  3. position offset 错误为何可能不触发 shape 异常?
  4. 为什么先用 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 处理方式:
  • 正确性对照(未运行可留空):
  • 时间/内存观察(未运行可留空):
  • 一个失败边界:
学完后,请把自己的理解、练习结果和仍不确定的问题写入文末“学后填写区”,再到唯一进度账本更新状态。预先阅读后续教材不会自动增加完成数。