Zer0e's Blog

去TM的AI其二:什么是多头注意力Multi-Head Attention

字数统计: 2.9k阅读时长: 12 min
2026/07/05 Share

多头注意力机制(Multi-Head Attention)是 Transformer 架构的核心组件,由 Vaswani 等人在 2017 年的论文《Attention Is All You Need》中首次提出。它是现代大语言模型(如 GPT、BERT 等)能够理解和生成高质量文本的关键所在。

为什么要用”多头”?

在理解多头注意力之前,我们先来思考一个问题:

想象你在阅读一句话:”银行在河的旁边”。这里的”银行”是什么意思?

  • 如果你关注”河”这个词,你会理解为”河岸”(river bank)
  • 如果你关注”钱”或”存款”这些词,你会理解为”金融机构”(financial bank)

人类在阅读时,会同时从多个角度理解词语之间的关系。多头注意力机制正是模仿了这种能力——让模型能够同时关注不同位置的不同表示子空间

单头注意力的局限性

单头注意力(Single-Head Attention)虽然能捕捉序列中元素之间的依赖关系,但它只能从一个角度(一个表示子空间)来建模这些关系。这就像只用一种方式理解世界,容易丢失重要信息。

多头注意力的核心思想

多头注意力机制的核心思想是:将模型的注意力分成多个”头”,每个头独立学习不同的注意力模式,最后将这些头的输出合并。

flowchart TD
    A[输入序列] --> B[分割成多个头]
    B --> C1[Head 1: 关注语法关系]
    B --> C2[Head 2: 关注语义关系]
    B --> C3[Head 3: 关注位置关系]
    B --> C4[Head N: 关注其他模式]
    C1 --> D[Concat 拼接]
    C2 --> D
    C3 --> D
    C4 --> D
    D --> E[线性变换]
    E --> F[输出]
    
    style B fill:#e1f5ff
    style D fill:#fff3e1
    style F fill:#e8f5e9

多头注意力的数学原理

1. 缩放点积注意力(Scaled Dot-Product Attention)

在理解多头注意力之前,我们需要先了解它的基础——缩放点积注意力。

给定查询向量 (Q)(Query)、键向量 (K)(Key)和值向量 (V)(Value),注意力计算公式为:
$$
[\text{Attention}(Q, K, V) = \text{softmax}\left(\frac{QK^T}{\sqrt{d_k}}\right)V]
$$

其中:

  • $$(QK^T) 计算查询和键之间的相似度$$
  • $$(\sqrt{d_k}) 是缩放因子,防止点积过大导致 softmax 梯度消失$$
  • $$(\text{softmax}) 将分数转换为概率分布$$
  • $$最后与 (V) 相乘得到加权求和的结果$$

2. 多头注意力的完整公式

多头注意力将 (Q)、(K)、(V) 分别通过 (h) 个不同的线性变换投影到低维空间:
$$
[\text{MultiHead}(Q, K, V) = \text{Concat}(\text{head}_1, \text{head}_2, …, \text{head}_h)W^O]
$$
其中第 (i) 个头的计算为:
$$
[\text{head}_i = \text{Attention}(QW_i^Q, KW_i^K, VW_i^V)]
$$
其中:

  • $$(W_i^Q \in \mathbb{R}^{d_{\text{model}} \times d_k})$$
  • $$(W_i^K \in \mathbb{R}^{d_{\text{model}} \times d_k})$$
  • $$(W_i^V \in \mathbb{R}^{d_{\text{model}} \times d_v})$$
  • $$(W^O \in \mathbb{R}^{hd_v \times d_{\text{model}}})$$

$$在原始论文中,(d_{\text{model}} = 512),(h = 8),因此每个头的维度 (d_k = d_v = d_{\text{model}}/h = 64)。$$

多头注意力的工作流程

让我们通过一个详细的流程图来理解多头注意力的完整计算过程:

flowchart TD
    subgraph Input["输入阶段"]
        A1[Input Embeddings
维度: seq_len × d_model] end subgraph Projection["线性投影层"] B1[Q = X·Wᵠ] B2[K = X·Wᴷ] B3[V = X·Wⱽ] end subgraph Split["分割为多个头"] S1[每个头维度: d_k = d_model / h] C1[Q₁, K₁, V₁] C2[Q₂, K₂, V₂] C3["..."] C4[Qₕ, Kₕ, Vₕ] end subgraph Attention["并行计算注意力"] D1["Head 1
softmax(Q₁K₁ᵀ/√dₖ)·V₁"] D2["Head 2
softmax(Q₂K₂ᵀ/√dₖ)·V₂"] D3["..."] D4["Head h
softmax(QₕKₕᵀ/√dₖ)·Vₕ"] end subgraph Merge["合并阶段"] E1[Concat
拼接所有头的输出] E2[Output = Concat·Wᴼ
线性变换] end Input --> Projection B1 --> C1 B2 --> C2 B3 --> C4 C1 --> D1 C2 --> D2 C4 --> D4 D1 --> E1 D2 --> E1 D4 --> E1 E1 --> E2 style Input fill:#e3f2fd style Projection fill:#f3e5f5 style Split fill:#fff3e0 style Attention fill:#e8f5e9 style Merge fill:#fce4ec

详细步骤解析

步骤 1:线性投影(Linear Projection)

输入序列首先通过三个不同的线性变换,生成 (Q)、(K)、(V) 三个矩阵:

flowchart LR
    A[输入 X
seq_len × d_model] --> B[线性层 Wᵠ] A --> C[线性层 Wᴷ] A --> D[线性层 Wⱽ] B --> E[Q 矩阵
seq_len × d_model] C --> F[K 矩阵
seq_len × d_model] D --> G[V 矩阵
seq_len × d_model] style A fill:#bbdefb style E fill:#c8e6c9 style F fill:#fff9c4 style G fill:#f8bbd0
  • Query(查询):表示”我在寻找什么”
  • Key(键):表示”我提供什么”
  • Value(值):表示”我的实际内容”

步骤 2:分割为多个头(Split into Heads)

$$将 (Q)、(K)、(V) 分别沿着特征维度切分成 (h) 个头。$$
$$假设 (d_{\text{model}} = 512),(h = 8),则每个头的维度为 (512 / 8 = 64)。$$

flowchart TD
    A[Q: 512 维] --> B[切分成 8 个头]
    B --> C1[Q₁: 64 维]
    B --> C2[Q₂: 64 维]
    B --> C3["..."]
    B --> C4[Q₈: 64 维]
    
    style A fill:#ffcc80
    style B fill:#a5d6a7
    style C1 fill:#90caf9
    style C2 fill:#90caf9
    style C4 fill:#90caf9

步骤 3:并行计算注意力(Parallel Attention Computation)

每个头独立计算注意力,这就是”多头”的核心——并行处理不同的表示子空间

flowchart LR
    subgraph Head1["Head 1"]
        A1[Q₁, K₁, V₁] --> B1["Attention
计算注意力分数"] B1 --> C1[Output₁] end subgraph Head2["Head 2"] A2[Q₂, K₂, V₂] --> B2["Attention
计算注意力分数"] B2 --> C2[Output₂] end subgraph HeadH["Head h"] A3[Qₕ, Kₕ, Vₕ] --> B3["Attention
计算注意力分数"] B3 --> C3[Outputₕ] end style Head1 fill:#e8f5e9 style Head2 fill:#fff3e0 style HeadH fill:#f3e5f5

步骤 4:拼接与输出变换(Concat & Output Transformation)

将所有头的输出拼接起来,再通过一个线性变换得到最终输出:

flowchart LR
    A[Output₁] --> B[Concat]
    C[Output₂] --> B
    D["..."] --> B
    E[Outputₕ] --> B
    B --> F[拼接后的向量
seq_len × h·dᵥ] F --> G[线性层 Wᴼ] G --> H[最终输出
seq_len × d_model] style B fill:#ffe0b2 style F fill:#c5e1a5 style H fill:#80deea

一个具体的计算示例

假设我们有一个简单的句子:”The cat sat on the mat”,共 6 个词。

参数设置

  • $$序列长度(seq_len)= 6$$
  • $$模型维度((d_{\text{model}}))= 512$$
  • $$头的数量((h))= 8$$
  • $$每个头的维度((d_k = d_v))= 64$$

计算过程

flowchart TD
    A[输入: 6 × 512 矩阵] --> B[生成 Q, K, V]
    B --> C[每个头处理 6 × 64 矩阵]
    C --> D["计算注意力分数: 6 × 6 矩阵"]
    D --> E["加权求和得到 6 × 64 输出"]
    E --> F[8 个头拼接: 6 × 512]
    F --> G[输出变换: 6 × 512]
    
    style A fill:#bbdefb
    style D fill:#fff59d
    style G fill:#a5d6a7

每个头会学习到不同的注意力模式。例如:

  • Head 1:可能关注主谓关系(cat → sat)
  • Head 2:可能关注冠词与名词的关系(the → cat, the → mat)
  • Head 3:可能关注位置关系(相邻的词)
  • 其他头:可能学习到更复杂的语义模式

多头注意力的优势

1. 多表示子空间学习

每个头可以专注于不同的特征模式,类似于 CNN 中不同的卷积核:

flowchart TD
    A[多头注意力] --> B1[语法模式]
    A --> B2[语义模式]
    A --> B3[位置模式]
    A --> B4[长程依赖]
    A --> B5[局部模式]
    
    B1 --> C[融合多种信息]
    B2 --> C
    B3 --> C
    B4 --> C
    B5 --> C
    C --> D[更强的表达能力]
    
    style A fill:#ff8a65
    style C fill:#4db6ac
    style D fill:#7986cb

2. 增强模型的表达能力

通过多个注意力头,模型可以:

  • 同时关注不同位置的信息
  • 捕捉不同层次的语义关系
  • 避免单一注意力头的信息瓶颈

3. 并行计算效率高

虽然计算多个头,但由于每个头的维度降低,且可以并行计算,整体效率反而更高。

自注意力 vs 交叉注意力 vs 因果注意力

在 Transformer 中,多头注意力有三种主要应用形式:

flowchart TD
    A[多头注意力机制] --> B1[自注意力 Self-Attention]
    A --> B2[交叉注意力 Cross-Attention]
    A --> B3[因果注意力 Causal Attention]
    
    B1 --> C1["Q, K, V 来自同一序列
编码器内部使用"] B2 --> C2["Q 来自解码器, K, V 来自编码器
解码器-编码器连接"] B3 --> C3["使用掩码防止看到未来信息
解码器内部使用"] style B1 fill:#e3f2fd style B2 fill:#f3e5f5 style B3 fill:#e8f5e9

代码实现(PyTorch)

下面是一个简化版的多头注意力实现:

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
import torch
import torch.nn as nn
import math

class MultiHeadAttention(nn.Module):
def __init__(self, d_model=512, num_heads=8):
super(MultiHeadAttention, self).__init__()
self.num_heads = num_heads
self.d_model = d_model
self.d_k = d_model // num_heads

# 线性变换层
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 scaled_dot_product_attention(self, Q, K, V, mask=None):
# 计算注意力分数
scores = torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(self.d_k)

# 应用掩码(如果提供)
if mask is not None:
scores = scores.masked_fill(mask == 0, -1e9)

# softmax 归一化
attention_weights = torch.softmax(scores, dim=-1)

# 加权求和
output = torch.matmul(attention_weights, V)
return output, attention_weights

def forward(self, x, mask=None):
batch_size, seq_len, _ = x.size()

# 线性投影并分割为多个头
Q = self.W_Q(x).view(batch_size, seq_len, self.num_heads, self.d_k).transpose(1, 2)
K = self.W_K(x).view(batch_size, seq_len, self.num_heads, self.d_k).transpose(1, 2)
V = self.W_V(x).view(batch_size, seq_len, self.num_heads, self.d_k).transpose(1, 2)

# 并行计算每个头的注意力
attention_output, attention_weights = self.scaled_dot_product_attention(Q, K, V, mask)

# 拼接所有头的输出
attention_output = attention_output.transpose(1, 2).contiguous().view(
batch_size, seq_len, self.d_model
)

# 输出线性变换
output = self.W_O(attention_output)

return output, attention_weights

实际应用场景

多头注意力在以下场景中发挥着关键作用:

1. 机器翻译

flowchart LR
    A[源语言:
I love AI] --> B[编码器
多头自注意力] B --> C[交叉注意力
连接编码器-解码器] C --> D[解码器
多头因果注意力] D --> E[目标语言:
我爱人工智能] style B fill:#ffcc80 style C fill:#ce93d8 style D fill:#81c784

2. 文本摘要

多头注意力可以:

  • 识别关键信息(某些头关注重要句子)
  • 捕捉上下文关系(某些头关注代词指代)
  • 理解文档结构(某些头关注段落关系)

3. 问答系统

flowchart TD
    A[问题: 谁发明了电话?] --> B[多头注意力匹配]
    C[文档: Bell发明了电话...] --> B
    B --> D1[Head 1: 实体匹配]
    B --> D2[Head 2: 关系识别]
    B --> D3[Head 3: 位置对齐]
    D1 --> E[答案: Bell]
    D2 --> E
    D3 --> E

可视化:注意力模式

研究表明,不同的注意力头确实会学习到不同的模式。以下是一些典型的注意力模式:

模式 1:关注相邻词(Local Attention)

某些头主要关注当前位置附近的词,类似于 n-gram 模型。

模式 2:关注句法关系(Syntactic Attention)

某些头会学习到句法依赖关系,例如:

  • 主语 ↔ 谓语
  • 冠词 ↔ 名词
  • 代词 ↔ 先行词

模式 3:关注全局信息(Global Attention)

某些头会均匀地关注整个序列,用于捕捉全局上下文。

总结

多头注意力机制是 Transformer 架构的核心创新,它通过以下三个关键步骤实现强大的特征提取能力:

  1. 投影:将输入映射到多个表示子空间
  2. 并行注意力:每个头独立计算注意力,捕捉不同的模式
  3. 合并:将所有头的输出拼接并变换,融合多视角信息
flowchart TD
    A[多头注意力] --> B1["多表示子空间
丰富的特征提取"] A --> B2["并行计算
高效的训练速度"] A --> B3["灵活的模式
自适应学习"] B1 --> C[Transformer
强大表现力] B2 --> C B3 --> C C --> D1[GPT 系列] C --> D2[BERT 系列] C --> D3[T5 系列] C --> D4[其他大模型] style A fill:#ff7043 style C fill:#42a5f5 style D1 fill:#66bb6a style D2 fill:#ab47bc style D3 fill:#26c6da style D4 fill:#ffa726

参考资料

  1. Vaswani, A., et al. (2017). “Attention Is All You Need”
  2. Clark, K., et al. (2019). “What Does BERT Look At? An Analysis of BERT’s Attention”
  3. The Illustrated Transformer - Jay Alammar
  4. PyTorch Documentation - nn.MultiheadAttention
CATALOG
  1. 1. 为什么要用”多头”?
    1. 1.1. 单头注意力的局限性
  2. 2. 多头注意力的核心思想
  3. 3. 多头注意力的数学原理
    1. 3.1. 1. 缩放点积注意力(Scaled Dot-Product Attention)
    2. 3.2. 2. 多头注意力的完整公式
  4. 4. 多头注意力的工作流程
  5. 5. 详细步骤解析
    1. 5.1. 步骤 1:线性投影(Linear Projection)
    2. 5.2. 步骤 2:分割为多个头(Split into Heads)
    3. 5.3. 步骤 3:并行计算注意力(Parallel Attention Computation)
    4. 5.4. 步骤 4:拼接与输出变换(Concat & Output Transformation)
  6. 6. 一个具体的计算示例
    1. 6.1. 参数设置
    2. 6.2. 计算过程
  7. 7. 多头注意力的优势
    1. 7.1. 1. 多表示子空间学习
    2. 7.2. 2. 增强模型的表达能力
    3. 7.3. 3. 并行计算效率高
  8. 8. 自注意力 vs 交叉注意力 vs 因果注意力
  9. 9. 代码实现(PyTorch)
  10. 10. 实际应用场景
    1. 10.1. 1. 机器翻译
    2. 10.2. 2. 文本摘要
    3. 10.3. 3. 问答系统
  11. 11. 可视化:注意力模式
    1. 11.1. 模式 1:关注相邻词(Local Attention)
    2. 11.2. 模式 2:关注句法关系(Syntactic Attention)
    3. 11.3. 模式 3:关注全局信息(Global Attention)
  12. 12. 总结
  13. 13. 参考资料