一句话总结: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 维度切分

合并 QKV 投影整理成多个 head,以及等价的逐头投影视角

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 两种切分实现方式

逐头小投影与一次合并大投影加 reshape 的等价关系

实际上,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 每个头计算自己的分数矩阵

四个 head 并行执行缩放、Mask、Softmax 和 Value 混合

切分后,每个头独立执行 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:拼接所有头

逐头输出经过转置、拼接、W_O 和 residual 路径

计算完所有头后,需要把它们合并回去:

各头输出: [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) 是一个可学习的矩阵,作用是:

  1. 融合不同头的信息
  2. 把拼接的表示转换到统一的空间
  3. 让模型学习如何组合各头的输出

11.4.3 为什么需要 W_O?

拼接后的向量是"机械地"把各头的输出放在一起,没有任何交互。

W_O 允许模型学习:

  • 哪些头的输出更重要
  • 不同头之间如何配合
  • 最终输出应该是什么样的

11.5 输出对比

11.5.1 A vs A @ W_O

拼接后的 A 与经过 W_O 混合后的 Attention 分支输出

图中对比了两种输出:

上方(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_modelnum_heads经典 MHA 的 d_head
GPT-2 Small7681264
GPT-2 Medium10241664
GPT-2 Large12802064
GPT-31228896128
LLaMA-7B409632128

这张表里的逐头宽度都是 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 章QKVQ 查询、K 键、V 值的完整计算流程
第 11 章Multi-Head多个头从不同角度理解句子

Attention 公式(完整版):

MultiHead(Q,K,V)=Concat(head1,...,headh)WO\text{MultiHead}(Q, K, V) = \text{Concat}(\text{head}_1, ..., \text{head}_h) W^O

其中:

headi=Attention(QWiQ,KWiK,VWiV)\text{head}_i = \text{Attention}(QW_i^Q, KW_i^K, VW_i^V)
Attention(Q,K,V)=softmax(QKTdk+M)V\text{Attention}(Q, K, V) = \text{softmax}\left(\frac{QK^T}{\sqrt{d_k}} + M\right)V

下一章预告

下一章还是 Part 3。我们会把 Attention 的输出、residual stream、embedding 参数表和反向传播分开看:前向传播到底生成了什么中间表示,训练时又究竟更新哪些参数。

引用本文 / Cite
Zhang, Wayland (2026). 第 11 章:Multi-Head Attention - 多视角理解. In Transformer 架构:从直觉到实现. https://waylandz.com/llm-transformer-book/%E7%AC%AC11%E7%AB%A0-Multi-Head-Attention-%E5%A4%9A%E8%A7%86%E8%A7%92%E7%90%86%E8%A7%A3/
@incollection{zhang2026transformer_11_-Multi-Head-Attention-,
  author = {Zhang, Wayland},
  title = {第 11 章:Multi-Head Attention - 多视角理解},
  booktitle = {Transformer 架构:从直觉到实现},
  year = {2026},
  url = {https://waylandz.com/llm-transformer-book/%E7%AC%AC11%E7%AB%A0-Multi-Head-Attention-%E5%A4%9A%E8%A7%86%E8%A7%92%E7%90%86%E8%A7%A3/}
}