- All-Reduce、All-Gather 和 Reduce-Scatter 都是分布式训练中的集合通信操作(collective communication)。它们不是两个 GPU 之间的点对点通信,而是由一个通信组中的所有 GPU 共同参与。
- 假设每张 GPU 上都有一个形状相同的张量:
GPU 0: x0
GPU 1: x1
GPU 2: x2
GPU 3: x3
- All-Reduce :聚合,每张卡都得到完整结果。执行求和形式的 All-Reduce Sum:
dist.all_reduce(x, op=dist.ReduceOp.SUM) 后会得到:
GPU 0: x0 + x1 + x2 + x3
GPU 1: x0 + x1 + x2 + x3
GPU 2: x0 + x1 + x2 + x3
GPU 3: x0 + x1 + x2 + x3
- All-Gather:收集所有分片,每张卡都得到完整张量。它可以理解为:每张 GPU 把自己的分片广播给其他所有 GPU。执行 All-Gather 后:
GPU 0: [X0, X1, X2, X3]
GPU 1: [X0, X1, X2, X3]
GPU 2: [X0, X1, X2, X3]
GPU 3: [X0, X1, X2, X3]
- 在FSDP中,平时每张 GPU 上只保存一部分参数:
GPU 0: parameter shard P0
GPU 1: parameter shard P1
GPU 2: parameter shard P2
GPU 3: parameter shard P3
- 某一层即将执行 forward 时,通过 All-Gather 临时恢复完整参数:
GPU 0: [P0, P1, P2, P3]
GPU 1: [P0, P1, P2, P3]
GPU 2: [P0, P1, P2, P3]
GPU 3: [P0, P1, P2, P3]
- Reduce-Scatter:先聚合,再把结果切分到不同 GPU,可以理解为Reduce + Scatter两个操作的组合。
- 假设每张 GPU 上都有一个完整张量,但把它逻辑上分成 4 段:
GPU 0: [a0, b0, c0, d0]
GPU 1: [a1, b1, c1, d1]
GPU 2: [a2, b2, c2, d2]
GPU 3: [a3, b3, c3, d3]
A = a0 + a1 + a2 + a3
B = b0 + b1 + b2 + b3
C = c0 + c1 + c2 + c3
D = d0 + d1 + d2 + d3
- 然后把结果 Scatter 到各 GPU。形式化地,第jjj张GPU得到yj=∑ixi(j)y_j=\sum_ix_i^{(j)}yj=∑ixi(j),其中xi(j)x_i^{(j)}xi(j)表示第iii张GPU输入的第jjj的分片。
GPU 0: A
GPU 1: B
GPU 2: C
GPU 3: D
- 在FSDP backward 时,每张 GPU 都可能计算出当前层的完整梯度:
GPU 0: full gradient G0
GPU 1: full gradient G1
GPU 2: full gradient G2
GPU 3: full gradient G3
- 但 FSDP 不希望每张 GPU 都长期保存完整聚合梯度,因此执行 Reduce-Scatter。这样既完成了数据并行中的梯度同步,又让梯度继续保持分片状态。
GPU 0: shard 0 of (G0 + G1 + G2 + G3)
GPU 1: shard 1 of (G0 + G1 + G2 + G3)
GPU 2: shard 2 of (G0 + G1 + G2 + G3)
GPU 3: shard 3 of (G0 + G1 + G2 + G3)
- All-Reduce = Reduce-Scatter + All-Gather,这是三者之间最重要的关系。
- 假设目标是让所有 GPU 都得到x0+x1+x2+x3x_0+x_1+x_2+x_3x0+x1+x2+x3,而目前有:
GPU 0: x0
GPU 1: x1
GPU 2: x2
GPU 3: x3
- 第一步:Reduce-Scatter,把最终求和结果切成 4 个分片。
GPU 0: sum(x)[shard 0]
GPU 1: sum(x)[shard 1]
GPU 2: sum(x)[shard 2]
GPU 3: sum(x)[shard 3]
- 第二步:All-Gather,收集所有聚合后的分片。
GPU 0: complete sum(x)
GPU 1: complete sum(x)
GPU 2: complete sum(x)
GPU 3: complete sum(x)
- NCCL 中常见的 Ring All-Reduce,本质上就可以看成:
Ring Reduce-Scatter
+
Ring All-Gather
- Continuous Batching(连续批处理)是一种用于大模型推理服务的动态调度方法。其核心思想是:不必等整个 batch 中所有请求都生成完,某个请求一旦结束,就立刻移出 batch,并马上加入新的请求。它也常被称为Iteration-level Scheduling,表示每生成一个 token,就重新调度一次请求集合。
- LLM 推理通常包含两个阶段:Prefill和Decode。其中Decode 阶段具有很强的自回归特性,每一轮 forward 通常只生成一个新 token。
请求输入
↓
Prefill:一次性处理整个 prompt
↓
Decode:每次生成一个 token
- 【为什么需要Continuous Batching?】不同用户请求的输出长度可能差异很大:请求 A:生成 5 个 token;请求 B:生成 50 个 token;请求 C:生成 200 个 token。如果使用传统静态 batching,把这些请求组成一个 batch A,B,CA, B, CA,B,C,那么 A 很快就结束了,但 batch 仍然要继续等待 B、C。A 结束后,它占用的 batch slot 无法立即被新请求使用,GPU 利用率会下降。batch 的有效大小从 3 逐渐下降到 1。GPU 的并行度越来越低。Continuous Batching 就是为了解决这个问题。
时间 →
A: █████ 完成
B: ██████████████████████████████████████████████
C: █████████████████████████████████████████████████████████...
- 假设有三个请求:A:需要生成 2 个 token;B:需要生成 4 个 token;C:需要生成 6 个 token。静态batching的处理过程如下:
时间步 → 1 2 3 4 5 6
A: A1 A2 - - - -
B: B1 B2 B3 B4 - -
C: C1 C2 C3 C4 C5 C6
- Continuous Batching 不要求整个 batch 一起开始、一起结束。假设等待队列中还有:D:生成 3 个 token;E:生成 2 个 token;F:生成 4 个 token。Continuous batching使得当 A 完成后,立刻加入 D,完整处理过程如下:
时间步 → 1 2 3 4 5 6 7 8
Slot 0: A1 A2 D1 D2 D3 F1 F2 F3
Slot 1: B1 B2 B3 B4 E1 E2 G1 G2
Slot 2: C1 C2 C3 C4 C5 C6 H1 H2
- batch 中的请求不断流动,因此叫 Continuous Batching。
请求完成
↓
立即释放对应资源
↓
从等待队列中选取新请求
↓
下一轮 decode 立即加入 batch
- LLM Decode 有一个天然特点:每个活跃请求在每一轮通常只需要生成一个 token。由于 KV Cache的存在,一次 decode forward 的输入不是所有历史 token,而通常只是每个请求最新的 token。所以每个 decode iteration 都天然形成一个 batch。因此 Continuous Batching 的调度粒度通常是 decode iteration,而不是一个完整请求。
- Continuous Batching 中一个重要问题是:新请求需要执行 Prefill,而已有请求只需要执行 Decode,两者如何放在一起?除了naive策略,即先 Prefill,再 Decode(当有新请求时,单独执行它的 Prefill,之后加入 Decode batch),常见的策略还有:
- Mixed Batching:在同一个 iteration 中同时处理若干 decode token + 若干 prefill token;
- Chunked Prefill:把长 prompt 切成多个较小块,然后与 Decode 交错执行。
- Mixed Batching 需要底层 attention kernel 支持不同长度、不同阶段的序列。
Batch:
A: 1 个 decode token
B: 1 个 decode token
C: 1 个 decode token
D: 512 个 prefill token
- Chunked Prefill:先切分,再混合执行。这样可以避免一个长 prompt 长时间阻塞已有请求。
原始 prompt:2048 tokens
切成:
chunk 0: 512 tokens
chunk 1: 512 tokens
chunk 2: 512 tokens
chunk 3: 512 tokens
执行:
Iteration 1: decode A/B/C + prefill D chunk 0
Iteration 2: decode A/B/C + prefill D chunk 1
Iteration 3: decode A/B/C + prefill D chunk 2
Iteration 4: decode A/B/C + prefill D chunk 3