· 分布式训练
Ring AllReduce 核心流程:梯度是怎么聚合的
梯度切成 N 块、N 张卡排成环,两轮各走 N−1 步:每卡只搬约 2 份梯度大小的数据,就能拿到全量梯度之和。
一句话:梯度切成 N 块,N 张卡排成环,两轮各走 N−1 步,最后人人拿到全量梯度和,而每卡只搬了约 2 份梯度的数据量。
1 流程
- 切块:梯度向量切成 N 块,每块
K/N。 - 第一轮 Reduce-Scatter,N−1 步:每步向右邻居发一块、从左邻居收一块(发和收是同一块编号),收到就地累加。结束时长这样:Rank n 完整持有第
(n+1) mod N块,其余是半成品。 - 第二轮 All-Gather,N−1 步:把手里那块沿环再传一圈,人人补齐所有块。
- 所有卡得到同一份全量梯度和 → 更新参数,进入下一步。
2 每个环节长什么样
展开:N=4 逐步传块顺序(可对照代码)
| 步骤 | R0 | R1 | R2 | R3 |
|---|---|---|---|---|
| ① 第 1 步 | c0 → R1 | c1 → R2 | c2 → R3 | c3 → R0 |
| ① 第 2 步 | c3 → R1 | c0 → R2 | c1 → R3 | c2 → R0 |
| ① 第 3 步 | c2 → R1 | c3 → R2 | c0 → R3 | c1 → R0 |
| ① 结束:R0 拥有 c1,R1 拥有 c2,R2 拥有 c3,R3 拥有 c0 | ||||
| ② 第 1 步 | c1 → R1 | c2 → R2 | c3 → R3 | c0 → R0 |
| ② 第 2 步 | c0 → R1 | c1 → R2 | c2 → R3 | c3 → R0 |
| ② 第 3 步 | c3 → R1 | c0 → R2 | c1 → R3 | c2 → R0 |
| ② 结束:人人都有 c0 c1 c2 c3 的完整结果 | ||||
3 省了多少、代价是什么
Ring(2(N−1)/N) 朴素全互联(N−1)
Ring:2(N−1) 步 树形:2log₂N 步
≈2K每卡发送量,与卡数无关
2(N−1)串行步数
10.2 ms1024 卡的起步费(每步 5µs)
0.03 ms搬 4MB 张量的纯传输时间
所以张量小的时候,Ring 基本在"排队握手"而不是搬数据。真实框架的做法是先把小张量分桶合并,再交给 Ring。
4 工程上真正起作用的一步:重叠
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 |
|---|---|---|---|
| 8 | 1.75 K | 0.250 × | 0.875 × |
| 64 | 1.97 K | 0.031 × | 0.984 × |
| 256 | 1.992 K | 0.008 × | 0.996 × |
| 时延场景(α=5µs,β=300GB/s) | Ring 总时延 |
|---|---|
| 8 卡 / 1M 元素 | 0.094 ms |
| 1024 卡 / 1M 元素 | 10.26 ms(起步费占 99%) |
| 1024 卡 / 1.07GB | 17.38 ms(传输 7.2ms 才开始与起步费打平) |
上面数字由 ring_allreduce/sim.py 跑出,N=2…256 均通过正确性校验(与 np.sum 逐元素一致)。