Search Tech Journey

Find topics, journeys and posts

back to blog
mladvanced 120m read

DL S051 · Distributed Training Basics — DDP and FSDP

Explain when you need each and what they cost. Part of the 'Deep Learning & LLMs From Scratch' 80-session self-study series.

🧠SoftwareM09 · Pretraining your own foundation model· Session 051 of 130 120 min

🎯 Understand DDP, FSDP, and ZeRO stages well enough to launch multi-GPU pretraining without wasting a weekend of GPU-hours.

Series: Deep Learning & LLMs From Scratch — 80 sessions · Session 51 / 80 · Module M09 · ~2 hours

The story we're starting with

Your 100M-param Llama-clone fits fine on one A100 40GB — model, optimizer states, activations, all in ~15 GB with room to spare. But scale up. A 1B-param model needs ~20 GB just for parameters + optimizer state in mixed precision. A 7B model needs ~140 GB — three A100 80GBs. A 70B model needs ~1.5 TB — twenty A100 80GBs. Suddenly "just fit it in one GPU" is not an option.

There are two big ideas for multi-GPU training. Data parallelism — replicate the model on every GPU and split the batch across them. Model parallelism — shard the model itself across GPUs. Modern LLM pretraining is a mix: DDP (data parallel, simple) works up to ~1B params. FSDP / ZeRO (sharded parameters + optimizer state) is what you need past that. Beyond ~70B you also need tensor parallelism and pipeline parallelism.

For your 100M M09 run you'll use one GPU and no parallelism — but you should understand DDP and FSDP because (a) the next model you scale up will need them, and (b) reading loss curves and throughput numbers from other people's runs is impossible without this vocabulary.

You will be able to
  • Explain the difference between DDP and FSDP in one paragraph each.
  • Write down the ZeRO stages 1/2/3 and what each one shards.
  • Estimate the memory footprint of a 7B model under DDP vs ZeRO-3.
  • Explain the comms overhead of FSDP and why you want fast interconnect (NVLink, InfiniBand).
  • Convert a single-GPU training script to `torchrun`-launched DDP in 4 lines.

Prerequisites

  • S018 (GPU + mixed precision) — you understand .to("cuda") and bf16.
  • S050 (architecture choices) — for parameter-count intuition.


1 · Memory accounting — where does the VRAM go?

Before any parallelism talk, you need to understand what's living in GPU memory during training. There are four things:

The four memory consumers

    Rule of thumb for a bf16 Adam-trained transformer: ~16 bytes per parameter of steady-state memory (ignoring activations).

    Concrete:

    ModelParamsSteady-state (16 B/param)Fits in 40GB A100?
    100M1e81.6 GBEasily
    1B1e916 GBYes, with activations tight
    7B7e9112 GBNo — need 2-3 GPUs
    70B7e101120 GBNo — need 15-20 GPUs

    Activations on top of the steady state can add another 10–100 GB depending on batch size and seq length. For a 7B model at 4k context, activations easily push you to 200 GB total memory needed. And you haven't fit a single training example yet.

    This is why distributed training exists.


    2 · DDP (DistributedDataParallel) — the easy one

    Mental model. Copy the whole model to every GPU. Split each batch into world_size micro-batches, one per GPU. Each GPU runs forward + backward on its slice. At the end of backward, all_reduce the gradients so every GPU sees the mean gradient across all GPUs. Optimizer step is identical on all GPUs (same params, same gradients, same result).

    The analogy
    🌍 Real world
    💻 Code world

    2.1 Memory: no savings

    Every GPU holds the entire 16 B/param footprint. DDP does not reduce per-GPU memory. It only lets you process a bigger effective batch by parallelizing.

    2.2 Compute: near-linear scaling

    If comms are fast (NVLink for intra-node, InfiniBand for inter-node), throughput scales ~90–95% of ideal up to 8 GPUs. The all_reduce cost is O(N) bytes per step, amortized across the compute of a step. For small models it can become the bottleneck; for real LLMs it's usually <15% of step time.

    2.3 Four-line switch

    Single-GPU training loop → DDP is genuinely 4 lines:

    import torch.distributed as dist
    from torch.nn.parallel import DistributedDataParallel as DDP
     
    dist.init_process_group("nccl")                       # 1
    local_rank = int(os.environ["LOCAL_RANK"])
    torch.cuda.set_device(local_rank)                     # 2
    model = model.to(local_rank)
    model = DDP(model, device_ids=[local_rank])           # 3
     
    # In your DataLoader, use DistributedSampler
    sampler = DistributedSampler(dataset)                 # 4
    loader = DataLoader(dataset, sampler=sampler, ...)

    Launch:

    torchrun --nproc_per_node=8 train.py

    That's it. DDP is the "just make it multi-GPU" button. Use it for anything ≤1B params on a single node.

    Try itConvert a single-GPU training loop to DDP and confirm scaling

    Start from your S040 nanoGPT training loop. Add the four lines from §2.3 and wrap your dataset in DistributedSampler. Time one epoch on 1 GPU, then 2, then 4. You should see roughly 1.85× speedup at 2 GPUs and ~3.6× at 4 (comm overhead eats the last ~10%). If you see <1.5× at 2 GPUs, either your batch is too small (comm dominates) or your model is tiny — both symptoms mean DDP overhead is out of proportion.

    💡 Hint · Use `torchrun --nproc_per_node=2` first — you can rent a 2×A10 box for ~$1/hr.

    3 · ZeRO — the memory-saving family

    DDP's limitation: replicated everything. ZeRO (Rajbhandari 2019, DeepSpeed team) asks: what if we shard the redundant stuff?

    Three stages, in escalating aggressiveness:

    3.1 ZeRO-1: shard optimizer states

    Each GPU owns 1/N-th of the optimizer state (Adam's m, v, master weights). Forward/backward still work on the full replicated params. During optimizer step, each GPU updates its slice; then all-gather the updated weights so every GPU has the new full model.

    Savings: 12 B/param optimizer state → 12/N B/param. On 8 GPUs, a 7B model drops from 84 GB optimizer state to 10.5 GB per GPU. Total per-GPU memory: 16 → ~5 B/param.

    3.2 ZeRO-2: also shard gradients

    Same idea, applied to gradients. Each GPU only stores its 1/N slice of gradients after backward.

    Savings: another 2/N B/param.

    3.3 ZeRO-3: also shard parameters (= FSDP)

    Now even the parameters themselves are sharded. During forward/backward, each GPU temporarily all-gathers the params it needs just for the current layer, then releases them.

    Savings: full 16/N B/param footprint. On 8 GPUs, a 7B model uses 14 GB per GPU — fits on one A100.

    PyTorch's FSDP = ZeRO stage 3. Same idea, different codebase.

    ZeRO stage cheat sheet

      4 · FSDP — how it actually works

      FSDP wraps each block (typically a transformer layer). During forward:

      1. All-gather this layer's params from every rank → each rank has the full layer.
      2. Run forward on the layer.
      3. Free the gathered params (keep only the local shard).
      4. Move to next layer, repeat.

      During backward, mirror it: all-gather params, compute gradients, reduce-scatter gradients (each rank ends up with its shard of the gradient), free the gathered params, move on.

      The math per step:

      • Forward all-gathers: N params × 2 bytes = 2N bytes moved across the network, per step.
      • Backward all-gathers + reduce-scatters: ~2 × 2N = 4N bytes.
      • Total comms: ~6N bytes per training step.

      Compare DDP: only 2N bytes (one all-reduce of gradients). FSDP has 3× the comms.

      Why anyone accepts 3× the comms: because otherwise the model doesn't fit at all. FSDP unlocks training models bigger than any single GPU's VRAM.

      4.1 PyTorch FSDP boilerplate

      from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
      from torch.distributed.fsdp.wrap import transformer_auto_wrap_policy
      from functools import partial
       
      wrap_policy = partial(transformer_auto_wrap_policy,
                            transformer_layer_cls={LlamaLikeBlock})
       
      model = FSDP(
          model,
          auto_wrap_policy=wrap_policy,
          device_id=local_rank,
          mixed_precision=torch.distributed.fsdp.MixedPrecision(
              param_dtype=torch.bfloat16,
              reduce_dtype=torch.float32,
          ),
      )

      Same launch: torchrun --nproc_per_node=8 train.py.


      FSDP's ~6N bytes/step comms budget has to complete inside the compute time of a step, or your GPUs stall waiting for data.

      Concrete example: a 7B model, 8 GPUs.

      • Comms per step ≈ 6 × 7e9 × 2 bytes = 84 GB moved.
      • PCIe 4.0: 64 GB/s bidirectional. Comms alone: 1.3 seconds.
      • NVLink 3.0 (A100): 600 GB/s. Comms alone: 0.14 seconds.
      • InfiniBand HDR: 200 Gbps = 25 GB/s. Comms alone: 3.4 seconds.

      Step compute time for a 7B model on 8×A100 at MFU 40%: ~0.5 seconds.

      • On NVLink: comms 0.14 s vs compute 0.5 s → mostly overlapped, ~10% overhead.
      • On PCIe only: comms 1.3 s vs compute 0.5 s → stall of 0.8 s per step. 60% throughput lost.
      • On IB inter-node: 3.4 s stall → 85% throughput lost.

      This is why serious training runs use NVLink for intra-node and InfiniBand (or fancier RDMA) for inter-node. Consumer GPUs with only PCIe are a bad fit for FSDP.

      The analogy
      🌍 Real world
      💻 Code world

      6 · When to use what — decision tree

      For our M09 100M run: one GPU, no parallelism. Total training time on one A100 = ~3 hours for 2B tokens. Adding a second GPU literally isn't worth it.

      If you scale to a 1B model at Chinchilla-optimal (20B tokens): ~150 A100-hours on one GPU. That's when DDP on 2-4 GPUs starts making sense to cut wall clock.


      7 · A real training-script pattern

      Here's a skeleton that works single-GPU today, DDP tomorrow, FSDP next year:

      import os, torch, torch.distributed as dist
      from torch.nn.parallel import DistributedDataParallel as DDP
       
      def setup():
          if "RANK" in os.environ:
              dist.init_process_group("nccl")
              rank = dist.get_rank()
              local_rank = int(os.environ["LOCAL_RANK"])
              torch.cuda.set_device(local_rank)
          else:
              rank = 0; local_rank = 0
          return rank, local_rank
       
      def build_model(cfg, local_rank):
          model = LlamaLike(cfg).to(local_rank)
          if dist.is_initialized():
              model = DDP(model, device_ids=[local_rank])
          return model
       
      def train():
          rank, local_rank = setup()
          model = build_model(CONFIG, local_rank)
          ...

      Launch single-GPU: python train.py. Launch multi-GPU: torchrun --nproc_per_node=4 train.py. Same code path.


      8 · Gradient accumulation — the poor person's DDP

      Can't afford more GPUs? Simulate a bigger batch by running k forward/backward passes and accumulating gradients before stepping:

      optimizer.zero_grad()
      for i, (x, y) in enumerate(loader):
          loss = model(x, y) / accum_steps
          loss.backward()
          if (i + 1) % accum_steps == 0:
              optimizer.step()
              optimizer.zero_grad()

      Effective batch size = micro_batch × accum_steps. Comes at zero comms cost — you're just running the training loop accum_steps times per optimizer update.

      For our 100M run on one A100, we might set micro_batch=32, accum_steps=8 → effective batch 256. Roughly what published Llama configs use per-GPU. Costs no extra memory (activations are freed between micro-batches).


      9 · War stories

      War story The FSDP wrap-policy that ate all my memory

      I FSDP-wrapped a 1B model, launched on 8 GPUs, expected ~2 GB per GPU. Actual: 12 GB per GPU. Turned out my auto_wrap_policy was too coarse — it wrapped the whole model as one FSDP unit, which meant the entire model was all-gathered at once during forward. FSDP degenerated to DDP under a fancier name.

      Lesson: always use transformer_auto_wrap_policy with your block class explicitly named. Verify with print(model) — you should see one FSDP wrapper per block.

      War story The DDP `find_unused_parameters` slowdown

      Set find_unused_parameters=True because I had a conditional path in my model. Training was mysteriously 30% slower. That flag adds a pass over all params after every backward to figure out which weren't used — expensive. Restructured the code to not have the conditional, dropped the flag, got the speed back.

      Lesson: find_unused_parameters=True is an emergency exit, not a default. If you're setting it, refactor your model instead.

      War story The BLOOM 176B chronicles — what breaking 384 A100s teaches you

      The BigScience BLOOM team open-sourced their entire training logbook (github.com/bigscience-workshop/bigscience/blob/master/train/tr11-176B-ml/chronicles.md) — required reading before you touch multi-node. What actually happened over their 4-month run on 384 A100-80GB:

      • Hardware failure rate: on average 1–2 GPUs per week hard-failed and had to be swapped. On a bad week, an entire node (8 GPUs) went down. Restart-from-checkpoint was invoked ~40 times over the run.
      • NCCL hangs: at least a dozen incidents where the collective ops silently hung with no traceback. Only symptom: nvidia-smi shows GPUs at 100% but no throughput. Cure was always the same — kill the job, restart from checkpoint.
      • Silent numerical corruption: at one point a single GPU on one node started producing occasionally-NaN gradients. AllReduce spread the NaNs to the whole cluster. Debugging took a week; the culprit was a marginally overheating HBM stack. Lesson: log per-rank grad norms, not just the reduced norm.
      • Loss spikes: they hit 12 significant spikes (>0.5 nats above trend). Their protocol: rewind ~200 steps, skip the offending data batch, resume. They never determined a single "cause" — a mix of unlucky batches (extremely long low-entropy sequences), tail-end fp16 range issues (they were on fp16, not bf16 — modern runs use bf16 for this exact reason), and occasional param corruption from silent hardware faults.

      Lesson: every 1B+ multi-node run needs (a) per-rank grad-norm logging, (b) automatic rewind-and-skip on spikes, (c) HBM temp monitoring, (d) NCCL_ASYNC_ERROR_HANDLING=1 set. Assume something will break every day.

      War story Meta OPT-175B's 'main issues' file — the same story

      Meta's OPT-175B team similarly open-sourced their chronicles.md (github.com/facebookresearch/metaseq/tree/main/projects/OPT/chronicles). Highlights that mirror BLOOM:

      • 35+ restarts over the 2-month run on 992 A100-80GB.
      • Hardware failures the dominant cause of downtime; NCCL hangs #2.
      • Their fix for repeated loss divergences at ~step 78k: lower LR by 10× for 5000 steps then resume schedule. Ad-hoc, worked.
      • Susan Zhang's talk on this run ("Training OPT-175B: A Story of Chaos") is on YouTube; watch it before your first multi-node.

      Lesson: frontier training is not a smooth optimization curve, it's a series of careful recoveries. Build your infra assuming this.


      10 · The 2025 receipts — FSDP2, DTensor, and Composer

      PyTorch's distributed story matured hard in 2024–2025. If you last read the FSDP docs in 2023, half of what you know is now outdated. Here's the diff.

      10.1 · FSDP2 (per-parameter sharding, PyTorch 2.4+)

      Original FSDP shards per FlatParameter — it concatenates all params in a wrapped module into one giant flat tensor and shards that. Simple, but has three problems: (1) you can't state_dict() a single layer without unsharding the whole group; (2) mixed precision is per-FlatParameter, not per-tensor — so a layer that needs fp32 forces the whole group to fp32; (3) LoRA and other partial-tuning setups are painful.

      FSDP2, shipped in PyTorch 2.4 (July 2024) and stable in 2.5+, shards per parameter tensor using the new DTensor (distributed tensor) primitive. The API is different — no wrap policy, you call fully_shard(module) and it does per-param sharding under the hood. Benefits:

      • Cleaner state_dict: model.state_dict() returns DTensors you can materialize per-param.
      • Per-tensor mixed precision (fp32 norms, bf16 linears, fp8 experts — all in one model).
      • Composes with tensor parallelism natively via parallelize_module.
      • 5–15% throughput improvement on 8+ GPU workloads (per PyTorch 2.5 release notes).

      When to use FSDP1 vs FSDP2 in mid-2025: if you're starting fresh, use FSDP2. If you have a working FSDP1 script, stay on it — the migration has sharp edges around checkpoint format (DCP works with both, but the on-disk layouts differ). See PyTorch's torch.distributed.fsdp.fully_shard docs.

      10.2 · DTensor — the new abstraction under everything

      torch.distributed.tensor.DTensor is a tensor that knows how it's sharded (Shard(dim=i), Replicate(), Partial()) across a DeviceMesh. Every 2025 PyTorch distributed feature is built on it:

      • FSDP2 = DTensor with Shard(0) on the FSDP mesh.
      • Tensor Parallelism (parallelize_module) = DTensor with Shard(0) or Shard(1) on the TP mesh, plus rules for how each operator (matmul, layernorm) handles the sharding.
      • 2D parallelism (FSDP × TP) = DTensor on a 2D DeviceMesh: mesh_2d = init_device_mesh("cuda", (dp, tp), mesh_dim_names=("dp", "tp")).
      • Pipeline parallelism (torch.distributed.pipelining, stable in 2.5) also uses DTensor for stage boundaries.

      The practical impact: for the first time, you can write a single training script that scales from 1 GPU to 8 to 512 without rewriting the parallelization glue. torchtitan (github.com/pytorch/torchtitan, Meta's reference LLM training repo) is the canonical example — read its train.py before your first multi-node run.

      10.3 · MosaicML Composer + LLM Foundry — batteries-included alternative

      If you'd rather not hand-roll FSDP + AMP + gradient accumulation + checkpointing + wandb, MosaicML's Composer (github.com/mosaicml/composer) wraps it all. Their LLM Foundry (github.com/mosaicml/llm-foundry) is the training script Databricks used for DBRX (132B MoE, 3072 H100s, March 2024). Highlights of what Composer buys you for free:

      • Automatic FSDP wrapping with sensible policies.
      • Auto-resume from the latest checkpoint on job restart (S3 or local).
      • Loss-spike detection with automatic batch-skip recovery (their write-up: mosaicml.com/blog/mpt-7b describes their spike protocol).
      • Deterministic dataloader that survives resume mid-epoch — the single most annoying thing to build yourself.
      • Native bf16 + fp8 mixed precision (via Transformer Engine on H100).

      DBRX's public training receipts (Databricks blog, March 2024): 132B total / 36B active MoE, 12T tokens, 3072 H100s, ~2 months wall-clock, ~$10M in compute. For comparison, their previous MPT-7B run was 8T tokens on 512 A100s in 9.5 days. Composer scaled the same code across both.

      Lesson for you: for a 100M run, hand-roll it — you'll learn more. For anything >1B, seriously consider Composer. The hand-rolled version is a career-long tax you don't need to pay.

      10.4 · The 2025 comms landscape

      A quick reality check on interconnects, mid-2025 pricing:

      InterconnectTypical GB/s (per GPU)Where you find it
      PCIe Gen4 x16~32Consumer / low-end cloud
      PCIe Gen5 x16~64Newer 4090/L40 boxes
      NVLink 3 (A100)600AWS p4d, GCP a2, Lambda A100
      NVLink 4 (H100)900AWS p5, GCP a3, Lambda H100
      NVLink 5 (H200/B200)1800AWS p5e, Nebius, CoreWeave 2025
      InfiniBand HDR (200Gb/s)25Multi-node A100
      InfiniBand NDR (400Gb/s)50Multi-node H100, standard 2025
      InfiniBand XDR (800Gb/s)100Blackwell clusters, late 2025

      Rule of thumb the numbers give you: FSDP is bandwidth-bound. If your AllGather + ReduceScatter per step exceeds ~30–50% of step time, you're bottlenecked. Fixes in order of impact: (1) larger micro-batch → more compute per gather; (2) activation checkpointing → fewer gathers per step; (3) upgrade interconnect (💸); (4) switch to 2D parallelism (TP inside the node, FSDP across nodes).

      Further reading (10.x)


      Common misconception
      ✗ What most people think

      "DDP and FSDP both let me use N GPUs, so they're basically the same thing with different memory tricks. FSDP is just DDP that saves VRAM."

      ✓ What is actually true

      They move different things over the wire, and that changes what hardware they need. DDP communicates gradients once per step — a fixed volume proportional to parameter count, and it can be overlapped with the backward pass almost entirely. FSDP communicates parameters too: it all-gathers each shard's weights before every layer's forward and again in backward, then reduce-scatters gradients. The memory saving is real, but it is bought with substantially more collective traffic per step.

      Why the myth is so sticky

      Because on a single node with fast interconnect, the extra traffic hides. NVLink is fast enough that FSDP's all-gathers overlap with compute and you see the memory drop with barely any throughput cost — so "FSDP is free memory" is a perfectly reasonable conclusion from your first 8-GPU experiment. The belief breaks when you cross to multiple nodes over slower fabric, where the parameter all-gather no longer hides behind compute and the step time visibly stretches. The myth is formed in the regime where communication is free and applied in the regime where it is the bottleneck.

      Prove it to yourself

      Instrument the step and watch where the time goes, not just the memory:

      import time, torch
      def timed_step(model, batch):
          torch.cuda.synchronize(); t0 = time.time()
          loss = model(**batch).loss
          loss.backward()
          torch.cuda.synchronize()
          return time.time() - t0, torch.cuda.max_memory_allocated() / 1e9
      
      # run identically under DDP and under FSDP, same global batch
      # then rerun with the SAME code across 2 nodes instead of 1
      # DDP step time should barely move; FSDP step time should grow
      From first principles
      Start with the question

      Why does ZeRO-3 / FSDP cut memory by roughly the world size while DDP cuts it by nothing at all — when both run the identical model on the identical GPUs? The asymmetry has one cause.

      1. 1
        Training state is four things: parameters, gradients, optimizer state, and activations. With Adam in mixed precision, optimizer state (two moments, usually fp32) plus an fp32 master copy dominates the parameter-related bill — several times the raw parameter bytes.
        forced by · Adam keeps first and second moment estimates per parameter, and mixed precision keeps a high-precision master weight alongside the low-precision compute copy
      2. 2
        DDP replicates the model on every rank. Each rank holds a full copy of all four categories and computes on its own slice of the batch.
        forced by · DDP parallelises data, not state — that is the entire definition of data parallelism
      3. 3
        So DDP's only collective is an all-reduce of gradients at the end of backward, keeping every replica's weights bit-identical. Memory per rank is therefore independent of world size.
        forced by · if every rank must be able to run the whole model alone, every rank must store the whole model
      4. 4
        ZeRO observes that this replication is pure redundancy: at any instant, a given optimizer-state entry is only read by the one update that touches its parameter. Nothing requires all ranks to hold it.
        forced by · the optimizer update is elementwise — parameter i's moments are needed only to update parameter i
      5. 5
        So shard it. But then a rank no longer holds the weights it needs for the forward pass, and must fetch them. Hence the all-gather before each layer, and the reduce-scatter of gradients so each rank only receives the gradient slice it owns.
        forced by · you cannot compute a layer without its full weight matrix, so sharded storage forces just-in-time reconstruction
      ⇒ Therefore

      Therefore memory-per-rank for the sharded categories falls as roughly 1/world_size, and the price is a parameter all-gather per layer per pass instead of a single gradient all-reduce per step. ZeRO-1, 2 and 3 are just how many of the four categories you have chosen to shard.

      And note what this predicts: activations are not sharded by ZeRO — they scale with per-rank batch size and sequence length, not world size. So a model that OOMs under FSDP at long sequence length will keep OOMing no matter how many GPUs you add, and the fix must be activation checkpointing or sequence/context parallelism, never more ranks. Verify it: hold per-rank batch fixed, double the world size, and watch peak memory fall for a short sequence but plateau for a long one.

      Mental modelOwn a slice, borrow the rest, give it back

      Picture the model's parameters as a long shelf of books and each GPU as a reader who permanently owns only 1/N of the shelf. To read chapter k, everyone shouts their piece of chapter k across the room, each reader assembles a temporary full copy, reads it, and immediately throws the borrowed pages away.

      DDP is the opposite arrangement: every reader owns the entire shelf, so nobody ever borrows — but the room is N times as full, and at the end of each round everyone shouts their annotations so all copies stay identical.

      • DDP: full replica per rank, one gradient all-reduce per step. Memory flat in N, comms low. FSDP/ZeRO-3: 1/N of params+grads+optimizer state, all-gather per layer. Memory ∝ 1/N, comms high.
      • ZeRO-1 shards optimizer state, ZeRO-2 adds gradients, ZeRO-3 adds parameters. Each stage saves more memory and adds more traffic. Take the lowest stage that fits.
      • Activations are outside this whole scheme. They scale with per-rank batch × sequence, and only activation checkpointing or a different parallelism axis touches them.
      • Gradient accumulation buys effective batch size with zero extra communication — it is the free knob, and the one to reach for before adding ranks.
      🔔 Fires when you see

      Fire this model the moment you see: an OOM that persists after adding GPUs · near-linear scaling on one node that collapses across nodes · a step time dominated by NCCL kernels in the profiler · someone recommending ZeRO-3 for a model that already fits · a loss curve that differs between 1-GPU and N-GPU runs at the same effective batch size.

      The tradeoff

      Your model does not fit comfortably on one GPU. Do you shard state with FSDP, keep DDP and shrink the per-rank footprint, or accept a smaller model?

      Stay on DDP + activation checkpointing + gradient accumulation
      + you gain the simplest possible failure surface — every rank is a complete model, so checkpoints are trivial, debugging is single-process reasoning, and comms stay at one overlappable all-reduce per step; scaling across nodes degrades gracefully
      − you pay activation checkpointing costs an extra forward pass worth of compute, and none of it helps if the parameters plus optimizer state alone exceed one GPU — at that point DDP simply cannot run
      pick when parameters + gradients + optimizer state fit on one device with room to spare, and your pressure is coming from activations rather than from state
      FSDP / ZeRO-3
      + you gain unlocks models whose state cannot fit on a single device at all; memory per rank falls roughly linearly in world size, so adding GPUs genuinely buys capacity rather than just speed
      − you pay parameter all-gathers on every forward and backward, which stop hiding behind compute once you leave a high-bandwidth intra-node fabric; sharded checkpoints need their own save/load path; and debugging is harder because no single rank ever holds a complete model
      pick when the state alone exceeds one device, or you measure that the memory freed lets you raise per-rank batch enough to recover the throughput the extra comms cost
      ZeRO-1 or ZeRO-2 as the middle setting
      + you gain optimizer state is the largest and most redundant category, so sharding just that (ZeRO-1) often recovers most of the memory for a communication pattern much closer to plain DDP
      − you pay if parameters themselves are the problem, this does not help; and you now have a configuration that is neither the well-trodden DDP path nor the well-trodden FSDP path
      pick when optimizer state is what pushed you over the edge — measurable by comparing peak memory with Adam against a momentum-free optimizer on the same model
      What a senior engineer actually does

      Take the least aggressive option that fits, and measure before escalating. The ordering that actually matters is: gradient accumulation → activation checkpointing → ZeRO-1 → ZeRO-2 → ZeRO-3, because each step down that list adds either compute or communication, and you should only pay when the previous rung has demonstrably run out.

      The trap to avoid is reaching for FSDP by default because it is what large-scale write-ups use. Those write-ups are describing models that genuinely cannot fit; if yours fits, FSDP hands you extra collectives, a harder checkpoint story, and a more confusing debugging experience in exchange for memory you were not short of. Always report peak memory and tokens/sec together — a configuration that halves memory and halves throughput has bought you nothing you could not have bought with gradient accumulation.


      11 · 🧠 Retention scaffold

      Quick recall · click to reveal
      ★ = stretch question

      One-line summary (write it in your own words): _______________________________

      Spaced review: re-read §1 (memory accounting) and §6 (decision tree) in 24 hours. Revisit the full session on day 7 focusing on §10 (2025 receipts) and the two chronicles.

      Next session (S052): rent an actual GPU — Lambda, RunPod, Modal walkthrough with 2025 dollar amounts.

      Sticky note (keep on your desk): "DDP if you fit. FSDP2 if you don't. TP inside the node. Multi-node = IB or don't bother."

      Legacy recall drills (kept for spaced-review continuity)

      1. Roughly how many bytes per parameter does bf16 Adam training need, and what breaks that down?

      ~16 B/param. 2 (bf16 params) + 2 (bf16 gradients) + 12 (fp32 optimizer master weights, m, v). Plus activations on top.

      2. What memory does DDP save vs single-GPU training?

      None. DDP replicates the full model on every GPU. It parallelizes compute (bigger effective batch, faster wall clock) but doesn't reduce per-GPU memory.

      3. What do ZeRO stages 1, 2, and 3 each shard?

      Stage 1: optimizer states. Stage 2: also gradients. Stage 3: also parameters. PyTorch FSDP = ZeRO stage 3.

      4. Why does FSDP need fast interconnect (NVLink/InfiniBand)?

      Because FSDP moves ~6N bytes per training step (3× more than DDP). If comms can't complete inside step compute time, GPUs stall. On slow PCIe interconnects, FSDP can lose 60%+ of throughput.

      5. What is gradient accumulation and when do you use it?

      Running k forward/backward passes and summing gradients before stepping. Simulates a larger batch on one GPU. Zero extra memory or comms. Use it whenever you want a bigger effective batch than fits in one micro-batch.

      Stretch

      Build a memory-estimator. Function estimate(N, world_size, mode) where mode ∈ {"single", "ddp", "zero1", "zero2", "fsdp"} returns per-GPU steady-state bytes. Add optional activation_checkpointing=True that shrinks activation memory by 4×. Use it to answer: "on 8×A100-40GB, what's the biggest model I can train under each mode?"

      In your own words

      "The three axes of scaling a training run are ____, ____, ____. DDP moves you along the first only; FSDP unlocks the second; TP/PP unlock the third."

      Spaced-review pointer

      • S018 (mixed precision) — bf16 vs fp32 costs matter to every calculation here.
      • S050 (architecture) — the per-param memory is fixed by the choices made there.

      Next-session teaser

      Enough theory. Session 052 is the actual GPU rental walkthrough — how to spin up an A100 on Runpod/Modal/Vast for under $2/hour, sync your training code, launch, monitor, and kill it before your credit card cries. Screenshot walkthrough, real dollar amounts, real commands.

      Bring back tomorrow

      • Number: 16 B/param bf16 Adam steady-state.
      • Rule: DDP = replicate + all-reduce; FSDP = shard + all-gather.
      • Command: torchrun --nproc_per_node=8 train.py.

      Previous: ← DL S050 · Next: DL S052 →