news 2026/9/30 3:54:44

Transformer输入输出与PyTorch代码实现:张量形状与Mask避坑

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
Transformer输入输出与PyTorch代码实现:张量形状与Mask避坑

做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_first=True的约定下(我强烈建议整个项目统一用这个约定),形状是这样:

  • src:(N, S),N是batch size,S是源端序列长度,元素是词表索引的整数
  • tgt:(N, T),T是目标端序列长度
  • src_key_padding_mask:(N, S),布尔或浮点,True表示该位置是padding
  • tgt_key_padding_mask:(N, T)
  • tgt_mask或叫attn_mask:(T, T),因果上三角

注意这里有个别扭的地方:padding mask在入口是二维的,但它最终要作用在四维的attention分数上。中间那次变形就是新手第一个大坑,后面第4节专门讲。

提示:如果你用的是nn.Transformer原生模块,默认batch_first=False,也就是(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)长度跟S
Decoder输出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_idx=0): super().__init__() self.embed = nn.Embedding(vocab_size, d_model, padding_idx=pad_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, 2i+1)} = \cos(pos / 10000^{2i/d_{model}})$$

偶数维度用sin,奇数维度用cos。为什么这么设计?直观理解是:不同维度对应不同频率的波,低维是高频(波长短,能区分相邻位置),高维是低频(波长长,能区分远距离位置)。这有点像用二进制编码表示数字,低位变化快、高位变化慢,组合起来就能唯一表示每个位置。

另外还有个数学性质:PE(pos+k)可以表示成PE(pos)的线性变换,这让模型有可能学会"相对位置"的概念——尽管位置编码本身是绝对的。

实现上有个细节值得说:位置编码矩阵应该用register_buffer注册,而不是普通属性。因为它是固定不变的常数,不该被优化器更新,也不该出现在state_dict里作为可训练参数。用register_buffer还自带一个好处:调用.cuda()或.to(device)时它会跟着一起搬,不用手动处理。

class PositionalEncoding(nn.Module): def __init__(self, d_model, max_len=5000, 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) # (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 来源mask
Encoder自注意力srcsrcsrcsrc padding mask
Decoder自注意力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}$把方差拉回1,softmax处于一个健康的温度区间。

实现上,缩放可以在两个位置做:一是在算完scores之后除,二是预先对Q做缩放。两种等价,我习惯在scores之后除,可读性更好:

def scaled_dot_product_attention(q, k, v, mask=None, dropout=None): 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了,整行都是-inf,softmax会产生NaN。这个问题在padding mask和因果mask叠加时可能触发,第4节细说。

另外注意softmax(dim=-1),最后一维是Key的序列长度。如果写成dim=1就会在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_idx=0): # 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, device=device), diagonal=1).bool().unsqueeze(0).unsqueeze(0)

torch.triu(..., diagonal=1)取的是严格上三角(不含对角线),也就是保留对角线以下为False、以上为True。diagonal=0会把对角线也屏蔽掉,那样每个位置连自己都看不到,违反直觉。

这里有个需要记住的约定差异:在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个词为止的历史",输出要预测第i+1个词。这就是所谓的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_index=pad_idx,让padding位置的损失被忽略。label_smoothing=0.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, dtype=torch.long, device=device) 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, keepdim=True) # (N, 1) ys = torch.cat([ys, next_word], dim=1) 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, mask=None, cache=None): # 训练时 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], dim=2) # 在序列维拼接 v = torch.cat([cache[1], v_new], dim=2) 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 PADT+2+padding
tgt_inputBOS w1 w2 EOS PADT+2
tgt_labelw1 w2 EOS PAD PADT+2

可以看到tgt_label里EOS后面的PAD被ignore_index忽略,只有EOS本身在参与计算,这正好对应"模型要学会在合适的时候输出EOS"。

6. 完整代码:一个能直接跑通的最小Transformer

前面拆了这么多细节,现在拼成一个完整可运行的实现。我刻意写得扁一些,不追求工程优雅,但求每一步的形状都清晰可查。

6.1 位置编码与掩码工具

这部分前面已经给过,这里补一个统一的掩码工具类,避免在多个地方重复构造:

def make_pad_mask(seq, pad_idx=0): # (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, device=device, dtype=torch.bool), diagonal=1 ).unsqueeze(0).unsqueeze(0) def make_tgt_mask(tgt, pad_idx=0): 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, dropout=0.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, mask=None): 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_attn

SublayerConnection用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是query,memory同时当key和value,mask用src_mask。三个参数、一个mask,都跟第3.1节的表格一一对应。

顶层模型:

class Transformer(nn.Module): def __init__(self, src_vocab, tgt_vocab, d_model=512, n_head=8, n_layer=6, d_ff=2048, dropout=0.1, max_len=5000, pad_idx=0): 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, factor=1.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_none=True)

注意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(比如N=2, S=5, T=7, d_model=16, n_head=4),尺寸互相不相等,能活捉所有"恰好相等"的隐藏错误。我特别推荐让S和T不相等、让N大于1,因为很多bug在N=1, S=T的对称设置下完全看不见。

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]]) # 目标端,1=BOS,2=EOS model = Transformer(src_vocab=16, tgt_vocab=16, d_model=32, n_head=4, n_layer=2, d_ff=64, dropout=0.0) criterion = nn.CrossEntropyLoss(ignore_index=0) opt = torch.optim.Adam(model.parameters(), lr=1e-3, betas=(0.9, 0.98), eps=1e-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_none=True) if step % 50 == 0: print(f"step {step:4d} loss {loss.item():.4f}")

我把dropout=0,因为做过拟合测试时要关掉一切随机性,否则loss的噪声会掩盖真实趋势。跑出来应该是loss从三四左右一路掉到0.01以内。如果掉不下去,检查清单是:tgt有没有右移、mask的True/False语义有没有反、softmax的dim对不对、权重共享有没有做错。

一个小细节:Adam的betas=(0.9, 0.98), eps=1e-9是Transformer论文里的配置,和PyTorch默认的(0.9, 0.999), 1e-8不同。区别不算大,但既然复现就用论文的值,能少一个变量。

8. 我踩过的坑:从维度对不上到loss不下降

把整个实现流程走几遍之后,我发现大部分坑其实可以归成四类。列出来供对照。

第一类是形状错误,会立刻报错。比如view的时候n_head和d_k顺序写反、contiguous忘了加、mask少一维广播不上。这类好修,因为报错信息会告诉你哪一行、哪个形状对不上。

第二类是静默的语义错误,不报错但结果不对。包括:softmax(dim=1)写成了在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错了"的排查时间。张量形状这东西,早一次报错,胜过事后三小时的猜测。

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/30 3:54:12

多Agent协作实战:Claude Code与Codex组队,Raven调度与Harness进化

1. 从单兵作战到团队协作:多Agent编排的必然趋势1.1 为什么单个AI助手已经不够用了过去大半年,我几乎把市面上主流的AI编程助手都深度用了一遍。最开始是Claude Code,后来是Codex,再后来各种Agent框架层出不穷。用久了会发现一个很…

作者头像 李华
网站建设 2026/9/30 3:53:20

亚马逊爆单增长闭环:数据化选品到复购的系统实操

行内做亚马逊的都知道,这两年“爆单”这个词已经从惊喜变成了焦虑的代名词。很多卖家还在靠老一套:看哪个类目火就冲进去,Listing抄优秀同行,广告预算拍脑袋定,结果要么是ACOS高得离谱,要么是单量起来之后被…

作者头像 李华
网站建设 2026/9/30 3:53:20

2分钟极速接入Claude Opus 5.5:Claude Code与AI Gateway配置实战

1. 为什么“2分钟接入”这件事值得认真拆解“2分钟上手,如何极速接入 Claude Opus 5.5”这个标题,乍一看像是一篇快餐式教程,但真正动手做过模型接入的人都知道,“2分钟”不是营销话术,而是一套被反复打磨过的路径设计…

作者头像 李华
网站建设 2026/9/30 3:53:01

RH134备考全解析:从LVM到SELinux,打造可交付的Linux系统

1. RH134核心知识地图:这次考试到底在考什么很多人看到RH134的第一反应是“RH124的进阶版”,这么说没错,但远远不够。RH124教的是“怎么用一台Linux服务器”,RH134教的是“怎么让一台Linux服务器稳定、安全、自动地跑起来”。这个…

作者头像 李华
网站建设 2026/9/30 3:52:35

电商平台分布式架构设计文档:从决策记录到容量测算落地

简介:方案文档围绕电商平台分布式架构设计,全面梳理从需求分析到技术落地的完整链路,面向系统架构师、后端开发及技术负责人等需要处理高并发、海量数据的从业者。文档先说明架构设计的必要条件和优势,再梳理购物、支付、物流、客…

作者头像 李华
网站建设 2026/9/30 3:51:57

小程序商城的首单,三个把客户劝退的细节

小程序商城的首单,三个把客户劝退的细节小程序商城上线后,最难的不是后面,而是第一单。老客户已经习惯在微信里报货,让他改变习惯的窗口只有一次。第一单不顺,后面就很难再推。看下来挫败首单的通常是三个很小的细节。…

作者头像 李华