ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

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

Accelerate Checkpointing 指南:用 save_state / load_state 实现模型训练状态的完整保存与恢复 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_state与skip_first_batches无缝恢复训练以及如何通过ProjectConfiguration和register_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_state与load_accelerator_state而Accelerator上的公开 API 定义在 src/accelerate/accelerator.pysave_state与 src/accelerate/accelerator.pyload_state。最小可用示例保存与恢复一个完整训练状态官方文档给出的最小示例完整呈现了注册调度器 → 保存初始状态 → 训练 → 恢复状态的全流程以下是继承原文并补充注释的版本from accelerate import Accelerator import torch accelerator Accelerator(project_dirmy/save/path) my_scheduler torch.optim.lr_scheduler.StepLR(my_optimizer, step_size1, gamma0.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 时有两条重要约束源码注释中反复强调保存与加载必须出自同一训练脚本。save_state/load_state是为训练中途断电续跑设计的其内部对模型、优化器的保存方式与prepare后的包装状态强相关并不保证跨脚本通用详见 src/accelerate/accelerator.py 中的 Tip 说明。如果目标是训练完成后导出模型权重供推理使用应使用accelerator.save_modelsrc/accelerate/accelerator.py支持max_shard_size分片与 safetensors 序列化。不注册的对象不会被保存。如果某个文件没有通过register_for_checkpointing注册即使它存在于保存目录中load_state也不会加载它。一次 save_state 到底写入了哪些文件save_accelerator_statesrc/accelerate/checkpointing.py会把状态组织成一组命名文件文件名的常量定义在 src/accelerate/utils/constants.py文件常量内容model.safetensors默认或pytorch_model.binsafe_serializationFalse时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.bin、optimizer_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_dirmy/save/path, automatic_checkpoint_namingTrue, total_limit3, # 只保留最近 3 个检查点 ) accelerator Accelerator(project_configurationproject_config) # 无需传目录自动保存到 my/save/path/checkpoints/checkpoint_0、checkpoint_1、... accelerator.save_state()源码层面的行为src/accelerate/accelerator.pysave_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_dict和load_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的数量是否与当前注册对象数量一致不一致会抛出RuntimeErrorsrc/accelerate/accelerator.py提示检查点必须由同一组注册对象产生或在目录中避免使用custom_checkpoint命名冲突文件。分布式后端下的保存差异FSDP / DeepSpeed / Megatron-LMsave_state内部会根据self.distributed_type分派不同的保存策略src/accelerate/accelerator.pyFSDP调用save_fsdp_model/save_fsdp_optimizersrc/accelerate/utils/fsdp_utils.py按 FSDP 分片规则保存模型与优化器DeepSpeed调用 DeepSpeed 引擎的save_checkpoint(output_dir, ckpt_id)模型与优化器状态由 DeepSpeed 统一管理ckpt_id为pytorch_model或pytorch_model_{i}加载时对应调用load_checkpointMegatron-LM委托模型的save_checkpoint/load_checkpoint处理其他场景DDP / 单卡等走通用路径用get_state_dict取出权重后交给save_accelerator_state。get_state_dictsrc/accelerate/accelerator.py同样按后端区分DeepSpeed ZeRO-3 场景要求配置stage3_gather_16bit_weights_on_model_saveTrue否则抛错并提示改用zero_to_fp32.py恢复权重FSDP 通过FULL_STATE_DICToffload_to_cpuTrue、rank0_onlyTrue只在 rank 0 汇总完整状态FSDP2 则基于torch.distributed.checkpoint的get_model_state_dict实现存在 offload 参数时走get_state_dict_offloaded_model。此外load_state对map_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_batchessave_state保存了采样器状态但如果你是在一个 epoch 的中间保存检查点恢复后 DataLoader 会从头开始取数导致同一批数据被重复训练。官方文档给出的解法是Accelerator.skip_first_batches它返回一个高效跳过前num_batches个 batch 的新 DataLoaderfrom accelerate import Accelerator accelerator Accelerator(project_dirmy/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会保留device、rng_types、iteration等属性重建一个新 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这段代码展示了三个关键工程细节可以平滑移植到任何训练脚本目录即元数据用epoch_{i}/step_{i}命名检查点目录恢复时从路径字符串解析出应从第几个 epoch、第几个 step 继续无需额外维护状态文件频率可控--checkpointing_steps可设为整数每 N 步存一次或epoch每个 epoch 存一次状态校验恢复后可立即评估模型性能并与断点前记录的值比对确认恢复无误。如何验证断点续训的正确性仓库提供了配套测试来保证该功能的可靠性src/accelerate/test_utils/scripts/external_deps/test_checkpointing.py 是一个带--resume_from_checkpoint参数的端到端脚本训练流程为每训练完一个 epoch 调用accelerator.save_state(output_dir)目录名为epoch_{epoch}同时把该 epoch 的准确率、调度器学习率、优化器学习率、epoch 号、总 step 数写入state_{epoch}.json断点续训时先从检查点目录解析出starting_epoch调用accelerator.load_state恢复恢复后立即跑一次评估然后断言恢复后评估的准确率 断点前记录的准确率、调度器/优化器学习率一致、epoch 号一致全部断言通过才认为加载成功src/accelerate/test_utils/scripts/external_deps/test_checkpointing.py。这套保存 → 重新启动 → 恢复 → 数值比对的验证思路同样适用于你自己的训练脚本恢复点上的任何数值loss、准确率、学习率、优化器状态都应该与中断前完全一致这是判断检查点系统是否正确的黄金标准。小结与最佳实践训练中断续跑用save_state/load_state并配合register_for_checkpointing注册调度器等对象二者只应在同一训练脚本内配对使用自动管理检查点通过ProjectConfiguration(automatic_checkpoint_namingTrue, total_limitN)自动命名并淘汰旧检查点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),仅供参考
返回列表