Skip to content

Collectives over processes: ring all-reduce

ModuleL11.2 · side · Python · Pass 9 (optional) · 4 h
You buildpython/tinyllm/dist/comm.py: chunk_bounds, Comm (send, recv, reduce_scatter, all_gather, all_reduce, broadcast, barrier), spawn
Contractcourse/contracts/py/tinyllm/dist/comm.pyi
Testscourse/tests/L11.2/ (what they check: section 4) · your own tests in python/tests/l11-2-comm/, rung R5, graded by mutation (threshold 0.80, every pitfall mutant required)
Needsnothing to build first · reading: M05.1 counting bytes
Used byL11.3 DDP and ZeRO run every collective through Comm
Milestonenone; this optional side module is not part of MS-L11
Optional depthPatarasuk and Yuan, “Bandwidth Optimal All-reduce Algorithms for Clusters of Workstations” (JPDC, 2009); Thakur, Rabenseifner, and Gropp, “Optimization of Collective Communication Operations in MPICH” (IJHPCA, 2005); the NCCL documentation on collective operations
  • An all-reduce is a reduce-scatter followed by an all-gather; on a ring each half takes p−1p - 1 steps of one chunk of n/pn/p entries (test_hand_example_ring_allreduce, test_reduce_scatter_and_all_gather).
  • Every rank sends 2(p−1)/p2(p-1)/p of the array whatever pp is, so adding ranks does not add traffic per link (test_bytes_moved_is_2_p_minus_1_over_p).
  • The ring adds in its own order: the result equals numpy.sum to rounding, and every rank holds the same bits (test_allreduce_matches_numpy_sum).
  • A blocking send on a ring deadlocks once messages outgrow the pipe buffer; even ranks send first and odd ranks receive first (test_large_messages_do_not_deadlock).
Terminal window
ol start L11.2 # stubs comm.py into your repo
ol tests L11.2 # read the test catalog first
ol check L11.2 # course tests, then your tests graded by mutation
ol mutate L11.2 # the full mutation grade of your tests
ol diff L11.2 # after passing: your code against the reference

L11.1 made one process train the capstone on a laptop. Every larger run splits the batch over several workers instead: each computes the gradient of its own slice, and before anyone takes a step the gradients must be averaged across all of them. That averaging is an all-reduce, and it is the operation every data-parallel step waits on. On a laptop the workers are processes on the same machine, which is enough to build the real algorithm: processes cannot share Python objects, so every byte moves through an explicit channel, and you can count those bytes. This optional module builds the collectives (all_reduce, reduce_scatter, all_gather, broadcast, barrier) over pipes with the bandwidth-optimal ring, and L11.3 builds data parallelism and ZeRO on top.

SymbolMeaningType / shape
pp (world)number of processesint
rr (rank)this process’s number, 0≤r<p0 \le r < pint
xrx_rrank rr‘s input array, nn entriesfloat64[n]
∑rxr\sum_r x_rthe elementwise sum all ranks wantfloat64[n]
cjc_jchunk jj of the flat array: entries [aj,bj)[a_j, b_j) from chunk_boundsslice
right, leftranks (r+1) mod p(r+1) \bmod p and (r−1) mod p(r-1) \bmod pint
β\betabytes per entry (8 for float64)int

Processes and channels. spawn(fn, p) starts pp processes with multiprocessing’s “spawn” method (a fresh interpreter each, so fn must be importable by name) and connects every pair of ranks with a full-duplex pipe. send(x, dst) writes a copy of the array, recv(src) blocks until the next array from src arrives; messages between two ranks arrive in order. A pipe holds only a few kilobytes: a larger send blocks until the receiver reads.

Why not gather to one rank. The obvious all-reduce sends every array to rank 0, adds, and sends the sum back: rank 0 receives (p−1)n(p-1)n entries and sends (p−1)n(p-1)n, so its link carries traffic that grows with pp, while the other links idle. The ring keeps every link equally busy.

Chunks. Split the flat array into pp contiguous chunks c0,…,cp−1c_0, \dots, c_{p-1}, sizes as numpy.array_split (the first n mod pn \bmod p chunks one entry larger). Rank rr owns chunk rr; ZeRO (L11.3) shards parameters by the same bounds, so they must agree everywhere.

Reduce-scatter on a ring. In step s=0,…,p−2s = 0, \dots, p-2, rank rr sends chunk (r−s−1) mod p(r - s - 1) \bmod p to its right neighbour, receives chunk (r−s−2) mod p(r - s - 2) \bmod p from its left neighbour, and adds what it received into its own copy of that chunk. Follow one chunk cjc_j: it starts at rank j+1j + 1, which sends its xj+1x_{j+1} part right; each rank adds its own part and passes the partial sum on; after p−1p - 1 steps it reaches rank jj holding all pp contributions. So every rank ends holding its own chunk of the sum, (∑rxr)[cr]\big(\sum_r x_r\big)[c_r], having sent p−1p - 1 chunks.

All-gather on a ring. Now each rank has one finished chunk and needs the others. In step ss, rank rr sends chunk (r−s) mod p(r - s) \bmod p right and receives chunk (r−s−1) mod p(r - s - 1) \bmod p from the left. After p−1p - 1 steps every rank has every chunk. All-reduce is reduce-scatter followed by all-gather.

Bytes. Each half sends p−1p - 1 chunks of about n/pn/p entries, so

bytes sent per rank=2 p−1p n β,\text{bytes sent per rank} = 2\,\frac{p-1}{p}\, n\, \beta,

which approaches 2nβ2n\beta and never grows with pp: doubling the ranks halves each chunk. A lower bound says no all-reduce can do better (each rank must send at least (p−1)/p(p-1)/p of its data out and receive the same in), which is why NCCL uses rings for large messages.

Rounding. The sum of chunk cjc_j is accumulated in ring order starting at rank j+1j + 1, not in rank order: floating-point addition is not associative, so the result differs from numpy.sum in the last bits. But each chunk is reduced once and then copied, so all ranks hold identical bits, which is what keeps data-parallel replicas identical.

Deadlock. If every rank does send then recv in a step, and the message is larger than the pipe buffer, every rank blocks in send waiting for its right neighbour to recv, which is itself blocked in send: a cycle of waits, forever. Breaking the cycle needs one rank that receives first. The contract’s rule: even ranks send then receive, odd ranks receive then send. With p≥2p \ge 2 rank 1 is always a receiver-first, so the chain unwinds from there, for odd and even pp alike. MPI calls the alternatives MPI_Sendrecv and non-blocking sends.

Broadcast and barrier. broadcast(x, src) passes xx around the ring from src: each rank receives from the left and forwards right, except the rank just before src, which would send it back to where it started (and leave a message in that pipe that the next collective would read by mistake). A barrier makes every rank wait until all have arrived; it is multiprocessing’s Barrier.

Failure. A rank that raises must stop the whole run with its traceback; otherwise its neighbours wait forever in a recv. spawn collects each rank’s result or traceback, terminates the rest on the first error, and raises TimeoutError if the ranks do not finish in time.

Three ranks hold x0=[1,2,3]x_0 = [1, 2, 3], x1=[10,20,30]x_1 = [10, 20, 30], x2=[100,200,300]x_2 = [100, 200, 300]; n=3n = 3, so each chunk is one entry: $c_0 = $ entry 0, $c_1 = $ entry 1, $c_2 = $ entry 2.

steprank 0 sendsrank 1 sendsrank 2 sendsrank 0 holds afterrank 1 holds afterrank 2 holds after
RS 0c2=3c_2 = 3c0=10c_0 = 10c1=200c_1 = 200c1c_1: 2+200=2022 + 200 = 202c2c_2: 30+3=3330 + 3 = 33c0c_0: 100+10=110100 + 10 = 110
RS 1c1=202c_1 = 202c2=33c_2 = 33c0=110c_0 = 110c0c_0: 1+110=1111 + 110 = 111c1c_1: 20+202=22220 + 202 = 222c2c_2: 300+33=333300 + 33 = 333
AG 0c0=111c_0 = 111c1=222c_1 = 222c2=333c_2 = 333gets c2=333c_2 = 333gets c0=111c_0 = 111gets c1=222c_1 = 222
AG 1c2=333c_2 = 333c0=111c_0 = 111c1=222c_1 = 222gets c1=222c_1 = 222gets c2=333c_2 = 333gets c0=111c_0 = 111

In step RS 0, rank rr sends chunk (r−1) mod 3(r - 1) \bmod 3 and adds what arrives into chunk (r−2) mod 3(r - 2) \bmod 3: rank 0 receives rank 2’s c1=200c_1 = 200 and adds its own 2. After reduce-scatter rank rr holds chunk rr of the sum: [111][111], [222][222], [333][333]; after all-gather every rank holds [111,222,333][111, 222, 333]. Each rank sent 4 entries of 8 bytes, 32 bytes, and 2⋅23⋅3⋅8=322 \cdot \tfrac{2}{3} \cdot 3 \cdot 8 = 32. This is test_hand_example_ring_allreduce.

python/tinyllm/dist/comm.py
def chunk_bounds(n: int, world: int) -> list[tuple[int, int]] # numpy.array_split sizes
class Comm:
rank: int; world: int; bytes_sent: int
def send(self, x, dst: int) -> None; def recv(self, src: int) -> NDArray
def reduce_scatter(self, x) -> NDArray # chunk `rank` of the sum
def all_gather(self, x) -> NDArray # every rank's chunk, in rank order
def all_reduce(self, x, op="sum") -> NDArray # "sum" or "mean", x's shape
def broadcast(self, x, src: int) -> NDArray
def barrier(self) -> None
def spawn(fn, world: int, *args, timeout: float = 60.0) -> list # fn(comm, *args) per rank

A worker is a module-level function fn(comm, *args); spawn returns the workers’ return values by rank. Collectives must be called by every rank in the same order.

TestKINDChecksWhy it matters downstream
test_hand_example_ring_allreduceunitsection 3: chunks, sums, 16 and 32 bytesyou and the test agree on the ring
test_chunk_boundsunitnumpy.array_split sizes for n<23n < 23, p≤5p \le 5ZeRO shards line up (L11.3)
test_allreduce_matches_numpy_sumdifferentialp=2,3,4p = 2, 3, 4 on 35 entries, sum and mean, bitwise-equal ranks, input untouchedDDP replicas stay identical
test_bytes_moved_is_2_p_minus_1_over_ppropertyexactly 2(p−1)nβ/p2(p-1)n\beta/p for p=2,4,5p = 2, 4, 5the bandwidth argument
test_reduce_scatter_and_all_gatherunitthe halves on their own, unequal chunksZeRO stage 2 and 3
test_broadcast_from_every_srcunitevery source, bytes per rank, a clean pipe afterwardsDDP’s initial weights
test_large_messages_do_not_deadlockboundary1 MiB messages on a ring of 3real gradient sizes
test_send_recv_point_to_pointunitorder, dtype, shape, a writable copy, bad ranksthe layer under everything
test_rank_error_propagatesboundarya crash on rank 1 raises its traceback quicklyno silent hangs
test_barrierunita late rank holds everyoneordering side effects
PitfallSymptomCaught by
1. adding the received chunk into the one you just sentwrong sums, different on every ranktest_hand_example_ring_allreduce (mutant s01)
2. one step too few in reduce-scatterpartial sums for p>2p > 2test_allreduce_matches_numpy_sum (mutant s02)
3. one step too few in all-gatherstale chunks on some rankstest_reduce_scatter_and_all_gather (mutant s03)
4. "mean" returning the sumgradients pp times too largetest_allreduce_matches_numpy_sum (mutant s04)
5. every rank sends firsta deadlock as soon as a message exceeds the pipe buffertest_large_messages_do_not_deadlock (mutant s05)
6. accumulating into the caller’s arraythe caller’s gradient changes under ittest_allreduce_matches_numpy_sum (mutant s06)
7. chunks of the wrong sizesZeRO shards and collectives disagreetest_chunk_bounds (mutant s07)
8. counting entries instead of bytesthe bandwidth numbers are off by β\betatest_bytes_moved_is_2_p_minus_1_over_p (mutant s08)
9. broadcasting back to the sourcea stray message corrupts the next collectivetest_broadcast_from_every_src (mutant s09)
10. ignoring a rank’s failurethe run “finishes” with a missing resulttest_rank_error_propagates (mutant s10)
11. a barrier that does not waitside effects racetest_barrier (mutant s11)
12. returning a read-only view of the messagethe receiver cannot write its own arraytest_send_recv_point_to_point (mutant s12)
DirectionModuleHow it uses this
BackM05.1counting bytes per parameter is the same arithmetic as counting bytes per link
ForwardL11.3DDP all-reduces gradients; ZeRO uses reduce-scatter and all-gather on parameter chunks

This module is optional (D20): a laptop capstone needs L11.1, not data parallelism. Without it, L11.3 cannot run.

Your pieceProduction equivalentWhat it addsWhere to look
Comm over pipesNCCLrings and trees over NVLink and InfiniBand, chunked and pipelined so every link is busy at onceNCCL source, src/collectives/
all_reducetorch.distributed.all_reduceprocess groups, backends (NCCL, Gloo, MPI), async handlestorch/distributed/distributed_c10d.py
the ringtree and recursive halving-doublinglower latency for small messages (log⁡p\log p steps instead of pp)Thakur et al. (2005)
spawntorchrunrendezvous across machines, restarts, elastic world sizestorch/distributed/run.py