- 推理引擎
- 大模型
【免费下载链接】FlexGen
Running large language models on a single GPU for throughput-oriented scenarios.
本指南以 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_checkpoint | deepspeed.runtime.engine | 保存模型权重、优化器状态、LR 调度器状态与用户自定义状态 |
| 加载训练检查点 | DeepSpeedEngine.load_checkpoint | deepspeed.runtime.engine | 恢复上述全部状态以继续训练 |
| ZeRO 检查点 fp32 权重恢复 | get_fp32_state_dict_from_zero_checkpoint | deepspeed.utils.zero_to_fp32 | 从 ZeRO 检查点重组出完整 fp32state_dict |
| ZeRO 检查点 fp32 权重恢复 | load_state_dict_from_zero_checkpoint | deepspeed.utils.zero_to_fp32 | 重组并直接写入一个给定模型 |
| ZeRO 检查点 fp32 权重恢复 | convert_zero_checkpoint_to_fp32_state_dict | deepspeed.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 可以看到完整的保存流水线:
- 若启用了 ZeRO 权重分区(
zero_optimization_partition_weights()),先调用self.optimizer.checkpoint_event_prologue()准备参数分区; - 创建目录并执行
dist.barrier()同步; - 确定
tag(缺省为global_step{global_steps})并统一为字符串,通过self.checkpoint_engine.create(tag)建立存储会话; - 调用
_create_checkpoint_file创建模型状态文件(mp_rank_XX_model_states.pt); - 调用
_save_checkpoint写入常规状态; - 若启用了 ZeRO 优化(
self.save_zero_checkpoint),调用_create_zero_checkpoint_files与_save_zero_checkpoint单独保存 ZeRO 优化器状态文件; - 若为 ZeRO 权重分区模式,调用
checkpoint_event_epilogue()收尾; self.checkpoint_engine.commit(tag)提交,若save_latest=True且为 rank 0,则把 tag 写入latest文件;- 最后一次
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 可以看到:
- ZeRO 权重分区模式下先执行
checkpoint_event_prologue(); - 调用
_load_checkpoint完成常规加载(模型权重、LR 调度器、非 ZeRO 优化器状态、global_steps/global_samples/skipped_steps等计数,见 L2781-L2794); - 若启用 ZeRO 或 bfloat16,调用
_load_zero_checkpoint加载优化器分片;失败时回退optimizer._restore_from_bit16_weights(); - 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 上的完整 PyTorch
state_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)中:
get_optim_files(L65-L75)按自然排序收集所有*_optim_states.pt;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列表;parse_model_state(L78-L97)从模型状态文件中恢复 buffer(必要时从 fp16 转回 fp32)与param_shapes;- 依据
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:检查点维度变形(供后续权重形状调整场景参考)。
八、使用注意事项与最佳实践
综合文档与源码实现,总结以下实操要点:
- 保存时所有 rank 必须协同调用
save_checkpoint,否则dist.barrier()会挂起;tag 在所有 rank 上必须一致。 - 加载时无 tag 依赖
latest文件,若该文件缺失且未传 tag,load_checkpoint只返回(None, None)并打印警告,需要显式传入 tag。 - ZeRO 检查点不允许跨 DP world size 加载优化器状态(会抛
ZeRORuntimeException);若只想迁移权重,可用load_module_only=True或改用zero_to_fp32系列函数。 - ZeRO-3 下
save_checkpoint后不能立刻load_checkpoint,需重新初始化 engine。 - fp32 权重恢复需要充足 CPU 内存;内存受限时优先用随检查点落盘的
zero_to_fp32.py命令行脚本离线转换。 - 恢复出的权重不依赖 DeepSpeed,但转换过程所在环境仍需安装 DeepSpeed(pickle 反序列化需要其数据结构定义)。
- 自定义训练现场状态(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.
相关推荐
DeepSpeed 模型检查点完全指南:save/load API、ZeRO fp32 权重恢复与 Universal Checkpoint
DeepSpeed 模型检查点完全指南:save/load API、ZeRO fp32 权重恢复与 Universal Checkpoint 导读 本文聚焦 D
人工智能大模型深度学习分布式训练预训练强化学习模型优化vLLM 模型权重流式加载实战:Run:ai Model Streamer 从对象存储加载模型与分片检查点
vLLM 模型权重流式加载实战:Run:ai Model Streamer 从对象存储加载模型与分片检查点 本文围绕 vLLM 对 Run:ai Model S
人工智能大模型模型推理服务推理引擎本地部署Flax 检查点保存与加载实战指南:基于 Orbax 的 Checkpointing 完整教程
Flax 检查点保存与加载实战指南:基于 Orbax 的 Checkpointing 完整教程 导读 本指南基于 Flax 官方文档 docs/guides/t
人工智能深度学习机器学习
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考