拓十年匠心定制 · 商业建站与技术教学双线并行 咨询热线:400-886-1026 service@lmnt.cn
ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

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

多头注意力中的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 的维度是 D,MHA 先通过三组权重把输入映射成 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 与 causality

MHA 里的 mask 其实只有两大类,职责完全不同。

第一类是 Attention Mask,最典型的用途是处理 padding。一个 batch 里的样本长度不一样,短的样本后面要补 padding token,凑成同一个长度才能堆成张量。但 padding token 是假的,不该参与注意力计算,所以要在 scores 里把接触到 padding 的位置遮掉。这类 mask 是"空间上的遮挡",跟时间先后无关。

第二类是 Causal Mask,也叫因果掩码,用在 decoder 或者任何自回归模型里。它的逻辑很简单:生成第 t 个 token 时,只能看见第 1 到第 t 个 token,不能看见第 t+1 个及之后的 token,否则就是作弊——相当于考试时卷子还没翻到后面,答案就出现在眼前了。这类 mask 只和位置的前后关系有关,跟 padding 无关。

理解这两类 mask 的本质区别后,你就能明白为什么实际代码里总是两个 mask 叠加使用,因为它们解决的是两个正交的问题。

1.3 我习惯用两个问题来定义 forward trace 和 back trace

"forward 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 列应该被整体遮掉。

如果不加 mask,softmax 之后第 4 列会有一定的概率值,也就是说 query 会从 padding token 里"吸收"信息。padding token 在 embedding 层通常是全 0,或者一个随机初始化的向量,模型训练时注意力就可能学到"去关注 padding token",因为它偶尔能提供错误的梯度信号。加了 -inf mask 之后:

exp(scores[:, 3] - 1e9) ≈ 0 exp(-inf) = 0

第 4 列的所有权重恒为 0,softmax 的归一化分母也自动避开它。前向传播时,每一个 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 的偏导直接为 0,V_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_grad=True) 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], True=visible 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 的位置严格是 0,softmax 权重在那些位置也严格是 0。这验证了一个重要的实操结论:mask 一旦正确加在了 softmax 之前,反向传播时被遮住的 logit 不会产生任何梯度,你不需要在 backward 里做任何额外处理。

3.3 最隐蔽的 bug:softmax 之后再乘 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 掉了、而当前行恰好只该看自己。

一旦出现 NaN,loss 迅速变成 NaN,梯度也全是 NaN,整个训练直接报废。我的防御手段有两层。第一层是数据侧,保证每个样本至少有一个有效 token;第二层是代码侧,在 softmax 前给加 -inf 的分母补一个极小值,或者干脆在构造 mask 时断言每一行至少有一个 True:

assert mask.any(dim=-1).all(), "mask has all-False row, will cause NaN in softmax"

这种防御性检查看着多余,但在大规模训练里能帮你省下半天定位时间。

4. Causal Mask 的前向与反向:单向视线里的两条 trace

4.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, dtype=torch.bool))

这里 True 表示可见。对角线上的位置默认是可见的,也就是每个 token 能看到它自己。绝大多数实现都保留对角线,因为一个 token 自己携带的信息通常是有用的。但也有少数场景会刻意 mask 掉对角线,比如某些对比学习或者去噪训练里,要求模型不依赖自身表示。这个选择对 forward trace 有直接影响:mask 掉对角线后,每个位置的信息来源少了一个,输出表示会更"独立",但也可能让训练变难。

4.2 forward trace:并行训练与自回归推理的不一致

因果 mask 带来的第一个结构性现象,是训练和推理的 forward trace 不一致。

训练时,整个序列是一次性并行喂给模型的。虽然 causal mask 限制了每个位置的可见范围,但所有位置的计算都是同时完成的,第 t 个位置的计算其实不需要等待前面的位置真正"生成完毕"。这也是 transformer 能高效并行训练的根本原因——我们要的不是过程串行,只是结果上保持因果性。

推理时,情况完全不同。模型必须一个 token 一个 token 地生成,第 t 个 token 生成完后,把它拼到输入末尾,再生成第 t+1 个。如果不做任何优化,每一步都要重新计算前面所有位置的 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=有效token,False=padding 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 矩阵里上三角是 False,padding 矩阵里那些位置即使是 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_grad=True) 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, dtype=torch.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_fill+softmax,而是直接用 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 mask,flash 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 全 False,softmax 除零断言每行至少一个 True,或给分母加 epsilon
mask 位置注意力权重非 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 但位置 t+1 有梯度causal mask 方向错误,用了上三角打印 mask,确认是真·下三角

排查 mask 问题的过程里,我最深的体会是:mask 属于"写起来三行、调起来一天"的代码。它不像模型结构那样有各种 fancy 的设计,但它直接决定了信息流的边界,而信息流的边界又决定了梯度流的范围。一个新模型到手,第一件事就是把 mask 打印出来看一遍;改任何涉及 attention 的结构,先画一个 forward/backward trace 的小图再动代码。这个习惯帮我挡掉了至少十次潜在的训练事故。

最后再分享一个小技巧:在 mask 相关的代码里,把assert写满。形状断言、单行非空断言、bool 值域断言,一个都别省。这些断言在训练脚本里看着啰嗦,但当你的模型跑到第 3 万步才发现 mask 有问题时,你会无比怀念这些"啰嗦"。

返回列表