本文基于 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 把梯度同步一致。

对于一个有 PP 个参数的模型、NN 张卡,每张卡的显存占用是:

显存项 DDP 每卡占用
参数(FP16 + FP32 master) 6P\approx 6P bytes
梯度 2P4P\approx 2P \sim 4P bytes
AdamW 优化器状态(m,vm, v 8P12P\approx 8P \sim 12P bytes
合计(粗算) 16P\approx 16P bytes

注意 DDP 下这 16P16P每张卡都全量持有的——增加卡数只能加大 batch,单卡显存完全不下降。所以 175B 的 GPT-3 仅权重+优化器就需要 TB 级显存,单卡 80G 根本装不下,DDP 直接出局。

1.2 ZeRO 的思路:把冗余切掉

DDP 的冗余在哪?在于「每张卡都存了一份一模一样的优化器状态/梯度/参数」。微软 DeepSpeed 的 ZeRO(Zero Redundancy Optimizer)论文核心洞察就是:这些冗余完全可以切分到不同卡上,需要时再通信拼回来

ZeRO 分三个阶段,依次切分:

  • ZeRO-1:只切分优化器状态(m,vm, v)→ 显存 4P+4P/N\approx 4P + 4P/N
  • ZeRO-2:切分优化器状态 + 梯度 → 显存 4P+2P+8P/N\approx 4P + 2P + 8P/N
  • ZeRO-3:参数 + 梯度 + 优化器状态全切 → 显存 16P/N\approx 16P/N

ZeRO-3 已经把单卡显存压到 16P/N16P/N,理论上加卡就能训更大模型。PyTorch 官方的 **FSDP(Fully Sharded Data Parallel)**本质上就是 ZeRO-3 思想在 PyTorch 原生生态里的工程化实现,名字来自「Fully Sharded」——所有参数张量都被完全切分(shard)。

面试要点:FSDP ≈ ZeRO-3。区别在于工程实现与生态:FSDP 是 PyTorch 原生、支持 meta device 延迟初始化、与 torch.compile/混合精度/activation checkpoint 深度集成。


二、FSDP 的核心机制:一切分片,按需聚合

2.1 分片粒度:intra-tensor(张量内分片)

这是 FSDP 与朴素 ZeRO-3 实现的关键区别之一。看 Tiny-FSDP 的 fsdp_partition_tensors

1
2
3
4
5
6
7
8
9
10
# tiny_fsdp/core/fsdp/partition.py
for name, tensor in tensors_dict.items():
original_shape = tensor.shape
full_shapes[name] = original_shape
# 沿 dim-0 切分,每个 rank 拿一个切片
dim0_size = original_shape[0]
shard_size = (dim0_size + world_size - 1) // world_size # 向上取整
start_idx = rank * shard_size
end_idx = min(start_idx + shard_size, dim0_size)
sharded_tensor = tensor[start_idx:end_idx] # 本 rank 的分片

张量内分片(intra-tensor sharding):每个参数张量都沿第 0 维均匀切成 NN 份,每卡只存一份。对应地,Tiny-FSDP 里 ZeRO-3 用的是 inter-tensor(张量间分片)——以整个 tensor 为粒度,把不同 tensor 分给不同 rank 持有:

1
2
3
# tiny_fsdp/core/zero3/partition.py —— 按 tensor 整块分配给 owner rank
parts[current_part].append((name, tensor.numel()))
part_assignment[name] = current_part
分片策略 粒度 负载均衡 适合场景
ZeRO-3(inter-tensor) 整个张量 依赖 tensor 大小分布,可能不均 层间大小差异大的模型
FSDP(intra-tensor) 张量 dim-0 切片 天然均匀 层内均匀的大模型

面试要点:张量内分片让负载天然均衡(每卡都拿到每个 tensor 的 1/N1/N),代价是每次 forward/backward 都要通信聚合整个 tensor。而 ZeRO-3 的张量间分片,小 tensor 不用切,但容易出现某卡持有很多大 tensor、另一卡空闲的不均衡。

2.2 前向:all-gather 拼回完整参数

参数被切分了,前向计算需要完整参数怎么办?用之前临时 all-gather 拼回来。看 gather_tensor

1
2
3
4
5
6
7
8
# tiny_fsdp/core/fsdp/partition.py
def gather_tensor(sharded_tensor, full_shape, ...):
# 为每个 rank 预留接收 buffer
all_shards = [torch.empty(shard_shape, ...) for r in range(world_size)]
dist.all_gather(all_shards, sharded_tensor) # 一次 all-gather
valid_shards = [s for s in all_shards if s.numel() > 0]
full_tensor = torch.cat(valid_shards, dim=0) # 拼成完整参数
return full_tensor

前向时每个 Linear 层的 forward_callback 都会调用 all_gather_param 把本 rank 的 weight 分片聚合成完整 weight,再做矩阵乘:

1
2
3
4
5
6
# tiny_fsdp/core/fsdp/module.py
def forward_callback(self, ctx, input, weight, bias, runtime_tuner):
weight_full = all_gather_param(weight, weight.full_shape) # all-gather
ctx.weight_param = weight # 保存对分片参数的引用,反向来 reduce-scatter
output = ops.linear_forward(input, weight_full, bias_full, runtime_tuner)
return ctx, output

2.3 反向:reduce-scatter 边聚边切

反向传播需要两件事:① 把本 batch 的局部梯度跨卡求和(reduce),② 结果只存自己负责的那片(scatter)。FSDP 用 reduce-scatter 把这两步合成一个集合通信原语:

1
2
3
4
5
6
# tiny_fsdp/core/fsdp/partition.py
def scatter_tensor(full_tensor, ...):
local_shard = full_tensor[start_idx:end_idx].clone() # 本 rank 的目标分片
all_shards = [full_tensor[r_start:r_end] for r in range(world_size)]
dist.reduce_scatter(local_shard, all_shards, op=ReduceOp.SUM) # 一步搞定
return local_shard

reduce_scatter 语义:把各 rank 的 all_shards 对应分片求和后,结果写到各 rank 的 local_shard。等价于「先 all-reduce 再切片」,但通信量减半——只传 P/NP/N 而不是 PP

反向 backward_callback 的核心流程:

1
2
3
4
5
6
# 1. 先 all-gather 拼回完整参数(反向也需要完整 W 来算 dW 和 dX)
# 2. 用完整张量算 grad_weight_full / grad_input
# 3. reduce-scatter 得到本 rank 负责的梯度分片 grad_weight
grad_weight = reduce_scatter_grad(grad_weight_full, ctx.weight_param)
# 4. 立即清掉缓存的完整参数,释放显存
clear_param_cache(ctx.weight_param)

2.4 优化器:只更新本 rank 的分片

这是 FSDP 显存省到底的关键:优化器状态也只为本 rank 持有的参数分片分配。看 FSDPAdamW._init_opt

1
2
3
4
# tiny_fsdp/core/fsdp/optim.py
for name, param in self.parameters.items(): # param 已经是分片
self.moments[name] = torch.zeros_like(param, device=device) # 跟分片一样大
self.velocities[name] = torch.zeros_like(param, device=device)

param 此刻是分片(P/NP/N 大小),所以 m,vm, v 也是 P/NP/Nstep 里每卡独立更新自己的分片,无需任何通信

1
2
3
4
def _step_fn(self):
for name, param in self.parameters.items():
param = self.one_step(name, param) # 本地更新
self._zero_grad(param)

面试要点: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,通信量 ΦP\Phi \approx P(参数量级)。
  • FSDP:每层前向 all-gather + 每层反向 reduce-scatter,通信量 2P\approx 2P,但被分摊到每层的计算之间,可以 overlap。
  • 关键 insight:FSDP 的通信量约为 DDP 的 2 倍,但单卡显存从 16P16P 降到 16P/N16P/N用通信换显存,且通信可被计算掩盖。

3.2 显存账(FP16 + AdamW,约 16 bytes/param)

策略 参数 梯度 优化器状态 每卡合计
DDP PP PP PP(含 m,vm,v 折算) 16P\sim 16P
ZeRO-3 P/NP/N P/NP/N P/NP/N 16P/N\sim 16P/N
FSDP P/NP/N P/NP/N P/NP/N 16P/N\sim 16P/N

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
2
3
4
# tiny_fsdp/core/fsdp/wrapper.py
with torch.device('meta'):
meta_tensors = OrderedDict(model.named_parameters()) # 不占真实显存
self.sharded_tensors, self.full_shapes = fsdp_partition_tensors(meta_tensors, ...)

meta tensor 只有 shape/dtype、不分配显存。先用它规划好分片方案,再由各 rank 只实例化自己那片真实参数。这是 PyTorch FSDP 训练百 B 模型的前提。

4.2 rank 0 初始化 + broadcast 保证一致性

各 rank 独立随机初始化会导致参数不一致。Tiny-FSDP 的做法:rank 0 算完整参数,broadcast 给所有 rank,再各自切片:

1
2
3
4
5
# tiny_fsdp/core/fsdp/module.py :: Linear.reinit_parameters
if dist.get_rank() == 0:
nn.init.kaiming_uniform_(weight_full, a=math.sqrt(5))
dist.broadcast(weight_full, src=0) # 同步
weight_shard = init_shard_from_full(weight_full) # 各 rank 取自己的片

注意 broadcast 全量再切,对超大模型仍有初始化峰值显存问题。生产级 FSDP 会用「分 tensor 逐个 broadcast + 立即 shard」来规避,这是常见的 follow-up。

4.3 全参数缓存与及时清理

all_gather_param 里有个细节:同一参数在一次 forward/backward 内缓存完整张量,避免重复 all-gather:

1
2
3
4
5
6
def all_gather_param(param, full_shape, async_op=False):
if hasattr(param, '_fsdp_full_tensor_cache'):
return param._fsdp_full_tensor_cache # 命中缓存
full_tensor = gather_tensor(param.data, full_shape)
param._fsdp_full_tensor_cache = full_tensor
return full_tensor

而反向结束后必须 clear_param_cache 立刻释放,否则完整参数常驻显存等于没省。缓存生命周期 = 该层一次前向/反向,这是 FSDP 显存控制的核心纪律。

4.4 用 torch.autograd.Function 钩住前向/反向

FSDP 的 all-gather/reduce-scatter 必须卡在前向和反向的精确位置。Tiny-FSDP 用自定义 autograd Function 把通信嵌进计算图:

1
2
3
4
5
6
7
8
9
10
11
# tiny_fsdp/core/module/linear.py
def _ApplyLinearFunc(runtime_tuner, forward_callback, backward_callback):
class LinearFunc(torch.autograd.function.Function):
@staticmethod
def forward(ctx, input, weight, bias=None):
ctx, output = forward_callback(ctx, input, weight, bias, runtime_tuner)
return output
@staticmethod
def backward(ctx, grad_output):
return backward_callback(ctx, grad_output, runtime_tuner)
return LinearFunc.apply

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
2
3
4
5
6
7
8
# 保存:临时 gather 全量 → state_dict → 恢复分片
for name, param in self.module.named_parameters():
full_param = all_gather_param(param, param.full_shape)
param.data = full_param
state_dict = self.module.state_dict(...)
# 恢复 sharded
for name, param in ...:
param.data = gathered_params[name]; clear_param_cache(param)

生产级 FSDP 通常存「分片 checkpoint」(每 rank 只存自己的 shard,加载更省显存),这是更进阶的优化点。


五、通信-计算 overlap:性能优化的主战场

FSDP 通信量比 DDP 大,不 overlap 的话吞吐会很难看。这是 Infra 岗的高频考点。

5.1 反向里的 overlap 模式

Tiny-FSDP 在 DDP/ZeRO-3 里展示了 overlap 范式——异步发起通信,同时算下一部分,最后 wait()

1
2
3
4
# tiny_fsdp/core/zero3/module.py :: backward_callback
handle_weight = sync_grad(grad_weight, async_op=True) # 异步 reduce
grad_input = ops.linear_input_grad(...) # 同时算 dX
handle_weight.wait() # 用前再等

发起 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
显存(每卡) 16P16P 16P/N16P/N 16P/N16P/N
通信量 P\sim P P\sim P(broadcast/reduce) 2P\sim 2P(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 一次参数?
反向算 L/W\partial L/\partial WL/X\partial L/\partial X 都需要完整 WW,而 WW 平时是分片存的,所以反向开始得重新 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 切成 NN 片分存各卡(省显存);前向 all-gather 拼回计算,反向 reduce-scatter 边求和边切回分片,优化器只更新本卡分片(零通信);通信被嵌进 autograd 的逐层 backward 中、用异步 overlap 藏进计算里。

把握住这几条主线——切分粒度、三阶段通信(all-gather/reduce-scatter/无)、显存账 16P/N16P/N、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 容易太多,是建立直觉的最佳起点。

参考资料