共计 2289 个字符,预计需要花费 6 分钟才能阅读完成。
传统序列模型的局限性
在 Transformer 出现之前,RNN 和 LSTM 是处理序列数据的主流方法。但它们存在两个致命缺陷:

- 顺序计算依赖 :必须逐个处理序列元素,无法并行计算。对于长度为 n 的序列,时间复杂度高达 O(n)
- 长距离依赖衰减 :信息通过隐藏状态逐步传递,超过 20 步后梯度消失严重(参考 [Hochreiter 1991])
自注意力机制原理
QKV 计算过程
给定输入矩阵 $X \in \mathbb{R}^{n \times d_{model}}$,通过三个权重矩阵得到查询 (Query)、键 (Key)、值 (Value):
$$
Q = XW^Q, \quad W^Q \in \mathbb{R}^{d_{model} \times d_k}
$$
$$
K = XW^K, \quad W^K \in \mathbb{R}^{d_{model} \times d_k}
$$
$$
V = XW^V, \quad W^V \in \mathbb{R}^{d_{model} \times d_v}
$$
缩放点积注意力
计算注意力权重的完整公式:
$$
Attention(Q,K,V) = softmax(\frac{QK^T}{\sqrt{d_k}})V
$$
缩放因子 $\sqrt{d_k}$ 防止点积结果过大导致 softmax 梯度消失
多头注意力
将 QKV 拆分成 h 个头并行计算:
$$
head_i = Attention(QW_i^Q, KW_i^K, VW_i^V)
$$
$$
MultiHead = Concat(head_1,…,head_h)W^O
$$
PyTorch 实现
单头注意力
import torch
import torch.nn as nn
import torch.nn.functional as F
class ScaledDotProductAttention(nn.Module):
def __init__(self, d_k):
super().__init__()
self.d_k = d_k
def forward(self, Q, K, V, mask=None):
# Q,K,V shape: (batch, seq_len, d_k)
scores = torch.matmul(Q, K.transpose(-2, -1)) / (self.d_k ** 0.5)
if mask is not None:
scores = scores.masked_fill(mask == 0, -1e9)
attn = F.softmax(scores, dim=-1)
output = torch.matmul(attn, V)
return output
完整多头注意力
class MultiHeadAttention(nn.Module):
def __init__(self, d_model=512, h=8):
super().__init__()
assert d_model % h == 0
self.d_k = d_model // h
self.h = h
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, mask=None):
batch_size = x.size(0)
# 线性投影
Q = self.W_Q(x) # (batch, seq, d_model)
K = self.W_K(x)
V = self.W_V(x)
# 分头处理
Q = Q.view(batch_size, -1, self.h, self.d_k).transpose(1,2)
K = K.view(batch_size, -1, self.h, self.d_k).transpose(1,2)
V = V.view(batch_size, -1, self.h, self.d_k).transpose(1,2)
# 计算注意力
scores = torch.matmul(Q, K.transpose(-2, -1)) / (self.d_k ** 0.5)
if mask is not None:
mask = mask.unsqueeze(1) # 广播到所有头
scores = scores.masked_fill(mask == 0, -1e9)
attn = F.softmax(scores, dim=-1)
context = torch.matmul(attn, V)
# 合并多头输出
context = context.transpose(1,2).contiguous()
context = context.view(batch_size, -1, self.h * self.d_k)
return self.W_O(context)
实践经验
超参选择策略
- 头数 h :通常 4 -16 之间,建议先用 8
- 模型维度 d_model:一般 512/768/1024,需能被 h 整除
内存优化技巧
- 梯度检查点:
torch.utils.checkpoint - 序列分块:将长序列拆分为多个子序列
- 混合精度训练:
torch.cuda.amp
梯度问题预防
- 初始化:使用 Xavier/Glorot 初始化
- 层归一化:每个子层后添加 LayerNorm
- 学习率预热:前 4000 步线性增加学习率
性能对比
| 模型 | 时间复杂度 | 空间复杂度 | 并行性 |
|---|---|---|---|
| RNN | O(n) | O(1) | 无 |
| LSTM | O(n) | O(1) | 无 |
| Transformer | O(n²) | O(n²) | 完全 |
实际测试中(n=512),Transformer 比 LSTM 快 3 - 5 倍
延伸思考
- 如何改造自注意力机制使其适用于图像数据?
- 当序列长度达到 10 万时,有哪些优化方法?
- 为什么说自注意力机制本质是一种图神经网络?
正文完
发表至: 未分类
近两天内
