← All Posts
Deep Learning · Popular Videos · Umar Jamil· Part 7 · Chapter 19

Distributed Communication Collectives

Six primitives cover everything. Every parallelism strategy in this series is a choice of where to place a collective and which one to use. Learn the six, learn their cost models, and a distributed training framework stops being mysterious and becomes an accounting exercise.

The setting

A communicator (a process group) is an ordered set of $P$ ranks, each holding a tensor. A collective is an operation that every rank in the group calls, with matching arguments, that transforms the group's data as a whole. Two facts shape everything:

Collectives are synchronizing

Every rank must call the same collectives in the same order. A rank that skips one, or calls them in a different order, deadlocks the group or silently pairs up mismatched buffers.

Ranks can belong to many groups

The same device can be in a data-parallel group along one mesh axis and a tensor-parallel group along another. Groups are how a device mesh is expressed in code.

The six primitives

Let each rank start with a tensor of $N$ bytes, and write $x_p$ for rank $p$'s data. Let $\oplus$ be an elementwise reduction, usually sum.

CollectiveBeforeAfterBytes each rank moves
Broadcastroot has $x$everyone has $x$$\approx N$
Reducerank $p$ has $x_p$root has $\bigoplus_p x_p$$\approx N$
All-reducerank $p$ has $x_p$everyone has $\bigoplus_p x_p$$2N\frac{P-1}{P}$
All-gatherrank $p$ has shard $x_p$ of size $N/P$everyone has $[x_0;\dots;x_{P-1}]$$N\frac{P-1}{P}$
Reduce-scatterrank $p$ has full-size $x_p$rank $p$ has shard $p$ of $\bigoplus_q x_q$$N\frac{P-1}{P}$
All-to-allrank $p$ has $P$ shards, one per destinationrank $p$ has shard $p$ from every rank$N\frac{P-1}{P}$
The identity that explains the table.
$$\textbf{all-reduce}=\textbf{reduce-scatter}\;+\;\textbf{all-gather}.$$
First reduce and shard, so each rank ends up owning the finished sum of one slice; then gather those finished slices everywhere. Each half moves $N\frac{P-1}{P}$ bytes, so all-reduce moves exactly twice as much. This decomposition is not merely a proof device — ZeRO and FSDP exploit it by keeping the intermediate sharded state instead of gathering it back.

All-to-all is the odd one out: it is a distributed transpose, moves no reduction, and is the primitive that expert parallelism is built on.

The cost model

Communication time is modelled with two constants: a per-message latency $\alpha$ and an inverse bandwidth $\beta$ measured in seconds per byte. A single message of $n$ bytes costs

$$T=\alpha+\beta n.$$

An algorithm that takes $s$ sequential communication steps and moves $n$ bytes per rank therefore costs $s\alpha+\beta n$. Two regimes follow immediately, and they call for different algorithms:

Small messages are latency-bound

The $s\alpha$ term dominates. You want an algorithm with few steps, so a tree with $O(\log P)$ depth wins even though it moves more data.

Large messages are bandwidth-bound

The $\beta n$ term dominates. You want the algorithm that moves the least data per rank, which is the ring, even though it takes $O(P)$ steps.

Gradient buffers in training are megabytes to gigabytes, so training is firmly in the bandwidth-bound regime. That is why the ring algorithm is the one worth understanding in detail.

Ring all-reduce, derived

Arrange the ranks in a logical ring where rank $p$ only ever sends to rank $p+1 \bmod P$. Split each rank's tensor into $P$ chunks of $N/P$ bytes. The algorithm runs in two phases of $P-1$ steps each.

Ring all-reduce on rank $r$
1phase 1 — reduce-scatter, for $s=0,\dots,P-2$:
2send chunk $(r-s)\bmod P$ to rank $r+1$
3receive chunk $(r-s-1)\bmod P$ from rank $r-1$ and add it into the local copy
4after $P-1$ steps, rank $r$ holds the fully reduced chunk $(r+1)\bmod P$
5phase 2 — all-gather, for $s=0,\dots,P-2$:
6send chunk $(r+1-s)\bmod P$ to rank $r+1$
7receive chunk $(r-s)\bmod P$ from rank $r-1$ and overwrite the local copy
8after $P-1$ more steps, every rank holds every fully reduced chunk

The only difference between the phases is add versus overwrite. Counting the traffic:

$$T_{\text{ring}}=\underbrace{2(P-1)}_{\text{steps}}\alpha+\underbrace{2\frac{P-1}{P}N}_{\text{bytes per rank}}\beta.$$
All-reduce does not get more expensive as the cluster grows. As $P\to\infty$ the bandwidth term tends to $2N\beta$, a constant. Doubling the number of ranks does not double the bytes each one must move. The latency term does grow linearly, which is why very large rings are replaced by hierarchical schemes, but the bandwidth cost is essentially free of $P$. This single fact is what makes data parallelism scale to thousands of devices.
1 contribution partially reduced fully reduced (all $P$)
phase
bytes sent per rank so far
fully reduced cells
Row $r$, column $c$ shows the state of chunk $c$ on rank $r$. Watch the diagonal of finished chunks appear at the end of phase 1, then sweep around the ring during phase 2. Bytes are in units of $N/P$, one chunk.

When the ring is the wrong choice

The ring is bandwidth-optimal but its latency grows as $2(P-1)\alpha$. For a small tensor across many nodes that term dominates, and a tree does much better:

AlgorithmLatency termBandwidth termBest for
Ring$2(P-1)\alpha$$2\frac{P-1}{P}N\beta$large tensors, the training case
Recursive halving/doubling$2\log_2 P\cdot\alpha$$2\frac{P-1}{P}N\beta$power-of-two $P$, medium tensors
Binary tree (reduce + broadcast)$2\log_2 P\cdot\alpha$$2N\beta$ at the root linkssmall tensors, latency-bound

Production libraries choose per call, based on message size and topology, and typically run hierarchically: reduce within a node over the fast intra-node fabric, then all-reduce across nodes over the slower network with only one representative per node, then broadcast back down inside each node. Since intra-node links are often an order of magnitude faster than inter-node ones, this turns a $P$-way problem into a much smaller one.

The slowest link sets the price. A collective spanning two nodes runs at inter-node speed no matter how fast the GPUs inside are. This is the reason the parallelism-ordering rules in the device-mesh note exist: put the chattiest strategy on the fastest axis.

Making communication free

The best collective is one you never wait for. Gradients become available progressively during the backward pass, starting with the last layer, so there is a long window in which to send them while the remaining layers are still computing.

1

Bucket. Issuing one collective per parameter tensor pays $\alpha$ hundreds of times. Group gradients into buckets of a few tens of megabytes and issue one collective per bucket. This is a pure latency optimization.

2

Launch asynchronously. Fire the collective as soon as a bucket is full and keep computing. The work is genuinely overlapped only if it runs on a separate stream from the compute kernels.

3

Wait as late as possible. Block on the handles only just before the optimizer step needs the values.

When overlap works, communication time vanishes from the critical path entirely and MFU stays high. When it does not — because buckets are too small, or the collective shares a stream with compute — you see it immediately as a gap in a profiler trace and as a disappointing utilization number.

Takeaway

All-reduce is reduce-scatter followed by all-gather, costs $2N\frac{P-1}{P}$ bytes per rank, and is therefore almost independent of cluster size. Everything else is choosing the right algorithm for the message size, putting the chattiest collective on the fastest link, and hiding it behind compute.

Check yourself