Multi-GPU and Distributed Training

When gradients are synchronised, and no_sync

DDP averages gradients during backward(), not at optimizer.step() — and no_sync() is the switch that skips the averaging so you can accumulate several batches for the price of one exchange.

On this page 5
  1. Why the timing matters
  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.

The workers agree on their corrections during the backward pass — and no_sync lets them stay quiet for a few rounds first.

Think of four surveyors measuring the same field from different corners. Every measurement, they radio each other and average their readings. The radio call costs time. If the field is huge and each measurement is slow, the call is worth it.

But if measurements are quick, the radio becomes the slow part. So they agree to take four measurements each, keep a running total, and radio once at the end. Same average, one quarter of the calls.

no_sync is that agreement.

Why the timing matters

Most people assume the workers talk when the weights are updated. They do not. The talking happens while the corrections are being worked out, layer by layer, back to front.

That is deliberate. The moment the last layer's correction is ready, it can be sent while the rest of the network is still being worked out. The radio call and the thinking happen at the same time, so the call is nearly free.

How it works

normal backward:
  last layer done  -> send it            \
  middle done      -> send it             }  sending overlaps with working
  first layer done -> send it, wait      /
  every worker now holds the SAME average correction

inside no_sync():
  last layer done  -> keep it
  middle done      -> keep it
  first layer done -> keep it
  every worker holds its OWN private correction, nothing was sent

A real example you have seen

A WhatsApp group for a shared expense. You can message the group after every single chai, or note them down and send one total in the evening. The final split is identical. The second way costs far fewer messages.

Remember this

  • Gradients are averaged inside backward(), not at optimizer.step().
  • Sending overlaps with computing, which is why DDP is fast.
  • no_sync() skips the sending so you can accumulate several batches cheaply.

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. The numbers below are small enough to check by hand, which is the point.

Watching the average happen

no_sync_demo.py
import torch
import torch.nn as nn
import torch.distributed as dist
from torch.nn.parallel import DistributedDataParallel as DDP

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

torch.manual_seed(0)
ddp = DDP(nn.Linear(4, 1, bias=False))
torch.manual_seed(10 + rank)                  # different data per rank
x = torch.randn(2, 4)

def grad():
    return ddp.module.weight.grad.flatten()[0].item()

# a plain backward: DDP averages the gradient across ranks before it returns
ddp.zero_grad(set_to_none=True)
ddp(x).sum().backward()
print(f"rank {rank}: plain backward -> grad[0] = {grad():+.4f}")

# inside no_sync, nothing is sent; each rank keeps its own private gradient
ddp.zero_grad(set_to_none=True)
with ddp.no_sync():
    ddp(x).sum().backward()
    print(f"rank {rank}: inside no_sync -> grad[0] = {grad():+.4f}   (local only)")

# the next backward OUTSIDE no_sync syncs the whole accumulated total
ddp(x).sum().backward()
print(f"rank {rank}: after the syncing backward -> grad[0] = {grad():+.4f}")

dist.destroy_process_group()
bash
torchrun --nproc_per_node=2 no_sync_demo.py
Output
rank 0: plain backward -> grad[0] = +0.2714
rank 0: inside no_sync -> grad[0] = +0.3184   (local only)
rank 0: after the syncing backward -> grad[0] = +0.5428
rank 1: plain backward -> grad[0] = +0.2714
rank 1: inside no_sync -> grad[0] = +0.2244   (local only)
rank 1: after the syncing backward -> grad[0] = +0.5428

The two processes print at the same time, so the lines above are grouped by rank for reading; your run will interleave them differently.

Do the arithmetic yourself. The private gradients are 0.3184 and 0.2244. Their mean is 0.2714 — exactly what the plain backward produced on both ranks. After accumulating two backward passes, both ranks hold 0.5428, which is twice that mean.

Three things follow from those six lines:

  1. DDP averages, it does not sum. Doubling the worker count does not double the gradient.
  2. Inside no_sync(), ranks genuinely diverge. The values differ.
  3. Leaving no_sync() syncs everything accumulated so far, in one exchange.

The accumulation pattern in a real loop

accumulate_ddp.py
import contextlib
import torch
import torch.nn as nn
import torch.distributed as dist
from torch.nn.parallel import DistributedDataParallel as DDP

dist.init_process_group(backend="gloo")
rank = dist.get_rank()
ACCUM = 4                                   # micro-batches per optimizer step

torch.manual_seed(0)
ddp = DDP(nn.Linear(4, 1))
opt = torch.optim.SGD(ddp.parameters(), lr=0.1)
torch.manual_seed(20 + rank)
batches = [(torch.randn(2, 4), torch.randn(2, 1)) for _ in range(ACCUM)]

opt.zero_grad(set_to_none=True)
for i, (x, y) in enumerate(batches):
    last = (i == ACCUM - 1)
    # sync on the final micro-batch only; stay silent for the rest
    ctx = contextlib.nullcontext() if last else ddp.no_sync()
    with ctx:
        loss = nn.functional.mse_loss(ddp(x), y) / ACCUM   # scale, or you get 4x
        loss.backward()
    print(f"rank {rank} micro-batch {i}: synced={last}")
opt.step()

w = ddp.module.weight
print(f"rank {rank}: weight checksum after the step {w.sum().item():.6f}")

dist.destroy_process_group()
Output
rank 0 micro-batch 0: synced=False
rank 0 micro-batch 1: synced=False
rank 0 micro-batch 2: synced=False
rank 0 micro-batch 3: synced=True
rank 1 micro-batch 0: synced=False
rank 1 micro-batch 1: synced=False
rank 1 micro-batch 2: synced=False
rank 1 micro-batch 3: synced=True
rank 0: weight checksum after the step -0.686250
rank 1: weight checksum after the step -0.686250

Four micro-batches, one exchange, and both ranks still land on the same weights. That is the whole benefit: 4x fewer collective calls with an identical result.

The / ACCUM matters. Accumulated gradients add up, so without the division you take a step four times too large. This is the same rule as single-GPU gradient accumulation.

When no_sync is worth it

no_sync trades communication for staleness. Nothing is stale here — the maths is exact — so the only cost is memory: gradients live in full precision across the whole accumulation window.

It pays when communication is a real share of step time: small models, slow interconnects, many machines. It pays nothing on one machine with NVLink and a large model, where the all-reduce was already hidden behind the backward pass. Measure before you complicate the loop — the profiler lesson shows how.

Common mistakes

Syncing on every micro-batch by forgetting no_sync. The loop still produces correct results, so nothing warns you. You paid for 4 exchanges and needed 1.

Forgetting to sync on the last micro-batch. Then the ranks never agree, the weights drift apart, and the run silently becomes 4 independent models. Note the last flag in the code above.

Calling zero_grad() inside the accumulation loop. It throws away everything accumulated. It belongs before the loop.

Expecting no_sync to help with find_unused_parameters=True. That flag forces a graph traversal every backward pass to work out which parameters got gradients. It is a separate cost, and it is expensive. Restructure the model to avoid it when you can.

Try it yourself

In the first script, replace the second ddp(x).sum().backward() with a second with ddp.no_sync(): block. Predict the printed values before running. The ranks should stay at their private numbers, doubled.

What to learn next

Researcher — Mathematics and papers.

Where the collective is issued

DDP registers an autograd hook on every parameter with requires_grad=True. Parameters are grouped into buckets, bucket_cap_mb=25 by default, ordered by approximately reverse registration order — an approximation of the order gradients become ready during backward. When every parameter in a bucket has fired its hook, the Reducer issues an asynchronous all_reduce on the flattened bucket. backward() returns only after every bucket's work handle has completed, so the synchronisation point is the end of the backward pass, never optimizer.step().

The reduction is ReduceOp.SUM followed by division by world_size, giving the arithmetic mean. Equivalently DDP computes

$$ \bar{g} = \frac{1}{R} \sum_{r=1}^{R} g_r $$

where $R$ is the world size and $g_r$ is rank $r$'s local gradient over its own micro-batch. When each rank holds $b$ samples and the per-rank loss is a mean over those $b$, $\bar{g}$ equals the gradient of the mean loss over all $Rb$ samples — the identity that learning-rate scaling rests on.

no_sync, precisely

DistributedDataParallel.no_sync() sets require_backward_grad_sync = False for the duration of the context. The autograd hooks still fire and still accumulate into param.grad; the Reducer does not launch its collectives. On the first backward outside the context, the buckets are reduced carrying the full accumulated sum, so $K$ micro-batches cost one all-reduce rather than $K$.

Communication volume per optimizer step drops from $K \cdot 2P\frac{R-1}{R}$ bytes per rank to $2P\frac{R-1}{R}$, where $P$ is the parameter payload in bytes and the $2\frac{R-1}{R}$ factor is the ring all-reduce constant. The catch is that overlap disappears for the synced pass: with $K-1$ silent passes, there is no computation left to hide the final exchange behind, so per-step wall time is not $1/K$ of the naive version. The gain is real but sublinear.

gradient_as_bucket_view=True makes param.grad a view into the communication bucket, removing one full copy of the gradients from memory. static_graph=True promises the autograd graph is identical every iteration; DDP then records the bucket-ready order on the first iteration and reuses it, which also makes find_unused_parameters unnecessary and permits activation checkpointing inside DDP. DDP.register_comm_hook replaces the collective entirely, which is how gradient compression (fp16 compress, PowerSGD) is implemented without touching the training loop.

References

  • Li et al. (2020), PyTorch Distributed: Experiences on Accelerating Data Parallel Training, VLDB 13(12) — §3.2 covers bucketing, the no_sync interface and the overlap measurements.
  • Vogels et al. (2019), PowerSGD: Practical Low-Rank Gradient Compression — the best-known DDP communication hook.
  • Ott et al. (2018), Scaling Neural Machine Translation — delayed updates as a communication-reduction technique, the idea no_sync implements.

What to learn next