PyTorch分布式训练如何容错

GPU
小华
2026-09-18

在 PyTorch 分布式训练(尤其是大规模、长周期训练)中,容错(Fault Tolerance)的核心目标是:某个进程 / 节点失败后,训练能恢复而不是从头再来。下面从「失败类型 → 原生机制 → 常见方案 → 最佳实践」系统说明。

一、分布式训练常见失败类型

  1. 节点 / 进程崩溃
  • OOM
  • 硬件故障(GPU、网卡)
  • 节点被调度系统回收
  1. 网络故障
  • NCCL 通信超时
  • 网络抖动
  1. 训练逻辑错误
  • 非确定性 bug(通常不靠容错解决)
  1. 资源抢占
  • 云上 Spot / 抢占式实例

二、PyTorch 原生容错能力(有限)

1. torch.distributed 本身不自带自动恢复

  • 一个 rank 死掉 → 整个 ProcessGroup 通常失效
  • NCCL 超时会直接抛异常

2. 弹性启动器(重点)

PyTorch 提供 Elastic Launch(torchelastic)

torchrun \
--nnodes=2:4 \
--nproc_per_node=8 \
--rdzv_id=job1 \
--rdzv_backend=c10d \
--rdzv_endpoint=localhost:29500 \
train.py

特性:

  • 支持 节点数量动态变化(min:max)
  • rank 失败后 等待重新加入
  • 适合 Spot 实例

⚠️ 但:

  • 只解决进程重启
  • 不自动恢复模型状态

三、真正实现容错的关键手段

✅ 1. 定期 Checkpoint(最核心)

必须保存:

  • model state_dict
  • optimizer state_dict
  • scheduler
  • epoch / step
  • 随机种子状态(可选)
torch.save({
'model': model.state_dict(),
'optimizer': optimizer.state_dict(),
'epoch': epoch,
}, 'ckpt.pt')

✅ 建议:

  • 每 N step 或每 epoch 保存
  • 异步写(避免阻塞)
  • 写临时文件再 rename(防半截文件)

✅ 2. 从 Checkpoint 恢复训练

ckpt = torch.load('ckpt.pt')
model.load_state_dict(ckpt['model'])
optimizer.load_state_dict(ckpt['optimizer'])
start_epoch = ckpt['epoch']

⚠️ 分布式注意:

  • 每个 rank 都加载
  • 或使用 rank 0 保存、广播

✅ 3. 使用 torch.distributed.checkpoint(推荐)

PyTorch 1.13+ 提供 分布式 checkpoint

from torch.distributed.checkpoint import (
save, load, FileSystemWriter, FileSystemReader
)
save(
{"model": model},
checkpoint_id="ckpt/",
storage_writer=FileSystemWriter("ckpt/")
)

优点:

  • 支持 分片存储(sharded)
  • 大规模模型友好
  • 恢复时只加载需要的分片

✅ 4. NCCL 超时与重试配置

import datetime
torch.distributed.init_process_group(
backend="nccl",
timeout=datetime.timedelta(seconds=1800)
)

配合:

  • 监控脚本 kill 卡死进程
  • 自动 restart

✅ 5. 外部调度系统容错(生产级)

系统容错方式
KubernetesPod 重启 + PVC 存 checkpoint
Slurm--requeue
Ray Train内置容错 + 重启

✅ 存活检测:

  • 心跳
  • 日志监控
  • NCCL watchdog

四、典型容错架构(推荐)

训练进程
↓
torchrun (elastic)
↓
checkpoint 到共享存储 (NFS / S3 / HDFS)
↓
失败 → 重启 → 加载 checkpoint → 继续

五、常见坑

❌ 只存 model,不存 optimizer → 精度崩
❌ checkpoint 写在本地盘 → 节点没了就没了
❌ NCCL 超时不处理 → 整个 job 卡死

❌ 弹性训练但没 checkpoint → 重启也白搭

六、总结一句话

PyTorch 分布式容错 = Elastic Launch + 稳定 Checkpoint + 外部重启机制
如果你愿意,我可以:
  • 给你一个 完整容错训练模板
  • 或针对 FSDP / DeepSpeed / Megatron 分别讲
  • 或结合 K8s / Slurm 实战配置

你现在是哪种训练场景?

亿速云提供售前/售后服务

售前业务咨询

售后技术保障

400-100-2938

7*24小时售后电话

官方微信小程序