You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.16 21:25:30