GQA: 为什么 Query head 和 KV head 可以不一样?
这一节单独把 GQA 拿出来讲,因为它特别适合帮助我们理解一个现实问题:
- 为什么 attention 的很多改进,看起来是在改结构,实际上是在为推理效率和 KV cache 服务?
这一节准备回答几个问题:
- 标准多头注意力里,Q/K/V 的 head 数通常是什么关系?
- GQA 在改什么?
- 为什么很多模型会让 Query head 数多于 Key/Value head 数?
- 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 推理时的吞吐
- 部署时能不能把模型跑起来
这一节之后最重要的收获是什么?
如果只保留最重要的几点,我觉得是:
- 标准多头注意力通常默认 Q/K/V head 数相同,但这不是不可动的规则。
- GQA 的核心是在尽量保留 Query 侧表达能力的同时,减少 Key/Value 侧的缓存成本。
- 它之所以重要,不只是因为“结构有点不一样”,而是因为它直接影响推理时最昂贵的 KV cache。
所以如果后面在代码里看到:
num_attention_headsnum_key_value_heads
这往往就是在告诉我们:
这个模型已经不再使用最朴素的标准 MHA,而是在朝着 GQA 这样的推理友好结构靠拢。