共计 2769 个字符,预计需要花费 7 分钟才能阅读完成。
背景痛点:为什么需要自注意力机制
传统 RNN(循环神经网络)在处理长序列时面临两个核心问题:

- 梯度消失 / 爆炸:随着序列长度增加,反向传播时梯度会指数级衰减或增长,导致模型难以训练。实验表明,当序列长度超过 50 时,LSTM 的准确率会下降约 30%
- 顺序计算限制:必须严格按时间步顺序计算,无法充分利用现代 GPU 的并行计算能力。在处理 1000 个 token 的序列时,RNN 的计算速度比 Transformer 慢约 200 倍
Transformer 架构通过自注意力机制 (Self-Attention) 解决了这些问题:
- 任意两个 token 间可直接建立联系,最大路径长度仅为 O(1)
- 计算过程天然适合并行化,理论 FLOPs 利用率可达 80% 以上
数学原理:缩放点积注意力
自注意力的核心计算公式如下:
$$\text{Attention}(Q,K,V) = \text{softmax}(\frac{QK^T}{\sqrt{d_k}})V$$
其中各变量的物理意义:
- $Q$ (Query): 查询向量,形状为 $[L_q, d_k]$
- $K$ (Key): 键向量,形状为 $[L_k, d_k]$
- $V$ (Value): 值向量,形状为 $[L_k, d_v]$
- $\sqrt{d_k}$: 缩放因子,防止点积结果过大导致 softmax 梯度消失
关键数学推导步骤:
- 计算相似度矩阵:$S = QK^T$(形状 $[L_q, L_k]$)
- 缩放处理:$S’ = S/\sqrt{d_k}$
- 归一化:$A = \text{softmax}(S’)$(注意力权重)
- 加权求和:$O = AV$(形状 $[L_q, d_v]$)
PyTorch 实现详解
基础注意力模块
import torch
import torch.nn as nn
import torch.nn.functional as F
from einops import rearrange
class ScaledDotProductAttention(nn.Module):
"""
输入形状:
q: [batch_size, n_heads, L_q, d_k]
k: [batch_size, n_heads, L_k, d_k]
v: [batch_size, n_heads, L_k, d_v]
mask: [batch_size, L_q, L_k]
"""
def __init__(self, dropout=0.1):
super().__init__()
self.dropout = nn.Dropout(dropout)
def forward(self, q, k, v, mask=None):
# 计算点积注意力分数 [batch_size, n_heads, L_q, L_k]
scores = torch.matmul(q, k.transpose(-2, -1)) / (q.size(-1) ** 0.5)
# 应用 mask(padding mask 或 sequence mask)if mask is not None:
scores = scores.masked_fill(mask == 0, -1e9)
# softmax 归一化
attn = F.softmax(scores, dim=-1)
attn = self.dropout(attn)
# 加权求和 [batch_size, n_heads, L_q, d_v]
output = torch.matmul(attn, v)
return output
多头注意力整合
class MultiHeadAttention(nn.Module):
def __init__(self, d_model=512, n_heads=8, dropout=0.1):
assert d_model % n_heads == 0, "d_model must be divisible by n_heads"
super().__init__()
self.d_k = d_model // n_heads
self.n_heads = n_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)
self.attention = ScaledDotProductAttention(dropout)
self.dropout = nn.Dropout(dropout)
self.layer_norm = nn.LayerNorm(d_model)
def forward(self, q, k, v, mask=None):
residual = q
# 线性变换并分头 [batch_size, L, n_heads, d_k]
q = rearrange(self.w_q(q), 'b l (h d) -> b h l d', h=self.n_heads)
k = rearrange(self.w_k(k), 'b l (h d) -> b h l d', h=self.n_heads)
v = rearrange(self.w_v(v), 'b l (h d) -> b h l d', h=self.n_heads)
# 计算注意力
output = self.attention(q, k, v, mask=mask)
# 合并多头 [batch_size, L, d_model]
output = rearrange(output, 'b h l d -> b l (h d)')
output = self.w_o(output)
# 残差连接 +LayerNorm
output = self.dropout(output)
output = self.layer_norm(output + residual)
return output
显存优化技巧
- 梯度检查点:
from torch.utils.checkpoint import checkpoint output = checkpoint(self.attention, q, k, v, mask) - 混合精度训练:
with torch.cuda.amp.autocast(): output = model(inputs) - 序列分块处理:当序列长度 >512 时,可采用分块计算注意力
常见错误与解决方案
- 忘记 LayerNorm:
- 现象:模型训练不稳定,loss 震荡
-
修正:确保每个子层都有残差连接 +LayerNorm
-
Mask 应用错误:
- 现象:模型在验证集表现异常
-
检查:确保 padding mask 正确传递给每一层
-
Dropout 设置不当:
- 建议:attention dropout 通常设 0.1,hidden dropout 设 0.3
延伸思考
- 位置编码如何影响注意力机制的效果?能否用相对位置编码替代绝对位置编码?
- 当序列长度达到 1024 时,如何优化注意力矩阵的内存占用?
- 多头注意力的头数是否越多越好?如何设计实验验证最优头数?
完整实现可参考 Colab Notebook:点击访问
正文完
