🎯 课程主题
多头注意力机制(Multi-Head Attention)——通过多组独立注意力头捕获不同子空间的信息。
📝 核心知识点
1. 多头注意力机制的定义
- 概念说明:做 $H$ 组独立的注意力机制,再融合输出。
- 关键细节:
- 论文中使用 $H = 8$ 个注意力头
- 每个头有独立的 $W_Q, W_K, W_V$ 权重矩阵(参数不同)
- 每个头独立计算自注意力
- 最后将所有头的输出拼接(Concat)并融合
2. 单个头的计算流程
- 概念说明:每个头执行标准的自注意力计算。
- 关键细节:
- 输入 $X$(带位置编码)→ 映射为 $Q, K, V$
- 计算注意力得分:$\text{Softmax}(\frac{QK^T}{\sqrt{d_k}})$
- 与 $V$ 相乘得到输出 $B$
- 每个头输出 $L \times d_k$
3. 多头融合过程
- 概念说明:将所有头的输出拼接后通过线性层融合。
- 关键细节:
- 8 个头的输出 Concat:$L \times (8 \times d_k)$
- 与权重矩阵 $W$ 相乘(全连接层):$(8 \times d_k) \times D$
- 最终输出:$L \times D$(与输入维度一致)
- 权重矩阵 $W$ 负责将拼接结果映射回原始维度
4. 多头机制的优势
- 概念说明:不同头在不同子空间捕获不同特征。
- 关键细节:
- 不同头关注不同方面:
- 有的头看短距离依赖
- 有的头看长距离依赖
- 有的头关注实体词
- 有的头关注标点、句法
- 有的头关注主谓关系
- 有的头关注介词关系
- 通过反向传播训练,各头自动学习不同关注点
- 效果优于单头注意力
- 不同头关注不同方面:
5. 输入输出维度一致性
- 概念说明:多头注意力输入输出维度相同,便于残差连接。
- 关键细节:
- 输入 $X$:$L \times D$
- 输出 $B$:$L \times D$
- 维度一致才能做残差连接和层归一化
6. 支持变长序列
- 概念说明:注意力机制支持不同长度的输入序列。
- 关键细节:
- 输入 $A_1, A_2$ → 输出 $B_1, B_2$
- 输入 $A_1, A_2, A_3, A_4$ → 输出 $B_1, B_2, B_3, B_4$
- 不需要定长,适合机器翻译等变长任务
🧮 核心公式与推导
多头注意力公式:
$$\text{MultiHead}(Q, K, V) = \text{Concat}(\text{head}_1, \ldots, \text{head}_H) W^O$$
单个头的计算:
$$\text{head}_i = \text{Attention}(Q W_i^Q, K W_i^K, V W_i^V) = \text{Softmax}\left(\frac{Q_i K_i^T}{\sqrt{d_k}}\right) V_i$$
其中:
- $H = 8$(头数)
- $W_i^Q \in \mathbb{R}^{D \times d_k}$,$W_i^K \in \mathbb{R}^{D \times d_k}$,$W_i^V \in \mathbb{R}^{D \times d_v}$
- $d_k = d_v = D / H = 512 / 8 = 64$
- $W^O \in \mathbb{R}^{(H \cdot d_v) \times D}$ 输出投影矩阵
维度变化:
- 每个头输出:$L \times d_k$
- Concat 后:$L \times (H \times d_k) = L \times (8 \times 64) = L \times 512$
- 经 $W^O$ 投影后:$L \times D = L \times 512$
🏗️ 模型架构与数据流向
💻 代码实战
import torch
import torch.nn as nn
import torch.nn.functional as F
import math
class MultiHeadAttention(nn.Module):
"""多头注意力机制实现"""
def __init__(self, d_model, n_heads):
"""
d_model: D, 词嵌入维度 (如512)
n_heads: H, 注意力头数 (如8)
"""
super().__init__()
self.d_model = d_model
self.n_heads = n_heads
self.d_k = d_model // n_heads # 每个头的维度 dk = D/H = 64
# 所有头的QKV权重合并为一个大矩阵 (效率更高)
self.W_Q = nn.Linear(d_model, d_model, bias=False) # D×D (包含8个头的D×dk)
self.W_K = nn.Linear(d_model, d_model, bias=False)
self.W_V = nn.Linear(d_model, d_model, bias=False)
self.W_O = nn.Linear(d_model, d_model, bias=False) # 输出投影
def forward(self, X, mask=None):
"""
X: [batch, L, D]
return: [batch, L, D]
"""
batch_size, L, _ = X.size()
# 1. QKV映射: [batch, L, D] → [batch, L, D]
Q = self.W_Q(X)
K = self.W_K(X)
V = self.W_V(X)
# 2. 分割多头: [batch, L, D] → [batch, L, H, dk] → [batch, H, L, dk]
Q = Q.view(batch_size, L, self.n_heads, self.d_k).transpose(1, 2)
K = K.view(batch_size, L, self.n_heads, self.d_k).transpose(1, 2)
V = V.view(batch_size, L, self.n_heads, self.d_k).transpose(1, 2)
# 3. 计算注意力: [batch, H, L, dk] @ [batch, H, dk, L] → [batch, H, L, L]
scores = torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(self.d_k)
if mask is not None:
scores = scores.masked_fill(mask == 0, float('-inf'))
attn_weights = F.softmax(scores, dim=-1)
# 4. 信息融合: [batch, H, L, L] @ [batch, H, L, dk] → [batch, H, L, dk]
output = torch.matmul(attn_weights, V)
# 5. 合并多头: [batch, H, L, dk] → [batch, L, H, dk] → [batch, L, D]
output = output.transpose(1, 2).contiguous()
output = output.view(batch_size, L, self.d_model)
# 6. 输出投影: [batch, L, D] → [batch, L, D]
output = self.W_O(output)
return output
# 示例
d_model = 512
n_heads = 8
batch_size = 2
L = 4
mha = MultiHeadAttention(d_model, n_heads)
X = torch.randn(batch_size, L, d_model)
output = mha(X)
print(output.shape) # torch.Size([2, 4, 512]) 输入输出维度一致
⚠️ 常见问题与避坑指南
- 每个头的 $W_Q, W_K, W_V$ 参数不同,这是多头机制的核心
- $d_k = D / H$,需保证 $D$ 能被 $H$ 整除(如 $512 / 8 = 64$)
- Concat 是在 $d_k$ 维度上拼接,不是在 $L$ 维度
- 输出维度需与输入维度一致($L \times D$),才能做后续的残差连接
- 多头注意力搞定后,Transformer 约 40% 的工作就解决了
💡 个人总结与延伸
多头注意力机制通过并行运行多个注意力头,让模型在不同子空间捕获不同维度的信息(短/长依赖、句法/语义等),显著提升了表达能力。这是 Transformer 区别于传统注意力机制的关键创新。现代大模型中,多头数量通常更多(如 32、96 头),且衍生出 GQA(分组查询注意力)、MQA(多查询注意力)等变体以平衡性能与效率。