从源码理解 FSDP:大模型分布式训练的显存破局之道
本文基于 Tiny-FSDP 项目源码整理。该项目用极简的 PyTorch 代码同时实现了 DDP、ZeRO-3、FSDP 三种分布式策略,是把大模型分布式训练「拆开看」的绝佳教材。我借它准备大模型 Infra 岗面试,把 FSDP 的每一个关键点都讲透。
一、为什么需要 FSDP:从 DDP 的显存瓶颈说起
https://zhuanlan.zhihu.com/p/694288870
https://zhuanlan.zhihu.com/p/2010127853522540210
面试时第一个常问的问题就是:「DDP 为什么训不动大模型?」
1.1 DDP 的显存模型
DDP(Distributed Data Parallel)是最朴素的数据并行:每张卡上都有一份完整的模型参数 + 梯度 + 优化器状态,只是各卡处理不同的数据 batch,反向传播后用 all-reduce 把梯度同步一致。
对于一个有 个参数的模型、 张卡,每张卡的显存占用是:
| 显存项 | DDP 每卡占用 |
|---|---|
| 参数(FP16 + FP32 master) | bytes |
| 梯度 | bytes |
| AdamW 优化器状态() | bytes |
| 合计(粗算) | bytes |
注意 DDP 下这 是每张卡都全量持有的——增加卡数只能加大 batch,单卡显存完全不下降。所以 175B 的 GPT-3 仅权重+优化器就需要 TB 级显存,单卡 80G 根本装不下,DDP 直接出局。
1.2 ZeRO 的思路:把冗余切掉
DDP 的冗余在哪?在于「每张卡都存了一份一模一样的优化器状态/梯度/参数」。微软 DeepSpeed 的 ZeRO(Zero Redundancy Optimizer)论文核心洞察就是:这些冗余完全可以切分到不同卡上,需要时再通信拼回来。
ZeRO 分三个阶段,依次切分:
- ZeRO-1:只切分优化器状态()→ 显存
- ZeRO-2:切分优化器状态 + 梯度 → 显存
- ZeRO-3:参数 + 梯度 + 优化器状态全切 → 显存
ZeRO-3 已经把单卡显存压到 ,理论上加卡就能训更大模型。PyTorch 官方的 **FSDP(Fully Sharded Data Parallel)**本质上就是 ZeRO-3 思想在 PyTorch 原生生态里的工程化实现,名字来自「Fully Sharded」——所有参数张量都被完全切分(shard)。
面试要点:FSDP ≈ ZeRO-3。区别在于工程实现与生态:FSDP 是 PyTorch 原生、支持
metadevice 延迟初始化、与torch.compile/混合精度/activation checkpoint 深度集成。
二、FSDP 的核心机制:一切分片,按需聚合
2.1 分片粒度:intra-tensor(张量内分片)
这是 FSDP 与朴素 ZeRO-3 实现的关键区别之一。看 Tiny-FSDP 的 fsdp_partition_tensors:
1 | # tiny_fsdp/core/fsdp/partition.py |
张量内分片(intra-tensor sharding):每个参数张量都沿第 0 维均匀切成 份,每卡只存一份。对应地,Tiny-FSDP 里 ZeRO-3 用的是 inter-tensor(张量间分片)——以整个 tensor 为粒度,把不同 tensor 分给不同 rank 持有:
1 | # tiny_fsdp/core/zero3/partition.py —— 按 tensor 整块分配给 owner rank |
| 分片策略 | 粒度 | 负载均衡 | 适合场景 |
|---|---|---|---|
| ZeRO-3(inter-tensor) | 整个张量 | 依赖 tensor 大小分布,可能不均 | 层间大小差异大的模型 |
| FSDP(intra-tensor) | 张量 dim-0 切片 | 天然均匀 | 层内均匀的大模型 |
面试要点:张量内分片让负载天然均衡(每卡都拿到每个 tensor 的 ),代价是每次 forward/backward 都要通信聚合整个 tensor。而 ZeRO-3 的张量间分片,小 tensor 不用切,但容易出现某卡持有很多大 tensor、另一卡空闲的不均衡。
2.2 前向:all-gather 拼回完整参数
参数被切分了,前向计算需要完整参数怎么办?用之前临时 all-gather 拼回来。看 gather_tensor:
1 | # tiny_fsdp/core/fsdp/partition.py |
前向时每个 Linear 层的 forward_callback 都会调用 all_gather_param 把本 rank 的 weight 分片聚合成完整 weight,再做矩阵乘:
1 | # tiny_fsdp/core/fsdp/module.py |
2.3 反向:reduce-scatter 边聚边切
反向传播需要两件事:① 把本 batch 的局部梯度跨卡求和(reduce),② 结果只存自己负责的那片(scatter)。FSDP 用 reduce-scatter 把这两步合成一个集合通信原语:
1 | # tiny_fsdp/core/fsdp/partition.py |
reduce_scatter 语义:把各 rank 的 all_shards 对应分片求和后,结果写到各 rank 的 local_shard。等价于「先 all-reduce 再切片」,但通信量减半——只传 而不是 。
反向 backward_callback 的核心流程:
1 | # 1. 先 all-gather 拼回完整参数(反向也需要完整 W 来算 dW 和 dX) |
2.4 优化器:只更新本 rank 的分片
这是 FSDP 显存省到底的关键:优化器状态也只为本 rank 持有的参数分片分配。看 FSDPAdamW._init_opt:
1 | # tiny_fsdp/core/fsdp/optim.py |
param 此刻是分片( 大小),所以 也是 。step 里每卡独立更新自己的分片,无需任何通信:
1 | def _step_fn(self): |
面试要点:FSDP 的 optimizer step 阶段零通信。因为梯度已经在反向时 reduce-scatter 到各卡,各卡手里的就是「全局梯度在本分片上的和」,直接更新即可。这正是 ZeRO 论文里 optimizer state 切分省显存的来源。
三、通信代价与显存账:一笔要算清的账
面试官最爱追的:「FSDP 省了显存,那代价是什么?」答:通信量增加了。
3.1 通信复杂度对比
| 阶段 | DDP | ZeRO-3 | FSDP |
|---|---|---|---|
| 前向 | 无 | broadcast(P) 每个参数张量 | all-gather(P) 每层 |
| 反向 | all-reduce(P) 梯度 | reduce(P) 梯度到 owner | reduce-scatter(P) 每层 |
| 优化器 | 无 | 无 | 无 |
- DDP:只在反向末尾做一次 all-reduce,通信量 (参数量级)。
- FSDP:每层前向 all-gather + 每层反向 reduce-scatter,通信量 ,但被分摊到每层的计算之间,可以 overlap。
- 关键 insight:FSDP 的通信量约为 DDP 的 2 倍,但单卡显存从 降到 。用通信换显存,且通信可被计算掩盖。
3.2 显存账(FP16 + AdamW,约 16 bytes/param)
| 策略 | 参数 | 梯度 | 优化器状态 | 每卡合计 |
|---|---|---|---|---|
| DDP | (含 折算) | |||
| ZeRO-3 | ||||
| FSDP |
ZeRO-3 与 FSDP 显存级别相同,差别在分片粒度与负载均衡。Tiny-FSDP README 实测 GPT-2 117M / 2×4090:DDP 2.1GB、FSDP 1.8GB,省显存明显;速度 4.5 vs 4.9 it/s,通信开销可接受。
四、工程实现里的关键细节(面试加分项)
光懂原理不够,Infra 岗会问你「实现层面踩过哪些坑」。Tiny-FSDP 里这些细节就是答案。
4.1 meta device 延迟初始化
大模型如果先在每卡实例化完整参数再切分,初始化那一下就会 OOM。FSDP 的标准做法是 torch.device('meta'):
1 | # tiny_fsdp/core/fsdp/wrapper.py |
meta tensor 只有 shape/dtype、不分配显存。先用它规划好分片方案,再由各 rank 只实例化自己那片真实参数。这是 PyTorch FSDP 训练百 B 模型的前提。
4.2 rank 0 初始化 + broadcast 保证一致性
各 rank 独立随机初始化会导致参数不一致。Tiny-FSDP 的做法:rank 0 算完整参数,broadcast 给所有 rank,再各自切片:
1 | # tiny_fsdp/core/fsdp/module.py :: Linear.reinit_parameters |
注意 broadcast 全量再切,对超大模型仍有初始化峰值显存问题。生产级 FSDP 会用「分 tensor 逐个 broadcast + 立即 shard」来规避,这是常见的 follow-up。
4.3 全参数缓存与及时清理
all_gather_param 里有个细节:同一参数在一次 forward/backward 内缓存完整张量,避免重复 all-gather:
1 | def all_gather_param(param, full_shape, async_op=False): |
而反向结束后必须 clear_param_cache 立刻释放,否则完整参数常驻显存等于没省。缓存生命周期 = 该层一次前向/反向,这是 FSDP 显存控制的核心纪律。
4.4 用 torch.autograd.Function 钩住前向/反向
FSDP 的 all-gather/reduce-scatter 必须卡在前向和反向的精确位置。Tiny-FSDP 用自定义 autograd Function 把通信嵌进计算图:
1 | # tiny_fsdp/core/module/linear.py |
forward_callback 里 all-gather 参数,backward_callback 里 reduce-scatter 梯度——通信与 autograd 反向传播顺序天然对齐(反正是从后往前逐层调用 backward)。
4.5 state_dict 的 gather/scatter 适配
保存 checkpoint 时,各 rank 只持有分片,得聚合成完整权重再存;加载时反向操作。Tiny-FSDP 的 FSDP.state_dict 临时 all-gather 出完整参数存盘,再恢复分片:
1 | # 保存:临时 gather 全量 → state_dict → 恢复分片 |
生产级 FSDP 通常存「分片 checkpoint」(每 rank 只存自己的 shard,加载更省显存),这是更进阶的优化点。
五、通信-计算 overlap:性能优化的主战场
FSDP 通信量比 DDP 大,不 overlap 的话吞吐会很难看。这是 Infra 岗的高频考点。
5.1 反向里的 overlap 模式
Tiny-FSDP 在 DDP/ZeRO-3 里展示了 overlap 范式——异步发起通信,同时算下一部分,最后 wait():
1 | # tiny_fsdp/core/zero3/module.py :: backward_callback |
发起 reduce 后立刻去算 grad_input,把通信藏在了计算里。FSDP 的 reduce-scatter 同理可异步化。
5.2 分层 all-gather 的 overlap
更重要的优化:不要一次 all-gather 整个模型,而是逐层 all-gather + 计算重叠。PyTorch FSDP 的 use_orig_params、prefetcher、forward_prefetch 都是干这个的——当前层计算时,预取 all-gather 下一层的参数。Tiny-FSDP 因为按层 forward_callback 天然是逐层 gather 的,已经具备分层 overlap 的雏形。
5.3 通信 bucketing 与 NCCL
生产级要点(面试常问):
- bucket:把多个小张量的 all-gather/reduce-scatter 合并成一次大通信,摊薄 launch 开销,撑满 NCCL 带宽。
- NCCL:NVLink/InfiniBand 上的 ring/tree all-gather,FSDP 的 all-gather 走的就是这条路。
- CPU offload(可选):把优化器状态甚至参数分片卸到 CPU 内存,进一步省 GPU 显存(FSDP
cpu_offload),代价是 PCIe 传输。
六、与其他策略的取舍:何时用 FSDP
| 维度 | DDP | ZeRO-3 | FSDP |
|---|---|---|---|
| 显存(每卡) | |||
| 通信量 | (broadcast/reduce) | (all-gather+reduce-scatter) | |
| 负载均衡 | 天然均衡 | tensor 间,可能不均 | 张量内,天然均衡 |
| 工程复杂度 | 低 | 中(需参数路由表) | 中高(需 autograd 钩子+分片管理) |
| 适合模型规模 | 中小(单卡装得下) | 大、层间不均 | 超大、层内均匀 |
| 生态 | 最成熟 | DeepSpeed | PyTorch 原生、与 compile/ckpt 集成 |
选型口诀:
- 单卡能装下完整模型 → DDP,简单高效。
- 模型大但层间大小差异大 / 需要 offload 灵活 → ZeRO-3(DeepSpeed)。
- 超大模型 + PyTorch 技术栈 + 想要原生 compile/activation-checkpoint 集成 → FSDP。
进阶:3D 并行里 FSDP/ZeRO-3 是数据并行维度,与 TP(Tensor Parallel,切层内矩阵)+ PP(Pipeline Parallel,切层间)正交。千卡训百 B 模型通常是 FSDP × TP × PP 三维组合。
七、面试高频问题速答
Q1:FSDP 和 ZeRO-3 是什么关系?
本质相同——都是把参数/梯度/优化器状态全切分。FSDP 是 PyTorch 原生实现,ZeRO-3 是 DeepSpeed 实现。工程细节(分片粒度、初始化、checkpoint 格式)有差异。
Q2:FSDP 为什么用 all-gather + reduce-scatter 而不是 broadcast + reduce?
all-gather 让每卡都拿到完整参数(前向每卡都要算),reduce-scatter 把「求和 + 取本卡分片」合成一步、通信量减半。broadcast/reduce 是 ZeRO-3 的 inter-tensor 模式,依赖 owner rank,负载不均。
Q3:FSDP 反向时为什么要再 all-gather 一次参数?
反向算 和 都需要完整 ,而 平时是分片存的,所以反向开始得重新 all-gather 拼回(FSDP2 的 use_orig_params 和重计算策略可减少这次 gather)。
Q4:FSDP 的 optimizer step 为什么不用通信?
反向时 reduce-scatter 已经把「全局梯度在本分片上的和」分发到各卡,每卡手里就是正确梯度分片,本地更新即可。
Q5:FSDP 通信量比 DDP 大,为什么还更快/可行?
① 通信被分层并 overlap 进计算;② 加卡带来的显存红利 > 通信代价,否则大模型根本训不了;③ 现代 NVLink/IB 带宽足以让 all-gather 被计算掩盖。
Q6:FSDP 的主要性能瓶颈?
- all-gather 没被 overlap → 通信 stall;
- bucket 太小 → NCCL launch 开销大、带宽利用率低;
- layer 太小/太多 → 通信次数多;
- 没开 activation checkpoint → activation 显存爆。
Q7:FSDP2 相比 FSDP1 改了什么?
FSDP2 基于 per-parameter DTensor + torch.compile 友好的细粒度分片,解决了 FSDP1 的「整模块 flatten 后无法 use_orig_params 友好、与 compile 兼容差」等问题。面试能提到 FSDP2/DTensor 会很加分。
八、总结
从源码层面,FSDP 的全貌可以用一句话概括:
平时把每个参数张量沿 dim-0 切成 片分存各卡(省显存);前向 all-gather 拼回计算,反向 reduce-scatter 边求和边切回分片,优化器只更新本卡分片(零通信);通信被嵌进 autograd 的逐层 backward 中、用异步 overlap 藏进计算里。
把握住这几条主线——切分粒度、三阶段通信(all-gather/reduce-scatter/无)、显存账 、meta 延迟初始化、autograd 钩子、通信-计算 overlap——大模型 Infra 岗的 FSDP 题基本就能稳住。剩下的就是 FSDP2/DTensor、3D 并行组合、CPU offload 这些进阶话题,留到面试里再展开。
Tiny-FSDP 这个项目的妙处在于:DDP/ZeRO-3/FSDP 三套实现共用同一套 ops/module 基类,把三种策略的差异压缩到了 forward_callback/backward_callback 几十行代码里。对照着读,比读 PyTorch FSDP 那几万行 C++/Python 容易太多,是建立直觉的最佳起点。






