Multi-GPU and Distributed Training
DistributedSampler and set_epoch
DistributedSampler hands each rank a different, non-overlapping slice of the dataset — and if you forget set_epoch, every epoch after the first is an exact rerun of the first.
- 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.
DistributedSampler is the rule that decides which samples each worker reads, so no two workers read the same one.
Picture four people asked to count a sack of a thousand coins. If nobody divides the sack, all four count all thousand coins. Four times the work, no extra speed. Instead you deal the coins into four piles first. Each person counts one pile, and you add the four answers.
DistributedSampler deals the piles. It is a small object that tells each worker which row numbers belong to it.
Without it, every worker trains on the whole dataset. Your loss curve looks normal, your run finishes, and you gained nothing — the exact failure that is hardest to notice.
Why set_epoch exists
The dealing is shuffled, so a worker does not always get the same block. But the shuffle uses a fixed recipe, and a recipe with the same ingredients gives the same result every time.
So on epoch two, every worker gets exactly the same pile it got on epoch one. And on epoch three. The model sees the same order forever, which is a slow, quiet way to make it memorise.
set_epoch(epoch) changes one ingredient in the recipe — the epoch number. Now each pass through the data is reshuffled.
How it works
dataset: 0 1 2 3 4 5 6 7 8 9 (ten samples, two workers)
epoch 0 shuffle -> 4 7 3 0 6 | 1 5 9 8 2
rank 0 gets ^^^^^^^^^ rank 1 gets ^^^^^^^^^
epoch 1 set_epoch(1) reshuffles
rank 0: 5 1 0 9 7 rank 1: 6 2 8 3 4
epoch 1 WITHOUT set_epoch
rank 0: 4 7 3 0 6 rank 1: 1 5 9 8 2 <- identical to epoch 0A real example you have seen
Dealing cards. You share out the deck so nobody holds the same card, and you shuffle between games. Skipping the shuffle means everyone gets the same hand every game.
Remember this
- Without the sampler, every worker trains on the whole dataset — wasted work.
- The sampler gives each worker a slice with no overlap.
- Call
set_epoch(epoch)at the top of every epoch, or the shuffle never changes.
What to learn next
- When gradients are synchronised, and no_sync — what happens after each rank has its slice.
- Scaling the learning rate with the number of GPUs — why the effective batch has grown.
- Samplers and weighted sampling — the single-process sampler this one extends.
Developer — Code and libraries.
Setup
pip install torchDistributedSampler accepts num_replicas and rank directly, so you can inspect its behaviour in one ordinary process. No launcher, no GPU. Every output below came from a plain python file.py run.
Seeing the slices, and what set_epoch changes
from torch.utils.data import DistributedSampler
data = list(range(10)) # ten samples, two workers
def slices(epoch, call_set_epoch):
out = []
for rank in range(2):
s = DistributedSampler(data, num_replicas=2, rank=rank, shuffle=True, seed=0)
if call_set_epoch:
s.set_epoch(epoch)
out.append(list(s))
return out
for epoch in range(3):
a, b = slices(epoch, call_set_epoch=True)
print(f"epoch {epoch} WITH set_epoch rank0={a} rank1={b} overlap={set(a) & set(b)}")
print()
for epoch in range(3):
a, b = slices(epoch, call_set_epoch=False)
print(f"epoch {epoch} WITHOUT set_epoch rank0={a} rank1={b}")epoch 0 WITH set_epoch rank0=[4, 7, 3, 0, 6] rank1=[1, 5, 9, 8, 2] overlap=set() epoch 1 WITH set_epoch rank0=[5, 1, 0, 9, 7] rank1=[6, 2, 8, 3, 4] overlap=set() epoch 2 WITH set_epoch rank0=[8, 1, 6, 0, 2] rank1=[7, 5, 9, 4, 3] overlap=set() epoch 0 WITHOUT set_epoch rank0=[4, 7, 3, 0, 6] rank1=[1, 5, 9, 8, 2] epoch 1 WITHOUT set_epoch rank0=[4, 7, 3, 0, 6] rank1=[1, 5, 9, 8, 2] epoch 2 WITHOUT set_epoch rank0=[4, 7, 3, 0, 6] rank1=[1, 5, 9, 8, 2]
Two facts, both visible. The overlap column is empty every time — the slices genuinely partition the data. And the bottom half repeats itself forever, which is the bug this lesson exists to prevent.
Notice also that the two ranks agree on the global shuffle without talking to each other. Both compute the same permutation from seed + epoch, then take different pieces of it. No communication is involved.
The uneven case: padding and drop_last
from torch.utils.data import DistributedSampler
data = list(range(7)) # seven samples do not split evenly by two
for drop_last in (False, True):
parts = [list(DistributedSampler(data, num_replicas=2, rank=r,
shuffle=False, drop_last=drop_last))
for r in range(2)]
seen = parts[0] + parts[1]
print(f"drop_last={drop_last!s:5} rank0={parts[0]} rank1={parts[1]}"
f" total={len(seen)} missing={sorted(set(data) - set(seen))}"
f" repeated={sorted({i for i in seen if seen.count(i) > 1})}")drop_last=False rank0=[0, 2, 4, 6] rank1=[1, 3, 5, 0] total=8 missing=[] repeated=[0] drop_last=True rank0=[0, 2, 4] rank1=[1, 3, 5] total=6 missing=[6] repeated=[]
Every rank must run the same number of steps, or the ranks that finish early hang waiting at the next collective. So the sampler forces an equal count one of two ways. With drop_last=False it repeats samples until the count divides evenly — sample 0 appears twice. With drop_last=True it drops the remainder — sample 6 is never seen.
This is the single most dangerous line in distributed evaluation. Padding a validation set means some samples are counted twice, so your reported accuracy is quietly wrong.
Wiring it into a DataLoader
import torch
import torch.distributed as dist
from torch.utils.data import DataLoader, DistributedSampler, TensorDataset
dist.init_process_group(backend="gloo")
rank, world = dist.get_rank(), dist.get_world_size()
dataset = TensorDataset(torch.arange(12).float().unsqueeze(1))
sampler = DistributedSampler(dataset) # reads rank and world size itself
loader = DataLoader(dataset, batch_size=2, sampler=sampler) # NOTE: no shuffle=True
for epoch in range(2):
sampler.set_epoch(epoch) # the one line people forget
seen = [int(b[0].item()) for batch in loader for b in batch[0]]
print(f"rank {rank} epoch {epoch}: {seen}")
dist.destroy_process_group()torchrun --nproc_per_node=2 ddp_loader.pyrank 0 epoch 0: [8, 5, 7, 2, 3, 11] rank 1 epoch 0: [0, 9, 4, 10, 1, 6] rank 0 epoch 1: [1, 6, 11, 5, 2, 10] rank 1 epoch 1: [4, 8, 0, 3, 9, 7]
Six samples each out of twelve, no overlap within an epoch, and a fresh split on epoch 1. Captured on CPU with the gloo backend and two processes; line order varies between runs.
Common mistakes
Passing shuffle=True to the DataLoader as well. PyTorch raises ValueError: sampler option is mutually exclusive with shuffle. Shuffling is the sampler's job now — leave the DataLoader alone.
Calling set_epoch once, outside the loop. Same effect as never calling it. It belongs on the first line of each epoch.
Using the padded sampler for validation. Either compute per-rank sums and counts and combine them with all_reduce — see saving and logging from one rank — or evaluate on rank 0 alone while the others wait at a barrier().
Assuming len(loader) is the dataset size. It is this rank's share. A rank-0 progress bar showing half the samples is correct, not broken.
Streaming datasets. DistributedSampler needs __len__ and indexing, so it cannot work with an IterableDataset. Those must split the stream themselves, usually by shard file or by rank-modulo filtering inside the worker.
Try it yourself
Change data in the first script to 11 items with drop_last=False and count how many samples get repeated. Then predict the answer for 13 items across 4 ranks before you run it.
What to learn next
- When gradients are synchronised, and no_sync — what happens after each rank has its slice.
- Scaling the learning rate with the number of GPUs — why the effective batch has grown.
- Samplers and weighted sampling — the single-process sampler this one extends.
Researcher — Mathematics and papers.
The construction
For a dataset of $N$ items and $R$ replicas, the sampler computes a per-rank length
$$ n = \begin{cases} \lceil N / R \rceil & \text{if } \texttt{drop_last=False} \ \lfloor N / R \rfloor & \text{if } \texttt{drop_last=True} \end{cases} $$
where $N$ is the dataset size and $R$ is num_replicas. The total index list is then forced to length $nR$: padded by cycling the head of the list, or truncated. Rank $r$ takes the strided slice indices[r : nR : R], so consecutive shuffled indices land on different ranks — an interleave, not a contiguous block. Under padding, at most $R - 1$ samples are duplicated per epoch, an $O(R/N)$ perturbation of the empirical distribution, negligible in training and not negligible in evaluation.
The permutation is generated by a fresh torch.Generator seeded with seed + epoch, so it is a pure function of those two integers. Every rank derives the identical global permutation with zero communication, which is what keeps the sampler free. It also means resuming mid-epoch cannot be done by seeding alone: you must record how many batches were consumed, or accept restarting the epoch.
Interaction with gradient statistics
Non-overlapping slices make the union of the $R$ per-rank minibatches a sample without replacement of size $Rb$, where $b$ is the per-rank batch size. The averaged DDP gradient is therefore an unbiased estimate of the gradient over that effective batch of $Rb$, with variance falling as $1/(Rb)$ — the fact underneath learning-rate scaling. Duplicated padding samples introduce a bias of order $(R-1)/N$ into that estimate, which vanishes for realistic $N$.
Alternatives at scale
For datasets too large to enumerate, the index-permutation model breaks down and sharding moves into the storage layer: WebDataset-style .tar shards assigned by rank and then by DataLoader worker id, or StatefulDataLoader from torchdata for exact mid-epoch resumption. Both replace "permute $N$ indices" with "permute shards, then permute within a shard buffer", trading perfect shuffling for constant memory. The big-dataset lesson covers the single-process form of the same trade.
References
- PyTorch documentation,
torch.utils.data.DistributedSampler— padding,drop_lastand theset_epochcontract. - Goyal et al. (2017), Accurate, Large Minibatch SGD — per-worker sharding and why shuffling must vary across epochs.
- Aizman et al. (2019), High Performance I/O For Large Scale Deep Learning — the shard-based alternative to index sampling.
What to learn next
- When gradients are synchronised, and no_sync — what happens after each rank has its slice.
- Scaling the learning rate with the number of GPUs — why the effective batch has grown.
- Samplers and weighted sampling — the single-process sampler this one extends.