共计 2009 个字符,预计需要花费 6 分钟才能阅读完成。
为什么需要 Transformer?
传统 RNN 在处理长序列时存在明显的梯度消失问题。对于时间步 $t$ 的隐藏状态 $h_t$,其梯度计算可以表示为:

$$\frac{\partial L}{\partial h_t} = \sum_{k=1}^{t} \left(\prod_{i=k+1}^{t} \frac{\partial h_i}{\partial h_{i-1}} \right) \frac{\partial L}{\partial h_k}$$
当 $\frac{\partial h_i}{\partial h_{i-1}}$ 的范数小于 1 时,随着 $t-k$ 增大,梯度会指数级衰减。而 Transformer 通过 Self-Attention 机制实现了:
- 并行计算 :所有位置的注意力权重可同时计算
- 长距离依赖 :任意两个位置的直接交互,不受序列长度限制
- 计算效率 :时间复杂度 $O(n^2 \cdot d)$,其中 $n$ 为序列长度,$d$ 为特征维度
Self-Attention 核心实现
Scaled Dot-Product Attention 的计算流程可分为四个阶段:
- Query-Key-Value 投影 :将输入 $X \in \mathbb{R}^{n \times d_{model}}$ 分别线性变换为 Q /K/V
# 使用 einops 进行清晰的可视化维度变换
q = einsum('n d, d h -> n h', x, W_q) # [n, d_k]
k = einsum('n d, d h -> n h', x, W_k) # [n, d_k]
v = einsum('n d, d h -> n h', x, W_v) # [n, d_v]
- 注意力分数计算 :
$$\text{Attention}(Q,K,V) = \text{softmax}\left(\frac{QK^T}{\sqrt{d_k}}\right)V$$
attn_scores = einsum('i d, j d -> i j', q, k) / sqrt(d_k) # [n, n]
-
Softmax 归一化 :对每行进行概率化处理
-
加权求和 :根据注意力权重聚合 value 信息
完整实现示例
以下是一个包含 LayerNorm 和残差连接的 Transformer 层实现(PyTorch 版本):
class TransformerLayer(nn.Module):
def __init__(self, d_model: int, n_heads: int):
super().__init__()
self.attn = MultiHeadAttention(d_model, n_heads)
self.ffn = PositionwiseFFN(d_model)
self.norm1 = nn.LayerNorm(d_model)
self.norm2 = nn.LayerNorm(d_model)
def forward(self, x: Tensor) -> Tensor:
# 建议在 GPU 上运行时启用 cudnn 优化
with torch.backends.cudnn.flags(enabled=True):
# 残差连接 + 层归一化
attn_out = self.attn(self.norm1(x))
x = x + attn_out # [batch, seq, dim]
# FFN 部分
ffn_out = self.ffn(self.norm2(x))
return x + ffn_out
训练避坑指南
- 梯度爆炸问题 :
- 解决方案:使用梯度裁剪(
torch.nn.utils.clip_grad_norm_) -
推荐值:设置 max_norm 在 1.0-5.0 之间
-
过拟合问题 :
- 添加 Dropout(注意力权重和 FFN 层)
-
使用 Label Smoothing(0.1 smoothing factor)
-
显存不足 :
- 采用梯度检查点(
torch.utils.checkpoint) - 降低 batch size 并累积梯度
延伸改进方向
- 稀疏注意力 :
- Local Attention:限制每个 token 只能关注周围窗口
-
Block Sparse:将注意力矩阵分块稀疏化
-
内存优化 :
- Flash Attention 算法(减少 HBM 访问次数)
- 混合精度训练(FP16+FP32)
性能评估示例
计算 FLOPs 的简易方法(以单头注意力为例):
- QK^T 计算:$2 \times n^2 \times d_k$
- Softmax:$3 \times n^2$
- AV 计算:$2 \times n^2 \times d_v$
对于 8 头 $d_{model}=512$ 的配置,处理 512 长度序列时:
$$\text{FLOPs} \approx 8 \times (2 \times 512^2 \times 64 + 3 \times 512^2 + 2 \times 512^2 \times 64) = 1.34 \text{GFLOPs}$$
经过这段时间的实践,我发现 Transformer 的实现虽然复杂,但只要理解清楚每个组件的设计意图,配合 PyTorch 的自动微分机制,就能逐步搭建出可用的模型。建议初学者从单头注意力开始,逐步扩展到完整架构,这样更容易定位问题。
