共计 2772 个字符,预计需要花费 7 分钟才能阅读完成。
传统 RNN 的困境与注意力机制的诞生
在自然语言处理领域,循环神经网络 (RNN) 曾是处理序列数据的标准方案。但实践中我们发现三个主要问题:

- 长程依赖丢失:随着序列长度增加,RNN 难以有效传递早期信息(梯度消失 / 爆炸问题)
- 顺序计算瓶颈:必须按时间步逐步计算,无法利用现代 GPU 的并行能力
- 固定编码局限:每个时间步的隐藏状态被迫包含所有历史信息,缺乏重点聚焦
自注意力机制 (Self-Attention) 原理详解
核心计算流程
-
输入表示:
对于输入序列 $X \in \mathbb{R}^{n \times d_{model}}$(n 为序列长度,$d_{model}$ 为特征维度),通过三个可学习矩阵投影得到:
$$
Q = XW^Q, \quad K = XW^K, \quad V = XW^V
$$
($W^Q, W^K \in \mathbb{R}^{d_{model} \times d_k}$, $W^V \in \mathbb{R}^{d_{model} \times d_v}$) -
注意力权重计算:
$$
\text{Attention}(Q,K,V) = \text{softmax}(\frac{QK^T}{\sqrt{d_k}})V
$$ - 除以 $\sqrt{d_k}$ 防止点积数值过大导致 softmax 梯度消失
-
softmax 沿每一行计算,保证权重和为 1
-
代码实现
import torch import torch.nn as nn import torch.nn.functional as F class SelfAttention(nn.Module): def __init__(self, d_model, d_k, d_v): super().__init__() self.W_q = nn.Linear(d_model, d_k) # Query 投影 self.W_k = nn.Linear(d_model, d_k) # Key 投影 self.W_v = nn.Linear(d_model, d_v) # Value 投影 self.scale = d_k ** -0.5 def forward(self, x): """ 输入: [batch_size, seq_len, d_model] 输出: [batch_size, seq_len, d_v] """ Q = self.W_q(x) # [B, L, d_k] K = self.W_k(x) # [B, L, d_k] V = self.W_v(x) # [B, L, d_v] attn = torch.matmul(Q, K.transpose(-1, -2)) * self.scale attn = F.softmax(attn, dim=-1) output = torch.matmul(attn, V) return output
多头注意力机制 (Multi-Head Attention) 进阶
设计动机
- 单一注意力头的局限:只能学习一种注意力模式
- 并行化优势:多个头可同时捕捉不同子空间的语义关系
实现关键
- 头部拆分与合并:
- 将 Q /K/ V 拆分为 h 份(h 为头数),每份维度 $d_k=d_v=d_{model}/h$
-
计算 h 个独立的注意力头后拼接结果
-
PyTorch 实现
class MultiHeadAttention(nn.Module): def __init__(self, d_model, num_heads): super().__init__() assert d_model % num_heads == 0 self.d_k = d_model // num_heads self.num_heads = 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 forward(self, x): batch_size = x.size(0) # 投影后拆分多头 [B, L, d_model] -> [B, L, h, d_k] Q = self.W_q(x).view(batch_size, -1, self.num_heads, self.d_k) K = self.W_k(x).view(batch_size, -1, self.num_heads, self.d_k) V = self.W_v(x).view(batch_size, -1, self.num_heads, self.d_k) # 转置为 [B, h, L, d_k] Q = Q.transpose(1, 2) K = K.transpose(1, 2) V = V.transpose(1, 2) # 计算缩放点积注意力 scores = torch.matmul(Q, K.transpose(-1, -2)) / (self.d_k ** 0.5) attn = F.softmax(scores, dim=-1) context = torch.matmul(attn, V) # 合并多头 [B, h, L, d_k] -> [B, L, d_model] context = context.transpose(1, 2).contiguous() context = context.view(batch_size, -1, self.num_heads * self.d_k) return self.W_o(context)
性能对比与实验观察
| 指标 | 单头注意力 | 多头注意力 (h=8) |
|---|---|---|
| 计算复杂度 | O(n²d) | O(n²d) |
| 参数量 | 3dd_k | 3dd |
| 并行度 | 低 | 高 |
| 语义捕获能力 | 单一模式 | 多样化模式 |
实际任务中(如机器翻译),多头注意力通常能带来 1.5-2.5 BLEU 值提升。
五大避坑指南
- 维度不对齐错误
- 现象:
RuntimeError: mat1 and mat2 shapes cannot be multiplied -
检查:确保 $d_{model}$ 能被头数整除,投影后张量形状匹配
-
softmax 数值溢出
- 现象:注意力权重出现 NaN
-
解决:必须进行缩放(除以 $\sqrt{d_k}$),对特别长的序列可分段计算
-
梯度消失问题
- 现象:模型难以学习远程依赖
- 对策:配合残差连接和 LayerNorm 使用
延伸思考方向
- 头数超参数选择:
- 实验发现不同任务最优头数不同(翻译常用 8 头,分类可能 4 头足够)
-
可通过注意力头可视化分析各头学习到的模式
-
计算效率优化:
- 稀疏注意力、局部注意力等变体如何权衡效果与速度
- FlashAttention 等优化技术原理探究
总结启示
通过实现完整的自注意力和多头注意力模块,我们深入理解了 Transformer 的核心设计思想。关键收获包括:
– 注意力机制通过动态权重实现序列元素的直接交互
– 多头设计类似 CNN 的多通道,能并行学习多样特征
– 实际使用时需注意数值稳定性和计算效率的平衡
建议读者在完成基础实现后,进一步尝试将其应用到具体 NLP 任务中,观察不同超参数配置下的性能变化,这将大大加深对机制的理解。
