跳转至

state_dict.py — Checkpoint State Dict 管理

源码路径:verl/third_party/torch/distributed/checkpoint/state_dict.py

文件概述

这个文件从 PyTorch 2.7.0 复制而来(修复了 2.6.0 的 OOM 问题),提供了分布式训练中模型和优化器 state_dict 的高层管理 API。它是 _state_dict_utils.py 的上层封装,面向用户提供简洁的接口。

这是一个非常大的文件(约 1400 行),包含大量处理 FSDP、DDP 等分布式训练框架 state_dict 的逻辑。我们重点介绍核心概念和主要 API。

为什么需要这个文件?

在分布式训练中,获取和设置模型 state_dict 并不像单卡那样简单。比如使用 FSDP 时: - 获取 state_dict 需要先收集所有分片 - 设置 state_dict 需要将完整参数切分到各 GPU - 还需要处理 DDP wrapper、checkpoint wrapper 等带来的参数名前缀

这个文件封装了所有这些复杂性,提供统一的 API。

关键代码讲解

1. 配置选项 StateDictOptions

@dataclass
class StateDictOptions:
    full_state_dict: bool = False
    cpu_offload: bool = False
    ignore_frozen_params: bool = False
    keep_submodule_prefixes: bool = True
    strict: bool = True
    broadcast_from_rank0: bool = False
    flatten_optimizer_state_dict: bool = False

这个数据类控制 state_dict 操作的行为:

  • full_state_dict:是否收集完整的 state_dict(而非保留分片形式)
  • cpu_offload:是否将张量移到 CPU。当 full_state_dict=True 时,只有 rank 0 获得完整数据,其他 rank 得到空字典,避免 OOM
  • ignore_frozen_params:是否忽略冻结参数(requires_grad=False)
  • strict:加载时是否严格匹配 key
  • broadcast_from_rank0:加载时是否从 rank 0 广播。适用于只有 rank 0 有完整数据的场景

2. 获取模型 state_dict get_model_state_dict()

这是最常用的 API 之一。它的核心逻辑是:

def get_model_state_dict(
    model: nn.Module,
    *,
    submodules: Optional[set[nn.Module]] = None,
    options: Optional[StateDictOptions] = None,
) -> dict[str, ValueType]:

内部流程(简化版):

  1. 处理 FSDP 上下文:如果模型使用了 FSDP,需要设置正确的 StateDictType
  2. full_state_dict=True → 使用 StateDictType.FULL_STATE_DICT
  3. full_state_dict=False → 使用 StateDictType.SHARDED_STATE_DICT

  4. 调用 model.state_dict():在正确的上下文中获取 state_dict

  5. 后处理:

  6. 如果需要 full_state_dict,调用 _gather_state_dict 收集所有分片
  7. 如果需要 cpu_offload,将数据移到 CPU
  8. 过滤子模块前缀、冻结参数等

3. 设置模型 state_dict set_model_state_dict()

def set_model_state_dict(
    model: nn.Module,
    model_state_dict: dict[str, ValueType],
    *,
    options: Optional[StateDictOptions] = None,
) -> _IncompatibleKeys:

这是 OOM 修复的重点所在。设置 state_dict 时需要处理多种场景:

场景 1:broadcast_from_rank0=True

# 只有 rank 0 有完整数据,广播给其他 rank
_broadcast_state_dict(
    model_state_dict,  # rank 0 的完整数据
    local_state_dict,  # 各 rank 的本地分片模板
    device,
    pg,
    cpu_offload=info.cpu_offload,
)

场景 2:full_state_dict=True(所有 rank 都有完整数据)

# 每个 rank 自己切分需要的部分
_distribute_state_dict(
    model_state_dict,
    local_state_dict,
    device,
    pg,
)

场景 3:分片 state_dict

# 直接用 model.load_state_dict() 加载
model.load_state_dict(model_state_dict, strict=info.strict)

PyTorch 2.6.0 的 OOM 问题正是出在场景 1 和 2 中的内存管理上。2.7.0 通过逐个张量广播(而非一次性全部加载)来解决。

4. 获取优化器 state_dict get_optimizer_state_dict()

def get_optimizer_state_dict(
    model: nn.Module,
    optimizers: Union[torch.optim.Optimizer, Iterable[torch.optim.Optimizer]],
    *,
    submodules: Optional[set[nn.Module]] = None,
    options: Optional[StateDictOptions] = None,
) -> OptimizerStateType:

优化器的 state_dict 比模型更复杂,因为它包含: - param_groups:参数组配置(学习率等) - state:每个参数的优化器状态(如 Adam 的 momentum、variance)

处理流程: 1. 获取 FSDP 优化器 state_dict(FSDP 有特殊的 optim_state_dict 方法) 2. 如果需要完整数据,收集所有分片 3. 将参数名从内部格式转换为用户可读格式 4. 可选展平为一级字典

5. FQN(Fully Qualified Name)处理

@functools.cache
def _get_fqns(model, name, dsd_fqn_modifiers="_fqn_modifiers",
              skip_ddp_prefix=True, skip_compiler_prefix=True):

在分布式训练中,同一个参数可能有多个名称: - 原始名称:layers.0.weight - DDP 包装后:module.layers.0.weight - FSDP 展平后:_flat_param(映射到多个原始参数) - torch.compile 后:_orig_mod.layers.0.weight

_get_fqns 函数负责将内部名称转换回用户可读的原始名称(FQN = Fully Qualified Name),并处理各种前缀。

6. FSDP state_dict 上下文管理

def _get_fsdp_state_dict_context(info, module):
    if info.full_state_dict:
        state_dict_config = FullStateDictConfig(
            offload_to_cpu=info.cpu_offload,
            rank0_only=info.cpu_offload,
        )
        return FSDP.state_dict_type(module, StateDictType.FULL_STATE_DICT,
                                      state_dict_config)
    else:
        return FSDP.state_dict_type(module, StateDictType.SHARDED_STATE_DICT,
                                      ShardedStateDictConfig())

FSDP 需要在特定上下文中获取 state_dict。FullStateDictConfig 配置了是否 offload 到 CPU 以及是否只在 rank 0 上保留完整数据。

核心 API 列表

名称 类型 说明
StateDictOptions 数据类 控制 state_dict 操作行为的选项
get_model_state_dict() 函数 获取模型的 state_dict
set_model_state_dict() 函数 设置模型的 state_dict(OOM 修复重点)
get_optimizer_state_dict() 函数 获取优化器的 state_dict
set_optimizer_state_dict() 函数 设置优化器的 state_dict
get_state_dict() 函数 同时获取模型和优化器的 state_dict
set_state_dict() 函数 同时设置模型和优化器的 state_dict
_get_fqns() 内部函数 参数名转换(内部名 → FQN)
_StateDictInfo 内部数据类 扩展的选项,包含 FSDP 模块信息

与其他模块的关系

  • _state_dict_utils.py:导入并使用其底层工具函数(_gather_state_dict、_broadcast_state_dict 等)
  • verl/trainer/:FSDP 训练时使用 get_model_state_dict / set_model_state_dict 保存和加载 checkpoint
  • PyTorch FSDP:深度集成 FSDP 的 state_dict API
  • PyTorch DDP:处理 DDP wrapper 带来的 module. 前缀

OOM 问题说明

PyTorch 2.6.0 中 set_model_state_dict 的 OOM 问题主要源于:

  1. 一次性加载所有参数:旧版本在广播时会同时在 GPU 上存储完整的 state_dict 和本地分片,导致显存翻倍
  2. 缺少及时释放:广播完成后没有及时释放完整张量

PyTorch 2.7.0 的修复策略: - 逐个张量广播,每个张量广播后立即释放 - 更好的内存管理,避免不必要的副本

verl 将修复后的代码复制到项目中,确保无论用户使用哪个 PyTorch 版本都能正确工作。

小结

checkpoint/state_dict.py 是 PyTorch 分布式 state_dict 管理的核心文件。它封装了 FSDP、DDP 等分布式策略的 state_dict 操作复杂性,提供了简洁统一的 API。verl 复制这个文件主要是为了修复 PyTorch 2.6.0 中的 OOM 问题,确保大模型训练的稳定性。

这个文件的代码量很大且涉及很多 PyTorch 内部 API,对于初学者来说不需要完全理解每一行代码。重要的是理解它提供的核心 API(get_model_state_dict / set_model_state_dict)的作用和使用场景。