新闻详情

新闻详情

首页 / 资讯中心 / 详情

PyTorch手写Transformer代码解读:张量形状、Mask与KV缓存

发布时间:2026/9/30 1:39:34来源:尧图网络
PyTorch手写Transformer代码解读:张量形状、Mask与KV缓存
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 四个贯穿全文的张量形状约定代码里所有形状相关的推理都建立在四个约定上。我把它们列出来后面每一节都会回来对照符号含义典型取值Bbatch size32、64L序列长度源或目标50、128Dd_model模型隐藏维度512H注意力头数n_head8d_k每个头的维度等于D // H64关键的约定是所有模块的输入输出统一是(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_model100, n_head8会在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, dropout0.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, maskNone): 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 的输入尺度不随头维度变化梯度分布也就稳定了。这个推导值得自己写一遍如果每个分量方差是1Var(Σ 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 tensors 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 实现里出的bugmask 相关的能占一半以上。它的问题在于论文里只用一句话带过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_id0): # (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, dtypetorch.bool, devicedevice), diagonal1)两个形状设计上的考虑值得说清楚。第一pad mask 生成时插了两个长度为1的维度是为了能和(B, H, L_q, L_k)的 scores 广播。如果不插(B, L)和(B, H, L_q, L_k)广播会按右对齐规则匹配很容易对错维度——而且 torch 的广播在某些情况下不会报错直接算出错误结果这是最阴险的一类bug。第二causal mask 用torch.triu(..., diagonal1)而不是diagonal0因为对角线上的 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被当作0True被当作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_idcausal mask(1, 1, L_tgt, L_tgt)解码器自注意力triu(diagonal1)解码器自注意力合并(B, 1, L_tgt, L_tgt)解码器自注意力上面两张按位或4. 位置编码与词嵌入两个看起来最简单却最容易埋雷的模块注意力机制本身是置换等变的——把输入序列的顺序打乱输出也只是被同样打乱模型完全感知不到位置。所以位置信息必须外挂进去。这一块的代码通常只有十几行但埋的雷一点不少。4.1 正弦位置编码的实现与register_bufferclass 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) 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 上时它还在 CPUx 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, dropout0.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_maskNone): 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_maskNone): 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, dropout0.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配合dropout0.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_indexpad_id, label_smoothing0.1, )reshape(-1, V)这一步的语义是把 batch 和序列维压平每个位置独立算一个分类问题。用reshape而不是view是因为经过前面的转置操作logits 可能不连续。ignore_indexpad_id让填充位置的损失不参与反向传播。这一步如果漏了模型会花大量容量去学预测 pad在短序列占比高的数据集上尤其明显。但这里有个冲突label_smoothing和ignore_index一起用时某些早期版本的 PyTorch 会把ignore_index对应的位置也做平滑导致填充位贡献了一个小的非零损失。我的做法是先不加label_smoothing跑通一遍确认 loss 曲线合理再加平滑对比。如果加了平滑之后 loss 不再下降到接近0就是这个问题。6.3 学习率预热与Adam的beta参数原论文的 warmup 调度def lr_lambda(step, d_model512, warmup4000): 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)、eps1e-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.98eps1e-91e-8配warmup时用1e-9warmup步数4000无按2~4 * 数据量/batch估weight_decay000 或 1e-47. 推理阶段自回归解码与KV缓存改造训练时所有位置并行计算推理时只能一个 token 一个 token 地生成。这个切换会让形状处理变复杂也有很多实现上的坑。7.1 greedy解码的最小实现torch.no_grad() def greedy_decode(model, src, src_mask, bos_id, eos_id, max_len50): model.eval() memory model.encode(src, src_mask) # (B, L_src, D) ys torch.full((src.size(0), 1), bos_id, dtypetorch.long, devicesrc.device) finished torch.zeros(src.size(0), dtypetorch.bool, devicesrc.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)], dim1) 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计算。改动集中在MultiHeadAttentiondef forward(self, query, key, value, maskNone, cacheNone): 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], dim2) v torch.cat([prev_v, v], dim2) new_cache (k, v) else: new_cache (k, v) scores torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(self.d_k) # 其余不变拼接的维度是dim2也就是序列维形状是(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设成了128self.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 的形状、每一个张量在设备之间搬对了地方。
网站建设高端定制企业官网
RELATED

相关资讯

更多精彩内容,欢迎继续阅读

较早相关资讯

最新相关资讯

YOLOv11安防异常行为识别:从检测到报警的实战指南 2026/9/30 5:22:57

YOLOv11安防异常行为识别:从检测到报警的实战指南

简介:目标检测作为计算机视觉的基础任务,通过定位与分类图像中的目标,为行为分析提供底层支撑。以YOLOv11为代表的单阶段检测算法,在推理速度与精度之间取得平衡,适合部署于实时监控场景。在安防异常行为识别中&#x…

阅读更多 →
模型压缩实战:结构化剪枝、INT8量化与知识蒸馏全流程解析 2026/9/30 5:22:57

模型压缩实战:结构化剪枝、INT8量化与知识蒸馏全流程解析

1. 项目起源与整体设计思路1.1 为什么需要 Model-Optimizer说句实在话,2023年之后做模型部署,最头疼的已经不是“训练不出好模型”,而是“模型跑不动”。我手头一个检测模型,在GPU上跑得飞快,可一旦要落到客户现场的Je…

阅读更多 →
STP/RSTP/MSTP与VLAN二层技术:从原理到排障 2026/9/30 5:22:57

STP/RSTP/MSTP与VLAN二层技术:从原理到排障

简介:《网工面试超级宝典》是一份面向网络工程师求职者的面试专项资料,聚焦二层交换技术中的STP与RSTP,从STP出现背景、基本概念、BPDU报文格式到拓扑计算过程均有系统梳理,也覆盖RSTP针对STP不足所做的改进。资源包共1个PDF文件&…

阅读更多 →
CentOS 7 yum源配置与离线源搭建实战:公网、内网、本地ISO全方案 2026/9/30 5:22:57

CentOS 7 yum源配置与离线源搭建实战:公网、内网、本地ISO全方案

1. 先把话说清楚:CentOS7 的源为什么必须自己配CentOS7 配 yum 源这件事,三年前是五分钟的活,现在却成了每台机器落地前都得先过一遍的"开机仪式"。原因不复杂:CentOS 7 在 2024 年 6 月 30 日走完了十年生命周期&#…

阅读更多 →
130.Agent-Agent设计模式-6种主流模式 2026/9/30 5:22:57

130.Agent-Agent设计模式-6种主流模式

摘要:本文介绍了 Agent 的 6 种主流设计模式,包括 ReAct、plan-and-execute、Multi-Agent、Workflow/DAG、Self-Ask 以及 Thinking and Self-Reflection。其中 ReAct 是 LangChain 1.0 之后的默认模式,通过思考与执行循环直至得到最终答案&am…

阅读更多 →
想拍出居延海日出大片?额济纳旗这处水天一色的观鸟秘境别错过 2026/9/30 5:22:50

想拍出居延海日出大片?额济纳旗这处水天一色的观鸟秘境别错过

西北小众文旅持续升温,观鸟摄影目的地迎来新热潮近年来,随着大众旅游需求从传统打卡向深度体验转变,西北原生态自然景观和生态科普类目的地越来越受到游客青睐,尤其是适合摄影创作、亲子研学、自驾休闲的小众秘境,持续…

阅读更多 →

今日资讯

本周资讯

本月资讯

看完文章仍有疑问?

联系尧图顾问,获取一对一建站咨询

立即免费咨询 📞 400-888-8888
📞 ✉