训练没有报错,梯度却悄悄变了
一次 FlashAttention、Checkpoint 与 Dropout 的排障记录
最近我排查了一个很有迷惑性的训练问题:训练没有立即报错,前期损失也看不出明显异常,但切换到 FlashAttention 后,梯度范数会慢慢高于对照实现,继续跑下去就开始发散。只要关闭 attention dropout,训练又能恢复正常。
沿着这个表象,我最先怀疑的是 FlashAttention 的反向计算。毕竟问题只在 Flash 路径上出现,又和 dropout 高度相关,很容易把它归到融合 kernel 的数值问题上。
但最后查下来,FlashAttention 的数学计算本身并没有出错。问题出现在这三个条件的交集里:
torch.compile
+ activation checkpoint(反向时重算前向,用计算换显存)
+ 融合 attention kernel 内部的 dropout
在我测试的 PyTorch 2.9.0 环境中,编译后的 checkpoint 在重算 attention 时,没有复现第一次前向所用的 dropout mask。前向与反向不再对应同一次随机采样,梯度也就悄悄走偏了。
最后的处理方式并不复杂:在调用融合 kernel 之前显式生成 seed,再把 seed 传给 FlashAttention。这样既保留了融合计算和完整重算,也没有在基准中测到额外的峰值显存,单步耗时只增加了约 0.5%。
先别急着给 FlashAttention 定罪
当时有三条现象都指向 FlashAttention:
- 换成 Flash 后,grad norm 明显变大;
- 问题会在训练一段时间后累积成发散;
dropout=0时训练正常。
不过,这些信息只能说明故障和 Flash 路径里的 dropout 有关,还不能说明 kernel 算错了。Dropout 同时牵涉前向采样、checkpoint 重算和 backward 重放,任何一环没有对齐,最后看起来都像是“梯度数值不稳定”。
所以我先把 kernel 单独拿出来测。在参考测试里,当前构建的前向和反向梯度都能与对照实现对齐。没有 checkpoint 时,FlashAttention 和 SDPA 的短训练结果也基本一致。
这一步排除了一个很有诱惑力的错误方向:问题虽然在 Flash 路径暴露,但不一定发生在 Flash 的数学计算里。
把变量一个个拆开
接下来我把 compile、checkpoint 和 dropout 三个因素分别开关,得到下面这组结果:
| 编译 | Checkpoint | Attention dropout | 结果 |
|---|---|---|---|
| 开 | 关 | 开 | 与对照一致 |
| 开 | 开 | 关 | 与对照一致 |
| 关(eager) | 开 | 开 | 与对照一致 |
| 开 | 开 | 开 | 梯度明显偏离 |
到这里,范围已经缩得很小:单独使用 compile 没问题,单独使用 checkpoint 没问题,FlashAttention 带 dropout 也没问题;只有三者同时出现时才会触发。
但这仍然不能回答一个关键问题:是 FlashAttention 特有的问题,还是所有融合 dropout 都会遇到?
真正让我改变排查方向的是下一组对照。
我让不同 attention 实现走同一套 compiled-checkpoint 测试,用相同输入、权重和初始随机状态,把 compiled + checkpoint 的梯度与同一编译路径下的 compiled + no-checkpoint 比较。下面的“相对误差”,就是两组梯度之差的范数除以参考梯度范数。
结果如下:
| Dropout 实现 | 梯度相对误差 | 结果 |
|---|---|---|
显式 F.dropout | 约 0 | 与参考一致 |
手写 matmul + softmax + F.dropout | 约 0 | 与参考一致 |
| SDPA 数学后端 | 约 0 | 与参考一致 |
| FlashAttention 内部融合 dropout | 约 0.30 | 明显偏离 |
| 融合 SDPA 内部 dropout | 约 0.32 | 明显偏离 |
这个结果很关键。出现问题的不只有 FlashAttention,融合 SDPA 也一样;反过来,只要 dropout 以普通 PyTorch 算子的形式明确出现在计算图里,梯度就是对的。
于是问题的问法变了:
不再是“FlashAttention 的 dropout 算对了吗”,而是“编译器知道融合 kernel 里面发生了一次随机采样吗”。
Checkpoint 重算为什么怕随机数
Activation checkpoint 的思路很简单:前向时少保存一些中间结果,反向需要时再把那段前向计算执行一遍,用计算换显存。
对普通的确定性运算来说,输入相同,重算结果自然相同。Dropout 不一样。第一次前向会随机生成一张 mask,重算时即使输入完全相同,只要随机数状态不同,就会得到另一张 mask。
所以,带 dropout 的 checkpoint 要正确,至少要满足:
重算时的 dropout mask = 第一次前向的 dropout mask
在本文的单设备测试里,eager checkpoint 会保存并恢复相应的随机数生成器状态,因此重算能够复现同一张 mask。
torch.compile 走的是另一套机制。在这组测试所走的编译路径里,显式的 F.dropout 能被识别并正确重放;它所使用的随机性对编译器是可见的。
FlashAttention 的 dropout 则发生在 C++/CUDA 融合 kernel 内部。对外层编译图来说,它只是一次不透明的自定义算子调用。编译器看得到 attention,却看不到 kernel 里面还消费了随机数。
导出计算图后还能看到另一个细节:这次 Flash 调用位于 checkpoint 包住的子图中,而测试版本的随机数分析没有继续进入这个子图。算子本身没有暴露随机语义,分析又看不到它,重算时自然不会替它恢复原来的随机状态。
于是实际发生的是:
第一次前向:随机状态 A → mask A → 中间结果 A
反向前重算:随机状态 B → mask B → 中间结果 B
两张 mask 单独看都合法,问题在于它们被放进了同一条前向—反向链路。重算得到的部分中间结果来自 mask B,而 backward 所依赖的随机信息仍描述原来的前向。两边不再匹配,梯度就失去了正确语义。
为什么几个直觉修复都没用
定位到这里之前,我试过一些看起来合理的办法。
我先调整了计算图的包装方式,包括去掉 allow_in_graph 和设置 fullgraph=False。但自定义算子仍留在编译路径中,kernel 内部的随机数也没有因此变得可见。
我也试过改变 rng_state 的保存位置,以及切换 checkpoint 的重入模式。前者只是在挪动状态,后者只是在换一种重算方式,都没有让重算复现原来的 mask。
另一个容易混淆的开关是 FlashAttention 的 deterministic=True。它主要影响 backward 算法的确定性,不会固定 forward dropout 使用的 seed。只要 seed 变了,mask 还是会变。
这些尝试有一个共同点:它们都在调整算子怎么被包住、状态放在哪里,却没有把真正缺失的信息补进来——重算进入融合 kernel 之前,必须拿到与第一次前向相同的 seed。
最终修复:把 seed 变成显式输入
既然编译器看不见 kernel 内部的 dropout,我最后选择不再要求它“猜”。做法是在 kernel 外部生成一个编译器能够识别的 seed,并把它作为输入传给自定义 FlashAttention 算子。
概念代码大致如下:
def flash_attention_with_replayable_dropout(q, k, v, dropout_p):
seed = torch.randint(
0,
2**62,
(1,),
dtype=torch.int64,
device="cpu",
)
return torch.ops.example.flash_attention_seeded(
q, k, v, dropout_p, seed
)
在 custom op 内部,再用这个 seed 初始化专用随机数生成器,并把它交给 FlashAttention 前向。
这样做能成立,靠的是两层保证:
第一,外层的 torch.randint 是编译器认识的随机操作。原始前向和 checkpoint 重算会拿到同一个 seed,但不同训练 step 之间的 seed 仍会正常变化。
第二,融合 kernel 使用这个 seed 重新初始化专用 generator。在同一 FlashAttention 构建、同一 kernel 路径,以及输入、设备、形状和配置相同的前提下,同一个 seed 会产生同一份 dropout 结果。
工程上需要注意,这层封装要建立真正的 custom op 边界,普通 Python 包装在这条编译路径里仍可能被继续追踪。Backward 也必须使用 seeded forward 返回的那份 rng_state,否则只是固定了入口,前后仍可能对不上。
测试里把 seed 放在 CPU,是因为 kernel 最终只需要一个整数,也可以避免为这个标量增加设备到主机的同步。使用 CUDA Graph 或不同设备编排时,需要重新评估这一选择,不能直接照搬。