Keyboard shortcuts

Press or to navigate between chapters

Press S or / to search in the book

Press ? to show this help

Press Esc to hide this help

KV Cache: 自回归推理为什么能避免重复计算?

这一节单独把 KV Cache 拿出来讲,因为它几乎是理解 LLM 推理效率时绕不过去的一步。
很多人第一次看生成代码时会疑惑:

  • 为什么模型不是每生成一个 token,就把整个前缀全部重新算一遍?

KV Cache 就是在回答这个问题。

这一节准备回答几个问题:

  1. KV Cache 到底缓存了什么?
  2. 为什么它能减少推理开销?
  3. 它会带来什么代价?
  4. 为什么它和 GQA 经常一起出现?

Q1: KV Cache 到底缓存了什么?

在 attention 里,给定某一层输入 \( X \),通常会先得到:

$$ Q = XW_Q,\qquad K = XW_K,\qquad V = XW_V $$

在自回归生成时,如果当前已经生成到第 \( t \) 个位置,那么第 \( t \) 步最关心的是:

  • 当前 token 的 Query
  • 历史所有 token 的 Key 和 Value

也就是说,当前步真正要拿来做 attention 的通常是:

$$ Q_t,\quad K_{1:t},\quad V_{1:t} $$

其中最适合缓存的是:

  • 历史位置已经算好的 \( K \)
  • 历史位置已经算好的 \( V \)

所以 KV Cache 本质上就是:

  • 把历史 token 的 K/V 保留下来
  • 下一步只增量追加当前 token 对应的新 K/V

Q2: 为什么它能减少推理开销?

如果没有 KV Cache,那么每生成一个新 token,都要把整个前缀重新过一遍 attention。
前缀越长,重复计算越多。

而有了 KV Cache 之后,生成第 \( t \) 个 token 时:

  • 不需要重新计算前 \( 1 \sim t-1 \) 个位置的 K/V
  • 只需要计算当前新位置的 K/V
  • 再把当前 Query 和历史缓存过的 K/V 做 attention

从直觉上看,它做的事情很简单:

  • 把“整段前缀重复算”变成“历史结果复用 + 当前步增量算”

所以 KV Cache 的直接收益通常就是:

  • 降低重复计算
  • 降低单步生成延迟
  • 让长文本生成更可接受

Q3: 它会带来什么代价?

KV Cache 不是白来的。
它省下了重复计算,但会把压力转移到缓存占用上。

如果某层的 K/V 形状写成:

$$ K, V \in \mathbb{R}^{B \times h \times n \times d_{\text{head}}} $$

那么随着生成长度 \( n \) 增加,缓存也会持续变大。
而且这是每一层都要存的,所以总占用会很可观。

这也是为什么在推理阶段,我们经常会同时关心:

  • 序列长度
  • 层数
  • head 数
  • \( d_{\text{head}} \)
  • dtype

因为这些都会直接影响 KV cache 的总大小。

Q4: 为什么它和 GQA 经常一起出现?

因为 KV Cache 里存的就是 K/V。
所以只要能减少 Key/Value 的 head 数,就能直接减少缓存开销。

这就是 GQA 重要的地方之一。
如果标准多头注意力缓存的是:

$$ K, V \in \mathbb{R}^{B \times h \times n \times d_{\text{head}}} $$

而 GQA 改成:

$$ K, V \in \mathbb{R}^{B \times h_{kv} \times n \times d_{\text{head}}} $$

并且 \( h_{kv} < h \),那 KV cache 的显存占用就会跟着下降。

所以很多时候:

  • KV Cache 回答的是“为什么推理不用反复重算”
  • GQA 回答的是“既然要缓存,怎样把缓存做得更省”

这一节之后最重要的收获是什么?

如果只保留最重要的几点,我觉得是:

  1. KV Cache 缓存的是历史 token 的 Key 和 Value,而不是整个模型的所有中间结果。
  2. 它的核心价值是把重复前缀计算改成增量计算,从而降低生成延迟。
  3. 它省的是算力,换来的是显存占用,所以推理优化常常会围绕 KV cache 展开。

如果后面在代码里看到:

  • past_key_values
  • use_cache=True

基本就可以直接联想到:
这里正在利用 KV cache 做自回归推理加速。