AllReduce 与 Ring 通信算法详解
AllReduce 是分布式训练中最核心的通信操作,Ring AllReduce 是其最高效的实现算法之一。理解 Ring AllReduce 的原理对优化分布式训练至关重要。
AllReduce 语义
AllReduce 将所有进程的数据进行归约操作(求和/平均/最大值等),并将结果分发给所有进程:
输入: [1, 2, 3, 4] (4个Rank)
Sum AllReduce 输出: [10, 10, 10, 10] (每个Rank获得全局求和)Ring AllReduce 算法
Ring AllReduce 将数据分为 N 份(N = 进程数),通过环形通信分两阶段完成:
Phase 1: Reduce-Scatter
每个 Rank 将数据分为 N 块,沿环传递并累加:
Step 1: Rank 0 → Rank 1 (块0) Rank 1 → Rank 2 (块1) Rank 2 → Rank 3 (块2) Rank 3 → Rank 0 (块3)
Step 2: Rank 0 → Rank 1 (块3) Rank 1 → Rank 2 (块0) Rank 2 → Rank 3 (块1) Rank 3 → Rank 0 (块2)
...
Step N-1: 每个Rank拥有一个完整归约的块Phase 2: All-Gather
沿环传递归约结果,使所有 Rank 获得完整数据:
Step 1-N: 与 Reduce-Scatter 方向相反,传递已归约的块python
# Ring AllReduce 伪代码
def ring_all_reduce(data, rank, world_size):
N = world_size
chunk_size = len(data) // N
# Phase 1: Reduce-Scatter
for step in range(N - 1):
send_idx = (rank - step) % N
recv_idx = (rank - step - 1) % N
send_chunk = data[send_idx * chunk_size : (send_idx + 1) * chunk_size]
recv_chunk = torch.zeros_like(send_chunk)
# 同时发送和接收
dist.send(send_chunk, dst=(rank + 1) % N)
dist.recv(recv_chunk, src=(rank - 1) % N)
# 累加接收到的数据
data[recv_idx * chunk_size : (recv_idx + 1) * chunk_size] += recv_chunk
# Phase 2: All-Gather
for step in range(N - 1):
send_idx = (rank + 1 - step) % N
recv_idx = (rank - step) % N
send_chunk = data[send_idx * chunk_size : (send_idx + 1) * chunk_size]
recv_chunk = torch.zeros_like(send_chunk)
dist.send(send_chunk, dst=(rank + 1) % N)
dist.recv(recv_chunk, src=(rank - 1) % N)
data[recv_idx * chunk_size : (recv_idx + 1) * chunk_size] = recv_chunk
return data通信量分析
Ring AllReduce 的总通信量:
- 每个 Rank 发送/接收量:
2 × (N-1)/N × 数据量 - 当 N 较大时:约等于
2 × 数据量
| 通信算法 | 每节点通信量 | 带宽利用率 | 适用场景 |
|---|---|---|---|
| 朴素 AllReduce | N × 数据量 | 低 | N ≤ 4 |
| Ring AllReduce | 2 × 数据量 | 高 | 通用 |
| Hierarchical | ~2 × 数据量 | 最高 | 多节点 |
Ring vs Hierarchical
单节点内(NVLink)使用 Ring AllReduce 即可;多节点场景推荐 Hierarchical AllReduce,先节点内 Ring 再节点间 Ring,减少跨节点通信量。
环断裂
Ring 算法要求所有 Rank 正常工作。任何一个 Rank 失败都会导致环断裂,训练挂起。生产环境建议使用容错机制或树形 AllReduce。