【Bug已解决】KeyError in optimizer.state_dict() under FSDP2 when using Adagrad optimizer 解决方案一、现象长什么样用 FSDP2 训练优化器是Adagrad在保存 checkpoint 调用optimizer.state_dict()或accelerator.save_state时抛KeyError sum (Adagrad 的状态键) # 或 KeyError optimizer.state.param_id (某参数 id 在 state 映射里找不到)最小判据触发FSDP2 Adagrad调用 optimizer.state_dict() / save_state 现象KeyErrorsum 或 param id 缺失 根因FSDP2 在序列化优化器状态时对 Adagrad 的 sum 键 / 参数 id 映射处理不当 影响无法保存 Adagrad 训练的 checkpoint最迷惑的是Adam 下state_dict()正常Adagrad 下 KeyError。因为 Adagrad 的状态键结构sum与 Adamexp_avg/exp_avg_sq不同FSDP2 的序列化逻辑对sum没覆盖。二、背景optimizer.state_dict()返回{state: {param_id: {...}}, param_groups: [...]}其中state是param_id - 该参数状态的映射。FSDP2 在分片下优化器状态也按 shard 分片序列化时需要把分片状态的 param_id正确映射到全局 param_id并把每个参数的状态键sum/exp_avg等原样保留。Adagrad 的状态键是sum逐元素梯度平方和。问题出在FSDP2 的 state_dict 序列化假设了特定状态键某些实现对状态键做了白名单 / 转换如只认exp_avg/exp_avg_sq遇到sum找不到对应处理 - KeyErrorparam_id 映射错位FSDP2 把分片参数的局部 id 映射回全局 id若 Adagrad 的某个 shard 参数在映射表里缺失比如该 shard 恰好为空 / 被优化器剔除state[param_id]查不到 - KeyError懒初始化 分片Adagrad 的sum在第一次step()时懒初始化若state_dict()在sum尚未对某些 shard 初始化时调用该 shard 的sum键缺失 - KeyError。根因是FSDP2 的优化器状态序列化对 Adagrad 的sum键 / 参数 id 映射处理不全。三、根因抽象成代码示意def flatten_optim_state(state, known_keys): out {} for pid, st in state.items(): # BUG只处理已知键sum 不在白名单 - KeyError out[pid] {k: st[k] for k in known_keys} # sum 缺失 - KeyError return out根因链条Adagrad 状态含sum键FSDP2 序列化逻辑对状态键做白名单 / 转换没含sumst[sum]找不到 - KeyError或 param_id 映射在某些 shard 缺失 -state[pid]KeyErrorAdam 状态键在白名单内所以正常Adagrad 的sum不在所以炸。一句话FSDP2 优化器状态序列化对 Adagrad 的sum键 / 参数 id 映射处理不全state_dict() 时 KeyError。四、最小可运行复现用纯 Python 模拟状态键白名单漏了 sum 导致 KeyError# repro_adagrad_statedict.py def flatten(state, known_keys): out {} for pid, st in state.items(): out[pid] {k: st[k] for k in known_keys} # sum 不在 - KeyError return out def main(): # Adagrad 状态sum 键 state {0: {sum: [1, 2, 3]}, 1: {sum: [4, 5, 6]}} try: flatten(state, known_keys[exp_avg, exp_avg_sq]) # 漏 sum except KeyError as e: print(复现成功 -, e) if __name__ __main__: main()运行输出复现成功 - sum白名单漏了sum序列化时 KeyError正是真实 bug 的抽象。五、解决方案第一层最小直接修复最小且必须的一步FSDP2 序列化优化器状态时保留全部状态键不白名单过滤只做分片 param_id 的映射# fix_layer1.py def flatten_optim_state(state, pid_map): # 修复保留每个参数的全部状态键含 Adagrad 的 sum out {} for local_pid, st in state.items(): global_pid pid_map.get(local_pid, local_pid) out[global_pid] dict(st) # 原样保留不过滤 return out要点不再对状态键做白名单过滤Adagrad 的sum自然被保留只做local_pid - global_pid映射键集合由优化器自身决定这一层改动最小但依赖所有状态键都被原样传递。六、解决方案第二层结构性改进把优化器状态序列化做成与优化器无关的通用逻辑状态键集合完全来自优化器实例本身动态读取FSDP2 只负责 param_id 映射与分片重建不假设任何具体键# fix_layer2.py from dataclasses import dataclass, field from typing import Dict, List dataclass class OptimStateSerializer: # 不预置任何键白名单键集合运行时从 optimizer 读取 def serialize(self, state: Dict, pid_map: Dict) - Dict: out {} for local_pid, st in state.items(): gpid pid_map.get(local_pid, local_pid) # 动态保留该参数所有状态键Adam 的 m/v、Adagrad 的 sum 等 out[gpid] {k: _maybe_slice(v) for k, v in st.items()} return out def _maybe_slice(v): # 若是分片张量保留其分片形态否则原样 return v # 用法 serializer OptimStateSerializer() sd serializer.serialize(optimizer.state, pid_map)要点serialize不预置键集合状态键完全来自st.items()动态新增任何优化器Adagrad / Adam / AdamW / 自定义都无需改序列化逻辑FSDP2 只管 param_id 映射与分片键语义交给优化器自身。七、解决方案第三层断言 / CI 守护写 pytest 验证Adagrad 的 sum 键在 state_dict 中被保留# test_adagrad_statedict.py import pytest def serialize(state, pid_map, known_keysNone): out {} for pid, st in state.items(): gpid pid_map.get(pid, pid) if known_keys is None: out[gpid] dict(st) # 修复不过滤 else: out[gpid] {k: st[k] for k in known_keys} return out def test_sum_key_preserved_no_whitelist(): state {0: {sum: [1, 2, 3]}} out serialize(state, pid_map{}, known_keysNone) assert sum in out[0], Adagrad 的 sum 键必须被保留 def test_sum_key_with_whitelist_fails(): state {0: {sum: [1, 2, 3]}} with pytest.raises(KeyError): serialize(state, pid_map{}, known_keys[exp_avg]) def test_pid_mapped(): state {0: {sum: [1]}} out serialize(state, pid_map{0: 7}) assert 7 in out and 0 not in outCI 一旦有人把状态键白名单加回来漏sumtest_sum_key_preserved_no_whitelist立刻变红。八、排查清单FSDP2 Adagrad 调state_dict()KeyError 时确认 KeyError 的键是不是sumAdagrad 状态键检查 FSDP2 序列化是否对状态键做了白名单过滤漏了sum检查 param_id 映射是否在某些 shard 缺失按第五 / 六节保留全部状态键、只做 pid 映射Adam 正常、Adagrad 异常几乎可断定是sum键未覆盖确认 Adagrad 的sum在state_dict()前已对所有 shard 初始化把第七节的 pytest 接进 CI守护sum 键被保留。九、小结FSDP2 下 Adagrad 调用optimizer.state_dict()报 KeyError根因是 FSDP2 的优化器状态序列化对状态键做了白名单/转换漏掉了 Adagrad 的sum键或 param_id 映射在某些 shard 缺失序列化时查不到即 KeyError。Adam 状态键在白名单内所以正常。三层层级第一层序列化时保留全部状态键不过滤只做 param_id 映射第二层用与优化器无关的通用OptimStateSerializer键集合运行时动态读取第三层pytest 验证 Adagrad 的sum键被保留锁进 CI。核心教训任何序列化优化器状态的逻辑都不应假设具体状态键m/v/sum。状态键语义属于优化器自身序列化层只负责 param 映射与分片——把键集合做成动态读取比白名单过滤稳也天然支持未来新优化器。本篇与第 546 篇互补546 聚焦 Adagrad 在分片下的数值正确性本篇聚焦其 state_dict 序列化。