Causal Mask

因果掩码把“当前位置不能偷看未来”写成 attention 的可见性约束。它不是一种正则化,也不是为了让模型更懂时间;它保证训练时的条件分布与生成时逐 token 展开的问题定义一致。

从目标函数到矩阵

自回归目标是:

1
P(x1...xT) = ∏t P(xt | x<t)

训练可以并行计算所有位置,但第 t 行的 attention 只能读取 0...t。长度为 4 时,可见矩阵为:

1
2
3
4
query 0: 1 0 0 0
query 1: 1 1 0 0
query 2: 1 1 1 0
query 3: 1 1 1 1

若将禁止位置加上负无穷,再进行 softmax,未来位置的权重在数值上变为 0;实现也可以用 fused causal kernel 隐式完成同一约束。

训练与 decode 的不同形状

训练时 Tq=Tk=T,三角 mask 让一个完整序列并行计算。decode 时新 query 通常只有一个 token:它可以读取 KV Cache 中的全部历史和自身,不需要重新构造完整 T×T 矩阵。这里的“只看过去”仍然成立,只是过去已被缓存为 K/V。

Padding 与其他可见性规则

causal mask 约束时间方向,padding mask 约束哪些 key 是真实 token;二者可能需要合并并广播到 [B,H,Tq,Tk]。Prefix LM、双向 encoder、span corruption 或 tool-call 特殊 token 可能使用不同的可见性图,不能把下三角矩阵当成所有 Transformer 的默认答案。

实现边界

  • 不同 API 的布尔 mask 可能是 True=禁止True=保留,必须读接口契约并用一个小矩阵断言。
  • mask 的 rank 错一维,广播仍可能成功,却把 batch/head/sequence 约束施加到错误轴。
  • additive mask 的负值要与 dtype 和 kernel 的数值范围兼容;“写一个极小常数”不等同于安全的负无穷。
  • packed sequence、prefix cache 和 paged KV 会改变物理布局,不改变逻辑可见性。

它向上约束 Autoregressive Generation,向下进入 Scaled Dot-Product Attention,横向对照 encoder 的双向 attention 和 prefix LM。

用可观测的 toy logits 验证 mask

不要只检查 mask tensor 的形状。可以令一个 query 的四个 key logits 分别为 [0, 1, 2, 3],在位置 1 应看到 key 0/1,而 key 2/3 的概率必须严格为零;再把同一行放到 batch 1、head 2,确认广播没有改变可见集合。这个测试能同时抓到布尔语义颠倒、上三角/下三角错位和 softmax 轴错误。

Cache、prefix 和滑动窗口的可见性

decode 的单行 mask 并不意味着“永远能读全部 cache”。如果请求采用滑动窗口,最早的 key 已被逻辑丢弃;如果共享 prefix cache,新 query 可以读共享前缀,但不能读另一请求的 suffix。prefix LM 还可能允许一段前缀双向可见、后缀保持因果,这时 mask 是分块矩阵而不是简单三角形。调度器和 kernel 都应以同一份逻辑可见性定义为准,否则会出现内容串请求或隐性信息泄漏。

Mask 的 bug 通常比 shape bug 更危险:数值能正常训练,loss 甚至下降,只是模型学会依赖了不该看到的未来。训练前用极小序列做梯度和可见性断言,训练后用位置置换测试检查是否真的保持因果。