
1. 先厘清概念MHA里的Attention Mask到底在“掩盖”什么1.1 从一个最小注意力打分过程说起MHA这三个字母在深度学习圈子里几乎天天见但我发现很多刚接触Transformer的朋友一看到“Attention Mask”“Causal Mask”还是会有点发怵。尤其是标题里还带着“with back forward trace”“with back trace”这种说法第一眼确实容易懵。先把概念拆开揉碎。先说多头注意力Multi-Head AttentionMHA。它的核心动作其实就一句话让序列里的每个token根据它与序列里其他token的相关性去聚合其他token的信息。这里的“相关性”是通过Query和Key的点积打分得到的——Query是“我在找什么”Key是“我能提供什么”两者点积越大说明越相关权重越高再用softmax把这些得分归一化成概率最后拿着概率去加权平均Value里的内容。这个动作展开成矩阵运算就是Attention(Q, K, V) softmax(QK^T / sqrt(d_k)) · V。QK^T会得到一个形状为目标序列长度, 源序列长度的得分矩阵。得分矩阵第 i 行第 j 列的含义就是“目标序列第 i 个token对源序列第 j 个token的关注度得分”。Attention Mask的作用就是在softmax之前的这个得分矩阵上人为地屏蔽掉一些“不允许看”的位置。1.2 back trace与forward trace标题里的这两个词怎么理解我第一次看到“with back forward trace”“with back trace”的时候还以为是某种调试追踪工具。但放在注意力机制的语境下这里的trace意思非常具体它描述的是信息在序列中“追溯——引用——聚合”的方向。back trace向后追溯。当模型处理到位置 i 的token时它可以回过头去查看位置 0 到 i-1 这些历史token把过去的上下文信息拿过来参与计算。这种追溯是“只看过去不看未来”。forward trace向前追溯。当模型处理到位置 i 的token时它可以“预知”位置 i1 到末尾这些未来token把未来的上下文也拉进当前的计算。这种追溯是“既要过去也要未来”。所以标题里的两句话翻译成大白话就是Attention Maskwith back forward trace 双向注意力掩码。当前位置 i 既能回看历史也能前瞻未来整个序列的信息对当前位置全部开放。Causal Maskwith back trace 因果掩码也叫自回归掩码。当前位置 i 只能回看历史包括自己不能看任何未来位置的token。注意力机制本身并不天然限制信息的方向它只是一个“想看哪就看哪”的加权汇总器是掩码人为地给模型装上了“信息可见范围”的约束。而这个约束直接决定了一个Transformer模块是编码器还是解码器。1.3 三种常用的掩码形态全可见、因果、双向在实际工程里我通常把掩码形态分成三大类掩码类型可见范围描述典型场景全可见None/全一所有token互相可见不设限制每个token都能关注整个序列包括未来普通编码器、跨模态交互前的基础注意力双向掩码Attention Mask with back forward trace过去当下未来完整上下文全部开放BERT、情感分析等需要双向语义理解的编码器任务因果掩码Causal Mask with back trace只看过去当下严格下三角结构防止未来信息泄漏GPT等自回归语言模型、时间序列预测从数学上看双向掩码矩阵是个全1矩阵或者不设限制的全0加法掩码因果掩码矩阵是个下三角矩阵——第 i 行只有第 0 到 i 列是可见的第 i1 到末尾全部被屏蔽。这个结构非常简单但背后的信息流约束含义非常深刻。理解了这一点后面去看Flash Attention、稀疏注意力、GQA这些进阶话题都会顺利很多。2. 掩码生效的精确位置与矩阵运算为什么是softmax之前加一个“负无穷”2.1 QK^T 打分矩阵的完整计算过程先设定一个具体例子假设输入序列长度为4特征维度为 d_model 8单头注意力维度 d_k 4。输入张量 X 的形状是batch_size2, seq_len4, d_model8。多头注意力会先把 X 线性投影成 Q、K、V每个头的形状是batch_size2, seq_len4, d_k4。我挑 batch0 来手动推一遍。Q 和 K 在最后一维上的点积结果形状是4, 4其中某个元素 score[i][j] Q[i] · K[j] / sqrt(d_k) 。除以 sqrt(d_k) 是为了防止点积值过大导致softmax梯度饱和这一步通常叫缩放点积注意力。如果当前是因果场景我们要求“第0个token能看第0个第1个token能看第0、1个第2个token能看第0、1、2个第3个token能看第0、1、2、3个”。这个要求翻译成得分矩阵就是要求矩阵的主对角线及以下所有元素保持原样主对角线上方的元素全部被屏蔽。2.2 掩码加在 softmax 之前理论上为什么不能后置很多人会问我能不能先算softmax再把那些不该看的位置的注意力权重设成0理论上可以但实践中几乎不会这么干原因有两个。第一softmax是对整个一行做归一化。如果不先在得分矩阵里把不该看的位置屏蔽掉这些位置的softmax概率就是非零的说明模型确实“分了一部分注意力”给不该看的位置。即使你之后把这一列的概率强行设成0也不得不重新归一化——否则整个权重行的和不为1后面对Value的加权平均就不再是一个概率意义上的混合。你当然可以重新归一化但这等于先算了一个错误的分布再强行修正既浪费计算语义也不干净。第二数值不稳定性。如果得分矩阵里有一行出现了异常大的正数softmax之后会把几乎全部概率都集中到这个位置其余位置的梯度会变得非常小。先加掩码、再进softmax是从源头保证“不该看的位置在softmax之前就已经从数值上被判了死刑”。所以标准实现一定是得分矩阵 → 加掩码把不可见位置设为 -inf 或一个极大的负数→ softmax → 对Value加权。2.3 从数值稳定性聊到 -inf 与 add_mask 的现实选择在实践中“不可见”位置的掩码值默认用 -inf。因为 -inf 进到 softmax 之后exp(-inf) 0对应位置的注意力权重是 0这恰好符合“完全不可见”的语义。但 -inf 会带来两个麻烦。第一个麻烦如果某一行的所有位置都被掩码了softmax里会出现 0/0得到 NaN。所以在实际训练中我们会保证每个序列至少有一个真实token或者额外加一点技巧不直接给掩码位置加 -inf而是给掩码位置加一个足够小的负数比如 -1e9。这样即使所有位置都被掩码分母也是 exp(-1e9)≈0不会出现NaN。最稳妥的做法是保证有价值的位置存在工程上一般不指望靠一个负无穷硬扛边界。第二个麻烦半精度。混合精度训练FP16/BF16里-inf 能表示但如果你用的是很小的负数比如 -1e4在半精度下的表示精度也够而如果用 -1e18有些硬件在极端情况下会出现奇怪行为。现在主流框架的Flash Attention实现里都不是先在内存里构造一个大掩码矩阵再相加而是把掩码信息作为分块索引或布尔矩阵传进去在kernel内部用条件判断跳过屏蔽位置。这样既节省了显存也避免了 -inf 叠加可能带来的精度抖动。我自己踩过的一个坑是在自定义注意力里直接对布尔掩码做 masked_fill(attn_weights, -inf)但在某些深度学习加速卡上-inf 与算子的交互不太稳定。后来统一改成 torch.where(mask, attn_weights, torch.finfo(dtype).min)跨设备行为稳定很多。注意如果训练时掩码忘了加负无穷而用了0通常不会立即报错但会表现为“训练loss偏低、生成严重异常”。这种bug没有exception可看只能靠注意力权重可视化来定位。3. 用 PyTorch 从零实现 Causal Mask并验证“只能向后追溯”3.1 生成因果掩码的三种写法先看最基本的实现。给定序列长度 n因果掩码是一个n, n的布尔矩阵其中第 i 行第 j 列为 True 表示位置 i 可以关注位置 j为 False 表示禁止关注。import torch n 4 # 方式一torch.tril causal_mask_1 torch.tril(torch.ones(n, n, dtypetorch.bool)) # tensor([[ True, False, False, False], # [ True, True, False, False], # [ True, True, True, False], # [ True, True, True, True]]) # 方式二手动构造 causal_mask_2 torch.tensor( [[i j for j in range(n)] for i in range(n)], dtypetorch.bool, ) # 方式三从整数索引生成 row_idx torch.arange(n).unsqueeze(1) # (n, 1) col_idx torch.arange(n).unsqueeze(0) # (1, n) causal_mask_3 row_idx col_idx assert causal_mask_1.equal(causal_mask_2) assert causal_mask_1.equal(causal_mask_3) print(三种方式生成结果一致)方式一是最常用的因为 torch.tril 是专门为下三角掩码设计的原语方式二适合自定义变体比如允许看前k个历史token的滑动窗口掩码方式三适合在动态维度下快速构造例如序列长度是动态时tensor运算比list快很多。3.2 把它接入多头注意力之前先验证掩码本身掩码造出来之后别急着塞进模型。我习惯先单独验证一下掩码的“可见性”对不对方法是对注意力权重矩阵做检查import torch def validate_causal(weights): # weights: (batch, heads, seq_len, seq_len) seq_len weights.shape[-1] mask torch.tril(torch.ones(seq_len, seq_len, dtypetorch.bool)) # 真实注意力权重中掩码位置必须全部为0 assert weights.masked_select(~mask).abs().max().item() 1e-6, 掩码位置出现非零注意力 n 4 q torch.randn(2, 4, 8) k torch.randn(2, 4, 8) v torch.randn(2, 4, 8) scores q k.transpose(-1, -2) / (8 ** 0.5) # (2, 4, 4) mask torch.tril(torch.ones(4, 4, dtypetorch.bool)).unsqueeze(0) # (1, 4, 4) scores scores.masked_fill(~mask, float(-inf)) weights torch.softmax(scores, dim-1) validate_causal(weights) for i in range(4): visible (weights[0, i] 0).sum().item() print(f位置 {i} 的非零注意力个数: {visible})运行结果会显示位置0有1个非零注意力位置1有2个位置2有3个位置3有4个。这正好对应“token i 只能追溯它自己以及它之前的所有token”。3.3 完整实验一个 token 到底能“看见”哪些 token为了让“back trace”这个抽象概念落地我做了一个更直观的实验。把4个token的Value向量设成可区分的结果然后看每个位置最终聚合出来的合成向量里到底混合了哪些位置的Value。import torch # 用可读的向量代替随机数让聚合关系可观察 values torch.stack([ torch.tensor([1.0, 0.0, 0.0]), torch.tensor([0.0, 1.0, 0.0]), torch.tensor([0.0, 0.0, 1.0]), torch.tensor([1.0, 1.0, 1.0]), ]) torch.manual_seed(0) scores torch.randn(4, 4) causal torch.tril(torch.ones(4, 4, dtypetorch.bool)) scores scores.masked_fill(~causal, float(-inf)) weights torch.softmax(scores, dim-1) aggregated weights values print(聚合后的向量形状:, aggregated.shape) for i in range(4): visible_indices [idx for idx in range(4) if weights[i, idx] 0] print(f位置 {i} 可见token: {visible_indices}) print(f位置 {i} 聚合向量: {aggregated[i].detach().numpy().round(3)})由于 forward 部分的权重全部为0最终的聚合向量必然等于 backward 部分的线性组合。位置0的聚合向量只会是 values[0] 的某个缩放位置1只会是 values[0] 和 values[1] 的线性组合以此类推。这就从数值上证明了Causal Mask只允许back trace不允许forward trace。3.4 代码运行时常见的几个低级错误在我带过的项目和社区问答里因果掩码实现出错的高频点主要有三个。第一个错误是掩码形状没对齐。很多新手直接用seq, seq形状的掩码去匹配batch, heads, seq, seq的注意力分数结果广播维度算错了。正确做法是用1, 1, seq, seq或者batch, 1, seq, seq形状参与运算让它在batch和heads维度上自动广播而不是把一个二维掩码直接硬塞进四维分数里。第二个错误是掩码用0而不是用负无穷。如果掩码位置填的是0而不是 -infsoftmax之后这些位置的注意力权重不会归零模型会“偷偷”看到未来从而破坏自回归约束。这个bug的隐蔽之处在于训练loss表面上可以正常下降一旦拿去做生成模型质量就会明显劣化而且很难排查。第三个错误是忘记给掩码匹配 dtype 和 device。要么掩码是在CPU上生成却和GPU张量运算要么用了浮点掩码去和半精度分数做 masked_fill导致warning甚至设备同步卡顿。养成好习惯掩码在生成后就调用 .to(scores.device)并且确认 dtype 与 scores 匹配——如果 scores 是 float16掩码最好是 bool 类型。4. 双向掩码与因果掩码如何影响 Transformer 的架构分工4.1 编码器为什么用“双追溯”BERT 的上下文表征逻辑编码器的任务是给输入序列中的每个位置生成一个“理解了上下文语义”的表示。什么叫理解上下文就是在判断“苹果”到底是水果还是公司的时候需要同时看它前后文说了什么——后文可能说“库克”也可能说“削皮”。如果只看前文信息严重不全。这就是双向注意力with back forward trace的价值让每个位置的向量同时在整合历史和未来信息。BERT预训练里的MLM任务把一部分token遮住让模型根据双向上下文去还原本质上是在逼模型把左右两侧的语义线索拼在一起。可以这么说编码器里的位置i语义表征是“整个序列在i处的投影”而不是“从0到i的序列在i处的投影”。MLM之所以拿不到生成能力也正是因为它依赖了未来信息——在做推断的时候待预测位置后面的token还没生成你根本没法提供forward trace。所以双向注意力适合做“理解”因果注意力适合做“生成”。4.2 解码器为什么必须用“单纯回看”自回归的物理约束自回归生成的过程是“一步步往外吐token”预测第 t 个token时模型手头只有 t-1 个已生成的token第 t 个之后的token根本不存在。如果训练时让它看到未来token它会学着“抄答案”——训练指标漂亮得吓人推理时因为未来token不可得预测瞬间垮掉。“让模型预测下一个词时看到答案等于考前把试卷答案念给考生听”这个比喻虽然调皮但特别准确。Causal Maskwith back trace是在训练阶段强制模型在做第 t 个位置的预测时只能使用第 1 到 t 个token的信息从而模拟推理阶段的真实条件分布。训练时由于teacher forcing可以并行算出所有位置的输出就必须靠掩码来手动干掉“未来信息”这是解码器自回归训练里最关键的物理约束之一。4.3 编码器-解码器注意力里还有一种“交叉追溯”除了自注意力Transformer解码器还有一层跨注意力cross-attention也叫编码器-解码器注意力。这层注意力的Query来自解码器当前正在生成的目标序列Key和Value来自编码器完整的源序列。关键点是跨注意力中的目标token通常可以看编码器的全部输出理论上不需要因果掩码来屏蔽未来但因为Key/Value来自另一条序列仍然需要padding mask去屏蔽编码器里的填充位。我在实际项目里见过不少把因果掩码错误地套在cross-attention上的情况结果发现目标序列之间的位置被错误屏蔽了模型的生成质量显著下降。记住一点因果掩码只针对 self-attention 的自回归建模cross-attention一般不加载因果掩码。4.4 训练与推理阶段掩码的一致性陷阱自回归Transformer在训练时输入是整个序列可能有填充注意力掩码是下三角推理时输入当前上下文时需要处理的关键是KV Cache的存在与否。在推理早期无KV Cache输入长度可能是1此时Q是1个tokenK也是1个token掩码是1,1就是一个标量1——实际上没有未来可看。推理中后期有KV CacheQ长度固定为1K长度是历史累计长度 i此时已经不存在“未来”因为k里只缓存了已经生成过的token未来的token根本不在内存里。于是推理阶段的因果掩码往往简化为全1或者不需要掩码——很多框架在解码时直接不再构造下三角掩码而是在prefill阶段一次性对整段prompt用下三角掩码在decode阶段用长度为1的Q自然避免未来。训练与推理掩码看似“不一致”其实是因果约束在不同阶段的最优实现方式。5. 从掩码细节延伸到真正的工程实践5.1 不要把 padding mask 和 attention mask 混在一起padding mask 解决的是“序列里有填充位填充位没有实际语义”的问题attention mask因果/双向解决的是“信息能否跨越某个方向”的问题。两者经常同时存在。正确的组合方式是两个布尔掩码按位与AND再转成add mask加到注意力得分上# padding_mask: (batch, 1, 1, seq_k)True表示有效 # causal_mask: (1, 1, seq_q, seq_k)True表示可看 combined_mask padding_mask causal_mask # 广播后 (batch, 1, seq_q, seq_k) attn_weights attn_weights.masked_fill(~combined_mask, float(-inf))值得注意的细节padding mask 的行维度取决于 query 位置是否有效。如果 query 位置本身是 padding它的整行注意力没有任何意义——最稳的做法是保证 loss 不会计算 padding 位置的输出或直接对 padding 行的 logits 做特殊处理。5.2 掩码的加法融合add_mask 还是 mask_fill在旧版实现里掩码通常被转成 add_mask一个和scores同形状的浮点张量有效位置0、无效位置一个很大的负数直接用scores scores add_mask。新版PyTorch里越来越多人改成scores.masked_fill(~mask, min_val)减少一次大矩阵加法的开销。近几年的趋势是Flash Attention一类kernel不支持外部传 -inf 浮点 add_mask要求以布尔掩码结构传入内部处理原因是它们把softmax分值分块计算掩码位置在块内直接跳过从而避免了整个 N×N add_mask 的显存占用。对长序列比如8K、16K来说省下来的显存非常可观。这也提醒我们写底层项目时掩码格式需要考虑目标kernel的特定接口不能拿着 old-style add_mask 硬套 Flash Attention。5.3 Flash Attention 下的掩码行为变化Flash Attention 通过 tiling 和 online softmax 把注意力计算拆成多个块在块内应用掩码。因此它的因果掩码使用通常只需要传入is_causalTrue这样的flag框架在kernel内部为每个块判断哪些 (row, col) 属于 causal 范围并跳过不可见块。相比手工构造 (N, N) 掩码再逐元素判断速度快、省显存。不过我用 Flash Attention 时踩过一个坑如果同时需要 padding mask 和 causal mask很多框架接口不直接支持同时传两个掩码。解决方案是把 padding mask 整合进attn_mask布尔或浮点和is_causal组合使用或者使用框架提供的attention_mask参数传一个batch, seq的布尔值让handler自动扩展成合法区域。不同的框架语义差异很大比如 PyTorch 的scaled_dot_product_attention与 HuggingFace 的_flash_attn相关接口参数含义就不一样使用前一定要核对文档和版本。5.4 推理阶段 KV Cache 与因果掩码的配合KV Cache让解码器在生成长文本时避免对历史token重复计算。因为每次只输入当前最新token作为query并且KV缓存中只包含历史的Key和Value因果掩码其实在KV Cache模式下是“天然满足”的——你根本拿不到未来token的Key和Value所以只要在构造KV缓存时保证历史顺序正确不需要额外蒙未来的下三角。但这里有一个隐蔽细节Batch Size。如果你用多路径采样或beam search每条路径都有独立的KV缓存掩码维度要跟随batch维度扩展。我调试多卡分布式注意力掩码时还碰到过一个 batch 维度通过all_gather聚合的情况掩码没有同步分配到所有卡导致最后几层注意力得分里出现了NaN。排查了半天最后的解决方案是确保每个数据并行rank上的因果掩码是设备本地的下三角掩码不要在跨卡通信后重新生成。6. 别把MHA搜成别的一次“集群混淆”与我的实操经验6.1 MHA在NLP与数据库语境下的不同含义先插一段题外话。很多刚入行的同学搜“MHA集群”时会搜到数据库领域的MHAMaster High Availability主库高可用相关内容这是一个常见的语境混淆。在 Transformer/注意力机制的上下文里MHA是 Multi-Head Attention多头注意力的缩写。两者的思维模型完全不同数据库MHA关心的是主从节点故障时的自动切换而注意力掩码关心的是序列信息在矩阵运算中的可见范围。搞清语境至少能省半天搜索时间。在这个语境下本期标题拆开解读MHA多头注意力Attention Maskwith back forward trace双向注意力掩码允许同时回看历史与前瞻未来Causal Maskwith back trace因果掩码只允许回看历史禁止前瞻未来6.2 业务落地经验双向掩码还是因果掩码场景说了算我自己的经验是选哪种掩码不能只看任务类型还要看是否依赖未来。以下是一个粗略的决策表供参考场景推荐掩码原因文本分类、情感分析、阅读理解双向back forward需要完整上下文理解关键词与指代关系文本生成、多轮对话、机器翻译解码端因果back建模自回归条件分布防止未来泄漏机器翻译编码端双向源语言句子的语义需要全局建模时间序列点预测因果预测期只能依赖过去观测不能引入未来语音识别、双向编码器双向语音帧的上下文理解需要左右两侧信息选错掩码的典型症状分类任务里用因果掩码模型F1普遍下滑几个点因为句首token得不到句末信息的语义补充生成任务里不小心去掉因果掩码训练loss极低但生成效果离谱推理阶段输出出现严重重复甚至复读训练样本。6.3 关于掩码调优的几个实战提醒第一注意力权重的可视化是验证掩码最直接的手段。把第t个位置的注意力权重输出出来检查如果t位置对大于t的位置出现非零权重掩码一定写错了。第二长序列训练时建议对掩码做分块处理。比如8K序列的掩码是64M个bool占64MB如果纯float32的add_mask就是256MB显存直接翻倍。用bool掩码在计算时转小数或直接用Flash Attention的is_causal能有效缓解。第三模型剪枝或蒸馏会导致掩码在子模块中的映射错位。我一次蒸馏实验里把某层自注意力替换成简化版本漏掉因果掩码的broadcast导致后面所有层看到未来信息、蒸馏loss居然更低吓出一身冷汗。现在我在架构改动后都会加一道“掩码完整性测试”对简化层单独前向打印可见范围确认与原始层一致。6.4 给后来者的最小清单基于我反复踩坑的经验这里列一份注意力掩码落地检查清单确认当前模块是编码器、解码器还是交叉注意力选择对应的双向、因果或不设限制掩码确认掩码用 -inf 而不是 0确保softmax后概率严格为0确认掩码形状能够广播不会在 batch 或 heads 维度产生错位确认 padding mask 与 attention mask 做了 AND 融合而非默认其一确认 Flash Attention 模式下使用框架支持的掩码格式而非直接传 -inf 矩阵确认 KV Cache 推理阶段是否还需要额外显式掩码确认训练和推理阶段的掩码语义一致——即使实现路径不同。很多注意力问题不在于概念不懂而在于工程实现中掩码与数据形状、设备、kernel之间的细微耦合。每一条在这些验证上吃过亏的人都值得留一份自己的checklist。