Transformer详解:从自注意力手算到PyTorch完整实现
上一篇 Seq2Seq 文章中,我们使用 GRU 构造了 Encoder-Decoder,并通过 Attention 缓解“把整句话压缩成一个向量”的信息瓶颈。
不过,RNN、LSTM 和 GRU 还有一个很难绕开的特点:
必须按照时间顺序逐步计算,后一个位置需要等待前一个位置的隐藏状态。
假设一句话有 100 个 token,RNN 需要从第 1 个 token 一直算到第 100 个 token。即使 GPU 很擅长矩阵并行,也不能完全消除这种前后依赖。
Transformer 改变了处理序列的方式:
不再使用循环逐步传递信息,而是让每个 token 直接通过 Attention 查看序列中的其他 token。
这样一来,训练时可以并行处理整段序列,也更容易建模相距很远的词之间的关系。
1. Transformer 要解决什么问题
先看一句话:
小明 把 苹果 放在 桌子 上,因为 他 刚刚 买了 它
理解这句话时,需要建立多处联系:
- “他”更可能指向“小明”。
- “它”更可能指向“苹果”。
- “放在”与“桌子上”共同表达一个位置关系。
在 RNN 中,前面的信息要经过多个时间步才能传到后面;序列很长时,远距离信息的传播路径也会很长。
Transformer 中,“它”可以直接对“小明”“苹果”“桌子”等所有位置计算相关程度,再把有用信息聚合回来。
1.1 RNN 与 Transformer 的主要区别
| 对比项 | RNN、LSTM、GRU | Transformer |
|---|---|---|
| 序列处理 | 按时间步递归 | 整个序列并行计算 |
| 信息交互 | 通过隐藏状态逐步传递 | token 之间直接计算 Attention |
| 位置信息 | 顺序天然包含在递归中 | 需要额外加入位置编码 |
| 长距离依赖 | 信息传播路径较长 | 任意两个位置可直接交互 |
| 主要计算 | 循环和矩阵乘法 | 大规模矩阵乘法 |
| 长序列代价 | 串行步骤多 | 标准 Attention 的矩阵为 $n\times n$ |
Transformer 并不是简单地“把 RNN 换成 Attention”。它还包含:
- Multi-Head Attention。
- Positional Encoding。
- Feed Forward Network。
- Residual Connection。
- Layer Normalization。
- Padding Mask 与 Causal Mask。
下面逐块拆开。
2. Transformer 的整体结构
原始 Transformer 用于机器翻译,因此仍然采用 Encoder-Decoder 结构。
源语言 token
-> Transformer Encoder
-> 上下文表示 Memory
-> Transformer Decoder
-> 目标语言下一个 token 的概率
2.1 Encoder 做什么
Encoder 接收完整输入序列,让每个输入 token 与其他输入 token 交换信息,最终得到一组包含上下文的表示:
$$
H=(h_1,h_2,\ldots,h_{T_x})
$$
与早期 Seq2Seq 只保留最后一个隐藏状态不同,Transformer Encoder 会保留每个位置的输出。
2.2 Decoder 做什么
Decoder 根据两部分信息预测下一个目标 token:
- 已经生成的目标 token。
- Encoder 输出的完整 Memory。
训练时,目标序列可以整体送入 Decoder,但必须使用 Causal Mask 遮住未来位置,防止模型偷看答案。
2.3 Encoder 和 Decoder 并非所有模型都同时使用
| 架构 | 使用部分 | 更适合的任务 |
|---|---|---|
| Encoder-only | 只使用 Encoder | 分类、序列标注、文本表示 |
| Decoder-only | 只使用 Decoder | 自回归文本生成 |
| Encoder-Decoder | 两者都使用 | 翻译、摘要、序列转换 |
因此,Transformer 是一套可组合的结构,不只是一种固定模型。
3. 输入进入 Transformer 前发生了什么
Transformer 不直接处理文字,而是处理三维张量。
假设:
- batch size 为 $B$。
- 序列长度为 $L$。
- 模型维度为 $d_{model}$。
输入 token id 的形状是:
[B, L]
经过 Embedding 后变成:
[B, L, d_model]
再加入位置编码:
$$
Z=X_{embedding}+P_{position}
$$
最终的 $Z$ 才会送入第一个 Transformer 层。
例如:
token id: [32, 20]
Embedding: [32, 20, 128]
位置编码相加后: [32, 20, 128]
位置编码与 Embedding 的维度必须相同,才能逐元素相加。
4. Attention 的直觉:查询、钥匙和值
Attention 最容易让人困惑的三个字母是:
- Query,简称 $Q$。
- Key,简称 $K$。
- Value,简称 $V$。
可以把一次检索理解为:
Query:我现在想找什么信息?
Key:我这里保存的信息适合被怎样检索?
Value:如果匹配成功,我真正提供什么内容?
以“它”为例:
Query("它") 与 Key("小明") 计算相似度
Query("它") 与 Key("苹果") 计算相似度
Query("它") 与 Key("桌子") 计算相似度
假设“它”与“苹果”的匹配分数最高,模型就会给“苹果”的 Value 更高权重。
4.1 Q、K、V 从哪里来
对于输入矩阵 $X$,通过三个不同的线性变换得到:
$$
Q=XW_Q
$$
$$
K=XW_K
$$
$$
V=XW_V
$$
其中 $W_Q$、$W_K$、$W_V$ 都是可学习参数。
需要注意:
Self-Attention 中 Q、K、V 都来自同一组输入,但使用不同参数投影,所以它们通常并不相等。
训练开始时,这些矩阵通常是随机初始化的。反向传播会逐渐学习“怎样提出查询”“怎样参与匹配”和“怎样提供内容”。
5. 缩放点积注意力
Transformer 使用的核心公式是:
$$
\operatorname{Attention}(Q,K,V)
=\operatorname{softmax}\left(\frac{QK^T}{\sqrt{d_k}}\right)V
$$
不要急着记公式,可以拆成四步。
第一步:计算匹配分数
$$
S=QK^T
$$
$S_{ij}$ 表示第 $i$ 个 Query 与第 $j$ 个 Key 的匹配程度。
第二步:进行缩放
$$
S’=\frac{S}{\sqrt{d_k}}
$$
$d_k$ 是每个 Key 和 Query 的向量维度。
第三步:转成注意力权重
$$
A=\operatorname{softmax}(S’)
$$
Softmax 按最后一个维度进行,所以 $A$ 的每一行之和都是 1。
第四步:对 Value 加权求和
$$
O=AV
$$
输出的每个位置都融合了所有可见位置的信息。
6. 手算一次完整 Self-Attention
为了看清计算过程,假设只有 3 个 token,每个 token 是二维向量:
$$
X=
\begin{bmatrix}
1&0\
0&1\
1&1
\end{bmatrix}
$$
为了简化计算,暂时令:
$$
W_Q=W_K=W_V=I
$$
因此:
$$
Q=K=V=X
$$
实际模型不会固定为单位矩阵,这里只是为了手算。
6.1 计算 $QK^T$
$$
QK^T=
\begin{bmatrix}
1&0&1\
0&1&1\
1&1&2
\end{bmatrix}
$$
第一行来自:
$$
[1,0]\cdot[1,0]=1
$$
$$
[1,0]\cdot[0,1]=0
$$
$$
[1,0]\cdot[1,1]=1
$$
它表示第一个 token 对三个 token 的原始匹配分数。
6.2 除以 $\sqrt{d_k}$
这里 $d_k=2$:
$$
\frac{QK^T}{\sqrt{2}}
\approx
\begin{bmatrix}
0.707&0&0.707\
0&0.707&0.707\
0.707&0.707&1.414
\end{bmatrix}
$$
6.3 对每一行做 Softmax
以第一行为例:
$$
\operatorname{softmax}([0.707,0,0.707])
\approx[0.401,0.198,0.401]
$$
完整权重矩阵近似为:
$$
A\approx
\begin{bmatrix}
0.401&0.198&0.401\
0.198&0.401&0.401\
0.248&0.248&0.503
\end{bmatrix}
$$
每一行都满足:
$$
\sum_j A_{ij}=1
$$
6.4 用权重聚合 Value
第一个位置的输出是:
$$
o_1
=0.401[1,0]+0.198[0,1]+0.401[1,1]
$$
$$
o_1\approx[0.802,0.599]
$$
三个位置的完整输出为:
$$
O=AV\approx
\begin{bmatrix}
0.802&0.599\
0.599&0.802\
0.751&0.751
\end{bmatrix}
$$
原来的第一个 token 是 $[1,0]$,经过 Attention 后变成 $[0.802,0.599]$。第二个维度从 0 变成了 0.599,说明它已经融合了另外两个位置的信息。
这就是 Self-Attention 最核心的事情:
根据内容计算 token 之间的关系,再把其他位置的信息按权重汇总到当前位置。
7. 为什么要除以 $\sqrt{d_k}$
如果 Query 和 Key 的每个分量都具有大致相同的方差,那么向量维度越大,点积结果的绝对值通常也越大。
例如 Softmax 的输入是:
[1, 2, 3]
得到的概率还比较平滑。如果变成:
[10, 20, 30]
Softmax 会非常接近 one-hot,最大值几乎占据全部概率。
这会让 Softmax 落入饱和区域,其他位置的梯度变得很小。除以 $\sqrt{d_k}$ 可以控制分数尺度,让训练更稳定。
缩放并不改变矩阵形状:
Q: [B, H, Lq, dk]
K^T: [B, H, dk, Lk]
score: [B, H, Lq, Lk]
weight: [B, H, Lq, Lk]
V: [B, H, Lk, dv]
output: [B, H, Lq, dv]
8. Self-Attention、Cross-Attention 和 Masked Self-Attention
三者使用同一个 Attention 公式,区别在于 Q、K、V 来自哪里,以及哪些位置允许被查看。
8.1 Encoder Self-Attention
Q 来自 Encoder 当前输入
K 来自 Encoder 当前输入
V 来自 Encoder 当前输入
输入序列中的每个 token 都可以查看其他非 PAD token。
8.2 Decoder Masked Self-Attention
Q、K、V 都来自 Decoder 当前输入
但第 $t$ 个位置不能看到 $t+1$ 及更后面的答案。
8.3 Encoder-Decoder Cross-Attention
Q 来自 Decoder
K 来自 Encoder 输出
V 来自 Encoder 输出
含义是:Decoder 在生成当前词时,拿自己的状态去检索源句中有用的信息。
| 类型 | Query 来源 | Key、Value 来源 | 典型作用 |
|---|---|---|---|
| Encoder Self-Attention | Encoder | Encoder | 理解输入内部关系 |
| Decoder Self-Attention | Decoder | Decoder | 理解已生成内容 |
| Cross-Attention | Decoder | Encoder | 从输入序列检索信息 |
9. 多头注意力为什么不是重复计算
单个 Attention 头只在一个投影空间中计算关系。Multi-Head Attention 会把特征拆成多个头,让不同头学习不同的关系。
$$
\operatorname{head}_i
=\operatorname{Attention}(QW_i^Q,KW_i^K,VW_i^V)
$$
$$
\operatorname{MultiHead}(Q,K,V)
=\operatorname{Concat}(\operatorname{head}_1,\ldots,\operatorname{head}_h)W^O
$$
例如:
d_model = 512
num_heads = 8
head_dim = 512 / 8 = 64
每个头处理 64 维特征,8 个头拼接后重新得到 512 维。
不同头可能学习到不同模式,例如:
- 某个头关注相邻词。
- 某个头关注主语与谓语。
- 某个头关注代词指代。
- 某个头关注句首或标点。
这些模式不是人工指定的,而是训练目标驱动下可能出现的分工。也不能简单认为“每个头一定对应一种人类可命名的语法关系”。
9.1 为什么 d_model 通常要能整除 num_heads
常见实现会把 $d_{model}$ 平均拆给各个头:
$$
d_{head}=\frac{d_{model}}{h}
$$
因此下面是合法的:
d_model=128, num_heads=4, head_dim=32
而下面无法平均拆分:
d_model=130, num_heads=8
PyTorch 的 nn.MultiheadAttention 也要求 embed_dim 能被 num_heads 整除。
10. 没有循环后,模型怎样知道顺序
Self-Attention 本身只根据内容计算关系。如果把输入 token 的顺序打乱,又没有任何位置信息,模型无法可靠区分:
猫 追 老鼠
老鼠 追 猫
所以需要向词向量加入位置编码。
原始 Transformer 使用正弦和余弦位置编码:
$$
PE(pos,2i)=\sin\left(\frac{pos}{10000^{2i/d_{model}}}\right)
$$
$$
PE(pos,2i+1)=\cos\left(\frac{pos}{10000^{2i/d_{model}}}\right)
$$
不同维度使用不同频率,使每个位置拥有不同模式。
10.1 手算一个位置编码
假设:
pos = 2
d_model = 4
那么:
$$
PE(2,0)=\sin(2)\approx0.9093
$$
$$
PE(2,1)=\cos(2)\approx-0.4161
$$
$$
PE(2,2)=\sin(2/100)\approx0.0200
$$
$$
PE(2,3)=\cos(2/100)\approx0.9998
$$
如果这个 token 的 Embedding 是:
$$
e=[0.3,-0.2,0.7,0.1]
$$
相加后:
$$
z=e+PE(2)
$$
$$
z\approx[1.2093,-0.6161,0.7200,1.0998]
$$
现在的表示同时包含 token 语义和“它位于第 2 个位置”的信息。
位置编码也可以设计为可学习参数。现代模型还会使用相对位置或旋转位置编码,但理解 Transformer 基础时,先掌握“内容表示需要注入位置信息”即可。
11. Transformer Encoder 层
一个经典 Encoder 层包含两个主要子层:
- Multi-Head Self-Attention。
- Position-wise Feed Forward Network。
每个子层外还包含残差连接与 LayerNorm。
原论文中的 Post-Norm 写法可以表示为:
$$
Z_1=\operatorname{LayerNorm}(X+\operatorname{MHA}(X))
$$
$$
Z_2=\operatorname{LayerNorm}(Z_1+\operatorname{FFN}(Z_1))
$$
多个 Encoder 层会堆叠起来:
输入 -> Encoder Layer 1 -> Encoder Layer 2 -> ... -> Encoder Layer N
第 1 层可能更多地组合局部信息,后续层继续在新的表示上建立更复杂关系。
12. 前馈网络 FFN 在做什么
Attention 负责不同 token 之间的信息交换,FFN 则对每个位置的特征进行非线性变换。
$$
\operatorname{FFN}(x)=\sigma(xW_1+b_1)W_2+b_2
$$
原始 Transformer 使用 ReLU:
$$
\operatorname{FFN}(x)=\max(0,xW_1+b_1)W_2+b_2
$$
很多后续模型使用 GELU 或门控形式,但基本作用仍然是增加每个位置内部的特征变换能力。
如果:
d_model = 512
dim_feedforward = 2048
则 FFN 的形状变化为:
[B, L, 512]
-> Linear(512, 2048)
-> 激活函数
-> Linear(2048, 512)
-> [B, L, 512]
因为 FFN 对每个位置独立使用同一组参数,所以称为 position-wise FFN。
13. 残差连接和 LayerNorm 为什么重要
13.1 残差连接
残差连接把子层输入直接加到输出:
$$
y=x+F(x)
$$
它提供了一条更直接的信息和梯度传播路径,使深层网络更容易优化。
注意相加要求形状一致,因此 Attention 和 FFN 子层最后都要回到 $d_{model}$ 维。
13.2 LayerNorm
LayerNorm 对单个样本、单个 token 的特征维进行归一化。
对向量 $x\in\mathbb{R}^{d}$:
$$
\mu=\frac{1}{d}\sum_{i=1}^{d}x_i
$$
$$
\sigma^2=\frac{1}{d}\sum_{i=1}^{d}(x_i-\mu)^2
$$
$$
\operatorname{LayerNorm}(x)
=\gamma\frac{x-\mu}{\sqrt{\sigma^2+\epsilon}}+\beta
$$
$\gamma$ 和 $\beta$ 是可学习参数。
13.3 Post-Norm 与 Pre-Norm
Post-Norm:
$$
y=\operatorname{LayerNorm}(x+F(x))
$$
Pre-Norm:
$$
y=x+F(\operatorname{LayerNorm}(x))
$$
原始论文使用 Post-Norm。PyTorch nn.Transformer 的 norm_first=False 表示 Post-Norm,设为 True 则使用 Pre-Norm。
14. Transformer Decoder 层
一个经典 Decoder 层包含三个主要子层:
- Masked Multi-Head Self-Attention。
- Encoder-Decoder Cross-Attention。
- Feed Forward Network。
每个子层同样配有残差连接和 LayerNorm。
目标前缀
-> Masked Self-Attention
-> Cross-Attention,查看 Encoder Memory
-> FFN
-> 预测下一个 token
Cross-Attention 中:
$$
Q=H_{decoder}W_Q
$$
$$
K=H_{encoder}W_K
$$
$$
V=H_{encoder}W_V
$$
因此 Decoder 的每个位置都能根据当前生成上下文,从源句不同位置提取信息。
15. Causal Mask:为什么训练时不能看未来
假设目标序列是:
<BOS> I like NLP <EOS>
训练时会错开一位:
Decoder 输入:<BOS> I like NLP
预测目标: I like NLP <EOS>
虽然 Decoder 输入可以一次性并行计算,但位置 I 不能看到后面的 like 和 NLP,否则相当于提前看答案。
长度为 4 的 Causal Mask 可以表示为:
$$
M=
\begin{bmatrix}
0&-\infty&-\infty&-\infty\
0&0&-\infty&-\infty\
0&0&0&-\infty\
0&0&0&0
\end{bmatrix}
$$
它会在 Softmax 之前加到分数矩阵:
$$
A=\operatorname{softmax}\left(\frac{QK^T}{\sqrt{d_k}}+M\right)
$$
因为 $e^{-\infty}=0$,被遮住的位置最终权重为 0。
15.1 为什么训练可以并行,推理仍要逐步生成
训练时,完整目标答案已知,因此可以把右移后的整个目标序列一起输入,并用 Mask 防止信息泄漏。
推理时,未来 token 尚不存在:
<BOS> -> I
<BOS> I -> like
<BOS> I like -> NLP
<BOS> I like NLP -> <EOS>
所以自回归 Decoder 的推理仍然通常逐步进行。
16. Padding Mask 与 Causal Mask 不要混淆
这两种 Mask 解决不同问题。
16.1 Padding Mask
Padding Mask 遮住为了对齐长度而添加的 <PAD>:
样本1:我 喜欢 NLP <PAD> <PAD>
样本2:我 非常 喜欢 学习 NLP
对应布尔 Mask:
样本1:[False, False, False, True, True]
样本2:[False, False, False, False, False]
在 PyTorch 的 Transformer API 中,布尔 padding mask 里的 True 通常表示该位置不参与注意力。
16.2 Causal Mask
Causal Mask 遮住当前位置之后的未来 token,通常是一个 $L\times L$ 的上三角矩阵。
| Mask | 典型形状 | 遮住什么 | 常见位置 |
|---|---|---|---|
| Padding Mask | [B, L] |
<PAD> |
Encoder 与 Decoder |
| Causal Mask | [L, L] |
未来 token | 自回归 Decoder |
实际训练 Decoder 时,两种 Mask 往往同时存在。
17. Transformer 的完整训练流程
以机器翻译为例:
1. 源句和目标句分词并转成 token id
2. 对 batch 内序列进行 padding
3. 构造 Decoder 输入和预测目标
4. 构造 src padding mask
5. 构造 tgt padding mask 与 causal mask
6. 前向传播得到每个位置的 logits
7. 使用交叉熵计算损失,忽略 PAD
8. 反向传播
9. 梯度裁剪
10. 优化器更新参数
模型输出形状为:
logits: [B, T, vocab_size]
target: [B, T]
计算损失前通常展平:
loss = loss_fn(
logits.reshape(-1, vocab_size),
target.reshape(-1),
)
其中:
loss_fn = nn.CrossEntropyLoss(ignore_index=PAD_ID)
CrossEntropyLoss 直接接收 logits,不要提前做 Softmax。
18. 从张量形状理解 PyTorch Transformer
当设置 batch_first=True 时,主要形状如下:
| 张量 | 形状 | 含义 |
|---|---|---|
src |
[B, S] |
源 token id |
tgt_in |
[B, T] |
右移后的目标 token id |
src_emb |
[B, S, E] |
源 Embedding 加位置编码 |
tgt_emb |
[B, T, E] |
目标 Embedding 加位置编码 |
tgt_mask |
[T, T] |
因果遮罩 |
src_padding_mask |
[B, S] |
源序列 PAD 遮罩 |
memory |
[B, S, E] |
Encoder 输出 |
decoder_output |
[B, T, E] |
Decoder 输出 |
logits |
[B, T, V] |
词表 logits |
其中:
- $S$ 是源序列长度。
- $T$ 是目标序列长度。
- $E=d_{model}$。
- $V$ 是词表大小。
19. PyTorch 的 nn.MultiheadAttention
先看最小例子:
import torch
from torch import nn
torch.manual_seed(0)
mha = nn.MultiheadAttention(
embed_dim=8,
num_heads=2,
batch_first=True,
)
x = torch.randn(2, 4, 8) # [batch, seq_len, embed_dim]
output, weights = mha(
query=x,
key=x,
value=x,
need_weights=True,
)
print(output.shape) # torch.Size([2, 4, 8])
print(weights.shape) # torch.Size([2, 4, 4])
默认情况下,weights 会对多个头取平均。如果希望保留每个头:
output, weights = mha(
x, x, x,
need_weights=True,
average_attn_weights=False,
)
print(weights.shape) # [2, 2, 4, 4]
这四个维度分别是:
[batch, num_heads, query_len, key_len]
20. 不依赖高级 API 实现缩放点积注意力
先按照公式直接实现:
import math
import torch
from torch import nn
def scaled_dot_product_attention(q, k, v, mask=None):
"""
q: [..., query_len, head_dim]
k: [..., key_len, head_dim]
v: [..., key_len, value_dim]
mask: 可广播到 [..., query_len, key_len],True 表示遮住
"""
scores = q @ k.transpose(-2, -1)
scores = scores / math.sqrt(q.size(-1))
if mask is not None:
scores = scores.masked_fill(mask, float("-inf"))
weights = torch.softmax(scores, dim=-1)
output = weights @ v
return output, weights
用前面的手算矩阵验证:
x = torch.tensor([
[1.0, 0.0],
[0.0, 1.0],
[1.0, 1.0],
])
output, weights = scaled_dot_product_attention(x, x, x)
print(weights)
print(output)
近似输出:
weights =
tensor([[0.4011, 0.1978, 0.4011],
[0.1978, 0.4011, 0.4011],
[0.2483, 0.2483, 0.5035]])
output =
tensor([[0.8022, 0.5989],
[0.5989, 0.8022],
[0.7517, 0.7517]])
这与手算结果一致。
21. 从零实现一个多头自注意力层
下面把线性投影、拆头、Attention、合并头完整连接起来:
class MultiHeadSelfAttention(nn.Module):
def __init__(self, d_model, num_heads, dropout=0.0):
super().__init__()
if d_model % num_heads != 0:
raise ValueError("d_model 必须能被 num_heads 整除")
self.num_heads = num_heads
self.head_dim = d_model // num_heads
self.q_proj = nn.Linear(d_model, d_model)
self.k_proj = nn.Linear(d_model, d_model)
self.v_proj = nn.Linear(d_model, d_model)
self.out_proj = nn.Linear(d_model, d_model)
self.dropout = nn.Dropout(dropout)
def split_heads(self, x):
batch, seq_len, _ = x.shape
x = x.reshape(
batch, seq_len, self.num_heads, self.head_dim
)
return x.transpose(1, 2)
def forward(self, x, mask=None):
q = self.split_heads(self.q_proj(x))
k = self.split_heads(self.k_proj(x))
v = self.split_heads(self.v_proj(x))
scores = q @ k.transpose(-2, -1)
scores = scores / math.sqrt(self.head_dim)
if mask is not None:
scores = scores.masked_fill(mask, float("-inf"))
weights = torch.softmax(scores, dim=-1)
weights = self.dropout(weights)
context = weights @ v
context = context.transpose(1, 2).contiguous()
batch, seq_len, _, _ = context.shape
context = context.reshape(batch, seq_len, -1)
return self.out_proj(context), weights
测试形状:
x = torch.randn(2, 5, 16)
attention = MultiHeadSelfAttention(d_model=16, num_heads=4)
output, weights = attention(x)
print(output.shape) # [2, 5, 16]
print(weights.shape) # [2, 4, 5, 5]
拆头前后形状变化是:
[B, L, d_model]
-> [B, L, num_heads, head_dim]
-> [B, num_heads, L, head_dim]
计算完成后再按相反顺序合并。
22. PyTorch 实现位置编码
import math
import torch
from torch import nn
class PositionalEncoding(nn.Module):
def __init__(self, d_model, max_len=5000, dropout=0.1):
super().__init__()
self.dropout = nn.Dropout(dropout)
position = torch.arange(max_len).unsqueeze(1)
div_term = torch.exp(
torch.arange(0, d_model, 2)
* (-math.log(10000.0) / d_model)
)
pe = torch.zeros(max_len, d_model)
pe[:, 0::2] = torch.sin(position * div_term)
pe[:, 1::2] = torch.cos(position * div_term)
self.register_buffer("pe", pe.unsqueeze(0))
def forward(self, x):
x = x + self.pe[:, :x.size(1)]
return self.dropout(x)
22.1 为什么使用 register_buffer
位置编码不是通过梯度学习的参数,所以不应该使用 nn.Parameter。但它需要:
- 跟随模型一起移动到 CPU 或 GPU。
- 跟随
state_dict保存。 - 不出现在优化器的可训练参数中。
register_buffer 正好满足这些要求。
23. 完整例子:Transformer 学习反转数字序列
为了不依赖外部数据集,下面让 Transformer 学习一个明确任务:
输入:[3, 7, 5, 9]
目标:[9, 5, 7, 3, EOS]
这个任务虽然简单,但完整包含:
- 变长序列。
- Padding Mask。
- Causal Mask。
- Encoder-Decoder Attention。
- 目标序列右移。
- 交叉熵训练。
- 自回归推理。
23.1 定义数据
import math
import random
import torch
from torch import nn
from torch.nn.utils.rnn import pad_sequence
torch.manual_seed(0)
random.seed(0)
PAD = 0
BOS = 1
EOS = 2
VOCAB_SIZE = 13 # 0、1、2是特殊token,3到12是普通数字token
def make_batch(batch_size=64, min_len=3, max_len=7):
src_list = []
tgt_list = []
for _ in range(batch_size):
length = random.randint(min_len, max_len)
src = torch.randint(3, VOCAB_SIZE, (length,))
tgt = torch.cat([
torch.tensor([BOS]),
src.flip(0),
torch.tensor([EOS]),
])
src_list.append(src)
tgt_list.append(tgt)
src = pad_sequence(
src_list, batch_first=True, padding_value=PAD
)
tgt = pad_sequence(
tgt_list, batch_first=True, padding_value=PAD
)
tgt_in = tgt[:, :-1]
tgt_out = tgt[:, 1:]
return src, tgt_in, tgt_out
tgt_in 与 tgt_out 错开一位:
tgt_in: BOS 9 5 7 3
tgt_out: 9 5 7 3 EOS
23.2 定义 Transformer 模型
class TinyTransformer(nn.Module):
def __init__(
self,
vocab_size,
d_model=64,
num_heads=4,
num_layers=2,
dim_feedforward=128,
dropout=0.1,
):
super().__init__()
self.d_model = d_model
self.src_embedding = nn.Embedding(
vocab_size, d_model, padding_idx=PAD
)
self.tgt_embedding = nn.Embedding(
vocab_size, d_model, padding_idx=PAD
)
self.position = PositionalEncoding(
d_model=d_model,
max_len=100,
dropout=dropout,
)
self.transformer = nn.Transformer(
d_model=d_model,
nhead=num_heads,
num_encoder_layers=num_layers,
num_decoder_layers=num_layers,
dim_feedforward=dim_feedforward,
dropout=dropout,
batch_first=True,
)
self.output = nn.Linear(d_model, vocab_size)
@staticmethod
def causal_mask(size, device):
return torch.triu(
torch.ones(size, size, dtype=torch.bool, device=device),
diagonal=1,
)
def forward(self, src, tgt_in):
src_padding_mask = src.eq(PAD)
tgt_padding_mask = tgt_in.eq(PAD)
tgt_mask = self.causal_mask(tgt_in.size(1), tgt_in.device)
scale = math.sqrt(self.d_model)
src_emb = self.position(self.src_embedding(src) * scale)
tgt_emb = self.position(self.tgt_embedding(tgt_in) * scale)
hidden = self.transformer(
src=src_emb,
tgt=tgt_emb,
tgt_mask=tgt_mask,
src_key_padding_mask=src_padding_mask,
tgt_key_padding_mask=tgt_padding_mask,
memory_key_padding_mask=src_padding_mask,
)
return self.output(hidden)
memory_key_padding_mask 与 src_key_padding_mask 通常使用同一个源序列 Mask:
- Encoder Self-Attention 不应查看源
<PAD>。 - Decoder Cross-Attention 也不应从源
<PAD>中取信息。
23.3 训练模型
device = torch.device(
"cuda" if torch.cuda.is_available() else "cpu"
)
model = TinyTransformer(VOCAB_SIZE).to(device)
optimizer = torch.optim.AdamW(
model.parameters(), lr=1e-3, weight_decay=0.01
)
loss_fn = nn.CrossEntropyLoss(ignore_index=PAD)
model.train()
for step in range(1000):
src, tgt_in, tgt_out = make_batch()
src = src.to(device)
tgt_in = tgt_in.to(device)
tgt_out = tgt_out.to(device)
logits = model(src, tgt_in)
loss = loss_fn(
logits.reshape(-1, VOCAB_SIZE),
tgt_out.reshape(-1),
)
optimizer.zero_grad()
loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
optimizer.step()
if (step + 1) % 100 == 0:
print(f"step={step + 1}, loss={loss.item():.4f}")
这是一个小型教学任务,真实训练还需要:
- 验证集与测试集。
- 学习率 warmup 和衰减。
- 更可靠的数据加载流程。
- 保存与恢复 checkpoint。
- 根据验证指标进行早停或选模。
23.4 自回归推理
@torch.no_grad()
def greedy_decode(model, src, max_len=10):
model.eval()
src = src.to(device)
generated = torch.full(
(src.size(0), 1),
BOS,
dtype=torch.long,
device=device,
)
finished = torch.zeros(
src.size(0), dtype=torch.bool, device=device
)
for _ in range(max_len):
logits = model(src, generated)
next_token = logits[:, -1].argmax(dim=-1)
next_token = torch.where(
finished,
torch.full_like(next_token, PAD),
next_token,
)
generated = torch.cat(
[generated, next_token.unsqueeze(1)], dim=1
)
finished = finished | next_token.eq(EOS)
if finished.all():
break
return generated[:, 1:]
测试:
test_src = torch.tensor([
[3, 7, 5, 9, PAD],
[4, 8, 6, 3, 7],
])
result = greedy_decode(model, test_src, max_len=7)
print(result.cpu())
训练充分后,结果应接近:
tensor([[9, 5, 7, 3, 2, 0],
[7, 3, 6, 8, 4, 2]])
其中 2 是 <EOS>,0 是已经结束样本后面补出的 <PAD>。
24. nn.Transformer 中的重要参数
nn.Transformer(
d_model=512,
nhead=8,
num_encoder_layers=6,
num_decoder_layers=6,
dim_feedforward=2048,
dropout=0.1,
activation="relu",
batch_first=False,
norm_first=False,
)
| 参数 | 含义 | 常见错误 |
|---|---|---|
d_model |
token 表示维度 | 与 Embedding 维度不一致 |
nhead |
注意力头数 | d_model 不能被它整除 |
num_encoder_layers |
Encoder 层数 | 层数增加后显存和计算量上升 |
num_decoder_layers |
Decoder 层数 | Encoder-only 任务不需要 Decoder |
dim_feedforward |
FFN 中间维度 | 误以为是词表大小 |
dropout |
Dropout 概率 | 推理时忘记 model.eval() |
batch_first |
是否使用 [B,L,E] |
与实际输入维度顺序不一致 |
norm_first |
是否使用 Pre-Norm | 误以为会改变输出形状 |
PyTorch 的 nn.Transformer 更接近原始架构的教学参考实现。理解基础组件后,再学习不同模型库中的预训练 Transformer 会容易很多。
25. Transformer 的计算复杂度
标准 Self-Attention 的分数矩阵形状是:
$$
[L,L]
$$
因此与序列长度相关的主要计算和内存代价通常为:
$$
O(L^2)
$$
如果序列长度从 512 增加到 1024,注意力矩阵元素数量会从:
$$
512^2=262144
$$
增加到:
$$
1024^2=1048576
$$
也就是约 4 倍。
这说明 Transformer 的并行能力很强,但标准全局 Attention 在超长序列上也会付出明显代价。
需要同时理解两句话:
- Transformer 训练时比递归模型更容易并行。
- 标准 Self-Attention 对序列长度具有二次复杂度。
它们并不矛盾。
26. Transformer 与 BERT、GPT 的关系
Transformer 是基础积木,BERT、GPT 等是基于这套积木形成的模型家族。
26.1 Encoder-only
只堆叠 Transformer Encoder,可以让每个 token 双向查看上下文,适合理解类任务。
完整输入 -> 双向 Self-Attention -> 上下文表示
26.2 Decoder-only
只堆叠带 Causal Mask 的 Transformer Decoder 风格模块,适合预测下一个 token。
历史 token -> Causal Self-Attention -> 下一个 token
26.3 Encoder-Decoder
Encoder 理解输入,Decoder 通过 Cross-Attention 读取 Encoder 输出,适合条件生成。
输入序列 -> Encoder -> Memory
目标前缀 -> Decoder + Memory -> 下一个 token
“Transformer”描述的是网络结构,而“预训练语言模型”还会涉及:
- 预训练目标。
- 训练数据。
- Tokenizer。
- 优化策略。
- 微调或对齐方式。
不要把这些概念混为一层。
27. 常见错误
27.1 忘记位置编码
如果没有循环、卷积或位置编码,Self-Attention 本身无法可靠表达 token 顺序。
27.2 把 Q、K、V 理解成三个固定输入
Self-Attention 中它们通常由同一个输入经过三组可学习线性投影得到。
27.3 Softmax 维度写错
应该沿 Key 所在的最后一个维度归一化:
weights = torch.softmax(scores, dim=-1)
这样每个 Query 对所有 Key 的权重之和为 1。
27.4 忘记除以 $\sqrt{d_k}$
这可能让点积分数过大,使 Softmax 过早饱和。
27.5 Decoder 训练时没有 Causal Mask
模型会直接看到未来答案,训练损失可能很好看,但推理时无法正常工作。
27.6 只设置 src_key_padding_mask
Encoder 不看源 PAD 还不够,Decoder 的 Cross-Attention 也要通过 memory_key_padding_mask 忽略源 PAD。
27.7 Mask 中 True 的语义理解反了
在 PyTorch nn.Transformer 常用的布尔 Mask 中,True 表示该位置被遮住。不同底层 API 的 Mask 约定可能不同,使用前应检查对应文档。
27.8 d_model 不能被头数整除
每个头需要获得相同的 head_dim,因此创建层时会直接报错。
27.9 对 logits 先做 Softmax
错误:
loss = loss_fn(torch.softmax(logits, dim=-1), target)
正确:
loss = loss_fn(logits, target)
27.10 推理时忘记 eval() 和 no_grad()
推理应关闭 Dropout 并避免构建梯度图:
model.eval()
with torch.no_grad():
output = model(src, tgt)
27.11 把 Attention 权重当成完整解释
Attention 权重能帮助观察模型的信息聚合方式,但不能自动等价为严格的因果解释或完整决策依据。
28. 推荐的学习与调试顺序
学习 Transformer 时,可以按以下顺序逐步验证:
1. 手算单头 Attention
2. 用 PyTorch 复现手算结果
3. 理解 Q、K、V 的形状
4. 实现拆头与合并头
5. 加入位置编码
6. 构造 Padding Mask
7. 构造 Causal Mask
8. 搭建单个 Encoder 层
9. 搭建 Encoder-Decoder
10. 用小任务验证完整训练和推理
调试时优先打印:
print("src:", src.shape)
print("tgt_in:", tgt_in.shape)
print("src_padding_mask:", src_padding_mask.shape)
print("tgt_mask:", tgt_mask.shape)
print("logits:", logits.shape)
对于 Attention 权重,检查:
print(weights.sum(dim=-1))
在没有把一整行全部遮住的情况下,结果应接近 1。
29. 总结
Transformer 的主线可以浓缩为:
token id
-> Embedding + Positional Encoding
-> Q、K、V 线性投影
-> 缩放点积 Attention
-> Multi-Head 拼接
-> 残差连接 + LayerNorm
-> FFN
-> 堆叠多个 Encoder / Decoder 层
-> 输出 token logits
需要真正掌握的核心包括:
- Attention 根据 Query 与 Key 的匹配程度,对 Value 加权求和。
- Self-Attention 的 Q、K、V 来自同一序列的不同投影。
- Multi-Head Attention 在多个表示子空间中并行建模关系。
- 没有递归后必须显式注入位置信息。
- Encoder 使用双向 Self-Attention 理解输入。
- 自回归 Decoder 必须用 Causal Mask 遮住未来。
- Padding Mask 与 Causal Mask 解决的是不同问题。
- Attention 负责 token 间交互,FFN 负责位置内特征变换。
- 残差连接和 LayerNorm 帮助深层网络稳定训练。
- 标准 Self-Attention 容易并行,但对序列长度具有二次复杂度。
当你能独立写出下面这行公式,并说清楚每个矩阵的来源、形状和作用时,就已经抓住了 Transformer 最重要的核心:
$$
\operatorname{Attention}(Q,K,V)
=\operatorname{softmax}\left(\frac{QK^T}{\sqrt{d_k}}\right)V
$$
参考资料
- Ashish Vaswani et al., Attention Is All You Need: https://arxiv.org/abs/1706.03762
- PyTorch Documentation,
torch.nn.Transformer: https://docs.pytorch.org/docs/stable/generated/torch.nn.Transformer.html - PyTorch Documentation,
torch.nn.MultiheadAttention: https://docs.pytorch.org/docs/stable/generated/torch.nn.MultiheadAttention.html - PyTorch Documentation,
torch.nn.TransformerEncoderLayer: https://docs.pytorch.org/docs/stable/generated/torch.nn.TransformerEncoderLayer.html
