PyTorch自定义模型分布式训练snapshot logic及SageMaker Spot实例恢复疑问
PyTorch分布式训练快照实现与SageMaker Spot实例恢复逻辑
一、实现PyTorch分布式训练的快照逻辑
分布式训练里没必要让每个节点都生成快照——经过参数同步后,所有节点的模型参数、优化器状态是完全一致的,只需要让主节点(rank=0)负责快照的保存与加载即可,既避免多节点写入冲突,也能保证快照的一致性。
具体实现要点:
- 仅主节点执行保存操作:通过
torch.distributed.get_rank()判断当前节点是否为主节点,仅主节点将快照写入共享存储(SageMaker中推荐用/opt/ml/checkpoint路径,该路径挂载了共享存储,所有节点均可访问)。 - 快照需包含核心状态:除模型参数外,还要保存优化器状态、当前训练的epoch/全局步数,这些是恢复训练的关键信息。
- 示例代码:
import os import torch import torch.distributed as dist def save_checkpoint(model, optimizer, epoch, global_step, save_dir="/opt/ml/checkpoint"): # 仅主节点执行保存 if dist.get_rank() == 0: os.makedirs(save_dir, exist_ok=True) checkpoint_path = f"{save_dir}/checkpoint_step_{global_step}.pt" checkpoint = { "model_state_dict": model.state_dict(), "optimizer_state_dict": optimizer.state_dict(), "epoch": epoch, "global_step": global_step } torch.save(checkpoint, checkpoint_path) def load_checkpoint(model, optimizer, load_dir="/opt/ml/checkpoint"): # 遍历目录找到最新的快照(按全局步数排序) if not os.path.exists(load_dir): return 0, 0 # 无快照,从头开始 checkpoint_files = [f for f in os.listdir(load_dir) if f.startswith("checkpoint_step_")] if not checkpoint_files: return 0, 0 # 按步数降序排序,取最新的快照 checkpoint_files.sort(key=lambda x: int(x.split("_")[-1].split(".")[0]), reverse=True) latest_checkpoint = os.path.join(load_dir, checkpoint_files[0]) # 所有节点加载同一个快照 checkpoint = torch.load(latest_checkpoint, map_location="cpu") model.load_state_dict(checkpoint["model_state_dict"]) optimizer.load_state_dict(checkpoint["optimizer_state_dict"]) return checkpoint["epoch"], checkpoint["global_step"]
二、SageMaker Spot实例的快照选择逻辑
首先明确:SageMaker不会自动帮你选择快照,所有快照选择逻辑需要你在训练脚本中实现,核心前提是必须把快照保存到共享存储(而非实例本地存储——Spot实例终止后本地数据会被彻底清除)。
针对你提到的“4台实例生成4个快照”的情况:
- 这种情况本身是不合理的,应该通过上述主节点单独保存的逻辑避免。如果真的出现多节点快照,你需要在训练启动时,通过脚本遍历共享存储中的快照文件,根据文件名中的epoch/步数、文件修改时间等维度,筛选出最新的有效快照进行加载。
- 当Spot实例被终止后,SageMaker会重新启动训练任务,新的实例组会挂载同一个共享存储。你的训练脚本需要在初始化阶段检查共享存储中是否存在快照,若存在则加载最新的快照继续训练,否则从头开始。
内容的提问来源于stack exchange,提问作者juvchan
相关产品推荐
相关产品推荐

