恒美微站 Logo 恒美微站
  • 首页
  • 关于我们
  • 建站服务
  • 主题模板
  • 案例展示
  • 资讯中心
  • 联系我们

多头注意力中的Attention Mask与Causal Mask:前向与反向传播解析

  • 首页
  • 资讯中心
  • /
  • 多头注意力中的Attention Mask与Causal Mask:前向与反向传播解析

相关资讯

Altium Designer PCB设计全流程:从原理图到Gerber的工程实践指南 2026/9/30 6:30:47
基于YOLOv5的智能人脸标注工具:三种模式实现高效预标注与格式导出 2026/9/30 6:30:47
Meta广告有机器人流量了吗?四项检查帮你排查垃圾流量 2026/9/30 6:30:47

最新资讯

聚环氧乙烷‑b‑聚甲基丙烯酸甲酯PEO-b-PMMA:从水相胶束到共混膜改性技术总结
欧司朗透镜怎么选?从光型到装车匹配
工业无线HMI品牌有哪些?从通信、安全到现场应用的选型分析
空窗期第一件事:给自己搭了套CRM
Wi-Fi 联盟认证项目费用构成明细解析
jevgrep 编码代理任务中的未知之未知扫雷:explore-unknowns Stage 4 系统性排查代码雷区实战指南

今日推荐

模型优化器实战:从FP32到INT8的推理加速与精度平衡
LangGraph+FastAPI构建可审计AI编码助手
基于图像预处理与几何特征的人脸脸型发型搭配系统实现

本周热门

从像素到笔画:srt-whiteboard-animation骨架笔迹追踪实现(Zhang-Suen细化+8邻接追踪)
网站建设的英语怎么说?别只背单词,看完这套安全完整流程才敢上线
新手入门看这篇:建设网站加盟避坑指南与SEO实操

本月精选

自研推理加速器Redwood:两周内实现PyTorch模型高效部署的实战教程
V4L2摄像头采集实战:从camera_client.rar到出图全流程解析
从“谁发明了钢琴键”到知识问答智能体:RAG与记忆工程实践

多头注意力中的Attention Mask与Causal Mask:前向与反向传播解析

发布时间:2026/9/30 6:30:47
多头注意力中的Attention Mask与Causal Mask:前向与反向传播解析 上周帮同事排查一个训练 loss 突然飙升的问题模型是标准的 Transformer decoder-only数据没换、超参没换唯一动过的地方是 attention mask 的生成逻辑。最后定位下来问题就出在布尔 mask 的取值约定上——在 PyTorch 的 bool mask 里True 到底是允许看到还是遮住不看很多人从来不关心但在 MHA 的 forward 和 backward 两条路径里这一反就直接引发信息泄漏和梯度错乱。这件事让我想认真写一篇关于 MHA 中 Attention Mask 和 Causal Mask 的文章。网上的教程大多直接甩一段代码说这是 causal mask抄就完了却很少有人讲清楚一个更底层的问题mask 在前向传播forward trace和反向传播back trace这两条路径里到底分别切断了什么。信息流在哪些位置被挡住梯度又在哪些位置归零哪些位置明明被遮了却还能偷到梯度——这些才是 mask 设计的核心。这篇文章会从 MHA 的基本计算流开始分别从 forward trace 与 back trace 两个视角拆解 Attention Mask 和 Causal Mask最后给出可以直接落地的实现、验证脚本和调参经验。适合正在啃 transformer 源码、自己训练小模型、或者和我一样被 mask bug 折磨过的人。读完之后你会明白三个关键问题mask 为什么要加在 softmax 之前而不是之后causal mask 的梯度到底回传到哪里训练和推理时 mask 最容易在哪一步悄悄出错。1. 先弄清 MHA 的计算流mask 插在哪个环节1.1 从 Q/K/V 到 attention 输出的一行行拆解多头的核心计算其实特别朴素。假设输入序列长度是 L每个 token 的维度是 DMHA 先通过三组权重把输入映射成 Q、K、V然后切成 H 个头每个头的维度是 D/H。之后对每个头独立做缩放点积注意力scores Q K^T / sqrt(d_head) weights softmax(scores, dim-1) output weights V我接触过的不少同学代码写了无数遍但问到 scores 的形状、softmax 是沿着哪一个维度归一化还是会卡壳。这里必须钉死scores 的形状是 [B, H, L, S]其中 L 是 query 序列长度S 是 key 序列长度。在自注意力里 L 等于 S在 cross-attention 里 L 是 decoder 长度S 是 encoder 长度。softmax 永远沿着最后一个维度也就是每个 query 对所有的 key 做归一化。而 mask 插入的位置就是在 softmax 之前、对 scores 做处理。这一步的时序很关键先 mask再 softmax最后加权求和。很多人代码里把 mask 写错了位置导致整个注意力分布悄悄变形模型还能train起来只是效果变差极难排查。1.2 两种 mask 的职责边界padding 与 causalityMHA 里的 mask 其实只有两大类职责完全不同。第一类是 Attention Mask最典型的用途是处理 padding。一个 batch 里的样本长度不一样短的样本后面要补 padding token凑成同一个长度才能堆成张量。但 padding token 是假的不该参与注意力计算所以要在 scores 里把接触到 padding 的位置遮掉。这类 mask 是空间上的遮挡跟时间先后无关。第二类是 Causal Mask也叫因果掩码用在 decoder 或者任何自回归模型里。它的逻辑很简单生成第 t 个 token 时只能看见第 1 到第 t 个 token不能看见第 t1 个及之后的 token否则就是作弊——相当于考试时卷子还没翻到后面答案就出现在眼前了。这类 mask 只和位置的前后关系有关跟 padding 无关。理解这两类 mask 的本质区别后你就能明白为什么实际代码里总是两个 mask 叠加使用因为它们解决的是两个正交的问题。1.3 我习惯用两个问题来定义 forward trace 和 back traceforward trace和back trace不是官方术语更像是我在实际调试中养成的一种思考方式。每次拿到一段 attention 代码我脑子里会自动跑两条路径。forward trace 问的是前向传播时位置 i 的输出到底聚合了哪些位置的信息把注意力权重矩阵画出来每一行非零的位置就是 forward trace 能到达的地方。mask 的作用就是提前把某条路封死让信息根本流不过去。back trace 问的是反向传播时某个位置的 loss 梯度能回传给哪些位置因为注意力是软性的、可微的梯度会沿着前向传播的路径逆流回去。前向没走过的地方反向自然没有梯度。但这里有个微妙的坑如果 mask 实现得不对比如在 softmax 之后才乘 0那么前向的信息虽然被削弱了反向的梯度却会因为归一化分母的关系绕一条小路影响到本不该影响的位置。下面两节就分别从这两条 trace 出发先把 Attention Mask 讲透。2. Attention Mask 的 forward trace信息是如何被切断的2.1 mask 矩阵的形状与构造从 [L, S] 到 [B, H, L, S]写代码之前先确定 mask 的形状。不同深度学习框架的约定略有差异但 PyTorch 生态里通常有两种形态。第一种是二维 mask形状 [L, S]直接描述 query 和 key 之间哪一对可见。这种 mask 的好处是不同 batch、不同 head 之间可以共享适合因果 mask 这种纯粹由位置决定的掩码。第二种是四维 mask形状 [B, H, L, S]或者至少带 batch 维度 [B, 1, L, S]。padding mask 必须用这种形态因为每个 batch 样本的 padding 位置都不一样无法用一个公共的二维矩阵描述。我在工程里习惯的做法是先用布尔矩阵表达可见性再用 masked_fill 把它转成浮点掩码。这样语义最清晰也不容易搞混。有一个约定必须提前钉死本文的 bool mask 里True 表示允许看到、参与计算False 表示遮住、不参与。这个约定和 PyTorch 自带的F.scaled_dot_product_attention保持一致后面写代码时不用来回切换心智模型。2.2 为什么必须用 -inf而不是乘 0这是几乎每个新手都会问的问题也是理解 forward trace 的分水岭。假设某个位置要遮住直观的想法是让它的注意力权重等于 0于是有人直接在 softmax 之后的权重矩阵上乘 0。这是错误的做法错得还很隐蔽。原因在于 softmax 是归一化操作它要除以所有 key 位置上的指数之和。如果在 softmax 之后乘 0前向传播时那个位置确实不再传递信息但它在 softmax 归一化时仍然贡献了分母把其他位置的注意力权重也一起稀释了。换句话说被遮住的位置虽然自己没输出信息却偷偷改变了其他位置的信息强度——这在语义上是错的。正确的姿势是在 softmax 之前把被遮住位置的 score 设为负无穷。这样exp(-inf) 0分子为零同时分母也没有它的贡献彻底切断 forward trace。从数学上看等价于把被遮住的位置从归一化里完整剔除。这是整个 mask 机制最核心的一句话mask 必须作用在 logits 上而不能作用在概率上。作用在 logits 上是把这条信息通路连根拔起作用在概率上只是给这条路盖了块布风一吹梯度一传就会露馅。2.3 一个典型例子padding 位置上真的没有信息吗来看一个具体的场景。假设一个 batch 里有一条样本真实长度是 3补了 1 个 padding token序列长度是 4。key 侧的 padding 位置是第 4 个索引 3。那么注意力分数矩阵是 4×4第 4 列应该被整体遮掉。如果不加 masksoftmax 之后第 4 列会有一定的概率值也就是说 query 会从 padding token 里吸收信息。padding token 在 embedding 层通常是全 0或者一个随机初始化的向量模型训练时注意力就可能学到去关注 padding token因为它偶尔能提供错误的梯度信号。加了 -inf mask 之后exp(scores[:, 3] - 1e9) ≈ 0 exp(-inf) 0第 4 列的所有权重恒为 0softmax 的归一化分母也自动避开它。前向传播时每一个 query 都完全看不到 padding 位置。这就是 forward trace 被切断的完整过程。这里我额外提一个容易忽略的点padding mask 同时要管 query 侧和 key 侧。通常我们在 key 侧遮掉 padding 列就够了但如果一个 padding token 作为 query 去查别的 key也会产生无意义的注意力行。在只计算 loss 在真实 token 上的场景里padding 行的输出不影响 loss所以很多实现只遮 key 侧。但如果你做的是需要完整序列输出的任务比如某些序列标注就得把 padding query 也遮掉否则模型会从 padding 行学到奇怪的统计规律。2.4 一个典型例子bool mask 与 float mask 的混用陷阱实际工程里最烦人的不是 mask 的形状而是 bool 和 float 两种形态的混用。PyTorch 的scaled_dot_product_attention接口里attn_mask既支持 bool 类型也支持 float 类型。bool 类型里 True 表示参与False 表示遮住float 类型里 0 表示不偏移-inf表示遮住。这两种语义很容易记反。我见过不止一次有人把 float mask 里应该填-inf的位置填成了 0于是被遮住的位置照样参与 softmax信息悄悄漏过去。还有人把 bool mask 从别的框架迁移过来忘了取反结果想遮住的没遮住想放开的全被遮了模型直接训练崩。所以我有一条规矩在团队代码里mask 统一用一种形态传递内部再显式转换。bool 是人的语义float 是机器的语义人的语义只出现一次剩下的都用masked_fill处理。3. Attention Mask 的 back trace梯度能被 mask 挡住吗3.1 软注意力的梯度传播路径前向传播时信息从 key/value 流向 query反向传播时梯度则从输出流回 query 和所有没有被遮住的 key/value。具体到公式上如果第 j 个 key 被 mask 掉了weights 矩阵里第 j 列就是 0那么输出对 V_j 的偏导直接为 0V_j 收不到任何梯度。对 Q 和 K 那边的梯度情况稍微复杂一点。虽然第 j 列的权重是 0但权重是 softmax 的输出softmax 的雅可比矩阵不是对角阵也就是说权重矩阵中某一行的各个元素之间会互相影响。关键在于被 mask 掉的那个 logit 是-inf它的梯度本身是 0而它对分母的贡献也是 0所以它不会传染给同一行的其他元素。最终结论实际上非常干净mask 掉的位置在前向没有信息流在反向没有梯度流。back trace 完全继承 forward trace 的边界。用大白话说这条路从来没存在过。3.2 被 mask 的位置梯度到底是不是 0说到这里可能有人会较真既然exp(-inf)在计算机里会被处理成 0而且梯度计算时 softmax 的公式里包含输出乘以某个差值的形式那被 mask 位置的梯度是不是严格等于 0我建议你用代码验证一遍而不是光看推导。下面的脚本构造了一个简单的注意力层用一个 bool mask 遮住部分位置然后观察scores的梯度import torch import torch.nn.functional as F torch.manual_seed(42) B, H, L, D 1, 1, 3, 8 q torch.randn(B, H, L, D, requires_gradTrue) k torch.randn(B, H, L, D) v torch.randn(B, H, L, D) mask torch.tensor([[[ [True, True, False], [True, True, False], [True, False, False], ]]]) # [B, H, L, S], Truevisible scores torch.matmul(q, k.transpose(-2, -1)) / (D ** 0.5) scores scores.masked_fill(~mask, float(-inf)) scores.retain_grad() weights torch.softmax(scores, dim-1) out torch.matmul(weights, v) out.mean().backward() print(softmax weights:\n, weights[0, 0]) print(scores grad:\n, scores.grad[0, 0])跑一下你会发现scores.grad在 mask 为 False 的位置严格是 0softmax 权重在那些位置也严格是 0。这验证了一个重要的实操结论mask 一旦正确加在了 softmax 之前反向传播时被遮住的 logit 不会产生任何梯度你不需要在 backward 里做任何额外处理。3.3 最隐蔽的 bugsoftmax 之后再乘 mask前文说了softmax 之后乘 mask 会让 forward trace 没被完全切断那 back trace 会怎样答案是会更糟。用一个具体的例子说明。假设某一行有三个 key 位置score 分别是 [1.0, 2.0, 3.0]第 3 个位置要被遮住。正确做法是把 score 改成[1.0, 2.0, -inf]softmax 后大约是[0.23, 0.63, 0.0]注意这个 0.23 和 0.63 是在只用前两个位置归一化的情况下得到的。错误的做法是在 softmax 之后把这个位置乘 0此时 softmax 是对三个位置归一化的结果是[0.09, 0.24, 0.67]再乘 0 变成[0.09, 0.24, 0.0]。你看前两个位置的权重被第三个位置偷走了从 0.23 缩水到 0.09。这意味着被遮住的位置虽然没有直接输出信息但它通过归一化分母改变了所有其他位置的注意力分布也改变了梯度回传的强度。这种 bug 在训练指标上很难察觉因为模型会慢慢适应这种被污染过的注意力分布但最终效果、特别是长序列上的泛化会明显比正确实现差。我排查过两起类似的 case最后都是用逐层对比 attention 权重分布的方式才定位到。3.4 全 mask 行的 NaN 陷阱还有一个跟 back trace 紧密相关的经典事故某一行被全部 mask 掉softmax 会变成 0 除以 0直接产出 NaN。这个场景在因果 mask 和 padding mask 叠加时特别容易触发。比如一个样本的真实长度是 0空样本某些数据清洗流程会产出这种脏数据或者 decoder 的某个 query 位置对应的可见范围为空又或者实现时把对角线也 mask 掉了、而当前行恰好只该看自己。一旦出现 NaNloss 迅速变成 NaN梯度也全是 NaN整个训练直接报废。我的防御手段有两层。第一层是数据侧保证每个样本至少有一个有效 token第二层是代码侧在 softmax 前给加 -inf 的分母补一个极小值或者干脆在构造 mask 时断言每一行至少有一个 Trueassert mask.any(dim-1).all(), mask has all-False row, will cause NaN in softmax这种防御性检查看着多余但在大规模训练里能帮你省下半天定位时间。4. Causal Mask 的前向与反向单向视线里的两条 trace4.1 下三角矩阵的构造与对角线之争因果 mask 的本质是一个下三角矩阵。长度为 L 的序列第 i 行第 j 列表示 query i 能否看到 key j规则是 j i 时可见j i 时不可见。也就是说第 0 行只有自己能看第 1 行能看 0 和 1最后一行能看到所有历史位置。用 PyTorch 一行就能生成causal_mask torch.tril(torch.ones(L, L, dtypetorch.bool))这里 True 表示可见。对角线上的位置默认是可见的也就是每个 token 能看到它自己。绝大多数实现都保留对角线因为一个 token 自己携带的信息通常是有用的。但也有少数场景会刻意 mask 掉对角线比如某些对比学习或者去噪训练里要求模型不依赖自身表示。这个选择对 forward trace 有直接影响mask 掉对角线后每个位置的信息来源少了一个输出表示会更独立但也可能让训练变难。4.2 forward trace并行训练与自回归推理的不一致因果 mask 带来的第一个结构性现象是训练和推理的 forward trace 不一致。训练时整个序列是一次性并行喂给模型的。虽然 causal mask 限制了每个位置的可见范围但所有位置的计算都是同时完成的第 t 个位置的计算其实不需要等待前面的位置真正生成完毕。这也是 transformer 能高效并行训练的根本原因——我们要的不是过程串行只是结果上保持因果性。推理时情况完全不同。模型必须一个 token 一个 token 地生成第 t 个 token 生成完后把它拼到输入末尾再生成第 t1 个。如果不做任何优化每一步都要重新计算前面所有位置的 Q/K/V复杂度是 O(L^2)长序列根本跑不动。由此引出 KV cache 的概念。因为 causal mask 保证第 t 个 token 只能看到前 t 个 token所以前面位置的 K 和 V 一旦算出来后面完全可以复用不需要重算。这就是为什么几乎所有推理框架都有 KV cache它正是利用了 causal mask 对 forward trace 的限制把重复计算缓存下来。理解了这个你就理解了为什么 KV cache 只在 decoder 的自注意力里有效而在 encoder 的 bidirectional attention 里没法直接用——后者每个位置都能看到全部位置缓存的收益大打折扣。4.3 back trace为什么前面的 token 会被训练得更充分causal mask 对 back trace 的影响比 forward trace 更值得琢磨。在一个 decoder-only 模型里如果 loss 是每个位置上交叉熵损失的加和那么第 t 个位置的 loss 梯度只能回传到位置 0 到 t 的 Q/K/V 上。顺着这个规则捋一遍位置 0 的输出会被位置 1、2、3 直到 L-1 全部看到所以位置 0 会收到来自所有后续位置 loss 的梯度位置 1 会收到位置 1 到 L-1 的梯度最后一个位置只会收到自己的梯度。这就产生了一个很实际的现象序列靠前的 token在训练中累积的梯度信号来源更多被优化得更充分越靠近序列末尾的 token梯度来源越少学习信号越稀疏。这也是为什么很多大模型在长文本上容易出现尾部遗忘或者对开头的引用更准确。如果你发现模型总是记不住开头段的信息除了位置编码的问题causal mask 的梯度分布差异也是重要嫌疑。理解这条 back trace 还有一个实际用途当你做梯度累积、或者给不同位置设计不同 loss 权重时你要意识到 mask 已经天然给了前面位置更多的梯度你再加权的时候要格外小心不要放大这种不平衡。4.4 Causal 与 padding mask 叠加AND 不是 OR实际训练里causal mask 和 padding mask 几乎总是同时出现。合并它们的原则是一个位置只要被其中一个 mask 遮住就应该被遮住。所以两者要用逻辑与AND而不是逻辑或OR。这里要特别小心两者的作用轴。causal mask 是二维的 [L, L]限制的是 query 行能看到的范围padding mask 需要广播到 [B, 1, L, S]限制的是 key 列是否是有效 token。合并时先给 causal mask 扩展 batch 和 head 维度再和 padding mask 做 AND# causal_mask: [L, L], True可见 # key_padding_mask: [B, S], True有效tokenFalsepadding causal causal_mask.unsqueeze(0).unsqueeze(0) # [1, 1, L, L] pad key_padding_mask.unsqueeze(1).unsqueeze(2) # [B, 1, 1, S] attn_mask causal pad # [B, 1, L, S]注意这里的 key_padding_mask 我用 True 表示有效和 PyTorch 官方某些接口里 True 表示 padding 的约定相反。这正是最容易踩坑的地方。我强烈建议在代码里统一用一个约定并在函数注释里写清楚而不是依赖记忆。合并之后还有一个细节padding 列的 causal 矩阵那一列是 False-inf 会覆盖掉 causal 里原本可能是 True 的位置。这没问题因为 padding 本来就不该被看到。反过来causal 矩阵里上三角是 Falsepadding 矩阵里那些位置即使是 True 也无效。AND 操作正好实现了两个条件都满足才算可见。4.5 一个小实验用 PyTorch 验证 masked 位置的梯度理论说了这么多不如亲手验证一次。下面这段代码演示了 causal mask 下位置 2 的 loss 只能回传到位置 0、1、2而位置 2 之后如果有的话不会有任何梯度传到前面。import torch L, D 4, 16 x torch.randn(1, L, D, requires_gradTrue) proj_q torch.randn(D, D); proj_k torch.randn(D, D); proj_v torch.randn(D, D) q x proj_q; k x proj_k; v x proj_v scores q k.transpose(-2, -1) / (D ** 0.5) causal torch.tril(torch.ones(L, L, dtypetorch.bool)) scores scores.masked_fill(~causal, float(-inf)) weights torch.softmax(scores, dim-1) out weights v # 只取第 2 个位置的输出算 loss loss out[:, 2, :].sum() loss.backward() # 梯度应该集中在第 0、1、2 个 token 上第 3 个 token 没有梯度 print(x.grad.abs().sum(dim-1)) # shape [1, 4]x.grad的最后一行对应第 3 个 token应该是 0。因为 causal mask 保证第 2 个位置的输出根本接触不到第 3 个 token 的信息。这个验证脚本我经常用它能在几分钟内确认你的 mask 方向没弄反、后向传播路径符合预期。5. 工程落地的坑与调试经验含集群场景5.1 PyTorch 标准实现与 F.scaled_dot_product_attention 的注意事项现在写 MHA我基本不再手写masked_fillsoftmax而是直接用 PyTorch 2.x 自带的F.scaled_dot_product_attention它会自动选择合适的 kernel包括 flash attention 和 memory-efficient attention。但注意这个函数对 mask 的形态有要求。attn_mask参数支持 bool 和 float 两种。bool 的 True 表示参与float 的值会直接加到 scores 上所以被遮住的位置要用-inf。另外当attn_mask是 float 类型时flash attention 的 kernel 不一定支持PyTorch 可能会回退到普通的 math 实现导致性能下降和额外显存占用。我的经验是如果能用 bool mask 表达就用 bool maskflash attention 对 bool mask 的支持通常更好如果需要 float mask比如想做相对位置 bias mask 叠加可以先把 bias 加到 QK^T 上再单独用 bool mask 遮。还有一个版本差异的问题早期 PyTorch 里attn_mask不支持值和 bias 同时传入最新的版本则支持把 attn_mask 的 float 值作为 bias 加进去。所以升级 PyTorch 版本后mask 相关的行为可能悄悄变化。我在项目里会固定 PyTorch 版本并写一个 mask 相关的单元测试防止这种隐形破坏。5.2 长序列与集群场景分块 causal mask 的边界效应最近的mha 集群讨论绕不开长序列下的分布式推理。当序列长度超过单卡显存能承载的范围注意力计算必须分块每块只负责一部分 query 或 key 的运算。这时候 causal mask 不再是一个简单的下三角而是一个分块下三角矩阵——每个 query 块只能看到它所在块及之前块的 key。这个设计会引入一个经典边界问题块与块之间的衔接处信息访问的粒度变粗了。比如把序列切成每块 1024 个 token第 1024 个 token 的 forward trace 覆盖 0 到 1024第 1025 个 token 在下一个块里它能访问当前块的 key 以及上一个块的 key看起来没问题。但如果某个 query 块只加载了有限的前序 KV比如滑窗 attention 只保留最近 4096 个 KV那么超过窗口的历史信息就会从 forward trace 里消失。很多 streaming 场景下模型忘事根源就在这种分块 mask 的边界截断。训练时同样有坑。长序列训练经常用到 sequence packing 或者 chunked attention如果 causal mask 按整个 packed 序列生成会把两条不同样本的 token 错误地关联起来造成跨样本信息泄漏。正确做法是维护一个序列归属 id同一个 id 内才允许注意力相连不同 id 之间即使位置相邻也要断掉。这是分布式长序列训练里最常见的隐性 bug 之一损失曲线看起来正常但模型总学不好怎么查都查不到原因。5.3 调试技巧三行代码看清 mask 和梯度最后分享一个我一直在用的调试套路。面对任何新的注意力实现我会在第一次跑通后做三件事。第一件打印 mask 矩阵。不要看代码推理直接打印前 8 行的 mask肉眼确认是下三角、上三角还是全 1。这一步能过滤掉一半的方向错误。第二件打印 attention weight 的对角线和最大值。检查对角线位置的权重是否偏高如果某个明显被 mask 的位置权重非零说明 mask 没生效不是加晚了就是数值类型错了。第三件对一个小样本做 backward检查被 mask 位置的 scores.grad 是否为 0。如果非 0说明 softmax 之后又做了某些污染操作或者 mask 根本没有真正切断这条通路的反向传播。把这三步剪成一个测试函数每次改模型结构都跑一遍能省下大量苦工。5.4 常见问题速查表症状常见原因解决办法loss 变 NaN某一行 mask 全 Falsesoftmax 除零断言每行至少一个 True或给分母加 epsilonmask 位置注意力权重非 0mask 加错位置或误用 0 代替 -inf确认在 softmax 之前 masked_fill(-inf)bool mask 语义搞反True/False 与框架约定不一致统一约定 True可见在代码注释里写明训练好推理差训练和推理 mask 不一致或推理时忘了 KV cache 边界检查推理时 mask 是否重新生成确认分块 mask 边界长序列跨样本泄漏sequence packing 忘记按样本 id 断开用序列归属 id 构造 mask保证跨样本不可见flash attention 性能骤降float attn_mask 不被 kernel 支持回退到 math尽量用 bool mask或者把 bias 加到 scores 再单独传 bool mask第 t 个 token 梯度非 0 但位置 t1 有梯度causal mask 方向错误用了上三角打印 mask确认是真·下三角排查 mask 问题的过程里我最深的体会是mask 属于写起来三行、调起来一天的代码。它不像模型结构那样有各种 fancy 的设计但它直接决定了信息流的边界而信息流的边界又决定了梯度流的范围。一个新模型到手第一件事就是把 mask 打印出来看一遍改任何涉及 attention 的结构先画一个 forward/backward trace 的小图再动代码。这个习惯帮我挡掉了至少十次潜在的训练事故。最后再分享一个小技巧在 mask 相关的代码里把assert写满。形状断言、单行非空断言、bool 值域断言一个都别省。这些断言在训练脚本里看着啰嗦但当你的模型跑到第 3 万步才发现 mask 有问题时你会无比怀念这些啰嗦。

关于恒美微站

恒美微站专注于为个体商户、工作室提供极简自助建站服务,让每个人都能轻松拥有专业网站。

快速链接

  • 关于我们
  • 建站服务
  • 主题模板
  • 案例展示
  • 资讯中心

服务项目

  • 可视化建站
  • 拖拽编辑
  • 主题定制
  • SEO 优化
  • 网站托管

联系方式

  • 📍 地址:北京市朝阳区建国路 88 号
  • 📞 电话:400-888-8888
  • ✉️ 邮箱:info@hmyw.cn
  • 🕐 时间:周一至周日 9:00-18:00

© 2024 恒美微站 hmyw.cn 版权所有 | 京 ICP 备 12345678 号