fsdp_checkpoint_manager.py — FSDP 检查点管理¶
文件路径: verl/utils/checkpoint/fsdp_checkpoint_manager.py
文件概述¶
FSDP 训练策略下的检查点管理器。FSDP 的检查点需要处理分片状态字典(sharded state dict)的保存和恢复。
核心功能¶
保存¶
def save(self, step):
# 1. 用 FSDP 的 full_state_dict 或 sharded_state_dict 获取模型状态
# 2. 只在 rank 0 保存(full state dict)或所有 rank 保存(sharded)
# 3. 保存优化器状态
# 4. 保存训练元数据(step, epoch 等)
加载¶
def load(self, step=None):
# 1. 找到检查点路径
# 2. 加载模型 state dict
# 3. 用 FSDP 的 load_state_dict 恢复分片
# 4. 恢复优化器状态
FSDP 检查点的特殊性¶
- Full state dict: 所有 rank 的参数合并为完整模型,只需 rank 0 保存。优点是兼容性好,缺点是需要额外显存来聚合。
- Sharded state dict: 每个 rank 保存自己的分片。优点是省显存,缺点是加载时需要相同的并行度。
与其他模块的关系¶
- 继承自
checkpoint_manager.py - 依赖
fsdp_utils.py进行状态字典操作 - 依赖
fs.py进行文件 IO
小结¶
处理 FSDP 分布式训练下检查点的保存和恢复,是训练容错和断点续训的关键。