Transformer 中有哪些 Attention Mask?如何正确实现?
简化版
Attention Mask 用于禁止某些 query-key 连接。常见有 Padding Mask(不看补齐 token)、Causal Mask(自回归位置不能看未来)、双向/局部/滑窗 Mask,以及 Prefix-LM、Packed Sequence 和多模态块状 Mask。实现时先统一“1 表示保留还是屏蔽”,再把布尔 Mask 转成加到 logits 的 0/负无穷;要处理广播、缓存偏移和整行全屏蔽,否则会泄漏未来或产生 NaN。
详细版
注意力 logits 为 S=QKᵀ/√d,加法 Mask 在允许位置加 0、禁止位置加足够小值,再 softmax。Padding Mask 通常按 key 屏蔽,Causal Mask 是下三角;两者用逻辑 OR 合并。
P = softmax(S + M)
M[i,j] = 0 if position j is visible to query i
= -infinity otherwise
增量解码使用 KV Cache 时,query 的绝对位置不是从 0 重启,causal mask 必须考虑 past_length。Packed 数据要阻止不同样本互相注意。测试应直接断言被屏蔽概率为 0,并做未来 token 改动不影响过去输出的因果测试,同时覆盖 fp16/bf16、空序列和全 padding。
完整版教学
一、Mask 修改的是连接关系
自注意力默认让每个 query 看到全部 key。Mask 将某些边从注意力图中删除,用同一算子表达因果、padding、局部窗口和分段隔离。
它不是简单删 token:同一个 token 对不同 query 可见性不同。Causal Mask 中,位置 5 能看 0~5,但位置 2 只能看 0~2。
记忆钩子:先画“query 行能看哪些 key 列”,再写 Mask,能避免大多数方向与广播错误。
二、Causal Mask 保证自回归因果性
训练 next-token prediction 时,位置 i 不能读取未来 j>i,否则标签泄漏。允许矩阵是下三角:
key 0 1 2 3
q0 ✓ ✗ ✗ ✗
q1 ✓ ✓ ✗ ✗
q2 ✓ ✓ ✓ ✗
q3 ✓ ✓ ✓ ✓
不同库对 is_causal 与显式 mask 的组合规则不同,重复应用通常不影响语义,但形状或优化路径可能改变,应读 API 契约。
三、Padding Mask 主要屏蔽 Key
batch 中句子长度不同,用 PAD 补齐。真实 query 不应把概率分给 padding key,所以按 key 维屏蔽。padding query 的输出也可能需要在损失和后续层置零。
| Mask | 依赖维度 | 作用 |
|---|---|---|
| Causal | query 与 key 相对位置 | 禁止未来信息 |
| Key padding | 每个样本的有效 key | 忽略补齐内容 |
| Query padding | 输出/损失位置 | 忽略补齐 query |
只屏蔽 query 而未屏蔽 key,会让真实 token 仍读取 PAD 表示。
四、布尔语义与加法语义容易相反
某些 API 中 bool True 表示保留,另一些表示禁止;attention mask 的 1/0 也可能相反。不要凭经验拼接,先用 2×3 小矩阵做单测。
加法 mask 在 softmax 前加入负大值。fp16 中直接构造超范围常量可能溢出,优先用框架提供的 dtype 最小值或原生 mask 接口。
禁止位置归一化后应为 0;在输出后乘 0 无法阻止它参与 softmax 分母。
五、全屏蔽行会产生 NaN
若某个 query 的所有 key 都被设为 -inf,softmax 计算 exp(-inf)/0,结果可能全 NaN。全 padding 样本、空窗口或错误组合 mask 容易触发。
解决方式根据语义选择:保证至少一个合法哨兵位置、跳过空样本,或使用安全 masked softmax 明确定义全屏蔽输出为零。不能随意放开一个真实 token,否则会泄漏信息。
单测要包含全屏蔽边界,而不只测正常句子。
六、Prefix-LM 与块状 Mask
Prefix-LM 允许前缀内部双向注意,生成区可以看全部前缀与自身过去,但不能看未来生成 token。多模态模型还可能让图像 token 双向、文本生成因果。
prefix -> prefix: bidirectional
decode -> prefix: visible
decode i -> decode j: j <= i
这类 Mask 不能用单个普通下三角表示,需要按 token 类型构造块矩阵或使用专用 kernel。
七、Packed Sequence 必须隔离样本
为了提高训练利用率,多个短样本可拼进一个长序列。若只使用 causal mask,后一个样本会看到前一个样本,造成数据串扰。
应记录 segment ID,只允许同 segment 且满足因果条件的连接。位置 ID 是否在段内重置,要与模型和位置编码设计一致。
假设两条长度 3 和 2 的样本拼成长度 5,位置 3 不能看 0~2,即使它们在“过去”。
八、KV Cache 下的偏移与性能
增量解码时只有新 query,key 包含过去缓存。若 past length=100,当前 query 的绝对位置为 100,不能按局部索引 0 构造只看 key0 的错误 Mask。
FlashAttention 等融合内核对因果、窗口和任意 mask 支持范围不同。通用四维 mask 可能使优化内核退化,需同时验证语义与性能。
测试比较无缓存全序列与逐 token KV 解码 logits,应在数值容差内一致。
九、常见误区与追问
- 误区:Mask 在 softmax 后乘零即可。 会改变归一化,允许位置概率不正确。
- 误区:Padding Mask 只屏蔽 padding query。 真实 query 仍可能读取 padding key。
- 误区:所有库中 1 都表示可见。 Bool 与数值语义随 API 不同,必须查契约。
- 误区:下三角可覆盖所有生成任务。 Prefix、多模态和 packed 需要块状/分段规则。
- 误区:KV Cache 每步从位置 0 构造 Mask。 必须加入 past length 的绝对偏移。
- 追问:为什么会出现 NaN? 一整行全是负无穷,softmax 分母为零。
- 追问:训练如何检查未来泄漏? 修改未来 token,断言过去位置 logits 不变。
- 追问:为什么任意 Mask 可能变慢? 融合内核只优化特定稀疏模式,通用 mask 会回退。
十、加强记忆
Attention Mask 记住“行是 query、列是 key、softmax 前屏蔽”:Causal 控制未来,Padding 忽略补齐,Prefix/局部/Packed 定义更复杂连接;先确认 API 中 True/1 的语义,再正确广播与合并。特别检查全屏蔽 NaN、segment 串扰和 KV Cache 位置偏移,并用因果不变性及缓存一致性测试证明实现没有泄漏。