PyTorch Distributed Training
When a model takes too long to train on one GPU, or doesn't fit in one GPU's memory at all, you distribute the work. PyTorch provides the building blocks. DistributedDataParallel (DDP) replicates the model on every GPU and splits the data. Fully Sharded Data Parallel (FSDP) shards parameters, gradients, and optimizer states across GPUs so models far larger than one device's memory can train. Tensor and pipeline parallelism split individual layers or layer stacks across devices for the largest models.
Distributed training is how large language models are pretrained and fine-tuned, and how vision and recommendation models scale to big datasets. The core concepts (processes, ranks, collective communication, and sharding) are the same whether you use raw PyTorch APIs or higher-level libraries like Hugging Face Accelerate, DeepSpeed, or Lightning.
TL;DR
- DDP: one process per GPU, each with a full model copy; gradients are all-reduced after
backward(). It's the default for models that fit on one GPU. - Launch with
torchrun, which setsRANK,LOCAL_RANK, andWORLD_SIZE; initialize the process group (NCCL on GPUs). - Use
DistributedSamplerso each rank sees a different data shard. - FSDP shards parameters, gradients, and optimizer state (ZeRO-3 style), which enables models larger than a single GPU.
- Tensor parallelism splits layers across GPUs; pipeline parallelism splits layer stacks. Combine them in 3D parallelism for the largest models.
- Save checkpoints from rank 0 (DDP) or with distributed checkpointing (FSDP), and scale the learning rate or batch size thoughtfully.
Quick Example
A DDP training script, launched with torchrun:
Core Concepts
Processes, Ranks, and Process Groups
Distributed PyTorch runs one process per GPU. Each has a global rank (0 to world size − 1) and a local rank (its GPU index on the node). torchrun launches processes, sets environment variables, and handles rendezvous and restarts (elastic training). init_process_group connects them using a backend: NCCL for NVIDIA GPUs, Gloo for CPU, and vendor backends for other accelerators.
Collective Communication
Processes coordinate through collectives: all-reduce (sum gradients across ranks), all-gather, reduce-scatter, and broadcast. Training speed at scale depends on interconnect bandwidth: NVLink within a node, and InfiniBand or RoCE between nodes. Communication overlapping with computation is what keeps GPUs busy.
DistributedDataParallel
- Every rank holds a full model replica and processes a different shard of each global batch.
- During
backward(), gradients are bucketed and all-reduced asynchronously while backpropagation continues, so every replica ends up with identical averaged gradients and identical weights. - The effective batch size is
per_gpu_batch × world_size. Larger batches often need a learning rate adjustment and warmup. - Use
no_sync()during gradient accumulation to skip unnecessary all-reduces.
DDP is simple and efficient, but each GPU must fit the whole model, its gradients, and its optimizer state.
FSDP: Sharded Data Parallelism
For large models, memory is dominated by parameters, gradients, and optimizer states (AdamW keeps two extra values per parameter). FSDP shards all three across ranks, following the ZeRO-3 approach:
- Before a layer's forward or backward pass, its full parameters are all-gathered.
- After use, the full parameters are freed.
- Gradients are reduce-scattered, so each rank keeps only its shard.
The result is that memory per GPU scales down with the number of GPUs, at the cost of extra communication. FSDP2 (fully_shard) in recent PyTorch offers a simpler per-parameter sharding API. Combine it with mixed precision, activation checkpointing, and CPU offload as needed.
Model Parallelism
- Tensor parallelism: split individual matrix multiplications (attention heads, MLP columns) across GPUs within a node. It requires fast interconnects.
- Pipeline parallelism: place consecutive layer groups on different GPUs and stream micro-batches through them.
- Sequence and context parallelism: split long sequences across devices.
- 3D parallelism combines data, tensor, and pipeline parallelism, as used for training the largest LLMs (Megatron-LM, DeepSpeed, TorchTitan).
Checkpointing
- DDP: save
model.module.state_dict()from rank 0 only, and load on all ranks. - FSDP: use
torch.distributed.checkpoint(DCP), where each rank writes its shard in parallel, and checkpoints can be resharded to a different world size. - Save optimizer and scheduler state and the data position for exact resumption. Long jobs will see failures, so checkpoint regularly and use elastic restarts.
Higher-Level Tools
Best Practices
Start With DDP, Move to FSDP When Memory Demands It
If the model and optimizer state fit on one GPU, DDP is simpler and usually faster. Switch to FSDP (or DeepSpeed ZeRO) when they don't, or when you need bigger batches than memory allows.
Get Single-GPU Training Right First
Verify correctness and throughput on one GPU, and ideally overfit a batch, before scaling out. Distributed bugs are much harder to debug.
Measure Scaling Efficiency
Track samples per second per GPU as you add GPUs. Poor scaling points to communication bottlenecks (slow interconnect, small batches), data loading limits, or stragglers.
Make Only Rank 0 Do Side Effects
Logging, saving, and evaluation reporting should happen once (rank 0), or be coordinated, to avoid duplicated files and corrupted checkpoints.
Common Mistakes
Forgetting DistributedSampler or set_epoch
Without DistributedSampler, every rank trains on the same data, wasting compute. Without set_epoch, the shuffle order repeats every epoch.
Rank-Dependent Code Paths That Skip Collectives
If one rank skips a forward pass or a collective (for example because of an if rank == 0 around model code), the others wait forever, and the job hangs. Keep model execution identical on all ranks.
Saving FSDP Models Like DDP
Calling state_dict() naively on sharded models saves only local shards, or gathers everything to one GPU and runs out of memory. Use the distributed checkpoint APIs or FSDP's state dict configuration.
FAQ
What's the difference between DataParallel and DistributedDataParallel?
nn.DataParallel uses one process with multiple threads, and it's slower and effectively deprecated. DistributedDataParallel uses one process per GPU with efficient all-reduce communication, and it works across machines. Always use DDP (or FSDP).
When do I need FSDP?
When the model's parameters, gradients, and optimizer states don't fit on a single GPU, which is common for models above a few billion parameters in full-precision training, or when you want larger batches. FSDP shards these across GPUs so memory per device decreases as you add devices.
How does the batch size change with more GPUs?
With data parallelism, each GPU processes its own batch, so the global batch size is per-GPU batch × number of GPUs. Larger global batches may need learning rate scaling and warmup to converge well, or you can keep the global batch constant by lowering per-GPU batches.
Do I need InfiniBand for multi-node training?
For large models and FSDP or tensor parallelism, fast interconnects matter a lot, and Ethernet-only clusters can be communication-bound. Small models with DDP can train acceptably over high-bandwidth Ethernet, especially with larger batches and gradient accumulation.
Related Topics
- PyTorch — The framework overview
- PyTorch Training Loop — The single-GPU foundation
- PyTorch Datasets & DataLoaders — DistributedSampler and input pipelines
- Large Language Models — Trained with these techniques
- Fine-Tuning — Distributed fine-tuning of pretrained models
- MLOps — Running training jobs reliably