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

GQA: 为什么 Query head 和 KV head 可以不一样?

这一节单独把 GQA 拿出来讲,因为它特别适合帮助我们理解一个现实问题:

  • 为什么 attention 的很多改进,看起来是在改结构,实际上是在为推理效率和 KV cache 服务?

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

  1. 标准多头注意力里,Q/K/V 的 head 数通常是什么关系?
  2. GQA 在改什么?
  3. 为什么很多模型会让 Query head 数多于 Key/Value head 数?
  4. GQA 和 KV cache 有什么关系?

Q1: 标准多头注意力里,Q/K/V 的 head 数通常是什么关系?

在最标准的 multi-head attention 里,通常会默认:

  • Query 的 head 数和 Key 的 head 数一样
  • Key 的 head 数和 Value 的 head 数一样

如果把 head 数记作 \( h \),那么常见写法是:

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

这里:

  • \( B \) 表示 batch size
  • \( n \) 表示序列长度
  • \( h \) 表示 head 数
  • \( d_{\text{head}} \) 表示每个 head 的维度

这也是最容易理解的形式,因为每个 Query head 都能直接对应一个 Key/Value head 去做 attention。

Q2: GQA 在改什么?

GQA 的核心想法是:

  • 不一定要让 Query head 数和 Key/Value head 数完全一样

更具体一点,可以记成:

$$ h_q > h_{kv} $$

其中:

  • \( h_q \) 表示 Query 的 head 数
  • \( h_{kv} \) 表示 Key/Value 的 head 数

这意味着模型会保留更多的 Query heads,但让多个 Query heads 共享同一组 Key/Value heads。

从“数学形式”上看,这不是把 attention 完全改写了。
从“工程含义”上看,它是在做一个非常现实的权衡:

  • 尽量保留 Query 侧的表达能力
  • 同时减少 Key/Value 侧的存储和缓存开销

Q3: 为什么很多模型会让 Query head 数多于 Key/Value head 数?

这里最值得记住的直觉是:

  • Query 是当前时间步动态计算出来的
  • Key / Value 在自回归推理时是可以缓存的

也就是说,在生成第 \( t \) 个 token 时:

  • 当前 token 的 \( Q_t \) 需要重新算
  • 历史 token 的 \( K_{1:t} \)、\( V_{1:t} \) 往往已经在 cache 里

所以很多实现会更愿意把“节省资源”的重点放在 K/V 上。
因为真正会被长期存起来、并且随着上下文变长不断累积的,是 K/V cache。

从这个角度看,保留更多 Query heads 的动机通常是:

  • 尽量保留多头注意力在不同子空间建模的能力

减少 Key/Value heads 的动机通常是:

  • 降低 KV cache 的显存占用
  • 降低推理时和 K/V 相关的访存与计算开销

所以一个很实用的记法是:

  • 保留更多 Q heads,主要是在保表达力
  • 减少 K/V heads,主要是在省缓存和推理成本

Q4: GQA 和 KV cache 有什么关系?

它们的关系其实非常直接。

如果标准多头注意力中:

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

那么缓存这些 K/V 时,显存占用会直接和 \( h \) 成正比增长。
而如果改成 GQA,让 Key/Value 的 head 数变成 \( h_{kv} \),那么缓存张量更接近:

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

只要 \( h_{kv} < h_q \),KV cache 的占用就会明显下降。

这也是为什么 GQA 经常和 KV cache 一起讨论。
它不只是一个“结构设计小技巧”,而是会直接影响:

  • 长上下文推理时的显存压力
  • 大 batch 推理时的吞吐
  • 部署时能不能把模型跑起来

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

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

  1. 标准多头注意力通常默认 Q/K/V head 数相同,但这不是不可动的规则。
  2. GQA 的核心是在尽量保留 Query 侧表达能力的同时,减少 Key/Value 侧的缓存成本。
  3. 它之所以重要,不只是因为“结构有点不一样”,而是因为它直接影响推理时最昂贵的 KV cache。

所以如果后面在代码里看到:

  • num_attention_heads
  • num_key_value_heads

这往往就是在告诉我们:
这个模型已经不再使用最朴素的标准 MHA,而是在朝着 GQA 这样的推理友好结构靠拢。