Skip to content

[Fix] Support sparse optimizer DCP resume - #2018

Closed
matrix72c wants to merge 1 commit into
InternLM:mainfrom
matrix72c:matrix72c/fix-lazy-optimizer-dcp
Closed

[Fix] Support sparse optimizer DCP resume#2018
matrix72c wants to merge 1 commit into
InternLM:mainfrom
matrix72c:matrix72c/fix-lazy-optimizer-dcp

Conversation

@matrix72c

Copy link
Copy Markdown

Motivation

PyTorch optimizers initialize per-parameter state lazily. A valid DCP checkpoint can therefore omit all optimizer-state leaves for a trainable parameter that did not receive a gradient before save. During resume, XTuner materializes a full optimizer load target, and DCP reports those absent lazy leaves as missing checkpoint keys.

Changes

  • Inspect DCP metadata before loading optimizer state and remove only parameter states whose leaves are wholly absent from the checkpoint.
  • Keep partially missing optimizer states fatal so corrupted AdamW state is not silently accepted.
  • Preserve model strictness while loading sparse optimizer state with the PyTorch 2.10+ strict=False behavior.
  • Add a regression test covering AdamW DCP resume, later lazy-state materialization, and rejection of partially saved state.

Impact

Model checkpoint validation remains unchanged. Existing optimizer states are restored normally; only wholly absent lazy per-parameter state is deferred until that parameter next receives a gradient.

Tests

  • XTuner pre-commit hooks passed for the changed files, except mypy (the hook environment exhausted local temporary storage while installing its pinned PyTorch 2.6 CUDA dependencies).
  • ruff check xtuner/v1/engine/train_engine.py tests/engine/test_train_engine_checkpoint.py
  • python3.12 -m py_compile xtuner/v1/engine/train_engine.py tests/engine/test_train_engine_checkpoint.py
  • PyTorch 2.10 DCP save/load/optimizer-step reproduction passed.

@YanhuiDua

Copy link
Copy Markdown
Collaborator

@claude review

@claude

claude Bot commented Aug 13, 2026

Copy link
Copy Markdown
Contributor

Claude: 代码 review(1/3:Summary 与流程图)

Summary

本 PR 修复「DCP resume 时,save 前未收到梯度的 trainable param 因 optimizer lazy state 缺失而导致 dcp.load 报 missing keys」的问题。做法是:load 前读取 checkpoint metadata,把整组 state 都不在 checkpoint 里的 param 从 load target 中摘掉(部分缺失仍视为致命错误),并把 optimizer 的 set_optimizer_state_dict 放宽为 strict=False;同时把原来 model/optimizer 共用的 set_options 拆成两套。

方向是对的,问题定位也准确,但有 1 个阻塞性问题strict 对 optimizer 侧生效是 PyTorch 2.10 才引入的行为,而本仓库声明支持 torch>=2.6.0 且明确特化过 2.7/2.8/2.9。在 2.6~2.9 上,被 pop 掉的 param 会让 _split_optim_state_dictKeyError,也就是 bug 没修好、报错信息还更难定位;而新增用例被 skipif(torch<2.10) 跳过,低版本 CI 发现不了。此外新增的回归测试没有走 TrainEngine.load_dcp,而是把它的实现在测试里复制了一遍,因此锁不住真正的修复路径。

ProduceBatchResult impact: not affected(本次 diff 不在 xtuner/v1/rl 下;但 xtuner/v1/rl/trainer/worker.py:1126 会调用 TrainEngine.load_dcp,因此 RL resume 同样受这次 optimizer 校验放宽的影响。)

Verdict: REQUEST_CHANGES(阻塞项:torch < 2.10 兼容性;回归测试需穿过 TrainEngine.load_dcp


Main Flowchart before this PR

flowchart TD
    subgraph CALLER["调用方"]
        T1["xtuner.v1.train.trainer.Trainer._load_checkpoint()"]
        R1["xtuner.v1.rl.trainer.worker.TrainingWorker.load()"]
    end

    subgraph ENGINE["xtuner.v1.engine.train_engine.TrainEngine"]
        E1["TrainEngine.load_dcp(weights_dir, load_states, load_args)"]
        E2["TrainEngine._get_dcp_state_dict(cpu_offload, save_optimizer)"]
        E3["TrainEngine.has_freeze_params"]
        E4["set_options = StateDictOptions(cpu_offload=True, strict = not has_freeze_params)<br/>model 与 optimizer 共用一套"]
    end

    subgraph TORCH["torch.distributed.checkpoint"]
        G1["state_dict.get_optimizer_state_dict()<br/>_init_optim_state() 为每个 trainable param<br/>materialize step/exp_avg/exp_avg_sq"]
        D1["dcp.load(state_dict, checkpoint_id)<br/>DefaultLoadPlanner strict=True"]
        D2["state_dict.set_model_state_dict(model, sd, options)"]
        D3["state_dict.set_optimizer_state_dict(model, optim, optim_state_dict, options)"]
    end

    T1 --> E1
    R1 --> E1
    E1 --> E2 --> G1
    E1 --> E3 --> E4
    E1 --> D1
    D1 -- "checkpoint 缺 optimizer.state.FQN.exp_avg" --> BAD["RuntimeError: missing keys<br/>resume 直接失败"]
    E1 --> D2
    E1 --> D3
    E4 -.-> D2
    E4 -.-> D3
Loading

关键点_get_dcp_state_dict() 通过 get_optimizer_state_dict() 得到的 load target 是稠密的(_init_optim_state() 会用零梯度跑一次 step,为所有 trainable param 建好 state),而 checkpoint 可能是稀疏的(save 时某些 param 还没拿到梯度)。DCP 以 load target 为准,于是稀疏 checkpoint 加稠密 load target 就等于 missing keys。

Main Flowchart after this PR

flowchart TD
    subgraph CALLER["调用方(未变)"]
        T1["Trainer._load_checkpoint() / TrainingWorker.load()"]
    end

    subgraph ENGINE["xtuner.v1.engine.train_engine"]
        E1["TrainEngine.load_dcp(weights_dir, load_states, load_args)"]
        E2["TrainEngine._get_dcp_state_dict(cpu_offload, save_optimizer)"]
        NEW1["NEW: dcp.FileSystemReader(weights_dir).read_metadata()<br/>saved_keys: set[str]"]
        NEW2["NEW: _prune_uncheckpointed_optimizer_state(optimizer_state_dict, saved_keys) -> list[str]"]
        NEW3["NEW: model_options = StateDictOptions(strict = not has_freeze_params)<br/>NEW: optimizer 侧无条件 StateDictOptions(strict=False)"]
    end

    subgraph TORCH["torch.distributed.checkpoint"]
        D1["dcp.load(state_dict, checkpoint_id)"]
        D2["set_model_state_dict(...) strictness 与改动前等价"]
        D3["set_optimizer_state_dict(...)"]
        D4["_split_optim_state_dict()<br/>2.10+: fqn 缺失且 strict=False 则跳过<br/>2.6~2.9: 无该分支,裸 KeyError"]
    end

    T1 --> E1 --> E2
    E1 --> NEW1 --> NEW2
    NEW2 -- "全组缺失" --> OK1["state.pop(param) 并 logger.warning(examples=...)"]
    NEW2 -- "部分缺失" --> BAD1["RuntimeError: Incomplete optimizer state<br/>显式失败,保留严格性"]
    NEW2 --> D1 --> D2
    D1 --> D3 --> D4
    NEW3 -.-> D2
    NEW3 -.-> D3
    D4 -- "torch < 2.10" --> BAD2["KeyError: FQN(未覆盖的版本)"]
    D4 -- "torch >= 2.10" --> OK2["resume 成功,param 下次拿到梯度时重新 materialize"]

    classDef changed fill:#fff3b0,stroke:#d4a017
    classDef broken fill:#ffd6d6,stroke:#c92a2a
    class NEW1,NEW2,NEW3 changed
    class BAD2 broken
Loading

文字说明(突出改动)

  1. 新增裁剪步骤:在 dcp.load 之前插入 metadata 比对,用扁平 key optimizer.state.<fqn>.<state_name> 判断 checkpoint 中是否存在;整组缺失则从 load target 摘除,部分缺失则抛错。
  2. strictness 拆分:model 侧保持原语义(strict = not has_freeze_params),optimizer 侧改为无条件 strict=False
  3. 未变:model 权重校验、load_states=False / load_args=False 的分支语义、save 路径。

@claude

claude Bot commented Aug 13, 2026

Copy link
Copy Markdown
Contributor

Claude: 代码 review(2/3:核心原理实现与单测、抽象与信息隐藏、线性业务流程)

核心原理实现与单测

核心实现是 xtuner/v1/engine/train_engine.py:62-99_prune_uncheckpointed_optimizer_state(),原理链条清晰、成立:

  • torch.optim.AdamW 的 per-parameter state 是惰性创建的(第一次 step() 才建 step/exp_avg/exp_avg_sq),没拿到梯度的 param 在 save 时不会出现在 get_optimizer_state_dict() 的输出里;
  • 而 resume 时 XTuner 用 get_optimizer_state_dict() 构造 load target,torch 内部的 _init_optim_state() 会先用零梯度跑一次 step,把所有 trainable param 的 state 都建出来,于是 load target 恒为稠密;
  • dcp.loadDefaultLoadPlanner 默认 strict=True(即 allow_partial_load=False),load target 里 metadata 中不存在的 fqn 一律报 missing key,resume 失败。

裁剪用的 key 拼接之所以能对上,依赖两个约定:_get_dcp_state_dict() 把 optimizer 放在顶层 key "optimizer" 下(train_engine.py:340),以及 DCP flatten_state_dict=True. 连接路径(本仓库 XtunerCacheSavePlanner 默认开启,xtuner/v1/patch/xtuner_cache_planner.py:29)。两个约定目前都成立,但都没有在 helper 的 Interface 上表达出来(见 inline 评论)。

「部分缺失 fatal」的判断也有道理:AdamW 无法只用 exp_avg 不用 exp_avg_sq 更新,静默接受会导致数值错误但不报错的训练。这个保证只在 key 粒度成立:shard 粒度的部分缺失(部分 rank 有 state、部分没有,metadata 里 key 存在但 chunk 覆盖不全)会被静默零填充,尤其在 skip_checkpoint_validation 已把 _validate_global_plan patch 成 no-op 的情况下(xtuner/v1/patch/torch_dcp_planner.py:19-24,inline 评论已详述)。

对应单测 tests/engine/test_train_engine_checkpoint.py

  • test_sparse_adamw_dcp_state_can_resume_and_materialize_later:真实 dcp.save/load 往返,断言 (a) 只有 bias 的三个 leaf 被裁剪,(b) weightexp_avg 被完整还原,(c) bias 下一步拿到梯度后重新 materialize。用真实 metadata 派生 saved_keys 是这个用例最有价值的地方,它把 DCP 扁平 key 格式这一契约钉住了。
  • test_partially_saved_optimizer_state_is_rejected:用合成 dict 验证部分缺失会抛 RuntimeError

覆盖缺口(重要):两个用例都没有调用 TrainEngine.load_dcp,第一个用例还把 load_dcp 的实现复制了一遍(连 StateDictOptions 都与生产代码不一致),因此 metadata 读取位置、prune 调用时机、model/optimizer 两套 options 的拆分一行都没被执行。按 .claude/CLAUDE.md「bug fix PR 必须包含复现原 bug 的回归测试」,这条还没真正满足。


抽象与信息隐藏评估

维度 评价
Depth 偏 Shallow。_prune_uncheckpointed_optimizer_state(dict, set[str]) -> list[str] 要求调用者自己知道:怎么拿 metadata、怎么把 metadata key 转成字符串集合、返回值代表什么、以及它会原地修改入参。Interface 的信息量和 Implementation 差不多。
Deletion test 内联回 load_dcp 后复杂度不会散落(只有一个调用者),说明它不是靠信息隐藏立足的 Module,而是「为了能单测而抽出的纯函数」。
信息隐藏 有泄漏:optimizer. 前缀由 _get_dcp_state_dict() 决定,却硬编码在 helper 内部;「saved_keys 必须是 DCP 扁平化后的 key」这一前置条件没有在签名或 docstring 中表达。
Locality 一条规则拆到三处:顶层 key 命名在 _get_dcp_state_dict()、key 拼接在 helper、metadata 获取与 strict 决策在 load_dcp()。改 key 布局要同时改三处,且失配时是静默降级。
测试面 当前测试穿过的 Seam(私有 helper)与调用者穿过的 Seam(load_dcp)不是同一个,正是附录 A 第 3 条要避免的情形。

改进意见:

  • Filesxtuner/v1/engine/train_engine.pyTrainEngine.load_dcp_prune_uncheckpointed_optimizer_state
  • Problem:为了可测试性抽出的纯函数丢掉了真实调用顺序(metadata 读取 -> 裁剪 -> load -> 两套 options)的 Locality,同时把顶层 key 命名这一调用方知识泄漏进 helper。
  • Solution:把它变成 TrainEngine 的私有方法,由它自己拥有 metadata 读取、前缀、裁剪、rank0 日志与版本分支;纯 key 判定逻辑若仍需单测,则把前缀作为显式参数传入:
def _drop_uncheckpointed_optimizer_state(self, state_dict: dict[str, Any], reader: dcp.FileSystemReader) -> bool:
    saved_keys = {str(k) for k in reader.read_metadata().state_dict_metadata}
    removed = _prune_uncheckpointed_optimizer_state(state_dict["optimizer"], saved_keys, key_prefix="optimizer")
    if removed:
        log_rank0.warning(f"Ignoring {len(removed)} lazy optimizer-state leaves; examples={removed[:5]}")
    return bool(removed)
  • Benefitsload_dcp 主干恢复成单一抽象层级的线性叙事;key 布局知识收敛到一处;测试可以穿过 load_dcp 这个真正的 Seam。

公开 Interface 的线性业务流程评估

TrainEngine.load_dcp() 是本次唯一被改动的公开 Interface。改动后 with profile_time_and_memory(...) 块内混进了三个抽象层级(存储层机制、数据搬运、业务意图、日志格式化):

with profile_time_and_memory(f"[Load DCP from {weights_dir}]"):
    if load_optimizer:
        metadata = dcp.FileSystemReader(weights_dir).read_metadata()     # 存储层机制
        saved_keys = {str(key) for key in metadata.state_dict_metadata}  # 数据搬运
        removed = _prune_uncheckpointed_optimizer_state(...)             # 业务意图
        if removed:
            logger.warning(f"...examples={removed[:5]}")                 # 日志格式化
    dcp.load(...)

建议收敛为业务意图级别的叙事(机制下沉到上面的私有方法):

with profile_time_and_memory(f"[Load DCP from {weights_dir}]"):
    if load_optimizer:
        self._drop_uncheckpointed_optimizer_state(state_dict, weights_dir)
    dcp.load(state_dict=state_dict, checkpoint_id=weights_dir)
    set_model_state_dict(self.model, state_dict["model"], options=model_options)
    if not load_optimizer:
        return
    self._load_optimizer_state(state_dict["optimizer"], load_states=load_states, load_args=load_args)

顺带一提:load_dcp 后半段处理 load_states=False / load_args=False 的 param_groups 拼装(train_engine.py:558-578)本来就是主干里最低层的机制细节,本 PR 未触及,属于可选清理。另外把 if/else 化简为 model_options = StateDictOptions(cpu_offload=True, strict=not self.has_freeze_params) 是好的改进。

@claude

claude Bot commented Aug 13, 2026

Copy link
Copy Markdown
Contributor

Claude: 代码 review(3/3:单测建议、其他 issues、Verdict)

单测建议

保留:

  • 「真实 dcp.save -> 读 metadata -> 裁剪 -> dcp.load -> 后续 materialize」这条主链路。它测的是真实代码路径,只依赖 torch 这个项目外依赖,是 Good Test 的骨架。

需要新增(优先级最高):

  • 一个穿过 TrainEngine.load_dcp() 的回归用例(单卡 gloo 即可,tests/engine/test_moe_train_engine.py:321 有现成范式)。这是本 PR 唯一真正缺失的测试。
  • torch < 2.10 的行为用例:期望是「明确的、带上下文的错误」而不是裸 KeyError(需配合生产代码补版本分支)。
  • 「全部 param 都被裁剪」的用例:这种情况几乎一定是 key 格式失配,应该报错而不是静默重置 optimizer。

建议精简与重构:

  • test_partially_saved_optimizer_state_is_rejected 当前用手写的合成 dict 与手写的 saved_keys,属于 Bad Test 倾向(测的是 helper 的内部约定,把 key 命名规则在测试里复制了一份)。建议改为从真实 metadata 派生 saved_keys 后主动 discard 掉一个 key 来模拟部分缺失。
  • 第一个用例一次断言了三件事(裁剪集合、还原值、后续 materialize),建议拆成三个用例,失败时定位更快。
  • 没有冗余测试、也没有过于简单的测试,这方面没问题。主要问题是组织格式不符合 .claude/CLAUDE.md:需要用 Test<Feature> 分组(TestClass docstring 说明测试类别)、补文件级两级 docstring(一级 TestClass、二级 test_func)、每个用例首行加一行中文行为注释。inline 评论里给了可直接套用的骨架。

其他 issues

  • Critical:torch 2.6~2.9 上裁剪会导致裸 KeyError。证据:_split_optim_state_dict 在 v2.6.0 / v2.9.0 中该分支只有 state[fqn] = optim_state_dict[_STATE][fqn] 两行,无 fqn in ... 判断也不读 info.strictelif info.strict: raise RuntimeError("Missing optimizer state for parameter ...") 是 torch main(2.10+)才加的。仓库声明 torch>=2.6.0pyproject.toml:40requirements/runtime.txt:12)且活跃支持 2.7/2.8/2.9(xtuner/v1/train/trainer.py:595xtuner/v1/patch/torch_shape_env_simplify_pt28.pyxtuner/v1/model/compose/qwen3_vl/modeling_qwen3_vl.py:21)。
  • Warningoptimizer. 前缀硬编码在 helper 内,失配时会把全部 param state 静默清空(只留一条 warning)。建议把前缀参数化,并加「全量裁剪」健全性检查。
  • Warning:optimizer 侧无条件 strict=False 关掉了 2.10 新增的唯一一条 optimizer 侧校验。建议按意图收紧,例如 strict = not removed and load_states
  • Warning:key 存在性判断无法覆盖 DTensor shard 粒度的部分缺失,PR 描述中「partial 一律 fatal」的保证在 FSDP/EP 场景下比看起来弱。
  • Nitlogger.warning 每个 rank 都会打印,本文件已有 log_rank0.warning 的既有写法(train_engine.py:111:280)。
  • Nit:额外多读了一次 .metadatadcp.load 内部还会再读),且把 storage backend 写死为 FileSystemReader。建议构造一次 reader 并同时传给 dcp.load(..., storage_reader=reader)
  • Nit:私有函数位置偏前(CLAUDE.md 要求公开接口靠前、私有接口靠后);docstring 缺 Args: / Returns:,也未说明「原地修改入参」这一副作用。
  • 安全性:无安全影响。日志只打印 optimizer state 的 key 名(最多 5 个),不含权重数据;未引入新的反序列化面(dcp.load 前的 add_safe_globals 逻辑未改动)。
  • load_dcp 的 docstring 声称「若 checkpoint 不含 optimizer state,则只加载模型权重」,代码里并没有这个实现,本 PR 的裁剪逻辑容易被误读为已经支持。补充确认:这种情况下 optimizer.param_groups 的 key 也不存在,dcp.load 仍会显式报错,所以不会静默降级;但建议顺手把这句过期 docstring 修掉或落实。

Verdict

REQUEST_CHANGES

阻塞项:(1) torch < 2.10 的版本兼容(当前会退化成裸 KeyError,且被 skipif 的单测掩盖);(2) 回归测试需要穿过 TrainEngine.load_dcp。其余为抽象收敛、strictness 收紧与测试组织格式建议。整体思路正确,补上版本分支和 engine 级回归测试后即可合并。

Comment on lines +544 to +547
if load_optimizer:
metadata = dcp.FileSystemReader(weights_dir).read_metadata()
saved_keys = {str(key) for key in metadata.state_dict_metadata}
removed = _prune_uncheckpointed_optimizer_state(state_dict["optimizer"], saved_keys)

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Claude: Critical — 该修复在 torch < 2.10 上会把原来的报错换成更难定位的 KeyError

本仓库声明支持 torch>=2.6.0pyproject.toml#L40requirements/runtime.txt:12),并且代码里明确对 2.7/2.8/2.9 做过特化(xtuner/v1/train/trainer.py:595torch.__version__.startswith("2.7.")xtuner/v1/patch/torch_shape_env_simplify_pt28.pyxtuner/v1/patch/xtuner_storage.py:29),说明 2.6~2.9 是活跃目标版本。

strict 对 optimizer 侧生效是 2.10 才有的行为

  • torch v2.9.0 / v2.6.0torch/distributed/checkpoint/state_dict.py::_split_optim_state_dict 中该分支只有两行,没有任何 fqn in ... 判断、也不读 info.strict
    if param.requires_grad:
        state[fqn] = cast(DictValueType, optim_state_dict[_STATE])[fqn]
  • torch main(2.10+) 才加上:
    if fqn in cast(DictValueType, optim_state_dict[_STATE]):
        state[fqn] = ...
    elif info.strict:
        raise RuntimeError("Missing optimizer state for parameter '{fqn}' in checkpoint. ...")

后果:在 2.6~2.9 上,这里 pop 掉的 param 会让下面的 set_optimizer_state_dict 抛出一个KeyError: '<fqn>'strict=False 在这些版本对 optimizer 完全无效)。也就是说,在这些版本上 bug 不但没修好,报错信息还从 DCP 的 “missing checkpoint keys” 退化成了无上下文的 KeyError;而新增单测被 skipif(torch<2.10) 跳过,CI 若跑 2.7/2.8 也发现不了。

建议按仓库既有习惯(Version(torch.__version__) >= Version(...))显式分版本,低版本走 “保留已 materialize 的零状态、只跳过缺失的读取” 这条 2.6 就可用的路径:

_TORCH_SUPPORTS_SPARSE_OPTIM_LOAD = Version(torch.__version__) >= Version("2.10")

if load_optimizer:
    if _TORCH_SUPPORTS_SPARSE_OPTIM_LOAD:
        # 2.10+: 可以从 load target 里删除,set_optimizer_state_dict(strict=False) 会跳过
        removed = self._prune_uncheckpointed_optimizer_state(state_dict, weights_dir)
        optim_options = StateDictOptions(cpu_offload=True, strict=False)
    else:
        # 2.6~2.9: _split_optim_state_dict 无 strict 分支,删除会 KeyError。
        # 保留 _init_optim_state 产生的零状态,只让 DCP 跳过 checkpoint 中不存在的 key。
        load_planner = DefaultLoadPlanner(allow_partial_load=True)
        optim_options = StateDictOptions(cpu_offload=True, strict=True)

DefaultLoadPlanner(allow_partial_load=True) 自 2.2 起就有,create_default_local_load_plan(..., strict=not allow_partial_load) 会直接跳过 metadata 中不存在的 fqn。若不想支持低版本,至少应在 load_optimizer 分支入口对 torch 版本做断言并给出明确的错误信息,而不是让它 KeyError。)

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

这个问题需要修复

Comment on lines +85 to +93
checkpoint_keys = {
state_name: f"optimizer.state.{parameter_name}.{state_name}" for state_name in parameter_state
}
present_state_names = {name for name, key in checkpoint_keys.items() if key in saved_keys}

if not present_state_names:
state.pop(parameter_name)
removed.extend(checkpoint_keys.values())
continue

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Claude: Warning — 硬编码的 "optimizer." 前缀把调用方的命名约定泄漏进了 helper,且失配时会「静默丢弃全部 optimizer state」

optimizer.state.{param}.{state_name} 这个扁平 key 由两个约定共同决定:

  1. _get_dcp_state_dict()xtuner/v1/engine/train_engine.py:340)把 optimizer state 放在顶层 key "optimizer" 下;
  2. DCP flatten_state_dict=True"." 连接路径(本仓库 XtunerCacheSavePlanner 默认 flatten_state_dict=Truextuner/v1/patch/xtuner_cache_planner.py:29,所以格式一致)。

但 helper 只拿到 optimizer_state_dict 这一个 dict,前缀却写死在函数体里。这形成一个隐性耦合:一旦顶层 key 改名、或 torch 改了 flatten 规则、或将来打开 flatten_optimizer_state_dictpresent_state_names 会对所有 param 都为空 → 走到 state.pop(...) 分支 → 全部 param state 被清掉,只留一条 warning;再叠加下面 set_optimizer_state_dict(strict=False),整个 optimizer 会被静默重置为初始状态继续训练,metric 上很难察觉。

两点建议:

  1. 把前缀变成 Interface 的一部分(或让 helper 同时接收整个 load target 与 metadata,见另一条关于线性流程的评论):
    def _prune_uncheckpointed_optimizer_state(
        optimizer_state_dict: dict[str, Any], saved_keys: set[str], *, key_prefix: str = "optimizer"
    ) -> list[str]:
        ...
        checkpoint_keys = {name: f"{key_prefix}.state.{parameter_name}.{name}" for name in parameter_state}
  2. 增加一个「全量裁剪」的健全性检查——所有 param 的 lazy state 同时缺失,几乎一定是 key 格式失配而不是真的都没拿到梯度:
    if pruned_parameters and not remaining_parameters:
        raise RuntimeError(
            f"No optimizer state matched checkpoint keys for any of {len(pruned_parameters)} parameters; "
            "the checkpoint layout is likely incompatible rather than lazily initialized."
        )

Comment on lines +95 to +97
if len(present_state_names) != len(checkpoint_keys):
missing = [key for name, key in checkpoint_keys.items() if name not in present_state_names]
raise RuntimeError(f"Incomplete optimizer state for parameter '{parameter_name}': missing {missing}")

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Claude: Warning — 「partial 一律 fatal」只在 key 粒度成立,shard 粒度的部分缺失仍会被静默零填充

这里的判定依据是 metadata 里 key 是否存在,而 metadata 是全局的(所有 rank 读到同一份)。问题在于分布式场景下 optimizer state 的缺失可能发生在 DTensor shard 粒度而不是 key 粒度:

若保存时 rank A 对某个 param 有 state、rank B 没有(例如 EP 下未被路由到的 expert、或某些 rank 侧 grad 为 None),get_optimizer_state_dict() 在两个 rank 上返回的 key 集合不同,DCP 的 global plan 取并集,metadata 里这个 key 是存在的,只是 chunk 覆盖不全。于是:

  • present_state_names 判定为「完整」,既不裁剪也不报错;
  • dcp.loadcreate_read_items_for_chunk_list 对没有 overlap 的区域直接产生 0 个 read item,rank B 的那部分 shard 保持 _init_optim_state 产生的零值,不报任何错
  • 原本能兜住这种不一致的 _validate_global_planskip_checkpoint_validation 时已被 patch 成 no-op(xtuner/v1/patch/torch_dcp_planner.py:19-24)。

所以 PR 描述里 “Keep partially missing optimizer states fatal so corrupted AdamW state is not silently accepted” 这一保证,在真实的 FSDP/EP 训练里比看起来要弱。麻烦确认一下触发本 issue 的真实场景(是纯 “param 全局都没拿到梯度”,还是 “部分 rank 没拿到梯度” 的 MoE/EP 情况):

  • 如果是前者,建议在注释/docstring 里明确写清「本函数只处理全局一致缺失,shard 级不一致不在覆盖范围内」,避免后续维护者误以为已经防住;
  • 如果是后者,仅靠 key 存在性判断无法修复,需要比对 metadata.state_dict_metadata[key].chunks 的覆盖范围,或者在 save 端就统一各 rank 的 state 集合(例如保存前对未初始化的 param 做一次全局 all-gather 对齐)。

Comment on lines +549 to +552
logger.warning(
f"Ignoring {len(removed)} lazy optimizer-state leaves absent from the checkpoint; "
f"examples={removed[:5]}"
)

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Claude: Nit(抽象层级 + 日志)

  1. 线性业务流程:load_dcp 的主干叙事应该是「准备 load target → 加载 → 写回 model → 写回 optimizer」,但这里把「构造 FileSystemReader、读 metadata、把 metadata key 转成字符串集合、格式化 warning」这些低一层的机制混进了主干,with profile_time_and_memory(...) 块内出现了 3 个抽象层级。建议下沉成一个私有方法,让主干保持线性;这个方法同时拥有「metadata 读取 + key 前缀 + 裁剪 + 日志」,正好把上一条评论提到的前缀知识收在一处:
with profile_time_and_memory(f"[Load DCP from {weights_dir}]"):
    if load_optimizer:
        self._drop_uncheckpointed_optimizer_state(state_dict, weights_dir)
    dcp.load(state_dict=state_dict, checkpoint_id=weights_dir)
    ...

def _drop_uncheckpointed_optimizer_state(self, state_dict: dict[str, Any], weights_dir: Path) -> None:
    saved_keys = {str(k) for k in dcp.FileSystemReader(weights_dir).read_metadata().state_dict_metadata}
    removed = _prune_uncheckpointed_optimizer_state(state_dict["optimizer"], saved_keys, key_prefix="optimizer")
    if removed:
        log_rank0.warning(...)
  1. logger.warning 会在每个 rank 都打印(千卡规模下就是上千行)。本文件已有 rank0-only 的既有写法(log_rank0.warning,见 xtuner/v1/engine/train_engine.py:111:280),这里建议改用 log_rank0.warning

  2. dcp.FileSystemReader(weights_dir).read_metadata() 额外读了一次 .metadatadcp.load 内部还会再读一次),且把 storage backend 写死成 FileSystemReader。建议构造一次 reader 并同时传给 dcp.load(..., storage_reader=reader),既省一次 IO 也避免两处对 backend 的假设不一致。

Comment on lines +583 to +585
# PyTorch 2.10+ treats wholly absent per-parameter optimizer
# states as valid lazy state when strict loading is disabled.
options=StateDictOptions(cpu_offload=True, strict=False),

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Claude: Warning — optimizer 侧无条件 strict=False 削弱了 model 之外的全部校验

改动本身有必要(load_states=Falseoptimizer_state_dict["state"] = {},在 torch 2.10 的 strict=True 下会对每个 trainable param 抛 “Missing optimizer state”),把 model / optimizer 的 options 拆开也比原来共用 set_options 更清晰。

但要注意:拆开后 optimizer 变成永远 strict=False,等于把 2.10 新增的那条唯一的 optimizer 侧校验完全关掉。结合上面裁剪逻辑一旦 key 前缀失配就会清空全部 state 的路径,两者叠加后「optimizer 被静默重置」不会有任何 error。建议按意图收紧,而不是一律关闭:

# 只有在确实裁剪掉了 lazy state(或显式不加载 states)时才放宽
optim_strict = not removed and load_states
options = StateDictOptions(cpu_offload=True, strict=optim_strict)

另外这两行注释说明的是 “2.10+ 的行为”,但代码没有任何版本分支(见我在 prune 调用处的 Critical 评论),注释与实际支持的 torch>=2.6 之间是有落差的,建议一并说明清楚。

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

这里是不是不要直接给默认参数为False

Comment on lines +1 to +17
from __future__ import annotations

import tempfile

import pytest
import torch
import torch.distributed.checkpoint as dcp
from packaging.version import Version
from torch.distributed.checkpoint.state_dict import (
StateDictOptions,
get_model_state_dict,
get_optimizer_state_dict,
set_model_state_dict,
set_optimizer_state_dict,
)

from xtuner.v1.engine.train_engine import _prune_uncheckpointed_optimizer_state

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Claude: Warning — 单测文件组织不符合 .claude/CLAUDE.md 的测试规范

规范要求:「Test functions must be named test_<description> and grouped in classes named Test<Feature>」,且本次 review 要求文件 docstring 按两级(TestClass → test_func)说明行为、每个 case 开头加一行中文注释说明测试哪种行为。当前文件是两个裸函数、无文件 docstring、无用例注释。建议重组为:

"""TrainEngine DCP resume 的稀疏 optimizer state 行为测试。

TestSparseOptimizerDcpResume —— 真实 DCP save/load 往返下的稀疏 optimizer state 行为
    test_absent_lazy_state_is_pruned_and_resume_succeeds —— 未拿到梯度的 param 其 state 全缺失时可正常 resume
    test_pruned_parameter_materializes_state_on_next_step —— 被裁剪的 param 在下次拿到梯度后重新 materialize
    test_restored_state_matches_saved_state —— 已保存 param 的 exp_avg 被完整还原

TestOptimizerStateValidation —— 非法/损坏 optimizer state 的拒绝行为
    test_partially_saved_parameter_state_is_rejected —— 单个 param 只存了部分 slot 时必须报错
"""


class TestSparseOptimizerDcpResume:
    """真实 DCP 往返下稀疏 optimizer state 的 resume 行为。"""

    def test_absent_lazy_state_is_pruned_and_resume_succeeds(self) -> None:
        # 验证:checkpoint 中完全缺失 lazy state 的 param 不会导致 resume 失败
        ...

另外建议把 saved_state / load_target 的构造抽成一个 fixture,三个断言点(裁剪结果、还原值、后续 materialize)分成三个用例,可读性和失败定位都更好。

Comment on lines +53 to +60
dcp.load(load_target, checkpoint_id=checkpoint_dir)
set_model_state_dict(restored_model, load_target["model"])
set_optimizer_state_dict(
restored_model,
restored_optimizer,
optim_state_dict=load_target["optimizer"],
options=StateDictOptions(strict=False),
)

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Claude: Critical(测试覆盖)— 回归测试没有走真实的修复路径 TrainEngine.load_dcp

这段把 load_dcp 里的 dcp.load → set_model_state_dict → set_optimizer_state_dict 在测试里重新实现了一遍,于是:

  • load_dcp 内部真正的修复代码(metadata 读取、prune 调用位置、model/optimizer 两套 StateDictOptions 的拆分)一行都没被执行;
  • 只要有人改坏 load_dcp(比如把 prune 放到 dcp.load 之后、或误传 strict=True 给 optimizer),这个测试照样通过
  • 测试里的 StateDictOptions(strict=False) 与生产代码的 StateDictOptions(cpu_offload=True, strict=False) 也已经不一致,测试无法锁住生产语义。

按 CLAUDE.md「Every bug fix PR must include a regression test that reproduces the original bug」,回归测试应当穿过 load_dcp 这个 public Interface。load_dcp 需要 dist 环境,但单卡 gloo 就够(本仓库已有 xtuner/v1/utils/test_utils.py / xtuner._testing 里的 dist 辅助设施,tests/engine/test_moe_train_engine.py:321 就是 engine 级 save/load 往返的现成范式)。建议至少补一个:

def test_load_dcp_resumes_when_a_parameter_never_received_grad(self, tmp_path):
    # 验证:某个 trainable param 在 save 前未拿到梯度时,load_dcp 不再因 missing keys 而失败
    engine = <minimal TrainEngine, world_size=1, gloo>
    engine.step_only_part_of_the_model()          # 让某个 param 的 optimizer state 保持 lazy
    engine.save_dcp(tmp_path / "weights")
    restored = <fresh engine>
    restored.load_dcp(tmp_path / "weights")       # 修复前在这里抛 missing keys
    assert <被裁剪的 param 不在 restored.optimizer.state >

如果确实无法在 CPU-only CI 上构造 TrainEngine,也请把当前这段改成直接调用 load_dcp(用最小 stub engine 持有 model/optimizer/has_freeze_params/model_cfg),而不是复制它的实现。

补充:整个用例被 skipif(torch < 2.10) 跳过,而仓库支持 torch>=2.6,所以在 2.6~2.9 的 CI 上该修复既未被验证、也确实是坏的(见我在 train_engine.py prune 调用处的 Critical 评论)。至少需要为低版本补一个「期望明确报错而不是裸 KeyError」的用例。

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

这个也需要修复下

Comment on lines +75 to +88
def test_partially_saved_optimizer_state_is_rejected() -> None:
optimizer_state = {
"state": {
"weight": {
"step": torch.tensor(1.0),
"exp_avg": torch.zeros(1),
"exp_avg_sq": torch.zeros(1),
}
}
}
saved_keys = {
"optimizer.state.weight.step",
"optimizer.state.weight.exp_avg",
}

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Claude: Nit(Bad Test 倾向)— 这个用例用手工合成的 dict 直接测私有函数

optimizer_statesaved_keys 都是手写的,等于把「DCP 扁平 key 的命名规则」在测试里复制了一份:如果生产端前缀写错(例如顶层 key 从 "optimizer" 改名),这个用例仍然绿。它测的是 helper 的内部约定,而不是模块的对外行为。

「partial state 必须 fatal」这个不变量本身值得测,建议让 saved_keys 来自真实的 DCP metadata,从而同时锁住 key 格式契约:

def test_partially_saved_parameter_state_is_rejected(self, tmp_path) -> None:
    # 验证:某个 param 只保存了部分 state slot 时必须报错,不能静默接受
    ...  # 真实 dcp.save 一个完整 checkpoint
    saved_keys = {str(k) for k in dcp.FileSystemReader(tmp_path).read_metadata().state_dict_metadata}
    saved_keys.discard("optimizer.state.bias.exp_avg_sq")   # 模拟 checkpoint 损坏/部分缺失
    with pytest.raises(RuntimeError, match="Incomplete optimizer state"):
        _prune_uncheckpointed_optimizer_state(load_target["optimizer"], saved_keys)

另外该用例没有 skipif,说明它在任何 torch 版本都跑;这也侧面说明它验证的只是纯 dict 逻辑,与真正的版本相关行为无关。

Comment on lines +62 to +74
def _prune_uncheckpointed_optimizer_state(optimizer_state_dict: dict[str, Any], saved_keys: set[str]) -> list[str]:
"""Remove wholly absent lazy optimizer states from a DCP load target.

``get_optimizer_state_dict`` materializes optimizer states for every
trainable parameter. A checkpoint may legitimately omit all state slots
for a parameter that had not received a gradient when it was saved. DCP
treats the materialized load target as authoritative and otherwise reports
those absent slots as missing checkpoint keys.

A partially saved parameter state is not lazy initialization and is kept as
an error: optimizers such as AdamW cannot step with only a subset of
``step``, ``exp_avg``, and ``exp_avg_sq`` restored.
"""

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Claude: Nit(.claude/CLAUDE.md 合规)

  1. 「When adding methods or functions to a module or class, declare public interfaces as early as possible and place private interfaces toward the end.」——这个私有模块级函数插在 DEVICE_MODULE 之后、所有公开类之前,位置偏前。建议下沉到文件末尾,或者按上面的建议直接变成 TrainEngine 的私有方法(它本质上只服务于 load_dcp)。
  2. Docstring 规范要求「参数类型必须在 docstring 中声明」「返回类型必须显式声明」。当前 docstring 只有正文说明,缺 Args: / Returns:saved_keys 期望的是扁平化后的 checkpoint key 集合这一关键约定完全没写出来,调用者很容易误传成 param 名集合)。建议:
def _prune_uncheckpointed_optimizer_state(optimizer_state_dict: dict[str, Any], saved_keys: set[str]) -> list[str]:
    """Remove wholly absent lazy optimizer states from a DCP load target.

    Args:
        optimizer_state_dict (dict[str, Any]): The optimizer load target; its ``state`` entry is pruned in place.
        saved_keys (set[str]): Flattened DCP checkpoint keys, e.g. ``optimizer.state.<fqn>.exp_avg``.

    Returns:
        list[str]: Flattened keys removed from the load target.
    """

(另外请在 docstring 或函数名上体现「原地修改入参」这一副作用——目前既返回值又原地改,调用方需要知道。)

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

建议作为trainengine的私有方法

set_options = StateDictOptions(cpu_offload=True, strict=False)
else:
set_options = StateDictOptions(cpu_offload=True, strict=True)
model_options = StateDictOptions(cpu_offload=True, strict=not self.has_freeze_params)

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

这里为啥要 set_options 改为 model_options

@matrix72c matrix72c closed this Aug 13, 2026
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants