数据并行与 ZeRO / FSDP
朴素数据并行每卡保存完整的参数、梯度与优化器状态,显存占用为参数量的 16 倍(混合精度 Adam)。ZeRO 三阶段依次分片这三者,FSDP 是 PyTorch 的原生实现。本章推导各阶段的显存与通信量。
- ZeRO-1/2 几乎不增加通信,ZeRO-3 把 all-reduce 换成 all-gather + reduce-scatter
- FSDP 的分片单元(wrap policy)决定通信与计算的重叠程度
- 梯度累积与全局 batch 的关系
朴素数据并行每卡保存完整的参数、梯度与优化器状态,显存占用为参数量的 16 倍(混合精度 Adam)。ZeRO 三阶段依次分片这三者,FSDP 是 PyTorch 的原生实现。本章推导各阶段的显存与通信量。