返回专辑
·Johan·3 分钟阅读

PyTorch SDPA 的两个静默坑:布尔 mask 与 eval dropout

拆清 scaled_dot_product_attention 的 keep-mask 语义、广播维度和 dropout_p 行为,用最小张量测试阻止注意力泄漏。

PyTorch SDPA 的两个静默坑:布尔 mask 与 eval dropout

把手写 attention 换成 torch.nn.functional.scaled_dot_product_attention 后,训练能降,验证结果却每次略有变化;再看 padding token,居然还分到注意力。两件事都不会制造 shape error:布尔 mask 的真假语义用反了,以及 model.eval() 并不会替函数式 SDPA 把 dropout_p 设成零。

SDPA 计算的核心是

这里 来自 causal bias 或传入的 attn_mask。对于布尔 attn_mask,PyTorch 的语义是 True 表示允许参与注意力;而 nn.MultiheadAttentionkey_padding_mask 中,True 表示 padding、需要屏蔽。把后者直接传给 SDPA,含义正好颠倒。

从 padding mask 构造 keep mask

假设 qkv 形状分别为 (N, H, L, E)(N, H, S, E)(N, H, S, Ev),上游给出的 key_padding_mask(N, S),其中 True 是 padding:

python
import torch.nn.functional as F


def attention(q, k, v, key_padding_mask, training, dropout):
    keep_mask = ~key_padding_mask[:, None, None, :]
    return F.scaled_dot_product_attention(
        q,
        k,
        v,
        attn_mask=keep_mask,
        dropout_p=dropout if training else 0.0,
        is_causal=False,
    )

keep_mask 的形状是 (N, 1, 1, S),可广播到注意力权重的 (N, H, L, S)。如果每个 head 或每个 query 有不同约束,就显式构造相应维度,别依赖碰巧能广播的二维张量。

浮点 mask 的语义不同:它会直接加到 attention score 上,通常被屏蔽位置填同 dtype 的负无穷。不要把 0/1 浮点张量当布尔 mask 使用;那是在给 score 加偏置,不是在做开关。

当前 API 不接受同时设置 attn_maskis_causal=True。既要 padding 又要因果约束时,可以合成一张布尔 keep mask:

python
causal = torch.ones(L, S, dtype=torch.bool, device=q.device).tril()
padding_keep = ~key_padding_mask[:, None, None, :]
keep_mask = causal[None, None, :, :] & padding_keep

out = F.scaled_dot_product_attention(
    q, k, v,
    attn_mask=keep_mask,
    dropout_p=self.dropout if self.training else 0.0,
)

非方形的 attention 需要重新确认因果对齐方向,不能默认 tril() 总符合缓存布局;增量解码带 KV cache 时尤其要用一个手算小例子核对。

eval() 不会改函数参数

函数式 SDPA 会按传入的 dropout_p 执行 dropout,不读取外层模块的 self.training。因此下面这行即使在 model.eval() 后也仍有随机性:

python
F.scaled_dot_product_attention(q, k, v, dropout_p=0.1)

模块的 forward() 应显式传 self.dropout if self.training else 0.0。这和 nn.Dropout 在 eval 模式自动关闭的行为不同,也是导出前后数值对不上的常见来源。

三个值就能验收 mask 方向

全零,未屏蔽位置的 softmax 权重就相等。第三个 value 故意放一个很大的数,并把它屏蔽:

python
import torch
import torch.nn.functional as F

q = torch.zeros(1, 1, 1, 2)
k = torch.zeros(1, 1, 3, 2)
v = torch.tensor([[[[1.0, 0.0],
                    [3.0, 0.0],
                    [100.0, 0.0]]]])
keep = torch.tensor([[[[True, True, False]]]])

out = F.scaled_dot_product_attention(
    q, k, v,
    attn_mask=keep,
    dropout_p=0.0,
)
torch.testing.assert_close(out, torch.tensor([[[[2.0, 0.0]]]]))

如果把 mask 方向写反,输出会被 100 主导,这个测试不会靠训练“慢慢发现”。再补一条 eval 测试:固定输入连续前向两次应按项目容差一致;若仍不同,先打印实际传给 SDPA 的 dropout_p

全为 False 的一行、混合精度下的负无穷,以及不同设备后端的数值误差也要单测。某个 query 没有任何可参与的 key 时,业务上通常意味着 mask 构造错误;不要只因为输出没有立刻变成 NaN 就接受它。最后用真实 batch 断言 padding 位置的 key 永远不可见,比只看最终 loss 更早发现注意力泄漏。

← 全部文章

johan's blog