Distributed Training (FSDP, DeepSpeed, Megatron-LM)
Distributed training frameworks split a model's parameters, gradients, and optimizer state across many GPUs so models too large for one GPU can still be trained.
Distributed training frameworks split a model's parameters, gradients, and optimizer state across many GPUsGPU (Graphics Processing Unit)GPUs, originally built for rendering graphics, turned out to be extremely well-suited to the parallel matrix multiplications deep learning requires. — sometimes thousands — for models too large to train on one machine. PyTorch FSDP (Fully Sharded Data Parallel) and DeepSpeed (Microsoft) shard model state across GPUs; Megatron-LM (NVIDIA) is built specifically for training very large transformersTransformerThe transformer is the neural network architecture built around self-attention, introduced in 2017, underlying essentially all modern LLMs. efficiently at scale. Most developers fine-tuning or experimenting won't touch these directly — they matter once a project outgrows a single machine.
How it works
Three kinds of parallelism get combined, often all at once:
- Data parallel — every device holds a full copy of the model and
processes a different slice of the batch. Gradients are averaged with
an
all-reducecollective before the optimizer step, so all replicas stay identical. - Sharded data parallel — FSDP and DeepSpeed ZeRO keep only a shard
of parameters, gradients, and optimizer state on each device,
all-gathering a layer's weights just before its forward pass and freeing them afterward. This trades extra communication for a much lower per-device memory footprint. - Model parallel — tensor parallelismTensor ParallelismTensor parallelism splits individual weight matrices across multiple GPUs so a model too large for one GPU's memory can still be trained or served. splits individual matrices across devices, while pipeline parallelism assigns different layers to different stages.
Underneath, collectives run over NCCL on CUDACUDACUDA is NVIDIA's programming platform that lets frameworks like PyTorch dispatch tensor computation to NVIDIA GPUs.
devices, using NVLink within a node and the network fabric between
nodes.
When it breaks
- The slowest rank sets the pace. Collectives are synchronous, so one straggler — a thermally throttled device, an uneven data shard — stalls every other rank at the barrier.
- Communication swamps compute. Sharding across nodes on ordinary Ethernet can spend more time moving weights than doing math. Keeping tensor parallelism inside a single node is the usual fix.
- Hangs instead of errors. If ranks disagree on how many collectives to run — a conditional branch, a rank-dependent early exit — the job blocks until a timeout rather than raising.
- Checkpoints are layout-bound. A sharded checkpoint often will not load under a different world size or sharding strategy without an explicit consolidation step, which bites when resuming or fine-tuningFine-tuningFine-tuning continues training a pretrained model on a smaller, curated dataset to teach it a specific behavior, such as following instructions. on a smaller cluster.
See also: GPUGPU (Graphics Processing Unit)GPUs, originally built for rendering graphics, turned out to be extremely well-suited to the parallel matrix multiplications deep learning requires., PyTorchPyTorchPyTorch is the dominant deep learning framework in both research and production, providing tensor computation, GPU dispatch, and automatic differentiation.
Learn more: Tooling & The Dev Stack