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.

On this page 5
  1. Why it exists
  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.

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 copies

Every 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

Developer — Code and libraries.

Setup

bash
pip install torch

Everything 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

ddp_train.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")          # 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()
bash
torchrun --nproc_per_node=2 ddp_train.py
Output
step 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

ddp_broadcast.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 = 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()
Output
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:

ddp_gpu.py — needs at least one CUDA GPU
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

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