【Bug已解决】Accelerate save_state() error using FSDP2/TP 解决方案

发布时间:2026/8/4 20:27:24
【Bug已解决】Accelerate save_state() error using FSDP2/TP 解决方案 【Bug已解决】Accelerate save_state() error using FSDP2TP 解决方案一、现象长什么样用accelerate的accelerator.save_state(output_dir)保存训练 checkpoint一旦底层是 FSDP2或叠加 TP保存阶段经常直接报错而不是训练阶段ValueError: optimizer state_dict contains sharded keys but save_state expected full keys或者RuntimeError: Cannot pickle DTensor; use distributed checkpoint format instead of torch.save又或者更隐蔽——保存成功了但 reload 时KeyError: optimizer.state.0.step # 或 shape mismatch on load还有一个常见但不报错的表现save_state在 FSDP2 下尝试用FULL_STATE_DICT把整个优化器状态 all-gather 到每个 rank瞬间 OOM报CUDA out of memory——你以为 OOM 是训练的事其实是保存时把全量状态拉回了单卡。这类问题共同点是save_state()默认把模型/优化器状态当成完整、可 pickle 的张量来存而 FSDP2/TP 下状态是分片的sharded且 TP 下参数是 DTensor二者都不满足torch.save友好的假设。二、背景accelerate的save_state()内部逻辑大致是收集model.state_dict()和optimizer.state_dict()然后torch.save成pytorch_model.bin/optimizer.bin之类。这个逻辑是为DDP / 单卡设计的——那时状态和参数都是完整的普通 Tensor。但 FSDP2 TP 改变了这个前提FSDP2 的参数/状态是分片的每个 rank 只持有1/N的参数和优化器状态。要存成完整状态必须做 all-gather开销巨大且需要正确的state_dict_type上下文SHARDED_STATE_DICT/LOCAL_STATE_DICT/FULL_STATE_DICT。save_state()若没进入SHARDED_STATE_DICT上下文就会按完整态去读读到的 key 是分片形态与 load 时期望的完整 key 对不上。TP 下参数是 DTensorparallelize把线性层权重切成ColwiseParallel/RowwiseParallel的 DTensor。torch.save无法 pickle DTensor它没有跨进程序列化的内置支持必须改用分布式检查点格式torch.distributed.checkpoint的save/load基于StateDictType.SHARDED_STATE_DICT。save_state与load_state的 context 必须对称保存用FULL、加载用SHARDED或反之都会 KeyError。下面用可运行代码复现sharded optimizer state 的 key 形态与 save_state 期望的完整 key 不匹配。三、根因根因一句话save_state()默认按完整、可 pickle 的torch.Tensor来存状态而 FSDP2/TP 下状态是分片参数sharded keys且 TP 下是 DTensor不可 pickle两者冲突导致报错或保存后无法加载。三个具体失配state_dict 类型上下文缺失保存没进入SHARDED_STATE_DICT上下文读到的分片 key 与 load 期望的完整 key 不符。DTensor 不可torch.save序列化TP 下参数是 DTensortorch.save直接拒绝 pickle。save 与 load 的 state_dict_type 不对称一边 FULL 一边 SHARDED导致 KeyError/shape mismatch。四、最小可运行复现用纯 Python 模拟分片优化器状态与完整优化器状态的 key 形态差异复现save_state的 KeyError 机制import torch def sharded_optim_state(rank, world): 模拟 FSDP2 下每个 rank 持有的分片优化器状态 key。 keys {} for pid in range(4): # 4 个参数 if pid % world rank: # 该 rank 负责的分片 keys[foptimizer.state.{pid}.exp_avg] torch.randn(2) return keys def full_optim_state_expected(): 模拟 load 时期望的完整 key。 return {foptimizer.state.{pid}.exp_avg: None for pid in range(4)} def buggy_save_then_load(): world 2 # 保存每个 rank 只存自己分片真实里是各 rank 写各分片 saved [sharded_optim_state(r, world) for r in range(world)] # 加载期望完整 key 集合 expected full_optim_state_expected() # 合并分片真实里还需 all-gather这里只看 key 是否齐全 merged {} for s in saved: merged.update(s) missing [k for k in expected if k not in merged] if missing: raise KeyError(f加载时缺失 key: {missing} f(save 用分片 keyload 期望完整 key)) return ok def main(): try: print(buggy_save_then_load()) except KeyError as e: print(复现到报错:, e) if __name__ __main__: main()运行会打出复现到报错: 加载时缺失 key: ...——对应 FSDP2 下用分片态保存、却按完整态加载的 KeyError。五、解决方案第一层最小直接修复最立竿见影的修复对 FSDP2/TP保存时必须进入SHARDED_STATE_DICT上下文并用分布式检查点 APItorch.distributed.checkpoint而不是torch.save。加载也必须对称使用同一类型。import torch from torch.distributed.fsdp import FullyShardedDataParallel as FSDP from torch.distributed.fsdp import StateDictType from torch.distributed.checkpoint import save, load, FileSystemWriter, FileSystemReader def save_fsdp2_checkpoint(model, optimizer, output_dir): # 关键修复用 SHARDED_STATE_DICT 上下文 分布式检查点 with FSDP.state_dict_type(model, StateDictType.SHARDED_STATE_DICT): model_sd model.state_dict() optim_sd optimizer.state_dict() save( {model: model_sd, optimizer: optim_sd}, checkpoint_idoutput_dir, storage_writerFileSystemWriter(output_dir), ) def load_fsdp2_checkpoint(model, optimizer, input_dir): # 加载必须与保存的 state_dict_type 对称 with FSDP.state_dict_type(model, StateDictType.SHARDED_STATE_DICT): model_sd model.state_dict() optim_sd optimizer.state_dict() load( {model: model_sd, optimizer: optim_sd}, checkpoint_idinput_dir, storage_readerFileSystemReader(input_dir), ) model.load_state_dict(model_sd) optimizer.load_state_dict(optim_sd)第一层修复让保存/加载都走分片态 分布式检查点消除 DTensor 不可 pickle 与 key 不匹配。六、解决方案第二层结构性改进把保存必须进入正确的 state_dict_type、且与加载对称收口成一个CheckpointSpec让save_state/load_state永远成对、永远带正确的 context避免有人手滑用FULL保存、SHARDED加载。import torch from dataclasses import dataclass, field from typing import Dict from torch.distributed.fsdp import StateDictType dataclass class CheckpointSpec: state_dict_type: StateDictType StateDictType.SHARDED_STATE_DICT # 记录保存时用的类型加载时必须一致 _saved_as: StateDictType field(defaultNone, initFalse) def begin_save(self, model): self._saved_as self.state_dict_type return torch.distributed.fsdp.FSDP.state_dict_type(model, self.state_dict_type) def begin_load(self, model): if self._saved_as is not None and self._saved_as ! self.state_dict_type: raise ValueError( f加载类型 {self.state_dict_type} 与保存类型 {self._saved_as} 不对称 f会导致 KeyError ) return torch.distributed.fsdp.FSDP.state_dict_type(model, self.state_dict_type) # 用法示意在真实分布式进程里 def main(): spec CheckpointSpec(state_dict_typeStateDictType.SHARDED_STATE_DICT) # 保存 # with spec.begin_save(model): # save(...) # 加载 # with spec.begin_load(model): # load(...) print(CheckpointSpec 已约束 save/load 类型对称) if __name__ __main__: main()第二层的关键是_saved_as记录保存类型加载时若不一致直接拒绝把save/load 不对称这个最容易犯的错误挡在运行前。七、解决方案第三层断言 / CI 守护加 pytest 守护(1) 分片保存的 key 集合经合并后必须等于完整 key 集合无缺失(2)CheckpointSpec拒绝 save/load 类型不对称。用单进程模拟 key 合并import pytest def sharded_keys(rank, world): return {foptimizer.state.{pid}.exp_avg: pid for pid in range(4) if pid % world rank} def merge_sharded(keys_list): merged {} for k in keys_list: merged.update(k) return merged def test_no_missing_keys_after_merge(): world 2 saved [sharded_keys(r, world) for r in range(world)] merged merge_sharded(saved) expected {foptimizer.state.{pid}.exp_avg for pid in range(4)} assert set(merged.keys()) expected # 合并不缺失任何 key def test_checkpoint_spec_rejects_asymmetric(): class FakeSpec: def __init__(self): self.saved_as None def begin_save(self, t): self.saved_as t def begin_load(self, t): if self.saved_as is not None and self.saved_as ! t: raise ValueError(save/load 不对称) spec FakeSpec() spec.begin_save(SHARDED) with pytest.raises(ValueError): spec.begin_load(FULL) # 不对称必须被拒 if __name__ __main__: pytest.main([__file__, -q])CI 里test_no_missing_keys_after_merge通过保证分片保存合并后 key 齐全test_checkpoint_spec_rejects_asymmetric守护 save/load 对称。八、排查清单accelerate.save_state()在 FSDP2/TP 下报错按此顺序查先看报错在哪个阶段是torch.save报 Cannot pickle DTensorTP 问题还是 KeyErrorkey 形态不匹配还是 OOM用 FULL 拉全量状态。确认是否在正确的 state_dict_type 上下文里保存FSDP2 应用SHARDED_STATE_DICT不要裸调state_dict()。TP 必须换分布式检查点DTensor 不能torch.save改用torch.distributed.checkpoint.save/load。保存与加载类型必须对称保存SHARDED就加载SHARDED不要一边 FULL 一边 SHARDED。检查 save_state 的底层后端确认accelerate版本支持 FSDP2 的分片保存老版本save_state只认 DDP 的完整态需升级或绕过自己写 save 逻辑。验证合并后 key 齐全多 rank 各存分片reload 时确认所有 key 都被合并回来无缺失。OOM 排查若保存时 OOM基本是误用了 FULL_STATE_DICT 触发全量 all-gather切到 SHARDED 即解。九、小结accelerate.save_state()在 FSDP2/TP 下报错根因不在保存这个动作本身而在**save_state默认按完整、可 pickle 的普通 Tensor来存状态而 FSDP2/TP 下状态是分片的、TP 下参数是 DTensor**分片 key 与完整 key 形态不符导致 KeyErrorDTensor 无法torch.save导致 pickle 失败误用 FULL_STATE_DICT 又会导致保存时 OOM。修复三层第一层保存/加载都进入SHARDED_STATE_DICT上下文并用分布式检查点 API第二层用CheckpointSpec记录保存类型、加载时若不对称直接拒绝第三层用 pytest 断言分片合并 key 齐全、save/load 类型对称。记住FSDP2/TP 的 checkpoint 不是torch.save能存的分片态保存、分布式格式存、对称加载三件事缺一不可。