多头注意力是 Transformer 的核心机制,它通过并行执行多个独立的注意力头(attention heads)来捕捉输入的不同表示子空间。以下是完整的数学推导(所有公式用 $ 包裹):
输入序列矩阵:
X∈Rn×dmodel
其中:
- n:序列长度
- dmodel:模型维度(如 512)
对每个头 i∈{1,2,…,h},使用独立的权重矩阵进行投影:
QiKiVi=X⋅WiQ,=X⋅WiK,=X⋅WiV,WiQ∈Rdmodel×dkWiK∈Rdmodel×dkWiV∈Rdmodel×dv
参数说明:
- h:头数(如 8)
- dk=dv=hdmodel:每个头的维度(如 64)
- 投影后维度:Qi,Ki∈Rn×dk,Vi∈Rn×dv
对每个头 i 独立计算缩放点积注意力:
headi=Attention(Qi,Ki,Vi)=softmax(dkQiKiT)Vi
其中:
- dkQiKiT∈Rn×n:缩放后的注意力分数矩阵
- 输出维度:headi∈Rn×dv
将所有头的输出在特征维度拼接:
MultiHead(Q,K,V)=Concat(head1,head2,…,headh)
拼接后维度:Concat(⋯)∈Rn×(h⋅dv)=Rn×dmodel
将拼接结果通过可学习权重矩阵 WO 投影回模型维度:
Output=Concat(head1,…,headh)⋅WO,WO∈Rdmodel×dmodel
最终输出维度:Rn×dmodel
MultiHead(X)=Concat(head1,…,headh)⋅WOwhereheadi=softmax(dk(XWiQ)(XWiK)T)(XWiV)with⎩⎨⎧dk=dv=hdmodelWiQ,WiK∈Rdmodel×dkWiV∈Rdmodel×dvWO∈Rdmodel×dmodel
输入 X: [n, d_model]
│
├─ Head 1 ──┐
│ Q1 = X·W1^Q [n, d_k] │
│ K1 = X·W1^K [n, d_k] → Attn1 [n, n] → head1 [n, d_v]
│ V1 = X·W1^V [n, d_v] │ │
├─ Head 2 ──┤ │
│ ... │ ├→ Concat [n, h·d_v] = [n, d_model]
├─ Head h ──┤ │
│ │ │
└──────────┘ │
│
Output = Concat·W^O [n, d_model] ←────────┘
子空间分解
每个头学习不同的投影:
WiQ,WiK,WiV 相互独立
使模型关注不同方面的信息(如语法/语义/位置)。
计算效率
虽然头数 h 增加,但单头维度 dk 减小:
单头计算量=O(n2dk)=O(n2hdmodel)
总计算量 O(n2dmodel) 与单头相同。
表达能力增强
输出是多个子空间的非线性组合:
Output=f(i=1∑hheadi⋅WiO)
(WO 隐含分解为子矩阵 WiO)
设 dmodel=4, h=2, 则 dk=dv=2
输入 X=[1−0.50.51−10.32−2]
头1计算:
W1Q=0.1−0.21.00.50.40.3−0.50.2, Q1=XW1Q=[0.85−0.29−0.050.74]
(类似计算 K1, V1 后求 head1)
头2计算:
W2Q=−0.31.1−0.40.70.20.60.8−0.1, Q2=XW2Q=[2.15−1.320.651.02]
拼接与输出:
Concat=[head1(1)head1(2)head2(1)head2(2)]
WO=0.50.1−0.30.9−0.20.80.7−0.51.1−0.40.20.30.30.61.0−0.8, Output=Concat⋅WO
| 特性 | 单头注意力 | 多头注意力 |
|---|
| 参数量 | 3dmodel2 | 3dmodel2 (相同) |
| 计算复杂度 | O(n2dmodel) | O(n2dmodel) |
| 表达能力 | 单一表示空间 | h 个正交子空间 |
| 并行性 | 低 | 高(头间完全并行) |
- 多样性:不同头关注不同模式(如局部依赖 vs 全局依赖)
headi 学习独立的关系表示
- 维度分解:将高维空间分解为低维子空间,降低学习难度
Rdmodel→⊕i=1hRdk
- 残差连接:多头输出通过残差连接传递到前馈网络,保留原始信息:
LayerOutput=LayerNorm(X+MultiHead(X))
多头注意力通过这种"分治策略",显著提升了模型捕捉复杂依赖关系的能力,成为 Transformer 架构的核心创新。� Transformer 架构的核心创新。