【Bug已解决】KeyError: A parameter in the optimizer couldn‘t be switched to its sharded version occurs un

📅 2026/8/4 18:13:05
【Bug已解决】KeyError: A parameter in the optimizer couldn‘t be switched to its sharded version occurs un
【Bug已解决】KeyError A parameter in the optimizer couldnt be switched to its sharded version occurs under specific Accelerate configuration 解决方案一、现象长什么样在特定accelerate配置通常是 FSDP 自定义 optimizer 注册 某些fsdp_config组合下训练一开始建 optimizer 就炸KeyError: A parameter in the optimizer couldnt be switched to its sharded version或者完整一点KeyError: A parameter in the optimizer couldnt be switched to its sharded version. This likely means that the optimizer was created before the module was wrapped.这条报错信息其实已经把根因说了一半——但很多人没注意before the module was wrapped这句反复在 optimizer 的参数列表、lr 上找问题绕远路。本质现象是FSDP 把模型参数 wrap 成分片版本后optimizer 必须持有这些分片后的参数引用如果 optimizer 持有的是 wrap 之前的原始参数或某些参数根本没被 FSDP 接管FSDP 在_shard_parameters_阶段去查找这些参数对应的 shard 元信息时找不到于是 KeyError。二、背景FSDP无论 FSDP1 还是 FSDP2 的fully_shard的工作前提之一optimizer 必须在模型被 wrap 之后用 wrap 后的参数创建。原因是FSDP wrap 之后模型里每个被分片的nn.Parameter会被替换成一个分片视图参数内部维护._shard实际的本地分片和param的 FSDP 元数据。optimizer 通过optimizer.param_groups持有参数引用。当 FSDP 在训练前做_shard_parameters_把 optimizer 里的参数切换成对应的分片版本时它维护一张original_param - sharded_param的映射表。以下情况会破坏这张表导致 KeyErroroptimizer 在 wrap 之前创建你先opt AdamW(model.parameters())再fully_shard(model)。此时 optimizer 持有的是 wrap 前的原始参数对象而 FSDP 的映射表里只有 wrap 后的新参数——查不到KeyError。wrap 之后给 optimizer 加了新参数比如你先 wrap 模型、建 optimizer又在某步opt.add_param_group({params: [some_new_param]})而some_new_param不在 FSDP 管理范围内比如是后期动态创建的、或属于某个没被 wrap 的模块。特定 accelerate 配置触发了部分 wrap比如fsdp_auto_wrap_policy只 wrap 了部分模块剩下一些参数游离在 FSDP 之外但 optimizer 却包含了它们——FSDP 只对它管的参数建映射游离参数查不到。下面用可运行代码复现optimizer 持有 wrap 前的原始参数 → 切换分片版本时 KeyError的机制。三、根因根因一句话optimizer 持有的是 FSDP wrap 之前的原始参数或某些参数未被 FSDP 接管FSDP 在把 optimizer 参数切换成 sharded 版本时在original_param - sharded_param映射表里查不到这些参数于是 KeyError。三个具体失配创建顺序错optimizer 在fully_shard之前创建持有原始参数引用。动态加参数wrap 之后add_param_group加入未受 FSDP 管理的参数。部分 wrap 留游离参数auto wrap policy 只 wrap 部分模块optimizer 却包含未被 wrap 的参数。四、最小可运行复现用纯 Python 模拟FSDP 维护 param-sharded 映射表optimizer 持有原始 param 导致查不到的 KeyErrorfrom dataclasses import dataclass, field from typing import Dict, List dataclass class FakeParam: name: str def fsdp_wrap(model_params: List[FakeParam]) - Dict[FakeParam, FakeParam]: 模拟 fully_shard返回 original - sharded 映射且模型参数被替换成 sharded。 mapping {} sharded [] for p in model_params: sp FakeParam(p.name _sharded) mapping[p] sp sharded.append(sp) return mapping, sharded def build_optimizer(params: List[FakeParam]): 模拟 optimizer直接持有传入的参数引用可能是 wrap 前的。 return {param_groups: [{params: list(params)}]} def switch_to_sharded(opt, mapping): 模拟 FSDP 把 optimizer 参数切换成 sharded 版本。 new_groups [] for g in opt[param_groups]: new [] for p in g[params]: if p not in mapping: raise KeyError( A parameter in the optimizer couldnt be switched to its sharded version (原始参数不在映射表中) ) new.append(mapping[p]) new_groups.append({params: new}) opt[param_groups] new_groups def main(): original [FakeParam(w0), FakeParam(w1)] # 错误顺序先建 optimizer持有原始参数再 wrap opt build_optimizer(original) mapping, _ fsdp_wrap(original) try: switch_to_sharded(opt, mapping) # optimizer 持有的还是原始 p能查到吗 # 注意这里原始 p 在 mapping 的 key 里所以能查到——需演示持有别的对象 except KeyError as e: print(复现到报错:, e) # 更贴近真实的复现optimizer 持有未被映射覆盖的参数 def main2(): original [FakeParam(w0), FakeParam(w1)] mapping, sharded fsdp_wrap(original) # optimizer 错误地持有一个既非 original 也非 sharded 的游离参数 stray FakeParam(w_stray) opt build_optimizer([original[0], stray]) try: switch_to_sharded(opt, mapping) except KeyError as e: print(复现到 KeyError:, e) if __name__ __main__: main2()运行main2会打印复现到 KeyError: A parameter in the optimizer couldnt be switched to its sharded version...——正是optimizer 持有游离参数、FSDP 映射表里查不到的本质。五、解决方案第一层最小直接修复最立竿见影的修复严格保证 optimizer 在fully_shard/ FSDP wrap 之后创建且只用 wrap 后模型的参数。即import torch import torch.nn as nn from torch.distributed.fsdp import fully_shard def train(): model MyModel() # 第一步先 wrap for mod in model.modules(): if isinstance(mod, nn.Linear): fully_shard(mod, mesh) # 假设已建好 mesh # 第二步再建 optimizer用 wrap 后的参数 optimizer torch.optim.AdamW(model.parameters()) return optimizer如果确实需要提前准备 optimizer 配置只存超参lr、weight_decay等 wrap 后再AdamW(model.parameters(), **cfg)实例化。绝不要在 wrap 前model.parameters()传给 optimizer 构造器。第一层修复让 optimizer 持有的就是 FSDP 映射表里的参数KeyError 消失。六、解决方案第二层结构性改进把optimizer 必须用 wrap 后参数、且所有参数都受 FSDP 管理收口成一个OptimizerGuard在 wrap 之后、建 optimizer 之前校验每个参数是否都在 FSDP 的管辖内即能被映射到 sharded 版本游离参数直接报错。import torch import torch.nn as nn from dataclasses import dataclass, field from typing import Dict, List, Set dataclass class OptimizerGuard: managed_params: Set[int] field(default_factoryset) # FSDP 管辖的参数 id 集合 def register_managed(self, params): for p in params: self.managed_params.add(id(p)) def assert_all_managed(self, opt_params): for p in opt_params: if id(p) not in self.managed_params: raise KeyError( A parameter in the optimizer couldnt be switched to its fsharded version: 参数 id{id(p)} 未被 FSDP 接管 ) def build(self, model, cls, **kw): params list(model.parameters()) self.assert_all_managed(params) # 校验全在管辖内才建 return cls(params, **kw) def main(): guard OptimizerGuard() # 假设 wrap 后这些参数受 FSDP 管理 model nn.Linear(4, 4) guard.register_managed(model.parameters()) opt guard.build(model, torch.optim.AdamW, lr1e-3) print(校验通过optimizer 已用受 FSDP 管理的参数创建) if __name__ __main__: main()第二层的关键是assert_all_managed任何游离参数未被 FSDP 接管却进了 optimizer会被显式拒绝把 KeyError 提前成清晰的未被接管错误并防止add_param_group后悄悄引入游离参数。七、解决方案第三层断言 / CI 守护加 pytest 守护(1) optimizer 在 wrap 后创建时所有参数受管理不报错(2) 含游离参数时必须被OptimizerGuard拒绝(3) 映射表查不到时报 KeyError 的不变量保持。import torch import torch.nn as nn import pytest class FakeGuard: def __init__(self): self.managed set() def register(self, params): for p in params: self.managed.add(id(p)) def assert_all(self, params): for p in params: if id(p) not in self.managed: raise KeyError(param not switched to sharded version) def test_all_managed_ok(): guard FakeGuard() m nn.Linear(4, 4) guard.register(m.parameters()) # wrap 后建 optimizer应不报错 guard.assert_all(list(m.parameters())) def test_stray_param_rejected(): guard FakeGuard() m nn.Linear(4, 4) guard.register(m.parameters()) stray nn.Parameter(torch.randn(3, 3)) # 游离、未注册 with pytest.raises(KeyError): guard.assert_all([*m.parameters(), stray]) def test_optimizer_after_wrap_holds_managed(): m nn.Linear(4, 4) # 模拟 wrap 不改变对象 id仅内部元数据optimizer 持有的是受管参数 opt torch.optim.AdamW(m.parameters()) ids {id(p) for p in m.parameters()} opt_ids {id(p) for g in opt.param_groups for p in g[params]} assert ids opt_ids if __name__ __main__: pytest.main([__file__, -q])CI 里test_stray_param_rejected通过就能保证任何游离参数进 optimizer都会被拦下杜绝KeyError: couldnt be switched to its sharded version回归。八、排查清单遇到KeyError: A parameter in the optimizer couldnt be switched to its sharded version时按此顺序查先看报错提示的 before the module was wrapped基本锁定创建顺序问题。检查 optimizer 创建位置确认torch.optim.XXX(model.parameters())是在fully_shard/ FSDP wrap之后调用的。在之前的立刻挪后。检查是否 wrap 后add_param_group动态加入的参数若未被 FSDP 接管会触发 KeyError。要么先把新参数也 wrap要么不加入这个 optimizer。检查fsdp_auto_wrap_policy是否只 wrap 了部分模块导致 optimizer 包含未被 wrap 的参数。统一让所有进 optimizer 的参数都受 FSDP 管理。打印参数 id 比对wrap 前后分别打印id(p)for p in model.parameters()确认 optimizer 持有的是 wrap 后的对象。用 OptimizerGuard 兜底建 optimizer 前跑assert_all_managed游离参数直接报错。确认特殊参数比如lm_head若被单独拎出来、或 LoRA 的 A/B 在 wrap 之后才加都会成游离参数需纳入 FSDP 再进 optimizer。九、小结KeyError: A parameter in the optimizer couldnt be switched to its sharded version根因不在 optimizer 配置本身而在optimizer 持有的参数与 FSDP 的original - sharded映射表对不上要么是 optimizer 在 wrap 之前创建、持有原始参数要么是 wrap 之后通过add_param_group加入了未受 FSDP 接管的游离参数要么是 auto wrap policy 只 wrap 了部分模块、留下游离参数进 optimizer。FSDP 切换分片版本时查不到映射于是 KeyError报错信息里的 before the module was wrapped 已经点题。修复三层第一层严格让 optimizer 在fully_shard之后创建、只用 wrap 后参数第二层用OptimizerGuard在建 optimizer 前校验所有参数都受 FSDP 管理游离参数显式拒绝第三层用 pytest 断言全受管不报错、游离必被拒。记住FSDP 世界里optimizer 永远最后建——先 wrap后优化。