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

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

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:

  1. Before a layer's forward or backward pass, its full parameters are all-gathered.
  2. After use, the full parameters are freed.
  3. 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

Checkpointing

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

References