Keyboard shortcuts

Press ← or → to navigate between chapters

Press S or / to search in the book

Press ? to show this help

Press Esc to hide this help

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:

splitscommunication per layergood for
tensor (TP)every weight matrix, across GPUs2 all-reduces of the activationslatency, fitting big models inside one NVLink domain
pipeline (PP)the layers, into stagesone hidden vector per token per stage boundaryfitting models across nodes with slower links
expert (EP)a MoE’s experts2 all-to-alls (or one all-reduce)big MoEs: DeepSeek-V3, Qwen3-235B, Kimi K2
data (DP)the requests, over whole replicasnone (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

  1. ★★ Load only the shard: give shard_tensor_parallel a 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_parallel in engine/parallel.py.
  2. ★★★ Fill the pipeline bubble: let the engine keep pp batches 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_loop in engine/parallel.py, with scheduling in engine/serve/engine.py.
  3. ★★ 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), adapting run.py’s cmd_dist and engine.parallel.parallel_engine.
  4. ★★ 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.py and drafter integration in engine/serve/spec.py.
  5. ★★★ 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_transfer and recv_transfer in engine/parallel.py, with per-layer completion hooks in engine/serve/model.py.

Check your understanding

  1. 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?
  2. Why does TP’s KV memory per GPU stop shrinking once there are more ranks than KV heads?
  3. A decode step’s all-reduces are about 1 MB each. Is that bandwidth- or latency-bound on NVLink? On PCIe?
  4. Why does pipeline parallelism with one batch in flight not make decoding faster, and what fixes it?
  5. When is all-to-all expert parallelism cheaper than tensor parallelism’s all-reduce for a MoE layer?
  6. Why does a prefix-aware router need a load term at all?
  7. 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) and vllm/model_executor/layers/linear.py; SGLang’s router (sgl-router) and its DP-attention implementation.