6. 理解 GQA

GQA(Grouped Query Attention,分组查询注意力)是 Google 在 2023 年论文《GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints》中提出的一种 Attention 优化方法。它的核心目标是:

在尽量不损失模型效果的前提下,大幅降低 KV Cache 的内存占用和推理成本。 (ACL Anthology)

如果你已经理解了 Transformer 的 Multi-Head Attention(MHA),那么 GQA 本质上就是:

让多个 Query Head 共享同一组 Key 和 Value Head。


先理解问题:为什么要优化 Attention?

标准 Transformer 的 Attention:

Attention(Q,K,V)=softmax(QKTd)V

对于每个 Head,都有独立的:

假设:

hidden_size = 4096
num_heads = 32

那么:

Q: 32个头
K: 32个头
V: 32个头

结构如下:

Head1: Q1 K1 V1
Head2: Q2 K2 V2
...
Head32: Q32 K32 V32

这就是:

MHA (Multi Head Attention)

MHA 最大的问题

在训练阶段问题不大。

但在推理阶段(尤其生成长文本):

模型要保存历史 Token 的

K Cache
V Cache

例如:

上下文长度 = 128K
Head数 = 32

KV Cache 会非常巨大。

实际上:

大模型推理时最大的瓶颈之一就是 KV Cache。 (IBM)


第一个解决方案:MQA

Google 之前提出:

MQA
(Multi Query Attention)

思路极其简单:

保留多个 Query Head

Q1
Q2
...
Q32

但是:

所有 Head 共用一个 K 和 V

K_shared
V_shared

变成:

Q1 ─┐
Q2 ─┤
Q3 ─┤
... ├── K_shared
Q32─┘   V_shared

即:

32个Q
1个K
1个V

KV Cache 直接缩小:

321

理论上节省:

32

的 KV 存储。

推理速度暴涨。 (ACL Anthology)


但 MQA 有副作用

虽然快了:

Q1
Q2
Q3
...
Q32

都在使用同一个:

K
V

很多 Head 的表达能力被压缩了。

效果通常会下降:

MHA > GQA > MQA

论文发现:

MQA 推理速度很好,但模型质量会有明显损失。 (ACL Anthology)


GQA 的核心思想

Google 想:

MQA太极端
MHA太昂贵

于是取中间方案:

Grouped Query Attention

假设:

32个Query Head

不要:

32个KV(MHA)

也不要:

1个KV(MQA)

而是:

8个KV

例如:

32个Q
8个KV

每4个 Query Head 共用一个 KV Head。


结构变成:

Q1
Q2
Q3
Q4
  ↓
 KV1

Q5
Q6
Q7
Q8
  ↓
 KV2

...

Q29
Q30
Q31
Q32
  ↓
 KV8

这就是:

Group Size = 4

数学表示:

设:

hq=QueryHeadshkv=KVHeads

那么:

group size=hqhkv

例如:

32/8=4

即:

4个Q共享1个KV

三种 Attention 对比

假设:

32个Query Head

MHA

32Q
32K
32V

结构:

Q1 -> K1,V1
Q2 -> K2,V2
...
Q32->K32,V32

KV Cache:

100%

效果:

最好

MQA

32Q
1K
1V

结构:

Q1
Q2
...
Q32

共享K,V

KV Cache:

1/32

效果:

下降明显

GQA

32Q
8K
8V

结构:

4个Q共享1个KV

KV Cache:

1/4

效果:

接近MHA

所以:

MHA ←→ GQA ←→ MQA

GQA 本质上是:

MHA 与 MQA 之间的折中方案

为什么 GQA 特别适合大模型?

因为推理时:

KV Cache 大小约为:

O(L×hkv×d)

其中:

所以:

如果:

32 KV Heads
→
8 KV Heads

则:

KV Cache:

75

对于:

32K
64K
128K

长上下文模型,

收益巨大。

这也是为什么:

等现代大模型广泛采用 GQA。


从代码角度理解

传统 Attention:

Q.shape = [B, L, 32, d]
K.shape = [B, L, 32, d]
V.shape = [B, L, 32, d]

GQA:

Q.shape = [B, L, 32, d]
K.shape = [B, L, 8, d]
V.shape = [B, L, 8, d]

然后:

K = repeat_interleave(K, 4)
V = repeat_interleave(V, 4)

逻辑上扩展成:

32 heads

供 Attention 使用。

例如:

KV1 → Q1,Q2,Q3,Q4
KV2 → Q5,Q6,Q7,Q8
...

从信息论角度理解

你之前问过:

RoPE 那些频率设计是不是拍脑袋?

GQA 反而比 RoPE 更容易理解。

GQA 背后的观察其实是:

Query Head 差异很大

不同 Head 学习不同模式:

语法
实体
位置
代码
推理

因此:

Q需要保留多样性

K/V Head 差异没那么大

论文实验发现:

很多 Head 的 Key/Value 存在冗余。

因此:

Q保持32个
KV压缩到8个

性能损失很小。 (Hugging Face)


一句话总结

GQA(Grouped Query Attention)可以理解成:

保留大量 Query Head 的表达能力,让多个 Query Head 共享较少数量的 Key/Value Head,从而大幅减少 KV Cache,占用更少显存,同时保持接近 Multi-Head Attention 的效果。 (ACL Anthology)

如果后面你准备从零实现 Llama/Qwen,我还可以继续讲:

  1. GQA 的完整数学推导
  2. Llama3 中 num_heads=32, num_kv_heads=8 的具体实现
  3. PyTorch 版 GQA 源码逐行解析
  4. KV Cache 为什么能从 GQA 中获得巨大收益
  5. FlashAttention 与 GQA 的关系。