参考了 (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)。它的思路是:每张卡没必要一直保存一模一样的模型状态,可以把这些状态分给不同的卡保管,需要时再通信。
\(\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-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 会把一组参数一起管理,这一组就叫一个 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,再对整个模型调用一次,接管剩余参数。这样仍然能合并通信,也能按组临时聚合参数。
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)。