读多模态模型源码时遇到这个 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?
最终答案由两部分决定:
- attn_mask:从 Attention 拓扑看,i 能否看到 j?
- 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 的 blockcumsum[:, :, None]→[B, S, 1]表示 Query 的 block- Broadcasting 后得到
[B, S, S]的关系矩阵
规则:
即: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 替换为 :
A B C D E F
masked 2.0 1.0 3.0 -∞ -∞ -∞Softmax 后, 位置的权重自动变为 0:
A B C D E F
attention 0.24 0.09 0.67 0 0 0最终输出:
虽然 本身没有被 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 3Prefix 内充分交互,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]如果不加处理:
- Attention 污染:真实 token 可能会读取 PAD 的内容
- 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]:
结果:
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)完整规则:
第一层(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 4Position 通过 RoPE 影响 Attention
q = apply_rope(q, positions)
k = apply_rope(k, positions)
scores = q @ k.TRoPE 使得 包含相对位置 。如果 PAD 占用 position,会人为拉大真实 token 的距离。
核心洞察:input_mask 服务于两个目的:
- 通过
valid_mask阻止读取 PAD - 通过
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_ar | Token 间的通信拓扑是什么? |
input_mask | 哪些位置是真实 token?如何计算逻辑位置? |
attn_mask | 基于拓扑的 Query-Key 关系 |
valid_mask | 过滤 padding 的 Query-Key 关系 |
positions | 真实 token 的逻辑位置(用于 RoPE) |
结论
make_attn_mask 看似简单,实际上统一表达了:
- Block-Causal:通过
cumsum(mask_ar)定义 block,统一实现 causal/prefix-LM/bidirectional - 信息流拓扑:Attention mask 不是"能不能看"的开关,而是定义 token 间的信息流图
- Padding 处理:既要阻止读取 PAD,又要避免 PAD 污染 position encoding
理解 Transformer 不是记住"causal mask 长这样、padding mask 长那样",而是理解:
Attention mask 定义了一张 token 间的信息流拓扑图,代码中的每个操作都在构造或过滤这张图的边。