Skip to content

AllReduce 与 Ring 通信算法详解

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 × 数据量
通信算法每节点通信量带宽利用率适用场景
朴素 AllReduceN × 数据量N ≤ 4
Ring AllReduce2 × 数据量通用
Hierarchical~2 × 数据量最高多节点

Ring vs Hierarchical

单节点内(NVLink)使用 Ring AllReduce 即可;多节点场景推荐 Hierarchical AllReduce,先节点内 Ring 再节点间 Ring,减少跨节点通信量。

环断裂

Ring 算法要求所有 Rank 正常工作。任何一个 Rank 失败都会导致环断裂,训练挂起。生产环境建议使用容错机制或树形 AllReduce。

相关资源

最近更新