torch.distributed中的DDP和FSDP的区别解释PyTorch的torch.distributed中的DDP和FSDP的区别
DDP 是数据并行:每张卡上有一份完整模型,每轮各卡用不同数据算 forward 和 backward;backward 结束后,各卡上的梯度要同步(通常 all-reduce),再在本卡用同一份梯度更新参数,所以每卡参数始终保持一致。显存占用 ≈ 单卡模型 + 单卡激活 + 梯度,模型必须能放进单卡。
FSDP 是分片数据并行:把模型参数、梯度、有时还有优化器状态按卡分片,每张卡只存 1/N(N=卡数);forward 时用 all-gather 把当前层所需参数临时拼起来算,算完丢掉;backward 同样 all-gather → 算梯度 → reduce-scatter 把梯度按片回写。这样单卡显存 ≈ 1/N 模型 + 当前层激活,可以训「单卡放不下」的大模型。
| 维度 | DDP | FSDP |
|---|---|---|
| 参数存储 | 每卡一份完整模型 | 每卡 1/N 参数(分片) |
| 梯度 | 每卡一份完整梯度,all-reduce | 分片,reduce-scatter 写回 |
| 优化器状态 | 每卡一份完整 | 可只存本卡分片(显存再省) |
| 单卡显存 | 约 1× 模型 + 激活 + 梯度 | 约 1/N 模型 + 激活(可训大模型) |
| 通信 | backward 后梯度 all-reduce | 每层 forward/backward 的 all-gather/reduce-scatter |
| 典型场景 | 模型能放进单卡、多卡加速 | 模型太大单卡放不下、大模型训练 |
当模型参数量大(如数十 B、上百 B),即使用上梯度 checkpoint,单卡也存不下「完整参数 + 梯度 + 优化器」。FSDP 用分片换通信:每次只 all-gather 当前层需要的参数,算完就丢,所以显存从「整模型」变成「1/N 模型 + 当前层激活」,能显著扩大可训模型规模;代价是每层多一次 all-gather/reduce-scatter,通信量比 DDP 的「一次梯度 all-reduce」大,需要好的通信和 overlap 设计。
| 返回模块 | 返回总览 |