19. Fast decode: syncs, CUDA graphs and compilation
In this chapter
- Why engine v1 reaches only a fraction of the bandwidth ceiling on a GPU: the CPU, not the GPU, is the bottleneck.
- Host synchronizations: what causes them, and how to keep the decode loop entirely on the device.
- Static shapes and a static KV cache, so that the same kernels run with the same addresses every step.
- CUDA graphs and
torch.compile: recording a whole decode step once and replaying it with one launch. - Reading a profile to tell launch-bound, memory-bound and compute-bound steps apart.
You will build
engine/fast.py: an on-device sampler and a FastDecoder with eager, CUDA-graph and compiled modes, plus StaticKVCache in engine/kv_cache.py.
Time: 4-6 hours. GPU: recommended (the eager mode and all tests run on the CPU; graphs and compilation need CUDA).
Where the time goes
Chapter 18 ended with a puzzle. The decode ceiling for Qwen3-0.6B in BF16 on an RTX 4090 is about 840 tokens/s (1.19 GB of weights read per token at 1,008 GB/s), but engine v1 runs far below it. Let’s count what one decode step asks the CPU to do.
Each of the 28 layers runs about 20 PyTorch operations: two RMSNorms (several kernels each without fusion), three projections, two QK-norms, RoPE on Q and K, the cache write, attention, the output projection, a residual add, the gate and up projections, SiLU, a multiply, the down projection, another residual add. That’s roughly 600 kernel launches per token. Each one costs the CPU several microseconds of Python and PyTorch dispatch, plus a few microseconds for the CUDA driver to launch it.
| per token | |
|---|---|
| GPU time if weights stream at full bandwidth (1.19 GB ÷ 1,008 GB/s) | 1.2 ms |
| CPU time to issue ~600 launches at ~5-10 µs each | 3-6 ms |
The GPU finishes each small kernel faster than the CPU can issue the next one, so it idles between them. A decode step of a small model is launch-bound: its speed is set by CPU overhead, not by memory or arithmetic. Larger models move toward the bandwidth ceiling, because each kernel does more work while the number of launches stays similar, but even an 8B model loses a noticeable fraction to overhead.
There are three fixes, in order of how much they help:
- Remove host synchronizations, so the CPU can run ahead and queue work while the GPU executes.
- Replay the whole step as a CUDA graph, so ~600 launches become one.
- Fuse kernels, so there are fewer of them and less memory traffic between them (Chapter 14).
Host synchronizations
CUDA kernel launches are asynchronous. When Python calls torch.matmul on CUDA tensors, the CPU enqueues the kernel and returns immediately, typically long before the kernel runs. While the queue is non-empty, the GPU never waits for the CPU: launch overhead overlaps with GPU execution.
Some operations break this. Anything that needs a value from the GPU on the CPU must wait for every queued kernel to finish:
int(token),token.item(),tensor.tolist(),print(tensor);if (tensor > 0).any():(theboolconversion is a sync);torch.nonzero, boolean-mask indexing likex[mask],torch.unique(output shape depends on data);- copying a CPU tensor to the GPU from non-pinned memory, and
torch.cuda.synchronize()itself.
Engine v1 calls int(token) every step, to check stop tokens and to yield the token. After that call, the GPU queue is empty, and the next step starts from scratch: the GPU waits for the first launch, the second, and so on. The CPU never gets ahead.
The fix is to keep everything the loop needs on the device:
- The next input token, its position and the output buffer are device tensors that the step updates in place.
- Sampling happens on the device and produces a device tensor.
- Stop tokens are checked every
check_everysteps, with one sync per check. A few tokens generated after a stop are trimmed afterwards. That’s wasted work of at mostcheck_every - 1steps, in exchange for removing almost all syncs.
Tip
torch.cuda.set_sync_debug_mode("warn")makes PyTorch print a warning at every synchronizing operation. Run one decode step with it on to find the syncs you didn’t know about.
Sampling without the CPU
Chapter 8’s sampler used torch.multinomial, which is fine, but top-p and min-p filters involve sorting and data-dependent cutoffs. For the graph-friendly path, the decoder uses the Gumbel-max trick, which draws an exact sample from $\operatorname{softmax}(z)$ with only elementwise operations and an argmax:
$$ \text{if } g_i = -\log(-\log u_i),\ u_i \sim \text{Uniform}(0,1) \text{ independently, then } \arg\max_i (z_i + g_i) \sim \operatorname{softmax}(z). $$
Why it works, in one line: $z_i + g_i$ is a Gumbel random variable with location $z_i$, and the probability that the $i$-th of several independent Gumbels is the largest is $e^{z_i}/\sum_j e^{z_j}$. Dividing $z$ by the temperature first gives temperature sampling.
def sample_on_device(logits, temperature):
"""Greedy when temperature == 0; otherwise the Gumbel-max trick, which draws exactly from
softmax(logits / temperature) using only elementwise ops and argmax (no host sync). (Your engine: Chapter 19)"""
if temperature == 0:
return logits.argmax(-1, keepdim=True)
uniform = torch.rand_like(logits, dtype=torch.float32).clamp_(1e-10, 1.0)
gumbel = -torch.log(-torch.log(uniform))
return (logits.float() / temperature + gumbel).argmax(-1, keepdim=True)
The milestone test draws 20,000 samples and checks their frequencies against the softmax. Top-k and top-p filters can be added before the argmax (set filtered logits to $-\infty$); top-k with a fixed k keeps shapes static, and so does top-p implemented with a sort and a mask rather than slicing.
Static shapes and a static cache
A CUDA graph records exact kernels with exact pointer arguments. Replaying it re-runs them on whatever data is at those addresses now. So everything a step reads and writes must live at fixed addresses with fixed shapes:
- The input token and position: one-element tensors updated in place with
copy_andadd_. - The KV cache: Chapter 16’s
KVCachereturns a view of the valid prefix, whose shape grows every step. That’s useless for a graph.
StaticKVCache solves the cache problem with the attention rule you’ve used since Chapter 5. It writes keys at their absolute positions in a preallocated buffer and always returns the full capacity, together with the position of each slot, key_positions = arange(capacity). Slots that haven’t been written yet sit at positions greater than every live query, so key_pos <= query_pos masks them, with no extra bookkeeping:
class StaticKVCache:
"""Fixed-shape cache indexed by absolute position. Each row is a slot that a request can own.
Writes go to explicit positions, so rows may hold different lengths. Reads return the full
capacity; unwritten or stale slots sit at positions greater than every live query, so the
causal rule (key_pos <= query_pos) hides them without any extra mask. Shapes never change,
which is what CUDA graphs and torch.compile need. (Your engine: Chapter 19)
"""
def __init__(self, layers, slots, kv_heads, capacity, head_dim, device="cpu", dtype=torch.float32):
self.capacity = capacity
shape = (slots, kv_heads, capacity, head_dim)
# zeros, not empty: masked weights are exactly 0, and 0 * garbage could still be NaN.
self.keys = [torch.zeros(shape, device=device, dtype=dtype) for _ in range(layers)]
self.values = [torch.zeros(shape, device=device, dtype=dtype) for _ in range(layers)]
self.key_positions = torch.arange(capacity, device=device)
def update(self, layer, k, v, positions, rows=None):
"""(Your engine: Chapter 19)"""
if positions.ndim == 1:
positions = positions.expand(k.shape[0], -1)
keys, values = self.keys[layer], self.values[layer]
if rows is None:
rows = torch.arange(k.shape[0], device=k.device)
# One scatter per tensor: row r, head h, position positions[r, t] <- k[r, h, t].
b, h, t, d = k.shape
row_index = rows[:, None].expand(b, t)
keys[row_index, :, positions] = k.transpose(1, 2)
values[row_index, :, positions] = v.transpose(1, 2)
return keys[rows], values[rows], self.key_positions
@property
def bytes(self):
return sum(t.numel() * t.element_size() for t in self.keys + self.values)
Two details matter:
- Zeros, not
empty. A masked attention weight is exactly 0, but $0 \times \text{NaN} = \text{NaN}$. Uninitialized memory can contain NaN bit patterns, so the buffers are zero-filled once. - Rows are slots. Each row of the cache can hold a different request at a different length. Writes go to
(row, position)pairs with one scatter. That’s exactly what continuous batching needs in Chapter 24, so the same class serves both.
The cost of a static cache: attention always reads capacity keys, even when only 50 are valid. For a short conversation in a large buffer, that wastes bandwidth. Production engines pass the valid length as a device tensor to a custom decode kernel that stops early (Chapter 25’s paged decode kernel does this). Choose a capacity near the longest request you expect, or capture one graph per capacity bucket.
The fast decoder
class FastDecoder:
"""Batch-1 decoder over static buffers. (Your engine: Chapter 19)"""
def __init__(self, model, capacity, mode="eager", temperature=0.0):
if mode not in ("eager", "graph", "compile"):
raise ValueError("mode must be eager, graph or compile")
if mode == "graph" and not torch.cuda.is_available():
raise RuntimeError("CUDA graphs need a CUDA device")
self.model, self.mode, self.temperature = model.eval(), mode, temperature
p = next(model.parameters())
layers, kv_heads, head_dim = model.cache_spec()
self.cache = StaticKVCache(layers, 1, kv_heads, capacity, head_dim, p.device, p.dtype)
self.capacity = capacity
self.token = torch.zeros(1, 1, dtype=torch.long, device=p.device) # next input token
self.position = torch.zeros(1, 1, dtype=torch.long, device=p.device) # its absolute position
self.outputs = torch.zeros(1, capacity, dtype=torch.long, device=p.device)
self.step_fn = self._step
self.graph = None
if mode == "compile":
self.step_fn = torch.compile(self._step, mode="reduce-overhead", fullgraph=False)
def _step(self):
"""One decode step that touches only static tensors: run, sample, record, advance. (Your engine: Chapter 19)"""
logits = self.model(self.token, self.cache, positions=self.position)[:, -1]
nxt = sample_on_device(logits, self.temperature)
self.outputs.index_copy_(1, self.position[0], nxt) # output i is stored at slot of its input
self.token.copy_(nxt)
self.position.add_(1)
return nxt
def _capture(self):
"""Record one step as a CUDA graph. Warm up on a side stream first, as PyTorch requires."""
stream = torch.cuda.Stream()
stream.wait_stream(torch.cuda.current_stream())
saved = (self.token.clone(), self.position.clone(), self.outputs.clone())
with torch.cuda.stream(stream):
for _ in range(2):
self._step()
torch.cuda.current_stream().wait_stream(stream)
self.graph = torch.cuda.CUDAGraph()
with torch.cuda.graph(self.graph):
self._step()
# Warmup and capture advanced the state; put it back. Cache slots they wrote lie at
# positions the real run will overwrite before reading.
self.token.copy_(saved[0]); self.position.copy_(saved[1]); self.outputs.copy_(saved[2])
@torch.inference_mode()
def generate(self, prompt_ids, max_new_tokens, stop_ids=(), check_every=16):
"""Return the generated token IDs (a Python list) for one prompt. (Your engine: Chapter 19)"""
ids = torch.as_tensor(prompt_ids, device=self.token.device).view(1, -1)
prompt = ids.shape[1]
if prompt + max_new_tokens > self.capacity:
raise ValueError("Prompt plus new tokens exceed the decoder's capacity")
positions = torch.arange(prompt, device=ids.device)[None]
logits = self.model(ids, self.cache, positions=positions)[:, -1] # prefill (eager)
first = sample_on_device(logits, self.temperature)
self.token.copy_(first)
self.position.fill_(prompt)
self.outputs.zero_()
if self.mode == "graph" and self.graph is None:
self._capture()
stop = torch.tensor(list(stop_ids) or [-1], device=ids.device)
produced = 1
while produced < max_new_tokens:
steps = min(check_every, max_new_tokens - produced)
for _ in range(steps):
if self.graph is not None:
self.graph.replay()
else:
self.step_fn()
produced += steps
window = torch.cat((first, self.outputs[:, prompt:prompt + produced - 1]), dim=1)
if stop_ids and bool(torch.isin(window, stop).any()): # one sync per check
break
tokens = torch.cat((first, self.outputs[:, prompt:prompt + produced - 1]), dim=1)[0].tolist()
for i, token in enumerate(tokens): # trim anything after a stop token
if token in stop_ids:
return tokens[:i + 1]
return tokens[:max_new_tokens]
Read _step first. It’s the whole decode step, and it touches only static tensors: run the model on self.token at self.position, sample on the device, store the token in self.outputs, copy it into self.token, and advance the position. No Python value depends on the GPU’s results, so nothing waits.
generate runs prefill eagerly (the prompt length varies, so it can’t be graphed without bucketing), samples the first token, then runs _step in groups of check_every, checking for stop tokens once per group.
CUDA graphs
A CUDA graph is a recording of a stream of GPU work. You capture it once and replay it as a single launch:
graph = torch.cuda.CUDAGraph()
with torch.cuda.graph(graph): # kernels are recorded, not executed
step()
graph.replay() # re-run every recorded kernel, one launch from the CPU
During capture, PyTorch records the kernels and their arguments; the memory they allocate comes from a private pool that stays reserved for the graph’s lifetime, so replays reuse the same addresses. The rules follow from “replay re-runs the recorded kernels on the same addresses”:
- No host syncs during capture. A capture can’t wait for results that don’t exist yet; PyTorch raises an error if you try.
- No data-dependent Python control flow. An
ifon a tensor’s value is evaluated once, at capture time, and its branch is baked in. - Inputs must be updated in place. Assigning
self.token = new_tensorgives a new address that the graph never sees. Useself.token.copy_(new). - Warm up first, on a side stream. The first call to many operations triggers lazy initialization (cuBLAS handles, autotuning, Triton compilation), which must not happen during capture.
_captureruns the step twice on a side stream before recording. - Random numbers work, because PyTorch registers the generator’s state with the graph and advances its offset each replay. Each replay draws fresh Gumbel noise.
The warmup and capture each executed a step, which advanced the token, the position and the outputs. _capture restores them afterwards. They also wrote K and V into one or two cache slots beyond the prompt; those slots sit at positions the real run will overwrite before any query can see them, so they’re harmless. Thinking through “what state did capture modify?” is a habit worth forming: it’s the most common source of graph bugs.
torch.compile
torch.compile(step, mode="reduce-overhead") does two jobs at once: it fuses chains of small operations into generated Triton kernels (fewer launches, less traffic), and it captures CUDA graphs automatically (“reduce-overhead” means “use graphs”). mode="max-autotune" additionally benchmarks matmul configurations. The first calls are slow (compilation can take a minute); measure only after warmup.
Compilation needs the same discipline as graphs: static shapes (or marked dynamic dimensions), no syncs inside the compiled region, and in-place state updates. Code written for mode="graph" is already compile-friendly, which is why the decoder supports both.
Measure it
python run.py fast --new-tokens 128 --device cuda
The command measures the plain generate() loop from Chapter 8, then the fast decoder in eager mode, then (on CUDA) in graph mode, and checks that greedy tokens agree. On a laptop CPU with the tiny test model, it printed:
{"mode": "generate()", "tokens_per_s": 384.0}
{"mode": "eager", "tokens_per_s": 359.6, "same_tokens": true}
No improvement, and that’s the right result: on a CPU, operations run synchronously, so there’s no queue to keep full and nothing for sync removal to win. On a GPU the picture changes completely. The shape of what you should expect for Qwen3-0.6B on a consumer GPU (illustrative, not a measurement from this book’s validation; run it and record yours):
| mode | what limits it | typical fraction of the bandwidth ceiling |
|---|---|---|
engine v1 (int(token) every step) | CPU launches, GPU idle between kernels | 10-25% |
| fast decoder, eager | CPU launches, but queued ahead | 15-35% |
| fast decoder, CUDA graph | GPU kernels, many still small | 40-70% |
torch.compile + graph | fused kernels | 50-80% |
The remaining gap comes from kernels that don’t reach peak bandwidth at batch 1 (a 1,024 × 3,072 mat-vec is too small to saturate the memory system), from attention reading the whole static capacity, and from unfused elementwise operations.
Read a profile
Guessing where time goes is unreliable. Profile one decode step:
from torch.profiler import profile, ProfilerActivity
with profile(activities=[ProfilerActivity.CPU, ProfilerActivity.CUDA]) as prof:
decoder.generate(prompt, 32)
print(prof.key_averages().table(sort_by="cuda_time_total", row_limit=15))
prof.export_chrome_trace("decode.json") # open in https://ui.perfetto.dev
In the trace, look at the GPU stream row:
- Gaps between kernels mean launch-bound. Fix with graphs.
- Back-to-back kernels, mat-vecs taking most of the time means memory-bound, which is the goal for decode. Check achieved bandwidth: bytes read ÷ kernel time.
- A
cudaStreamSynchronizeoraten::itemon the CPU row in the middle of a step is a sync you missed.
NVIDIA Nsight Systems (nsys profile python run.py fast ...) shows the same picture for the whole process, including Python, and is the tool GPU Mode lectures use most.
Build it
Engine milestone 19: fast decode. Implement sample_on_device, FastDecoder._step and FastDecoder.generate in engine/fast.py, and StaticKVCache.update in engine/kv_cache.py (the constructor and graph capture are provided).
pytest tests/test_ch19_fast.py # graph and compile tests run only on CUDA
python run.py fast --impl engine --device cuda
The tests check that the fast decoder produces exactly the tokens of sampling.generate under greedy decoding, that stop tokens are trimmed correctly whatever check_every is, that Gumbel-max samples match the softmax, that rows of a static cache can hold requests of different lengths, and (on a GPU) that graph and compile modes match eager.
Stretch exercises
- ★ Run one step of engine v1 under
torch.cuda.set_sync_debug_mode("warn")and list every sync it reports. Then do the same for the fast decoder. Where:experiments/ch19.py(create it), comparingengine.engine.LLMwithengine.fast.FastDecoder. - ★★ Fuse Q, K and V into one projection (concatenate the three weight matrices at load time) and gate and up into another. Count the kernels per step before and after with the profiler. Where:
Qwen3AttentionandSwiGLUinengine/qwen3.py, with weight conversion inengine/loaders.py. - ★★ Swap Chapter 14’s fused add+RMSNorm Triton kernel into
Qwen3Layer, keeping the PyTorch path as a fallback. Verify logits, then measure graph-mode tokens/s. Where:Qwen3Layer.forwardinengine/qwen3.py, callingengine.kernels.triton_basics.rmsnorm. - ★★ Graph the prefill too: capture one graph per prompt-length bucket (64, 128, 256, …) and pad prompts up to the next bucket. What must the padding positions be so that they don’t affect the real tokens? Where:
FastDecoder._capture/generateinengine/fast.py, with padding masks passed throughengine/qwen3.py. - ★★★ Add top-k and top-p filtering to
sample_on_devicewithout any operation whose output shape depends on data, and check the sample distribution against Chapter 8’s sampler. Where:sample_on_deviceinengine/fast.py.
Check your understanding
- Why does removing
int(token)from the loop make the GPU faster, even though the GPU does the same work? - Why can’t a CUDA graph contain
if logits.argmax() == stop_id:? - How does
StaticKVCachehide slots that haven’t been written yet, without a separate mask? - Why does the fast decoder show no speedup on a CPU?
- What happens if you write
self.token = nxtinstead ofself.token.copy_(nxt)inside a captured step?
Going deeper
- GPU Mode L1 (Profiling and integrating CUDA kernels in PyTorch) and L16 (Hands-on profiling) for the profiler and Nsight; L6 (Optimizing PyTorch optimizers) for horizontal fusion and why many tiny launches are slow; L35 (SGLang performance optimization) for CUDA graphs and overhead removal in a production engine.
- The PyTorch blog Accelerating Generative AI with PyTorch II: GPT, Fast and the
gpt-fastrepository: static KV cache,torch.compile(mode="reduce-overhead"), and int8/int4 weight-only quantization in under 1,000 lines. This chapter follows its approach. - PyTorch documentation: CUDA Graphs (capture rules, memory pools,
make_graphed_callables) and torch.compile troubleshooting (graph breaks, recompilations). - Gumbel (1954) and Maddison, Tarlow and Minka, A Sampling* (2014) for the Gumbel-max trick.
- vLLM’s and SGLang’s CUDA-graph runners, which capture one graph per batch size bucket.