←Home KnowML
Systems, Safety & InterviewChapter 28

GPU Architecture, CUDA & Distributed Training

What a GPU is, why its hardware fits neural network math almost by accident, and the distributed training layer built on top of it: SIMT cores, warps, tensor cores, CUDA and Triton, then DDP, ZeRO, FSDP, tensor and pipeline parallelism, and how to choose between them with arithmetic instead of taste.

42 min read Assumes: Efficient AI & Systems (23), Attention & Transformers (08)
Start reading
TL;DR

A GPU is thousands of simple cores built to execute the same instruction across thousands of pieces of data at once. That is SIMT: single instruction, multiple threads. It is exactly the shape of dense matrix multiplication, the operation that dominates deep learning compute. Warps of 32 threads execute in lockstep. Tensor cores fuse whole small-matrix multiply-accumulates into one hardware operation.

CUDA lets you write a kernel from a single thread's point of view while the hardware runs that same code across a grid of thread blocks. Triton exists because writing that CUDA well takes deep, specialized expertise: coalesced memory access, enough resident warps to hide memory latency. Most ML researchers don't have that expertise and, increasingly, don't need it.

Scale training beyond one GPU and the bottleneck moves from compute to communication. NVLink connects GPUs on the same node at very high bandwidth. InfiniBand connects nodes at meaningfully lower bandwidth. Ring all-reduce is the algorithm that keeps gradient synchronization from becoming the bottleneck as you add more GPUs, because each GPU only ever sends and receives a bounded amount of data, independent of cluster size.

Past one node the question becomes which thing you shard. Data parallel replicates everything and splits the batch. ZeRO and FSDP shard the optimizer, then the gradients, then the weights. Tensor parallel splits a single matrix multiply; pipeline parallel splits the model by depth. Add them in that order and stop at the first one that fits.

One sentence for an interview: GPUs win at deep learning because SIMT hardware matches dense matmul's shape almost perfectly, and distributed training scales because ring all-reduce turns gradient synchronization into a communication cost that doesn't grow with the number of GPUs.

01Intuition

A GPU is a factory floor with thousands of simple, identical workers, each doing the same small step of the same job at the same time. A CPU is a handful of extremely capable generalists, each able to do a completely different, complicated task, including deciding what to do next based on what just happened.

A CPU core is built to run a long, branchy, sequential program fast. It has deep pipelines, branch predictors, out-of-order execution, and a lot of on-chip logic dedicated purely to guessing what instruction comes next and recovering cheaply when the guess is wrong.

That is a great deal of silicon spent on control flow, and a CPU has only tens of these very capable cores, not thousands.

A GPU makes the opposite bet. It gives up almost all of that control-flow machinery and instead packs thousands of small, simple cores onto the same chip. All of them execute the exact same instruction at the exact same time, just on different pieces of data. That architecture has a name: SIMT, single instruction, multiple threads.

The bet cuts both ways:

  • Poor for a workload full of unpredictable branches, because every one of those thousands of cores has to be doing the same thing at once.
  • Close to perfect for a workload that repeats the same simple arithmetic an enormous number of times over different numbers.

Deep learning is exactly that second kind of workload.

A neural network's forward pass is, at the arithmetic level, matrix multiplications and elementwise operations repeated layer after layer.

A matrix multiply is thousands or millions of independent multiply-and-add operations: every output element depends only on one row and one column of the inputs, and none of those computations depend on each other, so there is almost no branching anywhere in it.

That makes it an embarrassingly parallel workload in the technical sense: parallelizing it needs no cleverness and no communication between the parallel pieces, only the same tiny operation many times over.

The accident underneath all of this

GPUs were originally built to compute pixel colors, which is also an enormous number of small, independent, identical computations.

So the hardware that turned out to suit neural networks was designed for something else, and the fit is close to an accident: the two workloads happen to have the same shape.

A GPU is fast at arithmetic, so what limits a program is often how many calculations it does per byte it moves. That ratio is called arithmetic intensity.

Try it Compare the arithmetic intensity of an add and a matrix multiply
n = 4096
add_flops, add_bytes = n * n, 3 * n * n * 4        # C = A + B, float32
mm_flops, mm_bytes = 2 * n ** 3, 3 * n * n * 4     # C = A @ B, each once

print(f"add:    {add_flops / add_bytes:8.2f} FLOP per byte")
print(f"matmul: {mm_flops / mm_bytes:8.2f} FLOP per byte")
Adding two 4096x4096 matrices does 0.08 floating-point operations per byte moved; multiplying them does 682.67. The add spends almost all its time waiting on memory, and the multiply has plenty of arithmetic to hide the waiting behind. That gap, about 8,000 times, is why the same GPU is memory-bound on one and compute-bound on the other.

02Timeline

Before

Neural networks trained on CPUs or on early GPUs built to rasterize triangles and shade pixels, so general computation meant disguising your math as a rendering problem, with no deep learning tooling or purpose-built matrix hardware.

→
Innovation

CUDA (2007) made GPUs a general-purpose compute platform, cuDNN and tensor cores (Volta, 2017) added library and hardware support for matrix-heavy deep learning, ring all-reduce (Baidu, 2017) let hundreds of GPUs average gradients without a central bottleneck, Horovod (2018) popularized it, and Triton (2019 paper, open-sourced by OpenAI in 2021) made most of hand-tuned CUDA's performance reachable.

→
After

Frontier training now runs on thousands of GPUs across many nodes, with parallelism chosen around the interconnect topology, and hardware- and kernel-level efficiency (page 23) is as central as any algorithmic idea.

03Architecture: from the code you write to the silicon that runs it

The factory-floor picture is the intuition. Here is the actual hierarchy, both in the code you write and in the hardware that runs it. The two mirror each other almost exactly, and that is the point of the CUDA programming model: you describe work in software terms, and the hardware underneath was built to match.

The CUDA execution hierarchy — software concepts mapped onto physical hardware
SOFTWARE — what you write HARDWARE — what runs it Grid the whole kernel launch GPU device the entire chip Thread block up to ~1024 threads + shared memory Streaming Multiprocessor (SM) one block runs on exactly one SM Warp 32 threads, one instruction, lockstep Warp scheduler issues an instruction, warp by warp Thread your kernel code, × thousands CUDA / tensor core executes one thread's instruction
Each level on the left has exactly one physical counterpart on the right. This is the literal execution model, not an approximation. A kernel launch is a grid of thread blocks. Each block is assigned to one streaming multiprocessor for its whole lifetime. Each block's threads are grouped into warps of 32, and the SM's warp scheduler issues one instruction to one warp per cycle. Each thread's instruction runs on one core. You write the kernel once, from a single thread's point of view; the hardware is what turns that into thousands of simultaneous executions.

That last point is the mental model to keep. CUDA code is written from one thread's perspective. Inside the kernel you ask "which piece of the problem am I", using built-in indices for your block and your thread within the block, and the hardware launches that same function across every thread in the grid at once.

You never write a loop over threads. The loop is the hardware, running your one function body thousands of times in parallel, with each copy seeing a different index.

04The memory hierarchy: where page 23's bandwidth story physically lives

Page 23 built its framing on one fact: LLM inference is usually bottlenecked by memory bandwidth and not by compute. This section is about where that ceiling comes from: a chain of physical memory that data has to move through, fastest and smallest at one end and slowest and largest at the other.

Four tiers, from closest to the arithmetic to furthest away:

  • Registers. Private to a single thread, sitting closest to the arithmetic units. The fastest memory on the chip, and there is very little of it per thread.
  • Shared memory / L1. Private to one streaming multiprocessor and shared among the threads of a block running on it. Still extremely fast, still small, but large enough to hold a working tile of data that many threads in the same block need to reuse. (The two are often described together, since they occupy the same physical on-chip pool on most architectures.)
  • L2 cache. Shared across the entire chip rather than one SM. Bigger than shared memory/L1, but a further step down in speed.
  • HBM. The GPU's main memory, holding everything: the model's full set of weights, activations, and the KV cache during inference. Enormous compared to what sits above it, and comparatively slow to reach.

HBM, the bottom tier, is what limits a memory-bound decode step. Shared memory is what FlashAttention tiles its computation to fit inside, so intermediate attention scores never have to leave it.

GPU memory hierarchy — small and fast at the top, huge and slow at the bottom
fastest, smallest ↑ Registers private to one thread Shared memory / L1 private to one SM L2 cache shared across the whole chip HBM / global memory huge — but the slow tier model weights + KV cache live here ↓ slowest, largest
Registers and shared memory/L1 live physically inside each streaming multiprocessor, private to it, and about as fast as on-chip memory gets. L2 is shared across the whole chip: larger, slower. HBM sits off the compute die entirely, with huge capacity and by far the slowest access. Every "memory-bound" claim on page 23 is a claim about the gap between the top of this stack and the bottom.

05Streaming multiprocessors, warps, and warp divergence

A streaming multiprocessor, or SM, is the basic physical building block of a GPU. One chip contains dozens of them. Each has its own chunk of cores, its own warp scheduler, its own register file, and its own shared memory/L1 pool. SMs operate largely independently of each other. A thread block, once assigned to an SM, stays there and runs using only that SM's local resources.

A warp is a group of 32 threads that the hardware always schedules and executes together. Same instruction, same time, different data. It is the actual physical unit of SIMT execution, one level below the thread block you write your code in terms of. Whatever a thread block's size is, the hardware internally splits it into warps of 32 and issues instructions warp by warp.

Warp divergence is what happens when threads inside the same warp need to take different branches of an if/else.

Say 16 of the 32 threads satisfy some condition and the other 16 don't. A warp can only execute one instruction stream at a time across all of its lanes, so the GPU can't run both branches simultaneously. It serializes them:

  • Run the if branch across the whole warp, masking off the 16 threads that shouldn't be doing that work so their results don't get written.
  • Run the else branch across the whole warp, masking off the other 16.

Every thread in the warp sits through both branches, so only half the work in each pass is useful.

That is why branch-heavy, data-dependent control flow is slow on a GPU. Dense matrix multiplication and elementwise operations avoid it almost entirely, since every thread in every warp runs the identical instruction on different numbers with almost nothing to diverge on.

06Tensor cores: dedicated hardware for matrix multiply-accumulate

A regular CUDA core does general-purpose arithmetic, which is useful for a huge range of operations but not specialized for any of them.

A tensor core is different. It is a separate, purpose-built piece of hardware sitting alongside the regular cores on the same SM, and it does exactly one thing extremely fast: a fused multiply-accumulate on small tiles of matrices, computed as a single hardware operation rather than as a sequence of individual multiplies and adds.

This is the hardware reason mixed-precision training, covered on page 23, runs faster and not merely smaller.

Reducing precision from FP32 to FP16 or BF16 shrinks the bytes you have to move. That is the memory-bandwidth half of the story. Tensor cores are the other half: they execute matrix multiplication at these reduced precisions at dramatically higher throughput than the general-purpose cores achieve, even at matched precision.

So training in FP16 or BF16 isn't only about fitting more into memory. It routes your matrix multiplies onto hardware built to be fast at exactly that operation at exactly that precision.

07The CUDA programming model: kernels, occupancy, and hiding latency

A kernel is a function you write once and launch across a grid of thread blocks. The diagram above showed the shape of it. Each block gets assigned to one SM, and within that block, many warps of 32 threads each get scheduled onto that SM's cores.

The mental shift, coming from ordinary sequential programming, is that you don't write a loop over the data. You write the body of the loop from one thread's point of view, and the hardware runs that body across every thread in the grid at once.

Occupancy is how many warps can be resident on an SM at the same time, out of some hardware maximum. It is determined by how much of the SM's limited register file and shared memory each thread block consumes. A kernel that uses a lot of registers or shared memory per thread fits fewer blocks, and therefore fewer warps, onto an SM at once. A lean kernel fits more.

Occupancy matters because of what happens when a warp stalls.

If a warp issues a memory request to HBM, it has to wait for that data to come back before it can keep computing. HBM is slow relative to the chip's arithmetic throughput, exactly as the memory hierarchy above lays out.

A warp scheduler with only one resident warp sits idle during that wait. A scheduler with many resident warps switches to a different warp that has work ready to go, right away, for that cycle. The first warp's memory latency gets hidden behind another warp's useful compute.

That is why occupancy matters. More resident warps give the scheduler more options to hide latency behind. Push a kernel's register or shared-memory usage too high and fewer blocks fit on an SM, occupancy drops, and the SM starts idling on HBM instead of hiding that wait.

08Triton: why most people no longer hand-write CUDA

Writing efficient CUDA requires the expertise from the last two sections, all at once:

  • Coalesced memory access, so a warp's threads read contiguous, aligned memory instead of scattering many small requests across HBM.
  • Register and shared-memory budgeting, to keep occupancy high.
  • Manual thread-indexing math, correct for every kernel you write.

Getting all of that right at the same time is a specialized skill, closer to hardware engineering than to the tensor-and-layer thinking most ML researchers work in, and building a new hand-tuned CUDA kernel for every new idea is a cost most research doesn't have time to pay.

Triton exists to remove that cost. It is a Python-embedded language and compiler where you write kernels at the level of tiles or blocks of data rather than individual threads. You describe the computation on a block of, say, 128 elements at once, using ordinary-looking array operations. Triton's compiler automatically handles the low-level thread scheduling, memory coalescing, and shared-memory allocation you would otherwise have to hand-tune in raw CUDA.

It doesn't always match a world-class hand-written CUDA kernel, but it routinely gets close: a large fraction of hand-tuned CUDA's performance for a fraction of the development time and expertise.

You have probably already used the result. FlashAttention's efficient reference implementation is written in Triton, and a large and growing share of PyTorch's own generated custom kernels are Triton kernels. OpenAI open-sourced Triton in 2021 because it delivered enough performance per engineering hour to be worth building a compiler around.

09The interconnect hierarchy: NVLink, InfiniBand, and Ethernet

Everything so far happens inside one GPU. Distributed training needs GPUs to talk to each other, and the physical link they talk over creates a different bottleneck from anything discussed above.

Three tiers, fastest to slowest:

  • NVLink, within one node. A direct, very high-bandwidth link built specifically for GPU-to-GPU communication, roughly hundreds of GB/s (approximate, and it varies by generation of hardware).
  • InfiniBand, or RoCE, between nodes. Still fast, purpose-built for high-performance computing and large training clusters, but roughly an order of magnitude or more below NVLink's bandwidth. (RoCE is RDMA over Converged Ethernet, which achieves similar low-latency direct-memory-access behavior over Ethernet-like hardware.)
  • Plain Ethernet without RDMA, the slow path. General-purpose networking, higher latency, lower throughput. Fine for cluster coordination and logging traffic, but not built to carry the volume of data distributed training needs to move between every synchronization step.

This drives how you place a training job's parallelism strategies onto physical hardware.

Why parallelism strategy is a wiring decision

Tensor-parallel shards exchange partial results within every layer, many times per forward pass and many times per backward pass, so they need to sit on GPUs connected by NVLink. Across an InfiniBand link, the communication cost would dwarf the compute almost immediately.

Data-parallel replicas only need to synchronize once per training step, on the gradient all-reduce. They can tolerate sitting across the slower inter-node InfiniBand link without it dominating step time.

That is why real training runs place tensor parallelism inside a node and data parallelism, typically with pipeline parallelism, across nodes. It is page 23's four parallelism strategies with the hardware reason for where each one goes.

10Collective communication and ring all-reduce

Multi-GPU training relies on a small set of standard communication patterns, called collectives, that every distributed training framework implements.

  • Broadcast. Sends the same data from one GPU to every other GPU.
  • Reduce. Combines data from every GPU, typically by summing, down onto one destination GPU.
  • All-reduce. Combines data from every GPU and gives every GPU the combined result. Exactly what gradient synchronization needs: every replica needs the same averaged gradient.
  • All-gather. Each GPU holds a different piece of data, and it ends with every GPU holding all the pieces, concatenated.
  • Reduce-scatter. Combines data from every GPU, but splits the combined result so each GPU ends up with only its own slice of it rather than the whole thing.

All-reduce is usually implemented as a reduce-scatter followed by an all-gather, which is exactly how ring all-reduce is built.

Arrange the N GPUs participating in an all-reduce in a logical ring. Each GPU has exactly two neighbors, one it sends to and one it receives from. Split each GPU's local data (its gradients) into N chunks.

The scatter-reduce phase runs for N−1 steps. At every step, every GPU simultaneously sends one chunk to its right-hand neighbor and receives a chunk from its left-hand neighbor, adding the received chunk into its own copy of that same chunk index. The chunk index being passed rotates each step. After N−1 steps, every GPU holds the fully-reduced sum for exactly one chunk index, a different one per GPU.

The all-gather phase then runs for another N−1 steps, using the same ring. Now each GPU just relays its already-final chunk onward instead of adding to it, until every GPU has received every chunk in its fully-reduced form.

The key property, at every step of both phases, is that each GPU only ever sends and receives one chunk's worth of data: a bounded amount, independent of how many GPUs are in the ring.

11The math: ring all-reduce's communication cost

$$\text{Data moved per GPU} \approx \frac{2(N-1)}{N} \times S$$
  • N the number of GPUs participating in the ring.
  • S the size of the data being all-reduced. For gradient synchronization, that is the total size of the gradients.
  • 2(N-1)/N N−1 steps of scatter-reduce plus N−1 steps of all-gather, each step moving 1/N of the data, for a total of $2(N-1) \times S/N$ per GPU.
  • The limit as N grows large, $2(N-1)/N \to 2$. A constant, not something that scales with the number of GPUs.
Try it Run both phases of a ring all-reduce and count the bytes
import numpy as np

N, S = 4, 12                            # ranks in the ring, elements each
start = [np.arange(S, dtype=float) + 100 * r for r in range(N)]
truth = sum(start)                      # what all-reduce must produce

buf = [d.copy() for d in start]
chunk = [slice(i * S // N, (i + 1) * S // N) for i in range(N)]
sent = 0

for step in range(N - 1):               # phase 1: scatter-reduce
    nxt = [b.copy() for b in buf]
    for r in range(N):
        c = (r - step) % N              # the chunk rank r forwards now
        nxt[(r + 1) % N][chunk[c]] += buf[r][chunk[c]]
        sent += buf[r][chunk[c]].nbytes
    buf = nxt

for step in range(N - 1):               # phase 2: all-gather
    nxt = [b.copy() for b in buf]
    for r in range(N):
        c = (r + 1 - step) % N          # now it relays, it does not add
        nxt[(r + 1) % N][chunk[c]] = buf[r][chunk[c]]
        sent += buf[r][chunk[c]].nbytes
    buf = nxt

print("every rank ends with the exact sum:",
      all(np.array_equal(b, truth) for b in buf))
print("rank 0 buffer:", buf[0].astype(int))
print("bytes sent per rank:", sent // N)
print("2(N-1)/N x S  predicts:", int(2 * (N - 1) / N * truth.nbytes))
The counted bytes and the formula agree exactly: 144 either way. That is the claim above turned into a measurement rather than an assertion. The other line worth staring at is the rank 0 buffer — every rank finishes holding the identical full sum, and no rank ever held more than one chunk of anybody else's data at a time. Raise N to 32 and re-run: the per-rank byte count barely moves, because $2(N-1)/N$ is already within a few percent of 2 by $N = 4$.

That limit is the elegant part. Doubling your GPU count doesn't double each GPU's communication burden: it stays roughly flat, approaching twice the data size and never exceeding it, however large the ring gets.

Compare a naive centralized parameter server, where every one of the N workers sends its gradients to the server and the server sends updated weights back to all N through that one network link. The link's traffic grows linearly with N, so it becomes the bottleneck you are trying to scale past, at exactly the point where you add more GPUs to go faster.

12Data parallel, exactly: what DDP does at every step

Eight GPUs, a model that fits on each of them, a global batch of 64. Every rank runs the same model on a different eight examples, and the only moment they speak to each other is the gradient all-reduce, so when that all-reduce happens is what the design turns on.

Once, at startup, rank 0 broadcasts its weights to every other rank. All eight now hold byte-identical parameters, optimizer state and gradient buffers. After that, no weight ever crosses the wire again.

Then, every step:

  • Forward. Each rank pushes its own eight examples through its own full copy of the model. No communication whatsoever. The loss each rank computes is a local loss over a local shard of the batch.
  • Backward. Also purely local arithmetic — but as each parameter's gradient is produced, DDP hands it to a bucket. Backward runs from the last layer to the first, so the last layer's gradients are ready long before the first layer's.
  • The all-reduce, in pieces. The moment a bucket is full, its all-reduce launches asynchronously. It travels the interconnect while the layers below it are still computing their own gradients.
  • Sync. At the end of backward, wait on the outstanding bucket all-reduces. Every rank now holds the same averaged gradient, bit for bit.
  • Optimizer step. Run locally on every rank. Nothing is communicated.
The invariant that keeps eight GPUs identical without ever sending a weight

Identical starting weights, plus identical optimizer state, plus a gradient that all-reduce makes bit-identical on every rank, means the optimizer step lands on identical weights everywhere. Synchronisation follows from those three facts; DDP doesn't add a separate step for it.

Break any one of them and the ranks drift apart silently. Batch norm computes its statistics per rank, so it needs SyncBatchNorm. A parameter that receives no gradient on some ranks leaves its bucket waiting, and the collective hangs.

Why buckets, and why the default is 25 MiB

Both obvious schemes are bad. All-reduce each parameter tensor on its own and you launch thousands of tiny collectives, every one paying a fixed launch and latency cost that its payload never amortises. All-reduce everything once at the end and the wire idles through the whole backward pass, then the GPUs idle through the whole transfer.

Bucketing is the compromise. Gradients are packed into contiguous buffers of a fixed size, and a bucket fires as soon as it fills. PyTorch's DDP defaults to 25 MiB per bucket: big enough to amortise the collective's fixed cost, small enough that the first bucket leaves early in the backward pass.

The same step, drawn to scale — one all-reduce at the end versus bucketed overlap
0 20 40 60 ms backward compute 8 layers × 6 ms, last layer first bucketed all-reduce 50.5 ms one bucket's flight time past the compute one big all-reduce 68 ms wire idle for the whole backward pass
Drawn to scale on a single time axis, from the numbers in the Try it below: eight layers at 6 ms of backward compute each, 50 MB of gradient per layer, 20 GB/s of usable bandwidth. Each green bar starts where a blue block ends, and travels the wire while the next block is still computing. What is left over at the end is one bucket in flight — 2.5 ms — instead of the full 20 ms of transfer.
Try it Sweep the bucket size and watch the overlap appear
# Backward runs last layer -> first. A layer's gradient is ready the
# moment its backward finishes, so it can go on the wire while the
# next layer down is still computing. That is the whole trick.
COMPUTE_MS = [6.0] * 8        # backward compute, per layer
GRAD_MB = [50.0] * 8          # gradient size, per layer
MB_PER_MS = 20.0              # 20 GB/s of usable link bandwidth

def wall_clock(bucket_mb):
    t = link_free = pending = 0.0
    for c, g in zip(COMPUTE_MS, GRAD_MB):
        t += c                            # compute this layer
        pending += g                      # its gradient joins the bucket
        while pending >= bucket_mb:       # bucket full -> fire it off
            pending -= bucket_mb
            link_free = max(link_free, t) + bucket_mb / MB_PER_MS
    if pending:                           # flush the last partial bucket
        link_free = max(link_free, t) + pending / MB_PER_MS
    return max(t, link_free)

compute = sum(COMPUTE_MS)
comm = sum(GRAD_MB) / MB_PER_MS
print("backward compute alone      %6.1f ms" % compute)
print("one all-reduce at the end   %6.1f ms  (compute + comm)"
      % (compute + comm))
for b in (400, 200, 100, 50, 25):
    print("bucket %3d MB               %6.1f ms" % (b, wall_clock(b)))
print("\nthe floor is the compute (%.1f ms) plus ONE bucket in"
      " flight: 2.5 ms for 50 MB at this bandwidth." % compute)
Read the 400 MB row against the 50 MB row. A bucket the size of the whole model is the no-overlap case — 68.0 ms either way, because nothing can be sent until the last gradient exists. At 50 MB a bucket closes after every layer and wall clock falls to 50.5 ms: 48 ms of compute, plus one bucket still in flight. Below that it stops helping, and in reality starts costing, because this model charges nothing per message and NCCL does.

Two consequences that catch people out. Effective batch size is per-rank batch times world size, so adding GPUs silently changes your optimisation problem and your learning rate should move with it. And no_sync() exists so that gradient accumulation skips the all-reduce on every micro-step but the last.

13The memory ledger, and the wall data parallel runs into

DDP scales throughput almost perfectly and capacity not at all. Every rank holds a complete copy of the training state, so eight GPUs give you eight times the batch and exactly as much room for the model as one GPU had.

Which raises the question of what "the training state" contains. Parameters are the smallest part of it.

Worked example Why a 7.5B-parameter model needs 120 GB before a single activation

Assume the standard recipe: mixed-precision forward and backward, AdamW, no sharding and no offload. Write $\Psi$ for the parameter count. Every line below is bytes per parameter.

  1. $$\text{fp16 parameters} = 2\Psi = 2 \times 7.5\times10^9 = 15\ \text{GB}$$
    Why reduced precision at all. Section 06: the matmuls have to land on tensor cores, and tensor cores want fp16 or bf16. This copy is the one the forward and backward passes actually read.
  2. $$\text{fp16 gradients} = 2\Psi = 15\ \text{GB}$$
    One gradient per parameter, in the same dtype the backward pass produced it in. This is the tensor DDP all-reduces.
  3. $$\text{fp32 master weights} = 4\Psi = 30\ \text{GB}$$
    The line people forget. fp16 carries about 10 bits of mantissa. Add an update of relative size $10^{-4}$ to an fp16 weight and it rounds straight back to where it started — the training stalls. So the optimizer keeps an fp32 copy, applies the update there, and casts down for the next forward pass.
  4. $$\text{Adam momentum} + \text{variance} = 4\Psi + 4\Psi = 60\ \text{GB}$$
    Adam keeps a running first and second moment per parameter, both fp32. This is the price of the optimizer that made large-model training work, and it is four times the size of the weights it is optimising.
  5. $$2\Psi + 2\Psi + 12\Psi = 16\Psi = 120\ \text{GB}$$
    Sixteen bytes per parameter. This is exactly the accounting in the ZeRO paper (Rajbhandari et al., 2020), which uses the same 7.5B example and reaches the same 120 GB.
The ratio is the point. Parameters are 15 GB of the 120. The optimizer and its master copy are 90 GB — three quarters of the bill — and they are touched exactly once per step, during the update. For the whole forward and backward pass, every rank is sitting on 90 GB of state it is not using. That observation is the whole of ZeRO.

Here is the wall. 120 GB does not fit on an 80 GB H100, and adding GPUs does not help: DDP's answer to eight of them is eight identical 120 GB copies, so you would buy 640 GB of HBM to hold 120 GB of distinct state, and it still would not fit.

Activations sit on top of all of this and none of the techniques in the next section shard them. Activation checkpointing is the orthogonal lever: store only layer boundaries, recompute the rest during backward, and trade roughly 30% more compute for a large drop in activation memory.

Try it The ledger, and what each ZeRO stage does to it
# Mixed-precision AdamW keeps five tensors for every parameter.
PARAM, GRAD = 2, 2           # bf16 weights, bf16 gradients
OPT = 4 + 4 + 4              # fp32 master copy + Adam m + Adam v
TOTAL = PARAM + GRAD + OPT   # 16 bytes per parameter

def per_gpu_gb(psi, n, stage):
    P, G, O = psi * PARAM, psi * GRAD, psi * OPT
    if stage == "DDP":  return (P + G + O) / 1e9
    if stage == "ZeRO-1": return (P + G + O / n) / 1e9
    if stage == "ZeRO-2": return (P + (G + O) / n) / 1e9
    return (P + G + O) / n / 1e9          # ZeRO-3

psi = 7.5e9
print(f"{TOTAL} bytes/param x {psi:.1e} params"
      f" = {psi * TOTAL / 1e9:.0f} GB of training state")
print(f"\n{'GPUs':>5}{'DDP':>9}{'ZeRO-1':>9}{'ZeRO-2':>9}{'ZeRO-3':>9}")
for n in (1, 8, 64, 512):
    cells = [per_gpu_gb(psi, n, s)
             for s in ("DDP", "ZeRO-1", "ZeRO-2", "ZeRO-3")]
    print(f"{n:>5}" + "".join(f"{v:>9.1f}" for v in cells))
print("\nGB per GPU. Weights + grads + optimizer only; activations"
      " are extra and are NOT sharded by any of these.")
Read down the DDP column first. It is 120.0 at every world size, which is the wall stated as a number: replication means more GPUs add zero capacity. Then read across the last row. At 512 ranks, ZeRO-1 still costs 30.2 GB and ZeRO-2 15.2 GB, because both keep a full copy of the parameters on every rank — the 4Ψ and 2Ψ terms never shard away. Only ZeRO-3 goes to zero, and section 14 is about what that costs.

14ZeRO, stage by stage: what is sharded and what it costs

The duplication in DDP buys exactly one thing: every rank can run the optimizer step without talking to anyone. ZeRO gives that up, one tier of the ledger at a time, and buys back memory in proportion to how much it gave up.

The unit of sharding is the parameter index. Rank $i$ owns slice $i$ of the flattened parameter vector, and "owning" a slice means being the only rank responsible for its optimizer state.

  • Stage 1, optimizer states. Each rank keeps the full parameters and full gradients, but only $1/N$ of the fp32 master weights and Adam moments. Per rank: $4\Psi + 12\Psi/N$.
  • Stage 2, plus gradients. Gradients are reduce-scattered rather than all-reduced, so each rank keeps only the slice it owns and frees the rest the moment the bucket is reduced. Per rank: $2\Psi + 14\Psi/N$.
  • Stage 3, plus parameters. The weights themselves are sharded. Before a layer can run, an all-gather reconstructs its full weights; after it runs, the non-owned shards are dropped. Per rank: $16\Psi/N$.
Per-GPU training state, 7.5B parameters on 64 GPUs — bars drawn to scale
80 GB — one H100 DDP 120.0 GB · 2Ψ on the wire ZeRO-1 31.4 GB · 2Ψ on the wire ZeRO-2 16.6 GB · 2Ψ on the wire ZeRO-3 1.9 GB · 3Ψ on the wire — 1.5× DDP 0 40 80 120 GB per GPU
Bar length is linear in gigabytes: 120.0, 31.4, 16.6 and 1.9, from the formulas above at $\Psi = 7.5\times10^9$ and $N = 64$. The pink line is an 80 GB H100's capacity, and it is the reason the first bar is not a viable configuration at any cluster size. Note that stages 1 and 2 are still dominated by their unsharded term — the bars would barely move if you doubled the GPU count. Only stage 3 keeps shrinking.
What each rank keepsPer-GPU bytesTraffic per step7.5B on 64 GPUs
DDPEverything, replicated$16\Psi$$2\Psi$120.0 GB
ZeRO-1Full params + grads, $1/N$ of optimizer$4\Psi + 12\Psi/N$$2\Psi$31.4 GB
ZeRO-2Full params, $1/N$ of grads + optimizer$2\Psi + 14\Psi/N$$2\Psi$16.6 GB
ZeRO-3$1/N$ of everything$16\Psi/N$$3\Psi$1.9 GB

The traffic column is the part to understand, because it is where intuition usually goes wrong. Stages 1 and 2 are free in communication terms: they move exactly as many bytes as plain DDP does.

The reason is section 10. A ring all-reduce is a reduce-scatter followed by an all-gather, $\Psi$ of traffic each, $2\Psi$ in total. ZeRO-1 and ZeRO-2 stop between those two halves, let each rank update the slice it now owns, and then all-gather the updated parameters instead of the summed gradients. Same two collectives, same bytes.

Stage 3 is the one that costs. Because no rank holds a full layer, the weights have to be all-gathered before the forward pass can use them, and again before the backward pass can use them. That is $\Psi + \Psi$ on top of the gradient reduce-scatter's $\Psi$: $3\Psi$, or 1.5× DDP. The figures in this section are the ZeRO paper's own.

Stage 3 moves the collective onto the critical path

The 1.5× understates it. DDP's one all-reduce hides behind the backward pass. Stage 3's all-gathers sit in front of every layer in both directions, so a layer cannot start until its weights arrive.

Implementations prefetch the next layer's weights while the current one computes, which hides the cost on NVLink and does not hide it on a slow fabric. The routine outcome: stage 2 is faster in wall clock, stage 3 is the only one that fits. Try stage 2 first and only escalate when it OOMs.

Both implementations also offload. DeepSpeed's ZeRO-Offload pushes optimizer state and the fp32 master copy to CPU memory, and ZeRO-Infinity extends that to NVMe. Both trade a much slower link for capacity you could not otherwise buy, so use them as a last resort.

FSDP versus DeepSpeed ZeRO: what is different

The algorithm is the same one. PyTorch's docs map FULL_SHARD onto ZeRO stage 3 and SHARD_GRAD_OP onto stage 2 explicitly, and NO_SHARD is DDP. Nothing in the memory or traffic table above changes depending on which library you run it through.

Just namingDifferent
"Sharding strategy" vs "ZeRO stage" — FULL_SHARD is stage 3, SHARD_GRAD_OP is stage 2, NO_SHARD is DDP.FSDP is in PyTorch core and composes with the rest of it. DeepSpeed is a separate engine with its own config JSON, launcher and checkpoint format.
"Shard" vs "partition"; "FSDP unit" vs "parameter group". Same object, two vocabularies.Sharding granularity. FSDP shards whatever you wrap, so a careless auto_wrap_policy — one unit for the whole model — gives you the memory profile of DDP and the communication profile of stage 3. DeepSpeed decides granularity for you.
Both overlap their collectives with compute; both prefetch; both support mixed precision and activation checkpointing.What ships around the edges. DeepSpeed has CPU and NVMe offload, its own pipeline engine and MoE support. FSDP has DTensor, device meshes and torch.compile.
Both reach the same loss curve. Neither is an approximation.Mixed-precision control. FSDP's MixedPrecision policy sets the parameter, reduction and buffer dtypes separately — reducing gradients in fp32 while computing in bf16 is a one-line change.

So the decision doesn't turn on quality. Already in PyTorch and want one dependency and torch.compile: FSDP. Need to offload to CPU or NVMe, or want the pipeline and MoE machinery in the same package: DeepSpeed. Pick on operational fit and stop reading comparison benchmarks.

15Tensor parallelism: splitting one matrix multiply across GPUs

ZeRO shards where weights are stored. Tensor parallelism shards where they are used. That difference is invisible until a single layer stops fitting on one GPU, and then it is the only thing that matters.

Stage 3 all-gathers a layer's complete weight matrix onto every rank before that layer runs. If the complete matrix does not fit in one GPU's memory, stage 3 cannot save you. Tensor parallelism never materialises it anywhere at all.

The mechanism, on a transformer's MLP block. Two matrices: $Y = \text{GeLU}(XA)$, then $Z = YB$.

  • Split $A$ by columns across the ranks. GeLU is elementwise, so rank $i$ can compute $\text{GeLU}(XA_i)$ using nothing but its own columns. No communication.
  • Split $B$ by rows, matching. Rank $i$ holds exactly the rows that pair with its columns of $A$, and produces $\text{GeLU}(XA_i)B_i$ — a partial sum over part of the hidden dimension.
  • One all-reduce adds the partial sums into $Z$. Splitting column-then-row is what reduces the communication to exactly one collective per block.

Attention splits the same way, by head: the Q, K and V projections go column-parallel so each rank owns whole heads, and the output projection goes row-parallel, ending in one all-reduce. Megatron-LM's accounting for a full transformer layer is two all-reduces in the forward path and two in the backward path.

Try it Shard an MLP block four ways and check it against the unsharded answer
import numpy as np
rng = np.random.default_rng(0)

T, H, F, TP = 4, 8, 32, 4      # tokens, hidden, ffn width, TP degree
X = rng.normal(size=(T, H))
A = rng.normal(size=(H, F))    # first MLP matrix
C = rng.normal(size=(F, H))    # second MLP matrix
gelu = lambda t: t * 0.5 * (1 + np.tanh(
    0.7978845608 * (t + 0.044715 * t ** 3)))

one_gpu = gelu(X @ A) @ C

# Column-parallel on A: rank i owns F/TP columns. GeLU is
# elementwise, so nothing has to be communicated here.
Acol = np.split(A, TP, axis=1)
# Row-parallel on C: rank i owns the matching F/TP rows, so its
# output is a PARTIAL sum over the full hidden dimension.
Crow = np.split(C, TP, axis=0)

partial = [gelu(X @ Acol[i]) @ Crow[i] for i in range(TP)]
tensor_parallel = sum(partial)          # <- the one all-reduce

print("max |1 GPU - %d-way TP| = %.2e"
      % (TP, np.abs(one_gpu - tensor_parallel).max()))
print("weights held per rank: %d of %d floats"
      % (Acol[0].size + Crow[0].size, A.size + C.size))
print("all-reduce payload per rank: %d floats (T x H, the"
      " activation)" % partial[0].size)
print("note the payload does NOT shrink as TP grows -- every"
      " rank all-reduces the whole activation, every layer.")
The error is 3.55e-15, which is floating-point noise. Tensor parallelism is an exact rearrangement, not an approximation — the same sum, associated differently. Each rank holds a quarter of the weights, which is the point. The last line is the catch: the all-reduce payload is the activation, sized $b \times s \times h$, and it does not shrink as you add ranks. Four-way tensor parallelism does not move a quarter of the data. It moves all of it, four ways.

That last property is why tensor parallelism is a wiring decision before it is a modelling one. Put numbers on it. Take a 7B-shaped model: hidden size 4096, 32 layers, batch 8, sequence 2048, bf16.

$$S = b \cdot s \cdot h \cdot 2\ \text{bytes} = 8 \cdot 2048 \cdot 4096 \cdot 2 = 134\ \text{MB}$$
  • on the wire a ring all-reduce moves $2(N-1)/N \times S \approx 2S = 268$ MB per GPU, from section 11.
  • per layer four all-reduces — two forward, two backward — so $4 \times 268\ \text{MB} \approx 1.07$ GB.
  • per step 32 layers, so $32 \times 1.07 = 34.4$ GB per GPU, per training step.

Now divide that by a link. NVIDIA gives an H100 SXM 900 GB/s of aggregate NVLink bandwidth, which is 450 GB/s in each direction. An InfiniBand NDR port runs at 400 Gb/s, or 50 GB/s. The same 34.4 GB therefore takes 76 ms on NVLink and 687 ms across InfiniBand.

Nothing about the arithmetic changed between those two numbers. The matmuls are identical. Only the wire is different, and it cost 611 ms per step. That is the concrete version of the placement rule in section 09: tensor parallelism stays inside a node, and the tensor-parallel degree must not exceed the number of NVLink-connected GPUs in it.

The second reason to reach for tensor parallelism

Everything above is about capacity. The other reason is latency. Data parallelism makes a batch finish sooner; it does nothing for a single forward pass, because one pass still runs end to end on one GPU.

Tensor parallelism splits that single pass, so it is the standard way to cut time-to-first-token when serving a large model. Page 31 is where that side of it lives.

16Pipeline parallelism, and the bubble you pay for it

Tensor parallelism cannot cross a node boundary. Pipeline parallelism is what you use when the model still does not fit and you have run out of NVLink.

Cut the model by depth. Layers 1–8 on GPU 0, 9–16 on GPU 1, and so on. Now the only thing that crosses a stage boundary is the activation tensor at the cut — once per microbatch, in each direction — rather than four all-reduces per layer. That is what makes it survivable over InfiniBand.

The problem is equally obvious. Stage 1 cannot start until stage 0 finishes, so with one batch in flight exactly one GPU is busy and $p-1$ of them are idle. A four-stage pipeline running a single batch achieves 25% utilisation.

The fix is to keep more than one thing in flight: split the batch into $m$ microbatches and feed them in back to back. Stage 0 starts microbatch 1 as soon as it has handed microbatch 0 down. The pipeline still has to fill and drain, and that is the bubble.

$$\text{wall clock} = (m + p - 1)\,t, \qquad \text{ideal} = m\,t$$
  • p pipeline stages; m microbatches; t time for one microbatch on one stage.
  • bubble $(p-1)\,t$ of dead time, which is $\dfrac{p-1}{m+p-1}$ of the wall clock.
  • as overhead against the ideal $m\,t$, the same bubble is $\dfrac{p-1}{m}$ — the form the Megatron-LM cluster paper quotes.

Read the second form as the design rule: the bubble depends on the ratio of microbatches to stages, not on either alone. GPipe's own guidance is that the overhead becomes negligible once $m \ge 4p$.

A four-stage forward sweep with eight microbatches — every idle cell drawn
11 time steps × 4 stages = 44 cells · 32 busy · 12 idle · 27.3% time → 0 1 2 3 4 5 6 7 stage 0 0 1 2 3 4 5 6 7 stage 1 0 1 2 3 4 5 6 7 stage 2 0 1 2 3 4 5 6 7 stage 3 dashed = idle · 6 cells filling the pipe, 6 draining it · p(p−1) = 12
Solid cells are a stage working on the numbered microbatch; dashed cells are a GPU with nothing to do. Every stage has exactly three — a leading triangle while the pipe fills, a trailing one while it drains — so twelve in all, which is $p(p-1)$ for any $p$. Only the forward sweep is drawn. Raising $m$ from 8 to 32 leaves those twelve unchanged while the grid grows to 140.
Try it Build the schedule, count the idle cells, check the formula
def grid(p, m):
    """Forward sweep only: stage s runs microbatch k at step s+k."""
    steps = m + p - 1
    g = [["." for _ in range(steps)] for _ in range(p)]
    for s in range(p):
        for k in range(m):
            g[s][s + k] = str(k % 10)
    return g

print("  p    m  steps   idle/cells    measured   (p-1)/(m+p-1)")
for p, m in [(4, 1), (4, 4), (4, 8), (4, 32), (8, 8), (8, 64)]:
    g = grid(p, m)
    cells = p * len(g[0])
    idle = sum(row.count(".") for row in g)
    print("%3d %4d %6d   %4d/%-6d %8.1f%% %13.1f%%"
          % (p, m, len(g[0]), idle, cells,
             100 * idle / cells, 100 * (p - 1) / (m + p - 1)))

print("\np=4, m=8.  '.' is a stage with nothing to do:")
for s, row in enumerate(grid(4, 8)):
    print("  stage %d  %s" % (s, "".join(row)))
print("\nidle cells are always p(p-1) = %d: a triangle while the"
      " pipe fills and another while it drains." % (4 * 3))
The measured and formula columns agree on every row, which is the point of counting rather than trusting. The first row is the pathology: p=4, m=1 wastes 75% of the cluster. Read down, and at $m = 4p$ the waste is under 10% at both depths — GPipe's rule of thumb falling out of the arithmetic. The idle count stays 12 on every $p=4$ row: microbatches dilute a bubble, they never remove one.

So why not push $m$ to 1000? Two reasons. Smaller microbatches mean smaller matmuls, which run further from peak on hardware that wants big tiles. And the naive GPipe schedule holds the activations of every in-flight microbatch, so activation memory grows with $m$.

The standard answer to the second is the 1F1B schedule: once the pipe is full, alternate one forward with one backward, so a microbatch's activations are released as early as possible. Same bubble, activation memory capped at $p$ microbatches instead of $m$. Interleaving further subdivides each stage to shrink the bubble, at the cost of more boundary crossings.

The traffic comparison is what justifies the whole approach. Summed over all microbatches, a stage boundary carries one full batch of activations forward and one of gradients back. For the 7B-shaped model above, four stages and three boundaries: $6 \times 134\ \text{MB} = 805$ MB per step, against tensor parallelism's 34.4 GB. Forty times less traffic, which is exactly why one crosses nodes and the other does not.

17Choosing: N GPUs, a model of size M

Every strategy above costs something, so the procedure is to add them in increasing order of cost and stop at the first one that fits. Nobody reaches for 3D parallelism because it is elegant.

Work the ledger first. Multiply the parameter count by 16 bytes for mixed-precision AdamW, then ask the questions in this order.

The decision, in the order the constraints bind
16 × params fits on one GPU? plus activations yes DDP across all N one bucketed all-reduce/step cheapest. stop here. no 16Ψ/N fits, one layer fits? the usual case yes ZeRO / FSDP stage 2 first, stage 3 only if it OOMs stage 3 costs 1.5× traffic no fits inside one node? i.e. across NVLink only yes + tensor parallel degree ≤ GPUs per node 4 all-reduces per layer no still doesn't fit frontier-scale only + pipeline across nodes m ≥ 4p microbatches bubble = (p−1)/m Nesting order, innermost first: tensor (NVLink) → pipeline (node boundaries) → data / ZeRO (InfiniBand)
Each "no" costs something the row above did not: ZeRO gives up the free local optimizer step, tensor parallelism gives up the node boundary, pipelining gives up a fraction of every GPU to the bubble. The bottom line is the placement rule from section 09 restated as a nesting order — the most talkative strategy gets the fastest wire, and the least talkative gets the slowest.

Activation checkpointing sits outside this ladder entirely and is almost always on. It costs roughly an extra forward pass of compute and removes most of the activation memory, which frequently turns a configuration that OOMs at step one into a configuration that runs.

Four concrete cases, on one node of 8 × 80 GB

ModelLedger at 16 bytes/paramChoiceWhy
1B16 GBDDP × 8Fits on one card with room for activations. Anything fancier is pure overhead.
7B112 GBFSDP SHARD_GRAD_OPOver 80 GB, so replication is out. Stage 2 gives $2\Psi + 14\Psi/8 = 26$ GB per card at no extra traffic.
70B1120 GBStage 3 across 2–4 nodes140 GB per card on 8, so one node cannot do it. The largest single weight is well under a card, so tensor parallelism is not required.
400B+6400 GBTensor × pipeline × dataStage 3's per-layer all-gathers become the bottleneck at this scale. Tensor parallel inside the node, pipeline across nodes, data parallel on top.

Every row was decided by a memory number compared with a capacity number, and then by a bandwidth number. The model family, the framework and the task never came into it.

18Why this hardware, and why this algorithm

Why SIMT hardware fits dense matrix math specifically: a matrix multiply's output elements are all independent of each other, computed by the same operation (multiply, accumulate, repeat) with no data-dependent branching anywhere in the computation.

That is the one constraint SIMT hardware is built around: thousands of simple cores, no branch-prediction machinery, all executing the same instruction across different data every cycle.

A CPU's expensive control-flow logic exists to handle the unpredictable, data-dependent branching that dense matrix math doesn't have. For this workload a CPU pays for capability it never uses, while a GPU spends its whole chip budget on the capability the workload needs.

Why ring all-reduce beats a naive centralized parameter server: a star topology, with every worker talking to one central server, funnels all N workers' traffic through a single network link.

That link has a fixed bandwidth ceiling, so the time to synchronize grows with N. The more GPUs you add to speed up training, the worse that one link's congestion gets, which is the opposite of what adding hardware is meant to do.

Ring all-reduce spreads the communication load evenly across every link in the ring instead of concentrating it on one node, so there is no single point where traffic piles up.

It isn't entirely free: with more GPUs in the ring there are more steps to complete, since scatter-reduce and all-gather each take N−1 hops. But the resource that was the bottleneck, the bytes any single GPU pushes through its own link, stays close to constant. That is the trade that lets large-scale data-parallel gradient synchronization scale gracefully instead of collapsing under its own communication cost.

19Tradeoffs, and what breaks

NVLinkInfiniBand / RoCEEthernet
BandwidthVery high — roughly hundreds of GB/s (approximate)High, but roughly an order of magnitude or more below NVLink (approximate)Lowest of the three, general-purpose
LatencyLowestLow — purpose-built for HPC / RDMAHigher, more protocol overhead
Typical scopeBetween GPUs on the same nodeBetween nodes in a training clusterGeneral networking
What runs over itTensor-parallel / expert-parallel traffic — every layerData-parallel gradient sync, pipeline-parallel activations — once per step / stageCluster management, logging — not the primary training path

Three failure modes are worth naming, because they show up in real kernels and real clusters and not just in theory.

  • Warp divergence from branch-heavy code silently halves or worse a kernel's effective throughput. Every diverging branch means the warp pays for both paths while only ever benefiting from one, and it is easy to write correct CUDA or Triton code that diverges without realizing it.
  • Low occupancy from register or shared-memory pressure is a quieter failure. The kernel still produces correct results, it is just slower than it should be, because the SM doesn't have enough resident warps to hide memory latency behind. Cores sit idle waiting on HBM instead of doing useful work for another warp.
  • Communication becoming the bottleneck when interconnect topology doesn't match your parallelism strategy is a cluster-design failure, not a kernel one. Put tensor-parallel shards across InfiniBand instead of NVLink, or fail to keep pipeline-parallel stage boundaries aligned with node boundaries, and you can end up with GPUs that are individually fast sitting mostly idle waiting on the network, which defeats the entire point of adding more of them.

And the ones that only appear above one node

  • ZeRO stage 3 on a slow fabric. Memory stops being the problem and throughput becomes it. The all-gathers sit in front of every layer, so when they cannot be prefetched fast enough the GPUs wait on weights instead of computing. Stage 2 is often the faster configuration that still fits.
  • An unbalanced pipeline. The slowest stage sets the rate for all of them. One extra layer on stage 0, or an embedding table that only lives on the first and last stage, and every other GPU idles for the difference on every microbatch.
  • Tensor-parallel degree larger than the node. Nothing errors. The collectives start crossing PCIe or InfiniBand, and the arithmetic in section 15 turns 76 ms per step into 687 ms.
  • One straggler rank. A collective is a barrier, so the slowest rank sets the step time for the entire job. A single GPU thermally throttling, or one node on a congested link, slows all 512 of them by the same fraction.
  • A parameter that receives no gradient. Its bucket never completes and the all-reduce hangs with no error. find_unused_parameters=True fixes it by scanning the graph every iteration, which is a real per-step cost — better to fix the model.
  • The learning rate nobody re-tuned. Effective batch size is per-rank batch times world size. Doubling the cluster doubles the batch and quietly changes the optimisation problem, and the run that diverged was a hyperparameter bug, not a systems one.

20Build this

You can read about coalesced memory access and still not believe layout beats arithmetic. Write one kernel that does no arithmetic at all, and the argument settles itself.

Project Write a naive transpose kernel, then profile your way out of it ~4 hours · CUDA C + Nsight Compute

Transpose a square matrix. The task is chosen for one property: a transpose performs zero floating-point operations. It only moves bytes. So any distance between what you measure and your GPU's bandwidth roof has to come from how those bytes are laid out. There is no arithmetic there to blame. You need an NVIDIA GPU for this. Colab's free tier is fine, and ncu runs there.

  1. Write the obvious kernel in CUDA C: one thread per element, out[x*N + y] = in[y*N + x]. Check it against NumPy before you time anything. A fast wrong kernel teaches nothing.
  2. Compute effective bandwidth by hand: bytes read plus bytes written, divided by measured time. Put it next to the peak bandwidth on your GPU's datasheet.
  3. Profile with ncu and go straight to the memory chart. Compare sectors requested per warp for the load against the same figure for the store. One of them is much worse than the other.
  4. Apply the single optimisation. Stage a tile in __shared__ memory: coalesced loads in, transpose inside the tile, coalesced stores out. Identical bytes, identical zero FLOPs, different order.
  5. Re-measure bandwidth and re-run ncu. Put the two sector counts side by side with the ones you started from.
  6. Now break the fix while keeping the shared memory. Leave the __shared__ tile in place, but restore the strided global writes so the tile buys you nothing. Measure a third time.
You'll know it worked when the two sector-per-warp figures, badly mismatched in the naive kernel, come out matched after tiling. Your measured bandwidth moves toward the datasheet roof by roughly the amount the profiler said it should. Track the gap you closed against the roof rather than the raw timings. A transpose computes nothing, so that gap was never arithmetic. It was warps asking HBM for scattered addresses when they could have asked for contiguous ones.
What the breakage teaches. Shared memory on its own is not the optimisation. Keep the tile and reintroduce the strided writes, and the win largely evaporates, because the lever was always the global access pattern that section 08 lists first. Shared memory is only the scratch space that lets you reorder before touching HBM. Once you are convinced, pad the shared tile by one extra column and measure again. Threads reading down a column of an unpadded tile land in the same bank, which is the shared-memory analogue of the problem you just fixed one tier down.

Where this runs in production

Say you are laying out a training job across a cluster with several GPUs per node, connected by NVLink within a node and InfiniBand between nodes. The placement decision follows directly from the communication-frequency argument on this page.

  • Tensor parallelism needs to communicate within every layer, so it goes on the GPUs connected by NVLink. It stays confined inside a single node, sized to however many GPUs that node has.
  • Pipeline parallelism only needs to pass activations at stage boundaries, so it can cross the slower InfiniBand link between nodes without that dominating step time.
  • Data parallelism replicates that whole tensor-parallel-plus-pipeline-parallel unit and synchronizes gradients once per step via ring all-reduce. It scales across as many of those units as you have, tolerating the InfiniBand link because ring all-reduce keeps each individual GPU's communication burden bounded regardless of how many replicas you add.

Get this placement backwards, say by splitting a single layer's weights tensor-parallel-style across an InfiniBand link, and the cluster's raw compute capacity stops mattering. The GPUs spend most of their time waiting on the network rather than computing.

21Interview questions

BeginnerWhy are GPUs so much faster than CPUs for deep learning?

A GPU uses SIMT architecture — thousands of simple cores executing the same instruction across different data simultaneously — instead of a CPU's handful of complex cores built with extensive branch-prediction and control-flow hardware. Deep learning is overwhelmingly dense matrix multiplication and elementwise operations: every output element is computed independently, by the same operation, with essentially no data-dependent branching. That's an embarrassingly parallel workload with no need for the control-flow machinery a CPU spends most of its silicon on, and it's exactly the shape SIMT hardware was built to execute — thousands of cores doing the same multiply-accumulate at once, rather than a few cores each handling their own complicated, branchy sequence of instructions.

BeginnerWhat is a warp, and what happens when threads inside one diverge?

A warp is a group of 32 threads that the GPU always schedules and executes together, running the identical instruction at the identical time on different data — the actual physical unit of SIMT execution, one level below the thread block. Warp divergence happens when threads within a warp need to take different branches of an if/else. Since a warp can only run one instruction stream across all its lanes at a time, the hardware serializes the branches: it runs the "if" path across the whole warp with the threads that shouldn't take it masked off, then runs the "else" path across the whole warp with the other threads masked off. Every thread sits through both passes, but only half the work in each pass is useful — which is why branch-heavy, data-dependent code is slow on a GPU.

IntermediateWhat is Triton, and why would you reach for it instead of hand-written CUDA?

Triton is a Python-embedded language and compiler for writing GPU kernels at the level of tiles or blocks of data rather than individual threads. Writing genuinely efficient CUDA by hand requires deep, specialized expertise in coalesced memory access, register and shared-memory budgeting for occupancy, and manual thread-indexing math — a real cost most ML researchers, who think in tensors and layers rather than threads and warps, don't have time to pay for every new idea. Triton's compiler automatically handles that low-level thread scheduling, memory coalescing, and shared-memory allocation, getting most of hand-tuned CUDA's performance for a fraction of the development time. FlashAttention's efficient reference implementation is written in Triton, and a large share of PyTorch's own generated custom kernels are Triton kernels, for exactly this reason.

IntermediateWhat's the practical difference between NVLink and InfiniBand, and why does it drive where you place different parallelism strategies?

NVLink connects GPUs on the same physical node at very high bandwidth, built specifically for direct GPU-to-GPU communication. InfiniBand (or RoCE) connects separate nodes — still fast and purpose-built for HPC clusters, but roughly an order of magnitude or more below NVLink's bandwidth. Tensor parallelism communicates within every single layer, many times per forward and backward pass, so it needs to live on NVLink-connected GPUs inside one node — putting it across InfiniBand would make communication dominate almost immediately. Data parallelism only synchronizes gradients once per step, so it can tolerate sitting across the slower InfiniBand link between nodes without that link becoming the bottleneck.

DeepExplain ring all-reduce, and why it scales better than a naive centralized approach.

Arrange N GPUs in a logical ring, split each GPU's gradient data into N chunks, and run two phases: scatter-reduce, N−1 steps, where each GPU sends one chunk to its ring-neighbor and accumulates a chunk received from its other neighbor, ending with every GPU holding the true sum for exactly one chunk; then all-gather, another N−1 steps, relaying each GPU's already-final chunk around the ring until every GPU has every chunk. At every step, each GPU only ever sends and receives one chunk — a bounded amount of data, independent of N. Total data moved per GPU works out to roughly 2(N−1)/N times the data size, which approaches a constant (2×) as N grows rather than scaling with cluster size. A naive centralized parameter server, by contrast, funnels every worker's traffic through one server's network link, whose bandwidth ceiling is fixed — so synchronization time grows linearly with the number of workers, becoming the exact bottleneck you're trying to scale past by adding more GPUs.

IntermediateWalk through exactly what DistributedDataParallel does in one training step.

Once at startup, rank 0 broadcasts its weights so every rank begins with byte-identical parameters and optimizer state. Then each step: every rank runs a forward pass on its own shard of the global batch with no communication at all, and computes a local loss. The backward pass is also local arithmetic, but as each parameter's gradient is produced it is appended to a bucket — PyTorch defaults to 25 MiB — and the moment a bucket fills, its all-reduce is launched asynchronously. Because backward runs from the last layer to the first, the last layers' gradients go on the wire while the earlier layers are still computing, which is what overlaps communication with compute. At the end of backward, DDP waits on the outstanding collectives; every rank now holds the same averaged gradient bit for bit. The optimizer step then runs purely locally on each rank. Weights never cross the wire after startup: identical inputs plus an identical gradient produce identical weights, so synchronisation is a consequence rather than an operation.

DeepWhat does each ZeRO stage shard, and what does each cost in communication?

Mixed-precision Adam costs 16 bytes per parameter: 2 for the fp16 weights, 2 for the fp16 gradients, and 12 for the fp32 master copy plus Adam's two moments. Stage 1 shards only that 12-byte optimizer tier, leaving 4Ψ + 12Ψ/N per rank. Stage 2 also shards gradients, reduce-scattering each bucket during backward instead of all-reducing it, leaving 2Ψ + 14Ψ/N. Stage 3 shards the parameters too, all-gathering each layer's weights just before it is used and dropping them afterwards, leaving 16Ψ/N. The communication answer surprises people: stages 1 and 2 are free. A ring all-reduce already is a reduce-scatter plus an all-gather, 2Ψ of traffic, and those stages just stop in between to update the slice each rank owns. Stage 3 costs 3Ψ — 1.5× — because the weights must be gathered in the forward pass and again in the backward pass. Worse than the ratio suggests, those gathers are on the critical path in front of every layer rather than hidden behind the backward pass, so on a slow interconnect stage 2 is frequently faster while stage 3 is the only one that fits.

IntermediateA four-stage pipeline with eight microbatches — how much of the cluster is idle, and what would you change?

With p stages and m microbatches, the schedule takes m + p − 1 time steps against an ideal of m, so the bubble is (p − 1) steps. As a fraction of wall clock that is (p−1)/(m+p−1) = 3/11 ≈ 27%; expressed as overhead against the ideal it is (p−1)/m = 3/8, which is the form the Megatron-LM paper quotes. The absolute number of idle cells is p(p−1) = 12 and it never changes — raising m dilutes the bubble rather than removing it. So the fix is more microbatches: GPipe's guidance is that the overhead becomes negligible around m ≥ 4p, which here means 16 or more. The limits on pushing m higher are that smaller microbatches run the matmuls further from peak, and that the naive schedule holds activations for every in-flight microbatch. A 1F1B schedule fixes the second by alternating forward and backward once the pipe is full, capping activation memory at p microbatches for the same bubble.

DeepFSDP or DeepSpeed ZeRO — what is genuinely different and what is just naming?

The algorithm is the same, and most of the apparent difference is vocabulary. PyTorch's own docs map FULL_SHARD onto ZeRO stage 3 and SHARD_GRAD_OP onto stage 2, with NO_SHARD being plain DDP; "shard" versus "partition" and "FSDP unit" versus "parameter group" name the same objects. Neither is an approximation, so they reach the same loss curve and the memory and traffic arithmetic is identical. What is actually different is operational. FSDP lives in PyTorch core and composes with DTensor, device meshes and torch.compile; DeepSpeed is a separate engine with its own config file, launcher and checkpoint format, and ships things FSDP does not — CPU and NVMe offload, a pipeline engine, MoE support. Sharding granularity differs too: FSDP shards whatever you wrap, so a careless auto_wrap_policy that makes the whole model one unit gives you DDP's memory profile with stage 3's communication, which is the classic FSDP performance bug. Choose on dependencies and on whether you need offload, not on expected quality.

DeepYou're handed a cluster with several GPUs per node on NVLink, and InfiniBand between nodes. How do you lay out tensor, pipeline, and data parallelism across it, and why?

Tensor parallelism goes inside a node, on the NVLink-connected GPUs, because it communicates partial results within every layer — many times per forward and backward pass — and needs the fastest possible link or that communication dominates the compute. Pipeline parallelism can cross the InfiniBand link between nodes, since it only passes activations at stage boundaries, which is far less frequent and far more tolerant of a slower link. Data parallelism replicates the whole tensor-parallel-plus-pipeline-parallel unit across as many groups of nodes as you have, synchronizing gradients once per step via ring all-reduce over InfiniBand — tolerable specifically because ring all-reduce keeps each GPU's communication burden bounded regardless of how many replicas you add. Getting this backwards — for instance splitting a single layer tensor-parallel-style across InfiniBand — turns the network into the bottleneck and leaves the GPUs' raw compute capacity mostly unused, since they'd spend most of their time waiting on data that has to cross the slower link on every layer.

22Go deeper

●Now write it yourself

Reading the derivation and being able to produce it are different skills. These are Deep-ML problems that exercise what this page covers — each one is checked against real test cases, not multiple choice.

Matched to this page from Deep-ML's catalogue of 1,380 problems. More at deep-ml.com, and Where to practise covers the other platforms and what each one trains.

My Notes — 28 GPU Architecture, CUDA & Distributed Training

Free notes

Highlights on this page