一句话总结:MHA、MQA、GQA 其实是同一个参数族——改变 KV 头的数量,就能在表示容量、KV Cache 和推理效率之间移动;哪一点最好,必须由具体模型和硬件上的测试决定。
23.1 问题不是“多少个张量”,而是头这一维有多大
第 22 章讲过,生成一个新 token 时,每一层都会算出新的 Q、K、V。Q 用完就可以丢掉,新 K、V 却要追加到这一层的 KV Cache 中。
先纠正一个很容易说错的地方:有 32 个头,不等于每层保存 64 个独立张量。常见实现中,每层的 Cache 仍是两个逻辑张量:
K cache: [B, H_KV, S, d_h]
V cache: [B, H_KV, S, d_h]
头数 是张量的一维。真正决定存储量的是这一维有多长:
其中 是 batch 大小, 是层数, 是已缓存 token 数, 是每头维度, 是每个元素的字节数。前面的 2 代表 K 和 V。
为了把三种结构放进同一张图,我们只需要两个数字:
- :Query 头数
- :Key 和 Value 的头数
并要求 能被 整除。每一组共享同一套 K、V 的 Query 头数是:
于是:
| 结构 | 条件 | 8 个 Query 头的例子 |
|---|---|---|
| MHA | 8 个 KV 头,每组 1 个 Q 头 | |
| GQA | 2 个 KV 头,每组 4 个 Q 头 | |
| MQA | 1 个 KV 头,8 个 Q 头共用 |
MHA 和 MQA 不是两座孤岛,GQA 也不是硬塞在中间的第四种魔法。它们就是这个参数族的两个端点和中间区域。
23.2 直觉:八位老师和几本笔记
继续用本书原来的教室例子。想象八位老师同时观察一堂课:
- MHA:八位老师各带一本笔记。每个人既能提出自己的问题,也能用自己的方式记录学生。
- MQA:八位老师仍然能提出不同的问题,但只能查同一本公共笔记。
- GQA:老师分成几组,每组共用一本笔记;组与组之间仍然保留不同记录。
这里的“提问方式”对应 Q 投影,“怎样索引和保存历史”对应 K、V 投影。共享 K、V 会减少保存历史所需的空间,但并不意味着各个 Query 头变成完全相同——它们仍有独立的 Q 行和不同的注意力输出。
在代码里,三者的形状可以统一写成:
q: [B, H_Q, T_q, d_h]
k: [B, H_KV, T_k, d_h]
v: [B, H_KV, T_k, d_h]
若 、,Query 头与 KV 头的映射就是:
Q0 Q1 Q2 Q3 -> KV0
Q4 Q5 Q6 Q7 -> KV1
概念上可以把每个 KV 头重复 次,再按普通 MHA 计算。但这个“重复”只适合解释或做参考实现;高效内核应直接把 Query 头映射到对应的 KV 头,避免真的把 Cache 扩大回去。
23.3 内存和参数到底省了多少
23.3.1 KV Cache
沿用第 22 章的数字:32 层、、1024 tokens、FP16/BF16(每元素 2 bytes)、batch 为 1。
| 结构 | 精确字节数 | 二进制单位 | 相对 MHA | |
|---|---|---|---|---|
| MHA | 32 | 536,870,912 | 512 MiB | 100% |
| GQA | 8 | 134,217,728 | 128 MiB | 25% |
| MQA | 1 | 16,777,216 | 16 MiB | 3.125% |
如果序列增长到 4096 tokens,三者分别是 2 GiB、512 MiB 和 64 MiB。这个比例是准确的,但不要立刻写成“并发一定提高 4 倍或 32 倍”。实际服务还要给模型权重、临时工作区、激活、内存碎片和调度留空间;上下文能否拉长,也受位置编码和模型训练长度约束。
23.3.2 投影参数
忽略 bias,令模型宽度为 。Q 和输出投影各有 个参数,K、V 投影各有 个参数:
所以:
| 结构() | 注意力投影参数 |
|---|---|
| MHA, | |
| GQA, | |
| MQA, |
32/8 的 GQA 相比 MHA 少了 37.5% 的注意力投影权重。它不是“整个 7B 模型少 37.5%”,更不是固定少 6%;全模型比例还取决于 FFN、词嵌入、层数和是否共享权重。
23.4 一个实现覆盖 MHA、GQA 和 MQA
下面的代码把张量统一成 [B, H, T, d_h]。grouped_attention_reference 会显式重复 K、V,便于验证;grouped_attention 使用 PyTorch 的原生 enable_gqa 路径。
import torch
import torch.nn.functional as F
def repeat_kv(x, num_query_heads):
"""Reference expansion: [B, H_kv, T, D_h] -> [B, H_q, T, D_h]."""
num_kv_heads = x.size(1)
if num_query_heads <= 0 or num_kv_heads <= 0:
raise ValueError("head counts must be positive")
if num_query_heads % num_kv_heads != 0:
raise ValueError("num_query_heads must be divisible by num_kv_heads")
return x.repeat_interleave(num_query_heads // num_kv_heads, dim=1)
def validate_gqa_shapes(q, k, v):
if q.ndim != 4 or k.ndim != 4 or v.ndim != 4:
raise ValueError("q, k, and v must be four-dimensional")
if q.size(0) != k.size(0) or q.size(0) != v.size(0):
raise ValueError("q, k, and v must have the same batch size")
if q.size(1) <= 0 or k.size(1) <= 0:
raise ValueError("head counts must be positive")
if q.size(1) % k.size(1) != 0 or k.size(1) != v.size(1):
raise ValueError("invalid GQA head counts")
if q.size(-1) != k.size(-1):
raise ValueError("q and k must have the same head dimension")
if k.size(-2) != v.size(-2):
raise ValueError("k and v must have the same sequence length")
def grouped_attention_reference(q, k, v, *, is_causal):
validate_gqa_shapes(q, k, v)
k_full = repeat_kv(k, q.size(1))
v_full = repeat_kv(v, q.size(1))
return F.scaled_dot_product_attention(
q, k_full, v_full, is_causal=is_causal, dropout_p=0.0
)
def grouped_attention(q, k, v, *, is_causal):
validate_gqa_shapes(q, k, v)
return F.scaled_dot_product_attention(
q,
k,
v,
is_causal=is_causal,
dropout_p=0.0,
enable_gqa=q.size(1) != k.size(1),
)
同一函数里:
q有 8 头、k/v有 8 头,就是 MHA;q有 8 头、k/v有 2 头,就是 GQA;q有 8 头、k/v有 1 头,就是 MQA。
这里故意把 is_causal 设成必填参数,因为 Prefill 和带 Cache 的 Decode 不能盲目共用同一个值:
- 同长度的训练或 Prefill,传
is_causal=True; - 单 token Decode 且 K/V 只包含过去与当前 token 时,传
is_causal=False,因为 Cache 里根本没有未来 token; - 一次解码多个 token 的 chunked Decode,要根据实际 cache position 构造右下对齐的 causal mask。
PyTorch 对非方阵的 is_causal=True 使用左上对齐遮罩。因此对 [T_q=1, T_k>1] 的单 token Decode 直接传 True,反而只会留下最早的 key。这是 mask 的位置边界,与 MHA、GQA 或 MQA 本身无关。
截至 2026 年 7 月,PyTorch 文档仍把 enable_gqa 标成实验性能力,并注明后端与张量类型限制。API 存在不等于你当前的 GPU、dtype 和 kernel 一定走最优路径;部署前要检查版本并做 profile。显式重复的参考实现一定会展开 K、V,而原生 API 在某个后端上是否避免展开,仍要以实际选中的 kernel 为准。
本章配套验证同时测试了 MHA、GQA、MQA 的 float32/float64 输出和梯度,还单独覆盖了单 token Cache Decode、非方阵 causal mask 和 K/V 不同特征维度的边界。
23.5 MHA 检查点怎样变成 GQA
GQA 论文提出的 uptraining 不是改一行配置就结束,而是两步:
- 把每组 MHA 的 K、V 投影头取平均,作为新的 GQA K、V 初值;Q 和输出投影不变。
- 从这个初值继续预训练,让模型适应共享后的 K、V。
论文中的“5%”指的是原始预训练计算量的 5%,不是一个对所有模型都成立的“使用 5% 数据微调”配方。论文在其 T5.1.1 实验中报告,uptrained GQA 取得接近 MHA 的质量和接近 MQA 的速度;这是一项有边界的实验结果,不是跨模型保证。
PyTorch 的 nn.Linear.weight 形状是 [out_features, in_features]。因此转换 K 或 V 权重时,头这一维在前面:
def mean_pool_kv_projection(weight, num_query_heads, num_kv_heads):
"""Pool [H_q * D_h, D_model] into [H_kv * D_h, D_model]."""
if weight.ndim != 2:
raise ValueError("weight must be two-dimensional")
if num_query_heads <= 0 or num_kv_heads <= 0:
raise ValueError("head counts must be positive")
if num_query_heads % num_kv_heads != 0:
raise ValueError("num_query_heads must be divisible by num_kv_heads")
if weight.size(0) % num_query_heads != 0:
raise ValueError("weight rows must be divisible by num_query_heads")
head_dim = weight.size(0) // num_query_heads
group_size = num_query_heads // num_kv_heads
return (
weight.reshape(num_query_heads, head_dim, weight.size(1))
.reshape(num_kv_heads, group_size, head_dim, weight.size(1))
.mean(dim=1)
.reshape(num_kv_heads * head_dim, weight.size(1))
)
若线性层带 bias,也要按同样的头分组平均。转换完成后仍要重新评估 loss、下游任务、长上下文质量和服务性能,不能把平均池化本身当作质量证明。
23.6 真实模型怎样选择
下面只列几个可以从论文或官方配置核对的历史快照;同一模型家族的不同尺寸和版本可能不同,最终应读取你实际加载的 config.json。
| 模型快照 | Q 头 | KV 头 | 结构 | 每组 Q 头 |
|---|---|---|---|---|
| Llama 2 7B | 32 | 32 | MHA | 1 |
| Llama 2 70B | 64 | 8 | GQA | 8 |
| Mistral-7B-v0.1 | 32 | 8 | GQA | 4 |
| Qwen2-7B | 28 | 4 | GQA | 7 |
Hugging Face 风格的配置通常长这样:
{
"num_attention_heads": 32,
"num_key_value_heads": 8
}
判断方法很简单:
- 两者相等:MHA;
num_key_value_heads为 1:MQA;- 介于两者之间:GQA;
- 还要确认 Q 头数能整除 KV 头数。
“8 个 KV 头是最佳答案”并不是定律。8 便于某些 2/4/8 路张量并行配置,但若 KV 头数小于并行度,某些系统需要复制 KV 头或采用不同切分;质量曲线也随模型、数据和训练预算变化。真正的选择步骤是:先定候选 ,再同时测验证质量、TPOT/吞吐、峰值显存以及目标并行策略。
23.7 MQA 和 GQA 的论文结论该怎样读
2019 年的 MQA 论文在它测试的翻译任务上报告了显著的解码效率提升和较小的质量下降。2023 年的 GQA 论文则在它的 T5.1.1 uptraining 实验中报告,GQA 可以靠近 MHA 的质量,同时接近 MQA 的速度。
这两句话都有“在它测试的设置中”。不能从中推出:
- MHA 在任何模型上都质量最高;
- MQA 在前沿模型上一定不可接受;
- GQA 一定不会掉点;
- 从 32 个 KV 头减到 8 个,Decode 就一定快 4 倍;
- 8 个 KV 头适合所有模型。
架构只改变了可能的权衡面,最终落点还由训练、内核、硬件、batch、上下文长度和服务框架共同决定。
23.8 章末总结
- KV Cache 通常是每层两个逻辑张量, 是它们的一维,不是“每头一个张量”。
- MHA、GQA、MQA 可统一为 参数族;两端分别是 和 。
- Cache 大小与 线性相关,但实际并发和速度不会只由这个比例决定。
- 原生 GQA 内核应直接映射 Query 头到 KV 头;显式
repeat_kv更适合作为参考实现。 - MHA 转 GQA 的 mean pooling 只是初始化,论文还继续使用了 5% 原始预训练计算量做 uptraining。
- 没有永恒的“8 头甜点”。读实际配置,在目标硬件上同时测质量与性能。
参考文献
- Attention Is All You Need(Vaswani et al., 2017)
- Fast Transformer Decoding: One Write-Head is All You Need(Shazeer, 2019)
- GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints(Ainslie et al., 2023)
- Llama 2: Open Foundation and Fine-Tuned Chat Models(Touvron et al., 2023)
- Mistral-7B-v0.1 官方配置
- Qwen2-7B 官方配置
- PyTorch scaled_dot_product_attention 文档
下一章
GQA 减少的是每个 token 要保存多少 K、V,却没有改变“当前 Query 要看哪些历史 token”。第 24 章我们继续追问:如果不再看完整历史,而只看一部分位置,能不能把 Attention 的计算也省下来?这就是 Sparse Attention 要解决的问题。