【Bug已解决】FSDP2 + SHARDED_STATE_DICT: optimizer checkpoint load fails with missing step key. (RuntimeE

📅 2026/8/1 16:59:32 👁️ 阅读次数 📝 编程学习
【Bug已解决】FSDP2 + SHARDED_STATE_DICT: optimizer checkpoint load fails with missing step key. (RuntimeE

【Bug已解决】FSDP2 + SHARDED_STATE_DICT: optimizer checkpoint load fails with missing step key. (RuntimeError: Missing key in checkpoint state_dict: optimizer.state.0.step.) 解决方案

一、现象长什么样

用 FSDP2 做全分片训练,并在保存 / 恢复时启用StateDictType.SHARDED_STATE_DICT,以便每个 rank 只存自己那一份分片。保存时一切正常,但重启加载时直接抛出:

RuntimeError: Missing key in checkpoint state_dict: optimizer.state.0.step.

也就是:优化器状态字典里的state子字典,期望有optimizer.state.0.step(第 0 号参数的步数计数器),可加载进来的分片里根本没有这个键。

最让人困惑的地方:

  • 保存阶段没有报错,torch.save成功落盘;
  • 优化器的step在训练时明明一直在自增(学习率调度依赖它);
  • 加载代码看起来和官方示例一致:optimizer.load_state_dict(ckpt["optimizer"])
  • 只要改用FULL_STATE_DICT就能正常加载,SHARDED_STATE_DICT一上就炸。

这说明问题不在"step 没被保存",而在"sharded 格式的键结构与加载端期望的键结构对不上"。

二、背景

FSDP2 通过torch.distributed.fsdp.FullyShardedDataParallel.state_dict_type这个上下文管理器来决定状态字典的形态:

  • FULL_STATE_DICT:聚合回完整、未分片的字典,键是全局 param 索引;
  • SHARDED_STATE_DICT:每个 rank 只持有自己分片的那部分,键是本地param 索引(从 0 开始),且state里每个参数组只含本 rank 负责的张量;
  • LOCAL_STATE_DICT:最原始的本地位姿,键也是本地索引。

关键点:优化器状态字典的state子字典,其键(param 索引)是在"当前这个 state_dict_type 上下文"里生成的。保存时若在SHARDED_STATE_DICT上下文里调用optimizer.state_dict(),得到的是按本地索引、step被切到对应 rank 的字典;加载时若不在相同上下文里调用optimizer.load_state_dict(),PyTorch 会用另一套键结构去对齐,于是找不到optimizer.state.0.step

更隐蔽的是:step是一个 Python 标量 / 小张量,sharded 保存时它属于"哪个 rank 负责该参数分片"的那一份。若保存端和加载端的 rank 数、分片方式、或者 state_dict_type 不一致,step这一项就会落在错误的分片里,加载端自然缺失。

三、根因

把不一致的上下文抽象成代码(示意,非照抄源码):

# 保存端(正确:在 SHARDED 上下文里取 optimizer state) with FSDP.state_dict_type(model, StateDictType.SHARDED_STATE_DICT): optim_sd = optimizer.state_dict() # 键是本地索引,step 随分片分布 torch.save({"optimizer": optim_sd}, path) # 加载端(错误:脱离了上下文,用默认 FULL 的键结构去对齐) ckpt = torch.load(path) optimizer.load_state_dict(ckpt["optimizer"]) # 键对不齐 -> Missing step

根因链条:

  1. optimizer.state_dict()的键结构由"调用它时所处的state_dict_type上下文"决定;
  2. 保存端在SHARDED_STATE_DICT内,键为本地索引,step被切到对应 rank;
  3. 加载端没有重新进入相同上下文,PyTorch 默认按FULL_STATE_DICT的全局索引去匹配;
  4. 两边键空间不一致,optimizer.state.0.step在加载端期望全局索引 0,但文件里它是某个本地分片里的键;
  5. 对齐失败,抛Missing key——这是典型的"保存 / 加载上下文不对称"导致的 KeyError。

另一个常见变体:保存时用了SHARDED,但optimizer.state_dict()之前没有modelstate_dict_type包裹(即 optimizer 状态其实还是 FULL 形态存进去的),加载端却按 SHARDED 去读,同样对不上。

四、最小可运行复现

用一段纯 Python 模拟"键空间不一致"导致step缺失:

# repro_sharded_step.py class OptState: def __init__(self): self.state = {} # param_index -> {"step": int, "exp_avg": ...} def save_sharded(self, rank, world): # 模拟 SHARDED:每个 rank 只存自己负责的本地索引 sharded = {} for idx in range(rank, len(self.state), world): sharded[idx] = self.state[idx] # 本地索引 return sharded def load_full_expect(self, sharded): # 模拟加载端用 FULL 键结构(全局索引 0..N-1)去对齐 for idx in range(len(self.state)): if idx not in sharded: raise KeyError(f"optimizer.state.{idx}.step") def main(): opt = OptState() opt.state = {0: {"step": 7}, 1: {"step": 7}, 2: {"step": 7}, 3: {"step": 7}} world = 2 rank0 = opt.save_sharded(0, world) # 本地索引 {0:.., 2:..} rank1 = opt.save_sharded(1, world) # 本地索引 {1:.., 3:..} merged = {**rank0, **rank1} # 合并后键是 0,1,2,3 -> 其实能对齐 # 但下面模拟"加载端只拿到了 rank0 的分片却按全局 0..3 期望" try: opt.load_full_expect(rank0) # 只给 rank0 分片 except KeyError as e: print("复现成功 ->", e) if __name__ == "__main__": main()

运行输出:

复现成功 -> optimizer.state.1.step

这正是真实 bug 的抽象:step被切进了不同 rank 的分片,加载端若没用对称的上下文 / 没合并完整分片,就会缺少某些step键。

五、解决方案(第一层:最小直接修复)

保证保存与加载都在同一个state_dict_type上下文里,这是最小且必须的一步:

# fix_layer1.py import torch from torch.distributed.fsdp import FullyShardedDataParallel as FSDP from torch.distributed.fsdp import StateDictType def save_sharded(model, optimizer, path): with FSDP.state_dict_type(model, StateDictType.SHARDED_STATE_DICT): sd = { "model": model.state_dict(), "optimizer": optimizer.state_dict(), # 与上下文一致 } torch.save(sd, path) def load_sharded(model, optimizer, path): ckpt = torch.load(path) # 关键:加载也进入完全相同的 SHARDED 上下文 with FSDP.state_dict_type(model, StateDictType.SHARDED_STATE_DICT): model.load_state_dict(ckpt["model"]) optimizer.load_state_dict(ckpt["optimizer"]) # 键结构对齐 -> step 不再缺失

只要 save / load 都用SHARDED_STATE_DICToptimizer.state.0.step这类键就会在两侧用同一套本地索引键空间生成与消费,缺失问题消失。

六、解决方案(第二层:结构性改进)

把"对称上下文"收敛成一个职责单一的 IO 助手,避免任何调用方忘记配对上下文;并显式确保step被纳入优化器状态:

# fix_layer2.py from dataclasses import dataclass from enum import Enum from typing import Any, Callable class StateDictKind(str, Enum): SHARDED = "sharded" FULL = "full" LOCAL = "local" @dataclass(frozen=True) class CkptSpec: kind: StateDictKind def make_state_dict_context(model, spec: CkptSpec): """单一来源:根据 spec 返回对应的 state_dict_type 上下文。""" from torch.distributed.fsdp import FSDP, StateDictType mapping = { StateDictKind.SHARDED: StateDictType.SHARDED_STATE_DICT, StateDictKind.FULL: StateDictType.FULL_STATE_DICT, StateDictKind.LOCAL: StateDictType.LOCAL_STATE_DICT, } return FSDP.state_dict_type(model, mapping[spec.kind]) def save_checkpoint(model, optimizer, path, spec: CkptSpec): with make_state_dict_context(model, spec): torch.save({ # type: ignore[name-defined] "model": model.state_dict(), "optimizer": optimizer.state_dict(), }, path) def load_checkpoint(model, optimizer, path, spec: CkptSpec): ckpt = torch.load(path) with make_state_dict_context(model, spec): # 强制对称 model.load_state_dict(ckpt["model"]) optimizer.load_state_dict(ckpt["optimizer"]) def assert_step_present(optimizer) -> None: """结构性守护:确保 step 计数器确实在优化器状态里。""" for grp in optimizer.param_groups: for p in grp["params"]: st = optimizer.state[p] assert "step" in st, "optimizer.state 缺少 step 计数器"

要点:

  • save_checkpoint/load_checkpoint共用make_state_dict_context,上下文对称由结构保证,无法被某个调用方漏掉;
  • assert_step_present在加载后立即校验每个参数的step,把"缺失"提前变成显式异常;
  • CkptSpec成为唯一真相来源,切换 FULL / SHARDED 只需改一处。

七、解决方案(第三层:断言 / CI 守护)

写一条 pytest,保存 -> 加载往返,断言step不丢且数值相等,把回归锁进 CI:

# test_sharded_optim_ckpt.py import pytest class FakeOpt: def __init__(self, n): self.state = {i: {"step": 7} for i in range(n)} self.param_groups = [{"params": list(range(n))}] def state_dict(self): return {"state": {k: dict(v) for k, v in self.state.items()}} def load_state_dict(self, sd): # 模拟 PyTorch:按 state 键对齐,缺失即报错 for idx in sd["state"]: pass for idx in self.state: if idx not in sd["state"]: raise KeyError(f"optimizer.state.{idx}.step") self.state = {int(k): v for k, v in sd["state"].items()} def test_step_roundtrip_sharded(): opt = FakeOpt(4) sd = opt.state_dict() opt2 = FakeOpt(4) opt2.load_state_dict(sd) # 对称:同一键空间 assert opt2.state[0]["step"] == 7 def test_missing_step_detected(): opt = FakeOpt(4) sd = opt.state_dict() del sd["state"][0] # 模拟 step 被切丢 with pytest.raises(KeyError): opt.load_state_dict(sd) def test_context_symmetry_required(): """加载端必须和保存端用同一 state_dict_type 键空间。""" saved = {0: 7, 2: 7} # 仅 rank0 分片(本地索引) missing = {1, 3} # 全局期望 0..3 for idx in missing: assert idx not in saved # 说明不对称会缺键

CI 一旦回归到"加载脱离上下文",test_missing_step_detected与往返测试会立即变红。

八、排查清单

遇到Missing key ... optimizer.state.0.step时:

  1. 确认保存端optimizer.state_dict()是否包在FSDP.state_dict_type(..., SHARDED_STATE_DICT)内;
  2. 确认加载端optimizer.load_state_dict()是否包在完全相同的上下文内;
  3. 比对两边 rank 数 / 分片方式是否一致(world size 变了也会错位);
  4. 若用accelerate,确认accelerator.load_statestate_dict_type与保存时一致;
  5. 加载后立即用assert_step_present(optimizer)校验step存在;
  6. 临时改用FULL_STATE_DICT验证"是不是键空间不对称"——能加载就坐实本 bug;
  7. 把第七节的 pytest 接进 CI,作为 sharded checkpoint 的回归护栏。

九、小结

FSDP2 +SHARDED_STATE_DICT下优化器加载报Missing key: optimizer.state.0.step,根因是保存与加载所处的state_dict_type上下文不对称:保存端在 SHARDED 上下文里按本地索引切分了step,加载端却脱离上下文、用全局索引键空间去对齐,于是键对不上。

三层层级:

  • 第一层:保存与加载都进入相同的SHARDED_STATE_DICT上下文,键结构对齐;
  • 第二层:用CkptSpec+make_state_dict_context把对称上下文收敛成单一入口,并加assert_step_present提前暴露缺失;
  • 第三层:写 pytest 做保存->加载往返与缺失检测,锁进 CI。

核心教训:任何"状态字典形态由上下文决定"的 API,保存与加载必须成对地处在同一上下文里;凡是不成对的,都是 KeyError / 静默错位的温床。