共计 2991 个字符,预计需要花费 8 分钟才能阅读完成。
传统序列建模的瓶颈
在处理文本、语音等序列数据时,传统 RNN/LSTM 面临两大核心痛点:

-
梯度消失问题:当序列长度超过 50 步时,反向传播的梯度会指数级衰减,导致模型难以学习长距离依赖关系。数学上可表示为:
$$\frac{\partial L}{\partial h_t} \approx \prod_{k=t}^{T} \frac{\partial h_{k+1}}{\partial h_k} \to 0 \quad (T\gg t)$$ -
顺序计算限制:RNN 的时序依赖性导致无法并行计算,处理长为 $n$ 的序列需要 $O(n)$ 时间步。这在处理万字长文或 DNA 序列时尤为致命
注意力机制的三要素
注意力机制通过 Query/Key/Value 分解实现动态权重分配,其核心计算流程为:
$$\text{Attention}(Q,K,V) = \text{softmax}(\frac{QK^T}{\sqrt{d_k}})V$$
- Query:当前需要计算的特征表示(如解码器当前词)
- Key:待检索的特征集合(如编码器所有词)
- Value:实际返回的特征信息(通常与 Key 维度相同)
常见变体对比
| 类型 | 计算公式 | 复杂度 | 适用场景 |
|---|---|---|---|
| 加性注意力 | $v^T\tanh(W_q q + W_k k)$ | $O(d^2)$ | 低维空间 |
| 点积注意力 | $q^Tk$ | $O(d)$ | 高维空间 |
| 缩放点积注意力 | $q^Tk/\sqrt{d}$ | $O(d)$ | Transformer 默认 |
PyTorch 实现带 mask 的多头注意力
import torch
import torch.nn as nn
import torch.nn.functional as F
class MultiHeadAttention(nn.Module):
def __init__(self, d_model=512, n_heads=8):
super().__init__()
assert d_model % n_heads == 0
self.d_k = d_model // n_heads
self.n_heads = n_heads
# 线性变换层 (batch_size, seq_len, d_model) -> (batch_size, seq_len, d_model)
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, q, k, v, mask=None):
batch_size = q.size(0)
# 线性变换 + 分头 (batch_size, seq_len, d_model) -> (batch_size, seq_len, n_heads, d_k)
q = self.w_q(q).view(batch_size, -1, self.n_heads, self.d_k)
k = self.w_k(k).view(batch_size, -1, self.n_heads, self.d_k)
v = self.w_v(v).view(batch_size, -1, self.n_heads, self.d_k)
# 转置为 (batch_size, n_heads, seq_len, d_k)
q, k, v = q.transpose(1, 2), k.transpose(1, 2), v.transpose(1, 2)
# 计算缩放点积注意力
scores = torch.matmul(q, k.transpose(-2, -1)) / torch.sqrt(torch.tensor(self.d_k))
# 应用 mask(如因果掩码)if mask is not None:
scores = scores.masked_fill(mask == 0, -1e9)
# 温度系数调节(训练稳定技巧)temperature = 1.0 # 可动态调整
attn = F.softmax(scores * temperature, dim=-1)
# 加权求和 + 合并头
output = torch.matmul(attn, v)
output = output.transpose(1, 2).contiguous().view(batch_size, -1, self.d_model)
return self.w_o(output)
关键实现细节:
– 温度系数:通过调整 softmax 前的数值尺度控制注意力分布的尖锐程度
– 数值稳定性:对 masked 位置用 -1e9 替代负无穷,避免 NaN
– 分头计算:将 d_model 拆分为 n_heads 个 d_k 子空间,并行计算
Transformer 架构实战
在标准 Transformer 中,自注意力层通过以下方式增强模型能力:
- 编码器自注意力:建立输入序列全局依赖
- 解码器自注意力:结合因果掩码实现自回归生成
- 交叉注意力:连接编码器 - 解码器信息流
KV 缓存加速推理
在自回归生成时,可通过缓存历史 Key/Value 避免重复计算:
class DecoderLayer(nn.Module):
def __init__(self, ...):
self.self_attn = MultiHeadAttention()
self.cross_attn = MultiHeadAttention()
def forward(self, x, encoder_out, past_kv=None):
# 自注意力(带因果掩码)self_attn_out = self.self_attn(
q=x, k=x, v=x,
mask=torch.tril(torch.ones(seq_len, seq_len)) # 下三角掩码
)
# 交叉注意力(使用编码器输出)cross_attn_out = self.cross_attn(
q=self_attn_out,
k=encoder_out,
v=encoder_out
)
# 更新 KV 缓存
new_kv = torch.cat([past_kv, current_kv], dim=1) if past_kv is not None else current_kv
return cross_attn_out, new_kv
生产环境避坑指南
- 内存溢出(OOM)
- 解决方案:采用梯度检查点 (gradient checkpointing) 或分块注意力
-
示例:将序列分成 64-128 的块处理
-
训练不稳定
- 现象:损失出现 NaN/INF
-
对策:
- 添加 LayerNorm
- 使用 Xavier 初始化
- 限制最大序列长度
-
长序列性能下降
- 优化方案:
- 稀疏注意力(如 Longformer 的滑动窗口)
- 线性注意力(Reformer 的 LSH 分桶)
IWSLT 实验对比
| 模型 | BLEU-4 | 参数量 | 推理速度(词 / 秒) |
|---|---|---|---|
| LSTM+Attention | 28.7 | 65M | 120 |
| Transformer(base) | 32.1 | 65M | 310 |
| Transformer(big) | 33.5 | 213M | 190 |
动手挑战
尝试实现 局部窗口注意力:
1. 修改注意力计算,使每个 token 只关注前后 $w$ 个邻居
2. 对比全局注意力的效果差异
3. 思考如何平衡局部与全局信息(提示:可参考 Swin Transformer)
参考文献
- Vaswani et al. Attention Is All You Need. NeurIPS 2017
- Dai et al. Transformer-XL. ACL 2019
- Kitaev et al. Reformer. ICLR 2020
