共计 2534 个字符,预计需要花费 7 分钟才能阅读完成。
背景痛点:RNN 的局限性与 Transformer 的革新
传统 RNN 在处理长序列时存在梯度消失 / 爆炸问题,LSTM 虽能缓解但仍受限于顺序计算模式。假设序列长度为 $n$,RNN 的时间复杂度为 $O(n)$,且难以并行化。相比之下,Transformer 的 self-attention 机制通过以下特性实现突破:

- 并行计算 :所有位置间的注意力权重可同时计算
- 长程依赖 :任意两个位置的直接交互不受距离限制
- 复杂度可控 :self-attention 的复杂度为 $O(n^2 \cdot d)$(d 为特征维度),当 $d \ll n$ 时优于 RNN
技术对比:计算复杂度分析
定义输入矩阵 $X \in \mathbb{R}^{n \times d}$,对比三种操作的复杂度:
-
卷积层 (kernel size=k)
$$O(n \cdot d^2 \cdot k)$$ -
循环层 (隐藏层 dim=h)
$$O(n \cdot d \cdot h)$$ -
Self-Attention
$$O(n^2 \cdot d)$$
实际应用中,当 $n > d$ 时 Transformer 更具优势。通过下式理解注意力机制的核心计算:
$$\text{Attention}(Q,K,V) = \text{softmax}\left(\frac{QK^T}{\sqrt{d_k}}\right)V$$
核心实现分步解析
位置编码实现
Transformer 通过以下三角函数为输入注入位置信息:
$$PE_{(pos,2i)} = \sin(pos/10000^{2i/d_{\text{model}}})$$
$$PE_{(pos,2i+1)} = \cos(pos/10000^{2i/d_{\text{model}}})$$
PyTorch 实现示例:
class PositionalEncoding(nn.Module):
def __init__(self, d_model, max_len=5000):
super().__init__()
pe = torch.zeros(max_len, d_model)
position = torch.arange(0, max_len).unsqueeze(1)
div_term = torch.exp(torch.arange(0, d_model, 2) * (-math.log(10000.0) / d_model))
pe[:, 0::2] = torch.sin(position * div_term)
pe[:, 1::2] = torch.cos(position * div_term)
self.register_buffer('pe', pe)
def forward(self, x):
return x + self.pe[:x.size(1)]
多头注意力实现
关键实现步骤:
-
线性投影生成 QKV:
self.q_linear = nn.Linear(d_model, d_model) self.k_linear = nn.Linear(d_model, d_model) self.v_linear = nn.Linear(d_model, d_model) -
分割多头与注意力计算 :
def split_heads(self, x, batch_size): return x.view(batch_size, -1, self.h, self.d_k).transpose(1, 2) # 计算缩放点积注意力 scores = torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(self.d_k) if mask is not None: scores = scores.masked_fill(mask == 0, -1e9) weights = F.softmax(scores, dim=-1) output = torch.matmul(weights, v) -
合并多头输出 :
output = output.transpose(1, 2).contiguous() output = output.view(batch_size, -1, self.d_model)
避坑指南
QKV 初始化方案
使用 Xavier 初始化防止梯度消失:
nn.init.xavier_uniform_(self.q_linear.weight)
nn.init.xavier_uniform_(self.k_linear.weight)
nn.init.xavier_uniform_(self.v_linear.weight)
注意力可视化
通过 plt.matshow 绘制注意力权重矩阵:
import matplotlib.pyplot as plt
plt.matshow(attention_weights[0, 0].detach().numpy())
plt.colorbar()
生产级训练建议
学习率热身策略
采用线性热身 + 逆平方根衰减:
optimizer = AdamW(model.parameters(), lr=0, betas=(0.9, 0.98), eps=1e-9)
# 热身阶段线性增加学习率
lr = min(step_num**-0.5, step_num * warmup_steps**-1.5)
混合精度训练
使用 Apex 库实现 FP16 训练:
from apex import amp
model, optimizer = amp.initialize(model, optimizer, opt_level="O1")
with amp.scale_loss(loss, optimizer) as scaled_loss:
scaled_loss.backward()
代码规范实践
关键张量维度验证示例:
assert q.size() == (batch_size, h, seq_len, d_k)
assert k.size() == (batch_size, h, seq_len, d_k)
assert v.size() == (batch_size, h, seq_len, d_v)
延伸改进方向
- 添加残差连接 :在每个子层实现 Add & Norm 操作
- 修改注意力计算 :尝试 ReLU 注意力替代 softmax
- 稀疏注意力 :实现局部窗口注意力降低计算复杂度
通过上述实现,读者可掌握 Transformer 的核心机制与工业级实现技巧。建议在完成基础版本后,逐步尝试改进方向以深入理解模型设计原理。
