共计 2550 个字符,预计需要花费 7 分钟才能阅读完成。
序列建模的痛点与 Transformer 的诞生
在自然语言处理领域,传统 RNN 和 LSTM 面临着两个核心挑战:

- 长距离依赖问题 :随着序列长度增加,RNN 难以有效捕捉远距离单词间的关系,梯度消失 / 爆炸现象频发。LSTM 通过门控机制缓解了该问题,但实验表明其在超过 100 个 token 的序列上表现仍会显著下降
- 并行计算限制 :RNN 的时序依赖性导致必须按顺序计算,无法充分利用 GPU 的并行计算能力。即便 LSTM 的单个 cell 计算仅需 $O(1)$ 时间,整个序列仍需 $O(n)$ 时间步完成计算
Transformer 通过完全基于注意力机制的架构解决了上述问题:
- 全局依赖性 :自注意力层使任意两个 token 都能直接建立联系,理论最大路径长度仅为 $O(1)$
- 并行计算 :所有位置的 attention 计算可同时进行,训练速度比 LSTM 快 5 -10 倍
Self-Attention 机制数学解析
核心计算公式
$$\text{Attention}(Q,K,V) = \text{softmax}(\frac{QK^T}{\sqrt{d_k}})V$$
其中:
– $Q \in \mathbb{R}^{n \times d_k}$ (Query)
– $K \in \mathbb{R}^{n \times d_k}$ (Key)
– $V \in \mathbb{R}^{n \times d_v}$ (Value)
计算流程分解
-
线性投影 :将输入 embedding $X$ 通过三个权重矩阵投影
$$Q = XW_Q, \quad K = XW_K, \quad V = XW_V$$ -
相似度计算 :通过点积衡量 query 与 key 的关联程度
$$S = QK^T \in \mathbb{R}^{n \times n}$$ -
缩放与归一化 :防止点积结果过大导致 softmax 梯度消失
$$S_{scaled} = \frac{S}{\sqrt{d_k}}$$ -
注意力权重 :通过 softmax 获得归一化的注意力分布
$$A = \text{softmax}(S_{scaled})$$ -
加权求和 :根据注意力权重聚合 value 信息
$$\text{Output} = AV$$
PyTorch 完整实现
import torch
import torch.nn as nn
import math
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
# 线性投影层
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):
"""
Args:
x: [batch_size, seq_len, d_model]
mask: [batch_size, seq_len, seq_len]
"""
batch_size = x.size(0)
# 1. 线性投影并分头
Q = self.W_q(x).view(batch_size, -1, self.n_heads, self.d_k).transpose(1,2)
K = self.W_k(x).view(batch_size, -1, self.n_heads, self.d_k).transpose(1,2)
V = self.W_v(x).view(batch_size, -1, self.n_heads, self.d_k).transpose(1,2)
# 2. 计算缩放点积注意力
scores = torch.matmul(Q, K.transpose(-2,-1)) / math.sqrt(self.d_k)
# 3. 应用 mask(解码器使用)if mask is not None:
scores = scores.masked_fill(mask == 0, -1e9)
# 4. softmax 归一化
attn_weights = torch.softmax(scores, dim=-1)
# 5. 加权求和
context = torch.matmul(attn_weights, V)
# 6. 合并多头输出
context = context.transpose(1,2).contiguous()\
.view(batch_size, -1, self.n_heads * self.d_k)
return self.W_o(context)
计算复杂度分析
训练阶段
-
自注意力层 :
$$4nd^2 + 2n^2d$$
(其中 $n$ 是序列长度,$d$ 是 embedding 维度) -
前馈网络 :
$$8nd^2$$
推理阶段
对于自回归生成任务(如 GPT),需要缓存之前的 key 和 value:
- 第 $t$ 步计算量 :
$$4d^2 + (2t+2)d$$
生产环境优化策略
显存优化
-
梯度检查点 :
from torch.utils.checkpoint import checkpoint def forward(self, x): return checkpoint(self._forward, x) -
激活值压缩 :使用混合精度训练
scaler = torch.cuda.amp.GradScaler() with torch.cuda.amp.autocast(): outputs = model(inputs) loss = criterion(outputs, targets) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update()
序列处理
- 动态 padding:按 batch 内最大长度 padding
- Bucket 策略 :将相似长度的样本分组处理
开放性问题
- 如何优化 attention 的 $O(n^2)$ 计算复杂度?(参考:稀疏注意力、局部窗口)
- 位置编码能否完全替代传统的位置信息建模?(参考:相对位置编码、旋转位置编码)
- 在多模态场景下如何统一设计 attention 机制?(参考:Cross-modality Attention)
