Accelerate Checkpointing 指南:用 save_state / load_state 实现模型训练状态的完整保存与恢复
2026/9/24 13:48:33 网站建设 项目流程

Accelerate Checkpointing 指南:用 save_state / load_state 实现模型训练状态的完整保存与恢复

【免费下载链接】accelerate🚀 A simple way to launch, train, and use PyTorch models on almost any device and distributed configuration, automatic mixed precision (including fp8), and easy-to-configure FSDP and DeepSpeed support项目地址: https://gitcode.com/gh_mirrors/ac/accelerate

在使用 Accelerate 训练 PyTorch 模型时,断点续训(Checkpointing)是工程落地的必备能力:训练中断后能精确恢复到之前的训练进度,而不是从头重来。本篇指南以 Accelerate 官方文档(docs/source/usage_guides/checkpoint.md)为核心骨架,结合当前仓库中 src/accelerate/checkpointing.py 与 src/accelerate/accelerator.py 的源码实现,讲解如何用Accelerator.save_state一键保存模型、优化器、RNG 随机数生成器、GradScaler 等全部训练状态,如何在断点后通过load_stateskip_first_batches无缝恢复训练,以及如何通过ProjectConfigurationregister_for_checkpointing定制保存策略。读完本篇,你将掌握一套可直接复制到训练脚本中的完整断点续训方案。

为什么要用 Accelerate 做 Checkpointing

在分布式训练中,直接调用torch.save(model.state_dict())会遇到一系列问题:

  • 模型经过accelerator.prepare()后可能被DistributedDataParallel、FSDP、DeepSpeed 等包装,直接保存的是包装层而非原始模型;
  • 训练状态远不止模型权重,还包括优化器动量、学习率调度器进度、随机数生成器状态、混合精度下的 GradScaler 缩放因子;
  • 分布式环境下,每个进程都要保存状态,何时同步、由谁落盘、如何避免进程间互相覆盖,都需要约定。

Accelerate 提供两个便捷函数一次性解决上述问题:

  • Accelerator.save_state:把模型、优化器、RNG 生成器、GradScaler(如启用混合精度)以及注册过的自定义对象保存到指定文件夹;
  • Accelerator.load_state:从save_state生成的文件夹中恢复全部状态。

这两者的实现位于 src/accelerate/checkpointing.py 的save_accelerator_stateload_accelerator_state,而Accelerator上的公开 API 定义在 src/accelerate/accelerator.py(save_state)与 src/accelerate/accelerator.py(load_state)。

最小可用示例:保存与恢复一个完整训练状态

官方文档给出的最小示例完整呈现了"注册调度器 → 保存初始状态 → 训练 → 恢复状态"的全流程,以下是继承原文并补充注释的版本:

from accelerate import Accelerator import torch accelerator = Accelerator(project_dir="my/save/path") my_scheduler = torch.optim.lr_scheduler.StepLR(my_optimizer, step_size=1, gamma=0.99) my_model, my_optimizer, my_training_dataloader = accelerator.prepare( my_model, my_optimizer, my_training_dataloader ) # 注册 LR scheduler:它不具备模型/优化器/采样器那样的"默认待遇",需要显式注册 accelerator.register_for_checkpointing(my_scheduler) # 保存初始状态 accelerator.save_state() device = accelerator.device my_model.to(device) # 执行训练 for epoch in range(num_epochs): for batch in my_training_dataloader: my_optimizer.zero_grad() inputs, targets = batch inputs = inputs.to(device) targets = targets.to(device) outputs = my_model(inputs) loss = my_loss_function(outputs, targets) accelerator.backward(loss) my_optimizer.step() my_scheduler.step() # 从之前的保存点恢复状态 accelerator.load_state("my/save/path/checkpointing/checkpoint_0")

使用这个 API 时有两条重要约束,源码注释中反复强调:

  1. 保存与加载必须出自同一训练脚本save_state/load_state是为"训练中途断电续跑"设计的,其内部对模型、优化器的保存方式与prepare后的包装状态强相关,并不保证跨脚本通用(详见 src/accelerate/accelerator.py 中的 Tip 说明)。如果目标是"训练完成后导出模型权重供推理使用",应使用accelerator.save_model(src/accelerate/accelerator.py,支持max_shard_size分片与 safetensors 序列化)。
  2. 不注册的对象不会被保存。如果某个文件没有通过register_for_checkpointing注册,即使它存在于保存目录中,load_state也不会加载它。

一次 save_state 到底写入了哪些文件

save_accelerator_state(src/accelerate/checkpointing.py)会把状态组织成一组命名文件,文件名的常量定义在 src/accelerate/utils/constants.py:

文件常量内容
model.safetensors(默认)或pytorch_model.binsafe_serialization=False时)SAFE_MODEL_NAME/MODEL_NAME模型权重,默认使用 safetensors 序列化
optimizer.binOPTIMIZER_NAME优化器state_dict()
scheduler.binSCHEDULER_NAME已注册调度器的state_dict()
sampler.binSAMPLER_NAMEDataLoader 采样器状态(仅当使用 Accelerate 的SeedableRandomSampler等自定义采样器时)
dl_state_dict.bin使用StatefulDataLoader时的完整 DataLoader 状态
scaler.ptSCALER_NAMEGradScaler 状态(仅混合精度训练时存在)
random_states_{process_index}.pklRNG_STATE_NAME当前进程的 Pythonrandom、NumPy、torch(含 CUDA/XPU/MLU/HPU 等各后端)RNG 状态,以及内部step计数
custom_checkpoint_{index}.pkl通过register_for_checkpointing注册的自定义对象

其中值得注意的实现细节:

  • 多对象场景下(多个模型/优化器/调度器),文件名会追加序号,如pytorch_model_1.binoptimizer_1.bin
  • 加载时(src/accelerate/checkpointing.py)会优先寻找model.safetensors,不存在则回退到pytorch_model.bin,因此两种序列化格式的检查点可以混用;
  • RNG 状态文件按process_index命名,保证每个分布式进程恢复各自的随机状态,从而让数据采样、dropout 等行为与中断前完全一致;
  • 加载 RNG 时若发现文件中带有step字段,会把该值写入override_attributes,最终由Accelerator.load_state恢复内部self.step计数(src/accelerate/accelerator.py)。

用 ProjectConfiguration 定制保存位置与自动命名

默认情况下,save_state(output_dir)要求你手动传入目录。若希望训练过程中自动迭代命名检查点,可通过ProjectConfiguration定制(对应文档原句:if automatic_checkpoint_naming启用后,每个检查点保存在Accelerator.project_dir/checkpoints/checkpoint_{checkpoint_number})。该配置类定义在 src/accelerate/utils/dataclasses.py,支持以下字段:

字段默认值说明
project_dirNone项目根目录,保存的数据存放于此
logging_dirNone本地日志目录,缺省时与project_dir相同
automatic_checkpoint_namingFalse是否启用检查点自动命名
total_limitNone最多保留的检查点数量,超出时删除最旧者
iteration0当前保存迭代计数(对应检查点编号)
save_on_each_nodeFalse多节点训练时,是否在每个节点都保存(否则仅主节点保存)

启用方式:

from accelerate import Accelerator from accelerate.utils import ProjectConfiguration project_config = ProjectConfiguration( project_dir="my/save/path", automatic_checkpoint_naming=True, total_limit=3, # 只保留最近 3 个检查点 ) accelerator = Accelerator(project_configuration=project_config) # 无需传目录,自动保存到 my/save/path/checkpoints/checkpoint_0、checkpoint_1、... accelerator.save_state()

源码层面的行为(src/accelerate/accelerator.py):

  • save_state内部检查project_configuration.automatic_checkpoint_naming,为真时把output_dir强制指向project_dir/checkpoints
  • total_limit有值且当前检查点数量超限时,由主进程按编号排序后删除最旧的检查点,日志会提示删除了多少个以腾出空间;
  • 每次保存后project_configuration.iteration += 1,检查点目录以checkpoint_{iteration}命名;
  • 若目标目录已存在同名检查点会抛出ValueError,提示手动调整save_iteration以跳过已存在的编号;
  • 启用自动命名后,load_state()可以不传input_dir,它会自动挑选编号最大的(即最新的)检查点加载(src/accelerate/accelerator.py)。

register_for_checkpointing:让调度器等自定义对象一并保存

模型、优化器、DataLoader 采样器会被 Accelerate 自动纳入检查点,但学习率调度器等对象不会。官方文档给出的解法是Accelerator.register_for_checkpointing只要对象同时具备state_dictload_state_dict方法,就能注册并随save_state/load_state自动存取。

其实现(src/accelerate/accelerator.py)会逐一校验传入对象是否同时具备两个方法,否则抛出ValueError并列出非法对象;通过校验后对象被追加到self._custom_objects。保存时调用save_custom_state写入custom_checkpoint_{index}.pkl,加载时调用load_custom_state读取,详见 src/accelerate/checkpointing.py。

accelerator.register_for_checkpointing(my_scheduler, my_custom_tracker) # 之后 save_state / load_state 会自动带上这两个对象

需要留意的是,load_state会严格校验目录中custom_checkpoint_*.pkl的数量是否与当前注册对象数量一致,不一致会抛出RuntimeError(src/accelerate/accelerator.py),提示检查点必须由"同一组注册对象"产生,或在目录中避免使用custom_checkpoint命名冲突文件。

分布式后端下的保存差异:FSDP / DeepSpeed / Megatron-LM

save_state内部会根据self.distributed_type分派不同的保存策略(src/accelerate/accelerator.py):

  • FSDP:调用save_fsdp_model/save_fsdp_optimizer(src/accelerate/utils/fsdp_utils.py),按 FSDP 分片规则保存模型与优化器;
  • DeepSpeed:调用 DeepSpeed 引擎的save_checkpoint(output_dir, ckpt_id),模型与优化器状态由 DeepSpeed 统一管理,ckpt_idpytorch_modelpytorch_model_{i};加载时对应调用load_checkpoint
  • Megatron-LM:委托模型的save_checkpoint/load_checkpoint处理;
  • 其他场景(DDP / 单卡等):走通用路径,用get_state_dict取出权重后交给save_accelerator_state

get_state_dict(src/accelerate/accelerator.py)同样按后端区分:DeepSpeed ZeRO-3 场景要求配置stage3_gather_16bit_weights_on_model_save=True(否则抛错并提示改用zero_to_fp32.py恢复权重);FSDP 通过FULL_STATE_DICT+offload_to_cpu=Truerank0_only=True只在 rank 0 汇总完整状态;FSDP2 则基于torch.distributed.checkpointget_model_state_dict实现;存在 offload 参数时走get_state_dict_offloaded_model

此外load_statemap_location有默认行为(src/accelerate/accelerator.py):多进程多设备且非MULTI_XPU时默认"on_device"(状态加载到各自设备),否则默认"cpu";你也可以通过load_model_func_kwargs里的map_location显式覆盖。load_accelerator_state会对非法值抛出TypeError,只接受None"cpu""on_device"三者。

恢复 DataLoader 进度:skip_first_batches

save_state保存了采样器状态,但如果你是在一个 epoch 的中间保存检查点,恢复后 DataLoader 会从头开始取数,导致同一批数据被重复训练。官方文档给出的解法是Accelerator.skip_first_batches,它返回一个高效跳过前num_batches个 batch 的新 DataLoader:

from accelerate import Accelerator accelerator = Accelerator(project_dir="my/save/path") train_dataloader = accelerator.prepare(train_dataloader) accelerator.load_state("my_state") # 假设检查点保存在第 100 个 step 处 skipped_dataloader = accelerator.skip_first_batches(train_dataloader, 100) # 恢复后的第一个 epoch 使用跳过版 for batch in skipped_dataloader: # 完成当前 epoch 的剩余部分 pass # 后续 epoch 回到原始 dataloader for batch in train_dataloader: pass

其底层实现在 src/accelerate/data_loader.py:非 IterableDataset 场景下,通过SkipBatchSampler包装原batch_sampler实现"从第 N 个 batch 开始取";对于DataLoaderShard/DataLoaderDispatcher会保留devicerng_typesiteration等属性重建一个新 DataLoader。官方文档特别提示,若原 DataLoader 是StatefulDataLoader(其自身支持完整状态存取),则不应使用skip_first_batches,恢复后直接沿用原 DataLoader 即可。

端到端实战:把断点续训接进真实训练脚本

仓库中的 examples/by_feature/checkpointing.py 是一个可直接运行的完整示例(GLUE MRPC 文本分类 + 断点续训),它演示了比官方文档示例更完整的工程化模式,包括两种检查点频率和目录解析逻辑:

# 训练开始时加载检查点(--resume_from_checkpoint 传入检查点目录) if args.resume_from_checkpoint: accelerator.load_state(args.resume_from_checkpoint) path = os.path.basename(args.resume_from_checkpoint) training_difference = os.path.splitext(path)[0] if "epoch" in training_difference: starting_epoch = int(training_difference.replace("epoch_", "")) + 1 resume_step = None else: resume_step = int(training_difference.replace("step_", "")) starting_epoch = resume_step // len(train_dataloader) resume_step -= starting_epoch * len(train_dataloader) for epoch in range(starting_epoch, num_epochs): # 恢复后的第一个 epoch 跳过已训练过的 step if args.resume_from_checkpoint and epoch == starting_epoch and resume_step is not None: if not args.use_stateful_dataloader: active_dataloader = accelerator.skip_first_batches(train_dataloader, resume_step) else: active_dataloader = train_dataloader overall_step += resume_step else: active_dataloader = train_dataloader ... # 按 step 或 epoch 频率保存 accelerator.save_state(output_dir) # 例如 "step_100" / "epoch_1"

这段代码展示了三个关键工程细节,可以平滑移植到任何训练脚本:

  1. 目录即元数据:用epoch_{i}/step_{i}命名检查点目录,恢复时从路径字符串解析出应从第几个 epoch、第几个 step 继续,无需额外维护状态文件;
  2. 频率可控--checkpointing_steps可设为整数(每 N 步存一次)或"epoch"(每个 epoch 存一次);
  3. 状态校验:恢复后可立即评估模型性能并与断点前记录的值比对,确认恢复无误。

如何验证断点续训的正确性

仓库提供了配套测试来保证该功能的可靠性:src/accelerate/test_utils/scripts/external_deps/test_checkpointing.py 是一个带--resume_from_checkpoint参数的端到端脚本,训练流程为:

  1. 每训练完一个 epoch 调用accelerator.save_state(output_dir)(目录名为epoch_{epoch}),同时把该 epoch 的准确率、调度器学习率、优化器学习率、epoch 号、总 step 数写入state_{epoch}.json
  2. 断点续训时,先从检查点目录解析出starting_epoch,调用accelerator.load_state恢复;
  3. 恢复后立即跑一次评估,然后断言"恢复后评估的准确率 == 断点前记录的准确率、调度器/优化器学习率一致、epoch 号一致",全部断言通过才认为加载成功(src/accelerate/test_utils/scripts/external_deps/test_checkpointing.py)。

这套"保存 → 重新启动 → 恢复 → 数值比对"的验证思路同样适用于你自己的训练脚本:恢复点上的任何数值(loss、准确率、学习率、优化器状态)都应该与中断前完全一致,这是判断检查点系统是否正确的黄金标准。

小结与最佳实践

  • 训练中断续跑:用save_state/load_state,并配合register_for_checkpointing注册调度器等对象;二者只应在同一训练脚本内配对使用;
  • 自动管理检查点:通过ProjectConfiguration(automatic_checkpoint_naming=True, total_limit=N)自动命名并淘汰旧检查点,load_state()不传参即可加载最新断点;
  • 恢复 DataLoader 进度:普通 DataLoader 用skip_first_batchesStatefulDataLoader直接加载状态;
  • 权重导出:训练完成后请使用save_model(支持 safetensors 与分片),不要复用save_state的产物做推理加载;
  • 分布式后端(FSDP / DeepSpeed / Megatron-LM)的检查点由 Accelerate 自动分派底层实现,但要注意 DeepSpeed ZeRO-3 需要开启stage3_gather_16bit_weights_on_model_save才能通过get_state_dict拿到 16bit 权重;
  • 每次恢复后做一次数值比对,让断点续训的正确性始终有据可查。

相关参考材料:官方文档原文 docs/source/usage_guides/checkpoint.md、API 参考 docs/source/package_reference/accelerator.md、核心实现 src/accelerate/checkpointing.py、配置类 src/accelerate/utils/dataclasses.py、实战示例 examples/by_feature/checkpointing.py。

【免费下载链接】accelerate🚀 A simple way to launch, train, and use PyTorch models on almost any device and distributed configuration, automatic mixed precision (including fp8), and easy-to-configure FSDP and DeepSpeed support项目地址: https://gitcode.com/gh_mirrors/ac/accelerate

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询