LESSON 6.2 · 推导 · 100 分钟

KV Cache:用显存换重复计算

生成第 t 个 token 时,过去 token 在每一层的 Key 和 Value 不会改变。KV Cache 保存这些张量,使模型只需为新 token 计算新的 Q、K、V,再让新 Query 读取缓存,而不用每一步重算整个前缀。对普通缓存,显存大致正比于 2 × 层数 × batch × 序列长度 × KV 头数 × head_dim × 每元素字节数

DIRECT ANSWER · VERIFIED SOURCES ·

KV Cache 为什么能加速自回归解码,又为什么会吃掉大量显存?

生成第 t 个 token 时,过去 token 在每一层的 Key 和 Value 不会改变。KV Cache 保存这些张量,使模型只需为新 token 计算新的 Q、K、V,再让新 Query 读取缓存,而不用每一步重算整个前缀。对普通缓存,显存大致正比于 2 × 层数 × batch × 序列长度 × KV 头数 × head_dim × 每元素字节数,其中 2 表示 K 和 V。

视频对齐

视频对齐点:把 prefill 与逐 token decode 分开画图,逐层标出新 Q、追加的新 K/V、读取的历史 K/V,并用模型配置实际代入缓存公式。

关键结论

  • KV Cache 主要避免 decode 阶段重复计算历史 K/V;它不会消除首次 prefill 的完整计算。
  • 缓存随活动序列长度和并发请求增长;长上下文服务经常首先受 KV 显存而非权重显存限制。
  • GQA/MQA 通过减少 KV 头数降低缓存,但 Query 头数可以保持更多。

边界与常见误解

公式假设每层保存完整、同精度的 K/V;滑动窗口、量化、offload、跨层共享和稀疏注意力都会改变结果。修改前缀、位置或 attention mask 时,旧缓存也未必可直接复用。

一手来源

学完你应该能够

  1. 能区分 prefill 一次写入整段 K/V 与 decode 每步追加一个位置。
  2. 能写出 KV cache 的 shape 与字节公式,并解释 GQA/MQA 如何改变 KV head 数。
  3. 能证明缓存版本与每步重算前缀的注意力输出一致。
  4. 能区分 KV cache 本身、分页分配和 prefix caching。

核心概念

Prefill 与 decode

Prefill 并行处理 prompt,产生每层所有 prompt token 的 K/V;decode 每步只为新 token 计算并追加一组 K/V,再让新 query 读取已有前缀。

缓存 shape

常见逻辑形状可写成 [layers, 2, batch, tokens, kv_heads, head_dim]。实现可能转置或分页,但元素数量仍由这些维度决定。

容量公式

字节数 = layers × 2(K,V) × batch × tokens × kv_heads × head_dim × 每元素字节。上下文翻倍时,其他条件不变,缓存容量线性翻倍。

MHA、GQA 与 MQA

MHA 通常让 query heads 与 KV heads 相同;GQA 让多组 query 共享较少 KV heads;MQA 只保留一组 K/V,因此直接降低缓存与读取量。

结果等价

正确缓存只复用已经算过的历史 K/V,不改变注意力数学。固定权重和输入时,逐步缓存输出必须与每步重算整个前缀一致。

实践任务

手写带 KV cache 的 attention 并验证等价性

  • 用三个二维 token 和固定 WQ/WK/WV,先手算第一个位置的 K/V 与输出。
  • 实现 decode_step:新 token 只计算一次 K/V,追加后让当前 query 读取整个 cache。
  • 实现 reference:每一步重算完整前缀,逐元素断言两种输出一致。
  • 填写 layers、tokens、kv_heads、head_dim、dtype 的容量表,再比较上下文翻倍和 GQA。

进入完整互动课程