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

33. Removing overhead: fusion, graphs and async scheduling

In this chapter

  • Where a serving step loses time outside the GPU's arithmetic: kernel launches, host bookkeeping, small matmuls and host syncs.
  • Fusing Q/K/V and gate/up projections into single matmuls, and residual-add + RMSNorm and SiLU-and-mul into single kernels.
  • A fused MoE: token-to-expert alignment with static shapes, and one grouped GEMM launch for every expert.
  • CUDA graphs for a serving engine: one per batch-size bucket, static buffers, padding rows and a scratch block.
  • Async scheduling: planning step s+1 while step s runs, with placeholder tokens resolved one step later.

You will build

FlatModel.fuse (engine/serve/model.py), silu_and_mul_kernel, moe_align and fused_moe_kernel (engine/kernels/triton_fused.py), GraphRunner.capture and run (engine/serve/graphs.py) and EngineCore.step_async (engine/serve/engine.py).

Time: 6-8 hours. GPU: needed for the speedups (every test runs on a CPU; real graph capture is tested on CUDA).

Where a step loses time

Chapter 19 made one request’s decode fast by removing host syncs and replaying a CUDA graph. The engine core of Chapter 31 brought those problems back, multiplied:

source of overheadwhere it comes fromsize, Qwen3-0.6B, batch 16, H100
GPU workweights + KV cache read once per step1.2 GB / 3.35 TB/s ≈ 0.36 ms
kernel launches28 layers × ~16 kernels, ~5 µs of CPU each≈ 2.2 ms
host bookkeepingscheduling, building the batch, Python0.3 ms measured below, growing with batch size
host syncs.tolist() of sampled tokens; MoE’s expert countsGPU drains its queue each time

The GPU could finish the step in about a third of a millisecond, but the CPU needs seven times that just to issue it. Small and mid-size models at moderate batch sizes are launch-bound, and the cure is the same as in Chapter 19: fewer, larger kernels, then graphs. But a serving engine adds two complications. Its batch changes shape every step, which graphs dislike, and the host-side work of scheduling grows with the number of requests, so it eventually costs as much as the GPU step, at which point the GPU idles while Python plans. This chapter removes each overhead in turn.

Fewer, larger matmuls

Every attention layer multiplies the same input $h$ by three matrices, $W_Q$, $W_K$ and $W_V$, and every MLP by two, $W_{\text{gate}}$ and $W_{\text{up}}$. Stacking each group into one matrix computes the same outputs with one matmul:

$$ [,q ;|; k ;|; v,] = h,[,W_Q ;|; W_K ;|; W_V,]^\top . $$

That’s three launches become one, the input is read once instead of three times, and the combined matrix has more output columns for the matmul’s tiles to divide among the SMs. For Qwen3-0.6B, whose K and V projections are only 1,024 × 1,024, the separate matmuls are too small to fill an H100; the fused 4,096-column one isn’t. The fusion is done once, at load time, and the originals are deleted so that it costs no memory:

    @torch.no_grad()
    def fuse(self, kernels=True):
        """Concatenate Q/K/V and gate/up weights so each becomes ONE matmul, and (kernels=True)
        use fused Triton norms, activation and MoE.  (Your engine: Chapter 33)

        The originals are deleted afterwards, so fusing costs no extra memory. Anything that
        reads the separate weights (Chapters 20-23's tools) must run before fusing.
        """
        if getattr(self, "adapters", None) is not None:
            raise ValueError("Fuse projections before installing an adapter bank")
        for layer in self.layers:
            attn = layer.self_attn
            if getattr(attn, "qkv_proj", None) is None and hasattr(attn, "q_proj"):
                parts = [attn.q_proj, attn.k_proj, attn.v_proj]
                attn.qkv_sizes = [p.out_features for p in parts]
                attn.qkv_proj = concat_linear(parts)
                del attn.q_proj, attn.k_proj, attn.v_proj
            mlp = layer.mlp
            if hasattr(mlp, "gate_proj") and hasattr(mlp, "up_proj"):
                mlp.gate_up_proj = concat_linear([mlp.gate_proj, mlp.up_proj])
                del mlp.gate_proj, mlp.up_proj
        self.fused = kernels
        return self

FlatModel routes through three small helpers so that the fused and unfused paths share one forward: project_qkv splits the fused output, mlp runs SwiGLU from gate_up_proj, and add_norm does a layer’s residual add and the next norm. With kernels=True, add_norm calls Chapter 14’s fused residual-add + RMSNorm kernel (one memory pass instead of three), and the MLP uses a fused SiLU-and-mul:

@triton.jit
def silu_and_mul_kernel(x_ptr, out_ptr, width, BLOCK: tl.constexpr):
    """out[r] = silu(x[r, :I]) * x[r, I:] for x [N, 2I].  (Your engine: Chapter 33)"""
    row, block = tl.program_id(0), tl.program_id(1)
    cols = block * BLOCK + tl.arange(0, BLOCK)
    mask = cols < width
    gate = tl.load(x_ptr + row * 2 * width + cols, mask=mask, other=0.0).to(tl.float32)
    up = tl.load(x_ptr + row * 2 * width + width + cols, mask=mask, other=0.0).to(tl.float32)
    out = gate / (1.0 + tl.exp(-gate)) * up
    tl.store(out_ptr + row * width + cols, out.to(out_ptr.dtype.element_ty), mask=mask)


def silu_and_mul(x, block=1024):
    check_device(x)
    x = x.contiguous()
    n, width = x.shape[0], x.shape[1] // 2
    out = torch.empty((n, width), device=x.device, dtype=x.dtype)
    silu_and_mul_kernel[(n, triton.cdiv(width, block))](x, out, width, BLOCK=block)
    return out

Other fusions production engines use, in order of payoff: RoPE applied inside the KV-write kernel; the attention output projection’s input quantized on the fly for FP8 matmuls; the LM head fused with sampling’s softmax. Each saves a pass over memory and a launch; none changes the math.

A MoE layer that never asks the host

Chapter 27’s forward_grouped sorts assignments by expert and runs one matmul per expert. That’s correct, and it has two problems in a serving loop. It launches $2E$ matmuls per layer (256 for Qwen3-30B-A3B’s 128 experts), and it calls bincount(...).tolist() to learn each expert’s count, a host sync that a CUDA graph can’t contain. The fix, used by vLLM, SGLang and MegaBlocks, is a grouped GEMM: one kernel launch in which each program multiplies a block of rows by its expert’s weights.

The rows must be arranged so that each block of BLOCK_M rows belongs to one expert. moe_align does that with device operations only, and with shapes that depend on the batch size but never on the routing:

def moe_align(expert_ids, num_experts, block_m):
    """Group assignments by expert, each group padded to a multiple of block_m rows.  (Your engine: Chapter 33)

    expert_ids [N, k] -> (sorted_ids [M_max], block_expert [M_max // block_m], num_padded [1]).
    sorted_ids[i] is the flat assignment index (token * k + slot) that row i of the grouped
    GEMM handles, or N * k for padding. Every shape depends only on N, k, E and block_m, never
    on the routing, and no value is read back to the host: a CUDA graph can capture this.
    """
    flat = expert_ids.reshape(-1)
    total = flat.numel()
    m_max = total + num_experts * (block_m - 1)
    m_max = triton.cdiv(m_max, block_m) * block_m
    counts = torch.zeros(num_experts, dtype=torch.int64, device=flat.device)
    counts.scatter_add_(0, flat.long(), torch.ones_like(flat, dtype=torch.int64))   # bincount, static shape
    padded = (counts + block_m - 1) // block_m * block_m
    padded_start = torch.cumsum(padded, 0) - padded            # where each expert's group begins
    start = torch.cumsum(counts, 0) - counts                   # where it begins in sorted order
    order = torch.argsort(flat, stable=True)                   # assignments, expert by expert
    sorted_expert = flat[order].long()
    rank = torch.arange(total, device=flat.device) - start[sorted_expert]
    sorted_ids = torch.full((m_max,), total, dtype=torch.int64, device=flat.device)
    sorted_ids[padded_start[sorted_expert] + rank] = order
    block_starts = torch.arange(0, m_max, block_m, device=flat.device)
    block_expert = torch.searchsorted(torch.cumsum(padded, 0), block_starts, right=True)
    num_padded = padded.sum().reshape(1)
    return sorted_ids, block_expert.clamp_max(num_experts - 1), num_padded

Each expert’s group is padded up to a multiple of BLOCK_M, so at most $E \times (\text{BLOCK_M} - 1)$ padding rows are added. The output buffer is sized for that worst case. Padding rows carry the index $N \cdot k$, one past the last real assignment, which the kernel masks. Even the “bincount” is a scatter_add_ into a fixed-size tensor, because torch.bincount returns a tensor whose size depends on the largest value, a shape the host would have to learn.

The kernel reads its block’s expert, gathers its rows’ inputs and multiplies them by that expert’s weight tiles:

@triton.jit
def fused_moe_kernel(a_ptr, w_ptr, c_ptr, sorted_ptr, block_expert_ptr, num_padded_ptr, weight_ptr,
                     N, K, total, top_k, s_am, s_we, s_wn, s_cm,
                     A_PER_TOKEN: tl.constexpr, MUL_WEIGHT: tl.constexpr,
                     BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr):
    """C[row(i)] = A[src(i)] @ W[expert(i)]^T for the BLOCK_M rows i of this program's group.  (Your engine: Chapter 33)

    A_PER_TOKEN: A has one row per token (first GEMM, src = assignment // top_k) or one row per
    assignment (second GEMM, src = assignment). MUL_WEIGHT scales each row by its router weight.
    """
    pid_m, pid_n = tl.program_id(0), tl.program_id(1)
    if pid_m * BLOCK_M < tl.load(num_padded_ptr):
        rm = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
        assignment = tl.load(sorted_ptr + rm)
        valid = assignment < total
        if A_PER_TOKEN:
            src = assignment // top_k
        else:
            src = assignment
        expert = tl.load(block_expert_ptr + pid_m).to(tl.int64)
        rn = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
        acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
        for k0 in range(0, K, BLOCK_K):
            rk = k0 + tl.arange(0, BLOCK_K)
            a = tl.load(a_ptr + src[:, None] * s_am + rk[None, :], mask=valid[:, None] & (rk[None, :] < K), other=0.0)
            w = tl.load(w_ptr + expert * s_we + rn[:, None] * s_wn + rk[None, :],
                        mask=(rn[:, None] < N) & (rk[None, :] < K), other=0.0)
            acc += tl.dot(a, tl.trans(w), input_precision="ieee")
        if MUL_WEIGHT:
            acc = acc * tl.load(weight_ptr + assignment, mask=valid, other=0.0).to(tl.float32)[:, None]
        tl.store(c_ptr + assignment[:, None] * s_cm + rn[None, :], acc.to(c_ptr.dtype.element_ty),
                 mask=valid[:, None] & (rn[None, :] < N))


def _grouped_gemm(a, w, out, sorted_ids, block_expert, num_padded, weights, top_k, per_token, mul_weight,
                  block_m, block_n=32, block_k=32):
    n_out, k_in = w.shape[1], w.shape[2]
    grid = (sorted_ids.numel() // block_m, triton.cdiv(n_out, block_n))
    fused_moe_kernel[grid](a, w, out, sorted_ids, block_expert, num_padded, weights,
                           n_out, k_in, out.shape[0], top_k, a.stride(0), w.stride(0), w.stride(1), out.stride(0),
                           A_PER_TOKEN=per_token, MUL_WEIGHT=mul_weight,
                           BLOCK_M=block_m, BLOCK_N=block_n, BLOCK_K=block_k)


def fused_moe(x, gate_up, down, weights, expert_ids, block_m=16):
    """x [N, D]; gate_up [E, 2I, D]; down [E, D, I]; weights, expert_ids [N, k] -> [N, D].

    Same math as Experts.forward_grouped (Chapter 27), with no Python loop over experts and no
    .tolist(): two grouped GEMMs, an activation, and a sum over each token's k results.
    """
    check_device(x, gate_up, down)
    n, k = expert_ids.shape
    sorted_ids, block_expert, num_padded = moe_align(expert_ids, gate_up.shape[0], block_m)
    x = x.contiguous()
    hidden = torch.empty((n * k, gate_up.shape[1]), device=x.device, dtype=x.dtype)
    _grouped_gemm(x, gate_up, hidden, sorted_ids, block_expert, num_padded, weights, k, True, False, block_m)
    act = silu_and_mul(hidden)
    out = torch.empty((n * k, down.shape[1]), device=x.device, dtype=x.dtype)
    flat_weights = weights.reshape(-1).contiguous()
    _grouped_gemm(act, down, out, sorted_ids, block_expert, num_padded, flat_weights, k, False, True, block_m)
    return out.view(n, k, -1).sum(1)

The first GEMM reads token rows (A_PER_TOKEN: assignment $i$ uses token $i // k$) and writes one row per assignment; the second reads those rows, multiplies by the router weight, and writes them back in assignment order, so the final view(n, k, -1).sum(1) adds each token’s $k$ expert outputs. Experts with no tokens cost nothing: they own no blocks. A block whose start lies past num_padded exits at once, so the grid can be sized for the worst case.

CUDA graphs for a changing batch

Chapter 19 captured one graph for one request. A serving engine’s decode batch has 13 requests in one step, 14 in the next, 12 after that, with different context lengths and block tables. Graphs replay fixed kernels on fixed addresses, so the engine captures one graph per bucket of batch sizes and pads each step up to the nearest bucket:

buckets:   1  2  4  8  16  24  32  40 ...  max_num_seqs
13 decode requests  ->  replay the 16-graph with 3 padding rows

Everything the graph reads lives in static buffers sized for the largest bucket. A replay copies this step’s token IDs, positions, slot mapping, lengths and block tables into them, fills the padding rows, and replays. The padding rows must be harmless, and the earlier chapters arranged for that:

  • their slot is -1, which write_kv (Chapter 31) redirects to a scratch block that the runner allocates and the block manager never hands out. Filtering the padding rows instead would require the host to know how many there are, a sync;
  • their length is 0, so the attention kernels (Chapter 32) produce zeros for them, never NaNs.
class GraphRunner:
    def __init__(self, runner, max_batch, max_blocks_per_seq, buckets=None, use_graphs=None):
        if getattr(runner.backend, "name", "") != "triton":
            raise ValueError("CUDA graphs need a backend that never reads host-side batch values (use triton)")
        self.runner, self.max_batch, self.max_blocks = runner, max_batch, max_blocks_per_seq
        self.buckets = buckets or default_buckets(max_batch)
        device = runner.device
        self.use_graphs = device.type == "cuda" if use_graphs is None else use_graphs
        # Static inputs, sized for the largest bucket. A replay reads whatever these hold.
        self.input_ids = torch.zeros(max_batch, dtype=torch.long, device=device)
        self.positions = torch.zeros(max_batch, dtype=torch.long, device=device)
        self.slot_mapping = torch.full((max_batch,), -1, dtype=torch.long, device=device)
        self.seq_lens = torch.zeros(max_batch, dtype=torch.int32, device=device)
        self.block_table = torch.zeros((max_batch, max_blocks_per_seq), dtype=torch.int32, device=device)
        self.query_start_loc = torch.arange(max_batch + 1, dtype=torch.int32, device=device)
        self.graphs, self.outputs, self.pool = {}, {}, None

    def meta(self, size):
        """BatchMeta over views of the static buffers. Host-side values are the bucket's, not the
        step's: the kernels' grids must not change between capture and replay."""
        return BatchMeta(self.query_start_loc[:size + 1], self.seq_lens[:size], self.block_table[:size],
                         self.slot_mapping[:size], self.runner.block_size, list(range(size + 1)),
                         [1] * size, 1, 1)

    def forward(self, size):
        flat = self.runner.flat
        hidden = flat(self.input_ids[:size], self.positions[:size], self.runner.kv_caches, self.meta(size),
                      self.runner.backend)
        return flat.compute_logits(hidden)

    @torch.inference_mode()
    def capture(self):
        """Record one graph per bucket, largest first, sharing one memory pool.  (Your engine: Chapter 33)"""
        if not self.use_graphs:
            return
        for size in sorted(self.buckets, reverse=True):
            stream = torch.cuda.Stream()
            stream.wait_stream(torch.cuda.current_stream())
            with torch.cuda.stream(stream):              # warm up: compile kernels, allocate workspaces
                for _ in range(2):
                    self.forward(size)
            torch.cuda.current_stream().wait_stream(stream)
            graph = torch.cuda.CUDAGraph()
            with torch.cuda.graph(graph, pool=self.pool):
                self.outputs[size] = self.forward(size)
            self.pool = graph.pool()
            self.graphs[size] = graph
        # Capture ran real steps on padding rows (slot -1): only the scratch block was written.

    def can_run(self, batch):
        meta = batch.meta
        return (meta.decode_only and meta.num_reqs <= self.max_batch
                and meta.block_table.shape[1] <= self.max_blocks and batch.logits_indices.numel() == meta.num_reqs)

    @torch.inference_mode()
    def run(self, batch):
        """Copy the step into the static buffers, pad to the bucket, replay.  (Your engine: Chapter 33)"""
        meta, r = batch.meta, batch.meta.num_reqs
        size = next(b for b in self.buckets if b >= r)
        self.input_ids[:r].copy_(batch.input_ids)
        self.positions[:r].copy_(batch.positions)
        self.slot_mapping[:r].copy_(meta.slot_mapping)
        self.seq_lens[:r].copy_(meta.seq_lens)
        self.block_table[:r].zero_()
        self.block_table[:r, :meta.block_table.shape[1]].copy_(meta.block_table)
        self.input_ids[r:size].zero_()
        self.positions[r:size].zero_()
        self.slot_mapping[r:size].fill_(-1)          # padding: written to the scratch block
        self.seq_lens[r:size].zero_()                # padding: attention returns zeros
        if size in self.graphs:
            self.graphs[size].replay()
            logits = self.outputs[size]
        else:                                        # CPU, or no graph for this bucket: same buffers, eagerly
            logits = self.forward(size)
        return logits[:r]

One subtlety: the BatchMeta built over the static buffers carries the bucket’s host-side values (num_reqs, max_query_len), not the step’s. The Triton launcher computes its grid from them, and a replay must launch exactly the grid that was captured. Anything else that reads host-side batch values, such as the reference backend’s Python loop over requests or split-KV’s choice of splits from max_seq_len, can’t be in a graph. GraphRunner refuses backends other than Triton, and the Triton backend uses the unified kernel for graphed steps.

Only decode-only batches are graphed. Prefill and mixed batches run eagerly: their token counts vary too much to bucket cheaply, and their large matmuls keep the GPU busy anyway, so launch overhead matters less. vLLM goes further with piecewise graphs: it compiles the model with torch.compile, splits the graph at every attention call, and captures the pieces between attentions for buckets of token counts, so even mixed batches replay most of their kernels from graphs (stretch exercise 3).

Graphs share one memory pool (pool=self.pool), captured from the largest bucket down, so the smaller graphs reuse the largest one’s intermediate buffers instead of each holding their own.

Async scheduling

After fusion and graphs, a decode step on the GPU is a few kernels and one replay. But between two steps the CPU still has work to do: read the sampled tokens back, check the stop conditions, publish blocks, plan the next step, build its metadata. While it does, the GPU idles:

synchronous:  CPU  [plan s][build s]               [read s, update][plan s+1][build s+1]
              GPU                  [====== s ======]                                   [== s+1 ==]

async:        CPU  [plan s][build s][plan s+1][build s+1][read s, update][plan s+2] ...
              GPU                  [====== s ======][====== s+1 ======][====== s+2 ...

The host time per step grows with the batch. The demo below measured 0.28 ms at 16 requests, 0.74 ms at 64 and 1.33 ms at 256 for this engine’s Python. On an H100, a decode step of Qwen3-0.6B at batch 256 takes well under a millisecond of GPU time, so a synchronous loop would leave the GPU idle more than half the time.

Async scheduling (SGLang’s “zero-overhead scheduler”, 2024; vLLM’s --async-scheduling, 2025) plans step $s+1$ before step $s$’s tokens are known. The scheduler assumes every request in step $s$ will produce one token. In the token lists, those tokens are placeholders:

  1. Launch step $s$; sample on the device (no .tolist()); for each request that samples, append a PLACEHOLDER and advance num_computed_tokens now.

  2. Plan step $s+1$ right away. A decode request’s input is its newest token, which is a placeholder. The runner copies its real value from step $s$’s token tensor into step $s+1$’s input IDs on the device:

        def prepare(self, scheduled, block_tables, pad_to=None, previous=None):
        """Build the batch. With async scheduling, a request's newest input token may still be a
        placeholder: its value is in `previous.tokens` on the device, so copy it there, on the
        device, without reading it back (Chapter 33)."""
        batch = build_batch(scheduled, block_tables, self.block_size, self.device, pad_to)
        if previous is not None:
            rows, sources = [], []
            for (request, n), start in zip(scheduled, batch.meta.query_start_loc_cpu):
                if request.token_ids[request.num_computed_tokens] == PLACEHOLDER:
                    rows.append(start)
                    sources.append(previous.row_of[request.request_id])
            if rows:
                index = torch.tensor(rows, device=self.device)
                batch.input_ids[index] = previous.tokens[torch.tensor(sources, device=self.device)]
        self.prepare_features(batch, scheduled)
        return batch
    
  3. Launch step $s+1$. Only then read step $s$’s tokens back (the one sync per step, now overlapping $s+1$’s GPU work), replace the placeholders, and check the stop conditions.

    def step_async(self):
        """Launch step s+1, then read step s's tokens while s+1 runs.  (Your engine: Chapter 33)

        The scheduler plans s+1 as if every request in s will produce its token: those tokens
        are PLACEHOLDERs in the token lists, num_computed_tokens advances at launch, and the
        runner copies each placeholder's real value into s+1's inputs on the device. Reading
        s's tokens back (the one host sync per step) then overlaps s+1's GPU work. A request
        that turns out to have stopped at step s wastes one token of s+1, which is dropped.
        """
        plan = self.scheduler.schedule()
        self.preemptions += len(plan.preempted)
        for request, blocks in plan.swap_out:
            self.runner.swap_out(request, blocks)
        for request, blocks in plan.swap_in:
            self.runner.swap_in(request, blocks)
        launched = None
        if not plan.empty:
            batch = self.runner.prepare(plan.scheduled, self.blocks.req_blocks, previous=self.inflight)
            logits = self.runner.execute(batch)
            sampling = [r for (r, _), k in zip(plan.scheduled, batch.sample_counts) if k]
            tokens = self.sampler.sample_device(logits, sampling, self.generators) if sampling else None
            for request, n in plan.scheduled:
                request.num_computed_tokens += n
                self.blocks.cache_full_blocks(request)
            slots = []
            for request in sampling:
                request.append(PLACEHOLDER)
                request.num_placeholders += 1
                slots.append((request, request.num_tokens - 1))
            launched = Inflight(tokens, slots, {r.request_id: i for i, r in enumerate(sampling)})
            self.steps += 1
        outputs = self.resolve(self.inflight) if self.inflight is not None else []
        self.inflight = launched
        return outputs

    def resolve(self, inflight):
        """Read a launched step's tokens back and do the bookkeeping sync mode does at once."""
        values = inflight.tokens.tolist() if inflight.tokens is not None else []
        outputs, now = [], time.monotonic()
        for (request, index), token in zip(inflight.slots, values):
            if request.status.finished:                 # stopped or aborted while this step ran
                continue
            request.token_ids[index] = token
            request.num_placeholders -= 1
            request.first_token_time = request.first_token_time or now
            status = self.check_stop(request, index)
            if status is not None:
                del request.token_ids[index + 1:]        # the speculative next placeholder, if any
                request.num_placeholders = 0
                self.scheduler.finish(request, status)
                self.generators.pop(request.request_id, None)
            else:
                self.blocks.cache_full_blocks(request)
            outputs.append(self.make_output(request, [token]))
        return outputs

Why is this safe? GPU work on one stream runs in launch order, so step $s+1$’s kernels see everything step $s$ wrote. Three consequences need care:

  • A request that stops at step $s$ is already in step $s+1$. Its extra token is computed and dropped; its blocks are freed when its stop is seen. A block freed this way may be handed to another request in step $s+2$, whose writes are ordered after step $s+1$’s stray one.
  • Placeholders must never be hashed. The block manager publishes full blocks by their tokens’ hash; a block containing a placeholder would be published under the wrong key. update_hashes stops at the first placeholder.
  • The scheduler must not run past max_tokens. A request whose last allowed token is in flight is skipped until it resolves.

Features that need a token on the host before the next step can be planned don’t fit: constrained decoding (Chapter 34) advances a grammar with the sampled token to build the next step’s mask, and speculative decoding (Chapter 37) must know how many drafts were accepted. Engines either run those requests synchronously or move the bookkeeping to the device.

Run it

python run.py overhead --requests 16 --slots 16 --new-tokens 64

On a laptop CPU with the 4-layer test model:

{"config": "eager", "tok_s": 1016.4, "steps": 64, "host_ms_per_step": 0.275, "same_tokens": true}
{"config": "fused", "tok_s": 965.5, "steps": 64, "host_ms_per_step": 0.323, "same_tokens": true}
{"config": "fused + async", "tok_s": 919.1, "steps": 64, "host_ms_per_step": 0.354, "same_tokens": true}

Nothing got faster, and that’s the expected result on a CPU: PyTorch’s CPU operations run synchronously, so there’s no asynchronous device whose idle time async scheduling could fill, launches cost little, and the fused path uses PyTorch rather than the Triton kernels (which would run in the slow interpreter). What the CPU run does show is that every configuration produces exactly the same tokens, and how much host time a step needs. With --requests 64 --slots 64 it was 0.74 ms per step, and with 256 requests 1.33 ms. On a CUDA machine the command adds the two graph configurations, and the differences appear: compare each configuration’s tokens/s with the bandwidth ceiling of Chapter 10.

Build it

Engine milestone 33: overhead. Implement FlatModel.fuse in engine/serve/model.py; silu_and_mul_kernel, moe_align and fused_moe_kernel in engine/kernels/triton_fused.py; GraphRunner.capture and GraphRunner.run in engine/serve/graphs.py; and EngineCore.step_async in engine/serve/engine.py (resolve, the runner’s placeholder copy, and the launchers are provided).

pytest tests/test_ch33_overhead.py
python run.py overhead --impl engine

The tests check SiLU-and-mul; fused dense and MoE models (with PyTorch and with Triton kernels) against unfused ones; that moe_align covers every assignment once with one expert per block; the fused MoE against Chapter 27’s loop for several batch, expert and top-k sizes; graph buckets with padding rows against solo outputs; async scheduling against synchronous scheduling, including preemption by recompute and swap, graph buckets, and stop tokens; and, on a GPU, real graph capture against eager execution.

Stretch exercises

  1. ★ Profile one decode step on a GPU with torch.profiler before and after fuse() and graphs. Count kernels per step and plot the GPU’s idle gaps. Where: experiments/ch33.py (create it), adapting run.py’s cmd_overhead.
  2. ★★ Build a persistent batch: keep the metadata tensors between steps and update only the rows of requests that joined, left or changed (vLLM V1’s InputBatch). Measure host time per step at 256 requests against build_batch. Where: persistent metadata in engine/serve/batch.py, updated from ModelRunner.prepare in engine/serve/runner.py.
  3. ★★★ Piecewise graphs: split FlatModel.forward at each attention call, capture each piece between attentions for buckets of token counts (64, 128, …, 2,048), and run attention eagerly in between. Mixed prefill and decode batches now replay most of their kernels from graphs. Measure time to first token under load. Where: FlatModel._forward in engine/serve/model.py and capture/replay in engine/serve/graphs.py.
  4. ★★ Write a fused kernel that applies RoPE to K and writes it, with V, into the paged pool in one pass. Test it against apply_rope + write_kv. Where: add a kernel in engine/kernels/triton_unified.py; integrate it in engine/serve/triton_backend.py and engine/serve/model.py.

Check your understanding

  1. Why is Qwen3-0.6B’s decode launch-bound at batch 16 on an H100, while Qwen3-32B’s isn’t?
  2. What three things does fusing Q/K/V into one matmul save?
  3. Why can’t forward_grouped run inside a CUDA graph, and which two changes in moe_align fix that?
  4. A batch of 13 decode requests runs in the 16-bucket graph. What do the 3 padding rows read and write, and why can’t they change the 13 real outputs?
  5. Why must the graph’s BatchMeta carry the bucket’s num_reqs rather than the step’s?
  6. In async scheduling, a request emits a stop token at step $s$. What happens to its token at step $s+1$ and to its blocks?
  7. Why can’t a block that contains a placeholder be published to the prefix cache?

Going deeper

  • vLLM V1: A Major Upgrade to vLLM’s Core Architecture (vLLM blog, 2025), on the persistent batch and piecewise CUDA graphs; vLLM’s vllm/v1/worker/gpu_model_runner.py and its async scheduling option.
  • SGLang v0.4: Zero-Overhead Batch Scheduler, Cache-Aware Load Balancer, Faster Structured Outputs (LMSYS blog, December 2024), for the overlapped CPU/GPU loop.
  • Gale et al., MegaBlocks: Efficient Sparse Training with Mixture-of-Experts (MLSys 2023), for block-sparse grouped expert computation; vLLM’s fused_moe Triton kernel, which follows the alignment and grouped-GEMM structure of this chapter.
  • PyTorch documentation: CUDA Graphs (memory pools, torch.cuda.graph(pool=...)) and torch.compile (mode="reduce-overhead"); GPU Mode L35 (SGLang performance optimization).