41. Many GPUs: tensor, pipeline, expert and data parallelism
In this chapter
- Why one GPU stops being enough, and the four ways to split serving across many: tensor, pipeline, expert and data parallelism, each with its communication bill.
- Tensor parallelism Megatron-style: column- and row-parallel layers, one all-reduce per sublayer, split KV heads and a split vocabulary.
- One scheduler driving many ranks in lockstep, as vLLM and SGLang do.
- Expert parallelism with all-to-all dispatch, a prefix-aware router for data-parallel replicas, prefill and decode on different machines with the KV cache shipped between them, and ring attention for contexts too long for one GPU.
You will build
engine/parallel.py: the parallel layers, shard_tensor_parallel, PipelineStage.forward, ExpertParallelMoE.forward, worker_loop, PrefixRouter.route, KVExporter.export, import_kv and ring_attention.
Time: 8-10 hours. GPU: not needed: every test runs several processes on the CPU with PyTorch's gloo backend. Two or more GPUs with NCCL run the same code.
When one GPU isn’t enough
There are three reasons to spread a model over several GPUs, and they lead to different designs:
- It doesn’t fit. Llama-3.1-70B in BF16 is 141 GB of weights; an H100 has 80 GB. Even when the weights fit, the KV cache needs room: every gigabyte left over is more concurrent requests (Chapter 31).
- It’s too slow. Decode reads every weight once per step (Chapter 10). Split the weights over 8 GPUs and each reads an eighth: the step takes (nearly) an eighth of the time, because the GPUs’ memory bandwidths add up.
- There’s too much traffic. One replica serves at most so many tokens per second. More replicas serve more.
The four kinds of parallelism split different things:
| splits | communication per layer | good for | |
|---|---|---|---|
| tensor (TP) | every weight matrix, across GPUs | 2 all-reduces of the activations | latency, fitting big models inside one NVLink domain |
| pipeline (PP) | the layers, into stages | one hidden vector per token per stage boundary | fitting models across nodes with slower links |
| expert (EP) | a MoE’s experts | 2 all-to-alls (or one all-reduce) | big MoEs: DeepSeek-V3, Qwen3-235B, Kimi K2 |
| data (DP) | the requests, over whole replicas | none (a router in front) | throughput |
Production deployments combine them: DeepSeek-V3 across 18 nodes runs attention data-parallel and experts expert-parallel; a Llama-405B deployment runs TP 8 inside each node and PP 2 across two. vLLM and SGLang expose all four with flags (--tensor-parallel-size, --pipeline-parallel-size, --enable-expert-parallel, --data-parallel-size); this chapter builds each one small enough to read.
The cost of talking
GPUs in one server talk over NVLink (900 GB/s per H100, each direction); servers talk over InfiniBand or RoCE (400 Gb/s = 50 GB/s per NIC, usually one NIC per GPU); consumer GPUs talk over PCIe (32-64 GB/s), often through the CPU. A ring all-reduce of $S$ bytes over $n$ GPUs sends $2\frac{n-1}{n}S$ bytes from each GPU and takes $2(n-1)$ latency-bound hops. For a decode step those messages are small (64 tokens × 8,192 features × 2 bytes = 1 MB), so latency, not bandwidth, sets their cost: a few microseconds per all-reduce on NVLink, which is why vLLM ships a custom one-shot all-reduce for small messages instead of NCCL’s ring, and why TP across PCIe or between nodes is slow.
The communicator
Four collectives cover the chapter: all-reduce (everyone gets the sum), all-gather (everyone gets everyone’s pieces), all-to-all (rank $i$ sends piece $j$ to rank $j$), and point-to-point send/receive.
class Comm:
"""The process group's collectives. With one process (or none initialized) each is a no-op."""
def __init__(self):
self.enabled = dist.is_available() and dist.is_initialized()
self.rank = dist.get_rank() if self.enabled else 0
self.world = dist.get_world_size() if self.enabled else 1
self.nccl = self.enabled and dist.get_backend() == "nccl"
self.traffic = {} # op -> [calls, bytes sent by this rank]
def count(self, op, t):
entry = self.traffic.setdefault(op, [0, 0])
entry[0] += 1
entry[1] += t.numel() * t.element_size()
def all_reduce(self, t, op=None):
"""Every rank ends with the elementwise sum (in place)."""
if self.world > 1:
self.count("all_reduce", t)
dist.all_reduce(t, op=op or dist.ReduceOp.SUM)
return t
def all_gather(self, t, dim=-1):
"""Every rank's tensor, concatenated along dim, on every rank."""
if self.world == 1:
return t
self.count("all_gather", t)
parts = [torch.empty_like(t) for _ in range(self.world)]
dist.all_gather(parts, t.contiguous())
return torch.cat(parts, dim=dim)
def all_to_all(self, chunks):
"""chunks[p] goes to rank p; returns the chunks every rank sent to this one. Row counts may
differ, so they are exchanged first. NCCL has a native all-to-all; gloo gets the same
result from pairwise sends and receives."""
if self.world == 1:
return [chunks[0]]
sizes = torch.tensor([c.shape[0] for c in chunks], dtype=torch.int64, device=chunks[0].device)
table = self.all_gather(sizes[None], dim=0) # table[src, dst] = rows src sends dst
tail, dtype = chunks[0].shape[1:], chunks[0].dtype
received = [torch.empty((int(table[src, self.rank]), *tail), dtype=dtype, device=chunks[0].device)
for src in range(self.world)]
if self.nccl:
dist.all_to_all(received, [c.contiguous() for c in chunks])
return received
received[self.rank] = chunks[self.rank]
ops = []
for peer in range(self.world):
if peer != self.rank:
ops.append(dist.isend(chunks[peer].contiguous(), peer))
ops.append(dist.irecv(received[peer], peer))
for op in ops:
op.wait()
return received
def broadcast_object(self, obj=None, src=0):
if self.world == 1:
return obj
box = [obj]
dist.broadcast_object_list(box, src)
return box[0]
def min(self, value):
t = torch.tensor([value], dtype=torch.int64, device="cuda" if self.nccl else "cpu")
return int(self.all_reduce(t, dist.ReduceOp.MIN).item())
def send(self, t, dst):
self.count("send", t)
dist.send(t.contiguous(), dst)
def recv(self, shape, dtype, src, device):
t = torch.empty(shape, dtype=dtype, device=device)
dist.recv(t, src)
return t
torch.distributed provides all of them on two backends: NCCL for GPUs and gloo for CPUs. Gloo lacks an all-to-all with uneven pieces, so all_to_all builds one from paired sends and receives. traffic counts calls and bytes, which the measurements below use.
spawn(fn, world) starts world processes on this machine, joined by one process group, the way torchrun does across machines; every test in this chapter runs through it.
Tensor parallelism
Megatron-LM’s observation (Shoeybi et al., 2019) is that a transformer’s matmuls come in pairs that split without communication in between. Take an MLP, $y = W_{down},\sigma(W_{up},x)$. Split $W_{up}$ by output rows across $n$ GPUs (column parallelism, after the column blocks of $W^T$): GPU $i$ computes features $[iI/n, (i+1)I/n)$ of the intermediate activation, needing all of $x$ but nothing from the other GPUs. The activation is elementwise, so each GPU applies it to its own slice. Then split $W_{down}$ by input columns (row parallelism): GPU $i$ multiplies its slice of the activation by its columns, producing a full-width partial sum. One all-reduce adds the partial sums, and every GPU holds $y$.
Attention splits the same way, by heads: Q, K and V are column-parallel (each GPU computes its heads), attention runs per head with no communication, and $W_o$ is row-parallel. So each layer costs two all-reduces, one after attention and one after the MLP, and every GPU holds $1/n$ of the weights and, because it only computes its own heads, $1/n$ of the KV cache.
@torch.no_grad()
def take_rows(linear, start, end):
"""Column parallelism: this rank computes output features [start, end) (no communication)."""
out = nn.Linear(linear.in_features, end - start, bias=linear.bias is not None,
device=linear.weight.device, dtype=linear.weight.dtype)
out.weight.copy_(linear.weight[start:end])
if linear.bias is not None:
out.bias.copy_(linear.bias[start:end])
return out
@torch.no_grad()
def take_cols(linear, start, end):
"""The matching row-parallel half: input features [start, end), a PARTIAL sum of the output."""
out = nn.Linear(end - start, linear.out_features, bias=False, device=linear.weight.device, dtype=linear.weight.dtype)
out.weight.copy_(linear.weight[:, start:end])
return out
class RowParallelLinear(nn.Module):
"""Partial products summed across ranks: the one all-reduce per sublayer."""
def __init__(self, linear, start, end, comm):
super().__init__()
self.local, self.comm = take_cols(linear, start, end), comm
self.bias = None if linear.bias is None else nn.Parameter(linear.bias.detach().clone())
def forward(self, x):
"""(Your engine: Chapter 41)"""
y = self.comm.all_reduce(self.local(x))
return y if self.bias is None else y + self.bias # added once, after the sum
def vocab_range(vocab, comm):
per = -(-vocab // comm.world) # padded: every rank holds `per` rows
return per, comm.rank * per, min(vocab, (comm.rank + 1) * per)
def vocab_slice(weight, comm):
per, lo, hi = vocab_range(weight.shape[0], comm)
local = weight.new_zeros((per, weight.shape[1]))
local[:hi - lo] = weight[lo:hi]
return nn.Parameter(local)
class VocabParallelEmbedding(nn.Module):
"""Each rank holds a slice of the table; tokens outside it look up zeros, and the all-reduce
assembles every row."""
def __init__(self, weight, vocab, comm):
super().__init__()
self.weight, self.comm = weight, comm # weight: this rank's slice (vocab_slice)
_, self.lo, self.hi = vocab_range(vocab, comm)
def forward(self, ids):
"""(Your engine: Chapter 41)"""
mine = (ids >= self.lo) & (ids < self.hi)
rows = F.embedding(torch.where(mine, ids - self.lo, 0), self.weight)
return self.comm.all_reduce(rows * mine[..., None].to(rows.dtype))
class ParallelLMHead(nn.Module):
"""Each rank scores its vocabulary slice; an all-gather assembles the full logits."""
def __init__(self, weight, vocab, comm):
super().__init__()
self.weight, self.vocab, self.comm = weight, vocab, comm
def forward(self, h):
"""(Your engine: Chapter 41)"""
return self.comm.all_gather(F.linear(h, self.weight), dim=-1)[..., :self.vocab]
Three details:
- The bias of a row-parallel layer is added once, after the all-reduce; added before, it would be counted $n$ times.
- The vocabulary is split too. Qwen3’s 151,936 × 4,096 embedding is 1.2 GB in BF16, worth splitting. Each rank holds a slice of rows; a token outside the slice looks up zeros, and the all-reduce assembles every token’s row. The LM head scores each rank’s slice and an all-gather builds the full logits. (vLLM all-gathers to one rank only, the one that samples; we gather everywhere, which is simpler and costs little.) The vocabulary is padded to a multiple of $n$.
- KV heads. Qwen3-8B has 32 query heads but 8 KV heads. Over 8 GPUs, each gets 4 query heads and 1 KV head. Over 16 GPUs, each gets 2 query heads, and they share a KV head with a neighbour: KV heads are replicated when there are fewer than ranks. Both ranks then cache the same KV head, so KV memory stops shrinking past $n = H_{kv}$.
@torch.no_grad()
def shard_tensor_parallel(model, comm):
"""Keep this rank's share of every layer. (Your engine: Chapter 41)
Attention: query heads split evenly; KV heads split too while there are at least as many as
ranks, and replicated beyond that (8 KV heads over 16 GPUs: each KV head on two ranks). Q, K
and V are column-parallel, o_proj row-parallel. MLP: gate/up column-parallel, down
row-parallel. Embedding and LM head: split by vocabulary. Norms and the router: replicated.
"""
from .moe import SparseMoeBlock
flat = model if isinstance(model, FlatModel) else FlatModel(model)
if hasattr(flat.model, "kv_spec"):
raise ValueError("Latent attention (Chapter 42) has one KV head: serve it with data-parallel attention")
c, tp, r = flat.cfg, comm.world, comm.rank
h, hkv, d = c.num_attention_heads, c.num_key_value_heads, c.head_dim
if h % tp or (hkv % tp and tp % hkv):
raise ValueError(f"{h} query / {hkv} KV heads can't be split over {tp} ranks")
hq, hk = h // tp, max(hkv // tp, 1)
kv0 = r * hk if hkv >= tp else (r * hq) // (h // hkv) # first KV head this rank needs
for layer in flat.layers:
a = layer.self_attn
if getattr(a, "qkv_proj", None) is not None:
raise ValueError("Shard before fuse(): fusing concatenates the projections this splits")
a.q_proj = take_rows(a.q_proj, r * hq * d, (r + 1) * hq * d)
a.k_proj = take_rows(a.k_proj, kv0 * d, (kv0 + hk) * d)
a.v_proj = take_rows(a.v_proj, kv0 * d, (kv0 + hk) * d)
a.o_proj = RowParallelLinear(a.o_proj, r * hq * d, (r + 1) * hq * d, comm)
m = layer.mlp
if isinstance(m, SparseMoeBlock):
layer.mlp = TensorParallelMoE(m, comm)
else:
lo, hi = shard_bounds(m.down_proj.in_features, comm)
m.gate_proj, m.up_proj = take_rows(m.gate_proj, lo, hi), take_rows(m.up_proj, lo, hi)
m.down_proj = RowParallelLinear(m.down_proj, lo, hi, comm)
embed = flat.backbone.embed_tokens.weight
tied = flat.model.lm_head.weight is embed
table = vocab_slice(embed, comm)
flat.backbone.embed_tokens = VocabParallelEmbedding(table, c.vocab_size, comm)
head = table if tied else vocab_slice(flat.model.lm_head.weight, comm)
flat.model.lm_head = ParallelLMHead(head, c.vocab_size, comm)
flat.cfg = replace(c, num_attention_heads=hq, num_key_value_heads=hk)
return flat
shard_tensor_parallel keeps only this rank’s share of each layer and records the local head counts in flat.cfg, so FlatModel.attention reshapes into the local heads, kv_spec reports the local KV heads, and the runner allocates a pool of the local size. Nothing else in the engine changes: the paged attention kernels from Chapter 32 run per head and never know other heads exist.
Sharding happens before fuse() (Chapter 33), because fusing concatenates exactly the projections that sharding splits. A real loader also never builds the full model on every rank: safetensors’ get_slice reads only the rows a rank keeps (stretch exercise 1).
MoE layers under tensor parallelism
A MoE layer could split each expert’s intermediate dimension like a dense MLP. With 128 experts of intermediate size 768 over 8 GPUs that leaves 96-wide slivers, too thin for efficient matmuls. vLLM’s --enable-expert-parallel instead gives each rank whole experts: every rank sees every token (attention is tensor parallel, so after its all-reduce every rank holds the same activations), runs the router, computes only the experts it owns, and the same single all-reduce adds everyone’s contributions. A shared expert splits like a dense MLP, and its partial sum joins the same all-reduce:
def local_experts(x, weights, experts, gate_up, down, first):
"""The routed experts this rank owns (first .. first + len(gate_up) - 1); others contribute
zero here and arrive in the all-reduce. A fused MoE kernel does the same with an expert map
that sends non-local experts to -1."""
out = torch.zeros_like(x)
owned = (experts >= first) & (experts < first + gate_up.shape[0])
for e in experts[owned].unique().tolist():
token, slot = torch.where(experts == e)
gate, up = F.linear(x[token], gate_up[e - first]).chunk(2, dim=-1)
y = F.linear(F.silu(gate) * up, down[e - first])
out.index_add_(0, token, y * weights[token, slot, None])
return out
class TensorParallelMoE(nn.Module):
"""A MoE layer under tensor parallelism: every rank sees every token, so experts are split
whole across ranks (expert parallelism inside the TP group) and the shared expert is split
by columns like a dense MLP. One all-reduce sums routed and shared partials together."""
def __init__(self, block, comm):
super().__init__()
e = block.experts.num_experts
if e % comm.world:
raise ValueError(f"{e} experts don't split over {comm.world} ranks")
per = e // comm.world
self.comm, self.first, self.gate = comm, comm.rank * per, block.gate
self.gate_up = nn.Parameter(block.experts.gate_up_proj[self.first:self.first + per].detach().clone())
self.down = nn.Parameter(block.experts.down_proj[self.first:self.first + per].detach().clone())
self.shared, self.shared_gate = None, block.shared_expert_gate
if block.shared_expert is not None:
s = block.shared_expert
lo, hi = shard_bounds(s.down_proj.in_features, comm)
self.shared = nn.ModuleDict({"gate_proj": take_rows(s.gate_proj, lo, hi), "up_proj": take_rows(s.up_proj, lo, hi),
"down_proj": take_cols(s.down_proj, lo, hi)})
def forward(self, x):
"""(Your engine: Chapter 41)"""
_, weights, experts = self.gate(x)
out = local_experts(x, weights, experts, self.gate_up, self.down, self.first)
if self.shared is not None:
s = self.shared
partial = s["down_proj"](F.silu(s["gate_proj"](x)) * s["up_proj"](x))
out = out + torch.sigmoid(self.shared_gate(x)) * partial # the gate is replicated: scaling commutes with the sum
return self.comm.all_reduce(out)
One scheduler, many workers
Who runs the scheduler? If every rank ran its own, they would have to agree on every decision (which requests, which blocks, which sampled token), which is fragile. vLLM and SGLang put the scheduler, block manager and sampler on one process and make the GPU workers dumb: each step, the driver broadcasts what to run, and every worker runs the same forward pass on its shard, meeting the others at each collective.
def portable(batch, device="cpu"):
"""The batch with its tensors on `device`, ready to pickle and broadcast."""
from .offload import meta_on
move = lambda t: t.to(device) if isinstance(t, torch.Tensor) else t # noqa: E731
extra = {k: move(v) for k, v in batch.extra.items()} if batch.extra else batch.extra
return replace(batch, input_ids=batch.input_ids.to(device), positions=batch.positions.to(device),
logits_indices=batch.logits_indices.to(device), meta=meta_on(batch.meta, device),
prompt_spans=None, extra=extra)
class _Swapped(dict):
def __init__(self, runner):
super().__init__()
self.runner = runner
def pop(self, key, default=None):
self.runner.comm.broadcast_object(("drop", key))
return self.runner.inner.swapped.pop(key, default)
class DistributedRunner:
"""Rank 0's runner. Every call that touches device state is broadcast first, so all ranks
run the same step in lockstep: the scheduler, block manager and sampler exist only here."""
def __init__(self, runner, comm):
self.inner, self.comm = runner, comm
self.swapped = _Swapped(self)
def __getattr__(self, name):
return getattr(self.inner, name)
def execute(self, batch):
self.comm.broadcast_object(("execute", portable(batch)))
return self.inner.execute(batch)
def swap_out(self, request, blocks):
self.comm.broadcast_object(("swap_out", request.request_id, blocks))
self.inner.swap_out(request, blocks)
def swap_in(self, request, blocks):
self.comm.broadcast_object(("swap_in", request.request_id, blocks))
self.inner.swap_in(request, blocks)
def shutdown(self):
self.comm.broadcast_object(("stop",))
def worker_loop(runner, comm):
"""Ranks 1..n-1: repeat whatever rank 0 does, until it says stop. (Your engine: Chapter 41)"""
while True:
command = comm.broadcast_object(None)
kind = command[0]
if kind == "execute":
runner.execute(portable(command[1], runner.device)) # collectives inside match rank 0's
elif kind in ("swap_out", "swap_in"):
getattr(runner, kind)(SimpleNamespace(request_id=command[1]), command[2])
elif kind == "drop":
runner.swapped.pop(command[1], None)
elif kind == "stop":
return
def parallel_engine(model, comm, config=None, mode="tensor", bounds=None, **engine_options):
"""Every rank calls this with the same model. Rank 0 gets an EngineCore to drive (call
engine.runner.shutdown() when done); the other ranks serve inside this call and get None."""
from .serve.engine import EngineCore, EngineConfig
from .serve.attention import get_backend
from .serve.runner import ModelRunner, num_blocks_for_memory
config = config or EngineConfig()
flat = shard_tensor_parallel(model, comm) if mode == "tensor" else PipelineStage(model, comm, bounds)
layers, kv_heads, head_dim = flat.kv_spec()
dtype = next(flat.parameters()).dtype
blocks = config.num_blocks or num_blocks_for_memory(config.kv_cache_bytes, layers, kv_heads, head_dim,
config.block_size, dtype)
config = replace(config, num_blocks=comm.min(blocks), cuda_graphs="off") # every rank: the same pool size
if comm.rank == 0:
engine = EngineCore(flat, config, **engine_options)
engine.runner = DistributedRunner(engine.runner, comm)
return engine
options = {"kv_cache_dtype": config.kv_cache_dtype} if config.kv_cache_dtype != "auto" else {}
worker_loop(ModelRunner(flat, config.num_blocks, config.block_size, get_backend(config.attention_backend, **options)),
comm)
return None
DistributedRunner wraps rank 0’s ModelRunner: execute, swap_out and swap_in are broadcast before they run locally, and everything else (prepare, the KV pool on rank 0) passes through. worker_loop is the whole of a worker. Every rank must have the same number of KV blocks, because block ids in the broadcast batch refer to all their pools at once, so parallel_engine takes the minimum over ranks.
Broadcasting a pickled batch costs a little CPU time per step: vLLM V1 writes the scheduler’s output into a shared-memory ring buffer that workers poll, and each worker builds its own input tensors. CUDA graphs (Chapter 33) capture NCCL all-reduces too, so the per-layer collectives replay with the rest of the step; parallel_engine turns graphs off because the gloo backend can’t be captured.
Pipeline parallelism
Tensor parallelism needs a fast link: two all-reduces per layer per step. Between servers, split the layers instead: rank $r$ owns a contiguous run of layers and forwards one hidden vector per token to rank $r + 1$.
class PipelineStage(FlatModel):
"""Rank r runs layers bounds[r] .. bounds[r+1] - 1. Rank 0 also embeds the tokens and, when
the residual stream comes back from the last stage, applies the final norm (the engine's
runner then applies the LM head there, where the sampler lives)."""
def __init__(self, model, comm, bounds=None):
super().__init__(model)
n, w = len(self.layers), comm.world
self.bounds = bounds or [round(i * n / w) for i in range(w + 1)]
self.comm = comm
mine = self.layers[self.bounds[comm.rank]:self.bounds[comm.rank + 1]]
self.layers = self.backbone.layers = nn.ModuleList(mine) # the other stages' layers are dropped
if comm.rank != 0:
self.backbone.embed_tokens = self.backbone.norm = self.model.lm_head = None
self.width, self.dtype = self.cfg.hidden_size, next(self.parameters()).dtype
def kv_spec(self):
return len(self.layers), self.cfg.num_key_value_heads, self.cfg.head_dim # this stage's layers only
def forward(self, input_ids, positions, kv_caches, meta, backend, embeds=None, **features):
"""This stage's layers; hidden states travel 0 -> 1 -> ... -> last -> 0. (Your engine: Chapter 41)"""
c, n = self.comm, positions.shape[0]
if c.rank == 0:
x = self.backbone.embed_tokens(input_ids) if embeds is None else embeds
else:
x = c.recv((n, self.width), self.dtype, c.rank - 1, positions.device)
rope = self.rope(positions, x.dtype)
delta = None
for i, layer in enumerate(self.layers):
h, x = self.add_norm(layer.input_layernorm, x, delta)
delta = self.attention(i, layer.self_attn, h, positions, rope, kv_caches[i], meta, backend)
h, x = self.add_norm(layer.post_attention_layernorm, x, delta)
delta = self.mlp(layer.mlp, h)
x = x if delta is None else x + delta # one tensor crosses each boundary: the residual stream
if c.world > 1:
c.send(x, (c.rank + 1) % c.world)
if c.rank != 0:
return None
x = c.recv((n, self.width), self.dtype, c.world - 1, positions.device)
h, _ = self.add_norm(self.backbone.norm, x, None)
return h
Only the residual stream crosses a boundary: x + delta, which the next stage’s first add_norm treats as a residual with nothing to add. The last stage sends it back to rank 0, which applies the final norm and the LM head where the sampler lives. (vLLM keeps the head on the last stage and returns sampled tokens instead: a few bytes rather than a hidden vector per sampled row. Sending the hidden state back keeps the engine unchanged here.)
As written, one batch is in flight at a time, so while stage 1 works, stage 0 waits: the pipeline bubble. With $p$ stages, each GPU is busy $1/p$ of the time. PP alone then buys memory, not speed. The fix is to keep $p$ batches in flight, each a step behind the one ahead: vLLM’s scheduler holds up to $p$ “virtual engines” whose batches never share a request. That is the async scheduling of Chapter 33, generalized from 1 to $p$ steps in flight (stretch exercise 2).
Measured
run.py dist runs a 4-layer model (hidden size 512, 8 query and 4 KV heads, 1,024-word vocabulary) with eight 64-token prompts and 16 new tokens each, on 1, 2 and 4 CPU processes over gloo, with the same 64 MB KV budget per rank:
{"mode": "tensor", "ranks": 1, "same_tokens": true, "steps": 16, "kv_blocks_per_rank": 512, "seconds": 0.27, "rank0_traffic": {}}
{"mode": "tensor", "ranks": 2, "same_tokens": true, "steps": 16, "kv_blocks_per_rank": 1024, "seconds": 0.6, "rank0_traffic": {"all_reduce": {"calls": 145, "MB": 11.65}, "all_gather": {"calls": 16, "MB": 0.26}}}
{"mode": "tensor", "ranks": 4, "same_tokens": true, "steps": 16, "kv_blocks_per_rank": 2048, "seconds": 0.79, "rank0_traffic": {"all_reduce": {"calls": 145, "MB": 11.65}, "all_gather": {"calls": 16, "MB": 0.13}}}
{"mode": "pipeline", "ranks": 2, "same_tokens": true, "steps": 16, "kv_blocks_per_rank": 1024, "seconds": 0.41, "rank0_traffic": {"all_reduce": {"calls": 1, "MB": 0.0}, "send": {"calls": 16, "MB": 1.29}}}
{"mode": "pipeline", "ranks": 4, "same_tokens": true, "steps": 16, "kv_blocks_per_rank": 2048, "seconds": 0.7, "rank0_traffic": {"all_reduce": {"calls": 1, "MB": 0.0}, "send": {"calls": 16, "MB": 1.29}}}
- Every configuration produces exactly the single-process tokens. (Exact equality holds here because FP32 sums in a different order still round to the same greedy choices on this model; in BF16 on GPUs, expect rare divergences after near-ties, as with any change of kernel.)
- KV capacity grows with the ranks. With a fixed budget per rank, TP 2 caches twice the blocks (each rank holds half the KV heads) and TP 4 four times. PP stages hold a fraction of the layers, so they gain the same way.
- The traffic matches the arithmetic. TP: 16 steps × (4 layers × 2 + 1 for the embedding) = 144 all-reduces, plus one to agree on the pool size. The first step (8 × 64 = 512 prompt tokens × 512 features × 4 bytes = 1 MB per all-reduce) dominates the 11.65 MB; a decode step moves 16 KB per all-reduce. PP sends one hidden vector per token per step: 1.29 MB in total.
- The times say nothing about GPUs. Four processes share a 4-core CPU, and each gloo all-reduce costs a millisecond or two. On GPUs with NVLink, TP 2 makes a bandwidth-bound decode step nearly twice as fast.
Expert parallelism with all-to-all
Tensor parallelism gives every rank every token. DeepSeek’s deployments do something else for big MoEs: attention runs data-parallel (each rank serves its own requests, with its own KV cache, so there’s no KV-head limit and no attention all-reduce), and experts run expert-parallel across all ranks. A token’s hidden vector travels to the ranks owning its $k$ experts and the results travel back: two all-to-alls per MoE layer.
class ExpertParallelMoE(nn.Module):
"""Experts split across ranks that hold DIFFERENT tokens (data-parallel attention, as in
DeepSeek's deployments). Each routed (token, expert) pair travels to the expert's rank and
back: two all-to-alls per MoE layer instead of an all-reduce of every token."""
def __init__(self, block, comm):
super().__init__()
e = block.experts.num_experts
self.per = e // comm.world
self.comm, self.first, self.gate = comm, comm.rank * self.per, block.gate
self.gate_up = nn.Parameter(block.experts.gate_up_proj[self.first:self.first + self.per].detach().clone())
self.down = nn.Parameter(block.experts.down_proj[self.first:self.first + self.per].detach().clone())
self.shared, self.shared_gate = block.shared_expert, block.shared_expert_gate # replicated
def forward(self, x):
"""Dispatch, compute, combine. (Your engine: Chapter 41)"""
_, weights, experts = self.gate(x)
k, flat = experts.shape[1], experts.reshape(-1)
owner = flat // self.per
order = owner.argsort(stable=True) # assignments grouped by destination rank
counts = torch.bincount(owner, minlength=self.comm.world).tolist()
rows = self.comm.all_to_all(list(x[order // k].split(counts)))
ids = self.comm.all_to_all(list(flat[order].split(counts)))
arrived, local = torch.cat(rows), torch.cat(ids) - self.first
y = torch.empty_like(arrived)
for e in local.unique().tolist(): # a grouped GEMM over the arrivals
idx = torch.where(local == e)[0]
gate, up = F.linear(arrived[idx], self.gate_up[e]).chunk(2, dim=-1)
y[idx] = F.linear(F.silu(gate) * up, self.down[e])
back = torch.cat(self.comm.all_to_all(list(y.split([r.shape[0] for r in rows]))))
out = torch.zeros_like(x).index_add_(0, order // k, back * weights.reshape(-1)[order, None])
if self.shared is not None:
out = out + torch.sigmoid(self.shared_gate(x)) * self.shared(x)
return out
Each rank sorts its (token, expert) assignments by destination rank, sends each destination its rows and expert ids, computes whatever arrives, and sends the results back in the same order, so the combine is an index_add_ with the router weights. The test gives two ranks different numbers of tokens and checks the result against the whole MoE on one process.
Moving $k$ hidden vectors per token twice per layer is the dominant cost at scale; DeepEP’s kernels overlap it with computation and use NVLink inside a node and RDMA between nodes. Routing is also never uniform: popular experts make their ranks the bottleneck, so DeepSeek’s EPLB replicates hot experts onto several ranks, rebalancing every few minutes from observed routing counts.
Data parallelism and prefix-aware routing
Data parallelism is the simplest: $N$ independent replicas behind a load balancer. Round-robin routing wastes the prefix cache, though: requests that share a long system prompt land on different replicas, and each replica prefills it. A router that remembers which prefixes it sent where can send the next such request to the replica that already holds its blocks, unless that replica is much busier.
class PrefixRouter:
"""Data parallelism: N independent engine replicas behind one router.
The router remembers which block hashes it has sent to each replica (an approximation of
that replica's prefix cache: it can't see evictions) and scores each replica by the prompt
tokens it would NOT have to prefill there, minus `balance` times the tokens it is already
working on. balance = 0 is pure cache affinity; a large balance is least-loaded routing.
"""
def __init__(self, num_replicas, block_size, balance=0.5, max_blocks=1 << 16):
self.block_size, self.balance, self.max_blocks = block_size, balance, max_blocks
self.seen = [OrderedDict() for _ in range(num_replicas)]
self.load = [0] * num_replicas
self.assigned = {}
def hashes(self, prompt, cache_key=()):
out, parent = [], None
for start in range(0, len(prompt) - self.block_size + 1, self.block_size):
parent = hash_block(parent, prompt[start:start + self.block_size], tuple(cache_key))
out.append(parent)
return out
def route(self, request_id, prompt, cost=None, cache_key=()):
"""Pick a replica for a request and charge it `cost` tokens. (Your engine: Chapter 41)"""
hashes = self.hashes(prompt, cache_key)
def score(r):
hit = 0
while hit < len(hashes) and hashes[hit] in self.seen[r]:
hit += 1
return hit * self.block_size - self.balance * self.load[r], -self.load[r], -r
best = max(range(len(self.seen)), key=score)
seen = self.seen[best]
for h in hashes:
seen[h] = None
seen.move_to_end(h)
while len(seen) > self.max_blocks:
seen.popitem(last=False)
cost = len(prompt) if cost is None else cost
self.load[best] += cost
self.assigned[request_id] = (best, cost)
return best
def finish(self, request_id):
replica, cost = self.assigned.pop(request_id, (None, 0))
if replica is not None:
self.load[replica] -= cost
class DataParallel:
"""Several EngineCores behind a PrefixRouter, stepped one after another (in a deployment,
each replica is its own process on its own GPUs, behind an HTTP router)."""
def __init__(self, engines, router):
self.engines, self.router = engines, router
self.owner = {}
def add_request(self, request_id, prompt, params):
replica = self.router.route(request_id, prompt, len(prompt) + params.max_tokens * params.n)
self.owner[request_id] = replica
return self.engines[replica].add_request(request_id, prompt, params)
def step(self):
outputs = []
for engine in self.engines:
for out in engine.step():
root = out.request_id.split(":")[0]
if out.finished and root in self.router.assigned and not any(
r.split(":")[0] == root for e in self.engines for r in e.scheduler.requests):
self.router.finish(root)
outputs.append(out)
return outputs
@property
def has_unfinished(self):
return any(e.has_unfinished for e in self.engines)
def generate(self, prompts, params):
results = {}
for rid, prompt in prompts.items():
self.add_request(rid, prompt, params)
while self.has_unfinished:
for out in self.step():
results.setdefault(out.request_id, []).extend(out.new_token_ids)
return results
The score of replica $r$ is the prompt tokens it would not have to prefill there minus balance × the tokens it is already working on. The router’s memory of cached blocks is approximate (it can’t see evictions); SGLang’s router keeps the same kind of approximate radix tree per worker, and llm-d subscribes to each vLLM replica’s KV-cache events to know exactly. DataParallel steps several in-process engines for the tests; in production each replica is its own server process (Chapter 36) and the router is an HTTP proxy.
Disaggregated prefill and decode
Prefill and decode want different things. A prefill chunk is compute-bound and makes every decode in the same step wait for it (Chapter 24’s interference, which chunked prefill only softens). Decode is bandwidth-bound and wants big batches. Disaggregation (DistServe, Splitwise, Mooncake) runs them on different GPUs: a prefill instance computes the prompt’s KV and the first token, ships the KV to a decode instance, and the decode instance carries on. Each pool is sized and tuned for its job, and the time to first token no longer depends on how busy decode is.
The engine already has every piece: preemption by swapping (Chapter 31) copies a request’s blocks to host memory and back. A prefill instance exports a finished one-token request exactly the way swap-out does, and a decode instance admits it exactly the way swap-in does:
@dataclass
class KVTransfer:
request_id: str
token_ids: list # prompt + the first output token
num_prompt_tokens: int
kv: list # per layer: (K blocks, V blocks[, scales]) in host memory
generator_state: object = None
class KVExporter:
"""On the PREFILL instance: when a request marked transfer_kv finishes its one-token run,
copy its blocks out before the scheduler frees them."""
def __init__(self, engine):
self.engine, self.ready = engine, {}
engine.scheduler.on_finish = self.export
def prefill(self, request_id, prompt, params):
request = self.engine.add_request(request_id, prompt, replace(params, max_tokens=1))
request.extra["transfer_kv"] = True
return request
def export(self, request):
"""(Your engine: Chapter 41)"""
if not request.extra.get("transfer_kv") or request.status is not Status.FINISHED_LENGTH:
return # stopped at its first token (EOS): nothing to decode
e, rid = self.engine, request.request_id
e.runner.swap_out(request, e.blocks.req_blocks[rid]) # the swap path already copies blocks out
generator = e.generators.get(rid)
self.ready[rid] = KVTransfer(rid, list(request.token_ids), request.num_prompt_tokens,
e.runner.swapped.pop(rid), generator.get_state() if generator else None)
def import_kv(engine, transfer, params):
"""On the DECODE instance: admit the request as if it had been swapped out after its prefill,
so the scheduler's swap-in path allocates blocks and copies the KV in. (Your engine: Chapter 41)"""
request = Request(transfer.request_id, transfer.token_ids[:transfer.num_prompt_tokens], params, engine.eos_token_id)
for token in transfer.token_ids[transfer.num_prompt_tokens:]:
request.append(token)
request.num_computed_tokens = transfer.num_prompt_tokens # the first output token's KV is computed next
request.swapped_out = True
engine.runner.swapped[request.request_id] = transfer.kv
if transfer.generator_state is not None: # seeded sampling continues the same stream
engine.generators[request.request_id] = torch.Generator(device=engine.device).set_state(transfer.generator_state)
elif params.seed is not None:
engine.generators[request.request_id] = torch.Generator(device=engine.device).manual_seed(params.seed)
engine.scheduler.add(request)
return request
def send_transfer(comm, transfer, dst):
"""Ship a transfer to another rank: metadata as a pickled object, KV as raw tensors."""
meta = replace(transfer, kv=[[(t.shape, t.dtype) for t in layer] for layer in transfer.kv])
dist.send_object_list([meta], dst)
for layer in transfer.kv:
for t in layer:
comm.send(t, dst)
def recv_transfer(comm, src):
box = [None]
dist.recv_object_list(box, src)
meta = box[0]
kv = [tuple(comm.recv(shape, dtype, src, "cpu") for shape, dtype in layer) for layer in meta.kv]
return replace(meta, kv=kv)
KVExporter hooks the scheduler’s on_finish, the moment before a finished request’s blocks are freed. The transfer carries the prompt, the first token, every layer’s blocks and, for seeded sampling, the random generator’s state, so the decode side continues the same random stream: the test checks that seeded sampling gives the same tokens as one engine. import_kv creates the request with num_computed_tokens set past the prompt and swapped_out = True, and the scheduler’s swap-in path allocates blocks and copies the KV in. The decode engine runs no prefill step at all.
Is shipping KV affordable? Qwen3-8B’s KV is 144 KB per token (Chapter 31), so a 2,000-token prompt is 295 MB: 6 ms over a 50 GB/s RDMA link, against roughly 100 ms to prefill it on an H100 (32 TFLOP at about 40% of peak). Real connectors (vLLM’s NIXL connector, Mooncake’s transfer engine, SGLang’s) move KV GPU-to-GPU over RDMA, layer by layer while the prefill is still running, and can transfer between instances with different TP sizes by re-slicing heads. send_transfer here goes over torch.distributed (gloo between CPU processes in the test).
Ring attention for very long contexts
A million-token prompt’s KV cache doesn’t fit on one GPU even for a small model, and its prefill is quadratic. Context parallelism splits the sequence: rank $r$ holds tokens $[rT, (r+1)T)$ and their Q, K and V. Attention needs every earlier key, so the K/V chunks travel around a ring of ranks: at each hop, each rank attends its queries to the chunk it holds, then passes the chunk on and receives the previous rank’s. After $n - 1$ hops every rank has seen every chunk.
def attention_partial(q, k, v, q_pos, k_pos, scale):
"""Causal attention of q [Tq, H, D] over k, v [Tk, H, D] -> (out, lse [Tq, H]); rows that see
no key get lse = -inf and contribute nothing when merged."""
scores = torch.einsum("qhd,khd->hqk", q.float(), k.float()) * scale
scores = scores.masked_fill(k_pos[None, None, :] > q_pos[None, :, None], float("-inf"))
lse = torch.logsumexp(scores, dim=-1) # [H, Tq]
probs = torch.exp(scores - torch.where(lse.isinf(), 0, lse)[..., None])
return torch.einsum("hqk,khd->qhd", probs, v.float()), lse.T
def merge_partials(out, lse, new_out, new_lse):
"""Chapter 32's split-KV merge, two at a time."""
top = torch.maximum(lse, new_lse)
top = torch.where(top.isinf(), 0, top)
a, b = torch.exp(lse - top), torch.exp(new_lse - top)
total = a + b
merged = (out * a[..., None] + new_out * b[..., None]) / torch.where(total > 0, total, 1)[..., None]
return merged, top + torch.log(total)
def ring_attention(q, k, v, comm, scale=None):
"""Each rank holds one contiguous chunk of a long sequence (q, k, v: [T, H, D]). K/V chunks
travel around the ring; after world - 1 hops every rank has attended its queries to every
earlier key without any rank holding the whole sequence. (Your engine: Chapter 41)"""
t, scale = q.shape[0], scale or q.shape[-1] ** -0.5
q_pos = torch.arange(comm.rank * t, (comm.rank + 1) * t, device=q.device)
out = torch.zeros(q.shape, dtype=torch.float32, device=q.device)
lse = torch.full((t, q.shape[1]), float("-inf"), device=q.device)
kv, owner = torch.stack((k, v)), comm.rank
for hop in range(comm.world):
if owner <= comm.rank: # chunks from later ranks are entirely masked
k_pos = torch.arange(owner * t, (owner + 1) * t, device=q.device)
o, l = attention_partial(q, kv[0], kv[1], q_pos, k_pos, scale)
out, lse = merge_partials(out, lse, o, l)
if hop + 1 < comm.world: # pass ours on, take the previous rank's
incoming = torch.empty_like(kv)
ops = [dist.isend(kv.contiguous(), (comm.rank + 1) % comm.world),
dist.irecv(incoming, (comm.rank - 1) % comm.world)]
for op in ops:
op.wait()
kv, owner = incoming, (owner - 1) % comm.world
return out.to(q.dtype)
Combining the partial results is exactly Chapter 32’s split-KV merge: each partial carries its log-sum-exp, and $o = \sum_s e^{\ell_s - \ell} o_s$ with $\ell = \log \sum_s e^{\ell_s}$. Causality makes the plain split unbalanced: rank 0’s queries see one chunk, the last rank’s see all of them. Production implementations split the sequence into $2n$ chunks and give rank $r$ chunks $r$ and $2n - 1 - r$ (“zigzag”), which evens out the work, and overlap each hop’s transfer with the previous chunk’s computation.
Build it
Engine milestone 41: many GPUs. Implement RowParallelLinear.forward, VocabParallelEmbedding.forward, ParallelLMHead.forward, TensorParallelMoE.forward, shard_tensor_parallel, PipelineStage.forward, ExpertParallelMoE.forward, worker_loop, PrefixRouter.route, KVExporter.export, import_kv and ring_attention in engine/parallel.py.
pytest tests/test_ch41_multi_gpu.py # 13 tests; several start 2-4 processes each
python run.py dist --impl engine
The tests check token-for-token equality with a single process under TP 2, TP 2 with tied embeddings, TP 4 (replicated KV heads), TP 2 for a MoE with seeded sampling, and PP 2; expert parallelism with uneven token counts; ring attention over three ranks; the router’s choices; and disaggregated serving in one process and across two.
Stretch exercises
- ★★ Load only the shard: give
shard_tensor_parallela path that reads each rank’s rows and columns straight from the safetensors files (safe_open(...).get_slice(name)[start:end]), so no rank ever holds the full model. Measure peak memory per rank. Where:shard_tensor_parallelinengine/parallel.py. - ★★★ Fill the pipeline bubble: let the engine keep
ppbatches in flight (each made of requests not in the others), using Chapter 33’s placeholder mechanism. Measure throughput on two GPUs against PP with one batch in flight. Where:DistributedRunner/worker_loopinengine/parallel.py, with scheduling inengine/serve/engine.py. - ★★ Two GPUs over PCIe: compare TP 2 and PP 2 on decode latency and throughput for a model that doesn’t fit on one GPU. Where does each win? Where:
experiments/ch41.py(create it), adaptingrun.py’scmd_distandengine.parallel.parallel_engine. - ★★ Tensor parallelism for Medusa and draft models (Chapter 37): which parts must be sharded, and which can stay replicated on rank 0? Where: paper first; implement sharding/wrappers in
engine/parallel.pyand drafter integration inengine/serve/spec.py. - ★★★ Layer-by-layer KV transfer: start sending layer $\ell$’s blocks as soon as the prefill finishes layer $\ell$ (a CUDA event per layer, a second stream for the copies). How much of the transfer hides behind the prefill? Where:
KVExporter,send_transferandrecv_transferinengine/parallel.py, with per-layer completion hooks inengine/serve/model.py.
Check your understanding
- Why can Q/K/V and the MLP’s gate and up projections be split by output without any communication, while $W_o$ and the down projection need an all-reduce?
- Why does TP’s KV memory per GPU stop shrinking once there are more ranks than KV heads?
- A decode step’s all-reduces are about 1 MB each. Is that bandwidth- or latency-bound on NVLink? On PCIe?
- Why does pipeline parallelism with one batch in flight not make decoding faster, and what fixes it?
- When is all-to-all expert parallelism cheaper than tensor parallelism’s all-reduce for a MoE layer?
- Why does a prefix-aware router need a load term at all?
- How does disaggregation reuse preemption by swapping, and what extra state does a seeded request need to carry?
Going deeper
- Shoeybi et al., Megatron-LM: Training Multi-Billion Parameter Language Models Using Model Parallelism (2019): the column/row-parallel construction.
- Huang et al., GPipe (NeurIPS 2019) and Narayanan et al., PipeDream and Efficient Large-Scale Language Model Training on GPU Clusters Using Megatron-LM (SC 2021), for pipeline schedules.
- DeepSeek-AI, DeepSeek-V3 Technical Report (2024) and the DeepEP, EPLB and “Day 6: inference system overview” releases (2025), for DP attention with EP.
- Zhong et al., DistServe (OSDI 2024); Patel et al., Splitwise (ISCA 2024); Qin et al., Mooncake (FAST 2025), for disaggregated serving.
- Liu, Zaharia and Abbeel, Ring Attention with Blockwise Transformers for Near-Infinite Context (2023).
- vLLM’s
vllm/distributed/(parallel state, custom all-reduce, KV connectors) andvllm/model_executor/layers/linear.py; SGLang’s router (sgl-router) and its DP-attention implementation.