Multi-GPU and Distributed Training

all_reduce, all_gather and broadcast

Six small operations are the entire vocabulary distributed PyTorch speaks — combine into one, copy from one to all, hand out pieces, collect pieces — and everything from DDP to FSDP is built out of them.

On this page 5
  1. Why they are built in
  2. How it works
  3. A real example you have seen
  4. Remember this
  5. What to learn next

One lesson, three depths. Pick the one that fits you today — you can switch any time.

Beginner — No maths. Plain English.

A collective is a sentence every worker says at the same time, so that they end up sharing information.

Think of four friends splitting a restaurant bill. "Everybody shout your share, and we all remember the total" — that is one kind of sentence. "I know the total, listen while I tell you" is another. "Here, each of you take one item from this list" is a third.

There are about six such sentences, and that is the entire language. Distributed training says nothing else.

The important rule is that everyone must say the sentence. If three friends shout and one stays silent, the other three stand there waiting forever.

Why they are built in

You could write this yourself with network sockets. Nobody should. Getting a sum across eight GPUs on four machines to be fast, and correct, and not deadlock, is a specialist's job.

So the collectives ship as library calls. NVIDIA's are called NCCL, said "nickel", and they know about the fast cables between GPUs. The CPU ones are called gloo.

How it works

                 rank0   rank1   rank2   rank3
start:             1       2       3       4

all_reduce SUM:   10      10      10      10     (combine, everyone keeps it)
reduce   to 0:    10       ?       ?       ?     (combine, one keeps it)
broadcast from 0:  1       1       1       1     (one tells everyone)
all_gather:     [1234]  [1234]  [1234]  [1234]   (everyone collects everything)

A real example you have seen

A cricket scoreboard. Each fielder does not track the total; the scorer collects every run and publishes one number that all the players then share. Collect, combine, publish — the same shape as all_reduce.

Remember this

  • A collective is one call that every worker must make, or the job hangs.
  • all_reduce = combine and everyone keeps the answer. This is what DDP uses.
  • broadcast = one worker's value replaces everyone else's.

What to learn next

Developer — Code and libraries.

Setup

bash
pip install torch

Captured on CPU with the gloo backend and two processes. No GPU needed for anything except the last note.

The four you will use most

collectives.py
import torch
import torch.distributed as dist

dist.init_process_group(backend="gloo")
rank, world = dist.get_rank(), dist.get_world_size()

def show(label, value):
    print(f"rank {rank} | {label:32} {value}")

# broadcast: one rank's value replaces everyone else's
t = torch.tensor([rank * 100])
show("before broadcast", t.tolist())
dist.broadcast(t, src=0)
show("after broadcast(src=0)", t.tolist())

# all_reduce: combine, everyone keeps the answer
t = torch.tensor([float(rank + 1)])
dist.all_reduce(t, op=dist.ReduceOp.SUM)
show("all_reduce SUM", t.tolist())

# reduce: combine, only the destination keeps the answer
t = torch.tensor([float(rank + 1)])
dist.reduce(t, dst=0, op=dist.ReduceOp.SUM)
show("reduce SUM to dst=0", t.tolist())

# all_gather: everyone ends up with everyone's tensor, in rank order
mine = torch.tensor([rank * 10 + 1])
bucket = [torch.zeros_like(mine) for _ in range(world)]
dist.all_gather(bucket, mine)
show("all_gather", [b.item() for b in bucket])

dist.barrier()          # nobody moves past this line until everyone arrives
show("past the barrier", "")

dist.destroy_process_group()
bash
torchrun --nproc_per_node=2 collectives.py
Output
rank 0 | before broadcast                 [0]
rank 0 | after broadcast(src=0)           [0]
rank 0 | all_reduce SUM                   [3.0]
rank 0 | reduce SUM to dst=0              [3.0]
rank 0 | all_gather                       [1, 11]
rank 0 | past the barrier
rank 1 | before broadcast                 [100]
rank 1 | after broadcast(src=0)           [0]
rank 1 | all_reduce SUM                   [3.0]
rank 1 | reduce SUM to dst=0              [2.0]
rank 1 | all_gather                       [1, 11]
rank 1 | past the barrier

Lines are grouped by rank for reading; the real run interleaves them.

Two rows deserve a second look.

after broadcast: rank 1 arrived holding 100 and left holding rank 0's 0. Broadcast overwrites. Its argument is named src because the data flows from that rank.

reduce SUM to dst=0: rank 0 has the correct total, 3.0. Rank 1 shows 2.0, which is its own untouched input. After reduce, only the destination's tensor is meaningful — the others hold whatever the backend left there. Read the wrong one and you get a number that looks plausible and is wrong.

Handing pieces out and collecting them back

scatter_gather.py
import torch
import torch.distributed as dist

dist.init_process_group(backend="gloo")
rank, world = dist.get_rank(), dist.get_world_size()

def show(label, value):
    print(f"rank {rank} | {label:26} {value}")

# scatter: rank 0 hands out one piece each
buf = torch.zeros(1)
pieces = [torch.tensor([float(i) * 11]) for i in range(world)] if rank == 0 else None
dist.scatter(buf, pieces, src=0)
show("scatter from 0", buf.tolist())

# gather: everyone sends one piece to the destination
mine = torch.tensor([float(rank) + 0.5])
bucket = [torch.zeros(1) for _ in range(world)] if rank == 0 else None
dist.gather(mine, bucket, dst=0)
show("gather to 0", [b.item() for b in bucket] if rank == 0 else "not the destination")

dist.destroy_process_group()
Output
rank 0 | scatter from 0             [0.0]
rank 0 | gather to 0                [0.5, 1.5]
rank 1 | scatter from 0             [11.0]
rank 1 | gather to 0                not the destination

scatter and gather are the one-sided halves of all_gather. The list argument exists only on the source or destination rank, and must be None everywhere else — passing a list on the wrong rank is a common TypeError.

The whole vocabulary

callwho provides datawho ends up with the answer
broadcastone (src)everyone
reduceeveryoneone (dst)
all_reduceeveryoneeveryone
gathereveryoneone (dst)
all_gathereveryoneeveryone
scatterone (src)everyone, a piece each
reduce_scattereveryoneeveryone, a piece of the combined result
barriernobodynobody — it only synchronises

all_reduce is what DDP calls on your gradients. all_gather and reduce_scatter are what FSDP calls on your parameters. Everything else in this section is a composition of these.

Backends do not support the same set

Not every backend implements every call. Running dist.reduce_scatter under gloo produces this, which is a real captured error and not a hypothetical:

Output
RuntimeError: ProcessGroupGloo does not support reduce_scatter

That is why FSDP is a GPU technique in practice: the operations it is built from exist in NCCL and not in the CPU backend. torch.distributed.is_nccl_available() and the backend support table in the PyTorch docs are worth checking before you design around an op.

Common mistakes

A collective inside an if. The classic is if rank == 0: dist.all_reduce(...). Rank 0 waits for partners who never come, and the job hangs until the timeout. Whole lesson: debugging distributed hangs.

Mismatched shapes or dtypes. Every rank must pass the same shape and dtype. A rank whose last batch is smaller will hang or corrupt memory. Pad, or drop the ragged batch.

Reading a non-destination tensor after reduce or gather. As the output above shows, rank 1's 2.0 looks like a real answer. It is not.

Assuming the operation is in-place when it is not. all_reduce modifies its tensor in place. all_gather writes into the list you supply and leaves the input alone. Mixing those up produces silently wrong results.

Calling collectives on tensors that need gradients. These are plain data movement operations with no autograd support. For differentiable versions, look at torch.distributed.nn.functional.

Try it yourself

Change ReduceOp.SUM to ReduceOp.MAX, then ReduceOp.PRODUCT, and predict each output before running. Then run with --nproc_per_node=3 and work out what all_gather prints before you look.

What to learn next

Researcher — Mathematics and papers.

Cost model

The two numbers that matter for a collective are latency $\alpha$ (per message) and inverse bandwidth $\beta$ (per byte). For $N$ ranks and a payload of $S$ bytes:

  • Ring all-reduce is two phases of $N-1$ steps each, moving $2S\frac{N-1}{N}$ bytes per rank in total: cost $\approx 2(N-1)\alpha + 2S\frac{N-1}{N}\beta$. The bandwidth term is independent of $N$ in the limit, which is the Patarasuk–Yuan bandwidth-optimality result. The latency term is linear in $N$, so rings are poor for small payloads on large clusters.
  • Tree / recursive-halving-doubling costs $\approx 2\log_2(N)\,\alpha + \dots\beta$ — better latency, worse bandwidth. NCCL picks between ring and tree by message size and topology at runtime, which is why the same code shows different scaling curves at different model sizes.
  • all_gather moves $S\frac{N-1}{N}$ bytes per rank, reduce_scatter the same. The identity $\text{all_reduce} = \text{reduce_scatter} + \text{all_gather}$ is exactly how ring all-reduce is implemented, and it is also why FSDP's communication volume is $1.5\times$ DDP's: DDP pays one all-reduce per step, FSDP pays an all-gather in forward, an all-gather in backward and a reduce-scatter.

Ordering and stream semantics

Collectives on a process group must be issued in the same order on every rank; NCCL matches operations positionally, not by name or tag. Two groups issued in opposite orders on two ranks deadlock. Under NCCL, calls are enqueued on the current CUDA stream and return immediately, so the CPU races ahead; correctness is maintained by stream ordering, and async_op=True returns a Work handle whose wait() inserts the dependency rather than blocking the host. This asynchrony is what lets DDP overlap bucket reductions with the remainder of the backward pass, and it is also why a NCCL error frequently surfaces at an unrelated later line.

ReduceOp covers SUM, PRODUCT, MIN, MAX, BAND, BOR, BXOR, plus PREMUL_SUM. Floating-point reduction is not associative, so the result depends on the topology-determined order; runs at different world sizes are not bitwise comparable.

Where the primitives reappear

Parameter-server designs use reduce plus broadcast; ring all-reduce replaced them because the server's bandwidth was the bottleneck (Horovod, 2018). Tensor parallelism uses all_reduce after a row-parallel matmul and all_gather after a column-parallel one. Sequence and context parallelism use all_to_all, the one primitive absent from the table above, exposed as dist.all_to_all_single. Expert routing in mixture-of-experts models is all_to_all twice per layer, which is why MoE throughput is dominated by interconnect quality.

References

  • Patarasuk and Yuan (2009), Bandwidth optimal all-reduce algorithms for clusters of workstations — the ring result and its cost model.
  • Thakur et al. (2005), Optimization of Collective Communication Operations in MPICH — the algorithm-selection-by-message-size approach NCCL also follows.
  • Sergeev and Del Balso (2018), Horovod — the move from parameter servers to ring all-reduce in deep learning.
  • NVIDIA, NCCL Developer Guide — supported operations, topology detection, and the environment variables that matter when it goes wrong.

What to learn next