- AI 技能
- 人工智能
- 大模型
- 深度学习
【免费下载链接】AI-Research-SKILLs
Comprehensive open-source library of AI research and engineering skills for any AI model. Package the skills and your claude code/codex/gemini agent will be an AI research agent with full horsepower. Maintained by Orchestra Research.
导读
本文围绕 PyTorch 官方 "Getting Started with Distributed Checkpoint (DCP)" recipe 展开,系统讲解torch.distributed.checkpoint的核心机制、Stateful封装模式、dcp.save/dcp.load基本用法,以及torch.distributed.checkpoint.state_dict下的分布式状态字典辅助函数。结合本仓库pytorch-fsdp2技能包中的 SKILL.md 与 pytorch_dcp_recipe.md 等参考文档,你将掌握:为什么 FSDP2 训练脚本应默认采用 DCP 而非朴素torch.save、如何用 DCP 在并行拓扑间自由迁移检查点、以及如何通过异步保存把检查点开销移出训练关键路径。
一、为什么 FSDP2 训练脚本需要 DCP
1.1 DTensor 分片状态字典无法朴素序列化
FSDP2 的核心特征是逐参数 DTensor 分片(per-parameter sharding)。在调用fully_shard()之后,模型的参数被转换为DTensor,张量数据分散在多张 GPU 上,每个 rank 只持有完整参数的某个分片。此时若直接执行torch.save(model.state_dict()),保存下来的只是每个 rank 本地的分片视图,既不是完整的参数张量,也无法表达分片元数据,加载时必然出错。
从 pytorch_fsdp2_tutorial.md 的对照可以看出,官方教程明确给出了两条状态字典工作流:
- 方案 A(DTensor 手动 API):保存时调用
DTensor.full_tensor()做 all-gather 汇聚成完整张量(可在 rank 0 上做 CPU offload 以避免 GPU 峰值内存);加载时先用distribute_tensor(full_tensor, meta_param.device_mesh, meta_param.placements)把完整张量重新分发,再model.load_state_dict(..., assign=True)。 - 方案 B(DCP 分布式状态字典辅助函数):保存用
get_model_state_dict(..., StateDictOptions(full_state_dict=True, cpu_offload=True)),加载用set_model_state_dict(..., StateDictOptions(full_state_dict=True, broadcast_from_rank0=True))。
方案 B 正是 recipe 推荐的最安全默认,也是本文的主题。
1.2 DCP 解决的三大痛点
根据 pytorch_dcp_overview.md 的总结,DCP 与朴素序列化相比有三个本质差异:
| 特性 | 朴素torch.save | Distributed Checkpoint |
|---|---|---|
| 保存/加载方式 | 单进程串行 | 多 rank 并行,每个 rank 只写自己持有的分片 |
| 拓扑适配 | 固定,无法跨集群拓扑迁移 | 加载时自动 resharding,可跨拓扑自由迁移 |
| 产物形态 | 单个.pt文件 | 多个文件(通常每 rank 至少一个) |
DCP 还是一种in-place 操作:模型先自行分配好存储空间,DCP 直接把数据加载进既有存储,而不是像load_state_dict那样整体替换状态。
⚠️ 重要边界:官方文档明确警告,DCP 保存的
state_dict不保证跨 PyTorch 版本的向后兼容。若你的工程需要严格跨版本恢复检查点,需要自行评估这一限制(见 SKILL.md 的 "Avoid" 清单)。
二、DCP 基本用法:Stateful 封装 + dcp.save / dcp.load
pytorch_dcp_recipe.md 给出的高层示例结构只有三个要点,但它们是理解全部 DCP 用法的骨架:
- 把应用状态包装进
Stateful对象,让 DCP 自动调用state_dict()/load_state_dict(); - 用
dcp.save(...)/dcp.load(...)完成落盘与恢复; - 用
get_state_dict/set_state_dict辅助函数,在分布式环境下正确取得并施加模型/优化器状态字典。
2.1 最小可运行骨架
以下代码综合了 recipe 的高层结构与 pytorch_fsdp2_tutorial.md、pytorch_fully_shard_api.md 中描述的模式,可直接作为给训练脚本接入 DCP 的起点:
import os import torch import torch.distributed as dist import torch.distributed.checkpoint as dcp from torch.distributed.checkpoint.state_dict import ( get_model_state_dict, get_optimizer_state_dict, set_model_state_dict, set_optimizer_state_dict, StateDictOptions, ) # ---------- 1. 初始化分布式环境(FSDP2 前置步骤) ---------- def init_distributed(): dist.init_process_group(backend="nccl") torch.cuda.set_device(int(os.environ["LOCAL_RANK"])) # ---------- 2. 保存:把“模型 + 优化器”打包成 Stateful 状态 ---------- def checkpoint_save(model, optimizer, path): # DCP 会自动对实现了 state_dict()/load_state_dict() 的对象调用对应方法。 # 这里直接把模型与优化器放进字典,即构成 recipe 所说的“应用状态”。 state = {"model": model, "optimizer": optimizer} dcp.save(state, checkpoint_id=path) # 所有 rank 都调用,各写各的分片 # ---------- 3. 加载:先让模型/优化器自行分配存储,再原位填充 ---------- def checkpoint_load(model, optimizer, path): # 先建立与保存时一致的“占位状态”(此时模型/优化器已在 meta 或真实设备上建好) state = {"model": model, "optimizer": optimizer} dcp.load(state, checkpoint_id=path) # in-place:加载进既有存储 # ---------- 4. 进阶:使用 state_dict 辅助函数(推荐默认) ---------- def checkpoint_save_with_helpers(model, optimizer, path): # full_state_dict=True 表示保存完整参数(而非分片视图) # cpu_offload=True 表示在 rank 0 上把完整张量落到 CPU,避免 GPU 峰值内存 model_sd = get_model_state_dict( model, options=StateDictOptions(full_state_dict=True, cpu_offload=True)) opt_sd = get_optimizer_state_dict( model, optimizer, options=StateDictOptions(full_state_dict=True, cpu_offload=True)) dcp.save({"model": model_sd, "optimizer": opt_sd}, checkpoint_id=path) def checkpoint_load_with_helpers(model, optimizer, path): # 加载时先按分片形状建立空状态,再由 DCP 原位填充 model_sd = get_model_state_dict(model) # 保持分片形态的“空壳” opt_sd = get_optimizer_state_dict(model, optimizer) dcp.load({"model": model_sd, "optimizer": opt_sd}, checkpoint_id=path) # broadcast_from_rank0=True 可让 rank 0 的完整状态广播给所有 rank set_model_state_dict( model, model_sd, options=StateDictOptions(full_state_dict=True, broadcast_from_rank0=True)) set_optimizer_state_dict(model, optimizer, opt_sd)关键点说明:
- 所有 rank 必须一起调用
dcp.save/dcp.load,它们内部按WORLD_SIZE协调写入与读取; - 加载是 in-place 的:模型必须已经完成内存分配(例如通过 meta 设备初始化流程 中的
to_empty(device="cuda")之后再load_state_dict),DCP 只负责填充数据; - 优化器必须在
fully_shard()之后创建,以确保其持有的是 DTensor 参数,否则辅助函数拿到的状态字典与模型分片形态不一致(见 SKILL.md 契约第 4 条)。
2.2 手动 Stateful 封装(替代直接传字典)
recipe 提到 "Wrap application state in aStatefulobject, so DCP automatically callsstate_dict()/load_state_dict()"。即你还可以定义自定义的Stateful类,把学习率调度器、随机数生成器状态等一并纳入检查点:
from torch.distributed.checkpoint.stateful import Stateful class TrainState(Stateful): def __init__(self, model, optimizer, lr_scheduler): self.model = model self.optimizer = optimizer self.lr_scheduler = lr_scheduler self.step = 0 def state_dict(self): return { "model": self.model.state_dict(), "optimizer": self.optimizer.state_dict(), "lr_scheduler": self.lr_scheduler.state_dict(), "step": self.step, } def load_state_dict(self, state_dict): self.model.load_state_dict(state_dict["model"]) self.optimizer.load_state_dict(state_dict["optimizer"]) self.lr_scheduler.load_state_dict(state_dict["lr_scheduler"]) self.step = state_dict["step"] # 保存/加载时直接把 Stateful 对象交给 DCP: # dcp.save({"train_state": TrainState(...)}, checkpoint_id=path) # dcp.load({"train_state": TrainState(...)}, checkpoint_id=path)这种做法的优势:DCP 对字典内每个条目统一走state_dict()/load_state_dict()协议,训练元数据(step、epoch)与模型参数天然同批落盘、同批恢复。
三、FSDP2 全流程中 DCP 的接入位置
为保证文章实战可用,下面给出一个与 SKILL.md "Minimal reference implementation outline" 对应的端到端流程,标出 DCP 在其中的确切位置:
1. init_distributed() # dist.init_process_group(backend="nccl") + set_device(LOCAL_RANK) 2. build_model_meta() # with torch.device("meta"): model = ... # → 对 TransformerBlock 子模块逐个 fully_shard(m, ...) # → 最后 fully_shard(model)(自底向上) # → model.to_empty(device="cuda") + model.reset_parameters() 3. build_optimizer() # 在 fully_shard 之后创建,持有 DTensor 参数 4. train_step() # model(inputs)(勿用 model.forward),DTensor 感知的梯度裁剪 5. checkpoint_save/load() # ← DCP 或 state_dict 辅助函数在这里接入3.1 自底向上分片为何是前提
pytorch_fully_shard_api.md 强调:"Users generally should not call fully_shard() only on the topmost root module."fully_shard会把已由前序调用分组过的参数排除在外,先分片子模块、再分片根模块,能形成更细粒度的通信组,带来更好的通信重叠与更低的峰值内存。这直接影响 DCP 保存的分片边界——分片粒度越合理,检查点文件的可重分片性(resharding)越稳定。
3.2 分片配置对检查点的影响
fully_shard的关键参数(同样记录在 pytorch_fully_shard_api.md):
mesh:1DDeviceMesh即经典 FSDP 分片(placement 为(Shard(0),));2D mesh 即 Hybrid sharding,placement 为(Replicate(), Shard(0)),跨一个维度分片、另一个维度复制。mesh 拓扑决定 DCP 保存时每个 rank 写入哪些分片;reshard_after_forward:None时非根模块默认True、根模块默认False;True在前向后释放非分片参数(内存优先),False保留(吞吐优先);mp_policy=MixedPrecisionPolicy(param_dtype=..., reduce_dtype=..., output_dtype=..., cast_forward_inputs=...):控制前向/反向中的参数与梯度 dtype,间接影响落盘数值精度;offload_policy=CPUOffloadPolicy():把参数/优化器状态放到 CPU,此时 DCP 保存会跨 CPU/GPU 边界搬运数据,需评估 PCIe/NVLink 流量开销。
从源码结构看,DCP 的 resharding 能力正是建立在 DTensor 的
DeviceMesh+Placement元数据之上的:保存时记录分片布局,加载时按目标拓扑重新计算放置方式,从而做到"保存用 8 卡拓扑、加载用 16 卡拓扑"。
四、深入理解:跨拓扑 Resharding 与检查点目录结构
4.1 Resharding 的实战含义
pytorch_dcp_recipe.md 的核心主张之一:"DCP saves/loads in parallel, and supports resharding across topologies at load time."
这意味着同一份检查点可以:
- 从4×GPU训练保存,在8×GPU上继续训练(数据并行维度扩大);
- 从FSDP2 分片保存,加载进Tensor Parallel + FSDP2 混合并行的配置;
- 在单卡上加载完整模型做评估或微调导出。
这正是 torchtitan/checkpoint.md 中"checkpoints saved with DCP can be resharded for different parallelism configurations"所述的场景:生产级 LLM 训练框架(TorchTitan)把 DCP 作为故障恢复与互操作检查点的标准方案。
4.2 DCP 的落盘形态
DCP 保存的目录结构(来自 torchtitan/checkpoint.md 的实际工程形态):
checkpoint/ ├── step-500/ │ ├── .metadata # 全局元数据(分片布局、版本信息) │ ├── __0_0.distcp # rank 0 写入的分片文件 │ ├── __0_1.distcp │ └── ... └── step-1000/ └── ...要点:
.metadata记录整体状态字典的分片方案,是 resharding 的依据;- 每个 rank 至少产生一个
__<rank>_<chunk>.distcp文件,因此检查点不是单文件; - 若需要把 DCP 分片检查点转换为单个
.pt文件(例如导出给单卡推理),可使用官方转换工具:
python -m torch.distributed.checkpoint.format_utils \ dcp_to_torch \ path/to/dcp/checkpoint \ checkpoint.pt4.3 进程组边界注意事项
pytorch_dcp_overview.md 给出两条硬性约束:
- 若显式传入 process group,只有该组内的 rank 才能调用 save/load;
- 所有参与的张量必须属于该进程组,混入其他组张量会导致协调失败。
在 FSDP2 + Tensor Parallel 组合(2D mesh)场景下,这要求你明确检查点逻辑作用于哪个维度的进程组,避免跨组混用。
五、异步保存:把检查点移出训练关键路径
当检查点体积大、保存耗时显著拖慢训练步进时,pytorch_dcp_async_recipe.md 提供了torch.distributed.checkpoint.async_save方案。
5.1 机制与代价
异步保存的本质:先把模型状态拷贝进内部 CPU 缓冲区,再在后台线程/进程完成落盘,训练循环无需等待磁盘写入完成。代价是:
- 引入额外内存开销——保存瞬间需要一份与模型状态等价的 CPU 缓冲;
- 若内存吃紧,可参考 recipe 中描述的pinned memory(固定内存)策略来改善拷贝与 DMA 性能。
5.2 使用模式
import torch.distributed.checkpoint as dcp state = {"model": model, "optimizer": optimizer} # 同步保存(默认):训练循环阻塞至写盘完成 dcp.save(state, checkpoint_id=path_sync) # 异步保存:立即返回,后台完成落盘 dcp.async_save(state, checkpoint_id=path_async)从 SKILL.md 的 Workflow B 看,无论同步还是异步,调用的都是同一套"先组装 state → 所有 rank 调用 → 恢复时 set_state_dict"的流程,只是把dcp.save换成dcp.async_save。
5.3 何时该用
- 检查点停顿显著(例如每 N 步全集群同步等待写盘)→ 用异步保存;
- CPU 内存有余量→ 可以承受拷贝缓冲开销;
- 保存频率高、单次写盘慢→ 异步收益最大。
TorchTitan 的生产配置也印证了这一取舍(见 torchtitan/checkpoint.md):
[checkpoint] enable = true folder = "checkpoint" interval = 500 async_mode = "async" # 可选: "disabled" / "async" / "async_with_pinned_mem"async_with_pinned_mem即对应 recipe 中提到的 pinned memory 优化路径。
六、实战清单:把 DCP 接入既有 FSDP2 脚本
综合 SKILL.md 的 Workflow B 与 pytorch_dcp_recipe.md 的指导,最小接入路径如下(可直接作为 Agent 的验收清单):
- 用
torchrun --nproc_per_node <gpus_per_node> ...启动,确保RANK/WORLD_SIZE/LOCAL_RANK可见; - 初始化进程组并
torch.cuda.set_device(LOCAL_RANK); fully_shard自底向上完成分片后再创建优化器(保证 DTensor 参数);- 用
Stateful包装或get_state_dict组装模型 + 优化器状态; - 所有 rank 调用
dcp.save(...)(或dcp.async_save(...))到共享路径; - 加载时先分配存储,再
dcp.load(...),最后用set_state_dict施加; - 若目标拓扑与保存时不同,显式验证 resharding 假设(mesh 维度、分片度是否匹配);
- 留意 DCP 的 PyTorch 版本兼容性警告,不要在同一个训练工程里混用 DCP 与临时
torch.save。
常见错误速查
| 症状 | 根因 | 修复 |
|---|---|---|
| 加载后参数全乱/形状不匹配 | 保存与加载的拓扑或分片配置不一致 | 核对两端的fully_shard策略与 mesh;利用 DCP 的 resharding 能力而非手动拼接 |
torch.save保存的只是本地分片 | 对 DTensor 直接朴素序列化 | 改用dcp.save,或先DTensor.full_tensor()汇聚(注意内存) |
| 优化器状态与模型对不上 | 优化器创建早于fully_shard | 把优化器创建移到所有fully_shard调用之后 |
| 异步保存后 OOM | 内部 CPU 缓冲占用过多 | 换同步保存,或采用 pinned memory / 降频保存策略 |
| 跨 PyTorch 版本加载失败 | DCP 不保证向后兼容 | 升级时重新导出检查点,或放弃跨版本恢复 |
七、结论与进一步阅读
给 FSDP2 训练脚本添加检查点时,DCP 模式是最安全的默认选择:它天然适配 DTensor 分片状态、支持多 rank 并行与加载时跨拓扑 resharding,并可通过async_save把落盘开销移出关键路径。核心 API 面只有三块——dcp.save/load、Stateful协议、torch.distributed.checkpoint.state_dict辅助函数,但三者组合起来即可覆盖从单卡导出到大规模集群故障恢复的全部场景。
本仓库pytorch-fsdp2技能包内的相关文档可继续深入:
- pytorch_dcp_recipe.md:本文主文档,DCP 入门 recipe;
- pytorch_dcp_overview.md:DCP 行为总览与重要注意事项;
- pytorch_dcp_async_recipe.md:异步保存 recipe;
- pytorch_fsdp2_tutorial.md:FSDP2 入门教程,含 DTensor vs DCP 两条状态字典工作流对照;
- pytorch_fully_shard_api.md:
fully_shardAPI 与分片语义细节; - pytorch_examples_fsdp2.md:官方
pytorch/examples中的 FSDP2 checkpoint 脚本入口; - SKILL.md:FSDP2 技能总览,含 DCP 接入契约与调试清单;
- torchtitan/checkpoint.md:TorchTitan 生产级 DCP 配置与目录结构示例。
若需在更高层训练编排器(如 Ray Train)中集成,可参考 ray_train_fsdp2_example.md;FSDP2 与 Tensor Parallel 的 mesh 组合细节见 pytorch_device_mesh_tutorial.md 与 pytorch_tp_tutorial.md。
- AI 技能
- 人工智能
- 大模型
- 深度学习
【免费下载链接】AI-Research-SKILLs
Comprehensive open-source library of AI research and engineering skills for any AI model. Package the skills and your claude code/codex/gemini agent will be an AI research agent with full horsepower. Maintained by Orchestra Research.
相关推荐
PyTorch Distributed Checkpoint (DCP) 实战解析:分片保存/加载、resharding 与异步 checkpoint
PyTorch Distributed Checkpoint DCP 实战解析:分片保存/加载、resharding 与异步 checkpoint 导读 Dis
人工智能机器学习深度学习分布式训练模型编译Megatron-LM 广义张量并行(GTP)深度解析:权重分片、异步预取与原生 DCP 重切分实战指南
Megatron LM 广义张量并行(GTP)深度解析:权重分片、异步预取与原生 DCP 重切分实战指南 本文基于 Megatron LM 官方 API 文档
人工智能大模型预训练分布式训练深度学习强化学习PyTorch FSDP2 全解:fully_shard 逐参数分片的全分片数据并行实现
PyTorch FSDP2 全解:fully_shard 逐参数分片的全分片数据并行实现 PyTorch FSDP2 以 torch.distributed.f
人工智能机器学习深度学习分布式训练模型编译
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考