☰
DeepSpeed 模型检查点(Model Checkpointing)实战指南:保存、加载与 ZeRO fp32 权重恢复
2026/9/25 10:22:05 网站建设 项目流程
  • 推理引擎
  • 大模型

【免费下载链接】FlexGen

Running large language models on a single GPU for throughput-oriented scenarios.

项目地址:https://gitcode.com/gh_mirrors/fl/FlexGen
点击查看免费下载

本指南以 DeepSpeed 官方文档 model-checkpointing.rst 为核心,系统讲解 DeepSpeed 引擎在训练期间保存与加载模型状态的完整机制:save_checkpoint/load_checkpoint两个引擎级 API 的用法与底层实现,以及从 ZeRO 检查点中恢复 fp32 权重的三种工具函数。读完本文,你将掌握在分布式训练(含 ZeRO-2/ZeRO-3)场景下安全地保存、恢复训练现场,以及将 DeepSpeed 检查点转换为标准 PyTorchstate_dict以便脱离 DeepSpeed 使用模型的完整实战方案。

一、总览:DeepSpeed 检查点 API 家族

DeepSpeed 将"模型检查点"划分为两套相互配合的能力,对应文档中的三大主题:

能力入口 API所属模块用途
保存训练检查点DeepSpeedEngine.save_checkpointdeepspeed.runtime.engine保存模型权重、优化器状态、LR 调度器状态与用户自定义状态
加载训练检查点DeepSpeedEngine.load_checkpointdeepspeed.runtime.engine恢复上述全部状态以继续训练
ZeRO 检查点 fp32 权重恢复get_fp32_state_dict_from_zero_checkpointdeepspeed.utils.zero_to_fp32从 ZeRO 检查点重组出完整 fp32state_dict
ZeRO 检查点 fp32 权重恢复load_state_dict_from_zero_checkpointdeepspeed.utils.zero_to_fp32重组并直接写入一个给定模型
ZeRO 检查点 fp32 权重恢复convert_zero_checkpoint_to_fp32_state_dictdeepspeed.utils.zero_to_fp32重组并落盘为单个.bin文件

其中前两个 API 的完整实现位于 engine.py(load_checkpoint始于 L2597,save_checkpoint位于 L2941),后三个工具的完整实现位于 zero_to_fp32.py。检查点文件内部的键名与命名约定统一定义在 constants.py 中。

二、保存训练检查点:save_checkpoint

2.1 函数签名与参数

def save_checkpoint(self, save_dir, tag=None, client_state={}, save_latest=True):

依据 engine.py L2941-L2953 的文档注释,各参数含义如下:

  • save_dir(必填):检查点保存目录。方法内部会调用os.makedirs(save_dir, exist_ok=True)确保目录存在。
  • tag(可选):检查点的唯一标识符。不传时默认使用当前全局步数,即tag = f"global_step{self.global_steps}"(见 L2967-L2968),例如global_step14。注意:所有 rank 上的 tag 必须一致,方法内部通过_checkpoint_tag_validation(L2923)校验这一点。
  • client_state(可选):由用户代码自定义的附加状态字典,例如 epoch 编号、已处理样本数等,会被合并进检查点一并保存。
  • save_latest(可选):是否在save_dir下写入一个名为latest的文件,内容为最新一次保存的 tag 字符串,供后续无 tag 加载时定位。

2.2 关键语义:所有进程都必须调用

save_checkpoint的文档注释明确强调(L2949-L2952):

Important: all processes must call this method and not just the process with rank 0. It is because each process needs to save its master weights and scheduler+optimizer states. This method will halt waiting to synchronize with other processes if it's called just for the process with rank 0.

每个进程都持有自己的模型主权重以及调度器/优化器状态的分片,因此不能只在 rank 0 上调用;方法内部会多次调用dist.barrier()做全局同步,如果只让 rank 0 调用而其他 rank 不调用,整个调用会挂起等待。这是一个极易踩坑的分布式约定。

2.3 内部执行流程

从 L2954-L3005 可以看到完整的保存流水线:

  1. 若启用了 ZeRO 权重分区(zero_optimization_partition_weights()),先调用self.optimizer.checkpoint_event_prologue()准备参数分区;
  2. 创建目录并执行dist.barrier()同步;
  3. 确定tag(缺省为global_step{global_steps})并统一为字符串,通过self.checkpoint_engine.create(tag)建立存储会话;
  4. 调用_create_checkpoint_file创建模型状态文件(mp_rank_XX_model_states.pt);
  5. 调用_save_checkpoint写入常规状态;
  6. 若启用了 ZeRO 优化(self.save_zero_checkpoint),调用_create_zero_checkpoint_files与_save_zero_checkpoint单独保存 ZeRO 优化器状态文件;
  7. 若为 ZeRO 权重分区模式,调用checkpoint_event_epilogue()收尾;
  8. self.checkpoint_engine.commit(tag)提交,若save_latest=True且为 rank 0,则把 tag 写入latest文件;
  9. 最后一次dist.barrier()同步后返回True。

2.4 检查点文件里到底存了什么

_save_checkpoint(L3164-L3180)构造的 state dict 包含:

state = dict(module=module, # 模型权重 state_dict buffer_names=self._get_buffer_names(), # 非持久 buffer 名列表 optimizer=self.optimizer.state_dict() if self.optimizer and not zero_optimizer_state else None, param_shapes=self._get_zero_param_shapes() if self.optimizer and zero_optimizer_state else None, lr_scheduler=self.lr_scheduler.state_dict() if self.lr_scheduler is not None else None, sparse_tensor_module_names=self.sparse_tensor_module_names, skipped_steps=self.skipped_steps, global_steps=self.global_steps, global_samples=self.global_samples, dp_world_size=self.dp_world_size, mp_world_size=self.mp_world_size, ds_config=self.config, ds_version=version) state.update(client_state) # 合并用户自定义状态

要点:

  • ZeRO 模式下不保存完整 optimizer state(zero_optimizer_state为真时optimizer为None),而是保存param_shapes—— 由_get_zero_param_shapes(L3207)生成的"参数名 → 形状"的有序映射,它是之后从 ZeRO 分片重组完整权重的关键元数据;
  • buffer_names由_get_buffer_names(L3186)遍历模块树收集,用于在恢复时把混在参数占位符中的 buffer 正确识别出来;
  • ds_config与ds_version会被写入检查点,便于日后回放训练配置。

而_save_zero_checkpoint(L3259-L3269)保存的是 ZeRO 优化器状态分片文件:

zero_sd = dict(optimizer_state_dict=self.optimizer.state_dict(), ds_config=self.config, ds_version=version)

同时在 rank 0 上会把恢复脚本zero_to_fp32.py拷贝到检查点目录(_copy_recovery_script,L3249-L3257),方便日后随时离线转换。

2.5 磁盘布局与命名约定

依据 constants.py L29-L36 的命名常量,一次保存产生的典型目录结构为:

save_dir/ ├── latest # 文本文件,内容为最新 tag,如 "global_step14" ├── zero_to_fp32.py # 自动拷贝的恢复脚本(ZeRO 模式下) └── global_step14/ # tag 目录 ├── mp_rank_00_model_states.pt # 常规模型状态(ZeRO-2) ├── zero_pp_rank_0_mp_rank_00_model_states.pt # ZeRO-3 模型状态 ├── zero_pp_rank_0_mp_rank_00_optim_states.pt # 优化器分片状态 ├── zero_pp_rank_1_mp_rank_00_optim_states.pt └── ...

其中OPTIM_FILE_SUFFIX = '_optim_states.pt'、MODEL_FILE_SUFFIX = '_model_states.pt',ZeRO-2 与 ZeRO-3 的模型状态文件名前缀不同(mp_rank_与zero_pp_rank_),zero_to_fp32.py中的get_model_state_file(zero_to_fp32.py L49-L62)正是依据这两条规则定位模型状态文件。

三、加载训练检查点:load_checkpoint

3.1 函数签名与参数

def load_checkpoint(self, load_dir, tag=None, load_module_strict=True, load_optimizer_states=True, load_lr_scheduler_states=True, load_module_only=False, custom_load_fn=None):

依据 engine.py L2597-L2622 的参数说明:

  • load_dir(必填):检查点所在目录;
  • tag(可选):唯一标识符;不传时尝试读取load_dir/latest文件中的 tag(ZeRO 通用检查点模式则读取latest_universal,见 L2624-L2641)。若latest文件不存在且未显式传 tag,会打印警告并返回(None, None);
  • load_module_strict(可选):是否严格要求模块state_dict键与检查点完全匹配;
  • load_optimizer_states(可选):是否加载优化器状态(如 Adam 的动量与方差);
  • load_lr_scheduler_states(可选):是否加载 LR 调度器状态;
  • load_module_only(可选):只加载模型权重(如用于 warm-start 迁移学习),此时不恢复优化器/调度器/步数;
  • custom_load_fn(可选):自定义模型加载函数。

返回值:二元组(load_path, client_state):

  • load_path:实际加载的检查点路径,加载失败时为None;
  • client_state:从检查点中提取出的用户自定义状态字典,供客户端代码恢复 epoch、样本数等训练现场。

3.2 内部执行流程与 ZeRO 特殊处理

从 L2643-L2667 可以看到:

  1. ZeRO 权重分区模式下先执行checkpoint_event_prologue();
  2. 调用_load_checkpoint完成常规加载(模型权重、LR 调度器、非 ZeRO 优化器状态、global_steps/global_samples/skipped_steps等计数,见 L2781-L2794);
  3. 若启用 ZeRO 或 bfloat16,调用_load_zero_checkpoint加载优化器分片;失败时回退optimizer._restore_from_bit16_weights();
  4. ZeRO 权重分区模式下执行checkpoint_event_epilogue()。

_load_zero_checkpoint(L2813-L2843)中有一条重要约束:ZeRO 检查点不允许在数据并行(DP)world size 变化时直接加载优化器状态,若dp_world_size != loaded_checkpoint_dp_world_size且需要加载优化器状态,会抛出ZeRORuntimeException。

3.3client_state的提取机制

client_state由 L2802-L2806 计算得出:先构建一个内部保留键集合deepspeed_states(包含module、sparse_tensor_module_names、skipped_steps、global_steps、dp_world_size、mp_world_size,并按需追加lr_scheduler、optimizer),然后返回检查点中不属于该集合的所有键——这正是保存时state.update(client_state)写入的自定义状态。因此,只要保存时把自定义状态放进client_state,加载时就能原样取回。

3.4 ZeRO-3 下的一个已知限制

load_checkpoint文档注释特别提醒(L2618-L2621):在 ZeRO-3 下,不能在save_checkpoint()之后立刻调用load_checkpoint(),因为此时engine.module仍是分区状态,而load_checkpoint()需要一个"干净"(未分区)的模型。如果确有需要,应在load_checkpoint()之前重新初始化 engine。

四、在训练代码中的标准用法

DeepSpeed 官方 BERT 预训练教程 bert-pretraining.md(L159-L199)给出了完整的客户端调用模式。保存时把自定义状态打包进client_state:

def checkpoint_model(PATH, ckpt_id, model, epoch, last_global_step, last_global_data_samples, **kwargs): """Utility function for checkpointing model + optimizer dictionaries The main purpose for this is to be able to resume training from that instant again """ checkpoint_state_dict = {'epoch': epoch, 'last_global_step': last_global_step, 'last_global_data_samples': last_global_data_samples} # Add extra kwargs too checkpoint_state_dict.update(kwargs) success = model.network.save_checkpoint(PATH, ckpt_id, checkpoint_state_dict) return

加载时通过返回值中的client_state恢复训练现场:

def load_training_checkpoint(args, model, PATH, ckpt_id): """Utility function for checkpointing model + optimizer dictionaries The main purpose for this is to be able to resume training from that instant again """ _, checkpoint_state_dict = model.network.load_checkpoint(PATH, ckpt_id) epoch = checkpoint_state_dict['epoch'] last_global_step = checkpoint_state_dict['last_global_step'] # ... 恢复 LR、样本计数并继续训练

这套"保存时打包、加载时拆包"的模式同样适用于本项目仓库中的 Megatron 教程 megatron.md(L256-L257 处同样声明了save_checkpoint/load_checkpoint的签名),是 DeepSpeed 训练脚本的标准范式。

五、ZeRO 检查点 fp32 权重恢复:三种工具函数

5.1 为什么需要"恢复"这一步

在 ZeRO 优化下,模型的 fp32 主权重被打散存放在各 rank 的优化器状态中(ZeRO-2 按参数组切分、ZeRO-3 按参数切分),检查点目录里是一堆*_optim_states.pt分片文件,无法直接用torch.load+load_state_dict得到完整模型。zero_to_fp32模块的作用就是把分片按保存时记录的param_shapes重新拼接成完整 fp32state_dict,得到的权重不依赖 DeepSpeed,可用于任何 PyTorch 应用或发布到模型库。

5.2get_fp32_state_dict_from_zero_checkpoint

签名与行为见 zero_to_fp32.py L360-L406:

def get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir, tag=None):
  • checkpoint_dir:检查点根目录(包含 tag 子目录的那一层);
  • tag:检查点标识,如global_step14;缺省时读取latest文件,找不到则抛出ValueError;
  • 返回值:一个已位于 CPU 上的完整 PyTorchstate_dict。

典型用法(文档注释中的示例):

from deepspeed.utils.zero_to_fp32 import get_fp32_state_dict_from_zero_checkpoint # do the training and checkpoint saving state_dict = get_fp32_state_dict_from_zero_checkpoint(checkpoint_dir) # already on cpu model = model.cpu() # move to cpu model.load_state_dict(state_dict) # submit to model hub or save the model to share with others

文档注释同时提醒:执行此操作后,model将不再适用于同一个 DeepSpeed 上下文(load_state_dict会移除模型上的 DeepSpeed 包装),如需继续 DeepSpeed 训练必须重新初始化 engine。

5.3load_state_dict_from_zero_checkpoint

签名与行为见 zero_to_fp32.py L425-L461:

def load_state_dict_from_zero_checkpoint(model, checkpoint_dir, tag=None):

它把三步打包成一个调用:将给定模型移到 CPU → 重组 ZeRO-2/3 检查点为完整 fp32state_dict→ 以strict=False加载进模型,然后返回模型本身:

from deepspeed.utils.zero_to_fp32 import load_state_dict_from_zero_checkpoint model = load_state_dict_from_zero_checkpoint(trainer.model, checkpoint_dir) # submit to model hub or save the model to share with others

文档注释特别提醒:调用前请确保有充足的 CPU 内存(重组过程会在内存中同时持有分片与完整权重);内存不足时改用zero_to_fp32.py命令行工具离线转换。

5.4convert_zero_checkpoint_to_fp32_state_dict

签名与行为见 zero_to_fp32.py L409-L422:

def convert_zero_checkpoint_to_fp32_state_dict(checkpoint_dir, output_file, tag=None):

在get_fp32_state_dict_from_zero_checkpoint基础上把结果torch.save成单个文件,例如path/checkpoint-12/pytorch_model.bin,之后即可用torch.load(file)+load_state_dict()在任意 PyTorch 应用中使用。

5.5 命令行离线转换:zero_to_fp32.py

save_checkpoint在 ZeRO 模式下会把 zero_to_fp32.py 自动拷贝到检查点目录(见 engine.py L3249-L3267),因此可以随时在任意机器上离线转换,无需原始训练进程。该脚本自带 CLI(L464-L482):

python zero_to_fp32.py . pytorch_model.bin # 在检查点目录内执行 python zero_to_fp32.py checkpoint_dir pytorch_model.bin -d # -d 开启调试输出

脚本头部的注释(L1-L8)说明了它的定位:虽然恢复过程本身不调用 DeepSpeed 做计算,但由于检查点是用 DeepSpeed 数据结构 pickle 的,当前 Python 环境仍需安装 DeepSpeed(至少可 importdeepspeed.utils与deepspeed.checkpoint.constants)。

六、fp32 权重重组原理:ZeRO-2 与 ZeRO-3 的差异

重组的核心流程在_get_fp32_state_dict_from_zero_checkpoint(zero_to_fp32.py L153-L181)中:

  1. get_optim_files(L65-L75)按自然排序收集所有*_optim_states.pt;
  2. parse_optim_states(L100-L150)读取optimizer_state_dict中的zero_stage与partition_count,据此选择分组键——ZeRO-2 用single_partition_of_fp32_groups,ZeRO-3 用fp32_flat_groups,并把各 rank 的分片按需拼接成fp32_flat_groups列表;
  3. parse_model_state(L78-L97)从模型状态文件中恢复 buffer(必要时从 fp16 转回 fp32)与param_shapes;
  4. 依据zero_stage分派到 ZeRO-2 或 ZeRO-3 的重组实现。

ZeRO-2的重组(_get_fp32_state_dict_from_zero2_checkpoint,L184-L277):每个 rank 保存的是完整 fp32 权重在各自优化器中的"单一分区",因此只需把各分片torch.cat成完整向量,再按param_shapes逐个narrow+view切回各参数。由于 ZeRO-2 出于 NCCL 性能考虑做过2 * world_size对齐(zero2_align,L256-L263),代码会按同样规则对齐后再做数值一致性校验,不匹配即抛ValueError。

ZeRO-3的重组(_get_fp32_state_dict_from_zero3_checkpoint,L287-L357):每个参数被按world_size均分并可能做了 padding(zero3_partitioned_param_info,L280-L284),因此需要"在参数边界处把分片重新拉链":对每个参数,从各 rank 的扁平向量中截取partitioned_numel拼接,再narrow掉 padding 部分并view成原始形状。

两个实现都带有XXX注释指出:大模型场景下内存开销会翻倍(拼接过程同时持有分片与完整权重),这正是文档建议"确保有充足 CPU 内存、否则用离线脚本"的原因。

七、配套测试与验证

仓库在 tests/unit/checkpoint 目录下提供了覆盖检查点各分支的单元测试,可作为行为验证的参考:

  • test_latest_checkpoint.py:验证latest文件驱动的无 tag 加载;
  • test_lr_scheduler.py:验证 LR 调度器状态存取;
  • test_moe_checkpoint.py:MoE(混合专家)模型检查点;
  • test_pipeline.py:流水线并行下的检查点;
  • test_tag_validation.py:跨 rank 的 tag 一致性校验;
  • test_zero_optimizer.py:ZeRO 优化器状态检查点;
  • test_reshape_checkpoint.py:检查点维度变形(供后续权重形状调整场景参考)。

八、使用注意事项与最佳实践

综合文档与源码实现,总结以下实操要点:

  1. 保存时所有 rank 必须协同调用save_checkpoint,否则dist.barrier()会挂起;tag 在所有 rank 上必须一致。
  2. 加载时无 tag 依赖latest文件,若该文件缺失且未传 tag,load_checkpoint只返回(None, None)并打印警告,需要显式传入 tag。
  3. ZeRO 检查点不允许跨 DP world size 加载优化器状态(会抛ZeRORuntimeException);若只想迁移权重,可用load_module_only=True或改用zero_to_fp32系列函数。
  4. ZeRO-3 下save_checkpoint后不能立刻load_checkpoint,需重新初始化 engine。
  5. fp32 权重恢复需要充足 CPU 内存;内存受限时优先用随检查点落盘的zero_to_fp32.py命令行脚本离线转换。
  6. 恢复出的权重不依赖 DeepSpeed,但转换过程所在环境仍需安装 DeepSpeed(pickle 反序列化需要其数据结构定义)。
  7. 自定义训练现场状态(epoch、样本数等)统一放进client_state,加载后从返回值的第二项取回,与 DeepSpeed 内部状态互不干扰。

九、总结

DeepSpeed 的模型检查点体系由两套 API 组成:引擎级save_checkpoint/load_checkpoint负责在分布式训练中完整地保存与恢复"模型 + 优化器 + 调度器 + 用户状态"的训练现场(engine.py),而zero_to_fp32工具族(zero_to_fp32.py)负责把 ZeRO 化的 fp32 权重从优化器分片中重组为标准的 PyTorchstate_dict,打通了"DeepSpeed 训练 → 通用模型发布"的最后一公里。理解本文所讲的目录布局、键名约定、ZeRO-2/3 重组差异与分布式同步约束,即可在各类训练脚本中安全地实现断点续训与模型导出。

  • 推理引擎
  • 大模型

【免费下载链接】FlexGen

Running large language models on a single GPU for throughput-oriented scenarios.

项目地址:https://gitcode.com/gh_mirrors/fl/FlexGen
点击查看免费下载

相关推荐

上一篇:番茄小说下载器:5分钟搭建个人数字图书馆,永久保存你的阅读时光
下一篇:用AI魔法将2D视频瞬间变立体3D:Deep3D深度解析

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

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

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

立即咨询