从MHA到MQA到GQA

发布于 20 天前  64 次阅读


好,我们从 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 组数),你可以自由选择想要的位置。


参考资料

  1. 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. PDFGQA 的原始论文
  2. Shazeer, N. (2019). Fast Transformer Decoding: One Write-Head is All You Need. arXiv:1911.02150. — MQA 的原始论文
  3. Vaswani, A., et al. (2017). Attention Is All You Need. NeurIPS 2017. arXiv:1706.03762. — Transformer 和 MHA 的原始论文
  4. Touvron, H., et al. (2023). Llama 2: Open Foundation and Fine-Tuned Chat Models.GQA 在工业界的应用验证

一沙一世界,一花一天堂。君掌盛无边,刹那成永恒。