Multi-GPU and Distributed Training
Your first DDP run with torchrun
torchrun starts one copy of your script per GPU and hands each copy a number; your job is to read that number, wrap the model in DDP, and let the gradients agree.
- 9 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.
You do not write a multi-GPU program; you write an ordinary program and let a launcher start several copies of it.
Think of a call centre. The manager does not write four different scripts. He prints one script, hires four agents, and gives each a headset numbered 1 to 4. Every agent reads the same page. The number on the headset is the only thing that differs.
torchrun is that manager. It is a command that starts your training file several times over, once per GPU, and stamps each copy with its own number.
Why it exists
Starting several processes by hand is painful. You would have to pick a port, pass every process its number, wait for all of them, and kill the survivors when one dies.
That work is identical for every project on earth. So PyTorch ships the manager as a command. You run torchrun, and it does the hiring, the numbering and the cleanup.
How it works
you type: torchrun --nproc_per_node=2 train.py
torchrun starts:
copy A RANK=0 ---\
}--- both run train.py, top to bottom
copy B RANK=1 ---/
inside each copy:
init_process_group() -> "hello, who else is here?"
DDP(model) -> copy A's weights are sent to everyone
loss.backward() -> gradients are averaged between the copiesEvery copy finishes the step holding the same weights. That is the whole trick.
A real example you have seen
A group of students revising from photocopies of one set of notes. Each reads a different chapter, then they meet and agree on the summary. Nobody keeps a private version of the answer.
Remember this
- You write one ordinary script; torchrun runs many copies of it.
- Each copy learns its own rank — the worker number it was given.
- After every backward pass, all copies hold the same weights again.
What to learn next
- DistributedSampler and set_epoch — giving each rank its own slice of the data.
- When gradients are synchronised, and no_sync — what
backward()sends, and when to stop it. - Random seeds and reproducibility — the seeding rules that still hold across ranks.
Developer — Code and libraries.
Setup
pip install torchEverything below runs on a plain CPU laptop with the gloo backend — the communication library PyTorch uses when there are no NVIDIA GPUs. The outputs shown were captured that way, with two processes. Nothing in this lesson needs a GPU.
The script
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") # nccl on GPUs, gloo on CPU
rank = dist.get_rank()
world = dist.get_world_size()
torch.manual_seed(0) # same starting weights on every rank
model = nn.Sequential(nn.Linear(10, 32), nn.ReLU(), nn.Linear(32, 1))
ddp_model = DDP(model) # add device_ids=[local_rank] on GPUs
opt = torch.optim.SGD(ddp_model.parameters(), lr=0.1)
torch.manual_seed(100 + rank) # each rank gets DIFFERENT data
x, y = torch.randn(16, 10), torch.randn(16, 1)
for step in range(3):
opt.zero_grad(set_to_none=True)
loss = nn.functional.mse_loss(ddp_model(x), y)
loss.backward() # gradients are averaged here
opt.step()
if rank == 0: # print from one rank only
print(f"step {step} loss on rank 0: {loss.item():.4f}")
w = ddp_model.module[0].weight
print(f"rank {rank}: first-layer weight checksum {w.sum().item():.6f}")
dist.destroy_process_group()torchrun --nproc_per_node=2 ddp_train.pystep 0 loss on rank 0: 1.0865 step 1 loss on rank 0: 0.9100 step 2 loss on rank 0: 0.8038 rank 0: first-layer weight checksum -4.635065 rank 1: first-layer weight checksum -4.635065
Two processes print at the same time, so line order shifts between runs. The two checksums are what matter: identical to six decimal places, even though the two ranks trained on completely different data. That is DDP working.
The five lines that make it distributed
dist.init_process_group(backend="gloo") — joins the team. It reads MASTER_ADDR, MASTER_PORT, RANK and WORLD_SIZE from the environment, which torchrun has already set for you. Use "nccl" on NVIDIA GPUs; it is far faster there.
dist.get_rank() — your worker number, from 0 to world - 1. Rank 0 is the one that prints, saves and logs.
dist.get_world_size() — how many workers exist in total, across every machine.
DDP(model) — the wrapper. At construction it broadcasts rank 0's parameters to everyone, then registers hooks that average gradients during each backward pass.
dist.destroy_process_group() — hangs up. Without it you can leave sockets and half-dead processes behind after a crash.
Proof that DDP synchronises at construction
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(rank) # deliberately DIFFERENT weights per rank
model = nn.Linear(4, 1)
print(f"rank {rank}: before DDP, weight[0][0] = {model.weight[0][0].item():+.4f}")
ddp = DDP(model) # rank 0's weights are broadcast to everyone
print(f"rank {rank}: after DDP, weight[0][0] = {ddp.module.weight[0][0].item():+.4f}")
dist.destroy_process_group()rank 0: before DDP, weight[0][0] = -0.0037 rank 0: after DDP, weight[0][0] = -0.0037 rank 1: before DDP, weight[0][0] = +0.2576 rank 1: after DDP, weight[0][0] = -0.0037
Rank 1 started at +0.2576 and ended holding rank 0's -0.0037. You do not have to seed every rank identically, because DDP overwrites them anyway. Seeding is still worth doing, since it keeps dropout and augmentation reproducible.
What changes on real GPUs
Four lines differ. This block needs NVIDIA hardware, so no output is shown for it:
import os
import torch
import torch.distributed as dist
from torch.nn.parallel import DistributedDataParallel as DDP
local_rank = int(os.environ["LOCAL_RANK"]) # my GPU index on THIS machine
torch.cuda.set_device(local_rank) # do this BEFORE creating any tensor
dist.init_process_group(backend="nccl") # the NVIDIA collective library
model = MyModel().to(local_rank)
ddp_model = DDP(model, device_ids=[local_rank], output_device=local_rank)LOCAL_RANK is the GPU index on the current machine; RANK is the global worker number. On one machine with four GPUs the two are equal. On two machines they are not, and mixing them up is the classic first-week bug.
Common mistakes
Creating tensors before torch.cuda.set_device(local_rank). They all land on GPU 0, which runs out of memory while the other cards sit idle. Set the device first, always.
Printing from every rank. Your log becomes four interleaved copies of itself. Guard prints with if rank == 0, and read saving and logging from one rank.
Giving every rank the same data. The script above does that on purpose, to prove the sync. Real training wants a different slice per rank — that is DistributedSampler.
Running the file with python ddp_train.py. Without torchrun there is no RANK in the environment, and init_process_group raises a ValueError about the env:// rendezvous. Use the launcher.
Try it yourself
Run it with --nproc_per_node=4. The checksums should still agree across all four ranks. Then delete the DDP(...) wrapper, keep everything else, and watch the checksums drift apart.
What to learn next
- DistributedSampler and set_epoch — giving each rank its own slice of the data.
- When gradients are synchronised, and no_sync — what
backward()sends, and when to stop it. - Random seeds and reproducibility — the seeding rules that still hold across ranks.
Researcher — Mathematics and papers.
What torchrun actually is
torchrun is the console entry point for torch.distributed.elastic, the replacement for torch.distributed.launch. It runs one agent process per node. The agent performs a rendezvous through a store (the c10d backend by default, a TCPStore hosted at the rendezvous endpoint), assigns global ranks, spawns nproc_per_node workers, and monitors them. If any worker exits non-zero, the agent tears down the rest — which is why a distributed job fails fast instead of leaving orphans behind.
The agent injects these variables into each worker: RANK, LOCAL_RANK, GROUP_RANK, WORLD_SIZE, LOCAL_WORLD_SIZE, MASTER_ADDR, MASTER_PORT, TORCHELASTIC_RESTART_COUNT and TORCHELASTIC_RUN_ID. Calling init_process_group() with no init_method reads that env:// set, which is why the call takes no arguments in practice.
With --max-restarts=N and an elastic range such as --nnodes=1:4, the agent restarts the whole worker group from the last rendezvous after a failure. Elasticity is group-level, not worker-level: there is no partial recovery, so the script must be able to resume from a checkpoint.
The process group and the Reducer
init_process_group builds a ProcessGroup; the default one is dist.group.WORLD. Backends are nccl (NVIDIA GPUs, ring and tree collectives over NVLink, PCIe or InfiniBand), gloo (CPU, plus a few GPU ops) and mpi (only when PyTorch was built against an MPI implementation). One job may hold several groups via dist.new_group(ranks=[...]), which is how FSDP and hybrid parallelism carve the world into meshes.
The DDP constructor does two things worth naming. It broadcasts state_dict() from src=0 — parameters and buffers. Then it builds the Reducer: parameters are assigned to buckets in approximately reverse registration order, defaulting to bucket_cap_mb=25, with one autograd hook per parameter marking its bucket ready. When a bucket fills, its all-reduce is enqueued at once, overlapping with the remainder of the backward pass.
broadcast_buffers=True, the default, re-broadcasts buffers from rank 0 at the start of every forward. For BatchNorm that keeps running statistics consistent across replicas, and it is also why per-rank buffers silently disappear unless you set the flag to False or convert to SyncBatchNorm.
Determinism
Gradient all-reduce is floating-point addition in a communication-dependent order, so DDP is not bitwise deterministic across different world sizes or interconnect topologies. Runs at a fixed world size on fixed hardware do reproduce, provided seeds are fixed and torch.use_deterministic_algorithms(True) is set.
References
- Li et al. (2020), PyTorch Distributed: Experiences on Accelerating Data Parallel Training, VLDB 13(12) — bucketing, overlap and scaling measurements.
- PyTorch documentation, Torch Distributed Elastic — the agent, rendezvous and restart semantics.
- Goyal et al. (2017), Accurate, Large Minibatch SGD: Training ImageNet in 1 Hour — the work that made multi-worker data parallelism routine.
What to learn next
- DistributedSampler and set_epoch — giving each rank its own slice of the data.
- When gradients are synchronised, and no_sync — what
backward()sends, and when to stop it. - Random seeds and reproducibility — the seeding rules that still hold across ranks.