Blog · 2026-08-19

从 make_attn_mask 理解 Attention 的完整机制

make_attn_mask 这个函数在做什么?为什么需要这么复杂的 mask 机制?

读多模态模型源码时遇到这个 Attention Mask 函数:

def make_attn_mask(input_mask, mask_ar):
    mask_ar = jnp.broadcast_to(mask_ar, input_mask.shape)
    cumsum = jnp.cumsum(mask_ar, axis=1)
    attn_mask = cumsum[:, None, :] <= cumsum[:, :, None]
    valid_mask = input_mask[:, None, :] * input_mask[:, :, None]
    return jnp.logical_and(attn_mask, valid_mask)

这个函数要解决的核心问题是:在 Transformer 中,Query i 能否读取 Key j

最终答案由两部分决定:

  1. attn_mask:从 Attention 拓扑看,i 能否看到 j?
  2. valid_mask:i 和 j 是否都不是 padding?

1. mask_ar:定义 Attention Block 边界

先看核心机制 mask_ar。它不是简单的"是否 causal"开关,而是定义 Attention Block 的边界

cumsum 构造 Block 编号

cumsum = jnp.cumsum(mask_ar, axis=1)

假设:

mask_ar = [1, 0, 1, 0, 1, 0, 0]
cumsum  = [1, 1, 2, 2, 3, 3, 3]

mask_ar[i] = 1 表示"从 token i 开始进入新 block",cumsum 把每个 token 映射到它所属的 block 编号:

token:  t0 t1 t2 t3 t4 t5 t6
cumsum:  1  1  2  2  3  3  3
         [blk1][blk2][blk 3]

Attention 规则:block(j) ≤ block(i)

attn_mask = cumsum[:, None, :] <= cumsum[:, :, None]

这里突然出现 [:, None, :] 扩展维度,是因为 Attention 本质上是 Query-Key 的二维关系,而不是单个 token 的一维属性。

  • cumsum[:, None, :][B, 1, S] 表示 Key 的 block
  • cumsum[:, :, None][B, S, 1] 表示 Query 的 block
  • Broadcasting 后得到 [B, S, S] 的关系矩阵

规则:

attn_maskij=[block(j)block(i)]\text{attn\_mask}_{ij} = [\text{block}(j) \leq \text{block}(i)]

即:Query 可以看到自己的 block 和之前所有 block,同一 block 内双向可见

2. 三种典型的 Attention 拓扑

标准 Causal(Decoder-only LM)

mask_ar = [1, 1, 1, 1, 1]  # 每个 token 一个 block
cumsum  = [1, 2, 3, 4, 5]

Attention 矩阵:

          A B C D E
Query A   1 0 0 0 0
      B   1 1 0 0 0
      C   1 1 1 0 0
      D   1 1 1 1 0
      E   1 1 1 1 1

经典下三角 causal mask。每个 token 独占一个 block,自然退化为 token-level causal。

Prefix-LM(多模态常见)

mask_ar = [0, 0, 0, 1, 1, 1]  # prefix 在 block 0
cumsum  = [0, 0, 0, 1, 2, 3]

Attention 矩阵:

          P1 P2 P3 A1 A2 A3
Query P1   1  1  1  0  0  0
      P2   1  1  1  0  0  0   ← Prefix 双向
      P3   1  1  1  0  0  0
      A1   1  1  1  1  0  0   ← 后续 causal
      A2   1  1  1  1  1  0
      A3   1  1  1  1  1  1

这解释了为什么多模态模型的 image/text token 可以互相看到:它们在同一个 block 0 内。

Bidirectional(BERT-style)

mask_ar = [0, 0, 0, 0, 0]  # 全部在同一 block
cumsum  = [0, 0, 0, 0, 0]

全连接,无 causal 限制。

3. 为什么 Mask 作用于 QK 而非 V?

Scaled Dot-Product Attention 分两阶段:

scores = Q @ K.T / sqrt(d) + mask  # mask 这里起作用
attn = softmax(scores)
output = attn @ V

一个具体的例子

假设 Query C 对所有 Key 的原始打分是:

             A    B    C    D    E    F
score       2.0  1.0  3.0  5.0  4.0  6.0

但 causal mask 要求 C 只能看到 A/B/C,不能看到未来的 D/E/F:

             A    B    C    D    E    F
mask         1    1    1    0    0    0

在 softmax 前,把非法位置的 score 替换为 -\infty

             A    B    C     D     E     F
masked      2.0  1.0  3.0  -∞    -∞    -∞

Softmax 后,-\infty 位置的权重自动变为 0:

             A     B     C     D    E    F
attention   0.24  0.09  0.67   0    0    0

最终输出:

OC=0.24VA+0.09VB+0.67VC+0VD+0VE+0VFO_C = 0.24V_A + 0.09V_B + 0.67V_C + 0V_D + 0V_E + 0V_F

虽然 VD,VE,VFV_D, V_E, V_F 本身没有被 mask,但它们的权重已经是 0,自然不会进入输出。

核心洞察:Mask 不是删除 V,而是切断 Query 到某个 Key/Value 位置的信息通路。

  • Q/K:决定从哪里读、权重多少(mask 在这里起作用)
  • V:携带实际内容(通过权重为 0 间接被屏蔽)

只要 Attention 权重为 0,对应的 V 就不会进入输出,无需单独 mask。

4. 多模态模型的信息流拓扑

在 PaliGemma / π0 这类模型中:

prefix = [image_tokens, text_tokens, state_tokens]
suffix = [action_tokens]
 
mask_ar = [0, 0, ..., 0,  1, 1, 1]
          └─── prefix ──┘ └suffix┘

信息流:

┌────────────────────────┐
│ image ↔ text ↔ state   │  block 0: 双向
└────────────────────────┘

    [action_1]               block 1

    [action_2]               block 2

    [action_3]               block 3

Prefix 内充分交互,Action 生成保持 causal。这不是"causal vs non-causal"的二选一,而是通过 cumsum(mask_ar) 定义的信息流拓扑图

5. Padding:最后要处理的实际问题

前面的 attn_mask 假设所有 token 都真实存在。但实际序列长度不一,需要 padding 对齐。

Padding 会破坏什么?

假设真实序列:

[A, B, C]

Padding 后:

[A, B, C, PAD, PAD]

如果不加处理:

  1. Attention 污染:真实 token 可能会读取 PAD 的内容
  2. Position 污染:PAD 占用物理位置,会影响 Position Encoding

valid_mask:过滤 Padding

valid_mask = input_mask[:, None, :] * input_mask[:, :, None]

其中 input_mask = [1, 1, 1, 0, 0] 表示前 3 个是真实 token。

这里再次用到 [:, None, :] 扩展维度,因为需要构造 Query-Key 二维关系

  • input_mask[:, None, :][B, 1, S]:Key 是否有效
  • input_mask[:, :, None][B, S, 1]:Query 是否有效
  • 相乘后 → [B, S, S]Mij=valid(i)valid(j)M_{ij} = \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

两层 Mask 的组合

return jnp.logical_and(attn_mask, valid_mask)

完整规则:

Mij=[block(j)block(i)]valid(i)valid(j)M_{ij} = [\text{block}(j) \leq \text{block}(i)] \land \text{valid}(i) \land \text{valid}(j)

第一层(attn_mask):Attention 拓扑允许吗?
第二层(valid_mask):这两个位置都真实存在吗?

6. Padding 还会污染 Position Encoding

仅仅阻止 Attention 读取 PAD 还不够。假设物理序列:

A B C PAD PAD D E
0 1 2  3   4  5 6  ← tensor index

如果直接用 index 做 position,模型会认为 C 和 D 相距 3 个位置,但逻辑上它们应该相邻。

逻辑位置的计算

positions = jnp.cumsum(input_mask, axis=1) - 1

例如 input_mask = [1, 1, 1, 0, 0, 1, 1]

cumsum     = [1, 2, 3, 3, 3, 4, 5]
positions  = [0, 1, 2, 2, 2, 3, 4]

PAD 不占逻辑位置,真实序列保持连续:

A B C D E
0 1 2 3 4

Position 通过 RoPE 影响 Attention

q = apply_rope(q, positions)
k = apply_rope(k, positions)
scores = q @ k.T

RoPE 使得 QiKjTQ_i' K_j'^T 包含相对位置 pipjp_i - p_j。如果 PAD 占用 position,会人为拉大真实 token 的距离。

核心洞察input_mask 服务于两个目的:

  1. 通过 valid_mask 阻止读取 PAD
  2. 通过 cumsum 计算逻辑位置(避免 PAD 污染 RoPE)

完整机制图

mask_ar → cumsum → block 编号

              [:, None, :] 扩展为二维关系

              attn_mask: block(j) ≤ block(i)

                      ├──────────┐
                      ↓          ↓
input_mask → valid_mask      positions
                ↓                ↓
         过滤 padding         RoPE(Q, K)
                ↓                ↓
                └────→ QK^T + M ←┘

                     Softmax

                      Attn @ V

                      Output

三个独立问题,一套统一机制:

组件解决的问题
mask_arToken 间的通信拓扑是什么?
input_mask哪些位置是真实 token?如何计算逻辑位置?
attn_mask基于拓扑的 Query-Key 关系
valid_mask过滤 padding 的 Query-Key 关系
positions真实 token 的逻辑位置(用于 RoPE)

结论

make_attn_mask 看似简单,实际上统一表达了:

  1. Block-Causal:通过 cumsum(mask_ar) 定义 block,统一实现 causal/prefix-LM/bidirectional
  2. 信息流拓扑:Attention mask 不是"能不能看"的开关,而是定义 token 间的信息流图
  3. Padding 处理:既要阻止读取 PAD,又要避免 PAD 污染 position encoding

理解 Transformer 不是记住"causal mask 长这样、padding mask 长那样",而是理解:

Attention mask 定义了一张 token 间的信息流拓扑图,代码中的每个操作都在构造或过滤这张图的边。