Infra: Data Parallel


参考了 (Hashimoto & Liang, 2026; Tazi et al., 2025)

Naïve Data Parallelism

最拿衣服的 DP 大概就是每张卡存个模型,每张卡然后算出 gradient 然后 All-Reduce。但 All-Reduce 的时候 GPU 也没啥别的事情做。一个简单的优化方法是,反向传播的时候,我们每经过一个 block 就 async 地触发那一层地 All-Reduce。使用 torch.Tensor.register_post_accumulate_grad_hook 就可以做。

然后如果每个小 block 都触发一次 All-Reduce,就比较零散。所以优化方案是几个 block bucket 起来做一次 All-Reduce。

然后对于 gradient accumulation 的时候,optimizer.step() 的前一步才需要执行 All-Reduce。

Zero Redundancy Optimizer (ZeRO)

ZeRO:从复制模型状态到逐步分片四列分别代表 GPU 0 至 GPU 3;每张卡的色块从上到下是参数、梯度、优化器状态。 Baseline 在每张卡复制全部三类状态。ZeRO-1 只分片优化器状态,ZeRO-2 进一步分片梯度,ZeRO-3 再分片参数。分片后 GPU i 仅保留第 i 份,四张卡的分片合起来是一份完整状态。虚线框标出分片前的范围,留白部分不在本卡存储。 右侧逐行给出每张 GPU 的模型状态显存公式。GPU 0GPU 1GPU 2GPU 3Memory / GPU模型状态 · bytesBaseline全部复制
\(\displaystyle 2\Psi+2\Psi+k\Psi\)
ZeRO-1分片优化器状态
\(\displaystyle 2\Psi+2\Psi+\frac{k\Psi}{N_d}\)
ZeRO-2进一步分片梯度
\(\displaystyle 2\Psi+\frac{2\Psi+k\Psi}{N_d}\)
ZeRO-3进一步分片参数
\(\displaystyle \frac{2\Psi+2\Psi+k\Psi}{N_d}\)
左右滑动查看各 GPU 和显存公式
Parameters \(2\Psi\)Gradients \(2\Psi\)Optimizer states \(k\Psi\)
图示 \(N_d = 4\),每张卡保留的分片位置不同;虚线框中的留白由其他 GPU 持有。
\(\Psi\) 为参数量,\(N_d\) 为数据并行度。参数与梯度按 16-bit 计,优化器状态为每参数 \(k\) bytes;色块高度按混合精度 Adam 的 \(k = 12\) 绘制(FP32 主权重及一、二阶矩)。仅计模型状态,不含激活与临时通信缓冲区。

我们计算一下通信量。我们就不管 \(\frac{N_d-1}{N_d}\) 的系数了,假设一次 Reduce-Scatter 或者一次 All-Gather 是 \(2\Psi\),一次 All-Reduce 是 \(4\Psi\)

对于 Naïve DP,最后需要 All-Reduce gradient 也就是 \(4\Psi\)

对于 ZeRO-1,切分了 optimizer states。首先单张卡走 forward 和 backward。这里有一个小技巧,就是对于 AdamW,如果知道一部分梯度,就可以算出那部分的一阶和二阶 moment 从而来更新模型参数。因此我们可以直接一个 Reduce-Scatter 计算平均并分发梯度,每张卡只负责更新自己那一小部分模型。最终把更新好的模型 All-Gather 起来。于是通信量是 \(2\Psi+2\Psi=4\Psi\),其实相比于 Naïve DP 没有特别多损失。对于 Muon 就比较复杂,因为他没有 AdamW 的性质,这边暂时不讨论。

对于 ZeRO-2,切分了 optimizer states 和 gradient。事实上和 ZeRO-1 是一样的。因为我们不难观察到 ZeRO-1 中,backward reduce-scatter 之后,每张卡只需要保存部分梯度,别的梯度可以立刻释放掉。所以通信仍然是一次 Reduce-Scatter 一次 All-Gather,也就是 \(4\Psi\)

对于 ZeRO-3,把模型参数也给切分了。这时候每次 forward 我们需要按层 All-Gather 聚合模型参数,用完就释放完整参数,所以 backward 前还需要再次 All-Gather 本层参数。梯度做 Reduce-Scatter 后,每张卡更新并保留自己的参数分片,等下次 forward 时再聚合。于是通信量是两次参数 All-Gather 加一次梯度 Reduce-Scatter,也就是 \(2\Psi+2\Psi+2\Psi=6\Psi\)

同一个训练步,从上到下对比
ZeRO-1Optimizer states 分片
ZeRO-1:训练流程与 GPU 间通信ZeRO-1 的四张 GPU 从上到下经历前向、反向、优化器更新和下一步。 完整参数常驻,前向和反向无需参数 All-Gather。反向产生的梯度通过 Reduce-Scatter 求平均,每张卡负责一个归约后的梯度分片。 ZeRO-1 保留完整梯度缓冲区;描边标出本卡负责的已归约分片,其余淡色部分是本地梯度,不表示已得到完整平均梯度。 AdamW 只更新本卡参数分片,然后 All-Gather 更新后的参数,使每张卡在下一步再次持有完整模型。 实线箭头传参数,虚线箭头归约梯度;深色路径强调 GPU 0 的通信,其余 GPU 同时执行相同操作。GPU 0GPU 1GPU 2GPU 3迭代开始 · 各卡持有的状态FORWARD · 层 1 → L完整参数已在本卡 · 无通信Forward完整参数常驻 → 下一层BACKWARD · 层 L → 1完整参数仍在本卡 · 无通信BackwardReduce-Scatter · 平均梯度保留完整梯度缓冲区描边:本卡负责的已归约分片OPTIMIZER STEPAdamW · 更新本卡分片All-Gather · 更新后的参数完整模型就绪 → 下一步前向
每步通信量\(2\Psi + 2\Psi = 4\Psi\)梯度 RS + 更新后 AG
ZeRO-2再分片 Gradients
ZeRO-2:训练流程与 GPU 间通信ZeRO-2 的四张 GPU 从上到下经历前向、反向、优化器更新和下一步。 完整参数常驻,前向和反向无需参数 All-Gather。反向产生的梯度通过 Reduce-Scatter 求平均,每张卡负责一个归约后的梯度分片。 梯度就绪后按 bucket 归约,及时释放非本卡的梯度,仅保留本卡负责的已归约分片。 AdamW 只更新本卡参数分片,然后 All-Gather 更新后的参数,使每张卡在下一步再次持有完整模型。 实线箭头传参数,虚线箭头归约梯度;深色路径强调 GPU 0 的通信,其余 GPU 同时执行相同操作。GPU 0GPU 1GPU 2GPU 3迭代开始 · 各卡持有的状态FORWARD · 层 1 → L完整参数已在本卡 · 无通信Forward完整参数常驻 → 下一层BACKWARD · 层 L → 1完整参数仍在本卡 · 无通信BackwardReduce-Scatter · 平均梯度逐 bucket 归约 · 只保留本卡梯度其他梯度及时释放OPTIMIZER STEPAdamW · 更新本卡分片All-Gather · 更新后的参数完整模型就绪 → 下一步前向
每步通信量\(2\Psi + 2\Psi = 4\Psi\)梯度 RS + 更新后 AG
ZeRO-3再分片 Parameters
ZeRO-3:训练流程与 GPU 间通信ZeRO-3 的四张 GPU 从上到下经历前向、反向、优化器更新和下一步。 前向按层 All-Gather 参数,计算后释放完整参数。反向按逆序再次 All-Gather 本层参数,计算梯度,再 Reduce-Scatter 平均梯度并只保留本卡分片,释放完整参数。 梯度就绪后按 bucket 归约,及时释放非本卡的梯度,仅保留本卡负责的已归约分片。 AdamW 只更新本卡参数分片,更新后不聚合完整模型;下一次前向再按层 All-Gather。 实线箭头传参数,虚线箭头归约梯度;深色路径强调 GPU 0 的通信,其余 GPU 同时执行相同操作。GPU 0GPU 1GPU 2GPU 3迭代开始 · 各卡持有的状态FORWARD · 层 1 → LAll-Gather · 本层参数Forward用完释放完整参数 → 下一层BACKWARD · 层 L → 1再次 All-Gather · 本层参数BackwardReduce-Scatter · 平均梯度逐 bucket 归约 · 只保留本卡梯度释放本层完整参数 → 上一层OPTIMIZER STEPAdamW · 更新本卡分片更新后保留分片 · 无通信下一步前向时再按层聚合
每步通信量\(2\Psi + 2\Psi + 2\Psi = 6\Psi\)前向 AG + 反向 AG + 梯度 RS
左右滑动对比 ZeRO-1 / ZeRO-2 / ZeRO-3
Parameters · 实线传输Gradients · 虚线归约Optimizer states
每条色带分为 \(N_d = 4\) 片,GPU i 负责第 i 片;深色箭头突出 GPU 0 的通信,其余路径淡化。所有方法中,优化器状态始终只保留本卡分片;更新处仅画出各卡新更新的参数分片。ZeRO-1 的淡色梯度是保留的本地缓冲区,不表示完整平均梯度。
ZeRO-3 的前向、反向框分别对每层重复,仅临时聚合当前层参数。通信量沿用正文的 16-bit 口径,忽略 \(\frac{N_d-1}{N_d}\);逐层 / bucket 的通信量按整个模型求和,不表示只有一次 API 调用。图中不展开预取与通信计算重叠。

References

Hashimoto, T., & Liang, P. (2026). CS336: Language Modeling from Scratch. Stanford University, Spring 2026. cs336.stanford.edu
Tazi, N., Mom, F., Zhao, H., Nguyen, P., Mekkouri, M., von Werra, L., Wolf, T., & HuggingFace, O. (2025). The ultra-scale playbook: Training LLMs on GPU clusters. Hugging Face. huggingface.co

Cite this post

@misc{pu2026mlrevisitinfrafsdp,
  author = {Pu, Fanyi},
  title  = {Infra: Data Parallel},
  year   = {2026},
  month  = {9},
  url    = {https://pufanyi.com/blog/ml-revisit-infra-fsdp}
}