# Strategies for Distributed Training of LLMs: A Complete Implementation Guide

> Master distributed LLM training with this guide. Learn data parallelism, ZeRO, pipeline parallelism, and fault-tolerant checkpointing for efficient large model development.

- Repository: [Rohit Ghumare/ai-engineering-from-scratch](https://github.com/rohitg00/ai-engineering-from-scratch)
- Tags: how-to-guide
- Published: 2026-07-19

---

**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`](https://github.com/rohitg00/ai-engineering-from-scratch/blob/main/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:

```python
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`](https://github.com/rohitg00/ai-engineering-from-scratch/blob/main/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:

```python
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`](https://github.com/rohitg00/ai-engineering-from-scratch/blob/main/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:

```python
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`](https://github.com/rohitg00/ai-engineering-from-scratch/blob/main/phases/19-capstone-projects/81-end-to-end-distributed-train/code/main.py) combines all strategies into a cohesive training pipeline:

1. Initialize a 4-rank Gloo process group
2. Instantiate a `MiniGPT` model with 112,640 parameters
3. Wrap with `ZeroOptimizer` for memory-efficient sharding
4. Execute 20 training steps with sharded checkpointing at step 10
5. 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 through `reduce_scatter` and `all_gather` operations.
- **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-scratch` repository under `phases/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`](https://github.com/rohitg00/ai-engineering-from-scratch/blob/main/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.