1. 先别急着看公式:Transformer代码解读(PyTorch)的入场姿势
我第一次把《Attention Is All You Need》的公式和一份PyTorch实现摆在一起看的时候,卡住的不是注意力机制本身,而是几行看起来毫无道理的代码:x.view(B, L, self.n_head, self.d_k).transpose(1, 2)、scores.masked_fill(mask, float("-inf"))、nn.LayerNorm被放在残差相加的外面。公式里明明写的是Attention(Q,K,V) = softmax(QK^T/√d_k)V,代码里却多了好几个中间变量和转置操作。
这篇 Transformer代码解读(PyTorch)想解决的问题就是这个:把一份可以跑通的最小实现,按张量形状这条主线拆开讲清楚。每个模块为什么这么写、形状怎么变、哪些地方是论文没写但工程上必须补的细节。适合已经看过原理、准备动手复现,或者正在读别人的开源实现却看不懂中间那几行的人。全文不依赖任何大型框架封装,纯torch.nn手写,改起来方便。
我会按"先定骨架、再拆注意力、然后处理mask、接着是位置编码和层结构、最后跑训练和推理"的顺序推进,中间穿插我自己踩过的坑。你可以把代码分块贴进一个.py文件,边读边跑,看到每个 tensor 的 shape 打印出来,比对着公式看效率高得多。
1.1 一份能跑的最小骨架包含哪几个类
一份结构清晰的实现,通常只有六个类:MultiHeadAttention、PositionwiseFeedForward、EncoderLayer、DecoderLayer、PositionalEncoding、Transformer。前两个是零件,中间两个是把零件组装起来的一层,最后一个是整体模型。
分的意义在于:注意力模块要被复用三次——编码器自注意力、解码器自注意力、解码器交叉注意力。如果把它塞进EncoderLayer里写死,交叉注意力就得复制一份代码出来,改bug要改三处。第一次写的时候我图省事没拆,后来调mask逻辑,三个位置改得满头大汗,第二次重构才拆出来。
PositionwiseFeedForward单独拆出来也是同理,它接受(B, L, D)形状的输入,内部只做逐位置的线性变换,不碰序列维度。这意味着它可以和注意力模块用同样的接口串起来,残差相加时形状天然对齐。
1.2 四个贯穿全文的张量形状约定
代码里所有形状相关的推理,都建立在四个约定上。我把它们列出来,后面每一节都会回来对照:
| 符号 | 含义 | 典型取值 |
|---|---|---|
| B | batch size | 32、64 |
| L | 序列长度(源或目标) | 50、128 |
| D | d_model,模型隐藏维度 | 512 |
| H | 注意力头数n_head | 8 |
| d_k | 每个头的维度,等于D // H | 64 |
关键的约定是:所有模块的输入输出统一是(B, L, D),头拆分只发生在MultiHeadAttention内部,拆完立刻转回(B, L, D)交出去。这个约定让残差连接写起来极其干净——只要形状都是(B, L, D),x + sublayer(x)永远合法。
提示:
d_model % n_head != 0是最常见的初始化报错来源。512/8、768/12、1024/16 都没问题,但如果随手写d_model=100, n_head=8,会在view那一步得到一个形状不匹配的错误,而且报错信息指向的是 view 而不是初始化,第一次遇到会找很久。建议在__init__里加一句assert d_model % n_head == 0。
2. 多头注意力:Transformer里唯一真正复杂的那段代码
整个 Transformer 里,只有这一段值得逐行读。其余的层堆叠、残差、LayerNorm 都是标准套路。注意力的核心就三件事:把输入投影成Q、K、V;算相似度并归一化;用权重加权V。多头则是在通道维上切分,让不同的子空间学不同的关系。
我见过不少初学者在这一段反复卡壳,原因是论文的公式是二维的(矩阵乘法),而代码是四维的(多了batch和head两维)。理解的关键是:把(B, H, L, d_k)看成一堆互不干扰的小矩阵,公式原封不动地作用在最后两维上。
2.1 QKV投影为什么合并成一个Linear更好
先把不合并的版本写出来,对照着看:
import math import torch import torch.nn as nn import torch.nn.functional as F 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_model = d_model self.n_head = n_head self.d_k = d_model // 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 _split_heads(self, x): # (B, L, D) -> (B, H, L, d_k) B, L, _ = x.size() return x.view(B, L, self.n_head, self.d_k).transpose(1, 2) def forward(self, query, key, value, mask=None): B = query.size(0) q = self._split_heads(self.w_q(query)) k = self._split_heads(self.w_k(key)) v = self._split_heads(self.w_v(value)) scores = torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(self.d_k) if mask is not None: scores = scores.masked_fill(mask, float("-inf")) attn = F.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), attn分开写三个Linear语义最直接,调试时也能单独检查w_q的梯度。但它有个实际代价:编码器的自注意力中 Q、K、V 全都来自同一个输入x,三个Linear意味着三次矩阵乘法((B*L, D) @ (D, D))。合并成一个nn.Linear(d_model, 3 * d_model),一次算完再切分,GEMM 的调用次数从三次降到一次,在 D 比较大的时候(比如 1024)能省下可观的时间。
省的不只是时间。分开的三个Linear各自独立初始化,虽然统计上等价,但合并写法是一次kaiming/xavier初始化作用于一个(D, 3D)的大矩阵,等价于三块共享同样的初始化分布,实际训练里前期的数值稳定性会稍微好一点。如果你想做成"合并计算 + 逻辑分开"的形式,可以这样:
self.w_qkv = nn.Linear(d_model, 3 * d_model) # forward 中: # qkv = self.w_qkv(query) # (B, L, 3D) # q, k, v = qkv.chunk(3, dim=-1) # 各 (B, L, D)chunk(3, dim=-1)是按最后一个维度均分,语义清晰,比手写切片好读。我的建议是:学习阶段用分开的三个Linear,部署或做大模型时换成合并版本,两者对训练结果的影响在同一个量级,不用纠结太久。
2.2 缩放点积里的除以根号d_k到底防的是什么
/ math.sqrt(self.d_k)这一行几乎每份实现都有,但真正想明白它防什么的人不多。说人话:假设 q 和 k 的每个分量都是均值0、方差1的独立随机变量,那么点积q·k是d_k个乘积之和,它的方差是d_k。d_k = 64的时候,点积的标准差就是8,量级在 ±20 上下浮动很常见。
问题出在后面那一步 softmax。softmax 对输入的尺度非常敏感:输入越大,分布越尖锐,接近 one-hot;输入很小,分布就趋于均匀。如果点积的方差随d_k线性增长,那么d_k一大,softmax 的输出就变成近似 one-hot 的极端分布,梯度几乎全部集中在最大值那个位置,其余位置梯度趋近0——这就是梯度消失的经典来源。
除以√d_k之后,点积的方差被拉回到1附近,softmax 的输入尺度不随头维度变化,梯度分布也就稳定了。这个推导值得自己写一遍:如果每个分量方差是1,Var(Σ q_i k_i) = d_k,除以√d_k后方差归一。
注意:这里说的
d_k是每个头的维度,不是d_model。经常看到有人写成/ math.sqrt(self.d_model),在小模型上可能看不出差别(512 的平方根和 64 的平方根差 2.8 倍),但这是错的,训练后期会明显感觉到收敛变慢。
2.3 view、transpose、contiguous三兄弟的顺序陷阱
_split_heads和它对应的逆操作,是新手最容易写出隐性bug的地方:
x.view(B, L, self.n_head, self.d_k).transpose(1, 2)这里view把最后一维 D 拆成(H, d_k),得到(B, L, H, d_k);transpose(1, 2)把 L 和 H 换位,得到(B, H, L, d_k)。逻辑上完全正确。但反过来合并头的时候,如果直接写:
out.transpose(1, 2).view(B, -1, self.d_model) # 有可能报错transpose只改 stride 不改内存布局,得到的张量在内存里不连续。view要求内存连续,所以会抛出RuntimeError: view size is not compatible with input tensor's size and stride。正确写法必须先.contiguous(),或者干脆改用reshape(它内部会在需要时自动复制)。
out = out.transpose(1, 2).contiguous().view(B, L, self.d_model)这里有个容易被忽略的细节:contiguous()是一次真实的内存拷贝,有开销。如果这段代码在你的瓶颈里,可以考虑用reshape代替view + contiguous——效果一样,写法更短。但要注意reshape在连续张量上返回视图、在不连续张量上返回拷贝,行为随输入变化,调试时不如显式contiguous()直观。
另一个陷阱是transpose(-2, -1)和transpose(1, 2)混用。前者是在最后两维上换位(用于 K 的转置,得到(B, H, d_k, L)),后者是换 L 和 H。这两个操作在同一段代码里出现,写到后面很容易写混。我的做法是全程用负数索引表示"最后两维"的运算,用正数索引表示"批和头"的维度操作,读的时候一眼能区分这是在算相似度还是在重整布局。
3. mask的三种形态:Transformer代码里翻车率最高的地方
如果统计一下 Transformer 实现里出的bug,mask 相关的能占一半以上。它的问题在于:论文里只用一句话带过("we mask out subsequent positions"),但代码里要处理三种完全不同的场景——源序列的 padding、目标序列的 padding、目标序列的因果可见性,而且它们还要能叠加。更麻烦的是,出错时往往不报异常,只是结果悄悄变差或者变成 nan。
所以这一节我打算把 mask 单独拎出来,讲清楚每张 mask 管什么、长什么样、怎么组合。
3.1 padding mask与causal mask的生成方式
padding mask 管的是"哪些位置是填充的,不要看它们"。假设 pad token 的 id 是0,源序列(B, L)里为0的位置就是无效位置:
def make_pad_mask(seq, pad_id=0): # (B, L) -> (B, 1, 1, L),方便广播到 (B, H, L_q, L_k) return (seq == pad_id).unsqueeze(1).unsqueeze(2) def make_causal_mask(size, device): # 上三角(不含对角线)为 True,表示"未来位置" return torch.triu(torch.ones(size, size, dtype=torch.bool, device=device), diagonal=1)两个形状设计上的考虑值得说清楚。第一,pad mask 生成时插了两个长度为1的维度,是为了能和(B, H, L_q, L_k)的 scores 广播。如果不插,(B, L)和(B, H, L_q, L_k)广播会按右对齐规则匹配,很容易对错维度——而且 torch 的广播在某些情况下不会报错,直接算出错误结果,这是最阴险的一类bug。第二,causal mask 用torch.triu(..., diagonal=1)而不是diagonal=0,因为对角线上的 token 应该能看到自己。
组合时用逻辑或:
def combine_masks(*masks): out = masks[0] for m in masks[1:]: out = out | m return out|是对布尔张量做逐元素或,True表示屏蔽。如果你的两张 mask 形状是(B, 1, 1, L)和(1, 1, L, L),广播后得到(B, 1, L, L),正好能用在 scores 上。
3.2 加性mask和乘性mask混用的后果
mask 的施加方式有两种主流写法:
# 写法A:布尔 mask + masked_fill,屏蔽位填 -inf scores = scores.masked_fill(mask_bool, float("-inf")) # 写法B:浮点 mask + 加法,屏蔽位是一个大负数 scores = scores + (1.0 - mask_float) * (-1e9)两种写法本身都能用,但混用会出事。我见过一份代码,上游生成了True/False的布尔 mask,下游却拿它去做乘法scores * mask。结果False被当作0,True被当作1,语义正好反过来——被屏蔽的位置乘1保留,有效位置乘0被清零。这种错误不会报错,loss 也能下降,只是模型在学一个完全错误的东西,你可能跑了两天才发现验证集指标不对。
写法B里那个-1e9也不是随便取的。如果 scores 本身量级很大(比如没做缩放),-1e9加上去之后可能被浮点精度"吃掉"一次有效数字,softmax 的结果不是严格0而是 1e-7 这种小量。在float16下这个问题更明显,-1e9会溢出成-inf,再参与后续计算可能产生 nan。所以我的习惯是:训练用布尔 mask +masked_fill(True, -inf),同时保证不存在整行全屏蔽的情况。下一节会讲为什么。
3.3 交叉注意力该用哪张mask
交叉注意力的 Q 来自解码器、K 和 V 来自编码器输出,所以 mask 由 K 的来源决定:应该用源序列的 padding mask,形状是(B, 1, 1, L_src)。
常见的错误是把解码器的 causal mask 顺手传进去。这样做的后果是:解码器第 i 个位置在关注编码器时,只能看到源序列的前 i 个 token。源序列的长度和目标序列长度往往不一样,广播之后行为更加诡异。更隐蔽的是,如果源和目标长度恰好相等,代码不会报错,模型照样能训,只是翻译质量莫名其妙地差。
我做过的检查是:在DecoderLayer.forward里加一行断言,把cross_mask的最后一维打印出来,确认它等于源序列长度。跑通之后删掉。这类"打印一次就知道对不对"的检查,比事后调参省时间得多。
| mask类型 | 形状 | 作用位置 | 来源 |
|---|---|---|---|
| 源 padding mask | (B, 1, 1, L_src) | 编码器自注意力、解码器交叉注意力 | 源序列== pad_id |
| 目标 padding mask | (B, 1, 1, L_tgt) | 解码器自注意力 | 目标序列== pad_id |
| causal mask | (1, 1, L_tgt, L_tgt) | 解码器自注意力 | triu(diagonal=1) |
| 解码器自注意力合并 | (B, 1, L_tgt, L_tgt) | 解码器自注意力 | 上面两张按位或 |
4. 位置编码与词嵌入:两个看起来最简单却最容易埋雷的模块
注意力机制本身是置换等变的——把输入序列的顺序打乱,输出也只是被同样打乱,模型完全感知不到位置。所以位置信息必须外挂进去。这一块的代码通常只有十几行,但埋的雷一点不少。
4.1 正弦位置编码的实现与register_buffer
class PositionalEncoding(nn.Module): def __init__(self, d_model, max_len=5000, dropout=0.1): super().__init__() 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) self.register_buffer("pe", pe.unsqueeze(0)) # (1, max_len, D) self.dropout = nn.Dropout(dropout) def forward(self, x): x = x + self.pe[:, :x.size(1)] return self.dropout(x)几个细节。div_term的写法是exp(arange(0, d_model, 2) * (-log(10000) / d_model)),等价于1 / (10000 ** (2i / d_model)),但用exp和log的组合能避免幂运算的数值误差,也让d_model很大时不会溢出。
register_buffer这一步不能省。它把pe登记为模块的缓冲区,随模型一起.to(device),但不会被optimizer当作参数更新。如果直接写成self.pe = pe.unsqueeze(0),它就成了一个普通属性,模型搬到 GPU 上时它还在 CPU,x + self.pe会因为设备不一致报错——或者更糟,在某些版本下静默做跨设备拷贝,拖慢每一次前向。
还有一点:self.pe[:, :x.size(1)]是按当前序列长度动态切片。这意味着即使max_len设成5000,实际只用了前面 L 行,没有多余计算。
4.2 训练长度外的外推与可学习位置编码的取舍
正弦编码的一个好处是理论上有外推能力:因为它是连续函数在整数点上的采样,位置 1000 的编码和位置 10 的编码之间存在平滑关系。但这个"理论外推"在实际里并不好用。我试过在训练长度128的模型上直接推理长度256,前面几十步还行,越往后越乱,因为注意力分布是在训练长度范围内学出来的,位置编码的形式没有变,但模型没见过长距离的相对关系。
另一个选择是nn.Embedding(max_len, d_model),把位置当作可学习的参数。它在训练长度内通常比正弦编码效果略好(参数可以自由拟合),但完全无法外推——输入长度超过max_len直接索引越界报错,比"效果变差"更难处理。
所以我的选择标准是:固定长度任务(如固定窗口的时序预测、定长分类)用可学习位置编码;变长任务(翻译、摘要)用正弦编码,并把max_len设成训练集最大长度的1.5倍左右留余量。这个余量不是为了外推,而是为了防止某条特别长的样本在训练中直接崩掉。
| 方案 | 参数量 | 外推能力 | 适用场景 |
|---|---|---|---|
| 正弦编码 | 0 | 有,但实际有限 | 变长序列、翻译 |
| 可学习位置编码 | max_len * d_model | 无 | 定长任务、分类 |
| 相对位置编码 | 额外参数 | 较好 | 长文本、需要长距离建模 |
5. 残差、LayerNorm与堆叠顺序:Post-LN还是Pre-LN
这一节讲的是层内的组织方式。它不涉及新的数学,但直接决定模型能不能训起来。原论文用的是 Post-LN,很多现代实现改用 Pre-LN,差别只有一行代码,训练稳定性却差很多。
5.1 Post-LN的结构与训练不稳定问题
Post-LN 的写法,也就是原论文的形式:
class EncoderLayer(nn.Module): def __init__(self, d_model, n_head, d_ff, dropout=0.1): super().__init__() self.self_attn = MultiHeadAttention(d_model, n_head, dropout) self.ffn = PositionwiseFeedForward(d_model, d_ff, dropout) self.norm1 = nn.LayerNorm(d_model) self.norm2 = nn.LayerNorm(d_model) self.drop1 = nn.Dropout(dropout) self.drop2 = nn.Dropout(dropout) def forward(self, x, src_mask=None): h, _ = self.self_attn(x, x, x, src_mask) x = self.norm1(x + self.drop1(h)) # 先加再归一化 h = self.ffn(x) x = self.norm2(x + self.drop2(h)) return x关键在x = self.norm1(x + h):残差相加的结果直接被归一化,然后传给下一层。这意味着每一层的输出都被 LayerNorm 重新缩放到标准正态附近,跨层的信息传递要经过多次归一化。
Post-LN 的问题是深层时梯度不稳定。第 N 层的梯度要穿过 N 次"归一化 + 残差",早期层收到的梯度容易被挤压。原论文靠的是warmup学习率调度(前4000步线性升温)来缓解,去掉 warmup 直接上大学习率,很容易看到 loss 在几百步后突然变成 nan。
5.2 Pre-LN的代码改动与warmup的关系
Pre-LN 把 LayerNorm 挪到子层之前:
def forward(self, x, src_mask=None): h, _ = self.self_attn(self.norm1(x), self.norm1(x), self.norm1(x), src_mask) x = x + self.drop1(h) # 残差通路是干净的 h = self.ffn(self.norm2(x)) x = x + self.drop2(h) return x改动很小,效果差别很大。残差通路上没有任何归一化操作,梯度可以从最后一层直接回传到第一层,早期层不会因为多次归一化而梯度消失。代价是各子层的输入被归一化过,表达能力和 Post-LN 略有不同。
实测下来,Pre-LN 在去掉 warmup、直接用固定学习率的场景下也能稳定收敛,而 Post-LN 基本必须配 warmup。这也是为什么大部分开源实现(尤其是近几年的)默认用 Pre-LN。
有个细节要注意:Pre-LN 的最后一层输出没有经过norm,送进输出投影前应该补一个self.norm,否则输出分布的尺度不一致:
# Transformer.forward 里 memory = self.encoder(src, src_mask) memory = self.norm_enc(memory) # Pre-LN 才需要我见过一份代码忘了这一步,训练 loss 能降但推理时输出概率整体偏移,调了很久才定位到。
5.3 前馈网络中间维度取4倍d_model的实际考量
class PositionwiseFeedForward(nn.Module): def __init__(self, d_model, d_ff, dropout=0.1): super().__init__() self.net = nn.Sequential( nn.Linear(d_model, d_ff), nn.ReLU(), nn.Dropout(dropout), nn.Linear(d_ff, d_model), ) def forward(self, x): return self.net(x)d_ff = 4 * d_model是原论文的设定(512 → 2048)。这个4倍不是精确调出来的,而是一个经验值:它让FFN的参数量约为8 * d_model²,与注意力部分的4 * d_model²大致同量级,两者在总参数里各占一半。如果你把d_ff设得太小,FFN 容量不足,模型整体表现会下降;设得太大,参数量暴涨,在小数据集上很快过拟合。
激活函数方面,原论文用 ReLU,后来的实现多用 GELU。GELU 在负值区域不是硬截断,梯度更平滑,在小批量训练时通常更稳。切换到nn.GELU()只需要改一行,代价是稍微慢一点。
提示:我在小数据集上试过把
d_ff降到2 * d_model,配合dropout=0.2,验证集指标反而比 4 倍好——因为模型容量超过数据量时,正则化的收益大于容量损失。这个参数值得在你的数据上扫一遍[2, 3, 4]倍。
6. 训练循环:从batch拼接到loss忽略位
模型搭好只是第一步,训练循环里的细节同样会决定结果。这一节讲三个最容易出错的地方。
6.1 teacher forcing与右移一位的拼接
序列到序列的训练用 teacher forcing:解码器的输入是"正确输出的前一位",预测目标是"正确输出的当前位"。实现上就是对同一个目标序列做两次切片:
tgt_in = tgt[:, :-1] # 从 BOS 开始,去掉最后一个 tgt_out = tgt[:, 1:] # 从第一个真实 token 开始,去掉 BOS假设目标序列是[BOS, 我, 爱, 学, 习, EOS](长度6)。tgt_in是[BOS, 我, 爱, 学, 习],tgt_out是[我, 爱, 学, 习, EOS]。解码器看到BOS时预测我,看到BOS 我时预测爱,以此类推。这就是因果语言模型的标准训练方式。
这里有个必须注意的点:因果 mask 要和tgt_in的长度对齐。如果 mask 用的是tgt的长度(6)而输入是5,形状不匹配要么报错要么广播成错误结果。我习惯在生成 mask 时始终用tgt_in.size(1)。
6.2 ignore_index、label smoothing与维度拉平的配合
损失计算是另一个坑区。模型的输出是(B, L, V),标签是(B, L),cross_entropy要求输入是(N, V)、目标是(N,):
logits = model(src, tgt_in) # (B, L, V) loss = F.cross_entropy( logits.reshape(-1, logits.size(-1)), # (B*L, V) tgt_out.reshape(-1), # (B*L,) ignore_index=pad_id, label_smoothing=0.1, )reshape(-1, V)这一步的语义是:把 batch 和序列维压平,每个位置独立算一个分类问题。用reshape而不是view,是因为经过前面的转置操作,logits 可能不连续。
ignore_index=pad_id让填充位置的损失不参与反向传播。这一步如果漏了,模型会花大量容量去学"预测 pad",在短序列占比高的数据集上尤其明显。
但这里有个冲突:label_smoothing和ignore_index一起用时,某些早期版本的 PyTorch 会把ignore_index对应的位置也做平滑,导致填充位贡献了一个小的非零损失。我的做法是先不加label_smoothing跑通一遍,确认 loss 曲线合理,再加平滑对比。如果加了平滑之后 loss 不再下降到接近0,就是这个问题。
6.3 学习率预热与Adam的beta参数
原论文的 warmup 调度:
def lr_lambda(step, d_model=512, warmup=4000): step = max(step, 1) return (d_model ** -0.5) * min(step ** -0.5, step * warmup ** -1.5)它是两条曲线取最小值:前warmup步按线性增长(step * warmup^-1.5),之后按step^-0.5衰减。前期的线性升温让参数在最开始几步不会因为随机初始化的大梯度被推得太远。
Adam 的超参方面,原论文用betas=(0.9, 0.98)、eps=1e-9。第二个 beta 从默认的0.999改成0.98,是因为 Transformer 训练前期梯度变化快,0.999 的动量太"黏",对二阶矩的估计滞后。eps设成1e-9而不是默认的1e-8,配合 warmup 在极小的学习率下更稳定。
实测中我把betas换回(0.9, 0.999)对比过,在小的翻译任务上差别不明显,但在深层模型(12层以上)上,0.98的收敛速度和最终指标都更好。
| 超参 | 原论文值 | PyTorch默认 | 建议 |
|---|---|---|---|
| betas | (0.9, 0.98) | (0.9, 0.999) | 深层模型用0.98 |
| eps | 1e-9 | 1e-8 | 配warmup时用1e-9 |
| warmup步数 | 4000 | 无 | 按2~4 * 数据量/batch估 |
| weight_decay | 0 | 0 | 0 或 1e-4 |
7. 推理阶段:自回归解码与KV缓存改造
训练时所有位置并行计算,推理时只能一个 token 一个 token 地生成。这个切换会让形状处理变复杂,也有很多实现上的坑。
7.1 greedy解码的最小实现
@torch.no_grad() def greedy_decode(model, src, src_mask, bos_id, eos_id, max_len=50): model.eval() memory = model.encode(src, src_mask) # (B, L_src, D) ys = torch.full((src.size(0), 1), bos_id, dtype=torch.long, device=src.device) finished = torch.zeros(src.size(0), dtype=torch.bool, device=src.device) for _ in range(max_len - 1): tgt_mask = make_causal_mask(ys.size(1), ys.device) logits = model.decode(ys, memory, tgt_mask, src_mask) # (B, L, V) next_token = logits[:, -1].argmax(dim=-1) # (B,) ys = torch.cat([ys, next_token.unsqueeze(1)], dim=1) finished |= next_token.eq(eos_id) if finished.all(): break return ys注意logits[:, -1]——只取最后一个位置。因为前面的位置在上一轮已经生成过了,因果 mask 保证最后一位的表示包含了全部历史信息。
另外要记得model.eval()和torch.no_grad()。前者关掉 dropout,后者关掉梯度记录。推理时忘了eval()是很常见的失误,表现为同一个输入每次生成的结果都不一样,找半天找不到原因。
7.2 把K与V缓存起来要改哪几行
上面这个实现在每次迭代里都要重新算整个ys的注意力,复杂度是O(L²)每步、总共O(L³)。生成100个 token 时,绝大部分计算是重复的。
KV 缓存的做法是:每一步只算新 token 的 Q、K、V,把新的 K 和 V 拼接到缓存里,注意力用"新的Q"对"全部的K、V"计算。改动集中在MultiHeadAttention:
def forward(self, query, key, value, mask=None, cache=None): q = self._split_heads(self.w_q(query)) k = self._split_heads(self.w_k(key)) v = self._split_heads(self.w_v(value)) if cache is not None: prev_k, prev_v = cache k = torch.cat([prev_k, k], dim=2) v = torch.cat([prev_v, v], dim=2) new_cache = (k, v) else: new_cache = (k, v) scores = torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(self.d_k) # 其余不变拼接的维度是dim=2,也就是序列维(形状是(B, H, L, d_k))。这里最容易错的是拼错了维度——拼到 head 维上,张量形状会变得很奇怪,而且往往不报错,只是注意力算错。
缓存之后每一步的复杂度从O(L²)降到O(L),生成长度100时提速非常明显。代价是显存占用随生成长度线性增长,长序列生成时要留意。
还有一个细节:用缓存时,mask 只需要最后一个位置的对应行,也就是形状(B, 1, 1, L_total)。如果还传完整的 causal mask,形状不匹配。用缓存时我一般直接给上三角 mask 的最后一行。
8. 三次nan与不收敛的排查记录
前面讲的都是"应该怎么写",这一节讲讲写错了会怎样。下面三个问题是我在实际调这份代码时真实遇到的,排查过程比结论更有用。
8.1 全屏蔽行的softmax
第一次 nan 出现在训练到第几十步的时候,loss 突然变成 nan。加了一个 hook 打印每层的输出,定位到注意力那一步。
原因是masked_fill(mask, -inf)之后,某一行的 scores 全是-inf。softmax 遇到整行-inf时,exp(-inf) = 0,分母是0,得到0/0 = nan。
哪来的全屏蔽行?目标序列[BOS, x, EOS, PAD, PAD]里,PAD位置在 causal mask 下,它能看到自己和之前的 R=位置。但如果这个位置本身是 PAD,而tgt_in又做了右移切片,就可能出现"某个位置能看到的全是 PAD"的情况,加上 pad mask 一叠加,整行都被屏蔽。
解决办法是在生成 mask 后加一句保护性检查:
def check_mask(mask): # 找出行内全部为 True 的位置 all_masked = mask.all(dim=-1) if all_masked.any(): print("警告:存在整行被屏蔽的位置,数量:", all_masked.sum().item()) return mask更根本的做法是保证 padding 位置不参与损失(ignore_index已经做到了),同时把它们的注意力输出用一个安全值替代。有些实现会在 softmax 之前把整行的值设为0而不是-inf,代价是 padding 位置会平均关注所有位置,但因为它们的输出不参与损失,不影响结果。
8.2 mask的dtype与device
第二次的问题不报错但结果不对:模型能训,loss 也在降,但验证集上的表现和随机猜差不多。
排查方式是造一个极端用例——只保留两个 token 的有效部分,其余全部 padding,看看模型的输出是否只依赖这两个 token。结果发现改变 padding 区域的内容也不影响输出,说明 mask 生效了;但把 mask 换成 numpy 生成的版本之后,结果又变了,这说明两次生成的 mask 不一样。
问题出在 dtype 上。seq == pad_id得到的是torch.bool,而某处从外部传进来的 mask 是torch.uint8。masked_fill在接收到uint8时,会把大于0的值当作 True,这一般没问题。但mask | causal_mask这种按位或操作,在uint8上做的是位运算而非逻辑运算,结果就可能出错。
修复方式是在 mask 进入注意力之前统一成布尔类型:
if mask is not None and mask.dtype != torch.bool: mask = mask.bool()device 也是一个问题源。mask 在 CPU 上生成、scores 在 GPU 上,masked_fill会隐式跨设备拷贝,在数据量大时明显拖慢,某些版本还会直接抛错。我的习惯是 mask 生成函数接收一个device参数,从源数据所在设备直接生成。
8.3 长序列外推与位置编码越界
第三次是推理时的问题。训练长度128,推理时输入了200个 token 的序列,直接报索引越界:
IndexError: index 150 is out of bounds for dimension 1 with size 128原因是max_len设成了128,self.pe只有128行,切片self.pe[:, :200]只拿到128行,和输入长度不匹配。如果形状匹配得上(比如输入正好128)就不会报错,但语义上已经错了。
修复方式是把max_len设大一些,并在 forward 里加长度检查:
def forward(self, x): if x.size(1) > self.pe.size(1): raise ValueError( f"输入长度 {x.size(1)} 超过位置编码最大长度 {self.pe.size(1)}" ) x = x + self.pe[:, :x.size(1)] return self.dropout(x)显式的报错比形状不匹配带来的隐式错误好得多。位置编码越界最麻烦的情况是输入长度恰好等于max_len的一部分,形状对得上、不报错,但位置编码和 token 的对应关系是错的。这类问题只能靠断言和单元测试发现。
| 现象 | 可能原因 | 排查手段 |
|---|---|---|
| loss 变 nan | 整行被屏蔽 /-1e9在 fp16 溢出 | 打印 mask 的行和、检查 dtype |
| 训练能降但效果差 | mask 语义反了 / 交叉注意力用错 mask | 造极端用例验证依赖关系 |
| 设备不匹配报错 | 位置编码没注册 buffer | 检查register_buffer |
| 推理结果每次不同 | 忘了model.eval() | 检查 dropout 是否关闭 |
| 生成变慢 | 没有 KV 缓存 | 检查每步是否重算全部 K、V |
写完这八个部分,我把这份实现放在小规模数据集上完整跑过:词表8000、d_model 256、4层编码器、4层解码器、8个头,单卡训练两小时左右能收敛到一个合理的水平。再往上加层数时,Pre-LN 加 warmup 的组合是我试过最省心的搭配,几乎不用调学习率就能稳定下来。真正花时间的从来不是写模型本身,而是确认每一个 mask 的形状、每一个张量在设备之间搬对了地方。