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 时,旧缓存也未必可直接复用。
一手来源
学完你应该能够
- 能区分 prefill 一次写入整段 K/V 与 decode 每步追加一个位置。
- 能写出 KV cache 的 shape 与字节公式,并解释 GQA/MQA 如何改变 KV head 数。
- 能证明缓存版本与每步重算前缀的注意力输出一致。
- 能区分 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。