Multi Head Attention
Multi-Head Attention
MHA 把一个 hidden vector 投影成多组 Q/K/V,让每个 head 在不同参数子空间里学习 token 关系,再把结果拼回原维度。它的价值不是“多算几遍同一个 attention”,而是把容量分配给多个可能互补的关系模式:局部搭配、长距离指代、句法结构或位置模式可以共存。
数据流
1 | hidden [B,T,D] |
标准同维实现的 Q/K/V 投影参数量约为 3D²,输出投影约为 D²(忽略 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 Attention 与 Tensor Shape,向上嵌入 Transformer Block,横向连接 GQA、MQA 与 KV 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,Tq 和 Tk 的来源不同;不能把 self-attention 的 causal mask 原样套上去。检查实现时把 self/cross、MHA/GQA、prefill/decode 各做一个 shape case,比只测方阵更能防止隐藏的转置错误。