一句话定义: 分组查询注意力(Grouped Query Attention,GQA)是一种「折中型」的注意力机制——它把多个 Query(查询)头分成若干组,让同组的 Query 头共享同一套 Key/Value(键值),在几乎不掉精度的情况下,大幅减少推理时的内存占用、提升速度。
为什么需要它:KV 缓存的烦恼
要理解 GQA 解决了什么问题,先得知道现代大模型推理时的一个隐形负担:KV 缓存。
大模型生成文字时,是一个 token 一个 token 往外蹦的。每生成一个新的 token,都需要用到「之前所有 token」的 Key 和 Value(你可以把它们理解为「之前每个词的特征记录」)。为了不算重复账,模型会把这些历史 K/V 存起来复用,这就是 KV 缓存。
问题在于:这个缓存会随序列长度线性增长,而且模型有多少个注意力头,缓存就要存多少份。序列一长、头数一多,显存很快就被吃光,成为推理速度的瓶颈。于是研究者们开始琢磨:能不能少存几份 K/V?
图 1:KV 缓存随序列长度线性增长——每多一个 token 就多存一份 K/V,序列越长显存压力越大,成为推理速度的瓶颈。
三个翻译小组的故事
想象一家翻译公司,有 8 个翻译员(对应 8 个 Query 头),他们工作时需要查阅词典(对应 K/V)。公司有三种运营方案:
- 方案 A(MHA,多头注意力):每个翻译员都配一本专属词典,共 8 本。最准,但书架最挤。
- 方案 B(MQA,多查询注意力):8 个人共用 1 本词典。书架最省,但抢着翻页、有些词还查不到,翻译质量下降。
- 方案 C(GQA,分组查询注意力):把 8 个人分成 4 组,每组 2 人共用 1 本词典,共 4 本。书架省了一半,质量也基本不掉。
GQA 就是方案 C——介于「全独立」和「全共享」之间的中间路线。
在 MHA → MQA 的光谱上定位
用一组公式把三者关系说清楚。设模型有
- MHA:
(每个 Query 头都有自己的 K/V) - MQA:
(所有 Query 头共享同一份 K/V) - GQA:
(Query 头分成 组,每组共享一份 K/V)
换句话说,GQA 是 MHA 和 MQA 之间的一个旋钮:把
图 2:GQA 的分组共享——8 个 Query 头分成 4 组,每组 2 个 Q 头共用 1 套 K/V(n_rep = 2),复制对齐到 8 个 Q 头后照常算标准注意力,KV 缓存只存 4 份。
算一算:KV 缓存到底省了多少
来看个小例子,直观感受 GQA 的省内存效果。假设:Query 头数
- MHA:
- GQA(G=4):
- MQA(G=1):
也就是说,GQA(4 组)相比 MHA 直接省了一半的 KV 缓存,组数越少省得越多。在大模型那种几十亿、上百亿参数、序列动辄几千上万的规模下,这个节省非常可观——既意味着更小的显存压力,也意味着更快的推理速度(毕竟要从显存里读取的数据变少了)。
图 3:KV 缓存大小对比(8 头、维度 128、序列长 1000)——MHA 约 205 万、GQA 约 102 万(省一半)、MQA 约 26 万(省 7/8);GQA 在精度几乎不掉的前提下砍掉一半缓存。
PyTorch 代码:GQA 长什么样
下面这段代码展示了 GQA 的核心结构。关键在于:Q 用 n_query_heads 组投影,而 K/V 只用 n_kv_heads 组投影(少于 Q 头数),然后把每个 KV 头「复制」若干次对齐到 Q 头数,再做标准注意力计算。
import torch
import torch.nn as nn
import torch.nn.functional as F
class GroupedQueryAttention(nn.Module):
def __init__(self, d_model, n_query_heads, n_kv_heads):
super().__init__()
self.n_q_heads = n_query_heads # Query 头数(如 8)
self.n_kv_heads = n_kv_heads # KV 头数(如 4,少于 Q 头数 → 省内存)
self.head_dim = d_model // n_query_heads
# Query 投影:输出仍是 n_q_heads 组
self.q_proj = nn.Linear(d_model, n_query_heads * self.head_dim)
# K/V 投影:只输出 n_kv_heads 组(这是省内存的关键)
self.k_proj = nn.Linear(d_model, n_kv_heads * self.head_dim)
self.v_proj = nn.Linear(d_model, n_kv_heads * self.head_dim)
self.out_proj = nn.Linear(d_model, d_model)
# 每个 KV 头被几个 Query 头共享
self.n_rep = n_query_heads // n_kv_heads
def forward(self, x):
B, L, _ = x.shape
# 投影并重塑:B, L, n_heads, head_dim -> B, n_heads, L, head_dim
q = self.q_proj(x).view(B, L, self.n_q_heads, self.head_dim).transpose(1, 2)
k = self.k_proj(x).view(B, L, self.n_kv_heads, self.head_dim).transpose(1, 2)
v = self.v_proj(x).view(B, L, self.n_kv_heads, self.head_dim).transpose(1, 2)
# 把 KV 头沿「头维度」复制 n_rep 次,对齐到 Q 头数(GQA 的核心操作)
k = k.repeat_interleave(self.n_rep, dim=1)
v = v.repeat_interleave(self.n_rep, dim=1)
# 标准 scaled dot-product attention
scores = (q @ k.transpose(-2, -1)) / (self.head_dim ** 0.5)
attn = F.softmax(scores, dim=-1)
out = attn @ v # B, n_q_heads, L, head_dim
out = out.transpose(1, 2).reshape(B, L, -1)
return self.out_proj(out)完整代码
把上面的模块拼起来,跑一个最小例子和一次反向传播:
import torch
import torch.nn as nn
import torch.nn.functional as F
class GroupedQueryAttention(nn.Module):
def __init__(self, d_model, n_query_heads, n_kv_heads):
super().__init__()
self.n_q_heads = n_query_heads
self.n_kv_heads = n_kv_heads
self.head_dim = d_model // n_query_heads
self.q_proj = nn.Linear(d_model, n_query_heads * self.head_dim)
self.k_proj = nn.Linear(d_model, n_kv_heads * self.head_dim)
self.v_proj = nn.Linear(d_model, n_kv_heads * self.head_dim)
self.out_proj = nn.Linear(d_model, d_model)
self.n_rep = n_query_heads // n_kv_heads
def forward(self, x):
B, L, _ = x.shape
q = self.q_proj(x).view(B, L, self.n_q_heads, self.head_dim).transpose(1, 2)
k = self.k_proj(x).view(B, L, self.n_kv_heads, self.head_dim).transpose(1, 2)
v = self.v_proj(x).view(B, L, self.n_kv_heads, self.head_dim).transpose(1, 2)
# KV 头复制对齐 Q 头(GQA 的核心:同组 Q 头共享同一份 KV)
k = k.repeat_interleave(self.n_rep, dim=1)
v = v.repeat_interleave(self.n_rep, dim=1)
scores = (q @ k.transpose(-2, -1)) / (self.head_dim ** 0.5)
attn = F.softmax(scores, dim=-1)
out = attn @ v
out = out.transpose(1, 2).reshape(B, L, -1)
return self.out_proj(out)
# 假数据:batch=2,序列长 16,模型维度 512
torch.manual_seed(0)
x = torch.randn(2, 16, 512)
# 8 个 Query 头,4 个 KV 头(每个 KV 头被 2 个 Q 头共享)
gqa = GroupedQueryAttention(d_model=512, n_query_heads=8, n_kv_heads=4)
out = gqa(x)
print(out.shape) # torch.Size([2, 16, 512])
# 训练一步:用 MSE 损失跑一次反向传播
target = torch.randn_like(out)
loss = F.mse_loss(out, target)
loss.backward()
print(f"loss: {loss.item():.4f}")小提示:实际推理引擎(如 vLLM、TensorRT-LLM)会做更聪明的优化——不真的复制 KV 头,而是让多个 Q 头直接指向同一块 KV 内存,这样连「复制」的开销都省了。上面的
repeat_interleave只是教学写法,便于理解「分组共享」这件事。
为什么它成了主流
GQA 最早由 Google 研究团队在 2023 年的论文《GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints》中正式提出。论文里还给出了一个实用技巧:可以从一个已经训练好的 MHA 模型出发,把部分 KV 头合并、再少量微调,就能得到一个 GQA 模型——不需要从头训练。
正是因为「精度几乎不掉 + 推理明显更快」这个甜点,GQA 迅速被主流大模型采纳:
- Llama 2(70B)、Llama 3 全系列
- Mistral 7B / Mixtral
- Phi-3、Gemini 等
可以说,今天你叫得出名字的新一代大模型,几乎都用 GQA。它和「注意力机制」本身一样,已经成为现代 LLM 的标配组件。
小结
一句话浓缩:GQA 让多个 Query 头「拼团」共享同一套 K/V,在 MHA 的精度和 MQA 的速度之间找到了一个几乎不丢精度的甜点位置,从而大幅削减 KV 缓存、提升推理效率。
参考资料
- GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints - Ainslie et al., 2023 https://arxiv.org/abs/2305.13245
- Grouped-Query Attention (GQA) - Sebastian Raschka, LLMs-from-Scratch https://sebastianraschka.com/llms-from-scratch/ch04/04_gqa/
- What is Grouped Query Attention (GQA)? - IBM Think https://www.ibm.com/think/topics/grouped-query-attention