好,我们从 0 开始,一步一步推导 GQA 是怎么来的。我会用公式说话,每个关键结论都会附上参考资料。
一、一切从一个问题开始:推理时显存爆了
Transformer 在训练时很快,但在推理生成文本时,每生成一个新词(token),都要把之前所有词的 Key 和 Value 重新从显存里读一遍。这些 Key 和 Value 被缓存起来,叫做 KV Cache。
对于标准的多头注意力(MHA),KV Cache 的大小是:
$$\text{KV Cache per layer} = 2 \times H \times L \times d_k \times \text{bytes}$$
其中 $H$ 是头数,$L$ 是序列长度,$d_k$ 是每个头的维度。层数一多、序列一长,这个数字就爆炸了。
这就是 GQA 要解决的核心问题:减少推理时 KV Cache 的显存占用和内存带宽消耗。
二、基础:标准 Scaled Dot-Product Attention
在讲 GQA 之前,我们必须先回到最原始的注意力公式。
给定输入序列的表示 $X \in \mathbb{R}^{L \times d_{model}}$,我们把它投影成 Query、Key、Value:
$$Q = X W_Q, \quad K = X W_K, \quad V = X W_V$$
其中 $W_Q, W_K, W_V \in \mathbb{R}^{d_{model} \times d_k}$。
然后计算注意力输出:
$$\text{Attention}(Q, K, V) = \text{softmax}\left(\frac{QK^T}{\sqrt{d_k}}\right) V$$
这就是 Vaswani 等人在 2017 年提出的原始注意力。
三、第一步演化:Multi-Head Attention (MHA)
原始 Transformer 把上述注意力做了 $H$ 次,得到 $H$ 个"头",每个头有自己的投影矩阵,这就是 MHA。
3.1 MHA 的公式
$$Q_i = X W_Q^{(i)}, \quad K_i = X W_K^{(i)}, \quad V_i = X W_V^{(i)}$$
其中 $i = 1, \dots, H$,每个 $W_Q^{(i)}, W_K^{(i)}, W_V^{(i)} \in \mathbb{R}^{d_{model} \times d_k}$。
每个头独立计算注意力:
$$\text{head}_i = \text{Attention}(Q_i, K_i, V_i) = \text{softmax}\left(\frac{Q_i K_i^T}{\sqrt{d_k}}\right) V_i$$
最后把所有头拼接起来,再投影:
$$\text{MultiHead}(Q, K, V) = \text{Concat}(\text{head}_1, \dots, \text{head}_H) W_O$$
其中 $W_O \in \mathbb{R}^{d_{model} \times d_{model}}$。
3.2 MHA 的问题
MHA 有 $H$ 个独立的 K 和 V 头。在自回归解码时,每生成一个新 token,都要把之前所有 token 的 $H$ 组 K 和 $H$ 组 V 从显存加载到计算单元。
KV Cache 大小(每层):
$$\text{Cache}_{\text{MHA}} = 2 \times H \times L \times d_k \times \text{bytes}$$
当模型很大(比如 70B 参数、80 层、64 头、8192 序列长度)时,KV Cache 可以达到 20 GB,光缓存就占掉一张 A100 的大部分显存。
四、第二步演化:Multi-Query Attention (MQA)
2019 年,Noam Shazeer 提出了 MQA(Multi-Query Attention)。他的想法很简单粗暴:
所有 Query 头共享同一套 K 和 V。
4.1 MQA 的公式
$$Q_i = X W_Q^{(i)}, \quad i = 1, \dots, H$$
$$K = X W_K, \quad V = X W_V$$
注意:这里只有 1 个 $W_K$ 和 1 个 $W_V$,所有 $H$ 个 Query 头共用同一组 $K$ 和 $V$。
每个头的注意力计算变成:
$$\text{head}_i = \text{softmax}\left(\frac{Q_i K^T}{\sqrt{d_k}}\right) V$$
4.2 MQA 的优势
KV Cache 直接减少到原来的 $1/H$:
$$\text{Cache}_{\text{MQA}} = 2 \times 1 \times L \times d_k \times \text{bytes}$$
推理速度大幅提升,因为每次解码只需要加载 1 组 K 和 1 组 V。PaLM、Falcon 等模型都采用了 MQA。
4.3 MQA 的问题
但 MQA 太极端了。所有 Query 头共用 1 个 K 和 1 个 V,相当于把模型的"注意力多样性"压缩到了极致。实验发现:
- 模型质量下降:困惑度(perplexity)比 MHA 高
- 训练不稳定:收敛困难
- 表达能力受限:不同的 Query 头无法学到不同的关注模式
五、第三步演化:Grouped Query Attention (GQA)
2023 年 5 月,Google Research 的 Joshua Ainslie 等人在论文 "GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints" 中提出了 GQA。
GQA 的核心思想是:
不要把所有 Query 头压到 1 个 KV 头上,而是分成若干组,每组共享 1 个 KV 头。
5.1 GQA 的公式
设 Query 头数为 $H$,KV 组数为 $G$(其中 $1 \leq G \leq H$)。
把 $H$ 个 Query 头分成 $G$ 组,每组有 $n_{rep} = H/G$ 个 Query 头。
$$Q_i = X W_Q^{(i)}, \quad i = 1, \dots, H$$
$$K_g = X W_K^{(g)}, \quad V_g = X W_V^{(g)}, \quad g = 1, \dots, G$$
每组内的 $n_{rep}$ 个 Query 头共享同一组 $K_g$ 和 $V_g$。
5.2 GQA 是 MHA 和 MQA 的统一框架
这是 GQA 最漂亮的地方——它是一个连续谱:
| 配置 | 含义 | 等价于 |
|---|---|---|
| $G = H$ | 每组 1 个 Query 头,各配独立 KV | MHA |
| $1 < G < H$ | 每组 $n_{rep}$ 个 Query 头共享 1 个 KV | GQA |
| $G = 1$ | 所有 Query 头共享 1 个 KV | MQA |
5.3 GQA 的 KV Cache
$$\text{Cache}_{\text{GQA}} = 2 \times G \times L \times d_k \times \text{bytes}$$
相比 MHA,减少了 $G/H$ 倍。例如 Llama 3 70B 用 $H=64, G=8$,KV Cache 减少为原来的 $1/8$。
六、GQA 的计算过程:一步一步推导
现在我们来走一遍 GQA 在前向传播中的完整计算流程,假设:
- 批次大小 $B = 1$
- 序列长度 $L$
- 模型维度 $d_{model} = 512$
- Query 头数 $H = 8$
- KV 组数 $G = 4$
- 每头维度 $d_k = d_{model} / H = 64$
- 每组 Query 头数 $n_{rep} = H/G = 2$
Step 1:线性投影
输入 $X \in \mathbb{R}^{L \times 512}$。
Query 投影($H = 8$ 个独立的投影):
$$Q = X W_Q \in \mathbb{R}^{L \times 512}$$
reshape 后:
$$Q \in \mathbb{R}^{L \times 8 \times 64}$$
Key/Value 投影(只有 $G = 4$ 个):
$$K = X W_K \in \mathbb{R}^{L \times 256}, \quad V = X W_V \in \mathbb{R}^{L \times 256}$$
reshape 后:
$$K \in \mathbb{R}^{L \times 4 \times 64}, \quad V \in \mathbb{R}^{L \times 4 \times 64}$$
注意:MHA 这里会有 $K \in \mathbb{R}^{L \times 8 \times 64}$,GQA 只有一半。
Step 2:扩展 KV 头(repeat_kv)
为了做矩阵乘法,K 和 V 的头数必须和 Q 对齐。GQA 的做法是把每个 KV 头复制 $n_{rep} = 2$ 次。
$$\tilde{K} = \text{repeat_kv}(K, n_{rep}=2) \in \mathbb{R}^{L \times 8 \times 64}$$
$$\tilde{V} = \text{repeat_kv}(V, n_{rep}=2) \in \mathbb{R}^{L \times 8 \times 64}$$
具体对应关系:
| Query 头 | 使用的 K/V 头 |
|---|---|
| $Q_1, Q_2$ | $K_1, V_1$ |
| $Q_3, Q_4$ | $K_2, V_2$ |
| $Q_5, Q_6$ | $K_3, V_3$ |
| $Q_7, Q_8$ | $K_4, V_4$ |
Step 3:计算注意力
对每个头 $i = 1, \dots, 8$:
$$\text{score}_i = \frac{Q_i \tilde{K}_i^T}{\sqrt{64}} \in \mathbb{R}^{L \times L}$$
$$\text{attn}_i = \text{softmax}(\text{score}_i) \in \mathbb{R}^{L \times L}$$
$$\text{head}_i = \text{attn}_i \cdot \tilde{V}_i \in \mathbb{R}^{L \times 64}$$
注意:虽然 $Q_1$ 和 $Q_2$ 共用 $K_1, V_1$,但它们的 $Q_1 \neq Q_2$,所以 $\text{attn}_1 \neq \text{attn}_2$,输出也不同。
Step 4:拼接与输出投影
$$\text{output} = \text{Concat}(\text{head}_1, \dots, \text{head}_8) W_O \in \mathbb{R}^{L \times 512}$$
七、三种机制的完整对比
| 指标 | MHA | GQA ($G=4, H=8$) | MQA |
|---|---|---|---|
| Query 头数 | $H=8$ | $H=8$ | $H=8$ |
| KV 头数 | $H=8$ | $G=4$ | $1$ |
| K/V 投影矩阵数 | 8 个 $W_K$, 8 个 $W_V$ | 4 个 $W_K$, 4 个 $W_V$ | 1 个 $W_K$, 1 个 $W_V$ |
| KV Cache(每层) | $2 \times 8 \times L \times 64$ | $2 \times 4 \times L \times 64$ | $2 \times 1 \times L \times 64$ |
| Cache 相对大小 | $1\times$ | $0.5\times$ | $0.125\times$ |
| 模型质量 | 最优 | 接近最优(损失 < 0.5%) | 明显下降 |
| 推理速度 | 最慢 | 接近 MQA | 最快 |
八、为什么 GQA 的质量损失很小?
核心直觉:KV 头之间存在冗余。不同的 K/V 头学到的是相似的位置和语义信息,没有必要每个 Query 头都配独立的 K/V。
GQA 保留了 $G$ 组独立的 KV 表示(而不是 MQA 的 1 组),所以模型仍然有足够的"注意力多样性"。实验表明,uptrained GQA 的困惑度和下游任务准确率非常接近 MHA。
九、从 MHA 转换到 GQA:Mean-Pooling
GQA 论文的一个重要贡献是:已经训练好的 MHA 模型不需要从头训练,可以通过 uptraining 转换成 GQA。
9.1 转换方法
对每个组的 KV 头做均值池化:
$$K_g^{\text{(GQA)}} = \frac{1}{n_{rep}} \sum_{i=(g-1)n_{rep}+1}^{g \cdot n_{rep}} K_i^{\text{(MHA)}}$$
$$V_g^{\text{(GQA)}} = \frac{1}{n_{rep}} \sum_{i=(g-1)n_{rep}+1}^{g \cdot n_{rep}} V_i^{\text{(MHA)}}$$
例如 $H=8, G=4, n_{rep}=2$:
$$K_1^{\text{(GQA)}} = \frac{K_1^{\text{(MHA)}} + K_2^{\text{(MHA)}}}{2}$$
$$K_2^{\text{(GQA)}} = \frac{K_3^{\text{(MHA)}} + K_4^{\text{(MHA)}}}{2}$$
9.2 Uptraining
转换后,用约 5% 的原始预训练计算量 继续微调,让模型适应新的 KV 共享结构。这比从头训练节省大量资源。
Llama 2 70B 就是用这种方式从 MHA 转换到 GQA 的。
十、总结:GQA 的演化脉络
2017 Vaswani et al. Scaled Dot-Product Attention → 基础
↓
2017 Vaswani et al. Multi-Head Attention (MHA) → H个独立Q/K/V头,质量最好,但KV Cache太大
↓
2019 Shazeer Multi-Query Attention (MQA) → 所有Q头共享1个K/V,Cache最小,但质量差
↓
2023 Ainslie et al. Grouped Query Attention (GQA) → 折中:分组共享,Cache中等,质量接近MHA
GQA 的本质:在 KV Cache 大小(推理效率)和注意力多样性(模型质量)之间做一个可调的权衡。通过超参数 $G$(KV 组数),你可以自由选择想要的位置。
参考资料
- Ainslie, J., Lee-Thorp, J., de Jong, M., Zemlyanskiy, Y., Lebrón, F., & Sanghai, S. (2023). GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints. EMNLP 2023. arXiv:2305.13245. PDF — GQA 的原始论文
- Shazeer, N. (2019). Fast Transformer Decoding: One Write-Head is All You Need. arXiv:1911.02150. — MQA 的原始论文
- Vaswani, A., et al. (2017). Attention Is All You Need. NeurIPS 2017. arXiv:1706.03762. — Transformer 和 MHA 的原始论文
- Touvron, H., et al. (2023). Llama 2: Open Foundation and Fine-Tuned Chat Models. — GQA 在工业界的应用验证

Comments | NOTHING