多头注意力机制(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 | import torch |
实际应用场景
多头注意力在以下场景中发挥着关键作用:
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 架构的核心创新,它通过以下三个关键步骤实现强大的特征提取能力:
- 投影:将输入映射到多个表示子空间
- 并行注意力:每个头独立计算注意力,捕捉不同的模式
- 合并:将所有头的输出拼接并变换,融合多视角信息
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
参考资料
- Vaswani, A., et al. (2017). “Attention Is All You Need”
- Clark, K., et al. (2019). “What Does BERT Look At? An Analysis of BERT’s Attention”
- The Illustrated Transformer - Jay Alammar
- PyTorch Documentation - nn.MultiheadAttention