在数据并行训练中,你通常只需要写一行 DDP(model),然后一切就像魔法一样跑起来。但有一次我好奇:每张卡上的 Adam 状态到底是各自为政,还是完全相同的?如果不同,模型参数怎么会保持一致?

先说结论:在标准的 PyTorch DDP 数据并行中,每张卡上的 Adam 状态(m 和 v)是完全一致的。 而在 ZeRO 或 FSDP 这类优化器状态分片策略中,它们不一致——但这是设计好的,参数依然能保持一致。这篇文章会把这两条路的原理讲透,顺便回答一个问题:如果状态真的不一致了,会发生什么?
先搞清楚数据并行在干什么
假设你要把单卡训练改成 8 卡。最朴素的想法是:把模型复制到 8 张卡上,每张卡喂不同的数据,各自前向、反向、更新参数。最后把 8 份模型平均一下。
但这样做有一个致命问题:每张卡从不同的数据批次中学到的东西不一样,它们的参数会逐渐漂移。哪怕你定期平均,训练过程也会非常不稳定。
标准的 DDP(DistributedDataParallel)不是这么做的。它的流程是:
- 每张卡持有模型的完整副本,处理不同的数据批次。
- 前向传播,计算损失,反向传播得到每个参数的梯度。
- 对所有卡上的梯度做 AllReduce(取平均),得到全局梯度。
- 每张卡用这个全局梯度调用优化器更新参数。
注意第 3 步:梯度被平均后,每张卡拿到的梯度是相同的。而第 1 步保证了每张卡的初始参数相同(通过启动时的广播)。所以第 4 步更新后的参数也相同。这就是 DDP 为什么能保持所有卡模型一致的底层逻辑。
Adam 状态是什么?
Adam 是深度学习里最常用的优化器。它除了参数本身,还维护两个额外的状态:一阶矩 m(梯度的指数移动平均)和二阶矩 v(梯度平方的指数移动平均)。每次更新时,它用 m 和 v 来调整参数更新方向和步长。详细公式可以看 Adam 原始论文。
如果你看过 Adam 的更新公式,你会发现它每步都要更新 m 和 v。比如 m = β1 * m + (1 – β1) * g,其中 g 是当前梯度。这里的 m 是历史梯度的累积,所以它依赖之前每一步的梯度。
这意味着:如果两张卡上的历史梯度序列不同,它们的 m 和 v 就会不同,最终更新参数的方向也会不同。
标准 DDP 下,Adam 状态为什么一致?
关键就在于 DDP 的第 3 步:所有卡拿到的是同一个全局梯度 g。假设所有卡的初始参数相同,初始 m 和 v 都为零(Adam 默认初始化),那么第一步更新后,m 和 v 都是基于同一个 g 算出来的,所以完全一致。第二步又基于同一个 g(新的全局梯度)更新,仍然一致。以此类推。
所以,只要你的代码是在 AllReduce 之后才调用 optimizer.step(),每张卡上的 Adam 状态就是一模一样的。这不是巧合,而是 DDP 设计的一部分。PyTorch 官方文档中对此有描述:DDP 文档。
这其实也是很多人困惑的点:每张卡的数据明明不同,为什么梯度平均后状态就能保持一致?因为优化器状态只跟平均后的梯度有关,跟每张卡自己的局部梯度无关。
状态不一致会怎样?
最典型的错误场景是:在 AllReduce 之前就用局部梯度更新优化器状态。比如你写了个自定义训练循环,先 optimizer.step() 再 allreduce,或者忘了 allreduce,那么每张卡都会用自己那批数据的梯度去更新 m 和 v。由于不同卡的数据不同,m 和 v 自然不同。
结果是什么?每张卡用不同的优化器状态去更新参数,哪怕它们当前的参数是相同的,更新方向和步长也不同。一步之后,参数就分叉了。再往后,每个卡在参数空间中走各自的路,loss 曲线会剧烈震荡,甚至直接发散。
我最早犯过这个错。当时我把 AllReduce 写在了 optimizer.step() 之后,结果训练两三千步后 loss 开始剧烈波动,我还以为是学习率太大了。最后打印了每个卡上的 m 值,发现它们已经差了好几个数量级。
ZeRO 和 FSDP:故意让状态不一致
DDP 虽然逻辑简单,但有一个显存问题:每张卡都要保存一份完整的模型参数、梯度和优化器状态。对于一个 10B 参数的模型,Adam 的 m 和 v 各占 8 字节(fp32),加上参数和梯度,一张卡可能就爆显存了。
ZeRO(ZeRO: Memory Optimizations Toward Training Trillion Parameter Models)提出了一种思路:优化器状态不需要每张卡都保存一份,可以分片存到不同卡上。这就是 ZeRO 阶段 1。PyTorch 的 FSDP 也是类似思想。
在 ZeRO 阶段 1 中,每张卡只保存一部分参数的 Adam 状态。更新时,每张卡用 reduce-scatter 拿到自己负责的那部分梯度,然后用自己保存的 m 和 v 更新对应参数。更新完后,通过 AllGather 把参数广播给所有卡,保证下一轮前向时所有卡都有完整的参数。
所以,在 ZeRO/FSDP 下,每张卡上的 Adam 状态天然不一致——因为它们分别负责不同的参数。但这种不一致是有意的,每个参数只由一个卡的优化器状态管理,不存在"同一个参数在不同卡上有不同状态"的问题。参数的一致性由 AllGather 保证。
这一点经常被误解。有人以为 ZeRO 和 DDP 一样,每张卡也有完整优化器状态,只是节省了梯度。其实不是,ZeRO 分片的是优化器状态本身。
| 策略 | 优化器状态存储 | 状态一致性 | 参数同步方式 |
|---|---|---|---|
| DDP | 每卡全量复制 | 一致 | 每步 AllReduce 梯度 |
| ZeRO 阶段1 | 分片到各卡 | 不一致(分片) | ReduceScatter 梯度 + AllGather 参数 |
| 模型并行 | 各卡只存自己层的参数 | 不适用 | 前向/反向中通信激活和梯度 |
即使理论上一致,实际也可能有微小偏差
上面说的是理想情况。现实中,即使你正确使用了 DDP,不同卡上的 Adam 状态也可能出现微小差异。原因有几个:
- AllReduce 的浮点求和顺序不同。比如卡 0 先和卡 1 求和,再和卡 2 求和;卡 1 可能先和卡 2 求和。不同顺序会带来舍入误差。
- 非确定性算法。某些 CUDA 算子(比如原子操作)在不同卡上可能产生不同的结果。
- 数据加载中的随机性,如果设置了不同的随机种子,Dropout 等层的输出不同,导致梯度不同(但经过平均后影响较小)。
这些微小差异在训练初期可能很小,但 Adam 的 m 和 v 是累积量,微小误差会不断累积。长期训练后,不同卡上的参数可能开始分叉,甚至影响收敛。PyTorch 官方在随机性文档中提到了这些问题,并提供了一些方法来保证可复现性。
如果你需要严格的可复现性,可以设置 torch.use_deterministic_algorithms(True),并且让所有卡使用相同的随机种子。
那么,不一致到底会带来多大影响?
取决于不一致的程度和策略。
对于标准 DDP,理想情况必须一致。如果不一致,说明实现有 bug,会导致训练崩溃或性能下降。
对于 ZeRO/FSDP,不一致是设计。每个参数只由一个卡更新,所以不会出现同一参数被不同状态更新的问题。只要通信正确,模型参数是一致的。
但对于模型并行(比如张量并行),每张卡负责不同的参数,优化器状态本来就不同。这和"不一致"不是一个维度。
所以,当你问"每张卡上的 Adam 状态一致吗"时,先想清楚你用的是哪种并行策略。
我的最后一条建议
如果你刚开始做分布式训练,不要自己造轮子。直接用 PyTorch 的 DDP 或 FSDP,它们已经处理好了所有同步细节。如果你好奇内部发生了什么,可以打印一下各卡的优化器状态对比,但别在生产环境做这种事。
记住:标准 DDP 下,状态不一致 = bug;ZeRO 下,状态不一致 = 设计。把这个区别搞清楚,你就不会再被分布式训练的优化器问题折磨了。
原创文章,作者:guanweilu,如若转载,请注明出处:https://guanweilu.cn/article/308.html