Infra: Data Parallel


· Updated

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

Naïve Data Parallelism

最拿衣服的 DP,大概就是每张卡存一份相同的模型,分别处理不同的数据,算出自己的 gradient,然后通过 All-Reduce 得到平均梯度。这样每张卡用同样的梯度更新参数,模型就能保持一致。但如果等整个 backward 结束才开始通信,通信就没法和这次 backward 的计算重叠。

一个简单的优化是:某些参数的梯度算好之后,就 async 地触发它们的 All-Reduce,同时继续计算其他梯度。不过,每个小 tensor 都单独通信会比较零散,所以可以把多个参数的梯度放进一个 bucket,等这个 bucket 的梯度都就绪后再一起通信。PyTorch DDP 已经实现了这套机制。

如果自己实现,可以用 torch.Tensor.register_post_accumulate_grad_hook 在某个参数的 .grad 就绪后收到通知。注意这是参数级的 hook,不是 block 级的;还需要自己管理 bucket,保证各张卡的通信顺序一致,并在 optimizer.step() 使用梯度前等待相关通信完成。

对于 gradient accumulation,可以先在本地累积几个 micro-batch 的梯度,只在最后一个 micro-batch 的 backward 中同步。在 DDP 中,可以把前几个 micro-batch 的 forward 和 backward 都放进 no_sync(),最后一次则正常执行。这里推迟的是梯度同步,不是由 optimizer.step() 自动触发通信。loss 的归一化也不能忘:例如 4 个等大的 micro-batch,每个 loss 都是样本均值,就可以先除以 4 再 backward。

Zero Redundancy Optimizer (ZeRO)

ZeRO 最初是在 (Rajbhandari et al., 2020) 中提出的,最初的实现来自 DeepSpeed (Rasley et al., 2020)。它的思路是:每张卡没必要一直保存一模一样的模型状态,可以把这些状态分给不同的卡保管,需要时再通信。

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 主权重及一、二阶矩)。仅计模型状态,不含激活与临时通信缓冲区。

我们计算一下通信量。这里 \(\Psi\) 是整个模型的参数量,\(N_d\) 是 data parallel 的卡数;参数和梯度都按 16-bit,也就是每个元素 2 bytes 来算。先只看一次 forward 和 backward,不考虑 gradient accumulation 或 hybrid parallelism。

按 ring 通信模型估算每张卡发送的数据量,忽略 \(\frac{N_d-1}{N_d}\) 的系数,对整个模型做一遍 Reduce-Scatter 或者 All-Gather,通信量约为 \(2\Psi\) bytes;一遍 All-Reduce 则约为 \(4\Psi\) bytes。这里的“一遍”是把所有层或 bucket 的通信量加起来,不是说只调用一次通信 API,也不把接收量再加一遍。

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

对于 ZeRO-1,只切分 optimizer states,模型参数仍然每张卡都有一份。这里有一个小技巧:AdamW 是逐元素更新的,只要有某部分参数、对应的平均梯度和历史 moment,就能更新这部分参数,不需要知道其他部分的梯度。因此可以通过 Reduce-Scatter 配合求平均,把梯度交给各自负责的卡;每张卡只更新自己的参数分片和 optimizer states,最后再 All-Gather 更新后的参数,让所有卡重新拿到同一份完整模型。通信量是 \(2\Psi+2\Psi=4\Psi\),和 Naïve DP 一样。

对于 Muon (Jordan et al., 2024) 就比较复杂,因为它的矩阵更新不能像 AdamW 那样,任意切成若干元素后各算各的。如何额外通信、还原所需的矩阵,(Liu et al., 2025) 有讨论,这边先不展开。

对于 ZeRO-2,进一步切分 gradient。它和 ZeRO-1 可以有相同的通信量,但梯度的保存方式不同:在 backward 过程中,一个 bucket 的梯度就绪后,就归约并只保留本卡负责的部分,其他部分及时释放。这样就不需要一直保留完整模型的梯度缓冲区。如果等整个 backward 结束、完整梯度都存下来了才释放,就已经承担了完整梯度的峰值显存。通信仍然是一遍 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 调用。图中不展开预取与通信计算重叠。

Fully Sharded Data Parallel (FSDP)

FSDP 的全分片模式和 ZeRO-3 的核心思路一样:把参数、梯度和 optimizer states 分给不同的卡,需要完整参数时再聚合。这边主要讨论 PyTorch FSDP (Zhao et al., 2023)。不过,FSDP 也支持其他分片策略,不能把所有配置都当成上面那套固定流程。

FSDP 主要接管参数、梯度的 sharding 和相关通信,尽量保持 PyTorch 原生的 training loop;DeepSpeed 则更偏完整的 training engine,因此做定制训练时,FSDP 通常更容易组合。但这不代表任意单卡代码都能原样运行,例如自定义 optimizer 仍要能正确处理分片参数。

FSDP1先打包,再切分
FSDP1:先拼成 FlatParameter,再切给两个 rank三个 shape 为 (4, 2) 的权重 W1、W2、W3 属于同一个 FSDP unit。拼接后共有 24 个元素,平均分给两个 rank。Rank 0 拿到完整 W1 和 W2 的前半部分;Rank 1 拿到 W2 的后半部分和完整 W3。每个本地 flat shard 有 12 个元素,分片边界穿过 W2。同组的 3 个权重 · shape (4, 2)W1W2W3flatten + concatenateW1W2W3FlatParameter · 从中间切成两段Rank 0W1W2flat shard · (12,)Rank 1W2W3flat shard · (12,)
一个 flat shard 可以跨过参数边界
FSDP2逐参数分片,按组通信
FSDP2:每个参数分别沿 dim 0 切给两个 rank同样的三个权重分别表示为 DTensor,沿 dim 0 切分。Rank 0 保存每个权重的前两行,Rank 1 保存后两行。每个 DTensor 的全局 shape 仍为 (4, 2),本地数据的 shape 为 (2, 2),每个 rank 总计 12 个元素。同组参数仍可合并通信。同组的 3 个权重 · shape (4, 2)W1W2W3每个参数分别用 DTensor 表示Shard(0) · 每个权重各取一半行Rank 0W1W2W3各取前 2 行 · local (2, 2)Rank 1W1W2W3各取后 2 行 · local (2, 2)
每个参数保留自己的 shape 与分片信息
左右滑动对比 FSDP1 / FSDP2
W1、W2、W3 分别代表三个 Linear 的 weight,同属一个通信分组;两边都分给 2 个 rank,每个 rank 都保存 12 个元素。虚线标出切分位置:FSDP1 的切口穿过 W2,FSDP2 则对每个权重分别按行切分,全局 shape 仍是 (4, 2)。图中只比较参数布局,省略 bias、padding 和计算时临时聚合的完整参数;箭头表示布局变化,不表示通信次数。

FSDP1:先打包,再切分

为了减少零散的小通信,FSDP1 会把一组参数一起管理,这一组就叫一个 FSDP unit。例如,我们可以把每个 Transformer block 划成一个 unit,将其中的参数 flatten、拼成一维的 FlatParameter,再平均切给不同的卡。这样,组内参数的 All-Gather 和梯度的 Reduce-Scatter 都可以合并进行 (Zhao et al., 2023)

unit 的大小决定了显存和通信之间的取舍:unit 越小,每次需要临时聚合的完整参数就越少,但通信次数也越多。Backward 之后,每张卡拿到自己负责的 gradient shard,再由 AdamW 使用本地的一、二阶 moment 更新对应的参数分片。每张卡的 optimizer 只需维护这一部分的 states。

这种“先打包再切分”的方式让通信布局很规整,但也把原始参数的边界打散了:一个 shard 可能包含几个参数的片段,而一个权重矩阵也可能被切到不同的卡上。因此,本地拿到的往往是一段一维数据,而不是原来形状的矩阵。

FSDP1 提供了两种访问这些参数的方式。use_orig_params=False 时,用户和 optimizer 直接看到 FlatParameter;设为 True 时,则保留原始 Parameter 对象,让它们的数据指向 flat buffer 中对应的部分。在 sharded 状态下,这些 view 可能只有原始参数的一截;某个 rank 没分到该参数的数据时,对应的 view 就是空 tensor。计算前,FSDP 再聚合参数并恢复原始 shape。

所以,FSDP1 很适合把一组参数打包后高效通信,但如果想按原始参数的 shape 做自定义操作,就需要处理这层 flat buffer 到原始参数的映射。FSDP2 的逐参数分片表示,让这件事更直接。

FSDP2:分别记录每个参数的分片

FSDP2 (Feng et al., 2025) 不再跨参数拼成一个 FlatParameter 后切分,而是给每个参数分别用 DTensor 表示分片。可以把 DTensor 理解成“本卡的数据,加上这份数据在完整 tensor 中的位置说明”。

例如,一个 shape 为 (8, 4) 的权重矩阵,沿第 0 维平均切给 4 张卡,每张卡保存 2 行。在分片状态下,param.shape 仍然是全局的 (8, 4),而 param.to_local().shape 是本地的 (2, 4)。保留的是完整参数的形状信息,不是每张卡仍然保存着完整参数。

另一个变化在 API 上:FSDP2 的 fully_shard 在原来的 module 上注册 hooks,不额外套一层 wrapper,所以参数的 FQN 保持不变。每个参数仍能按原来的名字找到,并带着自己的分片信息,这让参数操作和 checkpoint 更容易组合。

不过,分别记录参数不等于每个参数都单独发一次通信。一次 fully_shard 调用会建立一个通信分组;同组参数一起 All-Gather,梯度一起 Reduce-Scatter。通常先对各个 block 调用 fully_shard,再对整个模型调用一次,接管剩余参数。这样仍然能合并通信,也能按组临时聚合参数。

python
from torch.distributed.fsdp import fully_shard, FSDPModule
model = Transformer()
for layer in model.layers:
    fully_shard(layer)
fully_shard(model)

assert isinstance(model, Transformer)
assert isinstance(model, FSDPModule)
print(model)
#  FSDPTransformer(
#    (tok_embeddings): Embedding(...)
#    ...
#    (layers): 3 x FSDPTransformerBlock(...)
#    (output): Linear(...)
#  )

FSDP2 提供了 reshard_after_forward,控制一个 module 的 forward 结束后,是否立即释放 All-Gather 得到的完整参数。设为 True 时,算完就释放,回到只持有本地 shard 的状态,等 backward 需要这些参数时再 All-Gather。这里原来的 local shard 一直保留着,reshard 只是释放临时的完整参数并切回已有的 shard,本身不需要额外通信。

设为 False 时,则把完整参数留到 backward 继续使用,以更多显存换掉一次参数 All-Gather。在上面的默认用法中,各个 Transformer block 会在 forward 后 reshard,而最外层 model 管理的剩余参数会保留到 backward。

Checkpoint:让各张卡保存自己的分片

既然模型状态已经分散在不同的卡上,保存 checkpoint 时也不必先拼成完整模型。FSDP2 的 sharded state dict 用 DTensor 表示各个参数,可以交给 Distributed Checkpoint (DCP) 保存 (Zhang et al., 2025)

保存时,各 rank 共同参与,写出各自负责的 shards 和恢复所需的 metadata。加载时,DCP 根据目标模型的 sharding layout 读取数据,因此可以支持保存和恢复时使用不同的卡数。要恢复训练,还应保存 optimizer states 等训练状态,而不只是模型权重。

这不是 FSDP2 独有的能力,FSDP1 也支持 DCP。FSDP2 的便利在于逐参数的 DTensor 表示更容易操作。如果需要导出普通的完整 state dict,仍可以聚合参数后保存;这和直接保存分片 checkpoint 是两条不同的路径 (Feng et al., 2025)

References

Feng, W., Constable, W., & Mao, Y. (2025). Getting Started with Fully Sharded Data Parallel (FSDP2). PyTorch Tutorials. docs.pytorch.org
Hashimoto, T., & Liang, P. (2026). CS336: Language Modeling from Scratch. Stanford University, Spring 2026. cs336.stanford.edu
Jordan, K., Jin, Y., Boza, V., You, J., Cesista, F., Newhouse, L., & Bernstein, J. (2024). Muon: An optimizer for hidden layers in neural networks. kellerjordan.github.io
Liu, J., Su, J., Yao, X., Jiang, Z., Lai, G., Du, Y., Qin, Y., Xu, W., Lu, E., Yan, J., Chen, Y., Zheng, H., Liu, Y., Liu, S., Yin, B., He, W., Zhu, H., Wang, Y., Wang, J., … Yang, Z. (2025). Muon is Scalable for LLM Training. arxiv.org
Rajbhandari, S., Rasley, J., Ruwase, O., & He, Y. (2020). ZeRO: Memory Optimizations Toward Training Trillion Parameter Models. arxiv.org
Rasley, J., Rajbhandari, S., Ruwase, O., & He, Y. (2020). DeepSpeed: System Optimizations Enable Training Deep Learning Models with Over 100 Billion Parameters. Proceedings of the 26th ACM SIGKDD International Conference on Knowledge Discovery & Data Mining, 3505–3506. doi.org
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
Zhang, I., Kumpera, R., Huang, C.-C., & Pasqualin, L. (2025). Getting Started with Distributed Checkpoint (DCP). PyTorch Tutorials. docs.pytorch.org
Zhao, Y., Gu, A., Varma, R., Luo, L., Huang, C.-C., Xu, M., Wright, L., Shojanazeri, H., Ott, M., Shleifer, S., Desmaison, A., Balioglu, C., Damania, P., Nguyen, B., Chauhan, G., Hao, Y., Mathews, A., & Li, S. (2023). PyTorch FSDP: Experiences on Scaling Fully Sharded Data Parallel. arxiv.org

Cite this post

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