· 分布式训练

Ring AllReduce 核心流程:梯度是怎么聚合的

梯度切成 N 块、N 张卡排成环,两轮各走 N−1 步:每卡只搬约 2 份梯度大小的数据,就能拿到全量梯度之和。

一句话:梯度切成 N 块,N 张卡排成环,两轮各走 N−1 步,最后人人拿到全量梯度和,而每卡只搬了约 2 份梯度的数据量。

1 流程

两轮传递的流程第一轮边传边加,每张卡负责一块;第二轮把结果传一圈补齐① Reduce-Scatter(N−1 步)R0R1R2R3一边传一边加,每人负责 1 块② All-Gather(N−1 步)R0R1R2R3传一圈,人人拿到全部每步每卡只发 1 块、收 1 块(各 K/N),收发是同一块编号 —— 没人加班,也没人闲着
  1. 切块:梯度向量切成 N 块,每块 K/N
  2. 第一轮 Reduce-Scatter,N−1 步:每步向右邻居发一块、从左邻居收一块(发和收是同一块编号),收到就地累加。结束时长这样:Rank n 完整持有第 (n+1) mod N 块,其余是半成品。
  3. 第二轮 All-Gather,N−1 步:把手里那块沿环再传一圈,人人补齐所有块。
  4. 所有卡得到同一份全量梯度和 → 更新参数,进入下一步。

2 每个环节长什么样

4 张卡排成一个环每张卡只和左右邻居通信,没有中心节点Rank 0Rank 1Rank 2Rank 3梯度切成 N 块(每块 K/N);每张卡只跟左右邻居说话
没有"中心节点",所以没有哪张卡会成为瓶颈;N 张卡只要 N 条连线,不是 N² 条。
两轮结束后每张卡手上的数据第一轮后每张卡只完整拥有一块,第二轮后所有人拥有全部四块① 结束后② 结束后c0c1c2c3R0半成品已完成半成品半成品R1半成品半成品已完成半成品R2半成品半成品半成品已完成R3已完成半成品半成品半成品c0c1c2c3R0已完成已完成已完成已完成R1已完成已完成已完成已完成R2已完成已完成已完成已完成R3已完成已完成已完成已完成R3 拿到的是 c0 —— 环首尾相接,右下角绕回左上角;蓝色半成品最后扔掉,不额外占带宽
黄色 = 这块已加完,绿色 = 人人都有完整结果。蓝色半成品随后丢弃,这一部分不花任何通信成本,是这个算法划算的关键。
展开:N=4 逐步传块顺序(可对照代码)
步骤R0R1R2R3
① 第 1 步c0 → R1c1 → R2c2 → R3c3 → R0
① 第 2 步c3 → R1c0 → R2c1 → R3c2 → R0
① 第 3 步c2 → R1c3 → R2c0 → R3c1 → R0
① 结束:R0 拥有 c1,R1 拥有 c2,R2 拥有 c3,R3 拥有 c0
② 第 1 步c1 → R1c2 → R2c3 → R3c0 → R0
② 第 2 步c0 → R1c1 → R2c2 → R3c3 → R0
② 第 3 步c3 → R1c0 → R2c1 → R3c2 → R0
② 结束:人人都有 c0 c1 c2 c3 的完整结果

3 省了多少、代价是什么

Ring(2(N−1)/N) 朴素全互联(N−1)
141664256102481632641282565121024
纵轴对数刻度,单位是"份梯度"。蓝线几乎水平 —— 卡数从 8 涨到 1024,每卡发送量只从 1.75 份到 2 份;红线从 7 份涨到 1023 份。
Ring:2(N−1) 步 树形:2log₂N 步
116256409681632641282565121024
树形算法的总搬运量和 Ring 是同一量级(都≈2K),差别只在步数 —— 这正是 NCCL"小消息走树形、大消息走环"的原因。
≈2K每卡发送量,与卡数无关
2(N−1)串行步数
10.2 ms1024 卡的起步费(每步 5µs)
0.03 ms搬 4MB 张量的纯传输时间

所以张量小的时候,Ring 基本在"排队握手"而不是搬数据。真实框架的做法是先把小张量分桶合并,再交给 Ring。

4 工程上真正起作用的一步:重叠

算一层发一桶:通信藏进计算里梯度从后往前算出来,装满一桶立刻发出去通信,和剩下的计算重叠算梯度(从最后一层往前算)L6L5L4L3L2L1发出去做通信(凑够一桶就发)桶 1(L6-L5-L4)桶 2(L3-L2-L1)桶 1 的通信被 L3-L1 的计算完全盖住,只有桶 2 露在外面;桶越大带宽越满但开始越晚(默认 25MB)
"通信零成本"的真实含义不是通信变快,而是它被计算盖住了。

DDP 一边反向传播一边把算好的梯度装进桶(默认 25MB),桶满立刻发出去做 AllReduce,此时前面的层还在继续算。ZeRO-1/2 则把 AllReduce 退化成只做第一轮(≈K),因为每卡只需留下自己那 1/N 份交给优化器。

5 只需记住三件事

  • 省带宽:每卡 2(N−1)/N × K ≈ 2K,与卡数无关。朴素做法是 (N−1)K,256 卡差 128 倍。
  • 代价是步数2(N−1) 步。1024 卡 = 2046 步 ≈ 10ms 起步费,此时搬数据只要 0.03ms —— 小消息必须走树形。
  • 落地关键:分桶 + 反向重叠。算法本身没得选,把通信藏进计算才是性能来源。
SFT 最容易踩的坑:FSDP 下梯度裁剪必须用分片感知版本(否则每卡裁剪系数不同,权重静默发散);序列长度长尾会造成 straggler,靠 sequence packing 解决而不是调通信参数;梯度累积的中间步必须 no_sync(),否则通信量乘上累积步数。
展开:伪代码 + 实测数据
# K 个元素切成 N 块,每块 C = K/N;所有 Rank 执行同一段代码
for step in range(N - 1):                  # 第一轮:边传边加
    send = (rank - step)     % N
    recv = (rank - 1 - step) % N
    isend(buf[send*C:(send+1)*C], (rank + 1) % N)
    part = irecv((rank - 1) % N)
    wait()
    buf[recv*C:(recv+1)*C] += part          # 就地累加

for step in range(N - 1):                  # 第二轮:传一圈补齐
    send = (rank + 1 - step) % N
    recv = (rank - step)     % N
    isend(buf[send*C:(send+1)*C], (rank + 1) % N)
    buf[recv*C:(recv+1)*C] = irecv((rank - 1) % N)
    wait()

必须用非阻塞收发(isend/irecv)+ wait():写成先阻塞发、再阻塞收,环上会互相等死。

卡数 N每卡发送量对比朴素做法对比下界 2K
81.75 K0.250 ×0.875 ×
641.97 K0.031 ×0.984 ×
2561.992 K0.008 ×0.996 ×
时延场景(α=5µs,β=300GB/s)Ring 总时延
8 卡 / 1M 元素0.094 ms
1024 卡 / 1M 元素10.26 ms(起步费占 99%)
1024 卡 / 1.07GB17.38 ms(传输 7.2ms 才开始与起步费打平)

上面数字由 ring_allreduce/sim.py 跑出,N=2…256 均通过正确性校验(与 np.sum 逐元素一致)。

本笔记为个人学习用途的结构化整理与图解重绘,全部配图为程序化生成的自绘矢量图。

分享:
返回文章列表

相关文章

全部文章 »