Concept

Padding Mask

过滤序列中的 padding token,防止它们参与 Attention 计算和污染 Position Encoding

Padding Mask

Padding Mask 用于处理变长序列对齐时引入的 padding token,确保它们不会影响模型计算。

问题来源

Transformer 需要固定长度的 tensor,但实际序列长度不一:

真实序列:[A, B, C]
Padding后:[A, B, C, PAD, PAD]

Padding 会带来两个问题:

  1. Attention 污染:真实 token 可能读取 PAD 的内容
  2. 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]

Mijvalid=valid(i)valid(j)M_{ij}^{valid} = \text{valid}(i) \land \text{valid}(j)

结果:

          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   0

PAD 行列全为 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

Mij=[拓扑允许][Query 有效][Key 有效]M_{ij} = [\text{拓扑允许}] \land [\text{Query 有效}] \land [\text{Key 有效}]

相关概念

  • [[Attention Mask]]
  • [[Position Encoding]]
  • [[Block-Causal Attention]]