Multi-GPU and Distributed Training
Saving and logging from one rank only
Every rank runs the whole script, so a naive save writes the same file four times and a naive log reports one worker's slice as if it were the truth — rank 0 writes, everyone waits, and metrics are combined with all_reduce.
- 10 min read
- 3 reading levels
- Published
Read these first
On this page 5
One lesson, three depths. Pick the one that fits you today — you can switch any time.
Beginner — No maths. Plain English.
Every worker runs the whole script, so anything that writes a file or prints a line happens once per worker unless you stop it.
Picture four clerks given identical copies of the same form to process. Each finishes their share. Then all four walk to the filing cabinet and file a report in the same slot. Three of those reports were wasted effort, and if two arrive at once the file is a mess.
The fix is a rule everybody knows: clerk number one files the report; the others wait by the door. Then they all carry on together.
Why the log is a bigger trap than the file
The wasted file writing is visible. The logging problem is not.
Each worker only sees its own slice of the data. Worker 1's loss is the loss on worker 1's examples. Print it and you have published a number that describes a quarter of your validation set as though it described all of it.
With random slices the numbers are close, and you never notice. With sorted or grouped data they are far apart, and your dashboard has been lying for weeks.
How it works
saving:
rank 0 -> writes ckpt.pt
rank 1 -> waits at the barrier
rank 2 -> waits at the barrier
rank 3 -> waits at the barrier
...then all four continue together
logging:
rank 0 loss on its 3 samples: 0.59 \
rank 1 loss on its 3 samples: 2.12 }-- combine, then report: 1.35
(reporting 0.59 alone would be wrong)A real example you have seen
Exam results. A school does not announce "the class average" from one section's papers. It totals every section's marks and every section's headcount, then divides once. Combine, then divide — never average the averages of unequal groups.
Remember this
- Everything in the script runs once per worker, including
torch.saveandprint. - Let rank 0 write, and make the others wait at a barrier.
- Combine metrics with all_reduce before reporting, or your numbers describe one slice.
What to learn next
- all_reduce, all_gather and broadcast — the primitives used for both the metric and the barrier.
- What a state_dict really is — what you are actually writing to that file.
- Resuming training exactly where it stopped — the single-process rules that still apply here.
Developer — Code and libraries.
Setup
pip install torchCaptured on CPU with the gloo backend and two processes. No GPU needed.
The pattern, end to end
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))
opt = torch.optim.SGD(ddp.parameters(), lr=0.1)
torch.manual_seed(50 + rank) # each rank holds different data
x, y = torch.randn(6, 4), torch.randn(6, 1)
# --- one training step -------------------------------------------------
opt.zero_grad(set_to_none=True)
loss = nn.functional.mse_loss(ddp(x), y)
loss.backward()
opt.step()
# --- a metric that is honest across all ranks --------------------------
local = torch.tensor([loss.item() * len(x), float(len(x))])
dist.all_reduce(local, op=dist.ReduceOp.SUM) # sum losses AND sample counts
global_loss = (local[0] / local[1]).item()
print(f"rank {rank}: my loss {loss.item():.4f} | "
f"true loss over all {int(local[1])} samples {global_loss:.4f}")
# --- exactly one rank writes the file ----------------------------------
if rank == 0:
torch.save({"model": ddp.module.state_dict(), # .module strips the DDP wrapper
"opt": opt.state_dict(), "step": 1}, "ckpt.pt")
dist.barrier() # nobody reads before rank 0 wrote
ckpt = torch.load("ckpt.pt", map_location="cpu", weights_only=True)
print(f"rank {rank}: loaded keys {list(ckpt['model'])}")
dist.destroy_process_group()torchrun --nproc_per_node=2 ddp_save_log.pyrank 0: my loss 0.5934 | true loss over all 12 samples 1.3542 rank 0: loaded keys ['weight', 'bias'] rank 1: my loss 2.1151 | true loss over all 12 samples 1.3542 rank 1: loaded keys ['weight', 'bias']
Lines are grouped by rank for reading; the two processes interleave differently each run.
Look at the two per-rank losses: 0.5934 and 2.1151. Logging rank 0's number alone would have understated the loss by more than a factor of two. The combined figure, 1.3542, is the one that describes all 12 samples.
The four pieces that matter
ddp.module.state_dict() — .module is the original model inside the wrapper. Save without it and every key gains a module. prefix, which fails to load into an unwrapped model later. That specific failure has its own lesson.
dist.all_reduce on a [sum, count] pair — sum the weighted losses and the sample counts separately, then divide once. Averaging the per-rank averages is only correct when every rank holds the same number of samples, and the padded sampler means it often does not.
dist.barrier() — every rank blocks here until all have arrived. It guarantees rank 0's file is fully written before anyone opens it.
map_location="cpu" — loads onto the CPU first, then you move to the right device. Without it, a checkpoint saved from GPU 0 tries to restore straight onto GPU 0 on every rank, and that GPU runs out of memory.
Resuming a DDP run
import os
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()
torch.manual_seed(0)
model = nn.Linear(4, 1)
opt = torch.optim.SGD(model.parameters(), lr=0.1)
start_step = 0
if os.path.exists("ckpt.pt"):
# load BEFORE wrapping in DDP, on every rank, from CPU
ckpt = torch.load("ckpt.pt", map_location="cpu", weights_only=True)
model.load_state_dict(ckpt["model"])
opt.load_state_dict(ckpt["opt"])
start_step = ckpt["step"]
ddp = DDP(model) # broadcasts rank 0's weights, so all ranks agree anyway
print(f"rank {rank}: resuming from step {start_step}, "
f"weight checksum {ddp.module.weight.sum().item():.6f}")
dist.destroy_process_group()rank 0: resuming from step 1, weight checksum -0.381006 rank 1: resuming from step 1, weight checksum -0.381006
Two rules hide in that short file. Load before wrapping, so DDP's construction-time broadcast confirms every rank agrees. And load on every rank rather than rank 0 alone — that keeps optimizer state correct even though DDP would have fixed the weights for you.
Sharded checkpoints, briefly
One file written by rank 0 is right for DDP, where every rank holds the same complete model. It is wrong for FSDP, where each rank holds a different piece and gathering them onto one rank can exceed its memory.
For that case PyTorch ships torch.distributed.checkpoint, which writes one file per rank in parallel and can reload into a different world size. The API needs several GPUs to be worth demonstrating, so no output is shown here:
import torch.distributed.checkpoint as dcp
state = {"model": model.state_dict(), "optim": opt.state_dict()}
dcp.save(state, checkpoint_id="run42/step1000") # every rank writes its shard
dcp.load(state, checkpoint_id="run42/step1000") # reshards on loadCommon mistakes
Saving from every rank to the same path. Four processes writing one file concurrently can leave it truncated. The corruption shows up days later, at load time.
Forgetting the barrier after saving. Rank 1 reaches torch.load before rank 0 has finished writing and reads a partial file. It is a race, so it fails on maybe one run in twenty — the worst possible failure rate for finding a bug.
Barrier inside an if rank == 0 block. Rank 0 waits for ranks that will never arrive, and the job hangs forever. Every rank must reach every collective — see debugging distributed hangs.
Logging loss.item() from rank 0 as the run's loss. As the output above shows, that number can be off by a factor of two.
Early stopping computed per rank. If rank 0 decides to stop and rank 1 does not, half the job exits and the rest hangs at the next collective. Compute the decision from the all-reduced metric so every rank reaches the same conclusion.
Try it yourself
Delete dist.barrier() and run with --nproc_per_node=8 in a loop twenty times. Some runs will load a file rank 0 has not finished writing. That is what a race condition feels like.
What to learn next
- all_reduce, all_gather and broadcast — the primitives used for both the metric and the barrier.
- What a state_dict really is — what you are actually writing to that file.
- Resuming training exactly where it stopped — the single-process rules that still apply here.
Researcher — Mathematics and papers.
Why rank 0 is enough for DDP
DDP maintains the replica invariant: after every optimizer.step(), all ranks hold bitwise-identical parameters, because they applied identical averaged gradients to identical starting weights. Optimizer state is likewise a deterministic function of that gradient sequence, so Adam moments match across ranks too. Any single rank therefore holds a complete, sufficient snapshot. The invariant is what breaks under FSDP and tensor parallelism, where the parameter set is genuinely partitioned and a single-rank save is incomplete by construction.
Two things can silently break the invariant even under DDP. Buffers updated in the forward pass — BatchNorm running statistics being the common case — are computed per rank, and broadcast_buffers=True re-syncs them from rank 0 at the next forward, so rank 0's copy is authoritative but the others' work is discarded. And any parameter not touched by the loss on some rank receives no gradient there, which find_unused_parameters=True handles at the cost of a per-iteration graph traversal.
Metric reduction is a weighted mean
For per-rank losses $\ell_r$ over $n_r$ samples, the correct pooled value is
$$ \ell = \frac{\sum_{r} n_r \ell_r}{\sum_{r} n_r} $$
where $r$ indexes ranks. The unweighted mean $\frac{1}{R}\sum_r \ell_r$ coincides with it only when all $n_r$ are equal. Since DistributedSampler pads to equal length by duplicating up to $R-1$ samples, the counts are equal in training but the duplicates bias the estimate slightly; for evaluation, either drop the duplicates by tracking original indices, or evaluate on one rank. Metrics that are not means — AUC, F1, anything rank-based — cannot be reduced this way at all. They need all_gather of predictions, or a metric library that implements a distributed update/compute split, which is what torchmetrics provides.
Barriers and their cost
dist.barrier() is implemented as an all-reduce on a dummy tensor. Under NCCL it enqueues onto the current CUDA stream, so it synchronises stream order, not host order, and returns before the GPUs have actually met unless you also synchronise the device; device_ids should be passed to make the placement explicit. monitored_barrier() (gloo only) reports which rank failed to arrive, and is the single most useful debugging call in this section.
Filesystem visibility is a separate guarantee from the barrier. On a shared network filesystem, a barrier orders the processes but does not force metadata to propagate, so a torn read remains possible on NFS without proper close-to-open semantics. Writing to a temporary path and renaming atomically is the usual defence.
References
- Li et al. (2020), PyTorch Distributed, VLDB 13(12) — the replica invariant and buffer broadcast behaviour.
- PyTorch documentation, Distributed Checkpoint (DCP) — sharded save/load and resharding across world sizes.
- Rajbhandari et al. (2020), ZeRO: Memory Optimizations Toward Training Trillion Parameter Models — why a single-rank save stops being viable once state is partitioned.
What to learn next
- all_reduce, all_gather and broadcast — the primitives used for both the metric and the barrier.
- What a state_dict really is — what you are actually writing to that file.
- Resuming training exactly where it stopped — the single-process rules that still apply here.