PyTorch 分布式训练卡住不报错:用 monitored_barrier 定位失联 rank
分布式训练最耗时间的故障,往往不是明确报错,而是所有 GPU 利用率逐渐归零,进程既不退出,也没有新的日志。表面看像 NCCL 卡死,实际根因可能发生得更早:某个 rank 在 DataLoader 抛了异常、进入了不同的条件分支,或者保存 checkpoint 时停在文件系统上。其余 rank 直到下一次集合通信才开始等待,最后看到的堆栈离第一现场已经很远。
排查这类问题,第一步是让日志具备 rank 维度。每条关键日志至少带全局 rank、本地 rank、主机名、训练 step 和即将进入的集合操作。不要只让 rank 0 输出,因为失联的恰恰可能是非零 rank。日志写入独立文件更容易保留时间线,例如rank-0.log、rank-1.log,同时确保异常处理会刷新缓冲区。先找到各文件最后一个共同 step,再看哪个 rank 最早偏离。
第二步是在可疑阶段插入torch.distributed.monitored_barrier。它与普通barrier的区别,是能够在超时后由 rank 0 报告哪些 rank 没有按时确认。官方文档说明,这个同步过程通过主机侧点对点通信实现,需要 Gloo 进程组。若训练主后端是 NCCL,可以额外创建一个 Gloo group 专门做诊断,不要误以为直接在 NCCL 默认组上调用就会得到同样效果。
可复现的最小写法是先初始化主进程组,再创建诊断组:gloo_group = dist.new_group(backend="gloo")。在数据加载完成、前向结束、反向结束和优化器 step 之后分别放置带超时的 monitored barrier,例如dist.monitored_barrier(group=gloo_group, timeout=timedelta(seconds=30))。每个检查点前后打印阶段名。第一次加入时不要到处埋点,先用二分法把一个 step 分成前后两段,确定卡住区间后再细分,否则同步点会明显改变训练时序。
第三步是构造故障验证诊断是否可信。可以只让指定 rank 在某个 step 睡眠超过超时,观察 rank 0 是否准确报告;也可以让该 rank 在进入 barrier 前抛出受控异常,确认其他进程最终退出而不是永久等待。测试应在与生产相同的启动器下执行,因为torchrun、容器编排和作业平台对子进程退出的处理并不完全相同。完成验证后再去查真实问题,能避免把监控工具自身的配置错误当成训练故障。
如果失联发生在集合通信内部,还要补充 NCCL 的证据。开启 PyTorch 官方文档建议的分布式调试日志,记录进程组初始化参数、网卡选择和各 rank 的调用顺序。重点检查所有 rank 是否以相同顺序、相同张量形状进入 all-reduce 或 all-gather。某个条件分支只在 rank 0 执行一次集合操作,就足以让其他 rank 永久等待。动态 batch、最后一个不完整 batch、梯度累积条件和异常样本过滤,都是调用序列分叉的常见来源。
monitored_barrier也有边界。它是诊断工具,不该密集留在高频训练路径中;主机侧同步会增加开销,并可能掩盖竞态。它能指出谁没有到达,却不会自动说明该 rank 卡在数据、计算、网络还是存储。超时值也不能随便设成几秒,首次编译、数据预热和 checkpoint 本来就可能很慢。应根据阶段正常耗时设置阈值,并在日志中区分诊断超时与训练业务超时。
推荐的排查顺序可以固定下来:先按 rank 重建最后进度,再用 Gloo 的 monitored barrier 缩小阶段,接着复现一个受控失联,最后结合 NCCL 日志、DataLoader worker 日志和系统指标寻找根因。这样做的价值,不是给“卡死”换一个更漂亮的错误,而是把“全体都在等”还原成“哪一个 rank 从哪一步开始没有到达”。证据链一旦建立,分布式问题才从猜测变成可以复现的工程故障。
参考资料
- https://docs.pytorch.org/docs/stable/distributed.html
