跳转至

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 分布式训练下检查点的保存和恢复,是训练容错和断点续训的关键。