分组查询注意力(GQA)与多查询注意力(MQA)深度实战:从 KV 缓存瓶颈、头维度折叠到推理吞吐与显存占用的工程全解
自回归大模型推理的真实瓶颈,往往不在算力而在显存带宽与 KV 缓存占用。当模型层数、上下文长度、并发数同时放大,多头注意力(MHA)为每个查询头都保留一份独立的 Key/Value,KV 缓存体积随头数线性膨胀,最终把解码速度牢牢钉在 HBM 带宽上限上。分组查询注意力(Grouped Query Attention, GQA)与多查询注意力(Multi-Query Attention, MQA)用一种近乎暴力却极其有效的思路——让多个查询头共享同一组 K/V——把这道墙推后了一个数量级。本文从第一性原理出发,拆解 MHA 的 KV 缓存代价、MQA/GQA 的折叠机制、权重 upcycle 转换、推理期显存/带宽/吞吐的定量关系,并给出可插拔的参考实现与生产陷阱清单。
一、MHA 的 KV 缓存墙:为什么注意力头数成了推理瓶颈
Transformer 的自回归解码有一个根本特征:第 t 步生成一个 token,第 t+1 步的注意力却必须能看到第 1..t 步所有 token 的 Key/Value。于是工程上会把历史 K/V 缓存起来,形成 KV Cache。
KV 缓存的显存占用可以精确写出:
KV_bytes = 2 (K 和 V)
× num_layers
× num_kv_heads
× seq_len
× head_dim
× dtype_bytes
在标准的 MHA 配置下,num_kv_heads == num_heads。以 LLaMA-2-70B(80 层、num_heads=64、head_dim=128、FP16)为例,单条 4K 上下文的 KV 缓存约为:
2 × 80 × 64 × 4096 × 128 × 2 ≈ 1.07 GB
这还只是一条序列。若并发 64 条、上下文拉长到 32K,KV 缓存轻松突破 1.7 TB——远超模型权重本身(70B FP16 约 140GB)。更要命的是,解码每个新 token 都要把整段 KV 从 HBM 读进计算单元做 attention,因此解码吞吐受限于内存带宽而非算力(memory-bound)。头数越多,每次读取的字节越多,解码越慢。
这就是 GQA/MQA 要解决的"墙":在不显著牺牲质量的前提下,砍掉 KV 缓存里的头数维度。
二、第一性原理:多头注意力的 K/V 投影
标准 MHA 对每个注意力头 i ∈ [0, h) 都有独立的投影矩阵:
Q_i = X W_Q_i W_Q_i ∈ R^{d × d_h}
K_i = X W_K_i W_K_i ∈ R^{d × d_h}
V_i = X W_V_i W_V_i ∈ R^{d × d_h}
O_i = softmax(Q_i K_i^T / √d_h) V_i
其中 d 是模型隐藏维,h 是头数,d_h = d / h 是单头维度。W_K、W_V 各自被复制成 h 份,意味着 K/V 头与 Q 头一一对应。KV 缓存因此天然带有 h 这个因子——这正是第四节要消去的对象。
关键洞察:Query 负责"提问"(决定关注哪里),Key/Value 负责"被检索的内容表示"。 多个 Query 头去问同一组 Key/Value 表示,在表达上损失有限,却能让 KV 缓存骤减 h 倍。
三、MQA:所有查询头共享一组 K/V
MQA(Shazeer, 2019)把 K/V 的头数直接压到 1:
W_K, W_V ∈ R^{d × d_h} # 仅 1 个 K/V 头
K = X W_K # 所有 Q 头共用
V = X W_V
O_i = softmax(Q_i K^T / √d_h) V # 每个 Q 头 i 仍独立产出 O_i
- 收益:KV 缓存缩小
h倍(MQA 下num_kv_heads = 1),解码期 HBM 读取量同比例下降,memory-bound 场景解码速度显著提升。 - 代价:所有 Q 头被迫从同一份 K/V 表示里提取信息,模型capacity下降,训练更易不稳定、收敛质量略损,且生成多样性可能退化。
MQA 是"极端但干净"的端点:它是 GQA 在 g=1 时的特例。
四、GQA:分组共享——MHA 与 MQA 的连续谱
GQA(Ainslie et al., 2023)引入组数 g = num_kv_heads,把 h 个 Q 头均分成 g 组,每组共享一组 K/V:
g = num_kv_heads # 1 ≤ g ≤ h,且通常 g | h
每组 j 的 Q 头共享 K_j, V_j
O_i = softmax(Q_i K_j(i)^T / √d_h) V_j(i)
g = h→ 每个 Q 头独立 K/V → 退化为 MHA;g = 1→ 所有 Q 头共享一组 → 退化为 MQA;1 < g < h→ GQA,即二者的连续谱。
现代主流模型把 g 取为 2/4/8:LLaMA-2-7B/13B 用 g=32(即 MHA),70B 用 g=8;LLaMA-3 全系 g=8;Falcon-40B 用 MQA(g=1);Mistral-7B 用 g=8。经验上 g=8 能在"质量损失 < 1%"与"KV 缓存缩减 8 倍"之间取得甜点。
五、数学形式与 head 维度对齐
设 Q 头索引 i,其所属组 j = i // (h / g)。该组的 K/V 投影为:
K_j = X W_K_j W_K_j ∈ R^{d × d_h}, j ∈ [0, g)
V_j = X W_V_j
score_{i,t} = (Q_i · K_j,t) / √d_h
α_{i,t} = softmax_t(score_{i,t})
O_i = Σ_t α_{i,t} V_j,t
注意 head_dim 在所有头(Q/K/V)中保持一致,组间只是"共享 K/V 的来源",并不改变每个 Q 头输出 d_h 维向量,因此 GQA 不破坏 MHA 的输出张量形状,可无缝替换。
从参数看,GQA 仅把 W_K、W_V 的 head 维从 h·d_h 压到 g·d_h:
W_K ∈ R^{d × (g·d_h)} # 而非 h·d_h
推理期 KV 缓存头数因子由 h 变为 g,其余(层数、序列长、head_dim、dtype)不变。
六、权重转换:从 MHA 到 GQA/MQA 的 upcycle 与下采样
从头训练 GQA 最理想,但已有 MHA 检查点时可用下采样(down-sampling / upcycling)复用:
- 均匀下采样:将 MHA 的
h个 K/V 头按下标等距采样g个,直接作为 GQA 的 K/V 头。简单但可能丢失中间信息。 - 平均融合(mean pooling):对每组内
h/g个 MHA 的 K/V 头求平均,得到该组 GQA 的 K/V 头。这是实践中最常用的初始化,能更平滑地保留表示。 - Upcycling 继续训练:用平均融合初始化 GQA 权重后,再用较小学习率在自有语料上续训数千步,质量可恢复到接近原生 GQA(Ainslie 等报告 MQA/GQA upcycle 后困惑度接近从零训练)。
下采样公式(平均融合):
W_K_j = (1 / (h/g)) Σ_{i ∈ group_j} W_K_i^MHA
W_V_j = (1 / (h/g)) Σ_{i ∈ group_j} W_V_i^MHA
这给"已有 MHA 大模型想获得 GQA 推理收益"的团队一条低成本路径:不必重训,只需转换 K/V 投影并少量续训。
七、推理视角:KV 缓存显存、带宽与吞吐
把 GQA 代入第一节公式:
KV_bytes(GQA) = 2 × num_layers × g × seq_len × head_dim × dtype_bytes
对比 MHA(num_kv_heads = h):
压缩比 = h / g
以 h=64, g=8 为例,KV 缓存缩小 8 倍。在 memory-bound 的解码阶段,每步需读取的 KV 字节同步下降 8 倍,HBM 带宽压力骤减,单卡可支撑更长的上下文与更高的并发。
但注意一个常被忽视的边界:当 seq_len 较短、或 batch/并发很小,使每步计算量(受 compute-bound 主导)而非 KV 读取主导时,GQA 的加速会被稀释。GQA 的最大收益出现在长上下文 + 高并发解码的生产负载上,这与真实 LLM 服务(长文档、多轮对话、RAG)高度吻合。
另一方面,GQA 会轻微增加 prefill(首包)计算:因为 Q 头数 h 仍远大于 K/V 头数 g,Q·K^T 的矩阵乘中 Q 侧维度更大。但 prefill 通常是 compute-bound,且绝对耗时远小于长序列解码,净效应仍为正。
八、可插拔参考实现
下面给出一个支持 MHA / MQA / GQA 统一切换的 PyTorch 风格实现,核心是 num_kv_heads 参数与 repeat_interleave 的对齐逻辑:
import torch
import torch.nn.functional as F
class Attention(torch.nn.Module):
def __init__(self, d, num_heads, num_kv_heads=None, max_seq=4096):
super().__init__()
self.d = d
self.h = num_heads
self.d_h = d // num_heads
# GQA: num_kv_heads 可小于 num_heads;None 表示 MHA
self.g = num_kv_heads if num_kv_heads else num_heads
assert self.h % self.g == 0, "num_heads 必须能被 num_kv_heads 整除"
self.group_size = self.h // self.g
self.W_q = torch.nn.Linear(d, self.h * self.d_h, bias=False)
self.W_k = torch.nn.Linear(d, self.g * self.d_h, bias=False) # 仅 g 个 K/V 头
self.W_v = torch.nn.Linear(d, self.g * self.d_h, bias=False)
self.W_o = torch.nn.Linear(self.h * self.d_h, d, bias=False)
self.cache_k = torch.zeros(self.g, max_seq, self.d_h)
self.cache_v = torch.zeros(self.g, max_seq, self.d_h)
def forward(self, x, start=0):
B, T, _ = x.shape
q = self.W_q(x).view(B, T, self.h, self.d_h) # (B,T,h,d_h)
k = self.W_k(x).view(B, T, self.g, self.d_h) # (B,T,g,d_h)
v = self.W_v(x).view(B, T, self.g, self.d_h)
# 写入 KV 缓存(仅 g 个头)
self.cache_k[:, start:start + T] = k[0]
self.cache_v[:, start:start + T] = v[0]
Kc = self.cache_k[:, :start + T] # (g, L, d_h)
Vc = self.cache_v[:, :start + T]
# 把 g 个 K/V 头 repeat 到 h 个 Q 头(组内共享)
Krep = Kc.repeat_interleave(self.group_size, dim=0) # (h, L, d_h)
Vrep = Vc.repeat_interleave(self.group_size, dim=0)
scores = torch.einsum("bthd,hld->bhtl", q, Krep) / (self.d_h ** 0.5)
attn = F.softmax(scores, dim=-1)
out = torch.einsum("bhtl,hld->bthd", attn, Vrep)
return self.W_o(out.reshape(B, T, -1))
# 用法对照:
# MHA : Attention(d, num_heads=64, num_kv_heads=64)
# GQA : Attention(d, num_heads=64, num_kv_heads=8) # LLaMA-2-70B 风格
# MQA : Attention(d, num_heads=64, num_kv_heads=1) # Falcon 风格
要点:repeat_interleave 把 g 个 K/V 头按组摊平到 h 个 Q 头,使每个 Q 头在 attention 计算时拿到"正确的一份"共享 K/V,且与 MHA 的张量形状完全兼容——这正是 GQA 能直接替换 MHA 的工程基础。
九、变体谱系与对照
| 方案 | num_kv_heads | KV 缓存倍数 | 质量 | 典型代表 |
|---|---|---|---|---|
| MHA | h | 1× | 基准 | 早期 GPT、LLaMA-2-7B/13B |
| MQA | 1 | 1/h× | 略降、训练不稳 | Falcon-40B、PaLM |
| GQA | 1| g/h× |
接近 MHA |
LLaMA-2-70B、LLaMA-3 全系、Mistral |
|
| MLA | 低秩潜空间 | 更小 | 接近/更优 | DeepSeek-V2/V3(潜注意力) |
| 跨层 KV 共享 | — | 进一步降 | 略降 | GPT-J 风格层间复用 |
MLA(Multi-head Latent Attention) 走另一条路:不共享 K/V 头,而是把 K/V 压缩到低秩潜变量再展开,理论上 KV 缓存更小且质量更好,是 DeepSeek 系列的选择。GQA 胜在零结构改动、即插即用、与现有训练/推理栈完全兼容,因此成为当下开源模型的默认解。
十、12 项生产陷阱清单
- 整除约束:
num_heads必须能被num_kv_heads整除,否则分组无法均分,加载权重时易静默错位。 - 张量形状误配:实现里若忘记
repeat_interleave而直接用expand,某些框架的 in-place 操作会触发梯度/广播错误。 - KV 缓存只按 g 头分配:缓存维度应是
g而非h,按h分配会浪费h/g倍显存,背离 GQA 初衷。 - prefill 与 decode 不一致:prefill 阶段常 fused(如 FlashAttention),decode 阶段走缓存;两者 K/V 头摊平逻辑必须一致,否则数值漂移。
- upcycle 平均顺序:平均融合应按"组"聚合并保持 head_dim 对齐,逐元素均值若混入未初始化头会污染表示。
- 位置编码耦合:GQA 不改变 RoPE/ALiBi 的施加方式,但共享 K/V 后位置信息仍须逐头施加,避免把旋转位置误加在共享 K/V 上导致语义混乱。
- 量化 KV 缓存:GQA 让 KV 体积变小,但若仍对 KV 做 4-bit 量化,相对收益下降且可能引入精度坑;低 bit 量化与 GQA 可叠加但需分别验证。
- 并行切分(TP):张量并行下 K/V 头数
g也要能被 TP 度整除;否则某些 rank 分不到完整组,all-reduce 前需重排。 - MQA 训练不稳定:
g=1时单组 K/V 承载全部 Q 头,学习率与 warmup 要更保守,必要时退回g≥2。 - 注意力 mask 与共享头:因果 mask 在共享 K/V 时仍按 Q 头位置施加,不能因为 K/V 共享就放宽 mask,否则信息泄露。
- 长上下文外推:GQA 本身不影响位置外推(NTK/YaRN 在 RoPE 侧处理),但共享 K/V 下长序列 KV 读取减少,外推收益更明显——需联合测试。
- 服务侧 batch 拼接:高并发时不同请求 seq_len 不同,KV 缓存按
g头分配后,padding/ragged 管理要与 PagedAttention 等分页机制配合,避免碎片。
十一、可复现工具箱
Attention(d, h, g):上文统一实现,改g即可在 MHA/GQA/MQA 间切换做消融。- 显存速算:把模型配置代入
2 × layers × g × seq × d_h × 2,对比g=h即得压缩比。 - upcycle 脚本骨架:遍历 MHA 的
W_K/W_V,按group_size分段求均值写入 GQA 权重,再用小学习率续训。 - 对照基准:以 MHA 困惑度为基线,记录 GQA 在各
g下的质量退化曲线,选取"质量损失 < 1% 且压缩比最大"的g。
十二、总结
GQA/MQA 的本质是用表示容量的微小让步,换取 KV 缓存的成倍缩减与解码带宽压力的成倍释放。在长上下文、高并发的真实 LLM 服务负载下,这道权衡几乎总是划算的——这也是为何 LLaMA-3、Mistral、Gemma 等当代开源模型无一例外采用 GQA。工程上它的美妙之处在于零结构侵入:只需把 K/V 投影头数从 h 改为 g,用 repeat_interleave 摊平到 Q 头,即可与现有训练、量化、并行、分页缓存栈无缝对接。当你的推理服务被 KV 缓存显存或解码带宽卡住时,GQA 往往是用最小改动撬动最大吞吐的那颗螺丝。

发表评论 取消回复