Strategies for Distributed Training of LLMs: A Complete Implementation Guide
Training large language models across multiple GPUs requires coordinating data parallelism, ZeRO sharding, pipeline parallelism, and fault-tolerant checkpointing to overcome single-device memory and compute limitations.
The ai-engineering-from-scratch curriculum demonstrates how to implement these distributed training strategies from first principles in PyTorch. This guide examines the actual source code for building scalable LLM training systems that handle modern model sizes through gradient synchronization, optimizer state sharding, and atomic checkpointing.
Data-Parallel Training with DistributedDataParallel
The foundation of distributed training starts with data parallelism, where each GPU holds a complete model copy but processes different data shards. In phases/19-capstone-projects/77-data-parallel-ddp/code/main.py, the curriculum implements a minimal DistributedDataParallel wrapper that handles gradient synchronization without relying on high-level PyTorch DDP abstractions.
Broadcasting Parameters and Synchronizing Gradients
The wrapper initializes by broadcasting parameters from rank 0 to all workers, then synchronizes gradients after each backward pass:
class DistributedDataParallel:
def __init__(self, module: nn.Module, world_size: int):
self.module = module
self.world_size = world_size
self._broadcast_params()
def _broadcast_params(self) -> None:
for p in self.module.parameters():
dist.broadcast(p.data, src=0)
def sync_grads(self) -> None:
for p in self.module.parameters():
if p.grad is not None:
dist.all_reduce(p.grad.data, op=dist.ReduceOp.SUM)
p.grad.data.div_(self.world_size)
The training loop initializes a four-process group using the Gloo backend with file-based rendezvous, comparing per-rank losses against single-process baselines to verify correctness.
Memory-Efficient Training with ZeRO-1 Sharding
As models scale, optimizer states consume disproportionate memory. ZeRO-1 (Zero Redundancy Optimizer) shards optimizer states across ranks, reducing per-device memory to approximately one-third of standard Adam. The implementation in phases/19-capstone-projects/81-end-to-end-distributed-train/code/main.py creates a ZeroOptimizer class that partitions master weights, momentum, and velocity buffers.
Shard Distribution and All-Gather Operations
The optimizer pads parameter counts to ensure even divisibility by world_size, then assigns each rank a local shard:
class ZeroOptimizer:
def __init__(self, module, world_size, rank, lr=5e-3):
self.module = module
self.world_size = world_size
self.rank = rank
total = flat_param_numel(module)
pad = (-total) % world_size
self.chunk = (total + pad) // world_size
# create padded master shard and optimizer moments
...
During the step() method, the implementation uses dist.reduce_scatter to compute mean gradients for local shards, performs Adam updates on the shard, then employs dist.all_gather to reconstruct the full parameter vector for the next forward pass.
Pipeline Parallelism for Model Splitting
When models exceed single-GPU memory, pipeline parallelism divides the architecture into stages across devices. The curriculum demonstrates this in phases/19-capstone-projects/79-pipeline-parallel/code/main.py by splitting a transformer into two stages and streaming micro-batches through the pipeline. This approach overlaps forward and backward passes across ranks, maximizing utilization when the full model cannot reside on one device.
Production-Grade Checkpointing and Resume Verification
Production training runs require fault tolerance through sharded checkpointing that survives process crashes. The implementation provides atomic writes and cryptographic verification to ensure resume integrity.
Atomic Sharded Checkpoint Writes
The save_sharded function writes each rank's optimizer state to separate binary files with atomic rename semantics:
def save_sharded(per_rank_state, dest_dir, step):
dest = Path(dest_dir)
dest.mkdir(parents=True, exist_ok=True)
shards = []
for rank, state in enumerate(per_rank_state):
payload = _serialize(state)
sha = _sha(payload)
tmp_name = f"rank{rank}.bin.tmp"
final_name = f"rank{rank}.bin"
with open(dest / tmp_name, "wb") as f:
f.write(payload); f.flush(); os.fsync(f.fileno())
os.replace(tmp_name, final_name)
shards.append(ShardEntry(rank, final_name, sha))
manifest = {"world_size": len(per_rank_state), "step": step,
"shards": [asdict(s) for s in shards]}
manifest_tmp = dest / "manifest.json.tmp"
with open(manifest_tmp, "w") as f:
f.write(json.dumps(manifest, indent=2, sort_keys=True))
f.flush(); os.fsync(f.fileno())
os.replace(manifest_tmp, dest / "manifest.json")
return manifest
The JSON manifest records world_size, training step, and SHA-256 hashes for each shard, enabling corruption detection during resume operations.
Byte-Level Resume Verification
The verify_resume function loads sharded checkpoints using load_sharded, reconstructs the full optimizer state, and performs byte-for-byte comparison against in-memory snapshots taken at checkpoint time. This guarantees that distributed training can resume from exactly the previous state without silent data corruption.
End-to-End Integration
The capstone lesson in phases/19-capstone-projects/81-end-to-end-distributed-train/code/main.py combines all strategies into a cohesive training pipeline:
- Initialize a 4-rank Gloo process group
- Instantiate a
MiniGPTmodel with 112,640 parameters - Wrap with
ZeroOptimizerfor memory-efficient sharding - Execute 20 training steps with sharded checkpointing at step 10
- Verify resume integrity through byte-equality checks
Running python3 code/main.py produces verification output confirming that saved shards match in-memory snapshots exactly.
Summary
- Data parallelism replicates models across GPUs and synchronizes gradients via all-reduce operations, implemented in
DistributedDataParallel. - ZeRO-1 sharding partitions optimizer states across ranks using
ZeroOptimizer, reducing memory footprint throughreduce_scatterandall_gatheroperations. - Pipeline parallelism splits models into stages to handle oversized architectures that exceed single-device capacity.
- Sharded checkpointing provides fault tolerance through atomic file writes, JSON manifests with SHA-256 hashes, and byte-level resume verification.
- All implementations reside in the
ai-engineering-from-scratchrepository underphases/19-capstone-projects/, demonstrating production-grade distributed training from first principles.
Frequently Asked Questions
What is the difference between data parallelism and ZeRO-1 in distributed training of LLMs?
Data parallelism replicates the entire model on every GPU and only synchronizes gradients, while ZeRO-1 shards the optimizer states (master weights, momentum, and velocity) across ranks. ZeRO-1 reduces per-device memory usage to approximately one-third of standard data parallel training by distributing optimizer state via reduce_scatter and all_gather operations.
How does the curriculum verify that distributed training resumes correctly from checkpoints?
The verify_resume function in phases/19-capstone-projects/81-end-to-end-distributed-train/code/main.py loads sharded checkpoint files, reassembles the optimizer state, and compares master shard tensors byte-for-byte against snapshots captured at checkpoint time. This ensures exact state restoration without corruption.
Why use pipeline parallelism instead of larger batch sizes in data parallelism?
Pipeline parallelism becomes necessary when models exceed single-GPU memory limits. While data parallelism handles larger datasets by splitting batches, it cannot accommodate models that do not fit on one device. Pipeline parallelism splits the model architecture itself across stages, allowing simultaneous processing of micro-batches through different layers.
What backend does the curriculum use for distributed communication?
The implementations use the Gloo backend with file-based rendezvous points (init_method='file:///tmp/ddp_rendezvous'). This backend supports CPU and GPU operations across processes, enabling gradient synchronization through dist.all_reduce and dist.broadcast primitives without requiring specialized InfiniBand hardware.
Have a question about this repo?
These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →