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 得到空字典,避免 OOMignore_frozen_params:是否忽略冻结参数(requires_grad=False)strict:加载时是否严格匹配 keybroadcast_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]:
内部流程(简化版):
- 处理 FSDP 上下文:如果模型使用了 FSDP,需要设置正确的 StateDictType
full_state_dict=True→ 使用StateDictType.FULL_STATE_DICT-
full_state_dict=False→ 使用StateDictType.SHARDED_STATE_DICT -
调用
model.state_dict():在正确的上下文中获取 state_dict -
后处理:
- 如果需要
full_state_dict,调用_gather_state_dict收集所有分片 - 如果需要
cpu_offload,将数据移到 CPU - 过滤子模块前缀、冻结参数等
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 都有完整数据)
场景 3:分片 state_dict
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 问题主要源于:
- 一次性加载所有参数:旧版本在广播时会同时在 GPU 上存储完整的 state_dict 和本地分片,导致显存翻倍
- 缺少及时释放:广播完成后没有及时释放完整张量
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)的作用和使用场景。