Padding Mask
Padding Mask 用于处理变长序列对齐时引入的 padding token,确保它们不会影响模型计算。
问题来源
Transformer 需要固定长度的 tensor,但实际序列长度不一:
真实序列:[A, B, C]
Padding后:[A, B, C, PAD, PAD]Padding 会带来两个问题:
- Attention 污染:真实 token 可能读取 PAD 的内容
- Position 污染:PAD 占用物理位置,影响 Position Encoding
解决方案
1. 过滤 Attention(valid_mask)
input_mask = [1, 1, 1, 0, 0] # 1=真实, 0=padding
valid_mask = input_mask[:, None, :] * input_mask[:, :, None]构造二维关系矩阵 [B, S, S]:
结果:
A B C PAD PAD
Query A 1 1 1 0 0
B 1 1 1 0 0
C 1 1 1 0 0
PAD 0 0 0 0 0
PAD 0 0 0 0 0PAD 行列全为 0,不参与 Attention。
2. 计算逻辑位置(positions)
positions = jnp.cumsum(input_mask, axis=1) - 1例如 input_mask = [1, 1, 1, 0, 0, 1, 1]:
positions = [0, 1, 2, 2, 2, 3, 4]PAD 不占逻辑位置,真实 token 保持连续编号。这避免了 PAD 污染 RoPE 等相对位置编码。
为什么需要二维 mask?
Attention 需要回答 Query i 能否读取 Key j,这是两个位置的关系。
一维 [B, S] 只能表达"token i 有效吗",无法表达"i 和 j 之间能否通信"。
通过 [:, None, :] 扩展维度:
input_mask[:, None, :]→ Key 是否有效input_mask[:, :, None]→ Query 是否有效- 相乘 → Query-Key 关系矩阵
与其他 Mask 的组合
实际使用中,Padding Mask 通常和 Causal Mask / Block-Causal Mask 组合:
final_mask = attn_mask & valid_mask
相关概念
- [[Attention Mask]]
- [[Position Encoding]]
- [[Block-Causal Attention]]