PyTorch SDPA 的两个静默坑:布尔 mask 与 eval dropout
拆清 scaled_dot_product_attention 的 keep-mask 语义、广播维度和 dropout_p 行为,用最小张量测试阻止注意力泄漏。

把手写 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.MultiheadAttention 的 key_padding_mask 中,True 表示 padding、需要屏蔽。把后者直接传给 SDPA,含义正好颠倒。
从 padding mask 构造 keep mask
假设 q、k、v 形状分别为 (N, H, L, E)、(N, H, S, E)、(N, H, S, Ev),上游给出的 key_padding_mask 为 (N, S),其中 True 是 padding:
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_mask 与 is_causal=True。既要 padding 又要因果约束时,可以合成一张布尔 keep mask:
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() 后也仍有随机性:
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 故意放一个很大的数,并把它屏蔽:
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 更早发现注意力泄漏。
相关
也可以看看
- ·9 分钟阅读
torch.compile 图断裂:先定位重编译,再谈模式选择
从 Dynamo guard、graph break 与动态形状入手,建立可复现的编译诊断和冷启动、稳态验收方法。
- ·6 分钟阅读
CE target long dtype 断言
PyTorch 广播让错误 shape 的 loss 能跑起来——silent bug 产生错误梯度;写 forward 时 assert shape 或用 torchdim 习惯。
- ·5 分钟阅读
SDPA backend 与 FlashAttention
Self-attention 的 O(n²) 内存和算力随序列长度平方涨——长序列要算 FLOPs/activation 预算,FlashAttention 和线性近似是工程选项不是免费 lunch。
johan's blog