PyTorch 的分布式能力主要依赖:
torch.distributedDistributedDataParallel (DDP)NCCL / Gloo / MPI 通信后端核心问题:
✅ 结论:
PyTorch 分布式本身不提供容错,容错必须在“训练框架 + 调度 + 状态恢复”层面实现
周期性保存状态 + 外部重启 + 重初始化进程组
训练中断
↓
调度系统重启任务
↓
重新初始化 torch.distributed
↓
从 checkpoint 恢复模型 & 优化器
↓
继续训练def save_checkpoint(epoch, model, optimizer, scheduler):
state = {
"epoch": epoch,
"model": model.state_dict(),
"optimizer": optimizer.state_dict(),
"scheduler": scheduler.state_dict(),
}
torch.save(state, f"ckpt_epoch_{epoch}.pt")start_epoch = 0
if resume_path:
ckpt = torch.load(resume_path, map_location="cpu")
model.load_state_dict(ckpt["model"])
optimizer.load_state_dict(ckpt["optimizer"])
start_epoch = ckpt["epoch"]✅ 优点:
❌ 缺点:
PyTorch 官方容错 & 弹性训练方案
torch.distributedtorchrun \
--nnodes=1:2 \
--nproc_per_node=8 \
train.py| 特性 | 说明 |
|---|---|
| 容错 | worker 挂掉可重启 |
| 弹性 | 可增加 / 减少节点 |
| 协调 | 自动 re-rendezvous |
torchrun / elasticK8s Job
├─ Worker-0
├─ Worker-1
└─ Worker-NPyTorchJob✅ 适合:
FSDP 本身不支持自动容错,但:
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
from torch.distributed.fsdp import ShardedOptimizer
# 保存 sharded state
FSDP.full_state_dict(model)✅ 建议:
def main():
dist.init_process_group("nccl")
model = ...
model = DDP(model, device_ids=[local_rank])
optimizer = ...
start_epoch = 0
if resume:
load_checkpoint()
for epoch in range(start_epoch, max_epochs):
train_one_epoch()
save_checkpoint(epoch)
dist.destroy_process_group()配合:
torchrun --max_restarts=3 train.pyPyTorch 分布式本身不提供容错,容错 = TorchElastic + Checkpoint + 外部调度系统
| 场景 | 推荐方案 |
|---|---|
| 单机多卡 | DDP + torchrun |
| 多机训练 | TorchElastic + Slurm |
| 云上训练 | K8s + PyTorchJob |
| 大模型 | FSDP + TorchElastic |
| 工业训练 | Checkpoint + Restart |
如果你愿意,我可以:
你现在是 单机、多机,还是云上训练?