Prhub

#7283 [fsdp,veomni] fix: backfill missing state for DSD optimizer checkpoint

原始 PR 作者 wuxibin89 合并时间 2026-08-06 18:05 文件变更 1 提交数 1 评论 0 代码增减 +35 / -0

执行摘要

回填扁平化 DSD 检查点缺失优化器状态,修复恢复崩溃

PR body 说明:扁平化 DSD 检查点只持久化已拥有优化器状态的参数(即收到过梯度的参数),从未收到梯度的参数——如 DeepSeek-V4 稀疏注意力 indexer,其 forward 只返回 top-k 索引——在 torch unflatten 时被变成空 dict 而不是跳过,最终在 Adam.__setstate__ 中失败。回填新初始化优化器的状态正是模拟 DCP allow_partial_load 的行为:这些参数以默认状态恢复。

建议精读。改动小但处于训练恢复核心路径,设计上巧妙复用 optimizer.state_dict() 的初值来对齐 allow_partial_load 语义。值得关注的是缺少单元测试覆盖,可参考 tests/utils/ckpt/test_lora_custom_object_save_on_cpu.py 的写法,补一个 CPU 上的扁平化 state dict 回填校验。

讨论亮点

该 PR 无 review 评论与讨论线程。设计权衡体现在方法 docstring 中:用刚初始化的优化器补齐缺失条目,刻意对齐 DCP allow_partial_load 的语义,避免在恢复阶段为这些参数临时构造默认状态;同时用 state / param_groups 键区分格式,确保普通优化器检查点路径零行为变化。

实现拆解

  1. 背景与入口:该修复落在 verl/utils/checkpoint/fsdp_checkpoint_manager.pyFSDPCheckpointManager.load_checkpoint 优化器加载分支。此前该分支直接 torch.load 后调用 self.optimizer.load_state_dict(optimizer_state_dict),遇到扁平化 DSD 格式会崩溃。
  2. 新增回填方法:新增 _backfill_optimizer_state。首先通过顶层键 state / param_groups 区分普通优化器 state dict 与扁平化 DSD 格式:普通格式直接透传,因为 torch 本身容忍无状态参数;扁平化格式则取当前新初始化优化器的 state_dict() 作为基准,计算检查点中缺失的键。
  3. 合并与日志:当存在缺失键时,返回 {**current_state_dict,**state_dict},检查点中已有的状态优先,缺失条目保留初始值,并通过 log_with_rank 在 rank 0 打印缺失数量和示例键,便于排查。
  4. 调用点接线:在 load_checkpoint 的优化器分支中,torch.load 之后、load_state_dict 之前插入 optimizer_state_dict = self._backfill_optimizer_state(optimizer_state_dict)
  5. 配套改动:无测试文件、无配置变更,纯增量 35 行。建议后续在 tests/utils/ckpt/ 下补充 CPU 上的扁平化 state dict 回填测试。
文件 模块 状态 重要度
verl/utils/checkpoint/fsdp_checkpoint_manager.py 检查点 modified 7.01

关键符号

_backfill_optimizer_state

关键源码片段

verl/utils/checkpoint/fsdp_checkpoint_manager.py core-logic

唯一变更文件,新增 `_backfill_optimizer_state` 并在 `load_checkpoint` 优化器加载分支调用,修复扁平化 DSD 检查点恢复崩溃。

    def _backfill_optimizer_state(self, state_dict: dict) -> dict:
        """回填扁平化 DSD 优化器检查点中缺失的状态条目。        VeOmni 的 ``MultiOptimizer`` 走 ``torch.distributed.checkpoint.state_dict``
        且 ``flatten_optimizer_state_dict=True`` 保存时,只持久化收到过梯度的参数;
        从未收到梯度的参数(如 DeepSeek-V4 稀疏注意力 indexer)在 unflatten 后会被
        torch 变成空 dict,触发 ``Adam.__setstate__`` 报错。这里用新初始化优化器的
        状态补齐缺失条目,等价于 DCP ``allow_partial_load``:缺失参数从默认状态恢复。
        该方法由 ``load_checkpoint`` 在优化器加载分支中调用。
        """
        # 普通优化器 state dict(含 "state" / "param_groups" 顶层键):
        # torch 本身就容忍无状态参数,直接透传,无需回填。
        if "state" in state_dict or "param_groups" in state_dict:
            return state_dict
​
        # 扁平化格式:以当前(刚初始化)优化器的 state_dict 为基准,
        # 它的键集合代表所有可训练参数。
        current_state_dict = self.optimizer.state_dict()
        # 防御:若当前优化器返回普通格式,说明两侧结构不一致,保持原样。
        if "state" in current_state_dict or "param_groups" in current_state_dict:
            return state_dict
​
        # 检查点中缺失的键:这些参数从未收到梯度,需要回填初始状态。
        missing_keys = [key for key in current_state_dict if key not in state_dict]
        if not missing_keys:
            return state_dict
​
        log_with_rank(
            f"{len(missing_keys)} optimizer state entries are absent from the checkpoint "
            f"and keep their initial value, e.g. {missing_keys[:5]}",
            rank=self.rank,
            logger=logger,
            log_only_rank_0=True,
        )
        # 合并:检查点已有的状态优先(后覆盖),缺失条目自动退回初始值。
        return {**current_state_dict, **state_dict}
        if self.should_load_optimizer:
            remote_optim_path = os.path.join(local_path, f"optim_world_size_{self.world_size}_rank_{self.rank}.pt")
            local_optim_path = copy_to_local(remote_optim_path)
            optimizer_state_dict = torch.load(local_optim_path, weights_only=False)
            # 关键:在 load_state_dict 之前回填扁平化 DSD 检查点缺失的优化器状态,
            # 避免 unflatten 空 dict 导致 Adam.__setstate__ 崩溃。
            optimizer_state_dict = self._backfill_optimizer_state(optimizer_state_dict)
            self.optimizer.load_state_dict(optimizer_state_dict)

评论区精华

没有提炼出高价值讨论线程

当前评论区没有形成足够清晰的争议点或结论,后续有更多讨论时会体现在这里。

风险与影响

  1. 缺少测试覆盖:改动位于分布式恢复核心路径,但无对应测试文件;若 DSD 扁平化键空间与当前优化器 state_dict() 不一致,{**current_state_dict,**state_dict} 可能引入多余或冲突的键。
  2. 静默恢复风险:从未收到梯度的参数会被静默用初始状态恢复;若某个参数本应有状态却因存储问题缺失,该逻辑会掩盖真实数据问题。
  3. 兼容性:防御分支仅检查顶层键 state / param_groups,若后续 torch DCP 版本改变扁平化格式的键名或结构,missing_keys 计算可能失效。

影响范围集中在 FSDP 与 VeOmni 引擎的优化器检查点恢复路径,受益最大的场景是 DeepSeek-V4 这类含 sparse-attention indexer(前向只返回 top-k 索引、从不产生梯度)的模型训练恢复;对普通(非扁平化)优化器检查点无行为变化,因为回填前已短路。团队侧降低了 VeOmni 扁平化保存方案的恢复门槛,也为未来其他零梯度参数模型(如稀疏专家、固定 embedding)提供了默认恢复语义。

缺少测试覆盖 合并依赖两侧键空间一致 静默回填可能掩盖数据缺失

关联 Issue

未识别关联 Issue

当前没有检测到明确关联的 Issue 链接,后续同步到相关引用后会出现在这里。

完整报告

参与讨论