资讯详情

PyTorch 分布式训练容错实战:使用 torchrun 实现故障恢复与弹性训练

📅 2026/9/25 13:01:11 | 华诺云谱 👁 阅读
PyTorch 分布式训练容错实战:使用 torchrun 实现故障恢复与弹性训练
示例工程【免费下载链接】tutorialsPyTorch tutorials.项目地址https://gitcode.com/gh_mirrors/tuto/tutorials点击查看免费下载导读在分布式训练中任何单进程故障都可能中断整个训练任务而分布式环境恰恰是故障高发场景因此让训练脚本具备容错能力至关重要。本文聚焦 PyTorch 官方教程系列中的容错篇beginner_source/ddp_series_fault_tolerance.rst系统讲解如何借助torchrun启动多 GPU 训练任务、通过保存/加载训练快照snapshot实现优雅重启以及如何将训练脚本改造成即使节点动态加入/离开也能自动续跑的弹性结构。读完本文你将掌握用torchrun替代mp.spawn的标准启动方式、torchrun自动注入的环境变量、以及一套可直接复制的“快照保存—断点恢复”训练脚本骨架。本文属于 PyTorch DDP 视频教程系列的第 3 节系列其余部分包括 DDP 原理简介、单机多 GPU 训练、多节点训练 与 minGPT 实战建议按顺序阅读。为什么分布式训练需要容错在分布式训练中单个进程的失败就足以中断整个训练任务。与单机训练相比分布式场景的故障敏感性更高原因在于参与训练的进程分布在多个 GPU 甚至多台机器上任何一个节点宕机、网络闪断、显存耗尽或进程被杀都会导致整个 job 挂起训练通常持续数小时乃至数天从头重来代价高昂在实际生产环境中你可能还希望训练任务是弹性的elastic即计算资源可以在任务进行中动态地加入或离开。PyTorch 为此提供了名为torchrun的工具它带来两大核心能力容错Fault Tolerance当故障发生时torchrun会记录错误日志并尝试从训练任务最后一次保存的“快照”自动重启所有进程弹性训练Elastic Training当成员发生变化节点加入或移除时torchrun会在可用设备上终止并重新拉起进程让训练在无需人工干预的情况下继续。这里的“快照”保存的远不止模型权重。它还可以包含已经运行的 epoch 数、优化器状态以及任何训练任务为了连续性所必需的“有状态”属性stateful attribute。为什么用 torchrun它替你处理分布式训练的琐碎细节torchrun的目标是让你不必亲自处理分布式训练的“细枝末节”典型的收益包括无需手动设置环境变量你不必显式传递rank和world_sizetorchrun会自动分配同时还会注入LOCAL_RANK、MASTER_ADDR、MASTER_PORT等一系列环境变量无需在脚本中调用mp.spawn你只需要提供一个通用的main()入口然后用torchrun启动脚本即可。这样一来同一份脚本可以不加改动地运行在非分布式、单机多 GPU 和多机多 GPU 三种场景优雅重启故障后自动从最近一次保存的训练快照恢复训练。在仓库的 beginner_source/dist_overview.rst 中torchrun也被描述为“广泛使用的启动脚本launcher script它负责在本地和远程机器上为分布式 PyTorch 程序拉起进程”并且当训练跨多个节点时官方推荐使用torchrun启动多个 PyTorch 进程。优雅重启训练脚本的标准结构为了让torchrun能在故障后优雅恢复训练脚本应该按照“先加载快照、再初始化、后训练”的结构组织def main(): load_snapshot(snapshot_path) # 1. 从最近快照恢复状态 initialize() # 2. 初始化模型、优化器、数据加载器 train() # 3. 继续训练 def train(): for batch in iter(dataset): train_step(batch) if should_checkpoint: # 定期保存快照 save_snapshot(snapshot_path)这段结构背后的运行机制是一旦故障发生torchrun会终止所有进程并全部重启每个进程的入口点会先加载并初始化最近一次保存的快照然后从那里继续训练因此在任何一次故障中你只会丢失自上次快照以来的训练进度在弹性训练场景下每当成员发生变化节点加入或移除torchrun都会在可用设备上终止并重新拉起进程。有了上述脚本结构训练任务就可以无需人工干预地自动继续。从multigpu.py到multigpu_torchrun.py逐项改造本节对照本系列上一节中的multigpu.py见 beginner_source/ddp_series_multigpu.rst逐项展示为了启用torchrun需要做出的代码改动。1. 进程组初始化不再手工指定 rank 与 world_size原先用mp.spawn方式启动时需要手动设置MASTER_ADDR、MASTER_PORT并显式传入rank和world_size- def ddp_setup(rank, world_size): def ddp_setup(): - - Args: - rank: Unique identifier of each process - world_size: Total number of processes - - os.environ[MASTER_ADDR] localhost - os.environ[MASTER_PORT] 12355 - init_process_group(backendnccl, rankrank, world_sizeworld_size) init_process_group(backendnccl) torch.cuda.set_device(int(os.environ[LOCAL_RANK]))改动要点torchrun会自动为每个进程注入RANK和WORLD_SIZE此外还有LOCAL_RANK等环境变量因此init_process_group(backendnccl)无需任何显式参数即可正确初始化分布式进程组进程使用的 GPU 通过int(os.environ[LOCAL_RANK])获取LOCAL_RANK标识节点上的本地进程序号用于torch.cuda.set_device()把每个进程绑定到对应的 GPU。2. 使用 torchrun 提供的环境变量原先代码把gpu_id作为构造参数传入Trainer改造后直接从环境变量读取- self.gpu_id gpu_id self.gpu_id int(os.environ[LOCAL_RANK])gpu_id的用途包括把模型放在正确的 GPU 上以及控制“只在 rank 0 进程上保存检查点”。在多节点场景中LOCAL_RANK是节点内唯一的而RANK才是跨节点全局唯一的进程标识详见 intermediate_source/ddp_series_multinode.rst。3. 保存与加载快照定期把训练所需的全部相关信息写入快照是训练任务中断后无缝续跑的前提 def _save_snapshot(self, epoch): snapshot {} snapshot[MODEL_STATE] self.model.module.state_dict() snapshot[EPOCHS_RUN] epoch torch.save(snapshot, snapshot.pt) print(fEpoch {epoch} | Training snapshot saved at snapshot.pt) def _load_snapshot(self, snapshot_path): snapshot torch.load(snapshot_path) self.model.load_state_dict(snapshot[MODEL_STATE]) self.epochs_run snapshot[EPOCHS_RUN] print(fResuming training from snapshot at Epoch {self.epochs_run})要点说明快照是一个普通字典这里保存了MODEL_STATE模型权重和EPOCHS_RUN已完成的 epoch 数。快照内容完全可以扩展加入OPTIMIZER_STATE、学习率调度器状态、随机数生成器状态、当前 step 计数等凡是恢复训练所需的状态都应入库注意使用self.model.module.state_dict()当模型被DistributedDataParallel包装后实际权重保存在.module子模块中这一约定与本系列上一节的保存检查点方式一致self.epochs_run是Trainer的实例属性需要在__init__中初始化为 0加载快照后由快照覆盖。4. 在 Trainer 构造函数中加载快照当重启一个被中断的训练任务时脚本会首先尝试加载快照以便从断点续跑class Trainer: def __init__(self, snapshot_path, ...): ... if os.path.exists(snapshot_path): self._load_snapshot(snapshot_path) ...这段逻辑位于构造函数中保证任何一次进程拉起无论是首次启动还是故障后的重启都会自动检查并恢复快照。5. 恢复训练从上次 epoch 继续训练循环不再从 0 开始而是从快照记录的epochs_run处继续def train(self, max_epochs: int): - for epoch in range(max_epochs): for epoch in range(self.epochs_run, max_epochs): self._run_epoch(epoch)需要留意在使用DistributedSampler的情况下_run_epoch内部应当在每个 epoch 开始前调用self.train_data.sampler.set_epoch(epoch)以保证跨 epoch 的 shuffle 正常工作详见 beginner_source/ddp_series_multigpu.rst恢复训练时该机制同样适用不会因为中途重启而破坏数据划分的一致性。6. 运行脚本从mp.spawn切换到torchrun入口函数部分去掉mp.spawn与手工计算的world_sizeif __name__ __main__: import sys total_epochs int(sys.argv[1]) save_every int(sys.argv[2]) - world_size torch.cuda.device_count() - mp.spawn(main, args(world_size, total_epochs, save_every,), nprocsworld_size) main(save_every, total_epochs)启动命令相应地从- python multigpu.py 50 10 torchrun --standalone --nproc_per_node4 multigpu_torchrun.py 50 10参数含义--standalone单机模式torchrun自行管理本机上的所有进程不依赖外部 rendezvous 服务--nproc_per_node4本机每节点拉起 4 个进程即使用 4 块 GPU 训练50 10位置参数分别传入total_epochs总训练轮数和save_every每多少轮保存一次快照。在 intermediate_source/FSDP_advanced_tutorial.rst 中也可以看到同类用法torchrun --nnodes 1 --nproc_per_node 4 T5_training.py即通过--nnodes指定节点数、--nproc_per_node指定每节点进程数。多节点场景的完整配置每台机器运行相同的 rendezvous 参数、或借助 SLURM 等作业调度器可参考 intermediate_source/ddp_series_multinode.rst。快照设计的扩展实践原教程用snapshot.pt演示了最小可用快照实际生产中通常还需要快照字段建议内容说明MODEL_STATEmodel.module.state_dict()DDP 包装后须取.module的权重OPTIMIZER_STATEoptimizer.state_dict()恢复动量等优化器内部状态续跑时学习率曲线不跳变EPOCHS_RUN/STEPS_RUN整数计数决定训练循环从何处继续SCHEDULER_STATEscheduler.state_dict()恢复学习率调度位置其他自定义元数据如验证指标、随机种子按需扩展在容错训练之外本系列最后一篇 intermediate_source/ddp_series_minGPT.rst 还展示了把训练快照直接保存到云端存储的实践这样可以从集群中任何一个能访问云存储桶的节点继续训练进一步提升了训练任务的可迁移性与灵活性。使用前提与限制本文示例面向单机多 GPU场景需要一台配备多块 CUDA GPU 的机器原教程使用 AWS p3.8xlarge 实例含 4 块 GPU本地需安装支持 CUDA 的 PyTorch多节点场景需要在所有节点上安装 PyTorch 并保证节点间 TCP 互通同时注意多节点训练受节点间通信延迟限制——同一节点上 4 块 GPU 训练通常快于 4 台机器各用 1 块 GPU 训练详见 intermediate_source/ddp_series_multinode.rst容错能力依赖定期保存快照故障后最多丢失“自上次快照以来的进度”快照保存越频繁损失越小但保存本身也有 I/O 开销save_every需要根据实际训练速度权衡torchrun重启进程后不保证进程继续持有相同的LOCAL_RANK与RANK因此不要用RANK承担关键逻辑如数据划分决策——这是官方在 intermediate_source/ddp_series_multinode.rst 中明确给出的警告。进一步阅读本系列上一节Multi-GPU Training with DDP单机多 GPU 训练本系列下一节Multi-Node Training with DDP多节点训练本系列完结篇Training a GPT model with DDPminGPT 实战分布式训练概览Distributed Training Overview其中torchrun被列为官方推荐的分布式程序启动器赞分享示例工程【免费下载链接】tutorialsPyTorch tutorials.项目地址https://gitcode.com/gh_mirrors/tuto/tutorials点击查看免费下载相关推荐Awesome Agriculture中的机器学习与AI从作物预测到病虫害识别的应用案例Awesome Agriculture中的机器学习与AI从作物预测到病虫害识别的应用案例 Awesome Agriculture是一个专注于农业、 farmi文档知识库SkyPilot 弹性 Ray 分布式训练实战用 Job Group Dynamic Node Set 实现节点故障自动恢复SkyPilot 弹性 Ray 分布式训练实战用 Job Group Dynamic Node Set 实现节点故障自动恢复 导读 本篇文章基于 SkyP后端任务调度MLOps集群管理Lapce终极指南3步快速上手Rust开发的闪电级代码编辑器Lapce终极指南3步快速上手Rust开发的闪电级代码编辑器 还在为代码编辑器启动慢、占用内存多而烦恼吗Lapce发音/læps/是一款用纯Rust语代码编辑器桌面应用开发工具上一篇Windows 通知栏背单词ToastFish 三个真实场景完整指南下一篇Deepin升级工具高级配置自定义更新策略与通知频率设置创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
📝

华诺云谱内容团队

资深建站顾问 · 行业研究员

10年+企业数字化服务经验,专注智能建站、SEO优化与品牌营销,持续输出建站技巧、行业洞察与营销干货,已帮助5000+企业实现数字化增长。

你可能需要的服务

订阅华诺云谱资讯周报

每周一封,精选建站技巧、SEO与营销干货,直达邮箱。已有 8,000+ 企业主订阅,助你少走弯路。

↑