ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

all-reduce详解:多卡训练如何让每张GPU模型保持一致

all-reduce详解:多卡训练如何让每张GPU模型保持一致 最近在技术社区刷到 MiniMax-H3 的讨论很多人在聊它生成的竖屏短片效果也在聊模型下载后怎么加速推理。但我更关心的不是短片本身而是这些视频模型背后一个很容易被带过的底层问题多卡训练时all-reduce 到底怎样让每张卡上的模型保持一致这不是纯理论问题。写过分布式训练脚本的人都知道每次看 loss 曲线时真正想让所有 GPU 拿到的不是“各自的结果”而是一份完全一致的、聚合了所有卡信息之后的结果。all-reduce 就是负责这件事的算子。这篇文章想把它讲透从语义到 Ring 算法从 PyTorch/DDP 的落地到排查链路最后给一个判断框架。1. 先搞清楚一个容易被带偏的问题训练时每张卡为什么会有差异1.1 数据并行下每张卡看到的本来就是不同数据规模化训练里最常见的并行方式是数据并行。把训练集切成 N 份每张卡拿其中一份各自做前向、算 loss、再反向。这里的起始点就是不一致的每张卡输入的是不同 batch算出来的 loss 和梯度天然不同。不要觉得这有什么问题。数据并行本来就是用“数据分片”换取“算力扩展”每张卡看到的数据不同是设计的一部分。1.2 不一样从梯度开始最后会传导到参数如果每张卡各自按自己算出来的梯度直接更新参数那么从第一个 step 开始模型权重就已经分叉了。N 张卡跑一段时间后等于 N 个版本不同的模型在并行训练。最后保存模型的时候该信任哪一张卡的 checkpoint结果通常完全不可复现。所以需要一种机制每张卡算好自己的梯度之后先不要急着更新而是把 N 份梯度汇总成一份统一结果再广播回每一张卡。所有卡用同一份梯度去更新各自的参数。因为初始参数相同、更新规则相同、拿到的梯度也相同下一步参数自然一致。1.3 “每卡相同”具体指哪些东西相同这里需要分辨三层状态模型参数每个 step 结束之后N 张卡上的权重应当完全一致。这是最核心的要求。优化器状态如果每张卡初始状态一致、拿到的梯度一致、超参一致优化器状态也会保持一致。实际运行时会因为浮点运算顺序不同出现极小数值偏差但一般不会改变训练路径。随机状态dropout 这类训练期随机性可以不一致也不需要一致。如果要做严格复现需要统一 seed并控制 CUDA 层面引入的非确定性。一个很容易误解的点是很多人以为“每卡相同”是指计算过程完全一样。其实不是。输入不同、计算路径里可能有随机性这些差异都允许存在真正必须相同的是“聚合后的梯度”和“更新后的参数”。all-reduce 保证的是结果侧的收敛。1.4 先把主判断放在这里all-reduce 不是为了让每张卡“算得一模一样”而是让每张卡在拿到其他人的信息之后变成“最终结果一模一样”。它是把分布式训练的结果拉回等价于单卡训练的桥梁。2. all-reduce 的语义一次操作本质上要完成“聚合”与“分发”2.1 Reduce 和 All-Reduce 的区别在集合通信里reduce 指把多张卡上的张量聚合成一个结果这个结果可能只落在某一张卡上。比如 all-gather 是每个节点都拿到所有人的数据然后本地做 reduce这是一种可以实现目标的做法但通信成本高。all-reduce 把“聚合”和“分发”合并成了一个语义N 张卡各自持有一个张量经过一次调用之后每张卡都拿到完全相同的聚合结果。常见聚合操作有这么几种SUM直接把 N 份张量相加这是 DDP 梯度同步的默认做法。AVG相加后除以 N。MAX / MIN取最大或最小。PROD相乘。在梯度同步场景里最常用的是 SUM 和 AVG。PyTorch 的torch.distributed.all_reduce支持ReduceOp.SUM、ReduceOp.AVG等DDP 内部默认是 SUM所以使用时要留意 loss 的归约方式。2.2 集合语义带来的两个工程约束all-reduce 不是 A 到 B 的点对点通信而是要求某个集合内所有成员一起参与、结束时状态一致。这带来两个工程含义。第一个是调用必须“齐步走”所有 rank 都要执行相同的 all-reduce且参与的张量形态要匹配。如果一个 rank 没调用其他 rank 会一直卡在通信等待里。很多分布式脚本“挂住不动”的问题根源就在这里。第二个是结果需要确定性不管调用顺序怎样同一 batch 下各卡拿到的最终张量应当一致。这取决于底层实现的归约顺序也是后面讲数值精度时的一个重要背景。2.3 逻辑上可以拆成“先汇总再广播”不管底层怎么实现所有 all-reduce 都可以在逻辑上拆成两步汇总把所有卡的局部张量聚合成一份全局结果。广播把这份全局结果复制回每一张卡。Ring 算法之所以经典是因为它对这两步做了高效的流水化实现而不是先把数据汇总到某个中心再分发出去。高效点不在语义创新而在通信模式的重新设计。3. Ring AllReduce为什么它能高效地做到“让每卡相同”3.1 朴素做法的问题在哪里一个最直接的做法是每张卡把自己的张量广播给其他所有卡然后每张卡本地求和。两卡时很直观但到大规模集群就有几个问题。第一消息数随卡数快速增长变成接近 O(N²) 级别的事务数量。第二某一时刻多个来源同时发给同一目的地接收端网卡成为瓶颈。第三没有中心节点来统一汇聚很多实现会退化成“集线器模式”实际带宽浪费严重。3.2 Ring 的核心reduce-scatter all-gatherRing AllReduce 是当前多卡训练里最主流实现之一。它的思路是把 N 张卡看成一个环每张卡只和前后两个邻居通信。假设每张卡持有大小为 S 的梯度张量总共有 N 张卡。算法分两个阶段。第一阶段叫 reduce-scatter归约分散。把本卡大小为 S 的张量切成 N 份。每轮里每张卡把自己手里的某个 chunk 发给下一个邻居同时从上一个邻居收到一个 chunk并累加到自己对应的 chunk 上。重复 N-1 轮之后每张卡上有一份 chunk它是所有卡对应位置累加后的结果。第二阶段叫 all-gather全收集。每张卡把手里那份累加后的 chunk 继续沿环发给下一个邻居。重复 N-1 轮之后每张卡都收集齐了全部累加 chunk。把它们拼接起来就得到完整的 all-reduce 结果。这个阶段不需要计算只做搬运。3.3 通信量为什么近似“2 倍数据量”每张卡在第一阶段发送 N-1 次 chunk第二阶段同样发送 N-1 次 chunk。每张卡发送和接收的总数据量大致是每卡收发总量 2 × (N-1) × S / N当 N 较大时约等于 2S。这个数很关键无论集群里有几百张卡单个节点搬进搬出的数据量主要只由它自己的张量规模决定不会因为卡数变多而线性膨胀。所以在消息足够大的场景下Ring 是带宽最优的近似实现特别适合大模型、大梯度张量。如果每张卡把整份 S 都广播出去每卡收发量会变成 2(N-1)S比 Ring 差了 N 倍。Ring 厉害的地方就在这里同样一次 all-reduce它把“与卡数相关的开销”压到了只剩常数级别。3.4 四卡环的一个文字演示以四卡 A、B、C、D 组成环为例。reduce-scatter 阶段每轮大家同时做“发一个 chunk、收一个 chunk、累加”。三轮之后A 手里持有某一份完整累加 chunkB、C、D 各持有另外一份完整累加 chunk谁都不完整但每个人手里那份已经汇聚了所有卡的信息。all-gather 阶段A 把手里 chunk 发给 B同时从 D 收到所需 chunk。三轮之后A、B、C、D 都拥有全部四个累加结果拼接后就是完整张量。注意每个环节所有卡都在同时收发链路上没有闲置这是 Ring 能跑满带宽的核心原因。4. 从参数服务器到 NCCL为什么工业界默认选择 all-reduce4.1 参数服务器不是银弹在 all-reduce 流行之前参数服务器PS是分布式训练里很常见的一种方案。思路直接worker 卡负责计算梯度把梯度发给一组 serverserver 汇总后更新参数再把新参数广播回 worker。优点是实现直观天然支持异步更新。缺点是当模型变大、worker 变多时server 的入口带宽和出口带宽会变成瓶颈而且 server 节点一旦出问题整个训练就被拖住。从工程直觉看PS 更适合“通信不是主要矛盾、需要灵活调度”的场景数据并行 all-reduce 更贴近“算力足够多、想尽量压低通信开销”的场景。4.2 NCCL 把 Ring 做成了生产级NCCLNVIDIA Collective Communications Library把环、树等多种集合通信算法做成了生产级实现并针对单机 NVLink、多机 InfiniBand 或 RoCE 做了适配。几个关键优化点根据卡间拓扑选择算法比如单机内部用 NVLink 高带宽链路跨机场景用树或环降低链路压力。梯度分桶把多个参数的小梯度合并成更大的通信块减少消息数量。通信和反向计算重叠bucket 填满就先发不必等全部梯度算完。支持 RDMA 和 GPUDirect减少数据拷贝。这里要强调一点NCCL 不是只有 Ring。实际使用中它会根据设备数、节点数、数据规模自动或手动选择算法。使用者一开始不必手动干预但要知道有这些参数存在。方案通信方式主要瓶颈适合场景参数服务器中心节点汇聚再广播server 带宽、单点故障小规模、异步容忍度高朴素 AllReduce每卡广播再本地求和消息数多、接收端拥塞卡数非常少Ring AllReduce环上点对点流转延迟随环长度增加大梯度、大规模数据并行Tree AllReduce树状聚合再广播根节点仍是汇聚瓶颈消息较小、延迟敏感4.3 all-reduce 不只出现在梯度同步很多人把 all-reduce 等同于“梯度同步”。其实在大模型训练里它还频繁出现在张量并行中。以常见的 Megatron 风格模型并行举例一层 MLP 被切到多张卡上每张卡只算一部分神经元。前向过程中的某些位置需要把各卡的部分结果聚合起来才能继续算下一层。这个聚合动作往往就是 all-reduce。理解这一点后再看“让每卡相同”会更完整无论做数据并行还是张量并行最终都要保证关键张量在所有卡上是对齐的。5. 在 PyTorch 里DDP 是用 all-reduce 帮你怎么做的5.1 DDP 的三步固定流程实际项目里我一般不会手动写 all-reduce 来同步梯度因为DistributedDataParallel已经封装好了。封装不是魔法它只是做几件固定的事在进程组初始化时建立所有 rank 的通信上下文。前向结束后在反向传播过程中给每个参数的梯度注册钩子。梯度算出来后按 bucket 分批执行 all-reduce把多卡带来的梯度汇总回每张卡。对应到代码大致是这样一种结构import torch import torch.distributed as dist from torch.nn.parallel import DistributedDataParallel as DDP def train_worker(rank, world_size): dist.init_process_group(nccl, rankrank, world_sizeworld_size) torch.cuda.set_device(rank) model build_model().cuda(rank) model DDP(model, device_ids[rank], bucket_cap_mb25) dataset build_dataset(rankrank, num_replicasworld_size) loader torch.utils.data.DataLoader(dataset, batch_size8) optimizer torch.optim.AdamW(model.parameters(), lr1e-4) for batch in loader: optimizer.zero_grad() loss compute_loss(model, batch) # DDP 的梯度 all-reduce 默认是 SUM。 # 如果想等价于“全局平均梯度”需要在 loss 处除以 world_size。 loss loss / world_size loss.backward() optimizer.step()上面是通用写法具体模型和数据集要换成你自己的。写的时候注意两件事数据加载用DistributedSampler保证 batch 不重叠训练语义正确。loss 是否除以 world_size取决于你想让优化步等价于单卡哪种语义。这里没有唯一正确答案但必须明确自己选的是哪一种。5.2 几个会直接影响结果的 DDP 参数参数影响落地建议bucket_cap_mb控制梯度分桶大小影响 all-reduce 次数先用默认 25MB观察 GPU 利用率和通信耗时再调find_unused_parameters处理前向里没用到的参数设置不当会卡住有动态分支时设为 True会带来额外开销gradient_as_bucket_view让梯度视图直接指向 bucket省内存拷贝显存紧张时可开启但升级版本后要看 API 变化timeout控制通信等待超时长训练任务建议设置明确超时暴露卡住问题很多人忽略bucket_cap_mb但它是影响通信利用率的关键之一。bucket 太小通信次数多、消息碎bucket 太大通信占用的显存多、延迟发布。我会先跑一个稳定小任务再用torch.profiler看通信占比最后才动这个参数。5.3 混合精度里最容易忽略的一点在 AMP 或 bf16 训练里损失 scale 和梯度精度会影响 all-reduce 的结果。fp16 梯度做通信时数值范围有限。如果 loss scale 不当梯度可能下溢。bf16 指数范围更大但有效精度较少。直接用 bf16 梯度做 all-reduce结果会粗糙一些。更稳妥的做法是用 fp32 主权重和梯度做归约稳定之后再尝试低精度通信优化。这种细节放在“每卡相同”的话题里尤其重要。因为浮点归约顺序不同各卡拿到的最终值会有极小差异。多数情况下差异不影响训练但做严格复现和调优对齐时它就是需要关注的那一部分。5.4 如果你是冲着“视频生成模型加速”来的回到开头提到的 MiniMax-H3 讨论。社区里聊得多的往往是“能不能生成竖屏短片”“怎么加速推理”但很少有人聊模型背后的并行链路。我的体感是如果你只是用现成模型跑一次推理all-reduce 不一定出现在你面前但如果你要自己微调一个类似规模的视频生成模型或者想把推理拆到多卡并行all-reduce 迟早会出现。这类视频生成模型通常有几个共同点batch 内视频片段较长显存和算力需求高需要足够大的 batch 才能让训练稳定对时间维度的处理会带来更多激活值占用。用 DDP 或更上层框架跑这类训练时梯度 all-reduce 的通信占比往往很可观。优化方向不是调一个神秘参数而是先搞清楚通信发生的位置以及瓶颈是带宽还是延迟。6. 排查链路如果发现每卡结果不一致按什么顺序查很多人遇到“多卡训练结果不对”时第一反应是怀疑 all-reduce 坏了。但从工程经验看all-reduce 本身很少是根因绝大多数问题出现在它前后。6.1 先确认是不是真的“不一致”不要急着看通信。把同样数据、同样 seed 跑一次单卡再跑一次多卡对比 loss 曲线和最终指标。如果只是存在微小浮点差异很可能不是故障而是归约顺序、TF32、低精度算子导致的数值抖动。这种情况不需要消除只需要控制在可接受范围。如果差异很大比如 loss 发散或者从一个 step 开始就完全不同进入下一步。6.2 再查初始化和数据按顺序检查每张卡的 seed 是否一致。checkpoint 是否在初始化进程组之后正确加载还是每卡都从随机权重开始。DistributedSampler是否正确使用batch 是否重叠。数据预处理中是否有依赖进程号的随机操作没有固定。6.3 再查同步机制和 DDP 状态确认模型确实被DDP包住而不是仍然用了nn.DataParallel。检查模型是否有requires_gradFalse的参数。这类参数不参与梯度 all-reduce如果依赖它们变化步调就会出现偏差。检查是否有 buffer 依赖训练过程更新。DDP 默认在 forward 时从 rank 0 同步 buffer但 buffer 不参与梯度归约。如果各卡独立更新 buffer就可能不一致。检查是否有某个参数没有被任何 loss 反向使用尤其是有条件分支、动态结构的模型。此时需要find_unused_parametersTrue。6.4 再查通信环境如果是多机多卡还要看NCCL 版本和 PyTorch 自带的 NCCL 是否匹配。网卡配置比如 RDMA 是否启用网卡是否绑定到了正确的 CPU 和 NUMA 节点。NCCL_P2P_DISABLE、NCCL_IB_DISABLE这类环境变量是否被误设。日志里是否有 timeout、network error 等关键词。6.5 最后才查算子层面的非确定性如果以上都没问题再考虑算子层面是否设置了torch.backends.cudnn.deterministicTrue。是否关闭了 TF32。是否在低精度下启用了原子累加。张量并行或序列并行里的 all-reduce 次数和位置是否正确。这套排查顺序的核心逻辑是先排除“数据或初始化导致起点不同”再排除“同步机制失效”最后才进入“通信过程中数值不精确”的细节。否则很容易在最后一个环节浪费时间。7. 一个判断框架什么时候值得在 all-reduce 上花时间7.1 三个层次的需求判断不是所有项目都需要深入 all-reduce。这里给一个简单框架场景你需要关心什么建议学习 / 小实验跑通 DDP理解语义先用默认 DDP不手动调通信参数单机多卡生产通信占比、bucket、显存先用 profiler 分析再调 bucket 和重叠多机多卡 / 大规模带宽、拓扑、RDMA、故障恢复了解 NCCL 拓扑和并行策略必要时上 FSDP 或混合并行7.2 all-reduce 不解决什么写清楚边界能省掉很多调试时间它不解决数据不均衡。如果某张卡分到的 batch 明显更大或更复杂训练效率会被慢卡拖住。它不解决节点速度不一致。Ring 假设各节点速度接近遇到 straggler 时整体都会变慢。它不解决模型并行里的手动调度。在这类场景里all-reduce 只是众多算子之一关键是并行切分策略。它不保证推理阶段所有卡输出一致。推理一致性还需要模型本身确定性、精度设置和输入处理方式配合。7.3 长期优化方向如果你真的需要在一个长训练任务里持续压通信成本方向通常是这几条梯度压缩或量化用低精度表示梯度减少通信量但要在小规模先验证收敛。通信与计算重叠利用 bucket 和异步 all-reduce把通信藏进反向计算里。减少 all-reduce 次数把过小的参数合并进同一个 bucket。当模型大到单卡放不下要上 FSDP 或混合并行那是另一套权衡不能只在通信参数上打转。回到题目那句话all-reduce 怎样让每卡相同我的答案是它把“各自不同的局部结果”先聚合成一份“大家共享的全局结果”再让每一份副本回到每一张卡。机制本身看起来简单真正难的是背后的算法选择、工程封装和故障排查。如果你刚接触分布式训练不要急着手动实现 all-reduce。先用 DDP 把一个小模型跑通打开torch.profiler看一眼 GPU 利用率和通信耗时再尝试调 bucket 和 batch size。当你亲眼看到通信开销占了多少时间之后才会真正理解为什么 ring、bucket、低精度、拓扑优化这些东西值得被反复讨论。大模型训练和多卡推理越来越常见但能让整个系统稳定运行的从来不只是一个模型结构而是这些不起眼的集合操作在每一层正确地对齐。
返回列表