一个统一心智模型
三者都先把 hidden states 投影成 Query、Key、Value,再调用缩放点积注意力。差别不是“有没有多头”,而是 个 Query heads 共享多少组 K/V 投影:
| 结构 | 约束 | 一个 KV head 服务的 Query heads | K/V 逻辑 shape |
|---|---|---|---|
| MHA | 1 | ||
| GQA | 且 | ||
| MQA | 全部 |
统一投影公式可写成:
其中 ,。GQA 第 个 Query head 使用 KV head ;MQA 是 的极端情况。
具体 shape:
取 :
| 变体 | Q | K / V | attention output | 单层每 token K+V 元素 | |
|---|---|---|---|---|---|
| MHA | 4 | 各 | |||
| GQA | 2 | 各 | |||
| MQA | 1 | 各 |
Query shape 和最终 per-head output shape 没变;减少的是 K/V 投影参数、激活和持久缓存。若 ,head 映射为:
Q head: 0 1 2 3
KV head: 0 0 1 1
这个映射是逻辑共享,不要求内存中把每个 KV head 复制两份。教学 reference 常用 repeat_interleave 对齐 shape,生产 kernel 会直接按 group 索引读取。
可运行的参数与容量实验
实验问题:三种结构的 Q/K/V shape、投影参数量和缓存元素如何变化?GQA/MQA 的共享读取是否与显式复制 KV heads 数值一致?
import platform
import torch
import torch.nn.functional as F
torch.manual_seed(23)
DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
DTYPE = torch.float32
def run_variant(x, q_heads, kv_heads, head_dim):
hidden = x.shape[-1]; scale = hidden**-0.5
wq = torch.randn(hidden, q_heads * head_dim, device=DEVICE) * scale
wk = torch.randn(hidden, kv_heads * head_dim, device=DEVICE) * scale
wv = torch.randn(hidden, kv_heads * head_dim, device=DEVICE) * scale
batch, tokens, _ = x.shape
q = (x @ wq).view(batch, tokens, q_heads, head_dim).transpose(1, 2)
k = (x @ wk).view(batch, tokens, kv_heads, head_dim).transpose(1, 2)
v = (x @ wv).view(batch, tokens, kv_heads, head_dim).transpose(1, 2)
output = F.scaled_dot_product_attention(
q, k, v, is_causal=True, enable_gqa=(q_heads != kv_heads))
repeat = q_heads // kv_heads
explicit = F.scaled_dot_product_attention(
q, k.repeat_interleave(repeat, 1), v.repeat_interleave(repeat, 1),
is_causal=True)
torch.testing.assert_close(output, explicit, rtol=1e-5, atol=1e-6)
assert output.shape == (batch, q_heads, tokens, head_dim)
return q, k, v, output, (wq.numel(), wk.numel() + wv.numel())
def main():
b, t, hidden, hq, d = 1, 6, 32, 4, 8
x = torch.randn(b, t, hidden, device=DEVICE, dtype=DTYPE)
print(f"python={platform.python_version()} torch={torch.__version__}")
print(f"device={DEVICE} dtype={DTYPE} input={tuple(x.shape)} Hq={hq} D={d}")
for name, hkv in (("MHA", 4), ("GQA", 2), ("MQA", 1)):
q, k, v, output, params = run_variant(x, hq, hkv, d)
print(f"{name}: Hkv={hkv} Q={tuple(q.shape)} K=V={tuple(k.shape)} "
f"O={tuple(output.shape)} Q_params={params[0]} "
f"KV_params={params[1]} KV_elements={k.numel()+v.numel()}")
if __name__ == "__main__":
with torch.inference_mode(): main()
完整文件位于 examples/head_sharing_reference.py。2026-08-13 实测三种变体的 KV_elements 分别为 384、192、96,W^K+W^V 参数量分别为 2048、1024、512;每个变体都通过“GQA API 输出等于显式 KV head 展开”的断言。
这个实验没有证明 MHA、GQA、MQA 三个随机模型输出相等——它们权重 shape 不同,本就不是相同函数。它只证明给定某个共享参数化后,head 映射的两种执行表示等价。
对 KV Cache、计算与质量的影响
KV Cache 的标准 payload:
因此固定 时,GQA/MQA 容量相对 MHA 按 缩小。decode 每个新 Query 仍覆盖相同历史长度,但需要读取的独立 K/V heads 更少,通常减轻 KV 容量与带宽压力。
收益并非免费:
- KV sharing 降低 K/V 表示自由度,模型质量取决于架构、训练和任务;不能无条件说 MQA 与 MHA 等质。
- prefill 的 attention 计算仍有 个 Query heads;减少 K/V 投影与 storage 不等于把全部 attention FLOPs 除以 group size。
- backend 若先用
repeat_interleave物理扩展 K/V,会失去一部分内存/带宽收益;需要原生 GQA kernel 或不复制的索引方式。 - 张量并行要求 、 与 TP degree 的分片兼容;若每个 rank 的 KV heads 太少,可能出现复制 KV heads、特殊通信或 kernel 限制。
与 MLA、KV 量化的边界
MLA 不是 更小的 GQA。GQA/MQA 仍缓存标准 head-shaped K/V;MLA 缓存低维 latent,并另处理解耦位置编码路径,decode 中还涉及矩阵吸收。KV 量化则通常保留 MHA/GQA/MQA 的逻辑 head 结构,改变元素表示、scale metadata 与消费 kernel。
上下游关系
- 缩放点积注意力 定义每个 Query head 如何对其对应 K/V 序列计算权重;本页定义“对应哪个 KV head”。
- KV Cache 的容量与 decode 读取量直接使用本页的 ;KV 页反向给出 MHA/GQA/MQA 的具体 GiB 代入。
- 后续 Attention 总览应解释 QKV/O projections、head concat、RoPE 与 residual block,避免把 SDPA 当成完整 MHA 模块。
- 后续 MLA 页面应从本页的标准 head-shaped K/V 布局出发,再展示 latent 与 positional path 为何不同。
参考资料
- Vaswani et al., Attention Is All You Need — Multi-Head Attention。
- Shazeer, Fast Transformer Decoding: One Write-Head is All You Need — Multi-Query Attention。
- Ainslie et al., GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints — Grouped-Query Attention 与 uptraining。
- PyTorch — scaled_dot_product_attention —
enable_gqa的 API 约束与 backend 说明。