[Fix] Support sparse optimizer DCP resume - #2018
Conversation
|
@claude review |
|
Claude: 代码 review(1/3:Summary 与流程图) Summary 本 PR 修复「DCP resume 时,save 前未收到梯度的 trainable param 因 optimizer lazy state 缺失而导致 方向是对的,问题定位也准确,但有 1 个阻塞性问题:
Verdict: REQUEST_CHANGES(阻塞项:torch < 2.10 兼容性;回归测试需穿过 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
关键点: 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
文字说明(突出改动)
|
|
Claude: 代码 review(2/3:核心原理实现与单测、抽象与信息隐藏、线性业务流程) 核心原理实现与单测 核心实现是
裁剪用的 key 拼接之所以能对上,依赖两个约定: 「部分缺失 fatal」的判断也有道理:AdamW 无法只用 对应单测
覆盖缺口(重要):两个用例都没有调用 抽象与信息隐藏评估
改进意见:
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)
公开 Interface 的线性业务流程评估
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)顺带一提: |
|
Claude: 代码 review(3/3:单测建议、其他 issues、Verdict) 单测建议 保留:
需要新增(优先级最高):
建议精简与重构:
其他 issues
Verdict REQUEST_CHANGES 阻塞项:(1) torch < 2.10 的版本兼容(当前会退化成裸 |
| 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) |
There was a problem hiding this comment.
Claude: Critical — 该修复在 torch < 2.10 上会把原来的报错换成更难定位的 KeyError
本仓库声明支持 torch>=2.6.0(pyproject.toml#L40、requirements/runtime.txt:12),并且代码里明确对 2.7/2.8/2.9 做过特化(xtuner/v1/train/trainer.py:595 的 torch.__version__.startswith("2.7.")、xtuner/v1/patch/torch_shape_env_simplify_pt28.py、xtuner/v1/patch/xtuner_storage.py:29),说明 2.6~2.9 是活跃目标版本。
而 strict 对 optimizer 侧生效是 2.10 才有的行为:
- torch v2.9.0 / v2.6.0 的
torch/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。)
| 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 |
There was a problem hiding this comment.
Claude: Warning — 硬编码的 "optimizer." 前缀把调用方的命名约定泄漏进了 helper,且失配时会「静默丢弃全部 optimizer state」
optimizer.state.{param}.{state_name} 这个扁平 key 由两个约定共同决定:
_get_dcp_state_dict()(xtuner/v1/engine/train_engine.py:340)把 optimizer state 放在顶层 key"optimizer"下;- DCP
flatten_state_dict=True用"."连接路径(本仓库XtunerCacheSavePlanner默认flatten_state_dict=True,xtuner/v1/patch/xtuner_cache_planner.py:29,所以格式一致)。
但 helper 只拿到 optimizer_state_dict 这一个 dict,前缀却写死在函数体里。这形成一个隐性耦合:一旦顶层 key 改名、或 torch 改了 flatten 规则、或将来打开 flatten_optimizer_state_dict,present_state_names 会对所有 param 都为空 → 走到 state.pop(...) 分支 → 全部 param state 被清掉,只留一条 warning;再叠加下面 set_optimizer_state_dict(strict=False),整个 optimizer 会被静默重置为初始状态继续训练,metric 上很难察觉。
两点建议:
- 把前缀变成 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}
- 增加一个「全量裁剪」的健全性检查——所有 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." )
| 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}") |
There was a problem hiding this comment.
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.load侧create_read_items_for_chunk_list对没有 overlap 的区域直接产生 0 个 read item,rank B 的那部分 shard 保持_init_optim_state产生的零值,不报任何错;- 原本能兜住这种不一致的
_validate_global_plan在skip_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 对齐)。
| logger.warning( | ||
| f"Ignoring {len(removed)} lazy optimizer-state leaves absent from the checkpoint; " | ||
| f"examples={removed[:5]}" | ||
| ) |
There was a problem hiding this comment.
Claude: Nit(抽象层级 + 日志)
- 线性业务流程:
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(...)-
logger.warning会在每个 rank 都打印(千卡规模下就是上千行)。本文件已有 rank0-only 的既有写法(log_rank0.warning,见xtuner/v1/engine/train_engine.py:111、:280),这里建议改用log_rank0.warning。 -
dcp.FileSystemReader(weights_dir).read_metadata()额外读了一次.metadata(dcp.load内部还会再读一次),且把 storage backend 写死成 FileSystemReader。建议构造一次 reader 并同时传给dcp.load(..., storage_reader=reader),既省一次 IO 也避免两处对 backend 的假设不一致。
| # 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), |
There was a problem hiding this comment.
Claude: Warning — optimizer 侧无条件 strict=False 削弱了 model 之外的全部校验
改动本身有必要(load_states=False 时 optimizer_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 之间是有落差的,建议一并说明清楚。
| 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 |
There was a problem hiding this comment.
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)分成三个用例,可读性和失败定位都更好。
| 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), | ||
| ) |
There was a problem hiding this comment.
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」的用例。
| 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", | ||
| } |
There was a problem hiding this comment.
Claude: Nit(Bad Test 倾向)— 这个用例用手工合成的 dict 直接测私有函数
optimizer_state 和 saved_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 逻辑,与真正的版本相关行为无关。
| 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. | ||
| """ |
There was a problem hiding this comment.
Claude: Nit(.claude/CLAUDE.md 合规)
- 「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)。 - 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 或函数名上体现「原地修改入参」这一副作用——目前既返回值又原地改,调用方需要知道。)
| 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) |
There was a problem hiding this comment.
这里为啥要 set_options 改为 model_options
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
strict=Falsebehavior.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
ruff check xtuner/v1/engine/train_engine.py tests/engine/test_train_engine_checkpoint.pypython3.12 -m py_compile xtuner/v1/engine/train_engine.py tests/engine/test_train_engine_checkpoint.py