本文档从零讲解如何用纯 PyTorch 手写一个 Transformer,让它学会"把一串数字倒过来"。列表翻转是入门 Seq2Seq 的经典 toy task:任务规则极其简单(反转),但模型必须自己学会 "第 1 个输出对应最后一个输入"这种远距离的、依赖位置信息的映射关系,因此非常适合用来验证你手写的 Transformer 是否真的 work。
目录
任务定义与整体思路
全局配置与词表设计
数据集构造
Mask:Padding Mask 与 Causal Mask
模型组件
5.1 缩放点积注意力 Scaled Dot-Product Attention
5.2 多头注意力 Multi-Head Attention
5.3 位置编码 Positional Encoding
5.4 前馈网络 Feed Forward
5.5 Encoder Layer / Decoder Layer
5.6 组装完整的 Encoder-Decoder Transformer
5.7 Decoder-Only(GPT 风格)变体
训练流程
推理与贪心解码
超参数与训练效果
常见问题与思考
1. 任务定义与整体思路
任务 :输入一个数字列表,输出它的反转。
1 2 输入: [3, 5, 7, 2] 输出: [2, 7, 5, 3]
这本质上是一个序列到序列(Seq2Seq)问题,和机器翻译的结构一模一样:把"源序列"编码,再自回归地生成"目标序列"。因此我们可以原封不动地套用 2017 年论文 Attention Is All You Need 中的 Transformer 架构。
本文实现两种架构,改一行配置即可切换:
架构
说明
encoder_decoder
原始 Transformer:Encoder 编码源序列,Decoder 通过 cross-attention 读取源序列,自回归生成目标
decoder_only
GPT 风格:把 <BOS> src <EOS> tgt <EOS> 拼成一条序列,用 causal mask 做 next-token 预测
架构一:Encoder-Decoder 的数据流
1 2 3 4 5 6 7 8 9 10 11 12 13 源序列 (src) │ ▼ Embedding + 位置编码 │ ▼ ┌──────────────┐ memory │ Encoder │ ────────────────┐ └──────────────┘ ▼ ┌──────────────┐ ┌───────────┐ 目标序列前缀 (<BOS> ...) │ Decoder │ ──▶ │ Linear │ ──▶ 下一个 token 的概率分布 ────▶ │ (自回归生成) │ │(generator)│ └──────────────┘ └───────────┘
Encoder 把源序列编码成一组上下文表示 memory;Decoder 每一步把"已生成的目标前缀"和 memory 一起输入,预测下一个 token。
架构二:Decoder-Only 的数据流
1 2 3 4 5 6 7 8 9 10 <BOS> src... <EOS> tgt 前缀... │ ▼ Embedding + 位置编码 │ ▼ ┌──────────────────────────┐ ┌───────────┐ │ N × DecoderOnlyLayer │ ──▶ │ Linear │ ──▶ 每个位置的下一个 token 概率分布 │ (masked self -attention ) │ │(generator)│ └──────────────────────────┘ └───────────┘
Decoder-Only 没有 Encoder,也不分源/目标:把"prompt(<BOS> src <EOS>)+ 已生成的目标前缀"拼成一条序列 直接输入,模型在 causal mask 约束下对每个位置 预测下一个 token。源序列的信息全靠 self-attention 在序列内部传递——prompt 部分的 token 充当"上下文",生成位置的 token 通过注意力回看 prompt。
2. 全局配置与词表设计
把所有超参数集中在一个配置模块里,方便调整:
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 PAD_IDX = 0 BOS_IDX = 1 EOS_IDX = 2 VOCAB_SIZE = 13 TRAIN_NUM_SAMPLES = 200000 TEST_NUM_SAMPLES = 200 MIN_LEN = 3 MAX_LEN = 12 TRAIN_SEED = 0 TEST_SEED = 42 ARCHITECTURE = "decoder_only" D_MODEL = 128 NUM_HEADS = 4 NUM_ENCODER_LAYERS = 2 NUM_DECODER_LAYERS = 4 D_FF = 512 DROPOUT = 0.1 MODEL_MAX_LEN = 64 BATCH_SIZE = 64 EPOCHS = 9 LR = 3e-4 GRAD_CLIP = 1.0 CHECKPOINT_PATH = "checkpoint.pt"
设计要点:
PAD/BOS/EOS 占用 0/1/2 ,把 3~12 留给数字 0~9(即 token_id = digit + 3)。这样词表是固定 的,不需要根据数据动态构建,适合 toy task。
后续 loss 计算用 ignore_index=PAD_IDX 把 padding 位置的损失忽略掉。
训练/测试用不同的随机种子 生成数据,保证测试集不会出现在训练集中,评估的才是真正的"泛化能力"。
3. 数据集构造
3.1 Encoder-Decoder 格式的样本
为每个样本生成三个张量:
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 import randomimport torchfrom torch.utils.data import Datasetfrom torch.nn.utils.rnn import pad_sequence PAD_IDX = 0 BOS_IDX = 1 EOS_IDX = 2 class ReverseDataset (Dataset ): """ 输入: 一串数字 目标: 将这串数字反转 例如: src: 3 5 7 tgt_input: <BOS> 7 5 3 tgt_label: 7 5 3 <EOS> 通过传入不同的 seed,可以生成互不相同的训练集 / 测试集。 """ def __init__ (self, num_samples=2000 , min_len=3 , max_len=12 , seed=None ): super (ReverseDataset, self ).__init__() rng = random.Random(seed) self .samples = [] for _ in range (num_samples): length = rng.randint(min_len, max_len) digits = torch.randint(0 , 10 , (length,), dtype=torch.long) src = digits + 3 tgt = src.flip(0 ) tgt_in = torch.cat( [ torch.tensor([BOS_IDX], dtype=torch.long), tgt, ] ) tgt_out = torch.cat( [ tgt, torch.tensor([EOS_IDX], dtype=torch.long), ] ) self .samples.append((src, tgt_in, tgt_out)) def __len__ (self ): return len (self .samples) def __getitem__ (self, idx ): return self .samples[idx]
以 src = [3, 5, 7](token id 为 [6, 8, 10])为例:
1 2 3 src : 6 8 10 tgt_in : 1 10 8 6 # <BOS> 开头,整体右移tgt_out : 10 8 6 2 # 末尾补 <EOS>
为什么要右移(shift right)? 因为 Decoder 是自回归的:在位置 t,模型只能看到 <BOS> 和前 t-1 个已生成的 token,然后预测第 t 个 token。tgt_in[i] 的预测目标就是 tgt_out[i],两者恰好错开一位,这就是标准的 teacher forcing 训练方式。
3.2 Decoder-Only 格式的样本
GPT 风格的模型只吃一条序列,所以把 prompt 和目标拼接起来:
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 class DecoderOnlyReverseDataset (Dataset ): """ 与 ReverseDataset 相同的任务,但整理成 decoder-only(GPT 风格)的单序列格式: full: <BOS> src... <EOS> tgt... <EOS> input_ids: full[:-1] labels: full[1:],其中 prompt 部分(src token 和分隔 <EOS> 的预测位) 置为 PAD_IDX,训练时忽略这部分 loss,只学习生成 tgt。 __getitem__ 返回 (src, input_ids, labels),src 供评估时构造 prompt 使用。 """ def __init__ (self, num_samples=2000 , min_len=3 , max_len=12 , seed=None ): super (DecoderOnlyReverseDataset, self ).__init__() rng = random.Random(seed) self .samples = [] for _ in range (num_samples): length = rng.randint(min_len, max_len) digits = torch.randint(0 , 10 , (length,), dtype=torch.long) src = digits + 3 tgt = src.flip(0 ) full = torch.cat( [ torch.tensor([BOS_IDX], dtype=torch.long), src, torch.tensor([EOS_IDX], dtype=torch.long), tgt, torch.tensor([EOS_IDX], dtype=torch.long), ] ) input_ids = full[:-1 ] labels = full[1 :].clone() prompt_len = length + 2 labels[: prompt_len - 1 ] = PAD_IDX self .samples.append((src, input_ids, labels)) def __len__ (self ): return len (self .samples) def __getitem__ (self, idx ): return self .samples[idx]
关键技巧是 label masking :prompt 部分的 label 置为 PAD_IDX,让 loss 只计算在目标部分。模型只学习"看到 <BOS> src <EOS> 之后该如何续写反转序列",这正是 GPT 类模型做指令任务的标准做法。
3.3 Collate:batch 内对齐长度
一个 batch 里的序列长度不同,用 pad_sequence 右侧补 PAD:
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 def collate_fn_decoder_only (batch ): src_list = [item[0 ] for item in batch] input_ids_list = [item[1 ] for item in batch] labels_list = [item[2 ] for item in batch] src = pad_sequence(src_list, batch_first=True , padding_value=PAD_IDX) input_ids = pad_sequence(input_ids_list, batch_first=True , padding_value=PAD_IDX) labels = pad_sequence(labels_list, batch_first=True , padding_value=PAD_IDX) return src, input_ids, labelsdef collate_fn (batch ): src_list = [item[0 ] for item in batch] tgt_in_list = [item[1 ] for item in batch] tgt_out_list = [item[2 ] for item in batch] src = pad_sequence(src_list, batch_first=True , padding_value=PAD_IDX) tgt_in = pad_sequence(tgt_in_list, batch_first=True , padding_value=PAD_IDX) tgt_out = pad_sequence(tgt_out_list, batch_first=True , padding_value=PAD_IDX) return src, tgt_in, tgt_out
补齐产生的 PAD 之后会被两处机制处理掉:注意力中的 padding mask (不看它)和 loss 的 ignore_index (不惩罚它)。
4. Mask:Padding Mask 与 Causal Mask
Mask 是 Transformer 里最容易出错的部分,这里约定:1 表示可见,0 表示屏蔽 。
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 import torch PAD_IDX = 0 def create_padding_mask (seq, pad_idx=PAD_IDX ): """ seq: (batch_size, seq_len) return: (batch_size, 1, 1, seq_len) 1 表示真实 token 0 表示 padding token """ mask = (seq != pad_idx).float () return mask.unsqueeze(1 ).unsqueeze(2 )def create_causal_mask (size, device ): """ return: (1, 1, size, size) 下三角为 1,表示可见 上三角为 0,表示未来位置不可见 """ mask = torch.tril(torch.ones(size, size, device=device)) return mask.unsqueeze(0 ).unsqueeze(0 )def create_decoder_self_mask (tgt, pad_idx=PAD_IDX ): """ tgt: (batch_size, tgt_len) return: (batch_size, 1, tgt_len, tgt_len) Decoder self-attention mask: causal mask + padding mask """ pad_mask = create_padding_mask(tgt, pad_idx) causal_mask = create_causal_mask(tgt.size(1 ), tgt.device) mask = causal_mask * pad_mask return mask
下面用具体例子逐个展示每种 mask 长什么样。约定一个 batch(B = 2),其中第二条序列更短、右侧补了 PAD:
1 2 序列 1 : [BOS, 8, 5, 10] 长度 4 ,无 padding 序列 2 : [BOS, 7, PAD, PAD] 实际长度 2 ,补齐到 4
Padding Mask 的例子
create_padding_mask 对上面的 batch 生成形状 (2, 1, 1, 4) 的 mask:
1 2 序列 1: [1 , 1 , 1 , 1 ] 序列 2: [1 , 1 , 0 , 0 ]
这个 mask 可以广播到注意力分数 (B, H, L_q, L_k) 上,作用在最后一维(key 维度) :任何 query 位置都不允许注意 padding token。以序列 2 为例,无论哪个位置做 query,它算出的注意力权重只会分布在前 2 个真实 token 上,2 个 PAD 位拿到的权重恒为 0。
Causal Mask(因果掩码)的例子
create_causal_mask(4) 生成一个 4×4 的下三角矩阵(形状 (1, 1, 4, 4),与 batch 无关):
1 2 3 4 5 key: 0 1 2 3 query 0 [ 1 0 0 0 ] query 1 [ 1 1 0 0 ] query 2 [ 1 1 1 0 ] query 3 [ 1 1 1 1 ]
位置 i 只能看到 0..i,看不到未来 。这是自回归生成的核心约束:预测第 t 个 token 时不许偷看第 t 个及之后的答案。比如生成反转序列时,模型在预测第 2 个输出 token 的位置,注意力里绝不能出现第 2、3 个位置的真实 token,否则训练就成了抄答案。
Decoder 自注意力组合 Mask 的例子
create_decoder_self_mask 把上面两个 mask 逐元素相乘 。对序列 2([BOS, 7, PAD, PAD]),结果是形状 (1, 4, 4) 的矩阵:
1 2 3 4 5 key: 0 1 2 3 query 0 [ 1 0 0 0 ] query 1 [ 1 1 0 0 ] query 2 [ 1 1 0 0 ] query 3 [ 1 1 0 0 ]
既看不见未来(上三角为 0),也看不见 padding(后两列为 0)。
注意 Encoder 的自注意力不需要 causal mask——编码时整个源序列都是可见的,只需要 padding mask。对序列 2 来说,Encoder 用的 mask 就是:
1 2 key: 0 1 2 3 任意 query [ 1 1 0 0 ]
5. 模型组件
5.1 缩放点积注意力
对应论文公式:
Attention ( Q , K , V ) = softmax ( Q K T d k ) V \text{Attention}(Q, K, V) = \text{softmax}\left(\frac{QK^T}{\sqrt{d_k}}\right)V
Attention ( Q , K , V ) = softmax ( d k Q K T ) V
代码实现(作为 MultiHeadAttention 的一个方法):
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 def scaled_dot_product_attention (self, Q, K, V, mask=None ): """ Q, K, V: (batch_size, num_heads, seq_len, d_k) mask: 可广播到 (batch_size, num_heads, seq_len_q, seq_len_k) """ 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 ) attn_weights = F.softmax(scores, dim=-1 ) attn_weights = self .dropout(attn_weights) output = torch.matmul(attn_weights, V) return output, attn_weights
逐步拆解(以列表翻转为直觉):
QK^T :每个 query 位置与每个 key 位置算一个"相关度得分"。对翻转任务来说,模型要学会让输出的第 1 个位置去高度关注输入的最后一个位置。
除以 √d_k :d_k 较大时点积的方差会随之增大,softmax 进入饱和区、梯度消失。缩放让得分的方差与维度无关。
mask 填充 -1e9 :被屏蔽位置的得分变成一个极大的负数,softmax 后权重约等于 0。这里刻意不用 -inf,避免整行都被 mask 时 softmax 出现 NaN。
乘 V :按注意力权重对 value 加权求和,得到每个位置的输出表示。
5.2 多头注意力
多头注意力的思想是:把 d_model 切成 num_heads 份,每份独立做注意力,再拼回来 。不同的头可以学习不同类型的依赖关系(比如一个头关注"对称位置",另一个头关注"相邻位置")。
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 import mathimport torchimport torch.nn as nnimport torch.nn.functional as Fclass MultiHeadAttention (nn.Module): def __init__ (self, d_model, num_heads, dropout=0.1 ): super (MultiHeadAttention, self ).__init__() assert d_model % num_heads == 0 , "d_model 必须能被 num_heads 整除" self .d_model = d_model self .num_heads = num_heads self .d_k = 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 .W_o = nn.Linear(d_model, d_model) self .dropout = nn.Dropout(dropout) def scaled_dot_product_attention (self, Q, K, V, mask=None ): """ Q, K, V: (batch_size, num_heads, seq_len, d_k) mask: 可广播到 (batch_size, num_heads, seq_len_q, seq_len_k) """ 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 ) attn_weights = F.softmax(scores, dim=-1 ) attn_weights = self .dropout(attn_weights) output = torch.matmul(attn_weights, V) return output, attn_weights def forward (self, query, key, value, mask=None ): """ query: (batch_size, seq_len_q, d_model) key: (batch_size, seq_len_k, d_model) value: (batch_size, seq_len_v, d_model) """ batch_size = query.size(0 ) Q = self .W_q(query) K = self .W_k(key) V = self .W_v(value) Q = Q.view(batch_size, -1 , self .num_heads, self .d_k).transpose(1 , 2 ) K = K.view(batch_size, -1 , self .num_heads, self .d_k).transpose(1 , 2 ) V = V.view(batch_size, -1 , self .num_heads, self .d_k).transpose(1 , 2 ) attn_output, attn_weights = self .scaled_dot_product_attention( Q, K, V, mask ) attn_output = ( attn_output.transpose(1 , 2 ) .contiguous() .view(batch_size, -1 , self .d_model) ) output = self .W_o(attn_output) return output, attn_weights
forward 中的核心是张量形状变换——拆头与并头 。注意 view 只是"重排"没有新增参数,真正的可学习参数在 W_q/W_k/W_v/W_o 四个线性层里。另外 transpose 之后内存不连续,view 之前必须先调用 .contiguous()。
5.3 位置编码
注意力本身是置换不变 的——打乱输入顺序,输出只是跟着打乱,模型完全不知道"谁在第几个位置"。而翻转任务恰恰强依赖位置信息,所以必须显式注入位置。
使用论文中的正弦位置编码:
P E ( p o s , 2 i ) = sin ( p o s 10000 2 i / d m o d e l ) , P E ( p o s , 2 i + 1 ) = cos ( p o s 10000 2 i / d m o d e l ) PE_{(pos, 2i)} = \sin\left(\frac{pos}{10000^{2i/d_{model}}}\right), \quad PE_{(pos, 2i+1)} = \cos\left(\frac{pos}{10000^{2i/d_{model}}}\right)
P E ( p os , 2 i ) = sin ( 1000 0 2 i / d m o d e l p os ) , P E ( p os , 2 i + 1 ) = cos ( 1000 0 2 i / d m o d e l p os )
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 class PositionalEncoding (nn.Module): def __init__ (self, d_model, max_len=512 , dropout=0.1 ): super (PositionalEncoding, self ).__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 ) self .register_buffer("pe" , pe) def forward (self, x ): """ x: (batch_size, seq_len, d_model) """ x = x + self .pe[:, : x.size(1 )] return self .dropout(x)
不同维度对应不同频率的正弦波:低维波动快(编码局部位置),高维波动慢(编码全局位置)。
用 register_buffer 注册而非常量属性:buffer 不是参数、不参与训练,但 model.to(device) 和 state_dict 都会正确处理它。
forward 中直接加到 embedding 上:x = x + self.pe[:, :x.size(1)]。
5.4 前馈网络
1 2 3 4 5 6 7 8 9 10 11 12 13 14 class FeedForward (nn.Module): def __init__ (self, d_model, d_ff, dropout=0.1 ): super (FeedForward, self ).__init__() self .linear1 = nn.Linear(d_model, d_ff) self .linear2 = nn.Linear(d_ff, d_model) self .dropout = nn.Dropout(dropout) def forward (self, x ): x = self .linear1(x) x = F.relu(x) x = self .dropout(x) x = self .linear2(x) return x
注意力负责"混合不同位置的信息",FFN 负责"对每个位置独立地做非线性变换"。先升维再降维(d_ff 通常取 4 × d_model),提供模型的非线性表达能力。
5.5 Encoder Layer / Decoder Layer
每个子层都套着 残差连接 + LayerNorm 的壳(Post-LN 结构):
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 class EncoderLayer (nn.Module): def __init__ (self, d_model, num_heads, d_ff, dropout=0.1 ): super (EncoderLayer, self ).__init__() self .self_attn = MultiHeadAttention(d_model, num_heads, dropout) self .feed_forward = 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, src, src_mask ): """ src: (B, src_len, d_model) src_mask: (B, 1, 1, src_len) """ attn_out, _ = self .self_attn(src, src, src, src_mask) src = self .norm1(src + self .dropout(attn_out)) ff_out = self .feed_forward(src) src = self .norm2(src + self .dropout(ff_out)) return src
残差连接 :x + sublayer(x),让梯度可以绕过子层直接回传,深层网络才训得动。
LayerNorm :把每个 token 的特征分布稳定下来,加速收敛。
Decoder Layer 比 Encoder 多一个 cross-attention :
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 class DecoderLayer (nn.Module): def __init__ (self, d_model, num_heads, d_ff, dropout=0.1 ): super (DecoderLayer, self ).__init__() self .self_attn = MultiHeadAttention(d_model, num_heads, dropout) self .cross_attn = MultiHeadAttention(d_model, num_heads, dropout) self .feed_forward = 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, tgt, memory, tgt_mask, memory_mask ): """ tgt: (B, tgt_len, d_model) memory: encoder output, (B, src_len, d_model) tgt_mask: decoder self-attention mask (B, 1, tgt_len, tgt_len) memory_mask: cross-attention mask (B, 1, 1, src_len) """ attn_out, _ = self .self_attn(tgt, tgt, tgt, tgt_mask) tgt = self .norm1(tgt + self .dropout(attn_out)) attn_out, _ = self .cross_attn(tgt, memory, memory, memory_mask) tgt = self .norm2(tgt + self .dropout(attn_out)) ff_out = self .feed_forward(tgt) tgt = self .norm3(tgt + self .dropout(ff_out)) return tgt
三个子层各司其职:
masked self-attention :只看已生成的目标前缀(causal mask);
cross-attention :query 来自 Decoder,key/value 来自 Encoder 输出 memory——这是 Encoder 和 Decoder 之间的桥梁,Decoder 的每个位置通过它去"查询"源序列中相关的部分。翻转任务中,模型正是靠它学会"当前该输出源序列倒数第几个元素";
FFN :逐位置非线性变换。
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 class Transformer (nn.Module): def __init__ ( self, src_vocab_size, tgt_vocab_size, d_model=128 , num_heads=4 , num_encoder_layers=2 , num_decoder_layers=2 , d_ff=512 , dropout=0.1 , pad_idx=PAD_IDX, max_len=64 , ): super (Transformer, self ).__init__() self .d_model = d_model self .pad_idx = pad_idx self .src_embedding = nn.Embedding( src_vocab_size, d_model, padding_idx=pad_idx, ) self .tgt_embedding = nn.Embedding( tgt_vocab_size, d_model, padding_idx=pad_idx, ) self .pos_enc = PositionalEncoding(d_model, max_len, dropout) self .encoder_layers = nn.ModuleList( [ EncoderLayer(d_model, num_heads, d_ff, dropout) for _ in range (num_encoder_layers) ] ) self .decoder_layers = nn.ModuleList( [ DecoderLayer(d_model, num_heads, d_ff, dropout) for _ in range (num_decoder_layers) ] ) self .generator = nn.Linear(d_model, tgt_vocab_size) def make_src_mask (self, src ): """ src: (B, src_len) return: (B, 1, 1, src_len) """ return create_padding_mask(src, self .pad_idx) def make_tgt_mask (self, tgt ): """ tgt: (B, tgt_len) return: (B, 1, tgt_len, tgt_len) """ return create_decoder_self_mask(tgt, self .pad_idx) def encode (self, src, src_mask=None ): if src_mask is None : src_mask = self .make_src_mask(src) x = self .src_embedding(src) * math.sqrt(self .d_model) x = self .pos_enc(x) for layer in self .encoder_layers: x = layer(x, src_mask) return x def decode (self, tgt, memory, tgt_mask=None , memory_mask=None ): if tgt_mask is None : tgt_mask = self .make_tgt_mask(tgt) x = self .tgt_embedding(tgt) * math.sqrt(self .d_model) x = self .pos_enc(x) for layer in self .decoder_layers: x = layer(x, memory, tgt_mask, memory_mask) return x def forward (self, src, tgt ): """ src: (B, src_len) tgt: (B, tgt_len) tgt 是 decoder 的输入序列,通常以 <BOS> 开头。 """ src_mask = self .make_src_mask(src) memory = self .encode(src, src_mask) tgt_mask = self .make_tgt_mask(tgt) memory_mask = src_mask dec_out = self .decode( tgt=tgt, memory=memory, tgt_mask=tgt_mask, memory_mask=memory_mask, ) logits = self .generator(dec_out) return logits
几个关键细节:
Embedding 缩放 ——x = self.src_embedding(src) * math.sqrt(self.d_model):论文中的做法,embedding 乘以 √d_model,使其数值尺度与位置编码相当,避免位置信息被淹没。
encode / decode 分离 :把编码和解码写成独立方法,推理时可以只跑一次 Encoder ,然后循环调用 Decoder 逐 token 生成。
generator :最后一层线性把 d_model 维表示映射回词表大小的 logits。
5.7 Decoder-Only(GPT 风格)变体
GPT 风格的 block 没有 cross-attention,整条规定格式的序列 <BOS> src <EOS> tgt <EOS> 全靠 masked self-attention 自己消化:
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 class DecoderOnlyLayer (nn.Module): """ GPT 风格的 block:只有 masked self-attention + FFN,没有 cross-attention。 """ def __init__ (self, d_model, num_heads, d_ff, dropout=0.1 ): super (DecoderOnlyLayer, self ).__init__() self .self_attn = MultiHeadAttention(d_model, num_heads, dropout) self .feed_forward = 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 ): """ x: (B, seq_len, d_model) mask: causal + padding mask, (B, 1, seq_len, seq_len) """ attn_out, _ = self .self_attn(x, x, x, mask) x = self .norm1(x + self .dropout(attn_out)) ff_out = self .feed_forward(x) x = self .norm2(x + self .dropout(ff_out)) return xclass DecoderOnlyTransformer (nn.Module): """ 单序列输入(prompt + target 拼接),用 causal mask 做 next-token 预测。 序列格式: <BOS> src... <EOS> tgt... <EOS> """ def __init__ ( self, vocab_size, d_model=128 , num_heads=4 , num_layers=2 , d_ff=512 , dropout=0.1 , pad_idx=PAD_IDX, max_len=64 , ): super (DecoderOnlyTransformer, self ).__init__() self .d_model = d_model self .pad_idx = pad_idx self .embedding = nn.Embedding( vocab_size, d_model, padding_idx=pad_idx, ) self .pos_enc = PositionalEncoding(d_model, max_len, dropout) self .layers = nn.ModuleList( [ DecoderOnlyLayer(d_model, num_heads, d_ff, dropout) for _ in range (num_layers) ] ) self .generator = nn.Linear(d_model, vocab_size) def forward (self, x ): """ x: (B, seq_len),token id 序列(右 padding) return: (B, seq_len, vocab_size) """ mask = create_decoder_self_mask(x, self .pad_idx) h = self .embedding(x) * math.sqrt(self .d_model) h = self .pos_enc(h) for layer in self .layers: h = layer(h, mask) logits = self .generator(h) return logits
整个模型的 forward 极其简洁——mask 在内部自动构建,输入一条 token 序列,输出每个位置下一 token 的 logits。
6. 训练流程
6.1 损失函数与优化器
1 2 optimizer = torch.optim.Adam(model.parameters(), lr=3e-4 ) criterion = nn.CrossEntropyLoss(ignore_index=PAD_IDX)
逐 token 的交叉熵:logits (B, L, V) 与标签 (B, L) 都展平后计算。
ignore_index=PAD_IDX 是数据补齐方案的最后一块拼图:padding 位(以及 decoder-only 中被置 PAD 的 prompt 位)不产生任何损失和梯度。
6.2 训练循环
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 import torchimport torch.nn as nnfrom torch.utils.data import DataLoaderfrom tqdm import tqdmdef train (model, train_loader, device, epochs, lr, grad_clip, vocab_size, arch ): optimizer = torch.optim.Adam(model.parameters(), lr=lr) criterion = nn.CrossEntropyLoss(ignore_index=PAD_IDX) for epoch in range (1 , epochs + 1 ): model.train() total_loss = 0.0 total_tokens = 0 pbar = tqdm( train_loader, desc=f"Epoch {epoch} /{epochs} " , leave=True , ) for src, tgt_in, tgt_out in pbar: src = src.to(device) tgt_in = tgt_in.to(device) tgt_out = tgt_out.to(device) if arch == "decoder_only" : logits = model(tgt_in) else : logits = model(src, tgt_in) loss = criterion( logits.view(-1 , vocab_size), tgt_out.view(-1 ), ) optimizer.zero_grad() loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), grad_clip) optimizer.step() non_pad_tokens = (tgt_out != PAD_IDX).sum ().item() total_loss += loss.item() * non_pad_tokens total_tokens += non_pad_tokens pbar.set_postfix(loss=f"{total_loss / max (total_tokens, 1 ):.4 f} " ) avg_loss = total_loss / max (total_tokens, 1 ) tqdm.write(f"Epoch {epoch} finished, avg loss: {avg_loss:.4 f} " )
要点:
Teacher forcing :训练时把真实的 tgt_in(含 <BOS> 的前缀)喂给 Decoder,而不是喂模型自己上一步的预测,收敛更快更稳。
梯度裁剪 (clip_grad_norm_):限制梯度范数不超过 1.0,防止训练初期 loss 震荡。
tqdm 进度条上实时显示按有效 token 数加权的平均 loss。
6.3 评估:整序列完全正确率
对翻转任务,"逐 token 准确率"没有意义——错一个位置整条就错了,所以评估指标是整条序列完全匹配的比例 :
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 @torch.no_grad() def evaluate (model, test_dataset, device, batch_size, arch, collate, num_show=5 ): """ 在独立的测试集上评估(测试集与训练集使用不同种子生成,无重叠)。 返回整序列完全正确的准确率。 """ model.eval () test_loader = DataLoader( test_dataset, batch_size=batch_size, shuffle=False , collate_fn=collate, ) correct = 0 total = 0 for src, _, tgt_out in tqdm(test_loader, desc="Evaluating" , leave=False ): src = src.to(device) if arch == "decoder_only" : prompt = build_decoder_only_prompt(src) prompt_len = prompt.size(1 ) pred = greedy_decode_decoder_only( model, prompt, max_len=src.size(1 ) + 5 ) pred = pred[:, prompt_len:] for i in range (src.size(0 )): pred_str = tensor_to_str(pred[i]) tgt_str = expected_reverse_str(src[i]) if pred_str == tgt_str: correct += 1 total += 1 else : pred = greedy_decode(model, src, max_len=src.size(1 ) + 5 ) for i in range (src.size(0 )): pred_str = tensor_to_str(pred[i]) tgt_str = tensor_to_str(tgt_out[i]) if pred_str == tgt_str: correct += 1 total += 1 acc = correct / max (total, 1 ) print (f"Test accuracy: {acc:.4 f} ({correct} /{total} )" ) return acc
6.4 训练入口与 checkpoint 保存
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 def save_checkpoint (model, config_dict, path="checkpoint.pt" ): """训练结束后保存模型权重与超参数配置。""" torch.save( { "model_state_dict" : model.state_dict(), "config" : config_dict, }, path, ) print (f"Checkpoint saved to {path} " )def run (): device = torch.device("cuda" if torch.cuda.is_available() else "cpu" ) torch.manual_seed(0 ) train_dataset = build_dataset(num_samples=200000 , seed=0 ) test_dataset = build_dataset(num_samples=200 , seed=42 ) train_loader = DataLoader( train_dataset, batch_size=64 , shuffle=True , collate_fn=get_collate_fn(), ) model = build_model(device) train(model, train_loader, device) save_checkpoint(model, {...超参数...}) evaluate(model, test_dataset, device)
checkpoint 里同时保存 model_state_dict 和模型超参数(d_model、层数、架构类型等),因此推理端不需要猜配置。
7. 推理与贪心解码
7.1 token 序列转可读字符串
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 def tensor_to_str (x ): """ 将 token id 序列转成可读字符串。 遇到 EOS 停止。 """ result = [] for idx in x: idx = idx.item() if torch.is_tensor(idx) else int (idx) if idx == EOS_IDX: break if idx in (PAD_IDX, BOS_IDX): continue result.append(str (idx - 3 )) return " " .join(result)
7.2 Encoder-Decoder 的贪心解码
推理时没有真实目标可用,只能自回归 地一个 token 一个 token 生成:
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 @torch.no_grad() def greedy_decode (model, src, max_len=None ): """ src: (batch_size, src_len) """ model.eval () if max_len is None : max_len = src.size(1 ) + 5 src_mask = model.make_src_mask(src) memory = model.encode(src, src_mask) ys = torch.full( (src.size(0 ), 1 ), BOS_IDX, dtype=torch.long, device=src.device, ) for _ in range (max_len): tgt_mask = model.make_tgt_mask(ys) dec_out = model.decode( tgt=ys, memory=memory, tgt_mask=tgt_mask, memory_mask=src_mask, ) logits = model.generator(dec_out) next_token = logits[:, -1 :].argmax(dim=-1 ) ys = torch.cat([ys, next_token], dim=1 ) if (next_token == EOS_IDX).all (): break return ys
贪心解码 :每步取 argmax,简单直接。对翻转这种确定性任务足够好(beam search 主要提升开放性生成任务)。
memory 在循环外计算一次复用,是标准的推理优化。
每步只取 logits[:, -1:](最后一个位置的预测)作为新 token。
全部样本都生成 <EOS> 就提前停止。
7.3 Decoder-Only 的解码
把 prompt <BOS> src <EOS> 喂进去让模型续写:
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 from torch.nn.utils.rnn import pad_sequencedef build_decoder_only_prompt (src ): """ 把 encoder 输入的 src 序列转成 decoder-only 的 prompt。 src: (batch_size, src_len),含 PAD return: (batch_size, max_len + 2) 的 <BOS> src <EOS>(右侧 PAD 补齐) """ rows = [] for row in src: ids = row[row != PAD_IDX] bos = torch.tensor([BOS_IDX], dtype=torch.long, device=src.device) eos = torch.tensor([EOS_IDX], dtype=torch.long, device=src.device) rows.append(torch.cat([bos, ids, eos])) return pad_sequence(rows, batch_first=True , padding_value=PAD_IDX)@torch.no_grad() def greedy_decode_decoder_only (model, prompt, max_len=None ): """ decoder-only 模型的贪心解码。 prompt: (batch_size, prompt_len),格式 <BOS> src... <EOS>(可有右 PAD) return: (batch_size, prompt_len + gen_len),包含 prompt 在内的完整序列 """ model.eval () if max_len is None : max_len = prompt.size(1 ) + 5 ys = prompt for _ in range (max_len): logits = model(ys) next_token = logits[:, -1 :].argmax(dim=-1 ) ys = torch.cat([ys, next_token], dim=1 ) if (next_token == EOS_IDX).all (): break return ysdef expected_reverse_str (src_row ): """ 由一条 src 序列计算期望的反转结果字符串(用于 decoder-only 评估)。 src_row: (src_len,) token id 张量,含 PAD """ ids = [t.item() for t in src_row if t.item() != PAD_IDX] target = torch.tensor(ids[::-1 ] + [EOS_IDX], dtype=torch.long) return tensor_to_str(target)
生成完成后切掉 prompt 部分即得结果。
7.4 加载模型做反转
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 def load_model (checkpoint_path="checkpoint.pt" , device=None ): """加载保存的 checkpoint,按其中记录的架构构建并返回模型。""" device = device or torch.device("cuda" if torch.cuda.is_available() else "cpu" ) checkpoint = torch.load(checkpoint_path, map_location=device) cfg = checkpoint.get("config" , {}) arch = cfg.get("architecture" , "encoder_decoder" ) if arch == "decoder_only" : model = DecoderOnlyTransformer( vocab_size=13 , d_model=cfg["d_model" ], num_heads=cfg["num_heads" ], num_layers=cfg["num_decoder_layers" ], d_ff=cfg["d_ff" ], dropout=cfg["dropout" ], pad_idx=PAD_IDX, max_len=cfg["max_len" ], ).to(device) else : model = Transformer( src_vocab_size=13 , tgt_vocab_size=13 , d_model=cfg["d_model" ], num_heads=cfg["num_heads" ], num_encoder_layers=cfg["num_encoder_layers" ], num_decoder_layers=cfg["num_decoder_layers" ], d_ff=cfg["d_ff" ], dropout=cfg["dropout" ], pad_idx=PAD_IDX, max_len=cfg["max_len" ], ).to(device) model.load_state_dict(checkpoint["model_state_dict" ]) model.eval () return model, device, archdef digits_to_tensor (digits, device ): """把 0~9 的数字列表转成模型输入的 token id 张量 (1, len)。""" token_ids = torch.tensor([d + 3 for d in digits], dtype=torch.long) return token_ids.unsqueeze(0 ).to(device)@torch.no_grad() def reverse (model, digits, device, arch="encoder_decoder" ): """对一组数字列表做反转,返回预测出的数字列表。""" src = digits_to_tensor(digits, device) if arch == "decoder_only" : bos = torch.tensor([[BOS_IDX]], dtype=torch.long, device=device) eos = torch.tensor([[EOS_IDX]], dtype=torch.long, device=device) prompt = torch.cat([bos, src, eos], dim=1 ) pred = greedy_decode_decoder_only(model, prompt, max_len=src.size(1 ) + 5 ) pred_str = tensor_to_str(pred[0 , prompt.size(1 ):]) else : pred = greedy_decode(model, src, max_len=src.size(1 ) + 5 ) pred_str = tensor_to_str(pred[0 ]) return [int (x) for x in pred_str.split()] if pred_str else []if __name__ == "__main__" : model, device, arch = load_model() for digits in [[1 , 2 , 3 , 4 , 5 ], [9 , 8 , 7 , 6 ], [0 , 5 , 3 , 7 , 1 , 2 ]]: pred = reverse(model, digits, device, arch) print (f"input : {digits} " ) print ("expected :" , " " .join(str (d) for d in reversed (digits))) print ("pred :" , " " .join(str (d) for d in pred)) print ("correct :" , pred == list (reversed (digits)))
8. 超参数与训练效果
超参数
值
说明
D_MODEL
128
模型维度
NUM_HEADS
4
注意力头数,每头 d_k = 32
NUM_ENCODER_LAYERS
2
Encoder 层数
NUM_DECODER_LAYERS
4
Decoder 层数
D_FF
512
FFN 隐藏维度(4 × d_model)
DROPOUT
0.1
BATCH_SIZE
64
EPOCHS
9
LR
3e-4
Adam
TRAIN_NUM_SAMPLES
200,000
随机生成的训练样本数
序列长度
3 ~ 12
这个规模的模型(checkpoint 约 3MB)在 20 万随机样本上训练几个 epoch 后,测试集上的整序列准确率可以达到 ~100% ——模型完全学会了翻转规则,而不是背下了训练集(测试集种子不同、无重叠)。
9. 常见问题与思考
Q1:为什么训练收敛但推理时输出不对?
最常见的原因是 mask 错了。自查清单:① Decoder 自注意力是否用了 causal mask(训练时 teacher forcing 如果忘了 causal mask,模型会偷看答案,loss 极低但推理全错);② mask 的 1/0 约定是否和 masked_fill(mask == 0, -1e9) 一致;③ cross-attention 的 mask 应该用 src 的 padding mask,不是 tgt 的。
Q2:为什么 embedding 要乘 √d_model?
让 token embedding 与位置编码数值尺度相当。位置编码的值域是 [-1, 1],若 embedding 初始化得较小,位置信息会在相加时被淹没。
Q3:mask == 0 处为什么填 -1e9 而不是 -inf?
如果某一行所有 key 都被 mask(极端情况下),softmax(-inf, -inf, ...) 会得到 NaN 并沿计算图传播。-1e9 是一个足够小的有限值,softmax 后权重近似为 0 且不会产生 NaN。
Q4:翻转任务里模型学到了什么?
可以把注意力权重画出来观察:训练好的模型在 cross-attention 中,输出第 i 个位置时几乎把所有注意力权重放在输入的第 n-1-i 个位置上——它真正学会了"对称复制"这个算法,这正是 Transformer 位置外推能力的直观体现。
Q5:Encoder-Decoder 和 Decoder-Only 怎么选?
对本任务两者都能轻松收敛。Encoder-Decoder 结构更贴合"源/目标分离"的 Seq2Seq 设定,且推理时 Encoder 只跑一次,效率略高;Decoder-Only 实现更简单、与主流大语言模型(GPT 系列)同构,工程上更容易扩展。本文两种都实现了,改配置中的 ARCHITECTURE 即可对比。
Q6:可以继续改进的方向
把 MAX_LEN 调大(比如 50+),观察正弦位置编码的外推能力;
尝试 Pre-LN (LayerNorm 放在子层之前),训练更稳定,可以去掉 warmup;
实现 beam search 解码;
换成可学习的位置 embedding,对比效果;
用 nn.Transformer 或 F.scaled_dot_product_attention 对照验证自己手写的实现。