
Transformer 的论文我前后读了三遍PyTorch 版本的代码也重写过好几轮最开始是照着一份开源实现抄的训练跑起来 loss 死活不降折腾了整整两天才发现位置编码加在了 dropout 之后导致位置信息被随机丢掉了一大半。后来我养成一个习惯任何一份 Transformer 代码先不看模型结构先把forward里每一个张量的形状写在纸上对着走一遍出错的地方基本逃不掉。这篇 Transformer 代码解读就是把这些年攒下来的东西集中整理一遍从设计思路到逐行拆解再到训练跑通和踩坑排查全都摊开讲。这份代码适合谁看如果你已经能用 PyTorch 写个简单的 CNN 分类器但对 Transformer 的MultiHeadAttention里那几个view、transpose到底在干什么一直没搞明白或者想从零手写一份能跑起来的小型 Transformer那这篇会对你比较有用。我不打算讲太多数学推导公式该有的会有但重点放在“这行代码为什么这么写”“换个写法会出什么问题”上。环境层面Anaconda 配 PyTorch 环境、conda 和 pip 装包的区别、CPU 和 GPU 版本怎么选这些也会顺带说清楚毕竟我见过太多人卡在装环境这一步就没往下走。1. 整体设计思路这份代码为什么长这样先说结论一份能长期维护、方便调试的 Transformer 实现通常是按“组件化”拆的Embedding层、位置编码、注意力、前馈网络、编码器层、解码器层、整体模型每一层单独成类。这不是为了好看而是为了让出问题的时候能快速定位。1.1 从一次翻译任务看模型的输入输出契约在动手写代码之前我会先把输入输出的“契约”定死这一步比写代码本身重要。以最常见的机器翻译为例输入是一批 token id形状(batch_size, src_len)输出是另一批 token id形状(batch_size, tgt_len)。中间要经过的是一堆形状为(batch_size, seq_len, d_model)的浮点张量。这里的d_model是个关键超参原论文里 base 版本取 512big 版本取 1024。它决定了整个模型内部所有张量的“宽度”。为什么不是别的数因为它必须能被注意力头数整除512 能被 8 整除每个头分到 64 维如果取 5008 头就分不匀了。这是硬约束代码里通常会写一句assert d_model % num_heads 0别小看这一行它能省掉你后面无数次的形状报错。另一条契约是词表大小和输出维度。编码器和解码器共享同一套词嵌入权重是很常见的做法输出层再把d_model投回词表维度。这种权重共享在小数据集上能明显减少参数量还能让嵌入空间更一致我在几个中英翻译的小任务上都验证过共享之后收敛速度确实快一些。1.2 模块化拆分与“一个大类”的取舍我见过有人把整个 Transformer 写成一个巨大的nn.Moduleforward里几百行所有变量都是x、x2、x3。这种写法短期看着省事长期是灾难想单独测试注意力机制你没法测想看某一层输出你得在中间插 print。更合理的做法是这样拆PositionalEncoding只负责给嵌入加位置信息MultiHeadAttention只做 QKV 投影、缩放点积、拼接输出FeedForward就是两层线性加激活EncoderLayer把注意力、前馈、残差、LayerNorm 组装起来Encoder是 N 层EncoderLayer的堆叠。每个类的forward都控制在二三十行以内。这种拆法的另一好处是复用。比如想做只编码器的任务文本分类、序列标注直接拿Encoder堆几层加个池化头就能跑想做只解码器的任务语言模型拿DecoderLayer去掉交叉注意力那一块就是 GPT 风格的架构。一份拆得好的代码能覆盖后面大半年的实验需求。1.3 版本差异PyTorch 写法上的几个坑PyTorch 这几年版本迭代挺快有几个写法上的差异必须提。老版本里torch.nn.MultiheadAttention的默认 batch 维度在第一维早期版本是(seq_len, batch, dim)后来才在文档里明确batch_first参数如果你从网上抄了一段几年前的代码八成会在这上面栽跟头。另一个是torch.einsum的广泛应用。较新的实现里注意力计算常写成torch.einsum(bhqd,bhkd-bhqk, q, k)可读性好但要确认后端支持。我用下来多数情况下matmul配合transpose性能更稳尤其是序列比较长的时候。安装层面现在用 conda 建环境、pip 装 PyTorch 是最顺的路子。官网会给你一条带 CUDA 版本号的命令比如pip3 install torch torchvision --index-url https://download.pytorch.org/whl/cu121注意 CUDA 版本要和你显卡驱动匹配驱动太旧就装 CPU 版本功能上不影响只是慢。想省事就用conda install pytorch -c pytorch但有时候会拉一大堆依赖磁盘吃不消我一般还是 pip。2. 核心模块逐行拆解把代码拆到骨头缝里这一部分是全文的重头戏。我按数据流动的顺序来讲从最开始的词嵌入一路走到最后的输出投影中间每一行代码都尽量说透。2.1 词嵌入与位置编码为什么是加法而不是拼接嵌入层就一行nn.Embedding(vocab_size, d_model)。输入是(B, L)的 long 张量输出是(B, L, d_model)。有个细节是原论文会乘上sqrt(d_model)原因是嵌入初始化的方差比较小乘一个缩放因子能把它拉到和位置编码同一量级避免加起来之后位置信息被淹没。位置编码的实现长这样import torch import torch.nn as nn import math class PositionalEncoding(nn.Module): def __init__(self, d_model, max_len5000, dropout0.1): super().__init__() pe torch.zeros(max_len, d_model) position torch.arange(0, max_len, dtypetorch.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) # (1, max_len, d_model) self.register_buffer(pe, pe) # 随模型保存但不参与梯度 self.dropout nn.Dropout(dropout) def forward(self, x): # x: (B, L, d_model) x x self.pe[:, :x.size(1)] return self.dropout(x)几个要点。第一register_buffer是关键用普通的属性赋值张量不会跟着state_dict走模型保存再加载就丢了。第二pe用unsqueeze(0)加了 batch 维是为了利用广播不用手动expand。第三pe的切片只取前L个位置所以模型天然支持变长输入。为什么是“加法”而不是“拼接”拼接后维度翻倍后面所有线性层的参数量都要跟着涨。加法相当于把位置信息编码成一种“扰动”叠加在语义嵌入上模型有能力通过后续的线性变换把两者分开。这是原论文的取舍实测下来工作得不错后来几乎所有变体都沿用了这个思路。注意dropout 一定要放在加法之后。我最早那版代码把 dropout 加在x上位置编码加完后没做随机丢弃结果就是训练集和验证集的注意力图差异特别大。位置信息被丢掉之后模型对词序变得不敏感翻译出来的句子语序经常是乱的。2.2 多头注意力一次 reshape 就把 QKV 全搞定多头注意力的核心在于“分头”。最直观的写法是建 8 个独立的小注意力模块各自算完再拼接。但实际代码不会这么写因为那样效率太低而且每个头单独做一次矩阵乘法GPU 利用率上不去。标准做法是一次性做完投影然后 reshape 成多头class MultiHeadAttention(nn.Module): def __init__(self, d_model, num_heads, dropout0.1): super().__init__() assert d_model % num_heads 0 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 forward(self, query, key, value, maskNone): B, L, _ query.size() # (B, L, d_model) - (B, h, L, d_k) q self.w_q(query).view(B, -1, self.num_heads, self.d_k).transpose(1, 2) k self.w_k(key).view(B, -1, self.num_heads, self.d_k).transpose(1, 2) v self.w_v(value).view(B, -1, self.num_heads, self.d_k).transpose(1, 2) scores torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(self.d_k) if mask is not None: scores scores.masked_fill(mask 0, float(-inf)) attn torch.softmax(scores, dim-1) attn self.dropout(attn) out torch.matmul(attn, v) # (B, h, L, d_k) out out.transpose(1, 2).contiguous().view(B, -1, self.d_model) return self.w_o(out)view(B, -1, h, d_k).transpose(1, 2)这一串是理解重点。原始张量最后一维是d_model h * d_kview把它拆成(h, d_k)两维相当于把连续的 512 维切成 8 段各 64 维第 i 段就是第 i 个头的表示。transpose(1, 2)把头维度换到前面这样后面用matmul做批量矩阵乘法时(B, h, L, d_k) (B, h, d_k, L)天然就是每个头独立算。中间那个contiguous()是必须的。transpose只改元数据底层内存还是原来的顺序直接view会报错或者得到错误结果。contiguous()强制拷贝一份连续内存这个操作有开销但绕不过去。有个小优化是把transpose和contiguous合并写成.permute(0, 2, 1, 3).reshape(B, -1, d_model)效果一样。注意scores的缩放因子是1/sqrt(d_k)不是1/sqrt(d_model)。这一步很关键点积的方差随维度线性增长维度过大时 softmax 会进入饱和区梯度接近零。除以sqrt(d_k)把方差拉回 1 附近softmax 的输出更平缓梯度也更好。如果你改动过头的数量记得d_k跟着变缩放因子别写死。2.3 掩码的两种形态与 LayerNorm 的位置之争掩码有两种用途一定要分清。第一种是 padding mask用来屏蔽补齐的padtoken让模型不要关注这些无意义的位置。第二种是 causal mask也叫 look-ahead mask用在解码器的自注意力里保证第 t 个位置只能看到前 t 个位置不泄露未来信息。padding mask 的形状通常是(B, 1, 1, L)经过广播自动扩展到(B, h, L, L)。causal mask 是一个下三角矩阵形状(1, 1, L, L)用torch.tril(torch.ones(L, L))生成。两种掩码在训练时要做与运算合并def make_causal_mask(seq_len, device): return torch.tril(torch.ones(seq_len, seq_len, devicedevice)).bool() def combine_masks(pad_mask, causal_mask): # pad_mask: (B, 1, 1, L) causal_mask: (1, 1, L, L) return pad_mask causal_maskmasked_fill(mask 0, float(-inf))这一行的语义是mask 为 0 的位置填负无穷softmax 之后这些位置的概率就是 0。用-1e9也可以但float(-inf)更彻底不用担心数值溢出。有个坑是如果某一整行全是负无穷比如一个样本全被 pad 掉softmax 会输出 NaN。解决办法是给这一行留一个位置或者在后处理时把 NaN 置零。我在处理超长序列截断的时候踩过一次debug 了半天。LayerNorm 放哪儿是个老生常谈的问题。原论文是 Post-LN也就是LayerNorm(x Sublayer(x))后来大家发现 Post-LN 需要小心的学习率预热不然容易发散。Pre-LN 是x Sublayer(LayerNorm(x))训练稳定得多现在主流实现基本都用 Pre-LN。如果你是从零写代码我建议直接上 Pre-LN能省掉不少调参时间。代价是理论上表达能力略弱一点但实际任务里几乎看不出来。2.4 前馈网络与残差连接那些不起眼但重要的细节前馈网络结构很简单两层线性加一个激活class FeedForward(nn.Module): def __init__(self, d_model, d_ff, dropout0.1): super().__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): return self.linear2(self.dropout(torch.relu(self.linear1(x))))d_ff一般取4 * d_model。为什么是 4 倍原论文没说我理解是经验性的先升维到高维空间做非线性变换再降回来相当于给每个位置一个独立的“加工车间”。有研究比如 GLU 变体尝试过其他比例甚至门控结构效果差不多但 4 倍是默认起点。激活函数上原论文用 ReLU现在很多实现换成 GELU在几个任务上确实略好但差距不大。我一般先用 ReLU 跑通再对比 GELU不要一开始就纠结这个。残差连接在所有子层外面注意是x dropout(sublayer(x))dropout 加在子层输出上不是加在相加之后。位置错了会导致残差主干的信号被随机削弱训练后期 loss 震荡。这个点我在 review 别人代码时见过至少三次。3. 从零跑通完整实现与训练实操模块拆完了接下来是把它拼成能跑的完整流程。这一部分我会给出关键代码并说明每一步的选择理由。3.1 环境准备conda 建环境到验证 GPU先说环境因为卡在这的人实在太多。流程是装 Anaconda 或 Miniconda 作为环境管理器建一个独立环境再装 PyTorch。conda create -n transformer python3.10 -y conda activate transformer # 到 PyTorch 官网拿到匹配你机器的安装命令例如 pip3 install torch torchvision --index-url https://download.pytorch.org/whl/cu121Python 版本建议 3.9 到 3.11太新的版本有时候第三方包还没跟上。CUDA 版本看显卡驱动驱动版本较旧就装 CPU 版命令里把cu121换成cpu对应的索引。装完之后一定要验证import torch print(torch.__version__) print(torch.cuda.is_available()) # 期望 True print(torch.cuda.get_device_name(0))如果cuda.is_available()返回 False先查驱动再查是不是装成了 CPU 版本。nvidia-smi显示的 CUDA 版本是驱动支持的上限不是你必须装的版本装低于它的都行。编辑器我用 VS Code 配 Jupyter 插件比较多调试小片段方便写大项目就切 PyCharm。这俩都能识别 conda 环境在设置里选对解释器就行。3.2 数据管道与批处理padding 的坑最多数据这块最麻烦的是变长序列的 padding 和对应的 mask 生成。我一般写一个collate_fnPAD_ID 0 BOS_ID 1 EOS_ID 2 def collate_batch(batch, max_len128): srcs, tgts zip(*batch) src_lens [min(len(s), max_len) for s in srcs] tgt_lens [min(len(t) 2, max_len) for t in tgts] src_pad torch.full((len(srcs), max(src_lens)), PAD_ID, dtypetorch.long) tgt_pad torch.full((len(tgts), max(tgt_lens)), PAD_ID, dtypetorch.long) for i, (s, t) in enumerate(zip(srcs, tgts)): s s[:max_len] t t[:max_len - 2] src_pad[i, :len(s)] torch.tensor(s) tgt_pad[i, 0] BOS_ID tgt_pad[i, 1:len(t) 1] torch.tensor(t) tgt_pad[i, len(t) 1] EOS_ID return src_pad, tgt_pad这里解码器输入和目标输出要错开一位输入是bos w1 w2 ... wn目标是w1 w2 ... wn eos。这是 teacher forcing 的标准做法错位那一行代码写错了模型永远学不会生成loss 会降到一个平台期就卡住。我第一次写的时候就忘了错位loss 降到 2 左右就再也下不去查了两小时。mask 的生成要在 batch 之后def make_pad_mask(seq, pad_idPAD_ID): return (seq ! pad_id).unsqueeze(1).unsqueeze(2) # (B,1,1,L)注意unsqueeze两次把形状变成(B, 1, 1, L)这样和(B, h, L, L)的 scores 广播时才是按最后一维屏蔽也就是屏蔽掉 key 方向上的 pad。如果只 unsqueeze 一次形状变成(B, 1, L)广播结果完全不同报错或者静默算错。3.3 训练循环与学习率调度warmup 别省优化器用 Adambetas(0.9, 0.98)、eps1e-9这是原论文的设置跟默认值略有不同对稳定性有帮助。学习率调度用 Noam 那套带预热的方案class NoamScheduler: def __init__(self, optimizer, d_model, warmup_steps4000, factor1.0): self.optimizer optimizer self.d_model d_model self.warmup_steps warmup_steps self.factor factor self.step_num 0 def step(self): self.step_num 1 lr self.factor * (self.d_model ** -0.5) * min( self.step_num ** -0.5, self.step_num * self.warmup_steps ** -1.5 ) for group in self.optimizer.param_groups: group[lr] lr return lr这个公式的行为是前期学习率线性上升到warmup_steps附近达到峰值之后按步数的平方根倒数衰减。预热阶段很关键它让模型在参数还比较随机的时候不要迈太大的步子。如果把 warmup 去掉直接用固定的 1e-3我在小数据集上试过前几百步 loss 经常直接飙到 NaN之后就再也回不来了。一个完整的训练步大概是这样model.train() src, tgt src.to(device), tgt.to(device) tgt_in, tgt_out tgt[:, :-1], tgt[:, 1:] src_mask make_pad_mask(src) tgt_pad_mask make_pad_mask(tgt_in) causal make_causal_mask(tgt_in.size(1), device) tgt_mask tgt_pad_mask causal logits model(src, tgt_in, src_mask, tgt_mask) loss criterion( logits.reshape(-1, logits.size(-1)), tgt_out.reshape(-1) ) optimizer.zero_grad() loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0) optimizer.step() scheduler.step()clip_grad_norm_我基本没关过。Transformer 的梯度偶尔会冒尖裁剪到 1.0 能避免单步把参数带偏。损失函数用CrossEntropyLoss(ignore_indexPAD_ID, label_smoothing0.1)ignore_index让 pad 位置不参与损失标签平滑是原论文的做法能防止模型对训练标签过度自信验证集指标通常更好看。3.4 自回归推理greedy 与 beam search 的实现差异推理阶段没有目标序列得一个 token 一个 token 生成。最简单的贪心解码torch.no_grad() def greedy_decode(model, src, src_mask, max_len64, bos_idBOS_ID, eos_idEOS_ID): model.eval() memory model.encode(src, src_mask) ys torch.full((src.size(0), 1), bos_id, dtypetorch.long, devicesrc.device) for _ in range(max_len - 1): causal make_causal_mask(ys.size(1), src.device) out model.decode(memory, src_mask, ys, causal) next_token out[:, -1].argmax(dim-1, keepdimTrue) ys torch.cat([ys, next_token], dim1) if (next_token eos_id).all(): break return ys每次循环都重算整个解码序列的注意力效率不高但逻辑清晰。生产环境里会缓存 K、V也就是常说的 KV cache能显著减少重复计算。这个优化先不做等模型跑通、指标对了再上不然出问题的时候你分不清是缓存的锅还是模型的锅。beam search 是在贪心基础上的扩展每步保留 top-k 个候选路径最后选整体对数概率最高的。代码量大概翻三倍小任务上提升有限我一般贪心先跑指标不满意再换。4. 踩坑实录那些官方代码不会告诉你的问题前面都是“正确姿势”这一节讲翻车现场。我把这些年遇到的高频问题整理成一张速查表再挑几个典型的展开说。4.1 形状报错与静默算错的区分方法现象大概率原因排查动作viewsize mismatchtranspose 后没 contiguous换成reshape或补.contiguous()训练 loss 不降但无报错目标序列没错位 / mask 语义反了打印一组样本和标签逐 token 核对输出全是同一个词解码只用了最后一个位置检查out[:, -1]索引验证指标远差于训练dropout 位置错 / 忘了model.eval()逐层核对 dropout 调用点长句生成质量骤降超过 max_len 位置编码被截断打印pe切片后的实际长度形状报错好办报错信息里通常明确指出了哪两个维度对不上。真正难缠的是“静默算错”代码能跑loss 也降但模型学的不是你想让它学的东西。判断方法是拿一个极小的数据集比如 20 个句子去过拟合正常实现应该能在几百步内把 loss 压到接近 0如果压不下去结构上一定有 bug。这招我用得最多比读代码快得多。mask 的语义反了是经典错误。masked_fill(mask 0, -inf)的语义是“mask 为 0 的地方屏蔽掉”。如果你生成 mask 时写成了(seq ! pad_id)之外的条件或者传参时把 padding mask 和 causal mask 位置传反了模型可能看到未来信息训练 loss 会低得不正常但在真实推理时表现极差。4.2 训练不收敛的排查顺序我一般按这个顺序查先看数据再看 mask再看学习率最后看初始化。数据层面先确认标签和输入是否对齐尤其是自己写Dataset的时候。有一次我把源语言和目标语言在预处理阶段搞反了模型一直在学“把英文翻译成英文”loss 降到某个值就停了。mask 层面用一个小 batch 把scores打印出来看看被屏蔽的位置是不是负无穷没被屏蔽的位置是不是合理范围。如果发现整行都是负无穷那就是某个样本被全部 pad 掉了。学习率层面注意 warmup 是否生效。可以写个小脚本单独跑 scheduler把前几千步的学习率画出来看曲线是不是先升后降。如果一直是常数说明group[lr]没被真正更新可能是优化器创建时没传对参数组。初始化层面PyTorch 的nn.Linear默认用 Kaiming 均匀初始化对 Transformer 来说够用。但如果你的模型很深比如 12 层以上又用 Post-LN可能需要调小初始化的标准差。这部分不确定的话直接换 Pre-LN 更省心。4.3 显存与速度从注意力实现到混合精度显存不够是训练时的常见拦路虎。几个立竿见影的手段把batch_size减半把max_len砍到实际需要的最长值开启梯度累积模拟大 batch。更彻底的做法是换 attention 实现。torch.nn.functional.scaled_dot_product_attention在较新的 PyTorch 里已经有了底层会根据硬件自动选 FlashAttention 之类的融合算子速度快、显存省。用法上把 Q、K、V 和 mask 传进去就行注意它对 mask 的形状要求是布尔张量语义是 True 表示保留。混合精度训练也是常用手段torch.cuda.amp配合GradScaler一般能省 30% 到 40% 显存速度快一点。要注意的是 softmax 和 LayerNorm 这些对数值敏感的部分autocast 通常会自动处理但如果你自己写了自定义的算子最好手动指定float32。序列长度是显存的最大变量。注意力矩阵是L * L长度翻倍显存翻四倍。如果任务里句子普遍很长优先考虑是不是能做截断或者分块处理而不是一味加显存。5. 往后走这份代码还能怎么扩展代码跑通只是起点后面大概率要改。这一节说说扩展方向。5.1 从文本到视觉结构上要动哪些地方把 Transformer 用到图像上最典型的是 ViT。改动其实不大把词嵌入换成图像分块嵌入用卷积或者 unfold 把(B, 3, H, W)变成(B, N, d_model)N 是分块数量位置编码从固定正弦改成可学习参数因为图像块之间没有天然的序列顺序假设分类头从解码器换成简单的池化加线性。再往后是 Swin Transformer核心是引入窗口注意力和层级结构把注意力限制在局部窗口内配合窗口偏移实现跨窗口信息交换计算量从平方级降到线性级。最近还有 HGFormer 这类把超图学习引入视觉 Transformer 的工作思路是在注意力之外显式建模高阶关系代码结构上多了一个超图构建和消息传递模块但底层的注意力、残差、LayerNorm 这些组件和本文讲的完全一致。看懂这份基础实现再看这些变体会轻松很多。5.2 什么时候该扔掉手写实现手写实现的价值在于理解不在于生产。当你需要大规模训练、需要各种预训练权重、需要现成的分布式支持时直接用 HuggingFace 的transformers库更划算。它的BertModel、GPT2Model这些类已经把各种变体封装好了加载权重、保存、导出都很方便。但即使你后面一直用库我仍然建议至少手写一遍。原因很实际出问题的时候库的报错信息往往指向框架内部如果你知道forward里面每一步在做什么定位速度快很多。我自己就遇到过因为输入没有正确加特殊 token 导致效果差的情况库不会有任何提示但你知道模型的输入契约是什么一眼就能看出问题。另外手写实现是改结构的前提。比如你想在注意力里加一个相对位置偏置或者把前馈层换成门控结构用库的话得改源码或者写复杂的 hook用手写代码就是改几行的事。最后分享一个我在调试 Transformer 时的固定习惯写一个shape_check函数把每一层的输入输出形状打印一遍第一次跑通之前不关掉。听起来笨但它帮我省下的时间比我读一百篇解读文章都多。模型训不动的时候先别怀疑理论先把形状和 mask 过一遍九成的问题都在那里。