全部术语

GLOSSARY ENTRY

MHA、MQA 与 GQA

  • Multi-Head Attention
  • Multi-Query Attention
  • Grouped-Query Attention

三种 Query head 与 Key/Value head 参数共享方式;它们保持多 Query heads,但用不同的 KV head 数权衡模型质量、KV Cache 容量与 decode 读取带宽。

一个统一心智模型

三者都先把 hidden states XRB×T×HX\in\mathbb{R}^{B\times T\times H} 投影成 Query、Key、Value,再调用缩放点积注意力。差别不是“有没有多头”,而是 HqH_q 个 Query heads 共享多少组 K/V 投影

结构约束一个 KV head 服务的 Query headsK/V 逻辑 shape
MHAHkv=HqH_{kv}=H_q1[B,Hq,T,D][B,H_q,T,D]
GQA1<Hkv<Hq1<H_{kv}<H_qHqmodHkv=0H_q\bmod H_{kv}=0g=Hq/Hkvg=H_q/H_{kv}[B,Hkv,T,D][B,H_{kv},T,D]
MQAHkv=1H_{kv}=1全部 HqH_q[B,1,T,D][B,1,T,D]

统一投影公式可写成:

Q=XWQ[B,Hq,T,D]Q=XW^Q\rightarrow[B,H_q,T,D] K=XWK,  V=XWV[B,Hkv,T,D]K=XW^K,\;V=XW^V\rightarrow[B,H_{kv},T,D]

其中 WQRH×HqDW^Q\in\mathbb{R}^{H\times H_qD}WK,WVRH×HkvDW^K,W^V\in\mathbb{R}^{H\times H_{kv}D}。GQA 第 hh 个 Query head 使用 KV head h/g\lfloor h/g\rfloor;MQA 是 g=Hqg=H_q 的极端情况。

具体 shape:H=32,Hq=4,D=8H=32,H_q=4,D=8

B=1,T=6B=1,T=6

变体HkvH_{kv}QK / Vattention output单层每 token K+V 元素
MHA4[1,4,6,8][1,4,6,8][1,4,6,8][1,4,6,8][1,4,6,8][1,4,6,8]2×4×8=642\times4\times8=64
GQA2[1,4,6,8][1,4,6,8][1,2,6,8][1,2,6,8][1,4,6,8][1,4,6,8]3232
MQA1[1,4,6,8][1,4,6,8][1,1,6,8][1,1,6,8][1,4,6,8][1,4,6,8]1616

Query shape 和最终 per-head output shape 没变;减少的是 K/V 投影参数、激活和持久缓存。若 Hq=4,Hkv=2H_q=4,H_{kv}=2,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:

MKV=2LBTHkvDsM_{KV}=2LBTH_{kv}Ds

因此固定 Hq,D,L,B,T,sH_q,D,L,B,T,s 时,GQA/MQA 容量相对 MHA 按 Hkv/HqH_{kv}/H_q 缩小。decode 每个新 Query 仍覆盖相同历史长度,但需要读取的独立 K/V heads 更少,通常减轻 KV 容量与带宽压力。

收益并非免费:

  • KV sharing 降低 K/V 表示自由度,模型质量取决于架构、训练和任务;不能无条件说 MQA 与 MHA 等质。
  • prefill 的 attention 计算仍有 HqH_q 个 Query heads;减少 K/V 投影与 storage 不等于把全部 attention FLOPs 除以 group size。
  • backend 若先用 repeat_interleave 物理扩展 K/V,会失去一部分内存/带宽收益;需要原生 GQA kernel 或不复制的索引方式。
  • 张量并行要求 HqH_qHkvH_{kv} 与 TP degree 的分片兼容;若每个 rank 的 KV heads 太少,可能出现复制 KV heads、特殊通信或 kernel 限制。

与 MLA、KV 量化的边界

MLA 不是 HkvH_{kv} 更小的 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 读取量直接使用本页的 HkvH_{kv};KV 页反向给出 MHA/GQA/MQA 的具体 GiB 代入。
  • 后续 Attention 总览应解释 QKV/O projections、head concat、RoPE 与 residual block,避免把 SDPA 当成完整 MHA 模块。
  • 后续 MLA 页面应从本页的标准 head-shaped K/V 布局出发,再展示 latent 与 positional path 为何不同。

参考资料