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 1Diagram — 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 queue2. 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#
- IX.01 — Why distribute
- IX.03 — Tensor parallelism
- IX.07 — Collectives and NCCL
- IX.09 — Communication cost math
- IX.12 — When more GPUs hurt
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, samplingImplementation 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#
- Shard the weights offline from a normal checkpoint. Assert that concatenating the shards reproduces the original.
- TP MLP. Verify
partial_0 + partial_1 == fullwithin float tolerance. - TP attention. Heads split; each rank keeps a KV cache for its own heads only.
- Full TP forward and generation. Tokens identical to the single-device model.
- Real collectives (v2).
torch.distributedwith the NCCL backend. Sampling must agree across ranks: same seed, or sample on rank 0 and broadcast. - Time breakdown. Per decode step: compute vs all-reduce vs everything else.
- 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#
| Measurement | Expectation to write down first |
|---|---|
| Decode ITL: 1 GPU vs TP=2 | Less than 2× better |
| Fraction of step time in all-reduce, batch 1 | Large — tiny payload, fixed latency dominates |
| Same at batch 64 | Smaller fraction |
| Throughput: TP=2 vs 2 replicas, model fits on one GPU | Replicas win |
| Max concurrent sequences: TP=2 vs 2 replicas | Similar total KV; different failure modes |
| PP=2 ITL and throughput | ITL no better than 1 GPU; one hop per token |
| Predicted comm bytes/step vs measured | Should match the formula |
| All-reduce latency for a tiny message | The 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-testsall-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 2on the same model.
12. Interview questions this project answers#
- How is an MLP block split under tensor parallelism, and why column-then-row?
- How many all-reduces per layer, and how many bytes per decode step?
- Model fits on one GPU; you have two. TP=2 or two replicas? Justify.
- Why does TP scale poorly at small batch?
- Why is pipeline parallelism rarely used for latency-sensitive inference?