上一篇 Seq2Seq 文章中,我们使用 GRU 构造了 Encoder-Decoder,并通过 Attention 缓解“把整句话压缩成一个向量”的信息瓶颈。

不过,RNN、LSTM 和 GRU 还有一个很难绕开的特点:

必须按照时间顺序逐步计算,后一个位置需要等待前一个位置的隐藏状态。

假设一句话有 100 个 token,RNN 需要从第 1 个 token 一直算到第 100 个 token。即使 GPU 很擅长矩阵并行,也不能完全消除这种前后依赖。

Transformer 改变了处理序列的方式:

不再使用循环逐步传递信息,而是让每个 token 直接通过 Attention 查看序列中的其他 token。

这样一来,训练时可以并行处理整段序列,也更容易建模相距很远的词之间的关系。

Transformer整体结构

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:

  1. 已经生成的目标 token。
  2. 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 更高权重。

Self-Attention中的QKV

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矩阵手算

这就是 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 维。

Multi-Head Attention计算过程

不同头可能学习到不同模式,例如:

  • 某个头关注相邻词。
  • 某个头关注主语与谓语。
  • 某个头关注代词指代。
  • 某个头关注句首或标点。

这些模式不是人工指定的,而是训练目标驱动下可能出现的分工。也不能简单认为“每个头一定对应一种人类可命名的语法关系”。

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)
$$

不同维度使用不同频率,使每个位置拥有不同模式。

位置编码如何加入Embedding

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 层包含两个主要子层:

  1. Multi-Head Self-Attention。
  2. Position-wise Feed Forward Network。

每个子层外还包含残差连接与 LayerNorm。

Transformer Encoder层

原论文中的 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.Transformernorm_first=False 表示 Post-Norm,设为 True 则使用 Pre-Norm。

14. Transformer Decoder 层

一个经典 Decoder 层包含三个主要子层:

  1. Masked Multi-Head Self-Attention。
  2. Encoder-Decoder Cross-Attention。
  3. 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 不能看到后面的 likeNLP,否则相当于提前看答案。

长度为 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。

Decoder因果遮罩

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 时,主要形状如下:

PyTorch Transformer张量形状

张量 形状 含义
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_intgt_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_masksrc_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
$$

参考资料

  1. Ashish Vaswani et al., Attention Is All You Need: https://arxiv.org/abs/1706.03762
  2. PyTorch Documentation, torch.nn.Transformer: https://docs.pytorch.org/docs/stable/generated/torch.nn.Transformer.html
  3. PyTorch Documentation, torch.nn.MultiheadAttention: https://docs.pytorch.org/docs/stable/generated/torch.nn.MultiheadAttention.html
  4. PyTorch Documentation, torch.nn.TransformerEncoderLayer: https://docs.pytorch.org/docs/stable/generated/torch.nn.TransformerEncoderLayer.html