Attention Mask与Causal Mask:Transformer注意力掩码的back trace与forward trace
发布时间:2026/9/30 1:06:27来源:尧图网络
如果你在公司群里说“MHA 出问题了”做数据库的同事可能立刻以为是 MySQL 高可用集群要切换而搞 NLP 的同事脑子里冒出来的则是 Multi-Head Attention 多头注意力的矩阵算错了。缩写撞车这件事在我实际工作里真的闹过笑话。这次要聊的就是 MHA 里最不起眼、却最容易被忽略的一环——Attention Mask 和 Causal Mask也就是注意力掩码。这两类掩码通俗点说就是给注意力划一条“视线范围”一个 token 在做注意力计算时到底能看见哪些 token。标题里的 back trace 和 forward trace 描述的就是这件事——回看过去的信息或者前看未来的信息。Attention Mask 通常允许两个方向全部打开back forward traceCausal Mask 则把 forward trace 彻底封死只留 back trace。文章适合正在啃 Transformer 实现、准备训练 BERT/GPT 类模型或者被各种 mask 维度搞到头大的算法和工程同学我会把背后的原理、代码写法、以及实际训练中的各种坑一起梳理清楚。1. 先看整体mask 在 attention 里到底干了什么1.1 MHA 的两个世界以及我们要聊的哪一个MHA 这个缩写确实有点倒霉。在基础设施圈子它代表 MySQL High Availability是一套数据库高可用方案在深度学习圈子里它是 Multi-Head Attention 的缩写是 Transformer 家族的灵魂组件。写这个项目标题的人显然指的是后者否则也不会把 Attention Mask 和 Causal Mask 放在一起讨论。我之所以在文章开头提这个缩写撞车是因为它引出了一个很实际的问题当你和一个 DBA 讨论“MHA 的 mask 写错了”时他可能以为 MySQL 集群脑裂了而你说的其实是一张掩码矩阵。沟通成本这种事技术人应该都懂。回到正题。多头注意力的核心公式是Attention(Q,K,V) softmax(QK^T / sqrt(d_k)) V这里面的 Q、K、V 分别代表查询、键和值。对一整条序列来说每个 token 都会作为 query 去跟序列里所有 key 计算相似度再对 value 做加权求和。这里就出现了一个问题它凭什么都看训练的时候倒是无所谓但推理的时候未来的 token 根本还没生成出来。所以我们需要 mask 来告诉模型某些注意力位置是被禁止的。1.2 back trace 和 forward trace注意力的“视线”到底能看多远在使用 mask 之前得先理解 trace 这个词。如果你把注意力分数矩阵画出来横轴是 key 的位置纵轴是 query 的位置那么从某个 query 出发往矩阵的左下角方向看去是对应它左边那些 key 的注意力分数——这部分就是 back trace往右上角方向看去是对应右边那些 key 的分数——这就是 forward trace。之所以用 trace 这个词是因为在调试注意力的时候你经常会看到模型在哪些历史 token 上留下了“痕迹”。回看过去就是 back trace往前看未来就是 forward trace。一个没有 mask 的注意力矩阵每个 token 都可以同时回溯过去和预读未来而一个带因果掩码的注意力矩阵只有左下三角区域是有值的右上三角全部被压成负无穷。这个视角很有用。很多人把“双向/单向”当成抽象概念画成矩阵之后就非常直观了。Attention Mask 是整张矩阵都有值Causal Mask 是只有下三角有值。后面所有代码和问题排查其实都是围绕这张矩阵的形状和值在做文章。1.3 注意力矩阵上的“加法”mask 不是过滤而是偏置再深入一点要注意 mask 在实现上到底是怎么作用的。它不是在注意力矩阵上做“置零”操作而是在 softmax 之前给不允许的位置加一个极大的负数通常就是负无穷。# 伪代码核心就是 scores mask scores (Q K.transpose(-2, -1)) / math.sqrt(d_k) # [B, H, T, T] scores scores mask # 不允许的位置加 -inf probs torch.softmax(scores, dim-1) out probs V为什么必须是加在 softmax 之前因为 softmax 会把所有有限的数都变成非零概率。如果你在 softmax 之后再置零那些被“屏蔽”的位置在反向传播时梯度依然存在模型还会偷偷从这个位置拿信息。而exp(-inf) 0加在 softmax 之前被屏蔽的位置概率严格为 0梯度也为 0信息才真正断掉。这段逻辑是整个 mask 体系的基石。后面无论遇到什么杂七杂八的 mask 变体回到这张公式图上一切都能解释通。2. Attention Mask让每个 token 同时拥有 back 和 forward 两条视线2.1 双向掩码的适用场景理解任务需要全局视野Attention Mask在没有 padding 干扰的情况下其实是一张全零矩阵。每个位置都能看到序列里所有其他位置既能回看左边的 back trace也能前瞻右边的 forward trace。这种配置对应 Transformer 的 Encoder 部分也就是 BERT 这类双向编码模型的核心。为什么双向信息对理解任务这么重要因为分类、匹配、序列标注这类任务往往需要从整句话里提取语义。举个例子“这家餐厅虽然贵但是好吃”如果你只看“贵”这个词大概率会判断为负面评价但看到后面的“但是好吃”整体的情感极性就翻转了。这就是 forward trace 的价值——未来的 token 改变了当前词的语义判断。反过来一个出现在句尾的词它的含义也经常取决于前文早已出现的指代对象这就是 back trace 的价值。如果你在训练一个不需要生成的模型却错误地用了 Causal Mask表现通常不会立刻崩但会明显变差。因为模型被强行限制成只能从左往右看语义理解的上限被锁死了。选哪种 mask首先要回答的问题是这是个理解任务还是生成任务。2.2 从零写一个双向注意力的 mask 实现双向注意力在掩码层面其实很朴素。不考虑 padding 时mask 可以直接用全零矩阵import torch def build_bidirectional_mask(seq_len, dtypetorch.float32): # 双向可见非 padding 位置全部放行 return torch.zeros(seq_len, seq_len, dtypedtype)把这张 mask 加到 scores 上相当于什么都没做每个位置对序列里所有位置分配注意力权重。这也是为什么很多 Transformer 实现里Encoder 的attn_mask参数默认是None——双向场景下本来就没什么好屏蔽的。但真正到工程实现里你没法绕开 padding。一个 batch 里序列长短不一短序列后面会补[PAD]占位符。这些 padding token 是假的不应该参与注意力计算。所以实际使用的双向 mask 往往长这样def build_padding_mask(input_ids, pad_token_id0): # 有效位置为 Falsepadding 位置为 True valid input_ids ! pad_token_id # [B, T] padding_mask ~valid.unsqueeze(1).unsqueeze(2) # [B, 1, 1, T] return padding_mask # 使用示例 padding_mask build_padding_mask(input_ids) # [B, 1, 1, T] scores scores.masked_fill(padding_mask, float(-inf))注意padding_mask的形状是[B, 1, 1, T]它只屏蔽 key 方向也就是被 attend 的那个维度。masked_fill会自动广播到[B, H, T, T]让每个头上都应用同一个 mask。2.3 padding mask双向注意力里最容易踩的坑padding mask 是花式踩坑的高发区。我第一次自己写 Encoder 的时候忘了加 padding mask结果训练 loss 虽然能下降但下游任务效果始终差一截。把注意力权重打出来看了一下发现 padding token 分走了一部分注意力真实 token 之间的注意力被稀释了。更严重的是padding token 的 value 也会参与加权求和等于往输出里掺入噪声。还有几个具体细节值得注意。第一masked_fill的布尔 mask 是 True 表示要屏蔽跟某些框架里 True 表示允许的语义恰好相反换框架时千万要确认。第二padding mask 通常只需要屏蔽 key 方向因为 query 位置如果是 padding它的输出本来就会被丢掉不影响最终 loss。但如果你要做可视化或者严格对齐最好把 query 方向也屏蔽掉。第三用float(-inf)没问题但在混合精度训练下个别场景可能因为-inf参与加法导致梯度异常备选方案是-1e9。关于这一点后面第五节还会展开聊。2.4 为什么要“向后看”用任务反向理解 mask 设计双向注意力的“向后看”不是拍脑袋定的而是由任务的信息需求决定的。在分类、阅读理解、命名实体识别这类任务里当前 token 的表示需要被整句话的上下文修正。后面出现的词可以改变前面词的语义前面出现的词也会影响后面词的消歧。对比一下人是怎么理解的读一个句子时如果读到结尾发现前面的意思理解错了我们会回头重新解释前面的内容。双向注意力的 back forward trace本质上就是让模型在编码每个 token 时可以同时进行“回看”和“前瞻”一次性完成类似人的全局理解。而生成任务做不到这一点因为未来的词还没被写出来。3. Causal Mask把 forward trace 彻底封死只留 back trace3.1 因果掩码的语义位置 i 只能看见 j iCausal Mask 也叫因果掩码核心语义就是一句话第 i 个 token 只能 attend 到第 0 到第 i 个 token不能 attend 到任何 i 之后的位置。翻译成矩阵就是下三角为有效区域上三角为屏蔽区域。位置0: [O, X, X, X] 位置1: [O, O, X, X] 位置2: [O, O, O, X] 位置3: [O, O, O, O]每一行代表一个 query 位置每一列代表它能否访问某个 key 位置。O 表示允许X 表示屏蔽。这就是标题里说的“Causal Mask with back trace”——只允许回看历史不允许看未来。为什么生成模型必须这样因为自回归生成是从左到右一个一个吐 token 的。训练的时候如果你让第 3 个位置的 token 看到了第 5 个位置的 token模型等于提前偷看了答案推理的时候它没有这个条件训练和推理之间出现 gap生成的文本就会失控。3.2 两种标准实现triu/tril 与 -inf 的正确姿势PyTorch 里构造 Causal Mask 最经典的方式是torch.triu配合-infdef build_causal_mask(seq_len, dtypetorch.float32): mask torch.zeros(seq_len, seq_len, dtypedtype) mask torch.triu(mask, diagonal1).fill_(float(-inf)) return mask这段代码的意思是先生成一张全零矩阵然后取主对角线以上的上三角部分全部填充成负无穷。diagonal1这个参数很关键它表示从主对角线的上一行开始屏蔽所以对角线本身当前 token 看自己是允许的。如果你希望当前 token 也不能看自己就把diagonal改成2。另一种常见写法是用布尔矩阵def build_causal_mask_bool(seq_len): causal_mask torch.triu(torch.ones(seq_len, seq_len, dtypetorch.bool), diagonal1) return causal_mask # True 表示需要屏蔽的位置 scores scores.masked_fill(causal_mask, float(-inf))两种写法结果一样区别只是接口习惯。一些框架的attn_mask参数接受布尔矩阵True 表示屏蔽一些框架接受浮点矩阵直接加到 scores 上。建议在代码里统一成一种并写清楚注释不然两个月后自己看都容易懵。3.3 训练与推理的一致性prefill 阶段为什么也要 mask关于 Causal Mask一个常见的误解是推理时用 KV Cache 一个一个解码天然就是只看历史所以不需要 mask。这个说法只对了一半。逐个解码的时候确实只有当前这一个 token 的 query 在跟历史的 key/value 做点积之后没有更多 token不存在“看到未来”的问题。但推理不只有 decode 阶段还有 prefill 阶段。prefill 是要把用户输入的整段 prompt 一次性丢进模型并行算出每个 prompt 位置的中间状态和 KV Cache。这个时候多个位置是同时参与注意力计算的如果不用 Causal Mask第 3 个位置就会看到第 5 个位置的 prompt token信息就串了。所以结论是Causal Mask 不只是训练专用prefill 阶段也要。你在框架里看到的is_causalTrue这个参数就是为了让模型在并行计算时自动生成因果掩码。KV Cache 只是改变了解码时的计算方式并没有改变注意力的因果约束。3.4 进阶玩法prefix LM 与混合 maskCausal Mask 和 Attention Mask 不是非此即彼的关系它们可以组合成非对称的混合掩码。最典型的例子是 prefix LM也叫前缀语言模型。做法是把一个序列分成两段前半段是前缀允许双向可见后半段是生成区域只允许因果可见。UniLM 和部分多任务模型采用的就是这种思路。换句话说前缀部分的 token 可以看整段前缀但生成部分的 token 只能看自己以及它前面的所有 token。实现也简单就是手动画一张布尔掩码矩阵让对应区域为 False其余为 Truedef build_prefix_mask(seq_len, prefix_len): mask torch.zeros(seq_len, seq_len, dtypetorch.float32) # 前缀区域允许互相可见 mask[:prefix_len, :prefix_len] 0.0 # 生成区域只允许看左侧包括前缀 mask[prefix_len:, :] -float(inf) mask torch.tril(mask, diagonal0).fill_(0) return mask用户体验是前缀部分可以利用双向上下文进行语义理解生成部分又保持了自回归的特性。这种“想让你看到哪就直接画哪”的思路是理解所有高级注意力模式的关键——mask 矩阵本质上就是一张可见性画布。4. 一张表读懂两种 mask选型与背后信息流设计4.1 双向 vs 单向出片对照表把两种 mask 放在同一张表里对比选型的时候会清楚很多对比维度Attention Mask双向Causal Mask单向信息流方向back forward traceback trace only矩阵有效区域全矩阵下三角常见实现全零矩阵或 Nonetriu 上三角填 -inf适用模块Transformer EncoderTransformer Decoder典型模型BERT、ViTGPT、LLaMA典型任务分类、匹配、序列标注语言建模、文本生成是否需要处理 padding必须必须推理阶段是否需要基本不需要通常单次编码prefill 必须decode 视实现而定选型时只需要问自己两个问题我的任务是理解还是生成我训练出来的模型在推理时能不能拿到未来信息如果答案是“能”就用 Attention Mask如果答案是“不能”就用 Causal Mask。就是这么简单。4.2 为什么 mask 对所有 head 是共享的很多初学者会问多头注意力不是每个头各看各的吗为什么 mask 不能每个头不一样答案是mask 编码的是位置之间的可见性是全局的信息流规则而注意力头编码的是语义子空间是在“允许看到的位置”内部去关注不同的特征。两者不是同一个维度。用一个类比来解释所有观众坐在同一个电影院里mask 决定了哪面墙是屏幕、哪面墙是墙这是一条所有人都遵守的物理规则。不同注意力头就像不同观众虽然遵守同一条规则但有人关注主角的脸有人关注背景的风景。如果每个头各自画一块屏幕信息流的物理规则就乱了。不过也有刻意让 mask 因头而异的做法比如某些稀疏注意力模型会给不同 head 分配不同的窗口大小或跳步模式。这种设计属于结构化的先验不是 Transformer 默认行为除非你明确知道自己在做什么否则不要轻易打破共享规则。4.3 用 trace 的语言去设计自己的 mask掌握了 back trace 和 forward trace 这两个词之后你会发现设计新掩码其实是在设计信息流。滑动窗口注意力就是在 Causal Mask 基础上再叠加一个局部窗口只允许每个位置看最近 N 个 tokenLongformer 里则是在双向掩码上挖出几个全局 token 的特殊通道让特定位置可以跟全序列交互。从工程角度建议先画矩阵再写代码。无论你想做全局双向、严格因果、前缀混合、还是窗口稀疏先在草稿纸上把矩阵画出来标注哪些位置是 O 哪些是 X然后转换成triu或tril的组合表达式。大多数掩码 bug都是因为矩阵还没想清楚就开始写代码导致的。5. 实战排查mask 相关的坑我替你们踩过了5.1 mask 到底加在 softmax 前还是后这个问题看起来很基础但我真的见过同事把masked_fill写在 softmax 之后理由是“先把概率算出来再屏蔽不更直接吗”。结果就是被屏蔽的位置仍然有非零概率模型的行为训练时正常上线后开始胡言乱语。原因前面已经说过softmax 对输入是全连通的任何有限值都会得到非零概率。如果你先 softmax 再置零反向传播时置零位置的梯度虽然被截断了但前向计算时概率已经泄漏了。正确顺序必须是先算原始 scores加 mask再 softmax。# 正确 scores scores.masked_fill(mask, float(-inf)) probs torch.softmax(scores, dim-1) # 错误 probs torch.softmax(scores, dim-1) probs probs.masked_fill(mask, 0.0)判断一个 mask 实现是不是对的最快的方法是构造一个极端输入比如让所有位置相同然后看被屏蔽位置的注意力权重是不是严格等于 0。如果不是那就是顺序错了。5.2 -inf 与 NaNfp16、全行被屏蔽、FlashAttention 的边界问题-inf虽然是个标准做法但它在工程上有几个需要注意的边界情况。第一如果某一行所有位置都被屏蔽了softmax 算出来就是nan因为分母是 0。这种情况通常发生在序列长度大于实际有效长度、padding mask 和 causal mask 叠加不当的时候。解决方案是保证每一行至少有一个合法位置。比如因果掩码下对角线默认是允许的那么合法位置至少有一个但如果你把 padding 和 causal 叠加而某个 padding 之前的有效 token 全被覆盖了就可能出问题。第二混合精度训练里-inf的加法在个别硬件上可能触发异常值。我习惯的替代方案是-1e9它是一个足够大的负数exp(-1e9) 0在 fp16 下也不会溢出。但注意如果你的输入特征值很大使得原本的注意力分数也是十亿级别那-1e9就不够用了这种情况还是得回归-inf。第三使用 FlashAttention 时mask 的传递方式跟手动写法不太一样。torch.nn.functional.scaled_dot_product_attention提供了attn_mask和is_causal两种参数is_causalTrue会让底层 kernel 自动生成因果掩码显式的attn_mask则允许你传入自定义布尔矩阵或浮点矩阵。如果是布尔矩阵FlashAttention 内部能跳过被屏蔽的块性能和显存都会更好。用自定义 mask 时优先传布尔类型不要自己先去 fill-inf再传进去。5.3 常见问题速查表一页纸定位 mask 问题我把实际开发里遇到频率最高的几个 mask 问题整理成了一张表方便快速定位现象可能原因解决建议训练 loss 正常但生成内容大量重复Causal Mask 没生效模型跳过掩码直接全局可见检查 mask 是否加到 scores 上并打印矩阵确认上三角是 -inf分类任务效果上不去注意力可视化发现 pad 位置权重高padding mask 缺失或形状错误确认 key 方向掩码为[B, 1, 1, T]并用masked_fill正确屏蔽mask 维度广播报错mask 形状与 scores 形状不匹配scores 是[B, H, T, T]mask 需至少能广播到该形状常见是[T, T]fp16 训练出现 NaN某行全被屏蔽或者 -inf 参与计算导致异常检查每行合法位置数量考虑改用 -1e9 或调整 padding 策略相同代码在不同少批次下结果不一致mask 为 float 类型且包含 -inf某些底层 kernel 行为不同统一使用 bool mask 传给框架原生注意力 API似乎没有 mask但模型也能生成底层 API 可能自动套用了is_causalTrue传入的 mask 被忽略确认 API 参数是否冲突attn_mask与is_causal不要同时传这张表虽然不能覆盖所有花式问题但能覆盖我遇到过的高频问题。如果你碰到的情况不在表里优先去打印 mask 矩阵本身看看它是不是你想象中的形状和值。5.4 一个调试小技巧把 mask 画成热力图再训练最后分享一个我坚持了很久的习惯拿到一个新的掩码设计时不要在 GPU 上直接跑完整训练。先构造一个很小的人工序列比如 seq_len8把 mask 矩阵打印出来或者画成热力图然后手推一遍每行合法位置的数量跟预期对比。import torch def debug_mask(mask): # mask: [T, T] 或 [B, H, T, T]0 表示允许-inf 表示屏蔽 if mask.dim() 4: mask mask[0, 0] valid_counts (mask 0).sum(dim-1) print(每行合法位置数量:, valid_counts.tolist())如果是因果掩码合法位置数量应该是1, 2, 3, ..., T如果是双向掩码且没有 padding应该全是T如果有 padding最后一行的合法位置应该是有效长度而不是T。这一步能过滤掉八成以上的“看起来对但训练就崩”的掩码 bug比我见过的一些靠肉眼盯代码快得多。最后再分享一个小技巧。我在调试自注意力时会额外写一条断言在因果掩码下valid_counts必须严格等于torch.arange(1, seq_len 1)。一旦掩码因为某个上游逻辑被意外修改这条断言会第一时间把问题暴露出来而不是等到训练三个小时之后才发现 loss 异常。别小看这个动作它能帮你省下大量排查时间。顺便说一句如果哪天你在公司跟 DBA 说“MHA 的 mask 写错了”记得确认一下你们聊的是多头注意力还是 MySQL 集群——这种沟通成本能避免的还是尽量避免。
网站建设高端定制企业官网