
做Transformer这块我自己最早踩的坑恰恰不在注意力公式上而在张量形状上。公式背得滚瓜烂熟一上手写代码(N, L, d_model)和(L, N, d_model)来回切mask少了一维广播不上decoder的tgt忘了右移一位loss算出来还挺像那么回事只是永远不收敛。后来复盘才发现Transformer的难点从来不在多头注意力这四个字而在于输入输出的形状纪律数据从哪进、经过谁、变成什么、再从哪出。这篇就把Transformer的输入输出细节和pytorch代码实现从零拆一遍重点放在那些官方文档一笔带过、但真正写代码时必然要面对的地方。适合已经看过架构图、准备自己手撸一遍的中级读者也适合被维度错误折磨过、想系统理清形状关系的人。1. 从张量形状出发Transformer的输入输出是一条三进三出的流水线很多人学Transformer的顺序是先看QKV、再看attention、最后回头看输入输出结果就是每个模块单独都懂拼起来就散架。我建议的顺序反过来先把整条数据流水线的形状定死再往里填模块。因为模块是局部的形状是全局的形状错了模块写得再对也白搭。1.1 三个入口张量src、tgt、以及那两张mask一个标准的Encoder-Decoder Transformer前向传播的入口其实只有两类东西token id序列和掩码。token id序列有两条一条是源端src一条是目标端tgt。掩码同样有两条一条管源端padding一条管目标端padding加因果。在batch_firstTrue的约定下我强烈建议整个项目统一用这个约定形状是这样src(N, S)N是batch sizeS是源端序列长度元素是词表索引的整数tgt(N, T)T是目标端序列长度src_key_padding_mask(N, S)布尔或浮点True表示该位置是paddingtgt_key_padding_mask(N, T)tgt_mask或叫attn_mask(T, T)因果上三角注意这里有个别扭的地方padding mask在入口是二维的但它最终要作用在四维的attention分数上。中间那次变形就是新手第一个大坑后面第4节专门讲。提示如果你用的是nn.Transformer原生模块默认batch_firstFalse也就是(L, N, E)。翻官方文档时看到(S, N, E)别慌那只是序列维度在前面而已逻辑完全一样。混用两种约定是维度报错的第一大来源项目里选定一种就别改。1.2 输出端logits的形状决定了损失函数怎么写Encoder-Decoder的输出是(N, T, V)V是目标端词表大小。这个张量在代码里通常叫logits或output是未经softmax的原始分数不是概率。这一点决定了你会用nn.CrossEntropyLoss而不是nn.NLLLoss因为CrossEntropyLoss内部自带log_softmax。那为什么不是(N, T, d_model)因为d_model只是模型内部的隐藏维度要让模型说出词表里的某个词必须再经过一次线性投影把d_model映射成V。这层投影在经典实现里叫generator或output_projection并且和输入端的embedding共享权重——这是原论文的做法理由是输入端学到的词向量语义空间和输出端要预测的词空间应该一致共享能省参数量还能缓解小语料上的过拟合。所以完整链路是src (N,S) --embed-- (N,S,d_model) --encoder-- memory (N,S,d_model) tgt (N,T) --embed-- (N,T,d_model) --decoder(cross-attn取memory)-- (N,T,d_model) --output_projection-- logits (N,T,V)记住memory这个词PyTorch里Encoder的输出就叫memory它的形状是(N, S, d_model)序列长度跟源端走不跟目标端走。很多人写交叉注意力时把query和key搞混就是因为忘了memory的长度是S不是T。1.3 一张表看清所有中间张量的形状演变把关键节点列成表格写代码时对着查比在脑子里推快得多阶段名称形状batch_first备注输入src / tgt(N, S) / (N, T)整数id嵌入后src_emb / tgt_emb(N, S, d_model) / (N, T, d_model)已乘sqrt(d_model)位置编码后同上不变与pe逐元素相加Encoder输出memory(N, S, d_model)长度跟SDecoder输出dec_out(N, T, d_model)长度跟T最终logits(N, T, V)未softmax损失loss标量reshape成(N*T, V)我见过有人在Decoder输出后直接接argmax忘了过output projection结果在d_model维度上取最大值的索引得到一堆毫无意义的数字。这不是逻辑错误是纯粹漏了一步但现象上表现为模型输出全是垃圾非常难查。2. 词嵌入与位置编码输入端最容易埋雷的两步输入端的处理只有两步embedding查表和加位置编码。看起来简单但这两步各自都有一个隐藏设定漏掉任何一个模型都可能训练得起来但效果明显打折。2.1 为什么embedding要乘sqrt(d_model)原论文里有一个不起眼但很关键的细节词嵌入向量要乘以$\sqrt{d_{model}}$。原因和后面位置编码的幅度有关。正弦位置编码的取值在[-1, 1]之间而用标准初始化比如正态分布N(0,1)的embedding每个元素的方差是1维度累加后向量的模长随$\sqrt{d_{model}}$增长。如果embedding不缩放两者相加时位置编码的信息会被embedding的幅值淹没位置信号几乎不起作用。乘以$\sqrt{d_{model}}$之后两者的量级被拉到同一个区间相加才有意义。代码上就是一行class TokenEmbedding(nn.Module): def __init__(self, vocab_size, d_model, pad_idx0): super().__init__() self.embed nn.Embedding(vocab_size, d_model, padding_idxpad_idx) self.d_model d_model self.scale math.sqrt(d_model) def forward(self, x): # x: (N, L) - (N, L, d_model) return self.embed(x) * self.scale注意padding_idx参数不只是标记哪个id是padding它还会让这个位置的embedding永远不参与梯度更新并且初始化时置零。这正好符合我们的需求padding位置不携带语义。但代价是如果你把padding_idx设成了某个真实词的id那个词就永远学不到了。所以padding_idx必须是词表里专门预留的0号位。2.2 正弦位置编码的推导与register_buffer的意义位置编码的公式是$$PE_{(pos, 2i)} \sin(pos / 10000^{2i/d_{model}})$$ $$PE_{(pos, 2i1)} \cos(pos / 10000^{2i/d_{model}})$$偶数维度用sin奇数维度用cos。为什么这么设计直观理解是不同维度对应不同频率的波低维是高频波长短能区分相邻位置高维是低频波长长能区分远距离位置。这有点像用二进制编码表示数字低位变化快、高位变化慢组合起来就能唯一表示每个位置。另外还有个数学性质PE(posk)可以表示成PE(pos)的线性变换这让模型有可能学会相对位置的概念——尽管位置编码本身是绝对的。实现上有个细节值得说位置编码矩阵应该用register_buffer注册而不是普通属性。因为它是固定不变的常数不该被优化器更新也不该出现在state_dict里作为可训练参数。用register_buffer还自带一个好处调用.cuda()或.to(device)时它会跟着一起搬不用手动处理。class PositionalEncoding(nn.Module): def __init__(self, d_model, max_len5000, dropout0.1): super().__init__() self.dropout nn.Dropout(dropout) pe torch.zeros(max_len, d_model) position torch.arange(0, max_len, dtypetorch.float).unsqueeze(1) # (max_len, 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) # 奇数列 self.register_buffer(pe, pe.unsqueeze(0)) # (1, max_len, d_model) def forward(self, x): # x: (N, L, d_model) x x self.pe[:, :x.size(1)] return self.dropout(x)这里div_term用了指数形式的等价写法$10000^{-2i/d} \exp(-2i/d \cdot \ln 10000)$比直接做幂运算数值上更稳也更快。另外x.size(1)取的是当前序列长度实现按需切片这样max_len设大一点也不占额外显存。2.3 一个容易忽略的点位置编码作用于哪一端Encoder和Decoder都要加位置编码而且用的是同一个PositionalEncoding实例参数是常数共享完全安全。这一点常被忽略导致有人只在Encoder加了位置编码Decoder端就是一堆无序的词向量效果直接崩盘。还有一个问答社区里经常出现的疑问为什么Decoder端的输入也要位置编码因为它要自己做因果自注意力不给位置信息就不知道谁在谁前面因果mask也就失去了顺序的含义——mask只是机械地屏蔽右上角但模型得知道左边是过去、右边是未来这个语义得靠位置编码提供。3. 多头注意力的输入输出Q、K、V到底从哪来到这一节形状的复杂程度达到峰值。多头注意力的输入输出其实非常简洁——输入三个张量输出一个张量加一个注意力权重。但内部的维度变换有好几处任何一处顺序写错要么报错要么静默地算错。3.1 三种注意力在输入端的差别同一个MultiHeadAttention模块在Transformer里被用了三次差别全在Q、K、V来自哪里位置Query 来源Key 来源Value 来源maskEncoder自注意力srcsrcsrcsrc padding maskDecoder自注意力tgttgttgttgt padding 因果Decoder交叉注意力tgtmemorymemorysrc padding mask这三个用法决定了整个模型的语义Encoder自注意力让源端词互相看Decoder自注意力让目标端词看自己左边的历史Decoder交叉注意力让目标端每个位置去源端找相关信息。有个特别容易搞混的点交叉注意力的mask用的是源端的padding mask不是目标端的。因为被mask掉的是Key和Value所在的源端序列。如果你在这里传了目标端的mask形状对不上会报错但如果长度恰好相等比如src和tgt长度都是20它会静默地算错loss就是降不下去。这种能跑但不对的bug最恶心。3.2 view加transpose的维度变换为什么不能写错多头注意力的核心操作是把d_model拆成n_head × d_k让每个头在低维子空间里独立做注意力。变换顺序是# 输入 (N, L, d_model) q self.w_q(query) # (N, L, d_model) q q.view(N, -1, self.n_head, self.d_k) # (N, L, n_head, d_k) q q.transpose(1, 2) # (N, n_head, L, d_k)关键在view的维度顺序必须是(N, L, n_head, d_k)让d_model维度按头优先的方式切开即前d_k个元素属于第0个头接着d_k个属于第1个头。如果你写成(N, n_head, L, d_k)再直接view因为原始内存里d_model是连续的这样切出来的分组就完全乱了——每个头拿到的是一组不相邻的维度。这个错误的可怕之处在于它不会报错。形状对得上计算能走通只是注意力的语义被打乱了。而且因为多头注意力的各路输出最后会拼接回d_model即使分组错了线性层依然能学到一些东西loss也会下降只是最终指标比正确实现低几个点。这种慢性病在调试时几乎不可能通过看代码发现。所以我的习惯是先transpose再reshape或者用view时严格确认维度顺序。上面的写法先view成(N, L, n_head, d_k)再transpose是安全的。最后拼回来的时候要contiguous()x x.transpose(1, 2).contiguous().view(N, -1, self.n_head * self.d_k)因为transpose之后内存不连续直接view会报错必须contiguous()把它变回连续内存。3.3 除以sqrt(d_k)的位置与数值稳定性缩放点积注意力的公式是$\text{softmax}(QK^T / \sqrt{d_k})V$。为什么要除$\sqrt{d_k}$因为点积的结果是d_k个乘积之和每个乘积的方差约为1累加后方差变成d_k标准差变成$\sqrt{d_k}$。当d_k较大时比如64点积的数值会很大softmax之后会极其尖锐梯度趋近于0训练停滞。除以$\sqrt{d_k}$把方差拉回1softmax处于一个健康的温度区间。实现上缩放可以在两个位置做一是在算完scores之后除二是预先对Q做缩放。两种等价我习惯在scores之后除可读性更好def scaled_dot_product_attention(q, k, v, maskNone, dropoutNone): d_k q.size(-1) scores torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(d_k) if mask is not None: scores scores.masked_fill(mask, float(-inf)) p_attn scores.softmax(dim-1) if dropout is not None: p_attn dropout(p_attn) return torch.matmul(p_attn, v), p_attn注意masked_fill用的是float(-inf)而不是一个大负数。用-inf的好处是softmax之后严格为0不引入任何泄漏。但它有个坑如果某一行的所有位置都被mask了整行都是-infsoftmax会产生NaN。这个问题在padding mask和因果mask叠加时可能触发第4节细说。另外注意softmax(dim-1)最后一维是Key的序列长度。如果写成dim1就会在batch维度上做softmax结果是错的但形状完全正确——又是一次静默错误。4. 两张mask的形状陷阱padding mask和causal mask必须能广播mask这块我愿称之为Transformer实现里的深水区因为它的形状不是自然从数据流里长出来的而是从广播规则反推出来的。4.1 padding mask为什么是(N,1,1,L)先看attention分数的形状(N, n_head, L_query, L_key)。mask最终要和它逐元素作用所以mask必须能广播到这个形状。padding mask的原始信息是(N, L_key)——每个batch、每个key位置是不是padding。要广播到(N, n_head, L_query, L_key)需要把中间两个维度补成1(N, 1, 1, L_key)。def make_pad_mask(seq, pad_idx0): # seq: (N, L) - (N, 1, 1, L) return (seq pad_idx).unsqueeze(1).unsqueeze(2)补的这两个1是有语义的n_head维度补1表示所有头共享同一份mask正确padding跟头无关L_query维度补1表示所有query位置都屏蔽同一批key这个要看情况padding mask确实是这样一个性质——不管query是谁都不该关注padding的key。对比一下PyTorch原生nn.MultiheadAttention的设计它的key_padding_mask是(N, S)二维的内部帮你做广播。自己写的时候就没有这个便利得手动补维度。两条路线都对但自己写就得把广播规则想清楚。4.2 causal mask屏蔽的是右上角那一半因果mask让第i个位置只能看到第0到第i个位置。看attention矩阵(L_query, L_key)行是query、列是key那么被屏蔽的区域是列索引大于行索引的部分也就是右上三角。def make_causal_mask(size, device): # True 表示屏蔽, 形状 (1, 1, size, size) return torch.triu(torch.ones(size, size, devicedevice), diagonal1).bool().unsqueeze(0).unsqueeze(0)torch.triu(..., diagonal1)取的是严格上三角不含对角线也就是保留对角线以下为False、以上为True。diagonal0会把对角线也屏蔽掉那样每个位置连自己都看不到违反直觉。这里有个需要记住的约定差异在masked_fill里True应该表示屏蔽/丢弃。如果你用torch.tril生成下三角然后直接masked_fill(mask, -inf)那就把能看的部分全屏蔽了、该屏蔽的反而保留了结果是每个位置只能看到未来的词模型直接学废。我自己就因为烙印了下三角是因果mask这个印象翻过车因为掩码库比如HuggingFace里attention_mask的语义又是不一样的前者是1保留后者是True屏蔽。每次写mask前先确认当前框架的语义这一句话能省下几个小时的调试。4.3 两种mask合并与全屏蔽行的NaN问题目标端自注意力需要同时应用因果mask和padding mask。合并方式是逻辑或因为两者都是True即屏蔽# tgt: (N, T) pad_mask make_pad_mask(tgt, pad_idx) # (N, 1, 1, T) causal make_causal_mask(T, tgt.device) # (1, 1, T, T) tgt_mask pad_mask | causal # 广播到 (N, 1, T, T)广播是自动的(N,1,1,T)和(1,1,T,T)或运算后得到(N,1,T,T)再和注意力分数(N, n_head, T, T)作用时在head维度广播。整个过程不需要手动扩展。NaN问题的触发条件是某个query位置的所有key都被屏蔽。什么时候会发生如果序列以padding开头。假设一个batch里第一条样本长度是0纯padding那么第0个query位置面对的所有key都是padding被padding mask全部屏蔽加上因果mask之后仍然是全屏蔽softmax得到0/0 NaN。现实中完全长度为0的样本不常见但如果你的数据预处理里有截断逻辑、或者用了某种特殊的分桶策略就可能出现。稳妥的做法有两种一是用float(-inf)改为一个大负数比如-1e9这样softmax后是均匀分布而不是NaN二是在生成mask后检查并修正全屏蔽行。我一般直接用-1e9代价是理论上有一点点概率泄漏exp(-1e9)实际为0实践上完全够用。另一个细节混合精度训练时float(-inf)在fp16下容易溢出成NaN这也是我推荐-1e9的原因之一。如果你的训练脚本开了torch.cuda.amp这个坑很容易撞上。5. Encoder与Decoder的输入差异训练走teacher forcing推理走自回归前面都是单次前向传播的形状问题。真正让Transformer训练和推理走两条路的是Decoder的输入构造方式。5.1 训练时tgt为什么必须右移一位训练时我们有完整的参考译文tgt [BOS, w1, w2, ..., wn, EOS]。Decoder要做的是给定前面的词预测下一个词所以输入和标签来自同一句话只是错开一位tgt_input tgt[:, :-1] # [BOS, w1, ..., wn] 长度 T-1 tgt_label tgt[:, 1:] # [w1, w2, ..., EOS] 长度 T-1输入是[BOS, w1, ..., wn]标签是[w1, w2, ..., EOS]。第i个位置的输入是到第i个词为止的历史输出要预测第i1个词。这就是所谓的teacher forcing——用真实的前一个词作为输入而不是模型自己上一步的预测。配合因果mask这一句话的所有预测可以在一次前向传播里并行算出来第0个位置算P(w1|BOS)第1个位置算P(w2|BOS,w1)依此类推。因果mask保证了第1个位置的计算不会偷偷看到第2个位置的输入。这也正是Transformer相对于RNN训练效率高的核心原因。形状上Decoder的输出是(N, T-1, V)标签是(N, T-1)。算损失时把两者都展平loss criterion( logits.reshape(-1, logits.size(-1)), # (N*(T-1), V) tgt_label.reshape(-1) # (N*(T-1),) )criterion要带ignore_indexpad_idx让padding位置的损失被忽略。label_smoothing0.1也是常用配置原论文用了0.1能缓解模型对训练标签过度自信。提示如果忘记右移即拿tgt [BOS, w1, ..., wn]同时当输入和标签模型会学到直接把输入复制到输出这种平凡的恒等映射。表现是训练loss极低、快得离谱但推理时完全不会翻译。这是最隐蔽的一种错误因为所有指标都很好看。5.2 推理时的自回归循环与KV Cache的引入动机推理时没有参考译文只能一步步生成torch.no_grad() def greedy_decode(model, src, src_mask, max_len, bos_idx, eos_idx, device): model.eval() memory model.encode(src, src_mask) ys torch.full((src.size(0), 1), bos_idx, dtypetorch.long, devicedevice) for _ in range(max_len - 1): tgt_mask make_causal_mask(ys.size(1), device) out model.decode(memory, src_mask, ys, tgt_mask) # (N, len, d_model) logits model.generator(out[:, -1]) # 只取最后一步 (N, V) next_word logits.argmax(dim-1, keepdimTrue) # (N, 1) ys torch.cat([ys, next_word], dim1) if (next_word eos_idx).all(): break return ys注意这里的几个关键点。第一out[:, -1]只取最后一个位置的输出因为只有它对应下一个词的预测前面的位置在上一轮已经算过并且用过了。第二每生成一个词序列长度加一因果mask要重新构造以匹配新长度。第三循环次数上限是max_len。这个朴素实现有个明显的性能问题每生成一个词都要把整个前缀重新过一遍Decoder。生成长度为T的句子计算量是$O(T^2)$级别的重复。KV Cache的思路很直接自注意力里已经算过的位置的Key和Value是不会变的因为因果mask保证了它们只依赖自己及更早的位置。所以可以把每层、每个头的K和V缓存下来生成新词时只算新位置的Q、K、V其中K和V拼接到缓存上注意力只在新Q和全部缓存KV之间计算。class CachedMultiHeadAttention(nn.Module): def forward(self, query, key, value, maskNone, cacheNone): # 训练时 cache 为 None走正常路径 # 推理时把新算出的 k, v 与 cache 里的拼接 if cache is not None: k_new self.w_k(key).view(N, -1, self.n_head, self.d_k).transpose(1, 2) v_new self.w_v(value).view(N, -1, self.n_head, self.d_k).transpose(1, 2) k torch.cat([cache[0], k_new], dim2) # 在序列维拼接 v torch.cat([cache[1], v_new], dim2) new_cache (k, v) ...KV Cache能把推理从$O(T^2)$的重复计算降到接近$O(T)$的增量计算代价是显存占用随序列长度线性增长。现在主流的推理框架里KV Cache几乎是必选项理解了它的形状每层一份形状(N, n_head, 已生成长度, d_k)自己实现就不难。5.3 BOS、EOS、PAD三个特殊token在各处的角色三个特殊token在输入输出里各司其职混用会出问题PAD通常id为0。用于对齐batch内不等长的序列出现在输入端和标签里。损失函数里必须ignore_index掉否则模型会浪费容量去学预测padding。注意PAD在标签里出现时一定是句尾的一串不会出现在中间。BOS只有一个出现在Decoder输入的第0位作为生成起点。它从不出现在标签里因为没有任何词应该被预测成句子开始。EOS出现在标签的最后一位模型要学会在这里停下。在Decoder输入里EOS也出现在最后一位之前因为是teacher forcing模型需要看到前文完整的、包括还没结束的上下文。但在标签里它只是最后那个目标。整理成表格更清楚以[BOS, w1, w2, EOS]为例序列内容长度tgt原始BOS w1 w2 EOS PAD PADT2paddingtgt_inputBOS w1 w2 EOS PADT2tgt_labelw1 w2 EOS PAD PADT2可以看到tgt_label里EOS后面的PAD被ignore_index忽略只有EOS本身在参与计算这正好对应模型要学会在合适的时候输出EOS。6. 完整代码一个能直接跑通的最小Transformer前面拆了这么多细节现在拼成一个完整可运行的实现。我刻意写得扁一些不追求工程优雅但求每一步的形状都清晰可查。6.1 位置编码与掩码工具这部分前面已经给过这里补一个统一的掩码工具类避免在多个地方重复构造def make_pad_mask(seq, pad_idx0): # (N, L) - (N, 1, 1, L)True 表示屏蔽 return (seq pad_idx).unsqueeze(1).unsqueeze(2) def make_causal_mask(size, device): # (1, 1, size, size)True 表示屏蔽右上三角 return torch.triu( torch.ones(size, size, devicedevice, dtypetorch.bool), diagonal1 ).unsqueeze(0).unsqueeze(0) def make_tgt_mask(tgt, pad_idx0): T tgt.size(1) pad make_pad_mask(tgt, pad_idx) # (N, 1, 1, T) causal make_causal_mask(T, tgt.device) # (1, 1, T, T) return pad | causal # (N, 1, T, T)我把make_causal_mask的diagonal参数单独拎出来写成1是为了在代码里留下一个明显的锚点以后谁看到都能立刻确认这是严格上三角屏蔽未来。6.2 多头注意力与逐层连接class MultiHeadAttention(nn.Module): def __init__(self, d_model, n_head, dropout0.1): super().__init__() assert d_model % n_head 0, d_model 必须能被 n_head 整除 self.d_k d_model // n_head self.n_head n_head 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): N query.size(0) q self.w_q(query).view(N, -1, self.n_head, self.d_k).transpose(1, 2) k self.w_k(key).view(N, -1, self.n_head, self.d_k).transpose(1, 2) v self.w_v(value).view(N, -1, self.n_head, self.d_k).transpose(1, 2) # q,k,v: (N, n_head, L, d_k) scores torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(self.d_k) if mask is not None: scores scores.masked_fill(mask, -1e9) p_attn scores.softmax(dim-1) p_attn self.dropout(p_attn) x torch.matmul(p_attn, v) # (N, n_head, L, d_k) x x.transpose(1, 2).contiguous().view(N, -1, self.n_head * self.d_k) return self.w_o(x), p_attnSublayerConnection用Pre-LN结构也就是先LayerNorm再进子层残差直连。Pre-LN比原论文的Post-LN更好训练尤其是层数深的时候几乎不需要warmup也能收敛class SublayerConnection(nn.Module): def __init__(self, d_model, dropout): super().__init__() self.norm nn.LayerNorm(d_model) self.dropout nn.Dropout(dropout) def forward(self, x, sublayer): # Pre-LN: x Dropout(sublayer(LN(x))) return x self.dropout(sublayer(self.norm(x)))6.3 Encoder与Decoder的堆叠class EncoderLayer(nn.Module): def __init__(self, d_model, n_head, d_ff, dropout): super().__init__() self.self_attn MultiHeadAttention(d_model, n_head, dropout) self.ffn nn.Sequential( nn.Linear(d_model, d_ff), nn.ReLU(), nn.Linear(d_ff, d_model) ) self.sublayers nn.ModuleList( [SublayerConnection(d_model, dropout) for _ in range(2)] ) def forward(self, x, src_mask): x self.sublayers[0](x, lambda t: self.self_attn(t, t, t, src_mask)[0]) x self.sublayers[1](x, self.ffn) return x class DecoderLayer(nn.Module): def __init__(self, d_model, n_head, d_ff, dropout): super().__init__() self.self_attn MultiHeadAttention(d_model, n_head, dropout) self.cross_attn MultiHeadAttention(d_model, n_head, dropout) self.ffn nn.Sequential( nn.Linear(d_model, d_ff), nn.ReLU(), nn.Linear(d_ff, d_model) ) self.sublayers nn.ModuleList( [SublayerConnection(d_model, dropout) for _ in range(3)] ) def forward(self, x, memory, src_mask, tgt_mask): x self.sublayers[0](x, lambda t: self.self_attn(t, t, t, tgt_mask)[0]) # 交叉注意力Q 来自目标端K/V 来自 encoder 的 memory x self.sublayers[1]( x, lambda t: self.cross_attn(t, memory, memory, src_mask)[0] ) x self.sublayers[2](x, self.ffn) return x交叉注意力那行是整段代码里最值得盯的一行t是querymemory同时当key和valuemask用src_mask。三个参数、一个mask都跟第3.1节的表格一一对应。顶层模型class Transformer(nn.Module): def __init__(self, src_vocab, tgt_vocab, d_model512, n_head8, n_layer6, d_ff2048, dropout0.1, max_len5000, pad_idx0): super().__init__() self.pad_idx pad_idx self.src_embed TokenEmbedding(src_vocab, d_model, pad_idx) self.tgt_embed TokenEmbedding(tgt_vocab, d_model, pad_idx) self.pos_enc PositionalEncoding(d_model, max_len, dropout) self.encoder nn.ModuleList( [EncoderLayer(d_model, n_head, d_ff, dropout) for _ in range(n_layer)] ) self.decoder nn.ModuleList( [DecoderLayer(d_model, n_head, d_ff, dropout) for _ in range(n_layer)] ) self.generator nn.Linear(d_model, tgt_vocab) # 输出投影与目标端embedding共享权重 self.generator.weight self.tgt_embed.embed.weight self._init_weights() def _init_weights(self): for p in self.parameters(): if p.dim() 1: nn.init.xavier_uniform_(p) def encode(self, src, src_mask): x self.pos_enc(self.src_embed(src)) for layer in self.encoder: x layer(x, src_mask) return x def decode(self, memory, src_mask, tgt, tgt_mask): x self.pos_enc(self.tgt_embed(tgt)) for layer in self.decoder: x layer(x, memory, src_mask, tgt_mask) return x def forward(self, src, tgt, src_mask, tgt_mask): memory self.encode(src, src_mask) out self.decode(memory, src_mask, tgt, tgt_mask) return self.generator(out)权重共享那行self.generator.weight self.tgt_embed.embed.weight要注意顺序必须先创建embedding和generator再赋值绑定。绑定之后它们指向同一块内存优化器更新一次就是两次被更新——这是正确的行为不是bug。但如果d_model和tgt_vocab不匹配PyTorch会直接报shape错误这也算是个免费的检查。6.4 训练脚本与Noam学习率class NoamOpt: def __init__(self, d_model, warmup_steps, optimizer, factor1.0): self.optimizer optimizer self.d_model d_model self.warmup 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 ** -1.5 ) for group in self.optimizer.param_groups: group[lr] lr self.optimizer.step()这个学习率调度是原论文的经典设计前期线性预热后期按$step^{-0.5}$衰减。曲线形状像一个先升后降的小山包最高点出现在warmup_steps附近。为什么需要预热因为训练初期参数是随机初始化的梯度方向不稳定大学习率容易把模型带进一个坏区域小步慢走后再逐步放大。一步训练长这样model.train() tgt_input tgt[:, :-1] tgt_label tgt[:, 1:] src_mask make_pad_mask(src, pad_idx) tgt_mask make_tgt_mask(tgt_input, pad_idx) logits model(src, tgt_input, src_mask, tgt_mask) # (N, T-1, V) loss criterion( logits.reshape(-1, logits.size(-1)), tgt_label.reshape(-1) ) loss.backward() opt.step() opt.zero_grad(set_to_noneTrue)注意tgt_mask是用tgt_input构造的而不是tgt因为mask的序列长度必须和Decoder输入一致。长度差一位这种错误如果S和T恰好差一不一定报错但mask会整体错位。7. 跑通之后的验证形状、梯度、过拟合三步自查代码能跑不等于写对了。我自己的习惯是做三步自查按成本从低到高排列能在早期发现绝大多数实现错误。7.1 形状自检在forward里打形状日志最简单也最有效的办法在MultiHeadAttention.forward里加一段断言把关键张量的形状打出来。assert q.shape (N, self.n_head, L_q, self.d_k) assert k.shape (N, self.n_head, L_k, self.d_k) assert scores.shape (N, self.n_head, L_q, L_k) if mask is not None: # 检查 mask 能广播到 scores assert mask.shape[-2:] (L_q, L_k) or mask.dim() 4跑一个极小的dummy batch比如N2, S5, T7, d_model16, n_head4尺寸互相不相等能活捉所有恰好相等的隐藏错误。我特别推荐让S和T不相等、让N大于1因为很多bug在N1, ST的对称设置下完全看不见。7.2 梯度与权重初始化检查写完之后跑一次反向传播检查每个参数的梯度是不是None、是不是全0、是不是NaN。for name, p in model.named_parameters(): if p.grad is None: print(f[无梯度] {name}) # 绑定的权重会出现两次注意去重 elif torch.isnan(p.grad).any(): print(f[NaN梯度] {name}) elif p.grad.abs().sum() 0: print(f[零梯度] {name})几个常见现象generator.weight在权重共享后只有一份不会重复报pos_enc.pe不在参数列表里因为它是buffer不会出现在检查里这是对的如果某个LayerNorm的梯度全0可能是那一层的输入全是常量回去看数据预处理。权重的初始化也值得单独看Xavier uniform适合embedding之后的线性层nn.LayerNorm和nn.Embedding有自己的初始化方式不要粗暴覆盖padding_idx对应的embedding行通常在初始化后是零向量如果发现它非零说明你的初始化顺序把padding_idx的置零覆盖掉了。7.3 用极小语料做过拟合测试这是我最信任的一个测试构造10条左右的训练样本让模型过拟合。掉到接近0的loss就说明整条链路形状、mask、损失、优化器都对如果loss在一个高位平台震荡不下去那一定有实现错误。src torch.tensor([[1, 2, 3, 2, 0], [4, 5, 0, 0, 0]]) # 源端 tgt torch.tensor([[1, 6, 7, 2, 0], [1, 8, 2, 0, 0]]) # 目标端1BOS2EOS model Transformer(src_vocab16, tgt_vocab16, d_model32, n_head4, n_layer2, d_ff64, dropout0.0) criterion nn.CrossEntropyLoss(ignore_index0) opt torch.optim.Adam(model.parameters(), lr1e-3, betas(0.9, 0.98), eps1e-9) for step in range(500): tgt_in tgt[:, :-1] tgt_out tgt[:, 1:] logits model(src, tgt_in, make_pad_mask(src), make_tgt_mask(tgt_in)) loss criterion(logits.reshape(-1, logits.size(-1)), tgt_out.reshape(-1)) loss.backward() opt.step() opt.zero_grad(set_to_noneTrue) if step % 50 0: print(fstep {step:4d} loss {loss.item():.4f})我把dropout0因为做过拟合测试时要关掉一切随机性否则loss的噪声会掩盖真实趋势。跑出来应该是loss从三四左右一路掉到0.01以内。如果掉不下去检查清单是tgt有没有右移、mask的True/False语义有没有反、softmax的dim对不对、权重共享有没有做错。一个小细节Adam的betas(0.9, 0.98), eps1e-9是Transformer论文里的配置和PyTorch默认的(0.9, 0.999), 1e-8不同。区别不算大但既然复现就用论文的值能少一个变量。8. 我踩过的坑从维度对不上到loss不下降把整个实现流程走几遍之后我发现大部分坑其实可以归成四类。列出来供对照。第一类是形状错误会立刻报错。比如view的时候n_head和d_k顺序写反、contiguous忘了加、mask少一维广播不上。这类好修因为报错信息会告诉你哪一行、哪个形状对不上。第二类是静默的语义错误不报错但结果不对。包括softmax(dim1)写成了在head维度做、mask的True/False语义反了、交叉注意力传了错误的mask、多头切分时分组乱了。这类只能靠过拟合测试和形状断言来抓。第三类是输入构造错误。最典型的是tgt忘了右移、BOS/EOS的位置放错、padding的位置放在了序列中间。这类错误的表现是训练loss很低但生成质量差或者生成时输出一堆重复的EOS。第四类是数值与配置问题。包括用float(-inf)导致fp16下的NaN、warmup步数设得太小、学习率太高导致发散、ignore_index没设导致padding参与loss。我自己印象最深的一次是第二类和第三类的叠加模型训练loss降到了0.3看起来不错但生成结果全是重复的常见词。排查了半天先发现交叉注意力的mask用错了传了tgt的mask给src但那个位置长度恰好相等所以没报错修完之后效果只是略好再往上追发现是数据预处理时把源端和目标端搞混了源端填的是译文、目标端填的是原文。这种代码没错、数据错了的复合问题除非有端到端的过拟合测试否则很难定位。一个很实用的自检习惯是训练前先拿一批真实数据做前向把attention权重可视化成一个矩阵看形状和数值。如果某个位置的注意力分布完全均匀每列都是1/L说明mask把它屏蔽光了或者位置编码没加上如果分布极度尖锐且集中在第一个位置可能是padding没屏蔽干净模型学会了抄第一个词。attention权重是免费的可解释性信号别浪费它。最后分享一个我一直在用的小技巧把所有mask的构造都收拢到一个模块里然后在每个attention前用一个断言检查mask.shape能否广播到scores.shape。这个断言平时零成本只在维度出问题时触发能省下大量到底是哪一层的mask错了的排查时间。张量形状这东西早一次报错胜过事后三小时的猜测。