一句话总结:Multi-Head Attention 把总表示宽度整理成多个较小的“头”,每个头都有自己的 Q、K、V 投影视角,最后再把所有结果拼接并通过
W_O混合回来。它像多双眼睛看同一句话,但每双眼睛具体学什么,是训练出来的,不是我们提前指定的。
11.1 为什么需要多个头?
11.1.1 单头 Attention 的局限
在上一章,我们学习了 Attention 的完整计算流程。但那是单头 Attention——只有一个 Query-Key 匹配矩阵和一份 Value 混合结果。
单头并不是“只能理解一种关系”,一张注意力图也可以很复杂。它的限制更准确地说是:同一层里只有一套 Q、K、V 投影视角,所有关系都要挤在同一个匹配与混合空间里。
但语言是多维度的!考虑这个句子:
“皮影艺人发现驴皮影人已经开裂,所以把它放回箱子。”
理解这个句子需要关注多种关系:
- 语法关系:“放回”的动作发出者是“皮影艺人”
- 指代关系:“它”指的是前面的“驴皮影人”
- 位置关系:“放回箱子”说明动作和位置的关系
- 因果关系:“所以”连接开裂和收起皮影人
多头机制给了模型多个并行子空间,允许这些关系用不同方式表示;但它不保证每个头都会自动变成一个干净的“语法头”或“指代头”。
11.1.2 解决方案:多个头
Multi-Head Attention 的核心思想:并行运行多个较窄的 Attention 头,让它们有机会学习不同的匹配与信息路由模式。
Head 1:可能关注语法结构(主谓宾)
Head 2:可能关注指代关系(代词和名词)
Head 3:可能关注位置信息(相邻的词)
Head 4:可能关注语义相似性(同义词)
...
最后把所有头的结果拼接,再通过 W_O 学习怎样混合这些视角。
11.1.3 一个形象的类比
想象你在分析一幅画:
- 一双眼睛:只能关注一个方面(比如颜色)
- 多双眼睛:同时关注颜色、形状、纹理、构图...
Multi-Head Attention 就像给模型配了多双眼睛,每双眼睛专注于不同的特征。
11.2 切分多头:把总投影宽度整理成 num_heads 份
11.2.1 维度切分
Multi-Head 的核心操作,是把合并投影的最后一维 reshape 成 head 轴和逐头维度。它不是把原始 embedding 生硬切开;每个切片已经经过可学习的 Q、K、V 投影。
以 K(Key)为例,假设:
- d_model = 512
- num_heads = 4
- 在这个经典等宽例子里,
d_head = d_k = d_v = d_model / num_heads = 512 / 4 = 128
切分过程:
合并 K: [batch_size, seq_length, num_heads × d_k]
= [4, 16, 512]
↓
reshape: [batch_size, seq_length, num_heads, d_k]
= [4, 16, 4, 128]
↓
转置: [batch_size, num_heads, seq_length, d_k]
= [4, 4, 16, 128]
11.2.2 为什么要转置?
转置是为了方便后续的矩阵运算。
转置后的形状 [batch, num_heads, seq_len, d_k] 可以理解为:
- 对于每个 batch 中的样本
- 有 num_heads 个独立的 Attention 头
- 每个头处理 seq_len 个位置
- 每个位置的 Query / Key 用
d_k维向量表示
这样可以用一次批量张量运算并行算完所有头。打分阶段各头使用自己的切片;训练时它们仍通过 W_O、残差路径和共同损失一起学习,并不是互不影响的四个小模型。
11.2.3 Q、K、V 都要切分
同样的切分操作应用于 Q、K、V:
Q: [4, 16, 512] → [4, 4, 16, 128] # d_k
K: [4, 16, 512] → [4, 4, 16, 128] # d_k
V: [4, 16, 512] → [4, 4, 16, 128] # d_v
现在我们有了 4 组逐头的 (Q, K, V),可以并行完成 4 份 Attention 计算。
11.2.4 两种切分实现方式
实际上,Multi-Head 的"切分"有两种等价的实现方式:
方式一:在线性变换层进行多头"切分"(左图)
- 概念上每个头有自己的
W_i^Q、W_i^K、W_i^V投影切片 - 每个头计算
Q_i = X @ W_i^Q - 这是概念理解的方式
方式二:在向量上砍 N 刀均分"切分"(右图)
- 用一个合并的
W_Q矩阵生成完整的 Q - 然后把 Q 沿着 d_model 维度切成 num_heads 份
- 这是实际实现的方式(更高效)
在没有额外跨 head 结构时,把各个小矩阵沿输出维拼起来,就得到同一个大矩阵,所以两种写法数学上等价。实际代码通常使用方式二,因为:
- 一次大矩阵乘法比多次小矩阵乘法更高效
- GPU 更擅长处理大的连续矩阵运算
11.3 并行计算多个头
11.3.1 每个头计算自己的分数矩阵
切分后,每个头独立执行 Attention 计算:
对于每个 head h = 1, 2, 3, 4:
scores_h = Q_h @ K_h^T [4, 16, 128] @ [4, 128, 16] = [4, 16, 16]
weights_h = softmax(scores_h / √d_k + M_h)
output_h = weights_h @ V_h [4, 16, 16] @ [4, 16, 128] = [4, 16, 128]
11.3.2 维度追踪
让我们详细追踪一次计算:
Q @ K^T:
Q: [4, 4, 16, 128]
batch heads seq d_k
K^T: [4, 4, 128, 16]
batch heads d_k seq
Q @ K^T: [4, 4, 16, 16]
batch heads seq seq
Softmax(Q @ K^T / √d_k + M) @ V:
Attention Weights: [4, 4, 16, 16]
batch heads seq seq
V: [4, 4, 16, 128]
batch heads seq d_v
Output: [4, 4, 16, 128]
batch heads seq d_v
11.3.3 图中的数据示例
图中展示了第 1 批次的第 1 个头:
Q @ K^T先得到一个[16, 16]的原始匹配分矩阵,再经过缩放、mask 和 Softmax 变成权重- 乘以 V [16, 128] 得到 [16, 128] 的输出
- 这只是 4 个头中的一个!
11.4 合并多头输出
11.4.1 Concatenate:拼接所有头
计算完所有头后,需要把它们合并回去:
各头输出: [4, 4, 16, 128]
batch heads seq d_v
↓
转置: [4, 16, 4, 128]
batch seq heads d_v
↓
合并: [4, 16, 512]
batch seq d_model
合并操作就是把最后两个维度"拼接"起来:
- 这个例子里 4 个头 × 128 维 = 512 维
11.4.2 W_O:输出投影
合并后还有一个输出投影:
A @ W_O
[4, 16, 512] @ [512, 512] = [4, 16, 512]
W_O(Output Projection) 是一个可学习的矩阵,作用是:
- 融合不同头的信息
- 把拼接的表示转换到统一的空间
- 让模型学习如何组合各头的输出
11.4.3 为什么需要 W_O?
拼接后的向量是"机械地"把各头的输出放在一起,没有任何交互。
W_O 允许模型学习:
- 哪些头的输出更重要
- 不同头之间如何配合
- 最终输出应该是什么样的
11.5 输出对比
11.5.1 A vs A @ W_O
图中对比了两种输出:
上方(A):合并后、W_O 之前的输出
- 形状:[16, 512]
- 是各头输出的简单拼接
下方(A @ W_O):经过 W_O 投影后的 Attention 分支输出
- 形状:[16, 512]
- 是经过混合和投影的结果
11.5.2 数值对比
虽然形状相同,但数值分布完全不同:
- A 的数值是各头独立计算的结果
A @ W_O的数值是所有头维度经过学习投影后的结果
它不会替代或改写 token embedding 参数表。这个 Attention 分支会和 residual stream 合并;之后是 LayerNorm 还是 FFN,要看该模型采用 pre-norm 还是 post-norm,不能只凭这一张图固定顺序。
11.6 Multi-Head Attention 完整流程
11.6.1 流程图
输入 X [batch, seq, d_model]
↓
生成 Q, K, V(通过 W_Q, W_K, W_V)
↓
整理多头 [batch, num_heads, seq, d_head]
↓
并行计算 Attention(每个头独立)
↓
合并多头 [batch, seq, d_model]
↓
输出投影(@ W_O)
↓
输出 [batch, seq, d_model]
11.6.2 PyTorch 代码
# 代码示例
import torch
import torch.nn as nn
import torch.nn.functional as F
class MultiHeadAttention(nn.Module):
def __init__(self, d_model, num_heads):
super().__init__()
if d_model % num_heads != 0:
raise ValueError("d_model must be divisible by num_heads")
self.d_model = d_model
self.num_heads = num_heads
self.head_dim = d_model // num_heads
# 四个可学习的投影;这个教学例子取 d_k = d_v = head_dim
self.W_Q = nn.Linear(d_model, d_model)
self.W_K = nn.Linear(d_model, d_model)
self.W_V = nn.Linear(d_model, d_model)
self.W_O = nn.Linear(d_model, d_model)
def forward(self, x, allowed_mask=None):
batch_size, seq_len, _ = x.shape
# 1. 生成 Q, K, V
Q = self.W_Q(x) # [batch, seq, d_model]
K = self.W_K(x)
V = self.W_V(x)
# 2. 切分多头
Q = Q.view(batch_size, seq_len, self.num_heads, self.head_dim)
K = K.view(batch_size, seq_len, self.num_heads, self.head_dim)
V = V.view(batch_size, seq_len, self.num_heads, self.head_dim)
# 转置: [batch, num_heads, seq, head_dim]
Q = Q.transpose(1, 2)
K = K.transpose(1, 2)
V = V.transpose(1, 2)
# 3. 计算 Attention
scores = torch.matmul(Q, K.transpose(-2, -1)) / (self.head_dim ** 0.5)
if allowed_mask is not None:
# allowed_mask 可广播到 [batch, heads, query_len, key_len]
scores = scores.masked_fill(~allowed_mask, float('-inf'))
attention_weights = F.softmax(scores, dim=-1)
attention_output = torch.matmul(attention_weights, V)
# 4. 合并多头
attention_output = attention_output.transpose(1, 2) # [batch, seq, heads, head_dim]
attention_output = attention_output.contiguous().view(
batch_size, seq_len, self.d_model
)
# 5. 输出投影
output = self.W_O(attention_output)
return output
11.7 关键数字
11.7.1 参数量计算
Multi-Head Attention 有四个权重矩阵:
| 矩阵 | 形状 | 参数量 |
|---|---|---|
| W_Q | [d_model, d_model] | d_model² |
| W_K | [d_model, d_model] | d_model² |
| W_V | [d_model, d_model] | d_model² |
| W_O | [d_model, d_model] | d_model² |
只数这四个投影的权重,总参数量是 4 × d_model²。如果线性层带 bias,还要再加 4 × d_model;不同架构也可能用 MQA/GQA 等非等宽投影,后面第 23 章再讲。
以 GPT-2 Small 为例(d_model = 768):
- 每层 Attention 投影权重 =
4 × 768² = 2,359,296 - 再算 GPT-2 的 QKV 与输出 bias,一共是
2,362,368个参数,约 236 万
11.7.2 常见配置
| 模型 | d_model | num_heads | 经典 MHA 的 d_head |
|---|---|---|---|
| GPT-2 Small | 768 | 12 | 64 |
| GPT-2 Medium | 1024 | 16 | 64 |
| GPT-2 Large | 1280 | 20 | 64 |
| GPT-3 | 12288 | 96 | 128 |
| LLaMA-7B | 4096 | 32 | 128 |
这张表里的逐头宽度都是 64 或 128,但这是这些模型的设计选择,不是一条普适定律。现代模型还会把 Query 头数和 Key/Value 头数解耦。
11.8 多头的可视化理解
11.8.1 每个头关注什么?
在对 BERT 的一项分析中,研究者观察到固定位置、分隔符、句法和指代等模式;同一层的一些头也会表现得很相似。这里列的是训练后观察到的可能模式,不是给每个 head 预先分配的岗位(Clark et al., 2019)。
| 头的类型 | 关注的模式 | 例子 |
|---|---|---|
| 位置头 | 固定相对位置 | 总是关注前一个词 |
| 语法头 | 主谓宾关系 | 动词关注主语 |
| 语义头 | 相似含义 | 同义词互相关注 |
| 指代头 | 代词解析 | it 关注它指代的名词 |
| 分隔头 | 句子边界 | 关注标点符号 |
11.8.2 注意力模式示例
下面仍然用这句皮影例子做直觉说明。它不是某个真实模型的测量结果:
Head 1(位置头):
“放回”较多关注附近的“它”
Head 2(语法头):
“放回”较多关注“皮影艺人”(动作发出者)
Head 3(语义头):
“开裂”和“驴皮影人”之间有较强联系
Head 4(指代头):
“它”较多关注“驴皮影人”
11.8.3 头的冗余性
实际上,并非所有头都同等重要。Michel 等人在特定机器翻译模型和 BERT 任务上发现,测试时可以移除相当一部分头而不显著损伤当时评估的指标,有些层甚至能只保留一个头(Michel et al., 2019)。这说明冗余确实存在,但不能推出“任何模型的多数头都没用”;重要性取决于模型、层、任务和数据。
11.9 Multi-Head vs Single-Head
11.9.1 计算量对比
假设 d_model = 512、num_heads = 8、经典等宽 MHA 的 d_head = 64:
单头(宽度 512):
- Q @ K^T:[seq, 512] @ [512, seq] → O(seq² × 512)
多头(每头宽度 64):
- 8 个头,每个:[seq, 64] @ [64, seq] → O(seq² × 64)
- 总计:8 × O(seq² × 64) = O(seq² × 512)
只看 QK^T 这部分,乘加总量相同;weights @ V 也有同样结论。固定总投影宽度时,QKV 和 W_O 的主阶计算量也不因“切成几头”而改变。不过真实运行时间还受 kernel、内存布局和硬件影响,不能只用大 O 断言完全一样。
11.9.2 效果对比
多头的价值不在于凭空增加总宽度,而在于让同一层拥有多套独立的 Query-Key 投影和 Value 路由,再由 W_O 合并。它可以表达单个注意力加权平均难以直接表达的组合,但并不等同于传统“集成学习”。
11.9.3 为什么不用更多头?
在这个固定 d_model、等宽的经典设置里,增加头数意味着减小逐头宽度:
d_head = d_model / num_heads
如果 d_head 太小:
- 每个头的表达能力下降
- 可能无法捕捉复杂的模式
64 和 128 是上表几种经典模型采用的选择,不是通用“最佳点”。头数与逐头宽度仍然要靠架构设计和实验验证。
11.10 本章总结
11.10.1 核心概念
| 概念 | 解释 |
|---|---|
| Multi-Head | 把 Attention 分成多个头并行计算 |
| num_heads | 头的数量 |
| d_head | 这个等宽例子里每个头的维度 = d_model / num_heads |
| 切分 | 把 d_model 维度分成 num_heads 份 |
| 合并 | 把所有头的输出拼接回 d_model |
| W_O | 输出投影矩阵,混合各头维度 |
11.10.2 维度变化
输入: [batch, seq, d_model]
↓
切分: [batch, num_heads, seq, d_head]
↓
Attention: [batch, num_heads, seq, d_head]
↓
合并: [batch, seq, d_model]
↓
W_O 投影: [batch, seq, d_model]
11.10.3 核心认知
Multi-Head Attention 给同一层多套可学习的观察视角:逐头计算匹配与 Value 混合,拼接后再由
W_O重新组合。不同头可能出现不同模式,也可能彼此冗余;“多双眼睛”是理解结构的类比,不是保证每个头都有固定岗位。
本章交付物
学完这一章,你应该能够:
- 解释为什么需要 Multi-Head Attention
- 说出经典等宽 MHA 中
d_head = d_model / num_heads的关系 - 理解切分和合并的维度变化
- 知道
W_O矩阵的作用(混合各头维度) - 解释不同头可能学习到的不同模式
Part 3 阶段小结
到这里,我们已经走完 Multi-Head 的结构;Part 3 还剩第 12 章,要把“Attention 输出到底更新了什么”说清楚。
让我们回顾这三章学到的内容:
| 章节 | 主题 | 核心内容 |
|---|---|---|
| 第 8 章 | 线性变换 | 分清线性映射、点积、余弦与投影 |
| 第 9 章 | Attention 几何逻辑 | 用学出来的 Q-K 点积计算匹配分 |
| 第 10 章 | QKV | Q 查询、K 键、V 值的完整计算流程 |
| 第 11 章 | Multi-Head | 多个头从不同角度理解句子 |
Attention 公式(完整版):
其中:
下一章预告
下一章还是 Part 3。我们会把 Attention 的输出、residual stream、embedding 参数表和反向传播分开看:前向传播到底生成了什么中间表示,训练时又究竟更新哪些参数。