为什么说 AllReduce = ReduceScatter + AllGather?——结合 NCCL Ring AllReduce

参考

其实直接看https://arxiv.org/pdf/2507.04786的 fig.4就行了

NVIDIA NCCL 官方文档明确指出:

ReduceScatter 后执行 AllGather,在结果上等价于 AllReduce。

这个等式可以从两个层面理解:

  1. 语义上为什么等价
  2. NCCL 的 Ring AllReduce 实际上怎样把这两个阶段做出来

1. 先从最终目标看

假设有 4 个 GPU,每个 GPU 上都有 4 个 chunk:

1
2
3
4
GPU0: [a0, a1, a2, a3]
GPU1: [b0, b1, b2, b3]
GPU2: [c0, c1, c2, c3]
GPU3: [d0, d1, d2, d3]

假设 reduction 操作是 SUM。

最终应该得到:

1
2
3
4
S0 = a0 + b0 + c0 + d0
S1 = a1 + b1 + c1 + d1
S2 = a2 + b2 + c2 + d2
S3 = a3 + b3 + c3 + d3

所以 AllReduce 的最终目标是:

1
2
3
4
GPU0: [S0, S1, S2, S3]
GPU1: [S0, S1, S2, S3]
GPU2: [S0, S1, S2, S3]
GPU3: [S0, S1, S2, S3]

注意这里其实包含两个不同的问题:

问题 1:怎样把 a0、b0、c0、d0 合成 S0,把 a1、b1、c1、d1 合成 S1,…… 这是 Reduce

然后 问题 2:S0、S1、S2、S3 已经算出来以后,怎样让每一个 GPU 都拥有全部 S0~S3?这是 Gather / Dissemination

因此可以自然地把 AllReduce 拆成:

1
2
3
4
5
6
7
8
9
10
            AllReduce


┌───────┴───────┐
│ │
▼ ▼
ReduceScatter AllGather

把结果算出来 把结果传播给所有 GPU
每个 GPU 留一块

2. ReduceScatter 做了什么?

ReduceScatter 并不是让每个 GPU 都立刻得到完整结果。

它做的是:

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
所有 GPU 的输入


Reduce


[S0, S1, S2, S3]

│ Scatter


最后变成类似:

GPU0: [S1]
GPU1: [S2]
GPU2: [S3]
GPU3: [S0]

具体哪个 rank 最后拥有哪个 chunk,取决于 Ring 的调度和 chunk 编号方式;关键不是编号,而是:

ReduceScatter 完成后,每一个最终 reduced chunk 已经被完整计算出来,但它们分散在不同 GPU 上。

也就是说:

1
2
3
4
5
6
7
8
9
完整逻辑结果:

[S0, S1, S2, S3]


物理分布:

GPU0 GPU1 GPU2 GPU3
S1 S2 S3 S0

此时缺的已经不是 reduction——所有的 S0~S3 都已经算好了,缺的是:

“把这些 chunk 复制给其他 GPU。”

这正好就是 AllGather。

3. AllGather 接着做什么?

ReduceScatter 结束时:

1
2
3
4
GPU0: S1
GPU1: S2
GPU2: S3
GPU3: S0

执行 AllGather 后:

1
2
3
4
GPU0: [S0, S1, S2, S3]
GPU1: [S0, S1, S2, S3]
GPU2: [S0, S1, S2, S3]
GPU3: [S0, S1, S2, S3]

这已经和直接执行 AllReduce 的输出完全相同。

所以:

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
ReduceScatter

│ 先算出所有 reduced chunk
│ 但每个 GPU 只保存其中一块


GPU0: S1
GPU1: S2
GPU2: S3
GPU3: S0


│ AllGather


GPU0: S0 S1 S2 S3
GPU1: S0 S1 S2 S3
GPU2: S0 S1 S2 S3
GPU3: S0 S1 S2 S3

因此:

AllReduce = ReduceScatter + AllGather

4. Ring AllReduce 是怎么实现这个过程的?

论文 Fig. 4 给出了 NCCL 在 4 个 GPU 上执行 Ring AllReduce 的具体过程。

假设 Ring 为:

1
2
3
4
GPU0 ──→ GPU1
↑ │
│ ↓
GPU3 ←── GPU2

也就是:

1
GPU0 → GPU1 → GPU2 → GPU3 → GPU0

每个 GPU:

  • 从左侧 / predecessor 接收数据
  • 向右侧 / successor 发送数据

整个 Ring AllReduce 可以看成:

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
┌──────────────────────┐
│ ReduceScatter │
│ │
│ 数据一边绕 Ring │
│ 一边做 reduction │
└──────────┬───────────┘


每个 GPU 得到一个
完整 reduced chunk


┌──────────────────────┐
│ AllGather │
│ │
│ reduced chunk 继续 │
│ 沿 Ring 传播 │
└──────────┬───────────┘


每个 GPU 得到全部 chunk

5. 第一阶段:Ring ReduceScatter

仍然使用:

1
2
3
4
GPU0: [a0, a1, a2, a3]
GPU1: [b0, b1, b2, b3]
GPU2: [c0, c1, c2, c3]
GPU3: [d0, d1, d2, d3]

Ring:

1
GPU0 → GPU1 → GPU2 → GPU3 → GPU0

Round 1

各 GPU 同时发送一个不同的 chunk:

1
2
3
4
GPU0: a0 ─────────→ GPU1
GPU1: b1 ─────────→ GPU2
GPU2: c2 ─────────→ GPU3
GPU3: d3 ─────────→ GPU0

接收方拿到以后,与自己的同编号 chunk 做 reduction:

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
GPU1:

a0 + b0

GPU2:

b1 + c1

GPU3:

c2 + d2

GPU0:

d3 + a3

因此现在 Ring 上出现了 4 个 partial sum:

1
2
3
4
chunk0: a0 + b0
chunk1: b1 + c1
chunk2: c2 + d2
chunk3: d3 + a3

Round 2

这些 partial sum 继续向下一个 GPU 传:

1
2
3
4
GPU1 ── (a0+b0) ──→ GPU2
GPU2 ── (b1+c1) ──→ GPU3
GPU3 ── (c2+d2) ──→ GPU0
GPU0 ── (d3+a3) ──→ GPU1

每个接收 GPU 再把自己的对应 chunk 加进去:

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
GPU2:

a0 + b0 + c0

GPU3:

b1 + c1 + d1

GPU0:

c2 + d2 + a2

GPU1:

d3 + a3 + b3

现在每个 partial sum 已经包含 3 个 GPU 的贡献。

Round 3

再向前传一次:

1
2
3
4
GPU2 ── (a0+b0+c0) ──→ GPU3
GPU3 ── (b1+c1+d1) ──→ GPU0
GPU0 ── (c2+d2+a2) ──→ GPU1
GPU1 ── (d3+a3+b3) ──→ GPU2

接收方加入自己的最后一部分:

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
GPU3:

a0+b0+c0+d0 = S0

GPU0:

a1+b1+c1+d1 = S1

GPU1:

a2+b2+c2+d2 = S2

GPU2:

a3+b3+c3+d3 = S3

于是 ReduceScatter 阶段完成:

1
2
3
4
GPU0: S1
GPU1: S2
GPU2: S3
GPU3: S0

最重要的一点是:

到这里 reduction 已经全部完成了,此时不需要再做任何加法。

6. 第二阶段:Ring AllGather

现在 Ring 中分散着:

1
2
3
4
GPU0: S1
GPU1: S2
GPU2: S3
GPU3: S0

接下来只需要把这些完整的 reduced chunk 沿 Ring 传播。

注意:

  • ReduceScatter 阶段:receive + reduce + send
  • 而 AllGather 阶段:receive + copy + send不再进行 reduce

AllGather Round 1

1
2
3
4
GPU0 ── S1 ──→ GPU1
GPU1 ── S2 ──→ GPU2
GPU2 ── S3 ──→ GPU3
GPU3 ── S0 ──→ GPU0

现在每个 GPU 有两个结果 chunk。例如:

1
2
3
GPU0:

S1 + S0

AllGather Round 2

刚收到的数据继续向下传:

1
2
3
4
GPU0 ── S0 ──→ GPU1
GPU1 ── S1 ──→ GPU2
GPU2 ── S2 ──→ GPU3
GPU3 ── S3 ──→ GPU0

此时每个 GPU 有三个结果 chunk。

AllGather Round 3

最后再传播一轮:

1
2
3
4
GPU0 → GPU1
GPU1 → GPU2
GPU2 → GPU3
GPU3 → GPU0

最终:

1
2
3
4
GPU0: [S0, S1, S2, S3]
GPU1: [S0, S1, S2, S3]
GPU2: [S0, S1, S2, S3]
GPU3: [S0, S1, S2, S3]

这就是完整的 AllReduce。

7. 用一张图看完整 Ring AllReduce

输入:

1
2
3
4
GPU0: a0 a1 a2 a3
GPU1: b0 b1 b2 b3
GPU2: c0 c1 c2 c3
GPU3: d0 d1 d2 d3
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62




=========================================
ReduceScatter
=========================================

数据沿 Ring 流动:

GPU0 → GPU1 → GPU2 → GPU3 → GPU0

每经过一个 GPU:

partial_sum
+
local_chunk

new_partial_sum

经过足够多的 GPU 后:

GPU0: S1
GPU1: S2
GPU2: S3
GPU3: S0

其中:

S0 = a0+b0+c0+d0
S1 = a1+b1+c1+d1
S2 = a2+b2+c2+d2
S3 = a3+b3+c3+d3






=========================================
AllGather
=========================================

S0、S1、S2、S3 已经全部算好。

它们继续沿 Ring 传播:

GPU0 → GPU1 → GPU2 → GPU3 → GPU0

只:

recv + copy + send

不再 reduce。




GPU0: S0 S1 S2 S3
GPU1: S0 S1 S2 S3
GPU2: S0 S1 S2 S3
GPU3: S0 S1 S2 S3

所以:

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
                  AllReduce


┌────────────┴────────────┐
│ │
▼ ▼

ReduceScatter AllGather

recv + reduce + send recv + copy + send
│ │
▼ ▼
每人得到一个完整结果块 每人得到所有结果块

└────────────┬────────────┘



完整 AllReduce 结果

8. Fig. 4 中 NCCL primitive 到底是什么意思?

论文没有只停留在”ReduceScatter + AllGather”这个抽象描述,而是进一步展示了 NCCL Ring AllReduce 内部使用的 primitive。

对于 k 个 GPU,论文 Table V 给出的一个 loop iteration 是:

1
2
3
4
5
6
7
8
9
Step 0:      send

Step 1 ~ k-2: recvReduceSend

Step k-1: recvReduceCopySend

Step k ~ 2k-3: recvCopySend

Step 2k-2: recv

对于 4 个 GPU:k = 4,因此是:

1
2
3
4
5
6
7
8
9
10
11
Step 0: send

Step 1: recvReduceSend
Step 2: recvReduceSend

Step 3: recvReduceCopySend

Step 4: recvCopySend
Step 5: recvCopySend

Step 6: recv

总共:

1
2
3
2k - 1
=
7 primitive steps

9. 这些 primitive 分别在干什么?

send

1
2
3
4
读取本地 chunk


发送给下一个 GPU

用于启动 Ring 中的数据流动。

recvReduceSend

这是 ReduceScatter 阶段最关键的 primitive:

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
从前一个 GPU 收数据


recv


和自己的 local chunk 做 reduction


reduce


把 partial result 发给下一个 GPU


send

也就是:recv + reduce + send。数据每经过一个 GPU,就多融合一份输入。

recvReduceCopySend

这是 Fig. 4 里最值得注意的一步。它位于:

1
2
3
4
5
ReduceScatter

边界

AllGather

这个 primitive 同时做:

  • recv
  • 最后一次 reduce
  • copy 到自己的 output buffer
  • 把已经完全 reduce 好的 chunk 发给下一张 GPU

也就是:

1
2
3
4
5
6
7
8
recv

final reduce

得到完整 Sx
├────────→ copy 到本地 output

└────────→ send 给下一 GPU

因此它实际上把两个阶段的边界连接起来了:

1
2
3
ReduceScatter 最后一步
+
AllGather 第一次传播

recvCopySend

进入 AllGather 后已经没有 reduction,所以:

1
2
3
4
5
recv

copy 到自己的 output buffer

send 给下一 GPU

即:recv + copy + send。这里传播的是已经完整计算好的:

1
S0 / S1 / S2 / S3

recv

最后一块数据只需要接收并放到对应位置:

1
2
3
recv

AllReduce 完成

此时所有 GPU 都拥有完整结果。

10. 为什么 Fig. 4 是 2k-1 个 step,而我们常说两个阶段各 k-1 轮?

这里非常容易混淆。

从算法语义 / 通信轮次理解 Ring AllReduce:

1
2
ReduceScatter : k - 1 rounds
AllGather : k - 1 rounds

因此通常写成 2(k - 1) 个 Ring 通信轮次。

但 Fig. 4 分析的是 NCCL 内部 primitive 的执行序列。它把:

  • 最开始的 send
  • 最后的 recv

也分别列成 primitive step,并且在两个阶段的交界处使用 recvReduceCopySend,把 final reduction + copy + 下一阶段的 send 融合起来。

所以论文的 primitive 计数是:

1
1 + (k - 2) + 1 + (k - 2) + 1 = 2k - 1

这和 “ReduceScatter + AllGather” 并不矛盾——两者描述的粒度不同:

1
2
3
k-1 + k-1   描述:算法层面的 Ring rounds

2k-1 描述:Fig.4 中 NCCL primitive sequence

11. 所以”AllReduce = ReduceScatter + AllGather”应该怎么准确理解?

最准确的说法不是”NCCL 内部一定真的调用 ncclReduceScatter(...) 然后 ncclAllGather(...)“,而应该说:

从 collective 的数据语义来看,AllReduce 可以分解为 ReduceScatter 和 AllGather;Ring AllReduce 的实现也确实表现为前半段进行 ReduceScatter-like 的逐块归约,后半段进行 AllGather-like 的结果传播。

NCCL 的 Ring 实现会进一步做 primitive 融合(recvReduceCopySend),所以实际执行不是两个完全独立的 collective kernel 简单拼接。

可以理解为:

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
理论 / 语义:

AllReduce = ReduceScatter + AllGather

而 NCCL Ring 实现:

send

recvReduceSend

recvReduceSend

recvReduceCopySend ← 两阶段边界被融合

recvCopySend

recvCopySend

recv

12. 为什么 Ring 要把数据切成 chunk?

如果整个 tensor 都由一个 GPU 收集并做 reduction:

1
2
3
GPU0 ← GPU1
GPU0 ← GPU2
GPU0 ← GPU3

GPU0 很容易成为瓶颈。

Ring 的做法是把 N 个元素切成 k 份:

1
2
3
4
5
Tensor:

┌──────┬──────┬──────┬──────┐
│ C0 │ C1 │ C2 │ C3 │
└──────┴──────┴──────┴──────┘

让所有 GPU 同时 send + recv + reduce

1
2
3
4
GPU0 → GPU1
GPU1 → GPU2
GPU2 → GPU3
GPU3 → GPU0

四条链路可以同时工作。

因此不存在唯一的 reduction root,不同 chunk 的最终 reduction 分散到不同 GPU:

1
2
3
4
GPU0 负责一个 chunk
GPU1 负责一个 chunk
GPU2 负责一个 chunk
GPU3 负责一个 chunk

这正是 ReduceScatter 的含义。然后再利用同一个 Ring 做 AllGather。

13. Ring AllReduce 的通信量

假设每个 GPU 上的数据总量为 N bytes,一共有 k 个 GPU,每个 chunk 大约是 N/k。

ReduceScatter 有 k-1 轮,每一轮每个 GPU 发送一个 chunk:

1
ReduceScatter 每 GPU 发送量 = (k-1) × N/k = N(k-1)/k

AllGather 同样:

1
AllGather 每 GPU 发送量 = N(k-1)/k

因此 Ring AllReduce 每个 GPU 总发送量约为:

1
2N(k-1)/k

当 GPU 数量很大时,(k-1)/k ≈ 1,所以:

1
总发送量 ≈ 2N

这也是 Ring AllReduce 能较好利用链路带宽的重要原因之一:所有 GPU 都持续承担通信,而不是让一个 root 承担全部数据传输。

14. 最后总结

一句话:

ReduceScatter 负责”把每个结果 chunk 算完整”,AllGather 负责”让每个 GPU 都拿到这些已经算完整的 chunk”,两步接起来后,每个 GPU 都得到完整 reduction 结果,所以等价于 AllReduce。

结合 Fig. 4:

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
           Ring AllReduce




┌─────────────────────────────┐
│ ReduceScatter-like │
│ │
│ send │
│ recvReduceSend │
│ recvReduceSend │
│ │
│ partial sum 不断沿 Ring │
│ 流动并累加 │
└──────────────┬──────────────┘

│ recvReduceCopySend

│ 最后一次 reduce
│ + copy
│ + send

┌─────────────────────────────┐
│ AllGather-like │
│ │
│ recvCopySend │
│ recvCopySend │
│ recv │
│ │
│ 完整 reduced chunk 沿 Ring │
│ 传播,不再做 reduction │
└──────────────┬──────────────┘



每个 GPU 都得到完整 reduction result

所以最终可以记成:

AllReduce = ReduceScatter + AllGather

但在 NCCL Ring 的实际实现中,更准确地说是:

AllReduce = ReduceScatter-like reduction phase + AllGather-like dissemination phase

并且两阶段的边界通过 recvReduceCopySend 进行了融合。