在动手写Transformer的代码之前,我踩过最大的一个坑是:模型结构背得滚瓜烂熟,但真正把数据喂进去的那一刻,却被张量维度、位置编码的叠加方式、mask的形状这些"边角料"卡了整整两天。后来我才想明白,Transformer难的不是自注意力的公式本身,而是输入怎么变成模型能吃的张量、输出又怎么变回人能看懂的token这一整条链路。这篇就来把Transformer输入输出这条主干道彻底拆开,用PyTorch一行行实现给你看。不管你是刚学完注意力机制想动手跑通第一版模型的新手,还是已经会调库但没真正手写过底层的新同学,这篇都能让你对"数据进、结果出"这件事有清晰的掌控感,Transformer、PyTorch、代码实现三件事一次讲透。
1. 为什么值得从输入输出切入手撕Transformer
很多人学Transformer的路径是先啃论文里的缩放点积注意力公式,推导完softmax和除根号dk,然后一头扎进多头注意力的矩阵拆分里,最后被残差、层归一化的顺序搞晕。我不是说这条路径不对,而是它的门槛设得太高——你还没搞清楚模型"吃什么",就先研究它"怎么嚼",挫败感会很强。
1.1 输入输出是整条链路的骨架
Transformer本质是一个序列到序列的映射函数。输入端要做的事情只有三件:把离散的token变成连续向量、给向量注入位置信息、把不同长度的序列补齐对齐以便批量计算。输出端要做的事情也只有两件:把每个位置的隐状态投影回词表空间、用交叉熵把预测分布和真实标签对齐。你把这几件事拆清楚,中间那堆注意力层其实就是个黑盒,哪怕你不完全理解每个矩阵乘法,也能先把模型跑起来、把loss降下去。
我个人的经验是,先跑通输入输出、看到模型真的在学,再去深挖注意力的细节,学习动力会强很多。因为你能实时看到loss下降、看到生成的文本从乱码逐渐变成通顺的句子,这种正向反馈比单纯看公式有用得多。
1.2 输入输出藏着最多"看起来对其实错"的坑
真正让我栽跟头的,全是输入输出层面的细节。比如位置编码到底是在embedding之后加还是之前加、padding的mask方向和注意力mask的方向是否搞反、训练时输出要不要shift一位做teacher forcing、推理时的自回归循环怎么保持维度一致。这些坑有个共同特点:代码能跑通、不报错、loss也在动,但结果就是不对。你如果没系统拆过输入输出,debug时根本无从下手。
所以我建议的路线是:把输入输出的每一行都写清楚,每个张量的形状都标注出来,每步操作的意图都记下来。这样当模型表现异常时,你可以按"数据进→编码→注意力→输出→损失"这条线逐个环节排查,而不是对着报错发呆。这份掌控感才是手撕的价值所在。
1.3 一套代码覆盖三类典型任务
输入输出这条链路设计得好,是可以复用的。文本分类任务里,你取输出序列的特殊位置向量接个分类头;机器翻译任务里,你靠输出序列做自回归解码;语言模型任务里,你把输出和输入错开一位做下一词预测。三者共用同一套embedding、位置编码、mask构建逻辑,区别只在最后的输出头和损失函数。理解了这一点,你写一次底层,就能套用到大半的NLP任务上,性价比极高。
2. 输入端三大组件逐层拆解
输入这条链路看着简单,其实每一环都有讲究。我把它拆成embedding、位置编码、mask三块,逐块讲清楚为什么这么设计、参数怎么选、代码怎么写。
2.1 Token Embedding:词表到向量的第一跳
embedding层干的事情本质是一张查找表,把每个token的id映射成一个d_model维的稠密向量。词表大小vocab_size和模型维度d_model是两个必须提前定下来的超参数。我见过不少新手把vocab_size设得过小,导致大量token被迫映射到同一个unknown上,信息从一开始就丢了;也有人把d_model设成19、37这种奇数,后面拆多头的时候直接除不尽。
常见的配置是d_model取64、128、256、512、768这些2的幂或者4的倍数,这样除以注意力头数nhead能得到整数。比如d_model=512、nhead=8,每个头的维度就是64,干净利落。vocab_size一般根据你的分词器实际词表来定,中英文混合场景下几万到十几万都正常。我的建议是先用小词表小维度快速验证流程通不通,确认没问题再放大,否则一上来就上大模型,一个epoch跑半小时,调试效率极低。
还有一个容易被忽视的细节:embedding层的初始化。PyTorch默认用正态分布初始化,均值0标准差1。但Transformer论文里提到embedding要乘上一个缩放因子,也就是根号下d_model。这一步很多教程都一笔带过,其实很关键。因为位置编码的数值范围大约在[-1, 1]之间,而未经缩放的embedding数值方差接近1,两者量级差距大,直接相加会让位置信息被淹没。乘以根号d_model后,embedding的量级和位置编码匹配,两者才能有效融合。
import torch import torch.nn as nn import math class TokenEmbedding(nn.Module): def __init__(self, vocab_size, d_model, padding_idx=0): super().__init__() self.d_model = d_model self.embed = nn.Embedding(vocab_size, d_model, padding_idx=padding_idx) def forward(self, x): # x: [batch, seq_len] # 乘以 sqrt(d_model) 缩放,让量级和位置编码对齐 return self.embed(x) * math.sqrt(self.d_model)注意:padding_idx这一项一定要设,它会让padding对应的embedding向量在训练中保持为0,不参与梯度更新。如果你忘了设,padding位置会被当成正常token学习,序列越长,padding占比越大,对模型的污染越严重。
2.2 位置编码:让模型知道谁先谁后
Transformer的自注意力机制本身是位置无关的。换句话说,你把输入序列打乱顺序喂进去,注意力算出来的结果除了位置对调之外完全一样。但语言是有顺序的,"猫追老鼠"和"老鼠追猫"意思完全不同。所以必须显式地给每个位置注入顺序信息,这就是位置编码的作用。
原论文用的是正弦余弦位置编码,公式是固定的、不需要学习:
- 偶数维度:PE(pos, 2i) = sin(pos / 10000^(2i/d_model))
- 奇数维度:PE(pos, 2i+1) = cos(pos / 10000^(2i/d_model))
我第一次看到这个公式时也懵,为什么用正弦余弦、为什么底数是10000。用生活类比解释一下:这个设计让每个位置在每个维度上都有一个唯一的"指纹",而且不同位置之间的相对关系可以通过三角恒等式线性表达出来,模型只需要学一个线性变换,就能从任意位置的编码推算出相对距离。这就是它能泛化到训练时没见过的更长序列的原因。
具体实现时有个坑:用exp和log的组合来计算分母,比直接连乘幂次数值更稳定。下面这份实现我用了很多次,实测很稳。
class PositionalEncoding(nn.Module): def __init__(self, d_model, max_len=5000, dropout=0.1): super().__init__() self.dropout = nn.Dropout(p=dropout) pe = torch.zeros(max_len, d_model) position = torch.arange(0, max_len, dtype=torch.float).unsqueeze(1) # [max_len, 1] # 用 exp(log) 组合计算,避免大指数溢出 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] # register_buffer 让 pe 随模型保存但不参与梯度更新 self.register_buffer('pe', pe) def forward(self, x): # x: [batch, seq_len, d_model] x = x + self.pe[:, :x.size(1), :] return self.dropout(x)提示:pe用register_buffer注册而不是普通属性,这样模型保存时会带上它,加载时自动恢复,同时它不会被优化器当成参数更新。这个细节我一开始没注意,后来发现模型换设备时pe没跟着走,直接报device不匹配的错。
positional encoding还有一个变体是可学习的位置嵌入,直接建一个[max_len, d_model]的embedding表让模型自己学。两种方案各有取舍:固定编码泛化性好、参数少,适合序列长度不固定的场景;可学习编码更灵活,但需要见过足够长的序列才能学好对应的位置。我个人在序列长度比较固定的任务上更倾向可学习的方式,泛化要求高的场景用正弦编码。
2.3 归一化与Dropout的接入位置
输入端还有两个容易被忽略但影响很大的组件:层归一化和Dropout。原始论文里层归一化放在每个子层之后,也就是Post-LN结构,残差相加后再归一化。但后来大量实践发现Pre-LN更稳定,也就是先归一化再进子层。
Post-LN的问题在于,残差相加后的值会随着层数增加不断累积,导致深层梯度不稳定,训练时需要精心的warmup策略。Pre-LN则把归一化提前,主干的数值范围更可控,训练收敛快很多,对学习率也没那么敏感。我现在的习惯是一律用Pre-LN,PyTorch从1.12开始TransformerEncoderLayer就支持norm_first参数,直接设成True即可。
Dropout的放置也有讲究。输入端embedding和位置编码相加后会加一次dropout,每个注意力子层的输出后加一次,前馈网络的中间层也加一次。dropout率一般取0.1,小数据量时可以提到0.2到0.3防过拟合。手写的时候记得所有dropout层的训练和推理状态切换要一致,用model.train()和model.eval()统一控制,别手动去设p值。
3. 屏蔽掩码的构建细节
mask是Transformer里最容易被写错、又最难察觉的地方。它不控制计算,只控制注意力看哪里,所以错了也不报错,只是结果悄悄变差。
3.1 Padding Mask与Causal Mask的区别
Transformer里常见的mask有两类,用途完全不同,千万别混。
padding mask解决的是批量内序列长度不一的问题。短句补了padding,这些padding位置没有语义,其他位置不应该关注它们,所以要把注意力分数置成负无穷,softmax后自然变成0。它的形状通常是[batch, seq_len],标记每个位置是不是padding。
causal mask解决的是自回归生成时的"未来信息泄露"问题。预测第t个词时只能看到前t个词,不能偷看后面的答案。它是个上三角矩阵,右上部分置负无穷,把未来位置挡住。形状是[seq_len, seq_len],每条下三角为True。
我踩过的坑就是这两个mask的形状和方向搞混。padding mask在送进attn_mask或key_padding_mask时要转成合适的形状,PyTorch的TransformerEncoderLayer用的是src_key_padding_mask,形状是[batch, seq_len],而TransformerDecoderLayer用的是tgt_mask,形状是[seq_len, seq_len]。两个接口的命名风格不统一,很容易传错。
def create_padding_mask(seq, pad_id=0): # seq: [batch, seq_len] # 返回 [batch, seq_len],True 表示该位置是 padding,需要屏蔽 return (seq == pad_id) def create_causal_mask(seq_len, device): # 返回 [seq_len, seq_len],True 表示需要屏蔽(上三角) return torch.triu( torch.ones(seq_len, seq_len, device=device, dtype=torch.bool), diagonal=1 )注意:PyTorch的mask约定是"True的位置被屏蔽"。也就是说,你标记为True的地方,注意力会把它忽略掉。这和我一开始的直觉相反,我总想着True是保留,结果写反了,模型训练很久loss都不降。
3.2 mask在注意力中的生效路径
理解了mask的语义,还要理解它在哪里生效。注意力计算是Q乘K转置得到分数矩阵,再除以根号dk,然后加mask,最后softmax,再乘V。mask是在softmax之前加的,把需要屏蔽的位置加上一个极大的负数,softmax出来后这些位置的权重就接近0。
这里有个数值细节:加的是负无穷还是负的极大值。理论上是负无穷,但代码里直接用float('-inf')可能在某些情况下导致NaN,尤其是整行都被mask掉的时候。稳妥的做法是用一个足够大的负数,比如负1e9,或者用torch.finfo(dtype).min。此外还要注意,如果某一行全是mask,softmax后会出现0/0的NaN,一般要给padding位置或者加个安全值处理,但在正常的自回归和padding场景下不会出现整行全屏蔽,所以问题不大。
3.3 维度对齐的检查清单
我整理了一份mask维度速查,每次写新模型都对着核一遍,能省下大量debug时间。
| mask类型 | 典型形状 | 使用接口 | 作用 |
|---|---|---|---|
| padding mask | [batch, seq_len] | src_key_padding_mask | 屏蔽padding位置 |
| causal mask | [seq_len, seq_len] | tgt_mask / attn_mask | 屏蔽未来位置 |
| 合并mask | [batch, seq_len, seq_len] | attn_mask | 同时屏蔽两类 |
实际写码时我会在forward里加几行assert,把关键张量的shape打出来,跑第一个batch时确认一遍,之后就可以删掉。这个习惯帮我发现了无数次形状对不上的问题,强烈建议你也养成。
4. 输出端到损失函数的关键步骤
输入端搞定了,注意力层跑完,接下来是输出端。这里有两件事:把隐状态投影回词表,以及用损失函数把预测和标签对齐。
4.1 线性投影与权重共享
每个位置的隐状态是个d_model维向量,要变成词表上的概率分布,需要过一个线性层把维度从d_model映射到vocab_size,然后接softmax得到每个词的概率。这一步本身很简单,但有个优化技巧值得说:权重共享。
所谓权重共享,就是让输出线性层的权重直接复用embedding层的权重矩阵。理由是embedding做的是词表到向量的映射,输出层做的是向量到词表的映射,两者互为逆操作,共享权重能减少参数量,还能让两边的语义空间保持一致,实践中往往能带来性能提升。实现上只要一行代码把两个权重绑定即可。
self.fc_out = nn.Linear(d_model, vocab_size) self.fc_out.weight = self.embedding.embed.weight # 权重绑定提示:PyTorch的nn.Linear权重形状是[out_features, in_features],也就是[vocab_size, d_model],而nn.Embedding的权重形状是[num_embeddings, embedding_dim],也就是[vocab_size, d_model],两者一致,可以直接绑定。绑定后记得embedding层不要设padding_idx的向量为0后又被反传破坏,实践上没问题。
4.2 teacher forcing与标签移位
训练自回归模型时,输入序列和标签序列要错开一位。假设句子是[A, B, C, D],输入是[A, B, C],标签是[B, C, D]。模型看到A预测B,看到A B预测C,以此类推。这就是teacher forcing,用真实的上一个词而不是模型自己预测的词作为下一步输入,训练更快更稳。
代码实现时,如果你用的是nn.Transformer,输入和输出的embedding是分开处理的。输入过encoder得到记忆,输出右移一位后过decoder,再投影到词表,和标签算交叉熵。这里移位容易出错,我用的是切片方式:tgt_input = tgt[:, :-1],tgt_label = tgt[:, 1:]。注意还要给tgt_input前面加个起始符,或者直接用切片让decoder从第一个词开始。
# 假设 tgt: [batch, seq_len],包含 <bos> ... <eos> tgt_input = tgt[:, :-1] # 去掉最后一个,作为decoder输入 tgt_label = tgt[:, 1:] # 去掉第一个,作为预测目标交叉熵损失要注意ignore_index参数,把padding位置的损失忽略掉,否则模型会花大量精力去预测padding。loss算完后除以有效token数做平均,而不是除以总长度,这样不同批次的loss才可比。
criterion = nn.CrossEntropyLoss(ignore_index=0) # 0是padding id # logits: [batch, seq_len, vocab_size] loss = criterion(logits.reshape(-1, vocab_size), tgt_label.reshape(-1))4.3 完整的前向流程串一遍
把上面这些拼起来,一个最小的Transformer语言模型前向就是这样的:
class MiniTransformer(nn.Module): def __init__(self, vocab_size, d_model=256, nhead=8, num_layers=4, dim_ff=1024, dropout=0.1, max_len=512): super().__init__() self.d_model = d_model self.embedding = TokenEmbedding(vocab_size, d_model, padding_idx=0) self.pos_enc = PositionalEncoding(d_model, max_len, dropout) encoder_layer = nn.TransformerEncoderLayer( d_model=d_model, nhead=nhead, dim_feedforward=dim_ff, dropout=dropout, batch_first=True, norm_first=True ) self.encoder = nn.TransformerEncoder(encoder_layer, num_layers) self.fc_out = nn.Linear(d_model, vocab_size) self.fc_out.weight = self.embedding.embed.weight def forward(self, src, src_key_padding_mask=None, attn_mask=None): # src: [batch, seq_len] x = self.embedding(src) # [B, L, D] x = self.pos_enc(x) # [B, L, D] x = self.encoder(x, mask=attn_mask, src_key_padding_mask=src_key_padding_mask) logits = self.fc_out(x) # [B, L, vocab] return logits这份代码可以直接跑通一个字符级语言模型。我建议你拿一段中文文本、按字符建词表,跑上几十个epoch,看着它从生成乱码慢慢变成能拼出词,这个过程对理解整条链路帮助极大。
5. 训练与推理阶段的输入输出差异
训练和推理看起来只是调用方式不同,实际上输入输出的组织方式差别很大。搞不清这一点,很容易写出训练正常但推理拉胯的模型。
5.1 训练阶段的并行与推理阶段的串行
训练时因为有teacher forcing,整个目标序列可以一次性喂进去并行计算,所以Transformer训练效率高。但推理时你没有真实标签,只能一个词一个词地生成,生成第t个词需要前面t-1个词的结果,天然是串行的。这就导致推理速度远慢于训练,序列越长越明显。
如果你的任务是纯编码类,比如文本分类、序列标注,那推理也是并行的,没这个问题。但生成类任务就要注意了。
5.2 自回归推理循环的实现要点
自回归推理的核心循环是:维护一个已生成序列,每次把整个序列喂进去前向一次,取最后一个位置的logits,选出下一个词,追加到序列末尾,直到生成结束符或达到最大长度。
@torch.no_grad() def generate(model, start_tokens, max_new_tokens=50, temperature=1.0, top_k=None): model.eval() tokens = start_tokens # [1, seq_len] for _ in range(max_new_tokens): logits = model(tokens) # [1, L, vocab] next_logits = logits[:, -1, :] / temperature # 取最后一个位置 if top_k is not None: v, _ = torch.topk(next_logits, top_k) next_logits[next_logits < v[:, [-1]]] = float('-inf') probs = torch.softmax(next_logits, dim=-1) next_token = torch.multinomial(probs, num_samples=1) # [1, 1] tokens = torch.cat([tokens, next_token], dim=1) if next_token.item() == eos_id: break return tokens注意:这个朴素实现每步都要把整个序列重新前向一遍,重复计算极其严重,序列长到几百时慢得离谱。优化思路是KV Cache——把已经算过的key和value缓存下来,每步只对新token计算,这是所有推理加速方案的基础。手写阶段可以先不优化,跑通逻辑要紧,但要知道这个瓶颈在哪。
温度参数控制生成的随机性,temperature小于1让分布更尖锐、生成更确定,大于1让分布更平坦、生成更多样。top_k则是只在概率最高的前k个词里采样,避免采到长尾怪词。这两个是生成任务最常用的调节旋钮,建议都留出接口。
5.3 一个完整的训练循环模板
把前面的东西拼成一个训练循环,结构大概是这样:
def train_one_epoch(model, loader, optimizer, criterion, device): model.train() total_loss = 0.0 for batch in loader: src = batch['input_ids'].to(device) # [B, L] tgt = batch['labels'].to(device) pad_mask = create_padding_mask(src, pad_id=0) logits = model(src, src_key_padding_mask=pad_mask) loss = criterion(logits.reshape(-1, logits.size(-1)), tgt.reshape(-1)) optimizer.zero_grad() loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0) optimizer.step() total_loss += loss.item() return total_loss / len(loader)梯度裁剪这一步我单独拎出来说:Transformer对梯度爆炸比较敏感,尤其是学习率偏大的时候。clip_grad_norm_把梯度的整体范数限制在max_norm以内,是防止loss突然飞掉的保险丝。我在没加这行之前,遇到过loss从1.2直接跳到nan的情况,加上之后就稳了。
6. 常见问题与排查技巧实录
下面这些是我和身边朋友手撕Transformer时反复遇到的真问题,整理成速查表,遇到卡壳时对着看。
| 现象 | 可能原因 | 排查方法 |
|---|---|---|
| loss几乎不下降 | mask方向写反,或学习率过大 | 检查mask的True/False语义,打印一个batch的注意力权重看屏蔽对不对 |
| loss降了但生成乱码 | 训练和推理的标签移位不一致 | 对比训练时的tgt_input和推理时的起始输入 |
| 换设备报device不匹配 | 位置编码没注册成buffer | 用register_buffer,确认pe和输入在同一device |
| 长序列结果差 | 训练时max_len不够,位置编码越界 | 打印最大序列长度,确认没超PositionalEncoding的max_len |
| 内存爆掉 | batch和序列长度太大 | 减小batch,或改用梯度累积模拟大批量 |
| 维度报错 | d_model除不尽nhead | 确保d_model % nhead == 0 |
我重点讲几个不是一眼能看出来的。
第一个是mask写反。PyTorch的约定是True表示屏蔽,但很多人习惯用True表示保留,赋值时反着来。这个错误的隐蔽性在于,代码不报错,甚至loss也能缓慢下降,只是永远到不了理想水平。我的排查手段是构造一个极端的短序列,手动把mask打印出来,用眼睛核对屏蔽区域是不是我想要的。构造小样本验证是个万能技巧,遇到任何可疑行为都先缩小到能一眼看穿的最小例子。
第二个是标签位移。训练时我习惯tgt_input = tgt[:, :-1]、tgt_label = tgt[:, 1:],但推理时起始输入只给了起始符,模型看到的信息比训练时少。如果训练时序列开头没放起始符,模型会不适应。解决办法是训练和推理都统一带上起始符,保证分布一致。这类"训练和推理不一致"的问题在工程里非常常见,本质是数据管道的两套代码走了不同逻辑。
第三个是位置编码越界。PositionalEncoding的pe表大小是max_len,如果你喂进去的序列比它长,切片时虽然不会报错,但超出的位置拿不到位置编码,等于没位置信息。我建议在forward里加个长度断言,超出就报警或动态扩展,别让它静默出错。
实操心得:每次写完一个新模型,先拿两三条假数据过一遍forward,把每个中间张量的形状打出来核对。这个动作花不了两分钟,但能挡掉八成以上的低级错误。我早期省了这一步,结果花了几个小时在错误信息里绕圈。
最后一个经验是学习率调度。Transformer原论文用的是warmup加逆平方根衰减,先把学习率线性升上去再缓慢降下来。这个策略对深层Transformer很关键,warmup让模型在初期参数还乱的时候小步走,避免一开始就把参数带偏。我一般用固定步数warmup,warmup步数占总步数的5%到10%,之后用余弦退火或线性衰减。PyTorch的LambdaLR可以很方便地实现这套调度,配合前面说的梯度裁剪,训练稳定性会好很多。
把输入输出这条链路彻底拆透之后,我再去看那些注意力公式、多头拆分、前馈网络,就都是可以被顺畅嵌入到这条主干上的零件了。真正决定模型能不能跑起来、跑得对不对的,往往不是那些精妙的矩阵变换,而是embedding有没有缩放、mask有没有写反、标签有没有移位这些不起眼的地方。我现在的习惯是,每换一个新任务,先把数据管道的输入输出打印出来逐字段核对,确认这条链路干净了,再往中间堆结构。这个顺序反过来,debug的代价会大得多。