Below the API

Project 13 — Multi-GPU Inference

Expert 12h Difficulty 4/5

Prerequisites Projects 04-05, 11; Section IX (01-03, 07-09, 12)

Split one model across two devices by hand — tensor parallelism — and find out when it is worse than two independent copies.


1. What you build#

A tensor-parallel (TP) version of your transformer: each device holds a shard of every layer’s weights, computes a partial result, and the shards are combined with an all-reduce. Then a measured comparison of three ways to use two GPUs:

A. TP=2          one model sharded across both GPUs
B. 2 replicas    two full copies, requests load-balanced   (data parallelism)
C. PP=2          first half of the layers on GPU 0, second half on GPU 1

Diagram — Three ways to use two GPUs#

flowchart TB
  subgraph A["A. TP=2 - one model, sharded"]
    direction LR
    A0["GPU 0: half of every layer"] <-->|"2 all-reduces per layer"| A1["GPU 1: other half"]
  end
  subgraph B["B. Two replicas"]
    direction LR
    LB["Load balancer"] --> B0["GPU 0: full model"]
    LB --> B1["GPU 1: full model"]
  end
  subgraph C["C. PP=2 - layers split"]
    direction LR
    C0["GPU 0: first half of layers"] -->|"activations"| C1["GPU 1: second half"]
  end

  class A0,A1 memory
  class B0,B1 compute
  class C0,C1 io
  class LB queue

2. Why it matters#

“Add GPUs” is the most expensive default in the industry. TP exists to fit models that do not fit on one device and to cut per-token latency; it does not double throughput, and for models that already fit, replicas usually win. You should be able to show that with your own measurements and the communication arithmetic that explains it.


3. Read first#


4. Spec#

The Megatron layout, which needs exactly two all-reduces per layer:

MLP        W1 split by COLUMNS (output dim)  -> each rank computes half the hidden units
           activation applied locally
           W2 split by ROWS (input dim)      -> each rank produces a partial output
           ALL-REDUCE (sum)                  -> full MLP output on every rank

Attention  heads split across ranks: rank r owns heads [r·H/N, (r+1)·H/N)
           Q, K, V projections split by columns (per head) ; KV cache is sharded too
           output projection split by rows
           ALL-REDUCE (sum)

Replicated on every rank: embeddings, LayerNorm, residual stream, sampling

Implementation path:

v1  one process, two devices, explicit .to() and add   (understand the math)
v2  two processes, torch.distributed with NCCL all_reduce   (the real thing)

5. Milestones#

  1. Shard the weights offline from a normal checkpoint. Assert that concatenating the shards reproduces the original.
  2. TP MLP. Verify partial_0 + partial_1 == full within float tolerance.
  3. TP attention. Heads split; each rank keeps a KV cache for its own heads only.
  4. Full TP forward and generation. Tokens identical to the single-device model.
  5. Real collectives (v2). torch.distributed with the NCCL backend. Sampling must agree across ranks: same seed, or sample on rank 0 and broadcast.
  6. Time breakdown. Per decode step: compute vs all-reduce vs everything else.
  7. The comparison. A vs B vs C on throughput, ITL, TTFT, and max concurrency at fixed total memory.

6. Starter skeleton#

def shard_mlp(W1, b1, W2, n):               # W1: (hidden, d)  W2: (d, hidden)  — nn.Linear layout
    return [(w1, bb, w2) for w1, bb, w2 in
            zip(W1.chunk(n, dim=0), b1.chunk(n, dim=0), W2.chunk(n, dim=1))]

def tp_mlp_rank(x, w1, b1, w2):             # runs on ONE rank
    h = F.gelu(x @ w1.T + b1)               # this rank's slice of the hidden units
    return h @ w2.T                         # PARTIAL output, shape (B, T, d)

def tp_block(x, rank_params, dist):
    a = tp_attention_rank(ln1(x), rank_params.attn)     # this rank's heads only
    dist.all_reduce(a, op=dist.ReduceOp.SUM)            # all-reduce #1
    x = x + a + rank_params.attn_bias_once              # add shared bias exactly once
    m = tp_mlp_rank(ln2(x), *rank_params.mlp)
    dist.all_reduce(m, op=dist.ReduceOp.SUM)            # all-reduce #2
    return x + m + rank_params.mlp_bias_once

# Bytes on the wire per decode step:
#   2 all-reduces × n_layer × (batch × d_model × bytes_per_elem)

7. What to measure#

MeasurementExpectation to write down first
Decode ITL: 1 GPU vs TP=2Less than 2× better
Fraction of step time in all-reduce, batch 1Large — tiny payload, fixed latency dominates
Same at batch 64Smaller fraction
Throughput: TP=2 vs 2 replicas, model fits on one GPUReplicas win
Max concurrent sequences: TP=2 vs 2 replicasSimilar total KV; different failure modes
PP=2 ITL and throughputITL no better than 1 GPU; one hop per token
Predicted comm bytes/step vs measuredShould match the formula
All-reduce latency for a tiny messageThe floor that payload size cannot fix

8. Done when#

  • TP generation is token-identical to single-device.
  • You have the per-step time breakdown and can state the all-reduce fraction.
  • You have the A/B/C table and a one-paragraph recommendation for when to use each.
  • You can derive “two all-reduces per layer” from the layout on a whiteboard.
  • You can explain why TP wants NVLink and why PP tolerates slower links.

9. Common pitfalls#

Adding the bias on every rank. After the all-reduce it is counted N times. Add it on one rank, or after the reduce.

Ranks sampling different tokens. Generation diverges silently. Synchronize the choice.

Splitting W2 by columns instead of rows. The shapes work; the math doesn’t.

Head count not divisible by TP degree. With GQA, the KV heads must divide too.

Measuring TP on two GPUs connected only by PCIe and concluding TP is useless. Record the interconnect; it is the result.

Comparing TP=2 to a single GPU on throughput only. TP’s case is fit and latency.


10. No GPUs (or one)? Do this instead#

Run v2 with two CPU processes and the gloo backend. Correctness and the communication count are identical. To make the cost visible, inject a delay per all-reduce (e.g. 50 µs + bytes / 10 GB/s) and reproduce the A/B/C comparison in simulation. Then check your model against published TP scaling numbers.


11. Stretch goals#

  • TP=4 and the scaling curve; find where efficiency drops below 50%.
  • Run nccl-tests all-reduce across sizes and overlay it on your step-time model.
  • Combine TP=2 inside a node with replicas across nodes — the common production topology.
  • Compare your TP against vLLM --tensor-parallel-size 2 on the same model.

12. Interview questions this project answers#

  1. How is an MLP block split under tensor parallelism, and why column-then-row?
  2. How many all-reduces per layer, and how many bytes per decode step?
  3. Model fits on one GPU; you have two. TP=2 or two replicas? Justify.
  4. Why does TP scale poorly at small batch?
  5. Why is pipeline parallelism rarely used for latency-sensitive inference?

13. Next#

Project 14 — Distributed inference

↑↓ navigate ↵ open