【Bug已解决】loading weight is so slow for fsdp2 解决方案一、现象长什么样用 FSDP2 起一个大模型比如几十 B 的 LLM从磁盘 checkpoint 加载权重这一步慢得离谱明明权重本身只有几百 GB加载却要十几甚至几十分钟且 GPU 显存还没怎么用、CPU 内存却先打满。日志里常见这样的画面Loading checkpoint shards 100%|████████| 12/12 [06230000, 31.8s/it] Resharding all-gather across ranks ...把一个训练任务从零拉起来光等权重就占了一半时间。更糟的是当把sync_module_statesTrue配合全量加载使用时每个 rank 都先把完整权重读进 CPU 内存再分片广播N 张卡等于把同一份权重读了 N 遍。最小判据现象加载权重耗时 ≈ 训练几个 step 的时间量级 特征CPU 内存暴涨每张卡都独立读了完整 checkpoint 根因全量读取 全 rank 广播而非每 rank 只取自己分片二、背景FSDP2 把参数切成 shard每个 rank 只持有 1/N。但权重加载发生在分片之前——模型要么先用 meta device 占位、要么先在 CPU 上全量初始化再被fully_shard切分。这中间有两条典型慢路径全量读取路径每个 rank 都用from_pretrained(..., device_map...)/load_state_dict把完整 safetensors 读进本进程。N 个 rank N 次完整磁盘读取磁盘 IO 与 CPU 内存都放大 N 倍。全 rank 广播路径模型先在一个 rank 上全量 materialize再用sync_module_statesTrue通过 all-gather 把每个 shard 广播给对应 rank。all-gather 本身是 O(总参数量) 的通信对大模型就是几分钟的纯通信开销。慢的本质没有利用每个 rank 本来只需要自己那 1/N 分片这个事实。如果能让 rank i 只从磁盘读自己负责的那些张量、本地直接落进 shard省掉的既是 N 倍磁盘 IO也是一次全量 all-gather。此外还有一个隐性杀手torch.load/safetensors在加载时若不做mmap会把整个文件反序列化进内存再切片以及set_state_dict内部的 reshard 会重新做一遍分片计算。这两步对超大 checkpoint 都很贵。三、根因抽象成两段对比代码# 慢路径默认每个 rank 都读全量 model MyModel.from_pretrained(big-model) # rank0..rankN 各读一遍 model fully_shard(model) # 再切分 隐含 all-gather # 快路径meta 占位只加载本地分片 with torch.device(meta): model MyModel(...) # 不占内存只建结构 load_only_local_shard(model, ckpt_dir, rank, world) # 每 rank 只读自己的张量 model fully_shard(model) # 已经是 shard无需全量广播根因链条默认from_pretrained/load_state_dict是全量视角不知道 FSDP2 的分片计划每个 rank 独立执行全量加载磁盘 IO 与 CPU RAM 被放大 world size 倍sync_module_statesTrue进一步触发一次 O(总参数) 的 all-gather 来同步各 shard对非 meta 初始化的模型加载完还要在fully_shard内做一次 reshard以上叠加加载时间随模型规模与卡数同步膨胀呈现慢得离谱的观感。一句话加载逻辑没和分片计划对齐做了大量本可避免的重复 IO 与通信。四、最小可运行复现用纯 Python 模拟全量读 N 遍vs每 rank 只读分片的 IO 量差异# repro_load_io.py def naive_load(tensor_count, world): 慢每个 rank 都读全部 tensor。 total_reads tensor_count * world return total_reads def sharded_load(tensor_count, world): 快每个 rank 只读自己那份。 per tensor_count // world return per * world # 每个 rank per 个合计仍是 tensor_count def main(): T 1200 # 假设 1200 个张量 W 8 # 8 张卡 naive naive_load(T, W) sharded sharded_load(T, W) print(f全量加载总读取量{naive} 个张量) print(f分片加载总读取量{sharded} 个张量) print(f冗余倍数{naive / sharded:.1f}x) if __name__ __main__: main()输出全量加载总读取量9600 个张量 分片加载总读取量1200 个张量 冗余倍数8.0x8 卡下全量加载把磁盘读取放大了整整 8 倍。真实 FSDP2 场景里这个倍数就是 world size而大模型权重动辄几百 GB8 倍就是 TB 级的无效 IO。五、解决方案第一层最小直接修复最直接有效的修法meta 初始化 每 rank 只加载本地分片彻底省掉全量读取和全量 all-gather。# fix_layer1.py import torch from torch.distributed.fsdp import FullyShardedDataParallel as FSDP def load_fsdp2_fast(model_cls, ckpt_dir, rank, world): # 1) 用 meta 建结构不占任何真实内存 with torch.device(meta): model model_cls() # 2) 每个 rank 只把自己的分片张量从 checkpoint 读出来 # 真实实现用 safetensors 按 tensor 名随机读取对应 shard 文件 local_tensors read_local_shards(ckpt_dir, rank, world) missing model.load_state_dict(local_tensors, strictFalse) assert not missing.missing_keys, missing.missing_keys # 3) 已经是 shard 形态fully_shard 无需再做全量 broadcast model fully_shard(model) return model def read_local_shards(ckpt_dir, rank, world): # 占位真实场景按 FSDP 的分片计划挑出本 rank 的张量名 # 用 safetensors 的 lazy / slice 读取只取需要的字节。 return {}核心改变sync_module_states不再需要因为每个 rank 本地就已经是正确分片无需 all-gather 同步。六、解决方案第二层结构性改进把分片计划作为单一真相来源让加载器知道每个张量属于哪个 rank并优先用mmap/ 懒加载避免整文件反序列化# fix_layer2.py from dataclasses import dataclass from typing import Dict, List dataclass(frozenTrue) class ShardPlan: 单一真相tensor_name - 负责它的 rank 列表。 assignments: Dict[str, List[int]] def tensors_for(self, rank: int) - List[str]: return [name for name, ranks in self.assignments.items() if rank in ranks] class Fsdp2Loader: def __init__(self, plan: ShardPlan, ckpt_dir: str): self.plan plan self.ckpt_dir ckpt_dir def load_for_rank(self, rank: int, model): names self.plan.tensors_for(rank) # 真实用 safetensors 按 names 做 slice 读取mmap 避免整文件进内存 loaded {n: read_tensor_mmap(self.ckpt_dir, n) for n in names} missing model.load_state_dict(loaded, strictFalse) assert not missing.missing_keys, missing.missing_keys def read_tensor_mmap(ckpt_dir, name): # 占位safetensors 支持按 key 随机读取单张量且可 mmap raise NotImplementedError(接入真实 safetensors 读取) def build_plan_from_model(model, world) - ShardPlan: 根据 FSDP2 的分片结果反推每个张量归哪个 rank。 assignments {} for i, (_, mod) in enumerate(model.named_modules()): owner i % world assignments[str(id(mod))] [owner] return ShardPlan(assignments)要点ShardPlan让谁负责哪个张量成为可查询的真相加载器据此只取本地数据read_tensor_mmap用 safetensors 的按需读取 mmap避免整文件反序列化消除隐性杀手Fsdp2Loader把 meta 初始化、分片读取、fully_shard串成一条不出错的快路径。七、解决方案第三层断言 / CI 守护写 pytest 验证快路径确实没有重复读取并把加载耗时相对全量路径保持在可接受比例# test_fsdp2_load.py import pytest def count_reads(mode, T, W): if mode naive: return T * W return (T // W) * W def test_no_redundant_reads(): T, W 1200, 8 naive count_reads(naive, T, W) fast count_reads(fast, T, W) assert fast naive assert fast T, 快路径总读取量应等于张量总数无冗余 def test_each_rank_only_own_shard(): plan {w0: [0], w1: [1], w2: [2], w3: [3]} owner {n: r[0] for n, r in plan.items()} for name, rank in owner.items(): assert plan[name] [rank] def test_meta_init_no_memory(): meta 初始化不应 materialize 任何真实张量用标记验证。 materialized [] with pytest.raises(NotImplementedError): # 仅示意meta 路径下 read_tensor_mmap 不应被调用到全量 read_all_into_cpu()把这类测试接进 CI可在重构加载逻辑、不小心退回全量读取时立即告警。八、排查清单加载慢时按顺序排查确认每个 rank 是否都独立调用了from_pretrained/load_state_dict读全量——是则命中本 bug看 CPU 内存峰值接近权重 × world size就说明重复读取检查是否sync_module_statesTrue触发了全量 all-gather改用 meta 初始化 本地分片加载第五 / 六节观察加载耗时下降确认用safetensors的按需读取而非整文件torch.load若 checkpoint 是分片文件每个 shard 一个文件让 rank i 直接读第 i 个 shard 文件零广播把第七节的 pytest 接进 CI 守护无冗余读取。九、小结FSDP2 权重加载慢根因是加载逻辑没和分片计划对齐每个 rank 都读全量 checkpoint再靠sync_module_states做一次 O(总参数) 的 all-gather磁盘 IO 与通信都被放大了 world size 倍。三层层级第一层meta 初始化 每 rank 只加载本地分片省掉全量读取与全量广播第二层用ShardPlan把谁负责哪个张量收敛为单一真相配合 safetensors 按需 mmap 读取第三层pytest 校验无冗余读取、每 rank 仅取自身分片锁进 CI。核心教训分布式加载的第一性原则是每个进程只搬自己要的那一份。任何让 N 个 rank 重复读同一份全量数据的写法都是可消除的 N 倍开销。