共计 2642 个字符,预计需要花费 7 分钟才能阅读完成。
背景:为什么需要自注意力机制
在自然语言处理领域,传统的 RNN(循环神经网络)在处理长序列时存在明显的局限性。最突出的问题是:

- 梯度消失 / 爆炸 :随着序列长度的增加,RNN 难以有效捕捉远距离依赖关系
- 顺序计算 :无法并行处理序列,训练速度受限于序列长度
- 信息瓶颈 :最后一个隐状态需要压缩整个序列信息
而自注意力机制通过计算词与词之间的关联度,实现了:
- 直接建模任意距离的依赖关系
- 完全并行的序列计算
- 动态权重分配(不同位置关注不同重要性的上下文)
数学原理剖析
基础注意力计算
标准的缩放点积注意力公式为:
$$Attention(Q,K,V) = softmax(\frac{QK^T}{\sqrt{d_k}})V$$
其中:
– $Q \in \mathbb{R}^{n\times d_k}$(查询矩阵)
– $K \in \mathbb{R}^{m\times d_k}$(键矩阵)
– $V \in \mathbb{R}^{m\times d_v}$(值矩阵)
– $\sqrt{d_k}$ 缩放因子防止内积过大导致 softmax 饱和
多头注意力扩展
将 Q、K、V 通过不同的线性投影拆分成 $h$ 个头:
$$
\begin{aligned}
head_i &= Attention(QW_i^Q, KW_i^K, VW_i^V) \
MultiHead(Q,K,V) &= Concat(head_1,…,head_h)W^O
\end{aligned}
$$
每个头的维度通常为 $d_{model}/h$,这样拼接后能保持总维度不变。
PyTorch 实现详解
import torch
import torch.nn as nn
import torch.nn.functional as F
class MultiHeadAttention(nn.Module):
def __init__(self, d_model=512, num_heads=8):
super().__init__()
assert d_model % num_heads == 0, "d_model 必须能被 num_heads 整除"
self.d_model = d_model
self.num_heads = num_heads
self.d_head = d_model // num_heads
# 定义 QKV 的线性变换层
self.wq = nn.Linear(d_model, d_model) # [d_model, d_model]
self.wk = nn.Linear(d_model, d_model)
self.wv = nn.Linear(d_model, d_model)
# 最终输出的线性层
self.wo = nn.Linear(d_model, d_model)
def forward(self, x, mask=None):
"""
输入 x: [batch_size, seq_len, d_model]
输出: [batch_size, seq_len, d_model]
"""
batch_size, seq_len, _ = x.shape
# 1. 计算 QKV [batch_size, seq_len, d_model]
Q = self.wq(x)
K = self.wk(x)
V = self.wv(x)
# 2. 拆分为多头 [batch_size, num_heads, seq_len, d_head]
Q = Q.view(batch_size, -1, self.num_heads, self.d_head).transpose(1, 2)
K = K.view(batch_size, -1, self.num_heads, self.d_head).transpose(1, 2)
V = V.view(batch_size, -1, self.num_heads, self.d_head).transpose(1, 2)
# 3. 计算缩放点积注意力 [batch_size, num_heads, seq_len, seq_len]
attn_scores = torch.matmul(Q, K.transpose(-2, -1)) / torch.sqrt(torch.tensor(self.d_head))
# 应用 mask(如需要)if mask is not None:
attn_scores = attn_scores.masked_fill(mask == 0, -1e9)
# softmax 归一化
attn_weights = F.softmax(attn_scores, dim=-1)
# 4. 计算注意力输出 [batch_size, num_heads, seq_len, d_head]
attn_output = torch.matmul(attn_weights, V)
# 5. 合并多头 [batch_size, seq_len, d_model]
attn_output = attn_output.transpose(1, 2).contiguous()
attn_output = attn_output.view(batch_size, seq_len, self.d_model)
# 6. 最终线性变换
output = self.wo(attn_output)
return output
工程实践关键点
头数选择经验
头数 $h$ 与显存占用的关系近似为:
$$
显存 \approx batch_size \times seq_len^2 \times h \times \frac{d_{model}}{h}
$$
实践中建议:
- 基础模型(d_model=512):8 个头
- 大模型(d_model=1024):16 个头
- 头维度不应小于 64(保证每个头有足够表达能力)
常见陷阱规避
- 维度不匹配 :
- 确保
d_model % num_heads == 0 -
拼接前检查各头维度是否一致
-
掩码应用错误 :
- Decoder 需要严格的下三角掩码(因果注意力)
-
Padding 掩码应在 softmax 前应用
-
梯度不稳定 :
- 使用缩放因子 $1/\sqrt{d_k}$
- 初始化时适当减小线性层权重
延伸思考方向
- 多头机制的普适性 :
- 在浅层网络(如 3 层 Transformer)中,多头是否仍优于单头?
-
不同层是否需要不同数量的注意力头?
-
注意力头可解释性 :
- 能否通过聚类等方法量化不同头捕获的语义特征?
- 特定头是否专门处理语法 / 语义等不同层面的信息?
实践建议
在真实业务场景中使用多头注意力时,建议:
- 先用小批量数据验证维度变换的正确性
- 使用 TensorBoard 可视化注意力权重分布
- 对长序列考虑内存优化的稀疏注意力实现
多头注意力机制是 Transformer 架构的核心创新,理解其实现细节能帮助我们更高效地调试模型,也能启发新的结构改进思路。
