torch.distributed 本身不自带自动恢复ProcessGroup 通常失效PyTorch 提供 Elastic Launch(torchelastic)
torchrun \
--nnodes=2:4 \
--nproc_per_node=8 \
--rdzv_id=job1 \
--rdzv_backend=c10d \
--rdzv_endpoint=localhost:29500 \
train.py特性:
⚠️ 但:
必须保存:
torch.save({
'model': model.state_dict(),
'optimizer': optimizer.state_dict(),
'epoch': epoch,
}, 'ckpt.pt')✅ 建议:
ckpt = torch.load('ckpt.pt')
model.load_state_dict(ckpt['model'])
optimizer.load_state_dict(ckpt['optimizer'])
start_epoch = ckpt['epoch']⚠️ 分布式注意:
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/")
)优点:
import datetime
torch.distributed.init_process_group(
backend="nccl",
timeout=datetime.timedelta(seconds=1800)
)配合:
| 系统 | 容错方式 |
|---|---|
| Kubernetes | Pod 重启 + PVC 存 checkpoint |
| Slurm | --requeue |
| Ray Train | 内置容错 + 重启 |
✅ 存活检测:
训练进程
↓
torchrun (elastic)
↓
checkpoint 到共享存储 (NFS / S3 / HDFS)
↓
失败 → 重启 → 加载 checkpoint → 继续❌ 只存 model,不存 optimizer → 精度崩
❌ checkpoint 写在本地盘 → 节点没了就没了
❌ NCCL 超时不处理 → 整个 job 卡死
PyTorch 分布式容错 = Elastic Launch + 稳定 Checkpoint + 外部重启机制
如果你愿意,我可以:
你现在是哪种训练场景?