我很早就想写一篇“用 PyTorch 从零实现一个 Transformer”的经验总结,原因很简单:几乎每个入门深度学习的人都会卡在“看得懂图、写不出代码”这一步。Transformer 的架构图在网上随处可见,Encoder、Decoder、Multi-Head Attention 这些名词大家都能脱口而出,可真要关掉教程,自己从空目录开始写一个能训练、能推理的模型,很多人会在一开始就懵掉——位置编码怎么加?Mask 的形状到底是什么?Decoder 的输入往右移一位是什么意思?这些细节,看文章的时候都觉得“懂了”,动手才会发现全是坑。
这篇文章我尽量按实际项目的推进顺序来写,记录我从任务设计、数据处理到核心模块实现、训练调试的完整过程。代码全部只用 PyTorch 的基础组件,不碰 nn.Transformer、nn.MultiheadAttention 这些高层封装。适合两种人看:一种是已经看过 Transformer 图解、想动手验证理解的初学者,另一种是面试前需要快速把整个模型串起来的老手。我会把每个关键选择背后的原因也一并讲清楚,这样你不仅能照抄代码,还能知道每一步到底在做什么。
1. 为什么非得“手撕”一遍 Transformer
1.1 纸上得来终觉浅,注意力机制尤其如此
Transformer 的核心是自注意力机制,这个机制本身的公式很简单:Attention(Q, K, V) = softmax(QK^T / sqrt(d_k)) V。但我观察到一个现象:很多人能把公式默写出来,却回答不了几个特别基础的问题——为什么除的是 sqrt(d_k) 而不是 d_k?为什么 Q、K、V 要拆成多个头再拼回去?训练的时候 padding 位置和未来位置是怎么被“屏蔽”掉的?这些问题靠看是看不会的,只有自己实现一遍,才会在写错、调试、看 loss 变化的过程中真正留下肌肉记忆。
我这次实现还有一个原则:能自己写的地方绝不调现成 API。nn.MultiheadAttention 确实方便,但正因为方便,它把很多细节掩盖掉了。比如输入张量是 [batch, seq_len, embed_dim],进到这个模块内部后会被 reshape 成 [seq_len, batch, embed_dim],这种维度变化如果不自己处理一次,很难对 PyTorch 中 attention 的 tensor 流转建立直观感受。手写虽然代码量多一点,但换来的是对整个数据流清晰到每个维度的掌控感。
1.2 直接调用 nn.Transformer 差在哪
PyTorch 官方提供了 nn.Transformer,它可以一行代码搭出 Transformer,但这恰恰是问题所在。你用它完成一个简单任务后,除了会填参数,很难说真正理解了这个模型。参数填错了,报错信息又长又绕,你根本不知道问题出在 Encoder 还是 Decoder,更别提定位到某个维度 mismatch 的具体原因。
从我带过项目的经验来看,能徒手写出 Transformer 的人,遇到模型不收敛、loss 变成 NaN、推理结果全是一个 token 这类问题时,定位速度会快得多。因为他脑子里有一张完整的计算图,知道数据在每一层之后长什么样。这就好比老司机能通过发动机声音判断大致故障,而只会踩油门的人只能原地等救援。
1.3 这个项目适合谁、最终能带走什么
这个项目从零开始包括任务设计、数据构造、模型实现、训练评估,整套流程走下来大约半天到一天时间。不需要 GPU,CPU 上几分钟就能完成训练收敛,关键是能跑通、能验证、能观察注意力权重的变化。如果你是第一次接触 Transformer,建议把代码从头到尾敲一遍,不要直接复制粘贴。敲代码这个动作本身,就是逼迫大脑处理每个细节的过程。
需要模型API调用? 免费领10W Token,多模型网关一键接入 Claude、DeepSeek 等主流模型。
2. 任务选型:用“序列累加”验证模型
2.1 为什么选累加而不是翻译或文本生成
从零实现一个模型,最怕的是选错验证任务。一开始我考虑过用英文到中文的翻译,但翻译任务对数据量、词表、训练时间的要求太高,在 CPU 上等一次完整训练可能要按天计算,非常不利于快速迭代调 bug。后来我换成了语言模型式的人物对话生成,又发现评估比较主观,模型输出一句不通顺的话你很难判断是模型没训练好,还是实现本身有 bug。
最后我选择了一个特别简单的合成任务:给定一串数字,模型输出每个位置之前的累加和。举个例子,输入 [3, 1, 4, 1, 5],输出 [3, 4, 8, 9, 14]。这个任务有几个非常适合验证模型实现正确性的特点:
- 输入输出都是离散 token,天然适配分类式的交叉熵损失。
- 输出长度和输入长度完全相同,不需要复杂的长度控制逻辑。
- 任务本身要求模型“记住”前缀和,这意味着 Decoder 必须在每个时间步都 attend 到前序所有位置,能有效检验 attention 机制是否真的在起作用。
- 数据可以完全合成,批量生成只需要几行代码,不用折腾下载数据集。
- 训练速度快,CPU 上几分钟就能看到收敛。
如果你是一个有经验的工程师,可能会觉得这个任务“太简单了,一点都不像真实项目”。但请记住,我们的目标不是刷榜,而是用最短路径验证手写 Transformer 的正确性。模型能在这个任务上收敛到接近 100% 准确率,说明位置编码、多头注意力、mask 这些模块大概率没错;如果这些模块有问题,再复杂的任务也无法成功。
2.2 输入输出与词表设计
词表设计要同时考虑输入和输出。输入是 0~9 的数字,但输出是累加和,范围会超过 9。比如最长序列长度我设为 12,输入最大为 9,则最大的累加和是 12 × 9 = 108。所以输出词表至少要覆盖 0 到 108 这些数字。
为了简单,我不区分输入词表和输出词表,统一用一个包含 0~127 的数字 token 空间,再加上几个特殊 token 就够了。具体分配如下:
| token 含义 | 索引范围 |
|---|---|
| 数字 0~127 | 0~127 |
| BOS(begin of sequence) | 128 |
| PAD | 129 |
最大累加和是 108,低于 127,因此用这个词表不会出现越界。vocab_size = 130。输入数字只会用到 0~9 这部分,但词表留出更大的空间可以让输出 token 的范围覆盖所有可能的累加结果,实际训练中也更省心,省去了判断输出到底会不会越界的麻烦。
2.3 模型配置与参数量估算
模型规模不能太大,否则 CPU 上训练会让人失去耐心;也不能太小,否则 attention 机制发挥不出效果。我选择的配置如下:
| 超参数 | 取值 |
|---|---|
| d_model(embedding 维度) | 64 |
| num_heads(注意力头数) | 4 |
| d_ff(前馈层中间维度) | 128 |
| num_layers(编码器和解码器层数) | 2 |
| dropout | 0.1 |
| max_len(位置编码最大长度) | 64 |
| vocab_size | 130 |
关于头数的选择,我特意选了 4 而不是常见的 8。因为 d_model 是 64,如果分成 8 个头,每个头的维度只剩 8,信息容量有点紧张;4 个头时每个头有 16 维,相对合理。这个配置下,模型的参数量大概在 30 万左右,非常轻量。训练 30 个 epoch,CPU 上只需要两三分钟。
3. 数据准备:为 Transformer 造一份合适的数据集
3.1 合成数据的生成逻辑
数据准备是很多人容易轻视的环节,但它对训练效果的影响极其直接。我用 PyTorch 的 Dataset 类来封装数据生成逻辑,核心思路是:每次随机生成长度在 4 到 12 之间的数字序列,每个数字取值 0 到 9,然后计算前缀和作为目标序列。
python复制import random
import torch
from torch.utils.data import Dataset
PAD_TOKEN = 129
BOS_TOKEN = 128
MAX_VAL = 127
class CumSumDataset(Dataset):
def __init__(self, num_samples=10000, min_len=4, max_len=12, seed=42):
super().__init__()
self.num_samples = num_samples
self.min_len = min_len
self.max_len = max_len
random.seed(seed)
def __len__(self):
return self.num_samples
def __getitem__(self, idx):
length = random.randint(self.min_len, self.max_len)
src = [random.randint(0, 9) for _ in range(length)]
# 前缀和即为目标输出
tgt = []
s = 0
for x in src:
s += x
tgt.append(s)
return torch.tensor(src, dtype=torch.long), torch.tensor(tgt, dtype=torch.long)
这里有几个设计点需要说明。第一,序列长度是随机的,这样模型在推理时必须真正学会“处理任意长度”,而不是死记硬背某个固定长度下的模式。第二,目标值是前缀和而不是别的什么,这要求模型在每个解码位置都能获取“到目前为止所有 token 的信息”,如果 attention 的 mask 写错了,模型很难收敛到高准确率,任务就起到了验证作用。
3.2 shift-right:Decoder 输入是怎么构造的
Transformer 的 Decoder 是一个自回归模型,训练时常用 Teacher Forcing 技巧:目标序列的每个位置作为输入,预测下一个位置。这里有一个关键操作叫 shift-right,也就是把目标序列整体向右移一位,在最前面插入 BOS token,然后丢到 Decoder 里作为输入。
举个例子,某个样本的目标输出是 [3, 4, 8, 9]。Decoder 的输入应该是 [BOS, 3, 4, 8],而模型要预测的目标则是 [3, 4, 8, 9]。这样一来,Decoder 在预测第 2 个位置 4 的时候,输入已经看到了 [BOS, 3],符合自回归“只能看到过去”的约束。
在代码中,这个操作可以借助一个简单的 collate 函数实现,同时完成 batch 内 padding 和 mask 的构造。我写的 collate 逻辑如下:
python复制def collate_fn(batch):
src_list, tgt_list = zip(*batch)
src_lens = [len(s) for s in src_list]
max_len = max(src_lens)
batch_size = len(batch)
src_padded = torch.full((batch_size, max_len), PAD_TOKEN, dtype=torch.long)
tgt_padded = torch.full((batch_size, max_len), PAD_TOKEN, dtype=torch.long)
for i, (src, tgt) in enumerate(batch):
src_padded[i, :len(src)] = src
tgt_padded[i, :len(tgt)] = tgt
# decoder 输入:右移一位,开头插入 BOS
decoder_input = torch.full((batch_size, max_len), PAD_TOKEN, dtype=torch.long)
decoder_input[:, 0] = BOS_TOKEN
if max_len > 1:
decoder_input[:, 1:] = tgt_padded[:, :-1]
# padding mask:src 中 PAD 位置为 False
src_padding_mask = (src_padded != PAD_TOKEN).unsqueeze(1).unsqueeze(2) # [B, 1, 1, L]
# 目标序列 mask:padding 位置为 True(用于 CrossEntropyLoss ignore)
tgt_padding_mask = (tgt_padded != PAD_TOKEN)
return {
"src": src_padded,
"tgt": tgt_padded,
"decoder_input": decoder_input,
"src_padding_mask": src_padding_mask,
"tgt_padding_mask": tgt_padding_mask,
}
注意,在 batch 内长度不齐时,短序列用 PAD_TOKEN 补齐。PAD 位置的预测结果应该在损失函数里被忽略,否则模型会浪费大量参数在“学习输出 PAD”上,导致真实位置的预测质量下降。后面训练循环里面,我会用 CrossEntropyLoss(ignore_index=PAD_TOKEN) 来处理。
3.3 mask 的构造:Padding Mask 与 Target Mask
Mask 是 Transformer 实现中最容易出错的地方。Encoder 的注意力只需要屏蔽 padding 位置,也就是让模型在计算注意力权重时,忽略所有 PAD 位置。这个 mask 的形状通常是 [B, 1, 1, L],在多头注意力内部会广播成 [B, num_heads, L, L],其中每个 [i, j] 位置表示第 i 个 query 是否能看到第 j 个 key。
Decoder 里除了 padding mask,还需要一个因果 mask(causal mask),用来屏蔽未来位置。因为 Decoder 在训练时可以看到完整的目标序列,如果不加这个 mask,模型在预测第 2 个 token 的时候就能“偷看”第 10 个 token,这不合理。因果 mask 是一个下三角矩阵,形状为 [L, L],第 i 行第 j 列表示第 i 个 query 是否能看到第 j 个 key,当 j > i 时为 False。
我在代码中用一个工具函数生成下三角 bool mask,然后在 Decoder 层内部同时应用 padding mask 和因果 mask。初学者最容易忽略的是把 mask 的维度对齐到 scores 的 [B, heads, L_q, L_k],如果不对齐,masked_fill 时广播规则会报各种奇怪的错误。
4. 从零实现 Transformer 核心模块
4.1 位置编码:给序列注入顺序信息
注意力机制最大的问题是“无序”。如果不加位置编码,模型看到的 [3, 1, 4] 和 [4, 1, 3] 会被当成同样的输入。Transformer 的解决方法是给每个位置加上一个固定的向量,让模型能区分不同位置。原始论文用的是不同频率的正余弦函数:
python复制import math
import torch.nn as nn
class PositionalEncoding(nn.Module):
def __init__(self, d_model, max_len=64, dropout=0.1):
super().__init__()
self.dropout = nn.Dropout(dropout)
pe = torch.zeros(max_len, d_model)
position = torch.arange(0, max_len, dtype=torch.float).unsqueeze(1)
div_term = torch.exp(
torch.arange(0, d_model, 2).float() * (-math.log(10000.0) / d_model)
)
pe[:, 0::2] = torch.sin(position * div_term)
pe[:, 1::2] = torch.cos(position * div_term)
pe = pe.unsqueeze(0) # shape: [1, max_len, d_model]
self.register_buffer("pe", pe)
def forward(self, x):
# x: [B, L, d_model]
x = x + self.pe[:, :x.size(1)]
return self.dropout(x)
为什么用不同频率的正余弦?因为对于任意固定的偏移 k,PE(pos+k) 可以表示为 PE(pos) 的线性组合,这有助于模型学习位置之间的相对关系。简单来说,这种编码方式既能让模型感知到“这是第几个位置”,也能让模型在理论上更容易捕捉“两个位置相差多远”的信息。
这里有两个工程细节。第一,我用 register_buffer 而不是直接赋值给某个普通属性,这样位置编码会随着模型一起移动到 GPU 或 CPU,不会出现 device mismatch。第二,max_len 要设置得比训练时见过的最大序列长度稍大,推理时如果输入超过这个长度,self.pe[:, :x.size(1)] 会直接报错,所以我把 max_len 设成了 64,远大于训练时最长序列 12,留出余量。
4.2 多头注意力:最核心的 60 行
多头注意力是整个 Transformer 的心脏。它做的事情可以分成三步:把输入通过三个线性层映射成 Q、K、V;把 Q、K、V 拆成多个头,在每个头上分别计算缩放点积注意力;把所有头的结果拼起来,再通过一个输出线性层。
这里我直接贴出手写实现,并加上必要的注释:
python复制import torch
import torch.nn as nn
import math
class MultiHeadAttention(nn.Module):
def __init__(self, d_model, num_heads, dropout=0.1):
super().__init__()
assert d_model % num_heads == 0
self.d_model = d_model
self.num_heads = num_heads
self.head_dim = d_model // num_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.out_proj = nn.Linear(d_model, d_model)
self.dropout = nn.Dropout(dropout)
def forward(self, query, key, value, mask=None):
batch_size = query.size(0)
L_q = query.size(1)
L_k = key.size(1)
# 1. 线性投影,再拆成多头
# Q: [B, L_q, d_model] -> [B, L_q, num_heads, head_dim] -> [B, num_heads, L_q, head_dim]
Q = self.w_q(query).view(batch_size, L_q, self.num_heads, self.head_dim).transpose(1, 2)
K = self.w_k(key).view(batch_size, L_k, self.num_heads, self.head_dim).transpose(1, 2)
V = self.w_v(value).view(batch_size, L_k, self.num_heads, self.head_dim).transpose(1, 2)
# 2. 缩放点积注意力
scores = torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(self.head_dim)
if mask is not None:
# mask 需要扩展成 [B, num_heads, L_q, L_k]
mask = mask.unsqueeze(1) # [B, 1, L_q, L_k]
scores = scores.masked_fill(mask == 0, -1e9)
attn_weights = torch.softmax(scores, dim=-1)
attn_weights = self.dropout(attn_weights)
out = torch.matmul(attn_weights, V)
# 3. 拼接所有头,通过输出投影
out = out.transpose(1, 2).contiguous().view(batch_size, L_q, self.d_model)
out = self.out_proj(out)
return out
这里最关键的一个除法是 sqrt(self.head_dim),也就是 sqrt(d_k)。为什么要除这个数?因为 Q 和 K 的每个维度都是均值为 0、方差为 1 的随机变量,点积之后方差的量级会变成 d_k。如果不缩放,当 d_k 很大时,scores 的值会很大,softmax 会进入饱和区,梯度变得非常小。缩放之后,点积的方差回到 1,softmax 不至于饱和。这个细节如果不实现一遍,很难真正理解它为什么这么设计。
我刚才 mask 的处理方式是 mask.unsqueeze(1),这个要注意广播。因为在很多调用场景下,传入的 mask 是 [B, 1, L_q, L_k] 或者 [L_q, L_k],通常我会在更外层统一处理成 [B, 1, L_q, L_k],这样在函数内部再 unsqueeze(1) 变成 [B, 1, L_q, L_k] 会维度太夸张。不过这个函数内部已经 unsqueeze(1) 了,所以外部传入的 mask 一般是 [B, L_q, L_k](没有那个 1)。实际我在写代码时,不同模块对 mask 的形状要求不同,是最容易搞混的地方,后续在常见问题部分我会单独展开讲。
4.3 前馈网络、残差与 LayerNorm
每个 Encoder 或 Decoder 层除了注意力子层,还有一个逐位置的前馈网络(Position-wise Feed-Forward Network)。它是对每个位置独立做两次线性变换加一次 ReLU,中间维度 d_ff 通常比 d_model 大。我的配置里 `d_ff=128,而 d_model=64,就是为了让模型在这个子层有一定的非线性拟合能力。实现非常简单:
python复制class FeedForward(nn.Module):
def __init__(self, d_model, d_ff, dropout=0.1):
super().__init__()
self.net = nn.Sequential(
nn.Linear(d_model, d_ff),
nn.ReLU(),
nn.Dropout(dropout),
nn.Linear(d_ff, d_model),
)
def forward(self, x):
return self.net(x)
残差连接和 LayerNorm 是训练深层的保障。我实现中采用原始论文的 Post-LN 结构,也就是每个子层先做注意力或前馈,然后加上残差,最后做 LayerNorm。用代码表示就是 x = norm(x + sublayer(x))。这种结构实现简单,但训练时对学习率比较敏感,这也是我们后面要用 warmup 学习率的原因之一。
有一种变体叫 Pre-LN,是 x = x + sublayer(norm(x)),它训练更稳定,但原始论文用的是 Post-LN,为了尊重原版,这里采用 Post-LN。
4.4 组装 Encoder 与 Decoder
有了上面的基础模块,就可以组装完整的 Encoder Layer 和 Decoder Layer。
Encoder Layer 包含一个多头自注意力子层和一个前馈子层,每个子层后都跟残差和 LayerNorm。Decoder Layer 比 Encoder Layer 多一个 cross-attention 子层,它的 Q 来自 Decoder 自身上一层的输出,而 K、V 来自 Encoder 的输出。这样 Decoder 就能在生成每个 token 时,有选择地从输入序列中获取信息。
python复制class EncoderLayer(nn.Module):
def __init__(self, d_model, num_heads, d_ff, dropout=0.1):
super().__init__()
self.self_attn = MultiHeadAttention(d_model, num_heads, dropout)
self.ffn = FeedForward(d_model, d_ff, dropout)
self.norm1 = nn.LayerNorm(d_model)
self.norm2 = nn.LayerNorm(d_model)
self.dropout = nn.Dropout(dropout)
def forward(self, x, mask=None):
# self-attention + 残差 + LayerNorm
x = self.norm1(x + self.dropout(self.self_attn(x, x, x, mask)))
# 前馈网络 + 残差 + LayerNorm
x = self.norm2(x + self.dropout(self.ffn(x)))
return x
DecoderLayer 多了 cross-attention,我在这里用 memory 表示 Encoder 的输出。这个命名习惯来自很多开源实现,读起来比较直观。
python复制class DecoderLayer(nn.Module):
def __init__(self, d_model, num_heads, d_ff, dropout=0.1):
super().__init__()
self.self_attn = MultiHeadAttention(d_model, num_heads, dropout)
self.cross_attn = MultiHeadAttention(d_model, num_heads, dropout)
self.ffn = FeedForward(d_model, d_ff, dropout)
self.norm1 = nn.LayerNorm(d_model)
self.norm2 = nn.LayerNorm(d_model)
self.norm3 = nn.LayerNorm(d_model)
self.dropout = nn.Dropout(dropout)
def forward(self, x, memory, tgt_mask=None, memory_mask=None):
# masked self-attention
x = self.norm1(x + self.dropout(self.self_attn(x, x, x, tgt_mask)))
# cross-attention:Q 来自 x,K/V 来自 memory
x = self.norm2(x + self.dropout(self.cross_attn(x, memory, memory, memory_mask)))
x = self.norm3(x + self.dropout(self.ffn(x)))
return x
4.5 完整前向流程
把上面这些模块串起来,就得到一个完整的 Transformer 模型。构造函数里包含词嵌入、位置编码、多个 Encoder 层、多个 Decoder 层,最后还有一个输出线性层 generator,把 Decoder 最后一层的输出映射到词表大小的 logits。
python复制class Transformer(nn.Module):
def __init__(self, vocab_size, d_model=64, num_heads=4, d_ff=128, num_layers=2, dropout=0.1, max_len=64):
super().__init__()
self.embedding = nn.Embedding(vocab_size, d_model)
self.pos_encoding = PositionalEncoding(d_model, max_len, dropout)
self.encoder_layers = nn.ModuleList([
EncoderLayer(d_model, num_heads, d_ff, dropout) for _ in range(num_layers)
])
self.decoder_layers = nn.ModuleList([
DecoderLayer(d_model, num_heads, d_ff, dropout) for _ in range(num_layers)
])
self.generator = nn.Linear(d_model, vocab_size)
self.d_model = d_model
def encode(self, src, src_mask=None):
src_emb = self.embedding(src) * math.sqrt(self.d_model)
src_emb = self.pos_encoding(src_emb)
for layer in self.encoder_layers:
src_emb = layer(src_emb, src_mask)
return src_emb
def decode(self, tgt, memory, tgt_mask=None, memory_mask=None):
tgt_emb = self.embedding(tgt) * math.sqrt(self.d_model)
tgt_emb = self.pos_encoding(tgt_emb)
for layer in self.decoder_layers:
tgt_emb = layer(tgt_emb, memory, tgt_mask, memory_mask)
return tgt_emb
def forward(self, src, tgt, src_mask=None, tgt_mask=None):
memory = self.encode(src, src_mask)
out = self.decode(tgt, memory, tgt_mask, None)
return self.generator(out)
你可能注意到我在词嵌入之后乘了 sqrt(self.d_model)。这是原始论文里的一步操作,原因是 embedding 的数值量级与位置编码相加后,位置编码的信息不容易被淹没。这是一个很小的细节,但对训练稳定性有一定帮助。
前向流程概括起来就是:输入序列经过 embedding 和位置编码,进入 Encoder 得到 memory;Decoder 拿到目标序列右移后的输入,经过 self-attention 和 cross-attention,每一步都可以从 memory 中“提取”所需信息;最终输出层把 hidden state 变成词表大小的 logits,交给损失函数计算。
5. 训练循环与效果评估
5.1 损失函数与优化器配置
模型实现完成后,进入训练环节。损失函数直接用 CrossEntropyLoss,并设置 ignore_index=PAD_TOKEN,这样 PAD 位置的 logits 不会参与梯度计算。优化器用 Adam,学习率采用 warmup 策略——先线性上升,再按步数倒数的规律衰减。论文里用的是这个思路,实践中它确实能显著提升训练稳定性。
python复制import torch.optim as optim
criterion = nn.CrossEntropyLoss(ignore_index=PAD_TOKEN)
model = Transformer(vocab_size=130)
optimizer = optim.Adam(model.parameters(), lr=0.001, betas=(0.9, 0.98), eps=1e-9)
def lr_lambda(step, d_model=64, warmup_steps=4000):
if step == 0:
return 1.0
return (d_model ** -0.5) * min(step ** -0.5, step * (warmup_steps ** -1.5))
scheduler = optim.lr_scheduler.LambdaLR(optimizer, lr_lambda)
warmup 的思路很简单:训练初期模型的参数还没稳定,用太高的学习率容易把 loss 推到很大,甚至变成 NaN。先用小学习率“预热”几千步,等模型进入比较合理的参数空间后再提高学习率,然后逐步衰减。我在这个小型任务上把 warmup_steps 设成了 4000,由于我们的训练步数总量不大,这个参数其实稍微偏大,但结果仍然收敛得很好。
5.2 训练过程记录
训练循环的代码相对常规,但有几个细节值得注意。每次拿到 batch 后,decoder_input 已经是右移后的序列,我们把它传入模型,得到 [B, L, vocab_size] 的 logits;然后和 target(也就是未右移的原始目标序列)计算交叉熵。target 的 PAD 位置被 ignore_index 自动忽略。
python复制from torch.utils.data import DataLoader
dataset = CumSumDataset(num_samples=20000)
dataloader = DataLoader(dataset, batch_size=64, shuffle=True, collate_fn=collate_fn)
def train_one_epoch(model, dataloader, optimizer, criterion, device):
model.train()
total_loss = 0.0
for batch in dataloader:
src = batch["src"].to(device)
decoder_input = batch["decoder_input"].to(device)
tgt = batch["tgt"].to(device)
src_mask = batch["src_padding_mask"].to(device)
# 生成 causal mask,应用到 Decoder 的 self-attention
seq_len = decoder_input.size(1)
tgt_mask = torch.tril(torch.ones(seq_len, seq_len, device=device)).bool()
logits = model(src, decoder_input, src_mask, tgt_mask)
loss = criterion(logits.reshape(-1, logits.size(-1)), tgt.reshape(-1))
optimizer.zero_grad()
loss.backward()
optimizer.step()
scheduler.step()
total_loss += loss.item()
return total_loss / len(dataloader)
我的实测训练曲线大致如下:前 1~2 个 epoch loss 从 4.6 左右快速降到 1.5,之后下降速度放缓;到第 10 个 epoch 左右,loss 在 0.15 附近;第 20 个 epoch 后基本稳定在 0.02 以下。这意味着模型对训练集已经能达到非常高的准确率,说明模型实现本身没有致命 bug。
5.3 贪心解码推理
训练完成后,最激动人心的部分当然是看模型能不能用。推理时无法像训练那样一次性输入整个目标序列,因为推理时我们并不知道目标是什么。我们需要从一个 [BOS] 开始,把当前预测出来的 token 拼到输入序列后面,再次送入 Decoder,循环往复。这种方式叫贪心解码,每一步只取概率最大的 token。
python复制@torch.no_grad()
def greedy_decode(model, src, max_len=20, bos_token=BOS_TOKEN, device="cpu"):
model.eval()
src = src.unsqueeze(0).to(device)
src_mask = (src != PAD_TOKEN).unsqueeze(1).unsqueeze(2).to(device)
memory = model.encode(src, src_mask)
ys = torch.tensor([[bos_token]]).to(device)
for _ in range(max_len):
seq_len = ys.size(1)
tgt_mask = torch.tril(torch.ones(seq_len, seq_len, device=device)).bool()
out = model.decode(ys, memory, tgt_mask, None)
prob = model.generator(out[:, -1])
next_token = prob.argmax(dim=-1).item()
ys = torch.cat([ys, torch.tensor([[next_token]]).to(device)], dim=1)
return ys.squeeze(0).tolist()
一个典型的推理示例如下:
text复制输入: [7, 2, 9, 4]
真实: [7, 9, 18, 22]
预测: [7, 9, 18, 22]
能正确输出累加和,说明模型真正学会了在 Decoder 的第 i 个位置关注 Encoder 前 i 个位置的信息。你可以进一步尝试把输入改成训练时没见过的长度,比如长度为 15 的序列(训练时最长只有 12),如果模型仍然输出正确,说明它对长度具备一定的泛化能力。
6. 常见问题与排查实录
6.1 为什么 loss 不降或者乱跳
这个问题的原因通常不是模型结构本身,而是训练配置。我在调试早期踩过一个大坑:学习率设成固定的 0.001,没有 warmup,结果 loss 在 2.0 附近来回震荡,怎么都降不下去。后来换成论文的 warmup 策略才稳定收敛。
| 现象 | 常见原因 | 解决办法 |
|---|---|---|
| loss 不降 | 学习率过小或过大 | 使用 warmup 学习率调度器 |
| loss 震荡明显 | 学习率过大 | 调低初始学习率或增大 warmup_steps |
| loss 上升到 NaN | 数值不稳定 | 检查是否有除以零、-1e9 是否写成了 -float('inf') 导致溢出 |
| 训练收敛但推理全错 | mask 用错或 shift-right 没对齐 | 打印 mask 和输入输出张量逐维检查 |
还有一个容易踩的坑是 CrossEntropyLoss 输入形状。PyTorch 要求 logits 是 [N, C],target 是 [N],其中 C 为类别数。我这里把 [B, L, vocab_size] reshape 成 [-1, vocab_size],target reshape 成 [-1],顺序要对齐。如果不 reshape 或者顺序搞混,loss 会莫名其妙地偏高。
6.2 推理结果全是一个 token 或重复循环
这种情况大多是 Decoder 的 self-attention mask 没有正确加因果约束。训练时因为有 Teacher Forcing,模型可以看到完整目标序列,所以就算 mask 错了,loss 也未必高得离谱;但推理时目标序列是逐步生成的,如果 mask 没遮住未来位置,模型在生成当前 token 时会“看到”自己还没生成出来的 token,这就产生了信息泄漏。解决方法是检查每个 Decoder layer 的 self-attention 是否都传入了下三角 mask。我因为变量名重复,曾有一处把 tgt_mask 误传成了 memory_mask,导致模型在推理时怎么都不对。
6.3 Mask 形状对不上
这是初学者最容易卡住的报错点:The size of tensor a ... must match the size of tensor b ...。核心原因是 mask 的形状没有对齐到 scores 的 [B, num_heads, L_q, L_k]。我的经验是,在所有 mask 传入 MultiHeadAttention 之前,统一处理成 [B, 1, L_q, L_k],然后在函数内部再扩张到 [B, num_heads, L_q, L_k],这样代码路径清晰,不容易错。如果发现维度对不上,不要靠猜,直接在每个 mask 上打印 .shape,一步步确认。
6.4 PAD 位置被预测成有效数字
如果不设置 ignore_index=PAD_TOKEN,模型会花大量精力“预测 PAD”,短期看不出问题,但真实位置的准确率往往上不去。另外,在推理时如果输入长度小于 max_len,padding 部分的输出本来就不关心,所以我们在生成推理结果后,只取有效长度范围内的输出即可。
7. 扩展到更真实的场景
跑通累加任务只是第一步。当你用手写 Transformer 完整跑通一个简单的序列到序列任务后,再去看现在流行的大模型或者各种 Transformer 变体,你会发现自己已经能看懂它们的基础代码了。比如 Vision Transformer(ViT)把图片切块后当作 token 序列,核心的 self-attention 代码跟你手写的几乎一模一样;Swin Transformer 只是在自注意力的计算范围上做了窗口限制,本质上还是 Q、K、V 那套运算。
从我这个项目再往深处走,有几个自然的扩展方向值得尝试。第一,把任务换成字符级语言模型,用一批英文小说文本训练模型预测下一个字符,这能让你理解大语言模型里“自回归生成”的基本流程。第二,把固定正余弦位置编码换成可学习的位置编码,或者换成 RoPE、ALiBi 这类更现代的相对位置编码,观察不同位置编码对长序列泛化能力的影响。第三,引入 Beam Search 替代贪心解码,虽然代码量不大,但对生成质量的提升非常直观。
最后再分享一个小技巧:如果你想确认自己写的 attention 是不是真的在“看”正确的位置,可以在训练结束后把某几个样本的 attention 权重打印出来,画成热力图。你会非常直观地看到,Decoder 在生成第 i 个累加和时,把大部分注意力都放在了 Encoder 的前 i 个输入 token 上。这种“亲眼看到模型学到规律”的瞬间,才是手写实现最大的回报,也是“纸上得来终觉浅”这句话真正的分量所在。
