Multi-Head Attention

MHA 把一个 hidden vector 投影成多组 Q/K/V,让每个 head 在不同参数子空间里学习 token 关系,再把结果拼回原维度。它的价值不是“多算几遍同一个 attention”,而是把容量分配给多个可能互补的关系模式:局部搭配、长距离指代、句法结构或位置模式可以共存。

数据流

1
2
3
4
5
6
hidden [B,T,D]
→ Q/K/V linear projection [B,T,D]
→ split heads [B,H,T,Dh], Dh=D/H
→ SDPA:scores [B,H,Tq,Tk]
→ merge [B,Tq,H,Dh] → [B,Tq,D]
→ output projection [B,Tq,D]

标准同维实现的 Q/K/V 投影参数量约为 3D²,输出投影约为 (忽略 bias)。增加 head 数并不自动增加总 hidden 容量,因为通常 H×Dh=D;它改变的是参数如何被分到子空间。

为什么分头有效

单一 attention 必须用一套投影同时表达所有关系,容易在不同模式间竞争。多头让每个 head 使用自己的 Q/K/V 投影,再在输出投影处混合信息。因果链是:独立子空间 → 不同相似度与可见模式 → 多种上下文聚合 → 线性混合回统一表示。这里的“不同”是训练学出来的,不保证每个 head 都能被人解释成一个固定语义。

与 GQA/MQA 的关键差异

机制 Query heads KV heads Decode 的 KV 读取量 主要取舍
MHA Hq Hq 最大 表达容量和缓存成本都高
GQA Hq 1 < Hkv < Hq 中间 共享部分 KV,折衷质量与带宽
MQA Hq 1 最小 共享最强,可能损失表达能力

这不是只影响显存的实现细节:K/V 共享会改变每个 Q head 能读取的内容表示,因此必须在目标任务上验证质量。prefill 可能更受矩阵计算影响,decode 则更直接暴露 KV Cache 的内存带宽差异。

实现检查

  • 断言 D % H == 0,并明确 transpose 后的轴顺序。
  • Tq=Tk=T 是训练常见形状;单 token decode 时 Tq=1,历史 K/V 来自 KV Cache
  • GQA/MQA 的逻辑 head 扩展不等于物理复制缓存;查看 kernel 和 cache layout 才能判断实际带宽。
  • 用 attention 权重可视化不能证明 head 有稳定语义,最好结合删 head、任务回归和长上下文测试。

它向下依赖 Scaled Dot-Product AttentionTensor Shape,向上嵌入 Transformer Block,横向连接 GQAMQAKV Cache

参考资料

“每个 head 学一种语义”不是设计契约

多头提供的是多个参数化子空间,不保证一个 head 永远对应主语、括号或某种语言关系。训练中 head 可能冗余、协作或随层深变化;attention heatmap 只能说明某次输入上的权重,不能直接当作因果解释。更可靠的分析是屏蔽某个 head 后测任务变化,再结合不同长度、语言和随机种子重复,判断它是必要信息通路还是可替代容量。

MHA 的成本有两张账

H×Dh=D 固定时,增加 head 数通常不增加投影参数总量,却会改变 kernel 的并行粒度、每个 head 的 Dh 和 KV cache 的布局。MHA 的每个 KV head 都保留一份历史,所以 decode 的带宽成本随 Hkv 增长;GQA/MQA 通过共享 K/V 降低这项成本,却改变了可读取的表示。因而“参数量相同”不等于“推理成本相同”。

跨 attention 时,Q 来自 decoder,K/V 来自 encoder,TqTk 的来源不同;不能把 self-attention 的 causal mask 原样套上去。检查实现时把 self/cross、MHA/GQA、prefill/decode 各做一个 shape case,比只测方阵更能防止隐藏的转置错误。