Inference Zero To Hero
Build your own LLM inference engine from first principles, and understand every byte it moves.
When you send a message to a chatbot, a program reads billions of numbers from memory for every word it writes back. Somebody wrote that program. By the end of this book, you will have written one too. It will be a real inference engine that loads open-weight models from disk, runs them on your GPU (or CPU), streams answers, serves many users at once, and implements the architecture of Qwen3.8-Flash-Next: a 176-billion-parameter hybrid model with linear attention, sparse attention, a 512-expert mixture, hashed n-gram memory and multi-stream residuals.
You will also be able to change models: fine-tune them, attach LoRA adapters, quantize them, look inside them and edit what they compute.
Who this book is for
You are a working software engineer. You can read and write Python, you know what a function, a class and a loop are, and you are comfortable in a terminal. You do not need:
- prior machine-learning experience: we start from “what is a gradient?”;
- GPU or CUDA experience: we start from “what is a thread?”;
- heavy math: high-school algebra plus a willingness to follow small worked examples. When we need a derivative or a matrix identity, we derive it on the page with numbers.
If you already know some of this, skim the early chapters, but do the engine milestones. Later chapters assume your engine works.
What you will build
The book is organized around one project: your engine. Every chapter adds a piece and ends with a milestone test suite you make pass.
| Part | You build | Visible result |
|---|---|---|
| I. Neural networks | tensors, an autograd engine, a BPE tokenizer | a network that learns; a tokenizer trained on a novel |
| II. A GPT from scratch | attention, transformer blocks, training, sampling, a checkpoint loader | your GPT writes text; real GPT-2 runs in your code |
| III. GPU programming | CUDA and Triton kernels, FlashAttention | kernels you wrote, measured against PyTorch’s |
| IV. A dense engine | KV cache, Qwen3, engine v1, fast decode, quantization | chat with Qwen3 through your engine at near-hardware speed |
| V. Changing models | fine-tuning, LoRA/QLoRA, steering and abliteration | your own adapters and model edits |
| VI. Serving | continuous batching, paged KV cache, speculative decoding | many requests at once, sharing memory safely |
| VII. Frontier architectures | MoE, Gated DeltaNet, sparse attention, Flash-Next | a from-scratch Flash-Next that matches the official implementation |
| VIII. Putting the engine together | unified batching/paging, kernels, HTTP, speculation, formats, offload, parallelism, more models, images and adapters | one serving loop, an API, and correctness/latency evidence with explicit limits |
The book has 46 chapters (including setup), in eight parts. The reference implementation passes the milestone suites your engine will use; Appendix F records current counts and platform coverage. Its tiny FP32 Flash-Next agrees with the official Hugging Face implementation to about 1e-7. Real-checkpoint and production performance claims require separate validation.
How each chapter works
Every chapter follows the same rhythm, so you always know where you are:
- Why it matters: the problem, tied to the engine.
- Concepts: intuition first (a picture, an analogy, a small example worked with real numbers), then the precise version.
- See it: diagrams and interactive playgrounds you can poke.
- Code it: the real implementation, in tabs. Python is the main language. C++ and Rust versions sit one click away, and your choice is remembered:
def softmax(x):
e = [math.exp(v - max(x)) for v in x]
return [v / sum(e) for v in e]
void softmax(float* x, size_t n) {
float m = *std::max_element(x, x + n), sum = 0;
for (size_t i = 0; i < n; ++i) sum += (x[i] = std::exp(x[i] - m));
for (size_t i = 0; i < n; ++i) x[i] /= sum;
}
#![allow(unused)]
fn main() {
fn softmax(x: &mut [f32]) {
let max = x.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
let sum: f32 = x.iter_mut().map(|v| { *v = (*v - max).exp(); *v }).sum();
x.iter_mut().for_each(|v| *v /= sum);
}
}
- Build it: an engine milestone. You implement the chapter’s functions in your
engine/package and run its tests. The book supplies all the scaffolding (file layout, I/O, plotting, test harness), so you only write the ideas. - Stretch exercises: graded ★ to ★★★. Each Where note identifies the file to edit or a new experiment to create; Appendix A explains paths, imports and extension tests.
- Check your understanding: short questions, with answers in Appendix C.
- Going deeper: exact pointers into the source books, lectures and papers.
Callouts mark things worth stopping for:
Note
Background or a useful aside.
Tip
A practical shortcut or a debugging habit.
Warning
A pitfall that produces plausible-looking wrong answers.
Important
A platform-specific note (for example, DGX Spark), or something that changes how you should read the rest.
Pacing
At 10 to 12 hours a week, plan on roughly one chapter a week. The GPU chapters and the capstone take two. The full eight-part path is about a year, with something working at the end of every week. If you have more time, the milestones are designed to be done in one sitting each.
You can read the whole book without a GPU. The reference milestones have CPU paths, including Triton’s interpreter; CUDA kernels, graph capture and hardware-specific adapters need their target device. A GPU makes the performance chapters far more satisfying, though. The book was developed for an NVIDIA DGX Spark (GB10), and notes for that machine are marked.
Where this book comes from
The book draws on three excellent resources, and borrows from them freely:
- Sebastian Raschka, Build a Large Language Model (From Scratch) (Manning, 2024), cited as BALLM. Parts I and II follow its path from tokens to a trained GPT and reuse its corpus and several of its examples.
- Wen-mei Hwu, David Kirk and Izzat El Hajj, Programming Massively Parallel Processors, 5th edition, cited as PMPP. Part III follows its approach to CUDA, tiling, reductions and attention.
- The GPU Mode lecture series (github.com/gpu-mode/lectures), cited as Ln, for profiling, Triton, FlashAttention, quantization and serving.
You never need them open beside this book, but each chapter ends with exact pointers for a second explanation. Appendix E maps every chapter to its sources.
Start with Setup.
0. Setup: your workspace and engine
In this chapter
- Create a personal workspace with the book's code, a Python environment and PyTorch.
- Learn the loop you will repeat every chapter: run a test, implement, run it again.
- Optionally prepare the GPU, C++ and Rust toolchains.
Time: about 1 hour. GPU: optional.
The workspace
Everything you do happens in one folder: your copy of the book’s code. It contains:
inference-zero-to-hero-code/
├── engine/ YOUR engine. Every function you'll write is there, with its signature
│ and docstring; the body says raise NotImplementedError("TODO(Chapter 5) ...")
├── izh/ the complete reference engine: the answer key, same module and function names
├── tests/ one milestone test file per chapter (test_ch05_attention.py, ...)
├── run.py a runnable demo for each chapter (python run.py --help)
├── data/ training text and small datasets (no downloads needed)
├── rust/ cpp/ the Rust and C++ tracks
└── models/ checkpoints you download later (created when you need it)
engine/ and izh/ have identical structure. That’s deliberate. When you’re stuck, you can open the same file in izh/ and read the reference. When you want to skip ahead, every test and demo can run against the reference instead of your engine:
IZH_IMPL=izh pytest tests/test_ch05_attention.py # test the reference
python run.py attention --impl engine # run a demo with YOUR engine
1. Install uv
uv manages Python versions and packages quickly and reproducibly. These commands are for Linux and macOS (Windows users: use WSL2):
curl -LsSf https://astral.sh/uv/install.sh | sh
export PATH="$HOME/.local/bin:$PATH"
uv --version
uv can download Python for you, so you don’t need a separate Python installation.
2. Create your working copy
Download the code archive (or, from a clone of the book’s repository, use docs/inference-zero-to-hero/src/code-examples.zip) and extract it into a fresh folder:
mkdir -p "$HOME/izh" && cd "$HOME/izh"
unzip -n ~/Downloads/code-examples.zip
cd inference-zero-to-hero-code
ls # engine izh tests run.py data rust cpp ...
unzip -n never overwrites existing files, so re-extracting a newer archive won’t destroy your work. Put the folder under version control right away. Your engine is real code, and you’ll want its history:
git init && git add -A && git commit -m "Starting point"
3. Create the Python environment
uv venv --python 3.12
source .venv/bin/activate # do this in every new terminal
Now install PyTorch. Choose one route.
# Let uv pick the CUDA build that matches your driver, then verify below.
uv pip install torch numpy pytest --torch-backend=auto
uv pip install torch numpy pytest --torch-backend=cpu
# macOS on M-series: PyTorch's default build includes the MPS backend.
uv pip install torch numpy pytest
Check that it worked, and that a GPU is really doing work if you have one:
python - <<'PY'
import torch
print("torch", torch.__version__, "| CUDA available:", torch.cuda.is_available())
if torch.cuda.is_available():
print(torch.cuda.get_device_name(), torch.cuda.get_device_capability())
x = torch.ones(4, device="cuda")
print((x + x).cpu()) # tensor([2., 2., 2., 2.])
PY
Important
DGX Spark / GB10. Spark is an ARM64 machine with a Blackwell GPU of compute capability 12.1 (
sm_121). Use a PyTorch wheel built foraarch64with CUDA 13.--torch-backend=autonormally finds it. CPU and GPU share one pool of LPDDR5x memory (128 GB), which matters for the capacity planning in Chapters 18, 20 and 30.
4. The loop you will repeat every chapter
Run the first milestone’s tests. They fail, because you haven’t written anything yet:
pytest tests/test_ch03_autograd.py
E NotImplementedError: TODO(Chapter 3): implement __add__ in engine/autograd.py
That message is your to-do list. Each chapter’s Build it section tells you which functions to implement and gives hints. You edit engine/<module>.py, rerun the tests, and repeat until they’re green:
pytest tests/test_ch03_autograd.py # 5 passed
git commit -am "Chapter 3 milestone"
Then see your code do something:
python run.py autograd --impl engine
Tip
pytest -xstops at the first failure;pytest -k nameruns only matching tests;pytest --pdbdrops into a debugger where a test fails. The tests are short and readable, so open them: they are the specification.
Try it now with the reference, to confirm your environment works end to end:
IZH_IMPL=izh pytest -q # ~126 passed, 12 skipped (the GPU-only tests) on a CPU
python run.py tensors
5. Packages you’ll add later
Only PyTorch is needed until Chapter 9. Later chapters ask you to install more, with the command shown at the point of use:
| From | Packages | Why |
|---|---|---|
| Chapter 9 | uv pip install -r optional-requirements.txt | tokenizers and chat templates (transformers), comparing your models with independent implementations |
| Chapter 11 | CUDA Toolkit, uv pip install ninja setuptools | compiling your own CUDA kernels |
| Chapter 14 | triton (included with CUDA PyTorch; uv pip install triton on CPU) | writing kernels in Python |
| Chapter 21 | uv pip install -r workflow-requirements.txt | fine-tuning with PEFT |
6. Models you’ll download later
The engine never downloads anything by itself. When a chapter needs a real checkpoint, it shows an explicit command like this one, which saves a pinned snapshot into models/:
uv pip install -r optional-requirements.txt
hf download Qwen/Qwen3-0.6B --local-dir models/Qwen3-0.6B
| Chapter | Model | Disk | Purpose |
|---|---|---|---|
| 1, 17-23 | Qwen/Qwen3-0.6B | 1.5 GB | the dense model your engine runs |
| 9 | openai-community/gpt2 | 0.5 GB | your from-scratch GPT, loaded with real weights |
| 26 | Qwen/Qwen3-1.7B (optional) | 4 GB | a target for speculative decoding |
| 27 | Qwen/Qwen3-30B-A3B (optional, quantized) | 17-60 GB | a real mixture of experts |
| 30 | Qwen/Qwen3.8-Flash-Next (optional) | about 360 GB | the capstone at full scale |
The optional large models are never required to pass a milestone. Every milestone is verified on small configurations, against independent reference implementations.
7. Optional: C++ and Rust
Every core mechanism in the book also exists in C++ and Rust, shown in tabs. The Rust track grows into a CPU engine that loads real Qwen3 weights. To follow along:
# C++17 compiler and CMake (Ubuntu: sudo apt install build-essential cmake)
cmake -S cpp -B build/cpp -DCMAKE_BUILD_TYPE=Release
cmake --build build/cpp -j
build/cpp/izh all # self-checking demos, one per chapter
# Stable Rust from https://rustup.rs
cd rust
cargo test --release # unit tests + parity with the Python reference
cargo run --release -- generate --model-dir ../models/Qwen3-0.6B --new-tokens 32
Appendix B describes both tracks and the optional NVIDIA Rust GPU toolchains.
When something goes wrong
| Symptom | Fix |
|---|---|
No module named torch | source .venv/bin/activate in this terminal; check which python |
torch.cuda.is_available() is False on a GPU machine | Reinstall with --torch-backend=cu130 (or your CUDA version); check nvidia-smi |
No module named engine when running a script | Run commands from the workspace root, the folder that contains engine/ |
| Triton tests are slow on a CPU | Expected: the interpreter is about 100x slower than a GPU. Use pytest -k to run one |
nvcc missing (Chapter 11) | Install the CUDA Toolkit; a CUDA PyTorch wheel does not include the compiler |
| Out of memory | Use smaller batch or context; on Spark, check other processes with nvidia-smi |
Record uv pip freeze and nvidia-smi output alongside any result you want to keep. Next: the big picture.
1. The big picture: what happens when you press Enter
In this chapter
- Run a real language model on your machine in ten lines of Python.
- Take it apart: tokens, embeddings, layers, logits, sampling, and the loop that ties them together.
- Meet the two phases of inference, prefill and decode, and the one number that limits decode speed.
- Get a map of the engine you'll build and where each part of the book fits.
You will build
No engine code yet. You'll measure your machine's generation speed and predict it from first principles.
Time: 2-3 hours. GPU: optional.
A model in ten lines
Let’s start with the destination and work backwards. Install the libraries that read real checkpoints, and download a small but genuinely capable model, Qwen3-0.6B (about 1.5 GB):
uv pip install -r optional-requirements.txt
hf download Qwen/Qwen3-0.6B --local-dir models/Qwen3-0.6B
Now ask it something. Create experiments/ch01.py in the code root (the directory containing run.py), save the following example there, and run python -m experiments.ch01 from that root. Use this same script for the measurements and stretch exercises below:
from transformers import AutoModelForCausalLM, AutoTokenizer
import torch
path = "models/Qwen3-0.6B"
tokenizer = AutoTokenizer.from_pretrained(path)
model = AutoModelForCausalLM.from_pretrained(path, dtype=torch.bfloat16).eval()
messages = [{"role": "user", "content": "Explain what a GPU is in one sentence."}]
text = tokenizer.apply_chat_template(messages, tokenize=False, add_generation_prompt=True, enable_thinking=False)
ids = tokenizer(text, return_tensors="pt").input_ids
out = model.generate(ids, max_new_tokens=40, do_sample=False)
print(tokenizer.decode(out[0, ids.shape[1]:], skip_special_tokens=True))
You’ll get an answer along these lines (exact wording varies with library version and hardware):
A GPU (Graphics Processing Unit) is a specialized processor designed to rapidly perform
many calculations in parallel, originally for rendering graphics.
That’s the whole user-facing experience. A library you didn’t write did the work: model.generate. By Chapter 18, that library will be replaced by your engine, and you’ll know what every line of it does and why. The rest of this chapter opens the box.
Tip
No download possible right now? Everything below also works with the tiny random model in the reference package:
python run.py cacheandpython run.py fastexercise the same machinery. The text will be gibberish, but the mechanics are identical.
Step 1: text becomes token IDs
A model never sees characters. A tokenizer chops text into pieces from a fixed vocabulary and replaces each piece with its integer ID:
print(ids[0].tolist())
print([tokenizer.decode([i]) for i in ids[0]])
The output looks like this (middle IDs elided):
[151644, 872, 198, ..., 151645, 198, 151644, 77091, 198]
['<|im_start|>', 'user', '\n', 'Ex', 'plain', ' what', ' a', ' GPU', ' is', ..., '<|im_end|>', '\n', '<|im_start|>', 'assistant', '\n']
Notice three things. Common words are single tokens (' GPU', with its leading space). Rarer words may be split into pieces. And the chat template added special tokens like <|im_start|> (ID 151644) and <|im_end|> (ID 151645) that mark who is speaking. Qwen3’s vocabulary has 151,936 entries. You’ll build a tokenizer like this, trained on a short novel, in Chapter 4.
Step 2: IDs become vectors
The model’s first layer is a lookup table, the embedding, with one row of 1,024 numbers per vocabulary entry. Token 22670 (' GPU') becomes row 22670. The whole prompt becomes a matrix with one row per token: shape [tokens, 1024]. These rows are learned during training so that tokens used in similar ways end up with similar vectors.
Step 3: the vectors flow through layers
Qwen3-0.6B has 28 transformer layers. Each one does two things to every token vector:
- Attention lets each token gather information from the tokens before it. That’s how
' sentence'comes to “know” it’s part of a request about GPUs. (Chapter 5) - The MLP transforms each token’s vector on its own, applying what the model learned. (Chapter 6)
Each layer adds its result to the vector it received (the residual stream), so information accumulates as it flows upward. After 28 layers, each position holds a vector that summarizes everything relevant so far.
Step 4: the last vector becomes a prediction
The final vector, at the last position, is multiplied by a [1024 × 151936] matrix (for Qwen3-0.6B, the same matrix as the embedding table, “tied”). That produces 151,936 scores called logits, one per vocabulary entry: how plausible each token is as the next one.
with torch.no_grad():
logits = model(ids).logits[0, -1] # scores for the token after the prompt
probs = torch.softmax(logits.float(), dim=-1)
top = probs.topk(5)
for p, i in zip(top.values, top.indices):
print(f"{tokenizer.decode([i])!r:12} {p:.3f}")
A typical result (your probabilities will differ a little):
'A' 0.91
'GPU' 0.05
'The' 0.02
...
Softmax turned scores into probabilities that sum to 1. Here the model is quite sure the answer starts with “A”.
Step 5: choose, append, repeat
Sampling picks one token from that distribution. Always taking the most likely one is called greedy decoding. Then comes the trick that makes generation work: append the chosen token to the input and run the model again. Each step produces exactly one new token. A 40-token answer takes 40 trips through all 28 layers. This loop is the autoregressive generation loop, and it’s the heart of inference:
ids_so_far = ids
for _ in range(40):
with torch.no_grad():
logits = model(ids_so_far).logits[0, -1]
next_id = logits.argmax() # greedy
ids_so_far = torch.cat([ids_so_far, next_id.view(1, 1)], dim=1)
if next_id == tokenizer.convert_tokens_to_ids("<|im_end|>"):
break
Play with the whole pipeline here:
Two phases with very different costs
The loop above is wasteful. Every iteration re-processes the whole prompt, even though nothing about the earlier tokens has changed. Real engines remember each layer’s intermediate results for earlier tokens (the KV cache, Chapter 16), which splits generation into two phases:
- Prefill. Process all prompt tokens at once, filling the cache. Hundreds or thousands of tokens flow through each layer together, as one big matrix multiplication. GPUs are superb at this, and prefill is limited by arithmetic speed (it’s compute-bound).
- Decode. Produce one token at a time. Each step pushes a single token vector through every layer, reusing the cache. There’s little arithmetic per step, but every weight of the model must still be read from memory. Decode is memory-bound.
The user sees these as two numbers. Time to first token (TTFT) is mostly prefill. Inter-token latency, how fast the words stream, is decode.
The most important number in this book
Here is the back-of-envelope estimate that drives almost every design decision in modern inference engines.
To produce one token, decode multiplies one vector by every weight matrix in the model, so it must read every weight from memory once. Qwen3-0.6B has about 0.6 billion parameters. In BF16, each parameter takes 2 bytes, so each token requires reading about 1.2 GB.
How fast can memory deliver bytes? That’s the memory bandwidth of your device:
| Device | Memory bandwidth | Upper bound for Qwen3-0.6B (1.2 GB/token) |
|---|---|---|
| Laptop CPU (DDR5, 2 channels) | ~80 GB/s | ~65 tokens/s |
| DGX Spark (GB10, LPDDR5x) | 273 GB/s | ~230 tokens/s |
| RTX 4090 (GDDR6X) | 1,008 GB/s | ~840 tokens/s |
| H100 SXM (HBM3) | 3,350 GB/s | ~2,800 tokens/s |
$$ \text{tokens per second} ;\lesssim; \frac{\text{memory bandwidth (bytes/s)}}{\text{bytes read per token}} $$
That’s a ceiling: real engines reach 60-90% of it. Compare it with the arithmetic. One token needs about 2 floating-point operations per parameter, roughly 1.2 GFLOP. Even a laptop CPU does that in a few milliseconds, and a GPU in microseconds. During decode, the processor mostly waits for memory.
This one inequality explains a remarkable amount of the field:
- Quantization (Chapter 20) stores weights in 4 bits instead of 16. That’s 4x fewer bytes per token, so decode gets close to 4x faster.
- Batching (Chapter 24) decodes many users’ tokens in the same step. The weights are read once and used for every user, so throughput rises almost for free.
- Speculative decoding (Chapter 26) checks several guessed tokens in one pass over the weights.
- Mixture of experts (Chapter 27) reads only a few experts per token. Flash-Next stores 176B parameters but reads about 6B per token.
- KV-cache size (Chapters 16, 25, 28) adds to the bytes per token, which is why long contexts are slow and why linear attention exists.
What an inference engine does
An inference engine is the program between “a folder of weights” and “a fast, correct stream of tokens for many users”. Its jobs, and where you’ll build each one:
| Job | What it means | Where |
|---|---|---|
| Represent tensors and compute | matrix multiplications, attention, norms | Parts I-II |
| Load weights | read checkpoint files, map tensor names, check shapes | Chapters 9, 18 |
| Run the model correctly | exactly the architecture the weights were trained for | Chapters 6, 17, 27-30 |
| Manage state | KV cache, recurrent state, positions | Chapters 16, 25, 28 |
| Run it fast | kernels, fusion, CUDA graphs, quantization | Part III, Chapters 19-20 |
| Choose tokens | greedy, temperature, top-p, stop rules | Chapter 8 |
| Serve many requests | batching, scheduling, memory sharing, speculation | Part VI |
Production engines like vLLM, SGLang and llama.cpp are hundreds of thousands of lines, mostly hardware-specific kernels and model adapters. The core ideas fit in a few thousand lines, and those are the lines you’ll write.
The destination: Qwen3.8-Flash-Next
The capstone model is a deliberate stretch. Qwen3.8-Flash-Next (released August 2026) is a preview of the architecture behind Qwen4. Its 48 layers mix:
- Gated DeltaNet linear attention in 36 layers, which keeps a fixed-size memory instead of a growing KV cache (Chapter 28);
- Qwen Sparse Attention in the other 12 layers, where a small indexer picks which 2,048 past tokens each query reads (Chapter 29);
- a 512-expert mixture in every layer, 10 experts per token plus a shared one (Chapter 27);
- four parallel residual streams with gated reads and writes, and a 51-billion-parameter hashed n-gram memory (Chapter 30).
Every one of those components is a variation on ideas you’ll learn first in their simplest form. By Chapter 30, you’ll write the whole model from scratch and verify it against the official implementation.
Build it
There’s no engine code this week. Instead, build your intuition with a measurement.
- Time the Hugging Face generation above for 100 new tokens (call
torch.cuda.synchronize()before stopping the clock if you’re on a GPU), and compute tokens per second. - Look up your device’s memory bandwidth and compute the ceiling from the formula.
- Compute the ratio. Below 50% is normal for this naive loop. Chapter 19 gets you much closer to the ceiling.
- Run the reference engine’s demos to see prefill and decode separately:
python run.py cache # cached vs uncached generation time as output length grows
python run.py fast # tokens/s of the plain loop vs a decoder without host round trips
Stretch exercises
- ★ Change
do_sample=Falsetodo_sample=True, temperature=1.2. Run it three times. What changed, and what didn’t? Where: copy the example in A model in ten lines intoexperiments/ch01.py(create it); change itsmodel.generatecall. - ★★ Measure TTFT separately from decode speed by timing a 1-token generation and a 101-token one with a 500-token prompt. Which phase dominates for short answers? For long ones? Where: the same
experiments/ch01.py(create it), aroundmodel.generate. - ★★ Load the model in FP32 (
dtype=torch.float32) and repeat the speed measurement. Predict the slowdown from the bandwidth formula before you run it. Where: the same script’sAutoModelForCausalLM.from_pretrainedcall.
Check your understanding
- Why does generating a 40-token answer require at least 40 sequential passes through the model?
- Prefill of a 1,000-token prompt and decode of one token both read every weight once. Why is prefill not 1,000x slower than one decode step?
- A 7B-parameter model in BF16 runs on a 1 TB/s GPU. What’s the most tokens per second a single user can get? What if the weights were 4-bit?
Going deeper
- BALLM Chapter 1 (pp. 1-16): what LLMs are and how they’re built and trained.
- PMPP §20.1 (pp. 478-482): the decoder-only transformer from a systems angle, including the generation loop.
- GPU Mode L1 (profiling and integrating CUDA kernels in PyTorch) for a first look at where the time goes.
- Horace He, Making Deep Learning Go Brrrr From First Principles (2022): the compute-bound / memory-bound / overhead-bound framing used throughout this book.
2. Tensors: the data structure of deep learning
In this chapter
- What a tensor is: a flat block of memory plus a shape, a dtype, a device and strides.
- Elementwise operations, broadcasting and reductions, the vocabulary of every model.
- Matrix multiplication as "many dot products", and why it's the operation that matters most.
- Views versus copies, and why a correct shape can still be the wrong memory layout.
You will build
engine/tensors.py: stride arithmetic, the broadcasting rule, a matrix multiplication with explicit loops, and the linear-layer convention used by every checkpoint.
Time: 3-5 hours. GPU: not needed.
Why start here
Every number a language model touches lives in a tensor: the weights you load from disk, the token vectors flowing through layers, the attention scores, the KV cache. Every bug you’ll hit while building an engine is, at bottom, a tensor bug: the wrong axis, the wrong layout, the wrong dtype, or the wrong device. An hour spent getting fluent here saves days later.
A tensor is an array with four properties
Run this in your environment (python inside the activated venv):
import torch
x = torch.arange(24.0).reshape(2, 3, 4)
print(x.shape, x.dtype, x.device, x.stride())
torch.Size([2, 3, 4]) torch.float32 cpu (12, 4, 1)
Four properties describe this tensor completely:
- Shape
[2, 3, 4]: two blocks, each with three rows of four numbers. In language models the most common shape is[B, T, D]: a Batch of sequences, each with T tokens, each token a vector of D numbers. Ourxcould be two sentences of three tokens, each token a 4-number vector. - Dtype
float32: how each number is stored. FP32 takes 4 bytes, BF16 and FP16 take 2, INT8 takes 1. Token IDs are integers (int64by default). The dtype decides memory size, and as Chapter 1 showed, memory size decides decode speed. - Device
cpu: where the data lives and where operations on it run.x.to("cuda")copies it to GPU memory. Operations between tensors on different devices are errors, not silent copies. - Strides
(12, 4, 1): how to find an element in memory. That needs its own section.
Note
Axis meaning lives in your program, not in the tensor. PyTorch doesn’t know that axis 1 is “tokens”. Writing shapes in comments, like
# [B, T, D], is the single most effective habit for avoiding bugs in model code. This book does it everywhere.
Memory is flat: strides
RAM is one long line of bytes. A tensor stores its elements in a flat buffer, and the strides say how far to jump in that buffer to move one step along each axis. For our [2, 3, 4] tensor stored row by row, moving one step along the last axis moves 1 element, along the middle axis 4 elements (one row), and along the first 12 (one whole 3×4 block). The element at index (i, j, k) lives at
$$ \text{offset}(i, j, k) = \text{storage_offset} + 12i + 4j + 1k . $$
So x[1, 2, 3] is at 12 + 8 + 3 = 23, the last element. Strides like these, where the last axis moves fastest, are called contiguous or row-major.
Here’s why strides matter: many operations change only the strides, not the data. Transposing swaps two axes by swapping their sizes and strides. No numbers move:
t = x.transpose(1, 2)
print(t.shape, t.stride(), t.is_contiguous())
torch.Size([2, 4, 3]) (12, 1, 4) False
t is a view: another way of looking at the same memory. Views are free, which is why PyTorch uses them so much. The catch is that code which assumes row-major layout will read a view wrongly. A hand-written GPU kernel that computes row * width + col will happily read the wrong numbers from t, and the shapes will all look right. This exact bug shows up when loading checkpoints (Chapter 9) and writing kernels (Part III).
Some operations need contiguous memory. view refuses to reinterpret a non-contiguous tensor; reshape copies when it has to; contiguous() always produces a compact copy if needed:
t.view(2, 12) # RuntimeError: view size is not compatible with input tensor's size and stride
t.reshape(2, 12) # works: silently copies
t.contiguous().stride() # (12, 3, 1): a fresh row-major copy
A copy costs a full read and write of the tensor. In the performance chapters you’ll learn to spot hidden .contiguous() copies in a profile.
The same idea in C++ and Rust: a tensor struct is just a buffer, a shape and strides, and a transpose swaps two entries:
def contiguous_strides(shape):
strides, step = [], 1
for size in reversed(shape):
strides.append(step)
step *= size
return tuple(reversed(strides))
def element_offset(index, strides, storage_offset=0):
return storage_offset + sum(i * s for i, s in zip(index, strides))
// A tensor is a flat buffer plus shape plus strides; element (i, j, ...) is at
// offset + i*stride[0] + j*stride[1] + ...
struct Tensor {
std::vector<float> data;
std::vector<size_t> shape, strides;
size_t offset = 0;
static std::vector<size_t> contiguous(const std::vector<size_t>& shape) {
std::vector<size_t> s(shape.size(), 1);
for (size_t a = shape.size(); a-- > 1;) s[a - 1] = s[a] * shape[a];
return s;
}
Tensor(std::vector<float> d, std::vector<size_t> sh)
: data(std::move(d)), shape(sh), strides(contiguous(sh)) {}
float at(const std::vector<size_t>& index) const {
size_t flat = offset;
for (size_t a = 0; a < index.size(); ++a) flat += index[a] * strides[a];
return data[flat];
}
Tensor transpose(size_t a, size_t b) const { // swap sizes and strides; no data moves
Tensor t = *this;
std::swap(t.shape[a], t.shape[b]);
std::swap(t.strides[a], t.strides[b]);
return t;
}
};
#![allow(unused)]
fn main() {
/// A view over a flat f32 buffer. Element (i0, i1, ...) lives at
/// offset + i0*strides[0] + i1*strides[1] + ...
#[derive(Clone, Debug)]
pub struct Tensor {
pub data: Vec<f32>,
pub shape: Vec<usize>,
pub strides: Vec<usize>,
pub offset: usize,
}
impl Tensor {
/// Row-major ("contiguous") strides: the last axis moves fastest.
pub fn contiguous_strides(shape: &[usize]) -> Vec<usize> {
let mut strides = vec![1; shape.len()];
for axis in (0..shape.len().saturating_sub(1)).rev() {
strides[axis] = strides[axis + 1] * shape[axis + 1];
}
strides
}
pub fn from_vec(data: Vec<f32>, shape: &[usize]) -> Self {
assert_eq!(data.len(), shape.iter().product::<usize>(), "data length must match the shape");
Tensor { data, strides: Self::contiguous_strides(shape), shape: shape.to_vec(), offset: 0 }
}
pub fn get(&self, index: &[usize]) -> f32 {
let flat: usize = index.iter().zip(&self.strides).map(|(i, s)| i * s).sum();
self.data[self.offset + flat]
}
/// Swapping two axes swaps their sizes and strides. No data moves.
pub fn transpose(&self, a: usize, b: usize) -> Self {
let mut t = self.clone();
t.shape.swap(a, b);
t.strides.swap(a, b);
t
}
pub fn is_contiguous(&self) -> bool {
self.strides == Self::contiguous_strides(&self.shape)
}
}
}
Elementwise operations and broadcasting
Arithmetic between same-shape tensors works element by element: a + b, a * b, torch.exp(a). More interesting is what happens when shapes differ. Broadcasting lets a smaller tensor be reused across a larger one without copying:
b = torch.tensor([10., 20., 30., 40.]) # shape [4]
print((x + b)[0, 0]) # the same b added to every 4-vector in x
tensor([10., 21., 32., 43.])
The rule: align shapes from the right. Two sizes are compatible if they’re equal or one of them is 1; a missing axis counts as 1. The result takes the larger size on each axis. So [2, 3, 4] + [4] works (bias added to every token vector), and so does [3, 4] * [3, 1] (each row scaled by its own number):
s = torch.tensor([[1.], [2.], [3.]]) # shape [3, 1]
print(x[0] * s) # row i multiplied by s[i]
tensor([[ 0., 1., 2., 3.],
[ 8., 10., 12., 14.],
[24., 27., 30., 33.]])
But [2, 3, 4] + [2, 3] fails: aligned from the right, 4 meets 3. If you meant “one number per token”, the scalar tensor must be [2, 3, 1]. Use s.unsqueeze(-1) or s[..., None] to add that axis. You’ll write this rule yourself in the milestone.
Warning
Broadcasting can also succeed when you didn’t intend it.
[T, 1] - [T]gives a[T, T]matrix, not a[T]vector. If a loss value or an attention matrix suddenly has an unexpected extra axis, an accidental broadcast is the usual suspect. Assert shapes.
Reductions
A reduction collapses an axis: sum, mean, max, argmax, softmax’s denominator. You name the axis you want to collapse:
print(x.sum(dim=-1)) # sum each 4-vector: [2, 3, 4] -> [2, 3]
print(x.mean(dim=(0, 1))) # average over batch and tokens: [2, 3, 4] -> [4]
tensor([[ 6., 22., 38.],
[54., 70., 86.]])
tensor([10., 11., 12., 13.])
keepdim=True keeps the collapsed axis with size 1, which is exactly what you need to broadcast the result back: x - x.mean(-1, keepdim=True) centers every vector. That one line is the first half of LayerNorm (Chapter 6).
Matrix multiplication is many dot products
The dot product of two vectors multiplies them element by element and sums:
$$ [1, 2, 3] \cdot [4, 5, 6] = 1\cdot 4 + 2\cdot 5 + 3\cdot 6 = 32 . $$
A matrix product computes a dot product for every (row of A, column of B) pair:
$$ C_{ij} = \sum_{k} A_{ik} B_{kj}, \qquad [M, K] ;@; [K, N] ;\rightarrow; [M, N]. $$
The shared dimension K is summed away. That’s the shape rule to check first whenever a matmul fails. The work is $M \cdot N \cdot K$ multiply-adds, which we count as $2MNK$ floating-point operations (FLOPs). A transformer spends almost all of its arithmetic here.
def matmul_loops(a, b):
m, k = a.shape
k2, n = b.shape
assert k == k2
c = torch.zeros(m, n)
for i in range(m):
for j in range(n):
c[i, j] = sum(a[i, p] * b[p, j] for p in range(k))
return c
// C[M,N] = A[M,K] B[K,N], row-major. The i-k-j order streams rows of B and C sequentially.
inline std::vector<float> matmul(const std::vector<float>& a, const std::vector<float>& b, size_t m, size_t k, size_t n) {
std::vector<float> c(m * n, 0.f);
for (size_t i = 0; i < m; ++i)
for (size_t p = 0; p < k; ++p) {
float aip = a[i * k + p];
for (size_t j = 0; j < n; ++j) c[i * n + j] += aip * b[p * n + j];
}
return c;
}
#![allow(unused)]
fn main() {
/// C[m][n] = sum_k A[m][k] * B[k][n] for row-major A [M,K] and B [K,N].
/// The i-k-j loop order walks B and C along rows, which keeps memory access sequential.
pub fn matmul(a: &[f32], b: &[f32], m: usize, k: usize, n: usize) -> Vec<f32> {
let mut c = vec![0.0f32; m * n];
for i in 0..m {
for p in 0..k {
let a_ip = a[i * k + p];
let b_row = &b[p * n..(p + 1) * n];
let c_row = &mut c[i * n..(i + 1) * n];
for (c_ij, b_pj) in c_row.iter_mut().zip(b_row) {
*c_ij += a_ip * b_pj;
}
}
}
c
}
}
The Python triple loop is about a million times slower than a @ b. The C++ and Rust versions reorder the loops (i, then k, then j) so the inner loop walks memory sequentially. That’s your first glimpse of a theme that runs through Part III: the same arithmetic, arranged to respect memory, runs far faster.
Batched matmul and the linear layer
With more than two axes, @ treats the leading axes as a batch and multiplies the last two:
print((x @ torch.ones(4, 5)).shape) # [2, 3, 4] @ [4, 5] -> [2, 3, 5]
torch.Size([2, 3, 5])
One weight matrix transforms every token vector in every sequence. That’s a linear layer, the most common operation in a transformer. One convention matters enormously when loading real weights: PyTorch’s nn.Linear(in, out) stores its weight as [out, in] and computes
$$ y = x,W^{\top} + b . $$
Some checkpoints (GPT-2’s original format, Chapter 9) store [in, out] instead. For square matrices both orientations have the same shape, so a missing transpose passes every shape check and produces garbage. You’ll meet this bug in Chapter 9 and will be glad you know about it.
How big is a tensor?
Bytes = number of elements × bytes per element:
print(torch.tensor(1.0).element_size(), torch.tensor(1.0, dtype=torch.bfloat16).element_size()) # 4 2
Qwen3-0.6B’s embedding table is [151936, 1024]: 155.6 million numbers, which is 311 MB in BF16 and 622 MB in FP32. Get in the habit of doing this multiplication for every big tensor you meet. Memory capacity decides what fits; memory bandwidth decides how fast it runs.
Devices and the GPU
x.to("cuda") copies a tensor into GPU memory and returns the copy; operations on it then run on the GPU. Two facts to keep in mind until Chapter 10 explains them properly:
- GPU operations are asynchronous.
y = x @ xon CUDA returns immediately; the GPU works in the background. Anything that needs the value on the CPU (print(y),y.item(),y.tolist()) waits. Naive timing code therefore measures nothing useful. - Copies between CPU and GPU are slow compared with GPU memory itself. Keep data on the device; move only small results back.
Important
On DGX Spark the CPU and GPU share physical memory, but PyTorch still treats them as separate devices with separate allocations.
x.to("cuda")is still a copy, and.cpu()still synchronizes.
Build it
Engine milestone 2: tensor rules by hand. Implement in engine/tensors.py:
contiguous_strides(shape): row-major strides, e.g.(2, 3, 4) -> (12, 4, 1).element_offset(index, strides, storage_offset): where an element lives in flat storage.broadcast_shapes(a, b): the result shape ofa op b, orValueErrorif incompatible.matmul_loops(a, b): matrix multiplication with explicit loops.linear(x, weight, bias):nn.Linear’s rule for any number of leading axes.
pytest tests/test_ch02_tensors.py
The tests compare your offsets against real PyTorch views (transposes, slices with steps, permutes), so if they pass, you understand strides.
Tip
For
broadcast_shapes, walk the axes from the right with an indexi = 1, 2, ...and treat a missing axis as size 1. Forlinear, one line is enough: the leading axes broadcast through@on their own.
Stretch exercises
- ★ Predict, then check: the shapes of
torch.randn(8, 1, 6, 1) + torch.randn(7, 1, 5),torch.randn(2, 3).T @ torch.randn(2, 3), andtorch.randn(4, 5)[:, None, :] - torch.randn(4, 5)[None, :, :]. Where:experiments/ch02.py(create it). - ★★ Write
transpose_strides(shape, strides, a, b)andslice_view(shape, strides, offset, axis, start, step)that return the new (shape, strides, offset) without touching data. Verify againstx.transposeandx[:, 1::2]. Where: add both helpers toengine/tensors.py. - ★★ Time
matmul_loopsagainst@for 64×64 matrices. How many times slower is it? Usetorch.utils.benchmark.Timer. Where:experiments/ch02.py(create it), importingengine.tensors.matmul_loops. - ★★★ Implement
matmul_loopsfor any batched shapes[..., M, K] @ [..., K, N]with broadcasting on the leading axes, using yourbroadcast_shapes. Where: extendmatmul_loopsinengine/tensors.py.
Check your understanding
- Why does
[2, 3, 4] @ [4, 5]produce[2, 3, 5]? - A tensor has shape
[3, 2]and strides(1, 3). Is it contiguous? Which operation probably produced it? - Why can a missing transpose of a square weight matrix pass every shape check?
- How many bytes does a
[28, 2, 8, 4096, 128]BF16 tensor take? (It’s a KV cache you’ll meet in Chapter 16.)
Going deeper
- BALLM Appendix A, §§A.1-A.5 (pp. 251-271): PyTorch tensors and autograd for newcomers.
- PyTorch documentation, Tensor Views and Broadcasting semantics.
- Edward Yang, PyTorch internals (blog, 2019): how strides, storage and dispatch actually work inside PyTorch.
3. How networks learn: gradients, autograd, optimizers
In this chapter
- What "learning" means for a neural network: parameters, a loss, and a rule for improving them.
- Derivatives as slopes, the chain rule, and backpropagation through a computation graph.
- Why networks need nonlinearities, and what a multilayer perceptron is.
- The training loop, and the optimizers inside it: SGD, momentum, Adam and AdamW.
You will build
engine/autograd.py: a 100-line automatic differentiation engine that trains a small neural network.
Time: 5-7 hours. GPU: not needed.
Why an inference book teaches training
You’ll mostly run models in this book. But every weight you load was produced by the process in this chapter, and Parts II and V make you train and fine-tune models yourself. More immediately, the backward-pass machinery you build here is exactly what PyTorch runs when you call loss.backward(). Engine builders who know it can read memory reports, debug fine-tuning, and understand why training needs so much more memory than inference.
Learning is adjusting knobs to reduce a number
A neural network is a function with adjustable numbers inside it, called parameters or weights. Feed it an input, and it produces an output. A loss function turns “how wrong was that output?” into a single number. Learning means changing the parameters to make that number smaller.
The smallest possible example has one parameter w. It predicts w * x, and we want the prediction for x = 3 to be 8. The squared-error loss is
$$ L(w) = (w \cdot x - \text{target})^2 = (3w - 8)^2 . $$
At w = 2, the prediction is 6 and the loss is $(6-8)^2 = 4$. Should we increase or decrease w, and by how much?
The derivative says which way is downhill
The derivative $dL/dw$ is the slope of the loss as a function of w: how much the loss changes per small change in w. For our loss,
$$ \frac{dL}{dw} = 2x(wx - \text{target}) = 2 \cdot 3 \cdot (6 - 8) = -12 . $$
The slope is negative, so increasing w decreases the loss. Gradient descent takes a small step against the slope:
$$ w \leftarrow w - \eta \frac{dL}{dw} $$
where $\eta$ (eta) is the learning rate. With $\eta = 0.01$, the first steps are:
| step | w | loss | dL/dw |
|---|---|---|---|
| 0 | 2.0000 | 4.0000 | −12.0000 |
| 1 | 2.1200 | 2.6896 | −9.8400 |
| 2 | 2.2184 | 1.8085 | −8.0688 |
| 3 | 2.2991 | 1.2160 | −6.6164 |
| 5 | 2.4195 | 0.5498 | −4.4489 |
The loss falls, and the steps get smaller as the slope flattens near the minimum at w = 8/3. Try a learning rate twelve times larger ($\eta = 0.12$) and the loss grows: w goes 2.0 → 3.44 → 1.77 → 3.71 and the loss goes 4.0 → 5.38 → 7.24 → 9.75, because every step overshoots. Choosing the learning rate is the first hyperparameter you’ll tune.
Real networks have millions or billions of parameters. The gradient is the vector of all their partial derivatives, $\partial L / \partial w_i$, one per parameter, and gradient descent updates all of them at once. The idea is the same; only the bookkeeping grows.
The chain rule, and computing gradients mechanically
Nobody derives gradients for a billion-parameter model by hand. A network is built from simple operations (add, multiply, exp, …), and each one knows its own local derivative. The chain rule combines them: if $L$ depends on $d$, and $d$ depends on $a$, then
$$ \frac{\partial L}{\partial a} = \frac{\partial L}{\partial d}\cdot\frac{\partial d}{\partial a}. $$
Here’s a small expression worked through by hand. Let $a = 2$, $b = -3$, $c = 10$, $f = -2$, and compute
$$ e = a \cdot b = -6,\qquad d = e + c = 4,\qquad L = d \cdot f = -8 . $$
Work backwards from $L$, multiplying local derivatives as you go:
| node | local rule | gradient $\partial L/\partial(\text{node})$ |
|---|---|---|
| L | (start) | 1 |
| d | $L = d f \Rightarrow \partial L/\partial d = f$ | −2 |
| f | $\partial L/\partial f = d$ | 4 |
| e | $d = e + c \Rightarrow$ pass the gradient through unchanged | −2 |
| c | same as e | −2 |
| a | $e = a b \Rightarrow \partial e/\partial a = b$, times −2 | 6 |
| b | $\partial e/\partial b = a$, times −2 | −4 |
That procedure is backpropagation: a forward pass computes values and remembers how each was made, then a backward pass visits every operation once, in reverse order, multiplying the incoming gradient by the local derivative and passing it on. The cost of the backward pass is about twice the forward pass, regardless of the number of parameters. That’s why training billion-parameter models is possible at all.
Two details matter when you implement it:
- Order. A node’s gradient is complete only after every node that uses it has passed its contribution back. Sorting the graph topologically and walking it in reverse guarantees this.
- Accumulation. If a value is used twice (like
xinx * x + x), contributions from each use add up. Gradients are accumulated with+=, never assigned with=.
Build an autograd engine
The Value class below wraps one number, remembers the values it was computed from, and stores a small closure that applies its local chain-rule step. This design follows Andrej Karpathy’s micrograd. Each operation creates the output, then defines how gradient flows from the output to its inputs:
class Value:
def __init__(self, data, children=(), op=""):
self.data, self.grad = float(data), 0.0
self._children, self._op = tuple(children), op
self._backward = lambda: None
def __mul__(self, other):
other = other if isinstance(other, Value) else Value(other)
out = Value(self.data * other.data, (self, other), "*")
def backward():
self.grad += other.data * out.grad # d(ab)/da = b
other.grad += self.data * out.grad # d(ab)/db = a
out._backward = backward
return out
def backward(self):
order, seen = [], set()
def visit(node): # topological sort
if id(node) not in seen:
seen.add(id(node))
for child in node._children:
visit(child)
order.append(node)
visit(self)
self.grad = 1.0
for node in reversed(order):
node._backward()
// A scalar autograd node. shared_ptr because one value can feed many later operations.
struct Node {
double data, grad = 0;
std::vector<std::shared_ptr<Node>> children;
std::function<void(Node&)> backward = [](Node&) {};
explicit Node(double d) : data(d) {}
};
using Value = std::shared_ptr<Node>;
inline Value make(double d) { return std::make_shared<Node>(d); }
inline Value add(Value a, Value b) {
auto out = make(a->data + b->data);
out->children = {a, b};
out->backward = [a, b](Node& self) { a->grad += self.grad; b->grad += self.grad; };
return out;
}
inline Value mul(Value a, Value b) {
auto out = make(a->data * b->data);
out->children = {a, b};
out->backward = [a, b](Node& self) { a->grad += b->data * self.grad; b->grad += a->data * self.grad; };
return out;
}
inline Value power(Value a, double n) {
auto out = make(std::pow(a->data, n));
out->children = {a};
out->backward = [a, n](Node& self) { a->grad += n * std::pow(a->data, n - 1) * self.grad; };
return out;
}
inline void backward(const Value& root) {
std::vector<Value> order;
std::set<Node*> seen;
std::function<void(const Value&)> visit = [&](const Value& v) {
if (!seen.insert(v.get()).second) return;
for (auto& c : v->children) visit(c);
order.push_back(v);
};
visit(root);
root->grad = 1;
for (auto it = order.rbegin(); it != order.rend(); ++it) (*it)->backward(**it);
}
#![allow(unused)]
fn main() {
#[derive(Clone)]
pub struct Value(Rc<RefCell<Node>>);
struct Node {
data: f64,
grad: f64,
children: Vec<Value>,
// Given this node's gradient, add the chain-rule contribution to each child.
backward: Option<Box<dyn Fn(f64, &[Value])>>,
}
impl Value {
pub fn new(data: f64) -> Self {
Value(Rc::new(RefCell::new(Node { data, grad: 0.0, children: vec![], backward: None })))
}
fn op(data: f64, children: Vec<Value>, backward: impl Fn(f64, &[Value]) + 'static) -> Self {
Value(Rc::new(RefCell::new(Node { data, grad: 0.0, children, backward: Some(Box::new(backward)) })))
}
pub fn data(&self) -> f64 { self.0.borrow().data }
pub fn grad(&self) -> f64 { self.0.borrow().grad }
fn add_grad(&self, g: f64) { self.0.borrow_mut().grad += g; }
pub fn add(&self, other: &Value) -> Value {
Value::op(self.data() + other.data(), vec![self.clone(), other.clone()], |g, c| {
c[0].add_grad(g);
c[1].add_grad(g);
})
}
pub fn mul(&self, other: &Value) -> Value {
Value::op(self.data() * other.data(), vec![self.clone(), other.clone()], |g, c| {
let (a, b) = (c[0].data(), c[1].data());
c[0].add_grad(b * g);
c[1].add_grad(a * g);
})
}
pub fn powf(&self, n: f64) -> Value {
Value::op(self.data().powf(n), vec![self.clone()], move |g, c| {
c[0].add_grad(n * c[0].data().powf(n - 1.0) * g);
})
}
pub fn tanh(&self) -> Value {
let t = self.data().tanh();
Value::op(t, vec![self.clone()], move |g, c| c[0].add_grad((1.0 - t * t) * g))
}
/// Visit every node after all nodes that use it, then apply each local rule once.
pub fn backward(&self) {
let mut order = vec![];
let mut seen = HashSet::new();
fn visit(v: &Value, seen: &mut HashSet<usize>, order: &mut Vec<Value>) {
if seen.insert(Rc::as_ptr(&v.0) as usize) {
for child in &v.0.borrow().children {
visit(child, seen, order);
}
order.push(v.clone());
}
}
visit(self, &mut seen, &mut order);
self.0.borrow_mut().grad = 1.0;
for v in order.iter().rev() {
let node = v.0.borrow();
if let Some(f) = &node.backward {
f(node.grad, &node.children);
}
}
}
}
}
With it, the worked example above becomes:
from engine.autograd import Value
a, b, c, f = Value(2.0), Value(-3.0), Value(10.0), Value(-2.0)
L = (a * b + c) * f
L.backward()
print(L.data, a.grad, b.grad, c.grad, f.grad) # -8.0 6.0 -4.0 -2.0 4.0
Neurons, layers and why nonlinearity matters
A neuron computes a weighted sum of its inputs plus a bias, then applies a nonlinear function: $\tanh(w \cdot x + b)$. A layer is many neurons reading the same inputs, which is exactly the linear layer from Chapter 2 followed by an elementwise nonlinearity. A multilayer perceptron (MLP) stacks layers.
Why the nonlinearity? Without it, two linear layers collapse into one: $W_2(W_1 x) = (W_2 W_1)x$, a single matrix. You can stack a hundred linear layers and still only represent linear functions, which can’t even compute XOR. The nonlinearity between layers lets the network bend space. A classic demonstration is XOR, where the output is high when exactly one input is on:
| inputs | target |
|---|---|
| (0, 0) | −1 |
| (0, 1) | +1 |
| (1, 0) | +1 |
| (1, 1) | −1 |
No straight line separates the +1s from the −1s, but a 2-input MLP with two hidden tanh layers of 8 neurons learns it in a few dozen steps. This is the output of python run.py autograd --steps 200 with the reference engine:
{"step": 1, "loss": 5.63926}
{"step": 26, "loss": 2.49858}
{"step": 51, "loss": 0.00013}
{"step": 200, "loss": 0.0}
{"predictions": [-1.0, 1.0, 1.0, -1.0], "targets": [-1.0, 1.0, 1.0, -1.0]}
Transformer MLPs (Chapter 6) are exactly this structure, scaled up to thousands of neurons per layer, with smoother nonlinearities (GELU, SiLU) than tanh.
From scalars to tensors: PyTorch autograd
Your Value engine works on single numbers. PyTorch’s autograd does the same thing on whole tensors, so one node represents a million multiplications, and the local backward rules are themselves tensor operations running on the GPU. The interface is the same idea:
import torch
w = torch.tensor(2.0, requires_grad=True) # track operations on w
loss = (w * 3.0 - 8.0) ** 2
loss.backward() # fills w.grad
print(w.grad) # tensor(-12.)
Three switches control it, and they’re easy to confuse:
requires_grad=Trueon a tensor (parameters have it by default) means “record operations involving me”.torch.no_grad()andtorch.inference_mode()stop recording. Every inference path in your engine runs under one of them. Recording costs memory, because each operation keeps its inputs alive for the backward pass.model.train()/model.eval()switch layer behavior (dropout on or off). They do not turn gradient recording on or off. That’s a common misconception.
Warning
Gradients accumulate across
backward()calls, by design (it lets you sum gradients over several small batches). Forgettingoptimizer.zero_grad()makes every step use the sum of all previous gradients. The loss usually explodes after a few steps.
The training loop
Every training script you’ll ever read has the same five lines at its core:
for batch in data:
loss = loss_fn(model(batch.x), batch.y) # 1. forward
optimizer.zero_grad() # 2. clear old gradients
loss.backward() # 3. backward: fill .grad for every parameter
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) # 4. (optional) cap the gradient size
optimizer.step() # 5. update the parameters
The optimizer decides how to turn gradients into parameter updates.
Optimizers: SGD, momentum, Adam, AdamW
SGD (stochastic gradient descent) is the update rule from earlier, applied to gradients estimated on a random batch of data. It works, but one learning rate must suit every parameter, and noisy gradients make it zigzag.
Momentum keeps a running average of recent gradients and steps along it, which smooths the noise:
$$ m \leftarrow \beta m + (1-\beta) g, \qquad w \leftarrow w - \eta, m . $$
Adam also tracks a running average of squared gradients, $v$, and divides each parameter’s step by $\sqrt{v}$. Parameters with consistently large gradients take smaller steps, and rarely updated parameters take relatively larger ones:
$$ m_t = \beta_1 m_{t-1} + (1-\beta_1)g_t,\quad v_t = \beta_2 v_{t-1} + (1-\beta_2)g_t^2,\quad w \leftarrow w - \eta,\frac{\hat m_t}{\sqrt{\hat v_t} + \epsilon} $$
where $\hat m_t = m_t/(1-\beta_1^t)$ and $\hat v_t = v_t/(1-\beta_2^t)$ correct for starting both averages at zero. Here’s a worked first step with $g_1 = 2$, $\beta_1 = 0.9$, $\beta_2 = 0.999$: $m_1 = 0.2$ and $v_1 = 0.004$, and after correction $\hat m = 2$ and $\hat v = 4$. The step is $\eta \cdot 2/\sqrt 4 = \eta$. Adam’s first step has size equal to the learning rate, whatever the gradient’s scale. (PyTorch agrees: with $\eta = 0.1$, one Adam step moves the weight from 0 to −0.1.)
AdamW adds weight decay, a gentle pull of every weight toward zero, applied separately from the adaptive gradient step: $w \leftarrow (1 - \eta\lambda)w - \eta,\hat m/(\sqrt{\hat v} + \epsilon)$. It’s the default optimizer for transformers.
Two facts matter for engine builders:
- Optimizer state is big. AdamW stores $m$ and $v$ for every parameter, usually in FP32. Training a 0.6B model needs the weights (1.2 GB in BF16), plus gradients (another 1.2 GB), plus $m$ and $v$ (4.8 GB), plus FP32 master weights in mixed precision (2.4 GB), plus activations. Inference needs only the 1.2 GB. Chapter 22 shows how LoRA avoids most of that.
- Schedules. Learning rates usually ramp up from near zero for the first few hundred steps (warmup), then decay (often along a cosine curve). Your training loop in Chapter 7 does both.
Build it
Engine milestone 3: an autograd engine. In engine/autograd.py, implement the Value methods __add__, __mul__, __pow__, exp, log, tanh, relu and backward. Subtraction, division and negation are already written in terms of these. The Neuron, Layer, MLP and sgd_step helpers are provided.
pytest tests/test_ch03_autograd.py
python run.py autograd --impl engine # trains XOR with YOUR engine
The tests check every local derivative against a finite-difference estimate, check that reused nodes accumulate, and check that the MLP learns.
Tip
If a gradient is wrong, compare it with the finite-difference estimate $\big(L(w+\epsilon) - L(w-\epsilon)\big)/2\epsilon$ (
numerical_gradientin the module). It’s slow but independent of your backward code, which makes it the most useful debugging tool for gradients.
Stretch exercises
- ★ Add
sigmoidtoValue. Its derivative is $\sigma(x)(1-\sigma(x))$. Verify it withnumerical_gradient. Where: addValue.sigmoidinengine/autograd.py. - ★★ Train the XOR MLP with learning rates 0.005, 0.05 and 0.5. Plot the loss curves. Which converges, which is slow, and which oscillates? Where:
experiments/ch03.py(create it), adaptingrun.py’scmd_autograd. - ★★ Implement momentum and Adam for lists of
Values and compare the steps-to-convergence on XOR with plain SGD. Where: add optimizer helpers besidesgd_stepinengine/autograd.py; compare them inexperiments/ch03.py(create it). - ★★★ Re-implement the XOR network in PyTorch with
nn.Linearandtorch.optim.AdamW. Then count the parameters, gradients and optimizer state in bytes for a 1B-parameter model trained this way. Where:experiments/ch03.py(create it).
Check your understanding
- Why does gradient descent move against the gradient?
- Why must gradients be accumulated with
+=in the backward pass? - What is lost if you remove every nonlinearity from an MLP?
- Why is
model.eval()not a substitute fortorch.no_grad()? - Roughly how many bytes does AdamW training need per parameter, compared with BF16 inference?
Going deeper
- BALLM Appendix A §§A.3-A.7 (pp. 258-277): computation graphs, autograd and training loops in PyTorch. Appendix D (pp. 313-321): warmup, cosine decay and gradient clipping.
- Andrej Karpathy, The spelled-out intro to neural networks and backpropagation: building micrograd (video and code), the inspiration for this chapter’s engine.
- Kingma and Ba, Adam: A Method for Stochastic Optimization (2015); Loshchilov and Hutter, Decoupled Weight Decay Regularization (2019), the paper behind AdamW.
- GPU Mode L6 (Jane Xu, Optimizing PyTorch Optimizers): how optimizer steps are fused into a few kernels.
4. Text as numbers: tokenizers and embeddings
In this chapter
- Why models read tokens, and the trade-offs between characters, words and subwords.
- Byte-pair encoding (BPE): the algorithm behind GPT-2, Llama and Qwen tokenizers, step by step.
- Special tokens and chat templates: the protocol a chat model was trained on.
- Turning a text into training examples, and token IDs into vectors with an embedding table.
You will build
engine/tokenizer.py and engine/data.py: a byte-level BPE tokenizer you train on a short novel, and the sliding-window dataset you'll train your GPT on in Chapter 7.
Time: 4-6 hours. GPU: not needed.
A model predicts IDs, not text
A language model is a function from a sequence of integers to a probability distribution over the next integer. A tokenizer is the contract that connects those integers to text. It splits text into units from a fixed vocabulary, maps each unit to an ID, and maps IDs back to text.
The contract is part of the model. A checkpoint’s embedding row 22670 was trained to mean whatever text the tokenizer maps to 22670. Feed it IDs from a different tokenizer and the model receives nonsense, even if every shape is right. When you load real models (Chapters 9 and 18), you’ll always use the checkpoint’s own tokenizer files.
Characters, words or something in between?
Characters are the simplest option: the vocabulary is every character in your training text, and each character gets an ID. It never meets an unknown word. But sequences get long. “tokenization” is 12 tokens, and attention’s cost grows with sequence length (Chapter 5).
Words give short sequences, but the vocabulary explodes (every name, typo, number and inflection is a separate entry) and any word not in it is unrepresentable.
Subwords split the difference: frequent words are single tokens (" the"), rare words are built from pieces (" token" "ization"). A good subword vocabulary of 50k-250k entries covers any text in a few tokens per word. The standard way to build one is byte-pair encoding.
Bytes first
Text in a computer is bytes. UTF-8 encodes ASCII characters as one byte each and other characters as two to four: "é" is the bytes 195 169, and "✓" is 226 156 147. If a tokenizer starts from the 256 possible byte values, every possible string is representable, with no unknown tokens ever. That’s byte-level BPE, used by GPT-2, Llama 3 and Qwen.
Byte-pair encoding, by hand
BPE builds the vocabulary bottom-up from a training text:
- Start with the 256 single-byte tokens.
- Count every adjacent pair of tokens in the training text.
- Merge the most frequent pair into a new token everywhere it occurs, and record the merge.
- Repeat until the vocabulary is as large as you want.
Here it is on a tiny corpus: “low” five times, “lower” twice, “newest” six times and “widest” three times. Before counting, a pre-tokenizer splits text into chunks (roughly: words with their leading space, punctuation and numbers), and pairs are never counted across chunk boundaries. These are the first eight merges the reference implementation learns:
| new ID | merge | new token | why |
|---|---|---|---|
| 256 | e + s | es | “es” appears in newest ×6 and widest ×3 |
| 257 | es + t | est | now es t is the top pair |
| 258 | l + o | lo | low ×5, lower ×2 |
| 259 | lo + w | low | |
| 260 | ␣ + low | ␣low | the leading space is part of the token |
| 261 | ␣ + n | ␣n | |
| 262 | ␣n + e | ␣ne | |
| 263 | ␣ne + w | ␣new |
Now encoding a word it never saw works by applying the learned merges, earliest first:
"lowest" -> ['low', 'est'] IDs [259, 257]
" newer" -> [' new', 'e', 'r'] IDs [263, 101, 114]
" lowly" -> [' low', 'l', 'y'] IDs [260, 108, 121]
Unknown words fall back to smaller pieces, down to single bytes if necessary. The order matters: encoding must apply merges in the order they were learned (lowest new-ID first), because later merges were learned on text where earlier merges had already happened.
The two core helpers are short in every language:
def pair_counts(ids, counts=None):
counts = Counter() if counts is None else counts
for pair in zip(ids, ids[1:]):
counts[pair] += 1
return counts
def merge_pair(ids, pair, new_id):
out, i = [], 0
while i < len(ids):
if i + 1 < len(ids) and (ids[i], ids[i + 1]) == pair:
out.append(new_id)
i += 2
else:
out.append(ids[i])
i += 1
return out
inline std::vector<int> merge_pair(const std::vector<int>& ids, std::pair<int, int> pair, int new_id) {
std::vector<int> out;
for (size_t i = 0; i < ids.size();) {
if (i + 1 < ids.size() && ids[i] == pair.first && ids[i + 1] == pair.second) { out.push_back(new_id); i += 2; }
else out.push_back(ids[i++]);
}
return out;
}
inline std::map<std::pair<int, int>, int> pair_counts(const std::vector<int>& ids) {
std::map<std::pair<int, int>, int> counts;
for (size_t i = 0; i + 1 < ids.size(); ++i) ++counts[{ids[i], ids[i + 1]}];
return counts;
}
#![allow(unused)]
fn main() {
/// Count adjacent pairs, remembering the order in which each pair first appeared (for ties).
pub fn pair_counts(ids: &[u32], counts: &mut HashMap<(u32, u32), (usize, usize)>, weight: usize) {
for w in ids.windows(2) {
let next_rank = counts.len();
let entry = counts.entry((w[0], w[1])).or_insert((0, next_rank));
entry.0 += weight;
}
}
/// Replace non-overlapping occurrences of `pair`, scanning left to right.
pub fn merge_pair(ids: &[u32], pair: (u32, u32), new_id: u32) -> Vec<u32> {
let mut out = Vec::with_capacity(ids.len());
let mut i = 0;
while i < ids.len() {
if i + 1 < ids.len() && (ids[i], ids[i + 1]) == pair {
out.push(new_id);
i += 2;
} else {
out.push(ids[i]);
i += 1;
}
}
out
}
/// Learn `num_merges` merges from whitespace-separated words (a simplified pre-tokenizer).
pub fn train(text: &str, num_merges: usize) -> Vec<((u32, u32), u32)> {
let mut words: Vec<(Vec<u32>, usize)> = {
let mut freq: Vec<(String, usize)> = vec![];
for word in text.split_inclusive(' ') {
match freq.iter_mut().find(|(w, _)| w == word) {
Some(entry) => entry.1 += 1,
None => freq.push((word.to_string(), 1)),
}
}
freq.into_iter().map(|(w, n)| (w.bytes().map(u32::from).collect(), n)).collect()
};
let mut merges = vec![];
for step in 0..num_merges {
let mut counts = HashMap::new();
for (ids, n) in &words {
pair_counts(ids, &mut counts, *n);
}
// Most frequent pair; ties go to the pair seen first.
let Some((&best, _)) = counts.iter().max_by(|a, b| a.1 .0.cmp(&b.1 .0).then(b.1 .1.cmp(&a.1 .1))) else { break };
let new_id = 256 + step as u32;
for (ids, _) in words.iter_mut() {
*ids = merge_pair(ids, best, new_id);
}
merges.push((best, new_id));
}
merges
}
}
Training on a real text
The book’s training corpus is Edith Wharton’s 1908 short story The Verdict: 20,479 characters and 3,634 words, public domain, and the same text BALLM uses. Training a 400-token vocabulary on it takes under a second (python run.py bpe --vocab 400):
first merges: [' t', 'he', ' a', 'in', ' h', ' s', ' w', ' o', ' the', 'ou', 're', 'it']
last merges: ['est', 'elf', 'ce', 'qu', ' sh', 'ind', ' Str', ' Stroud']
The first merges are the most common English letter pairs and " the". The last ones already include a character’s name, “Stroud”. With a 1,024-entry vocabulary, the whole story becomes 6,926 tokens instead of 20,479 bytes, about 3 bytes per token. Production tokenizers trained on trillions of tokens of diverse text reach about 4 bytes per token on English with 100-250k entries.
Note
Real tokenizers differ from ours in the pre-tokenizer. GPT-2 and Qwen split text with a regular expression that uses Unicode classes (
\p{L}for letters,\p{N}for numbers), which Python’s built-inremodule doesn’t support. Ours is a simplified ASCII version. Qwen also splits numbers into single digits, so that arithmetic sees consistent pieces. The merge algorithm itself is the same.
Special tokens and chat templates
Some IDs don’t represent text at all. They mark structure:
- End of text (
<|endoftext|>) separates documents during training, so the model learns where text ends. - Role markers in chat models, such as Qwen’s
<|im_start|>and<|im_end|>(the “ChatML” format), mark who is speaking and where a turn ends.
A chat template turns a list of messages into the exact token sequence the model was fine-tuned on. Qwen3’s looks like this:
<|im_start|>user
Explain what a GPU is in one sentence.<|im_end|>
<|im_start|>assistant
The final <|im_start|>assistant\n is the generation prompt: it tells the model that it’s now the assistant’s turn. Omit it and the model may continue the user’s message instead. When the model emits <|im_end|>, its turn is over and your engine should stop. Many “the model rambles forever” bugs are a missing generation prompt or a missing stop token.
Warning
Special tokens must be inserted as IDs, not tokenized as text. If
<|im_end|>is typed into a prompt as ordinary characters and tokenized byte by byte, the model sees<,|,im, … instead of its trained marker. Tokenizer libraries handle this with an “allowed special tokens” option; your BPEencodehas anallow_specialflag.
From a token stream to training examples
A language model learns to predict the next token at every position. From one stream of IDs, a sliding window of length context creates input/target pairs whose targets are the inputs shifted by one:
ids: 10 11 12 13 14 15 16 17 18 19
input: 10 11 12 13 target: 11 12 13 14
input: 12 13 14 15 target: 13 14 15 16 (stride 2: windows overlap)
input: 14 15 16 17 target: 15 16 17 18
One window gives context training examples at once: position 0 learns “after 10 comes 11”, position 1 learns “after 10 11 comes 12”, and so on. The stride controls overlap: stride = context gives disjoint windows, and a smaller stride reuses text.
Warning
Split before you window. If you make overlapping windows first and then split them randomly into training and validation sets, nearly identical windows land in both. Validation loss then measures memorization, not generalization. Split the token stream (or better, whole documents) first, then window each part.
split_then_windowdoes this.
Embeddings: from IDs to vectors
The model’s first operation turns each ID into a vector by looking up a row of an embedding table of shape [vocab_size, D]:
emb = torch.nn.Embedding(3, 2)
emb.weight.data = torch.tensor([[1., 0.], [0., 1.], [1., 1.]])
print(emb(torch.tensor([[2, 0]]))) # rows 2 and 0 -> [[[1., 1.], [1., 0.]]]
A lookup is mathematically the same as multiplying a one-hot vector (all zeros except a 1 at the ID) by the table, which is why embeddings train like any other weight. The lookup just skips the multiplication by zeros. Before training, the rows are random; training moves rows of tokens that behave alike closer together.
Two observations matter for engines:
- The table is big. Qwen3-0.6B’s is
[151936, 1024]: 156M of the model’s 600M parameters. Many small models tie the output projection to this table (they reuse it to turn the final hidden vector back into vocabulary scores), saving another 156M parameters. - A lookup reads one row per token, not the whole table. That’s why Flash-Next can afford a 51-billion-parameter hashed n-gram embedding (Chapter 30): each token touches only a few rows of it.
Position
An embedding lookup gives the same vector for " GPU" wherever it appears, so the model can’t tell “dog bites man” from “man bites dog”. GPT-2 adds a second learned table indexed by position: x[t] = token_embedding[id[t]] + position_embedding[t]. That table has a fixed number of rows, which caps the context length. Modern models use rotary position embeddings instead (Chapter 17), which encode position inside attention and generalize better to long sequences.
Build it
Engine milestone 4: a trainable tokenizer and a dataset. Implement:
- in
engine/tokenizer.py:pair_counts,merge_pair,BPETokenizer.trainandBPETokenizer._encode_chunk(the character tokenizer, pre-tokenizer, special-token handling and save/load are provided); - in
engine/data.py:windows(ids, context, stride).
pytest tests/test_ch04_text.py
python run.py bpe --vocab 512 --impl engine
The tests train on The Verdict, check that encode/decode round-trips any text (including accented characters and emoji), and check that your tokenizer compresses the text to under half its byte length.
Tip
Training is fast if you count each distinct pre-tokenized chunk once, weighted by how often it occurs. “the” appears hundreds of times but needs to be merged only once. In
_encode_chunk, the pair to merge next is the one with the lowest merge ID among pairs that have a merge.
Stretch exercises
- ★ Encode
"Hello, world!","hello world"and"HELLO WORLD"with your 512-token tokenizer. Count the tokens. Why do the counts differ so much? Where:experiments/ch04.py(create it), importingengine.tokenizer.BPETokenizer. - ★★ Plot bytes-per-token on The Verdict for vocabulary sizes 300, 500, 1,000, 2,000 and 4,000. Where do the returns diminish? Where:
experiments/ch04.py(create it), training ondata/the-verdict.txt. - ★★ Install
tiktokenand compare GPT-2’s tokenization of a paragraph of The Verdict with yours. Which words does GPT-2 keep whole that yours splits? Where:experiments/ch04.py(create it). - ★★★ Make
_encode_chunkfast: instead of rescanning all pairs after each merge, keep a priority queue of mergeable pairs keyed by merge rank. Measure the speed-up on the whole story. Where:BPETokenizer._encode_chunkinengine/tokenizer.py.
Check your understanding
- Why can a byte-level BPE tokenizer encode any string, while a word tokenizer cannot?
- Why must merges be applied in the order they were learned?
- What goes wrong if a chat prompt omits the final
<|im_start|>assistant\n? - Why does splitting overlapping windows into train and validation sets inflate validation scores?
- Why does an embedding table, by itself, carry no information about word order?
Going deeper
- BALLM Chapter 2 (pp. 17-49): tokenizing text, special tokens, byte-pair encoding with
tiktoken, sliding-window data loading, token and position embeddings. Our corpus and window examples follow it. - Sennrich, Haddow and Birch, Neural Machine Translation of Rare Words with Subword Units (2016): the paper that brought BPE to NLP.
- Andrej Karpathy, Let’s build the GPT Tokenizer (video) and the
minbperepository: byte-level BPE in depth, including the GPT-2 and GPT-4 regex pre-tokenizers. - The Hugging Face
tokenizersdocumentation, for howtokenizer.jsonstores pre-tokenizers, merges and special tokens.
5. Attention from first principles
In this chapter
- The problem attention solves: letting every token use information from every earlier token.
- Attention built up in four steps: plain dot-product weights, learned queries/keys/values, scaling, causal masking.
- Multi-head attention, and the reshapes that are easy to get subtly wrong.
- A position-aware formulation that later serves caching, batching and sparse attention unchanged.
You will build
engine/attention.py: split_heads, merge_heads and causal_attention, the function every model in this book calls.
Time: 5-7 hours. GPU: not needed.
The problem: context
After Chapter 4, each token is a vector that describes the token in isolation. But meaning depends on context. In “The animal didn’t cross the street because it was too tired”, the vector for “it” should end up carrying information about “animal”. Predicting the next token needs the whole preceding text, not just the last word.
Before 2017, the standard answer was the recurrent neural network (RNN). It reads tokens one at a time and squeezes everything seen so far into a single fixed-size vector, its hidden state. That has two problems. Information from far back must survive many squeezes and fades. And processing is inherently sequential, so training can’t use a GPU’s parallelism across positions.
Attention takes the opposite approach: when processing token $t$, look directly at every earlier token, decide how relevant each one is, and take a weighted mixture of their information. Nothing is squeezed, and all positions can be computed at once. That one idea, from Attention Is All You Need (Vaswani et al., 2017), is the core of every model in this book. (Recurrence comes back, in a much improved form, in Chapter 28.)
Step 1: attention with no parameters at all
We’ll follow BALLM’s example: six tokens, “Your journey starts with one step”, each already embedded as a 3-dimensional vector:
x = torch.tensor([[0.43, 0.15, 0.89], # Your
[0.55, 0.87, 0.66], # journey
[0.57, 0.85, 0.64], # starts
[0.22, 0.58, 0.33], # with
[0.77, 0.25, 0.10], # one
[0.05, 0.80, 0.55]]) # step
Let’s compute an enriched, context-aware vector for “journey”. The recipe has three steps.
Score every token by its similarity to “journey”. The dot product is a natural similarity measure: it’s large when two vectors point the same way.
scores = x @ x[1] = [0.9544, 1.4950, 1.4754, 0.8434, 0.7070, 1.0865]
Normalize the scores into weights that are positive and sum to 1, using softmax:
$$ w_j = \frac{e^{s_j}}{\sum_k e^{s_k}} \quad\Rightarrow\quad w = [0.1385,\ 0.2379,\ 0.2333,\ 0.1240,\ 0.1082,\ 0.1581]. $$
“journey” weighs itself highest, then “starts”, whose vector is nearly identical.
Mix: the new vector is the weighted sum of all token vectors:
$$ z_{\text{journey}} = \sum_j w_j, x_j = [0.4419,\ 0.6515,\ 0.5683]. $$
That’s attention. Doing it for every token at once is two matrix multiplications and a softmax:
$$ \text{weights} = \operatorname{softmax}_{\text{row}}(X X^{\top}), \qquad Z = \text{weights}; X . $$
Row $i$ of the [6, 6] weights matrix says how much token $i$ draws from each token $j$.
Step 2: learned queries, keys and values
Plain dot products only find vectors that are already similar. A model needs to learn what to look for, and that depends on the role a token plays. So attention gives each token three different learned views of itself, each produced by its own weight matrix:
- Query $q_i = x_i W_q$: what token $i$ is looking for.
- Key $k_j = x_j W_k$: what token $j$ offers, which is matched against queries.
- Value $v_j = x_j W_v$: what token $j$ hands over if it’s chosen.
$$ \text{scores} = Q K^{\top}, \qquad Z = \operatorname{softmax}(\text{scores}), V . $$
Separating keys from values is the important design choice. A token can be found for one reason and contribute something else. For “it”, the query might look for “animate noun earlier in the sentence”; the key of “animal” advertises exactly that; and its value passes along information useful for predicting what comes after “it”. These are intuitions, not literal labels. Training simply finds whatever projections lower the loss. The mechanism makes such behavior possible.
Step 3: scale the scores
Dot products of $d$-dimensional vectors with independent, unit-variance components have a standard deviation of about $\sqrt d$. Measured on random vectors:
| head dimension $d$ | std of $q\cdot k$ | std of $q\cdot k/\sqrt d$ |
|---|---|---|
| 2 | 1.44 | 1.02 |
| 64 | 8.08 | 1.01 |
| 1024 | 31.6 | 0.99 |
Large scores make softmax saturate: softmax([1, 2, 3]) is [0.09, 0.24, 0.67], but softmax([8, 16, 24]) is [0.0000, 0.0003, 0.9997]. A saturated softmax puts all its weight on one token, and its gradient is nearly zero, so training stalls. Dividing scores by $\sqrt{d}$ keeps them in a healthy range at any width. This is scaled dot-product attention:
$$ \operatorname{Attention}(Q, K, V) = \operatorname{softmax}!\left(\frac{Q K^{\top}}{\sqrt{d}}\right) V . $$
Step 4: no peeking at the future
A language model learns to predict token $t+1$ from tokens $\le t$. If position $t$ could attend to position $t+1$ during training, it would simply copy the answer. The loss would drop to near zero and the model would learn nothing useful.
The fix is a causal mask: before the softmax, set the scores for future positions to $-\infty$. Since $e^{-\infty} = 0$, they get exactly zero weight, and the remaining weights still sum to 1. For our six tokens (using the plain dot products of Step 1):
Your journey starts with one step
Your 1.000 0.000 0.000 0.000 0.000 0.000
journey 0.368 0.632 0.000 0.000 0.000 0.000
starts 0.228 0.389 0.382 0.000 0.000 0.000
with 0.205 0.296 0.292 0.208 0.000 0.000
one 0.175 0.225 0.227 0.157 0.216 0.000
step 0.139 0.218 0.213 0.142 0.099 0.190
The first token can only attend to itself. The last row is unchanged from the unmasked version, because nothing comes after “step”.
Warning
Mask before the softmax. Masking after it (zeroing weights of future tokens) leaves their scores in the denominator, so the remaining weights no longer sum to 1, and information about the future leaks through the normalization.
The best test of a mask doesn’t compare numbers with a library. It tests the rule directly: change only the last input token, and check that every earlier output stays exactly the same. Your milestone tests do this.
Multiple heads
One attention pattern per layer is limiting. A token may need to track its subject, its previous token and the most recent comma all at once. Multi-head attention runs $H$ smaller attentions in parallel, each with its own projections into a $D_h$-dimensional space (typically $D = H \cdot D_h$), then concatenates their outputs and mixes them with an output projection $W_o$.
In code, nobody runs $H$ separate projections. One big projection produces all heads’ queries at once, and a reshape splits them:
x: [B, T, D]
q = x @ Wq.T [B, T, H*Dh]
q.reshape(B, T, H, Dh) [B, T, H, Dh] split the last axis into heads
.transpose(1, 2) [B, H, T, Dh] heads become a batch axis
scores = q @ k.transpose(-2, -1) [B, H, T, T] one T×T matrix per head
y = weights @ v [B, H, T, Dh]
y.transpose(1, 2) [B, T, H, Dh]
.reshape(B, T, H*Dh) [B, T, D] heads concatenated per token
out = y @ Wo.T [B, T, D]
Warning
The transposes are not optional.
reshape(B, H, T, Dh)directly on a[B, T, H*Dh]tensor produces the right shape with the wrong data: it interleaves tokens and heads. Nothing crashes, and the model just computes garbage. Your milestone test checks that head 1 of token 2 really holds features 4-8 of token 2.
Fewer key/value heads: grouped-query attention
Modern models often use fewer key/value heads than query heads. In grouped-query attention (GQA), each KV head is shared by a group of query heads: Qwen3-0.6B has 16 query heads and 8 KV heads, and Flash-Next’s attention layers have 24 query heads and just 2 KV heads. Each query head still computes its own attention pattern, but the K and V tensors are smaller. That matters enormously for the KV cache’s memory (Chapter 16). In code, query head $h$ reads KV head $h \div (H_q/H_{kv})$. The simplest implementation repeats each KV head to match the query heads before the matmul.
Positions, not just “the last T keys”
There’s one more design decision in your causal_attention, and it pays off for the rest of the book. Instead of building a fixed triangular mask, it takes the absolute position of every query and key and applies the rule
$$ \text{key } j \text{ is visible to query } i \iff \text{pos}(k_j) \le \text{pos}(q_i). $$
During ordinary training, queries and keys are both positions $0..T-1$, and this reproduces the triangle. But the same rule also handles, unchanged:
- cached decoding (Chapter 16): 1 new query at position 57 against 58 cached keys;
- chunked prefill: queries 40-47 against keys 0-47, a rectangular mask that a square triangle gets wrong;
- batches of different lengths (Chapters 19 and 24), where each row has its own positions;
- sparse attention (Chapter 29), via the extra
allowedmask.
def causal_attention(q, k, v, query_positions=None, key_positions=None, allowed=None, scale=None):
"""Scaled dot-product attention with a causal rule. (Your engine: Chapter 5)
query_positions: [T] or [B, T]; defaults to the last T positions of the S keys.
key_positions: [S] or [B, S]; defaults to 0..S-1.
allowed: optional extra boolean mask broadcastable to [B, Hq, T, S] (sliding windows,
sparse selections). True means "may attend".
Grouped-query attention: Hq must be a multiple of Hkv; each KV head serves Hq/Hkv query heads.
"""
batch, q_heads, t, d = q.shape
kv_heads, s = k.shape[1], k.shape[2]
if q_heads % kv_heads:
raise ValueError(f"Query heads ({q_heads}) must be a multiple of KV heads ({kv_heads})")
if kv_heads != q_heads:
k = k.repeat_interleave(q_heads // kv_heads, dim=1)
v = v.repeat_interleave(q_heads // kv_heads, dim=1)
scale = 1.0 / math.sqrt(d) if scale is None else scale
scores = (q.float() @ k.float().transpose(-2, -1)) * scale # [B, Hq, T, S]
qp = _positions(query_positions, batch, t, s - t, q.device) # [B, T]
kp = _positions(key_positions, batch, s, 0, q.device) # [B, S]
mask = kp[:, None, None, :] <= qp[:, None, :, None] # [B, 1, T, S]
if allowed is not None:
mask = mask & allowed
scores = scores.masked_fill(~mask, float("-inf"))
weights = torch.softmax(scores, dim=-1)
return (weights @ v.float()).to(q.dtype)
// q [T, Hq*D] for positions S-T..S-1; k, v [S, Hkv*D]. Query head h reads KV head h/(Hq/Hkv).
inline std::vector<float> causal_attention(const std::vector<float>& q, const std::vector<float>& k, const std::vector<float>& v,
size_t T, size_t S, size_t hq, size_t hkv, size_t d) {
std::vector<float> out(T * hq * d, 0.f), scores(S);
float scale = 1.f / std::sqrt(float(d));
for (size_t i = 0; i < T; ++i)
for (size_t h = 0; h < hq; ++h) {
size_t kvh = h / (hq / hkv), visible = S - T + i + 1;
const float* qv = &q[(i * hq + h) * d];
for (size_t j = 0; j < visible; ++j) {
const float* kv = &k[(j * hkv + kvh) * d];
scores[j] = std::inner_product(qv, qv + d, kv, 0.f) * scale;
}
softmax(scores.data(), visible);
for (size_t j = 0; j < visible; ++j)
for (size_t e = 0; e < d; ++e) out[(i * hq + h) * d + e] += scores[j] * v[(j * hkv + kvh) * d + e];
}
return out;
}
#![allow(unused)]
fn main() {
/// q: [T, Hq*D] for the new tokens at positions start..start+T; k, v: [S, Hkv*D] for every
/// cached position 0..S. Each query reads only keys at positions <= its own (causality).
/// Query head h reads KV head h / (Hq/Hkv): grouped-query attention.
pub fn causal_attention(q: &[f32], k: &[f32], v: &[f32], t: usize, s: usize,
q_heads: usize, kv_heads: usize, d: usize) -> Vec<f32> {
let start = s - t;
let group = q_heads / kv_heads;
let scale = 1.0 / (d as f32).sqrt();
let mut out = vec![0.0; t * q_heads * d];
let mut scores = vec![0.0; s];
for i in 0..t {
let visible = start + i + 1; // keys 0..=start+i
for h in 0..q_heads {
let kvh = h / group;
let qv = &q[(i * q_heads + h) * d..][..d];
for j in 0..visible {
let kv = &k[(j * kv_heads + kvh) * d..][..d];
scores[j] = qv.iter().zip(kv).map(|(a, b)| a * b).sum::<f32>() * scale;
}
softmax(&mut scores[..visible]);
let o = &mut out[(i * q_heads + h) * d..][..d];
for j in 0..visible {
let vv = &v[(j * kv_heads + kvh) * d..][..d];
for (oo, vvv) in o.iter_mut().zip(vv) {
*oo += scores[j] * vvv;
}
}
}
}
out
}
}
The C++ and Rust versions process one sequence with explicit loops and no batch axis. Read them to see the computation without tensor broadcasting. They implement the same rule: query $i$ of $T$ new tokens sees keys $0 \ldots S-T+i$.
What attention costs
For a sequence of $T$ tokens with model width $D$, the score matrix has $T^2$ entries per head, and computing scores and outputs costs about $4T^2D$ FLOPs per layer. At $T = 8{,}192$ and FP32, one head’s score matrix alone is 256 MiB. This quadratic growth is attention’s central problem at long context. FlashAttention (Chapter 15) removes the memory cost without changing the result. Linear attention (Chapter 28) and sparse attention (Chapter 29) change the computation itself.
Build it
Engine milestone 5: causal attention. In engine/attention.py, implement:
split_heads(x, heads):[B, T, H*Dh]→[B, H, T, Dh].merge_heads(x): the inverse.causal_attention(q, k, v, query_positions=None, key_positions=None, allowed=None, scale=None): grouped-query attention with the position rule above. The_positionshelper that fills in defaults is provided.
pytest tests/test_ch05_attention.py
python run.py attention --impl engine
The tests compare with PyTorch’s scaled_dot_product_attention, check this chapter’s worked example, perturb a future token to prove causality, and exercise GQA and rectangular queries.
Tip
Compute scores in FP32 (
q.float() @ k.float().transpose(-2, -1)) and cast the result back toq.dtype. Build the mask by broadcasting positions:key_pos[:, None, None, :] <= query_pos[:, None, :, None]has shape[B, 1, T, S]and broadcasts over heads. Usemasked_fill(~mask, float("-inf"))beforetorch.softmax.
Stretch exercises
- ★ Recompute this chapter’s six-token causal weight table with your function (use
xas q, k and v with a single head, andscale=1.0). Where:experiments/ch05.py(create it), importingengine.attention.causal_attention. - ★★ Write a test that catches the “reshape without transpose” bug: build
qso that each head’s features are distinguishable and assert thatmerge_heads(split_heads(x))equalsx, but that a wrong reshape doesn’t. Where: createtests/test_ch05_stretch.py, importingengine.attention. - ★★ Implement a sliding-window variant using the
allowedargument: each query sees only the last $w$ keys. (You’ll meet this again in Chapter 29.) Where: add a window-mask helper toengine/attention.pyand pass its result tocausal_attention(allowed=...). - ★★★ Measure the runtime and peak memory of
causal_attentionfor $T$ = 512, 1,024, 2,048 and 4,096. Confirm the quadratic growth and estimate when it stops fitting in your device’s memory. Where:experiments/ch05.py(create it), importingengine.attention.causal_attention.
Check your understanding
- Why does attention normalize over keys (each row) rather than over queries (each column)?
- Why mask before the softmax?
- Which inputs can affect the output at position 0 of a causal attention layer?
- What does dividing by $\sqrt{d}$ fix, and what goes wrong without it?
- With 16 query heads and 8 KV heads, which KV head does query head 11 read?
Going deeper
- BALLM Chapter 3 (pp. 50-91): attention from simplified weights to causal multi-head attention. This chapter’s worked example comes from §3.3.
- PMPP §20.2-20.3 (pp. 482-488): multi-head attention as matrix operations, and a first CUDA implementation.
- Vaswani et al., Attention Is All You Need (2017); Ainslie et al., GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints (2023).
- Jay Alammar, The Illustrated Transformer: the classic visual walkthrough.
6. The transformer block and a complete GPT
In this chapter
- The residual stream: why each layer adds to its input instead of replacing it.
- LayerNorm and the GELU MLP, the other two ingredients of a transformer block.
- Assembling embeddings, blocks and an output head into a GPT, and counting its parameters exactly.
You will build
engine/gpt.py: a GPT-2-compatible model. In Chapter 7 you train it from scratch; in Chapter 9 it runs OpenAI's real GPT-2 weights unchanged.
Time: 4-6 hours. GPU: not needed.
The shape of a GPT
A GPT (“generative pre-trained transformer”) is a stack of identical blocks between an input embedding and an output head:
token IDs [B, T]
│ token embedding + position embedding
▼
x [B, T, D] ──► block 1 ──► block 2 ──► ... ──► block L ──► final LayerNorm ──► head ──► logits [B, T, V]
Every block has the same shape in and out, [B, T, D], so blocks stack like Lego. GPT-2 small has $L = 12$ blocks of width $D = 768$; Qwen3-0.6B has 28 blocks of width 1,024. Once you’ve built one block, you’ve built the model.
The residual stream
Each block contains two branches, attention and an MLP. Each branch reads the current vector, computes something, and adds its result back:
$$ \begin{aligned} a &= x + \operatorname{Attention}(\operatorname{LN}_1(x)) \ y &= a + \operatorname{MLP}(\operatorname{LN}_2(a)) \end{aligned} $$
Picture the vector $x$ as a residual stream that flows from the embedding to the output. Each branch reads from it and writes a small update into it. This has three big consequences:
- Training works at depth. The gradient of $x + f(x)$ with respect to $x$ is $1 + f’(x)$. Even if a branch’s gradient is tiny, the “1” carries the signal straight back to early layers. Before residual connections (He et al., 2015), very deep networks barely trained.
- A block can do nothing. If both branches output zero, the block is the identity. A new block starts as a small perturbation of a working network, not a random scrambling of it. Your milestone test checks exactly this: zero the output projections and the block must return its input.
- Interpretability. Every branch reads and writes the same space, so you can inspect, edit or steer what’s in the stream. Chapter 23 does exactly that.
This “pre-norm” arrangement (normalize the input of each branch) is used by GPT-2 and every modern model. The original 2017 transformer normalized after the addition (post-norm), which is harder to train at depth.
LayerNorm: keep the numbers in range
As updates accumulate in the stream, its vectors can drift to very different scales from token to token. Layer normalization rescales each token vector, independently of every other token and every other example, to mean 0 and variance 1, then applies a learned per-feature scale $\gamma$ and shift $\beta$:
$$ \operatorname{LN}(x)_i = \gamma_i ,\frac{x_i - \mu}{\sqrt{\sigma^2 + \epsilon}} + \beta_i, \qquad \mu = \frac1D\sum_i x_i,\quad \sigma^2 = \frac1D\sum_i (x_i-\mu)^2 . $$
For $x = [1, 2, 3]$: $\mu = 2$, $\sigma^2 = 2/3$, and the normalized vector is $[-1.2247, 0, 1.2247]$ (before $\gamma$ and $\beta$). Note that the variance divides by $D$, not $D-1$; the “sample variance” correction from statistics would be a different operation and would break parity with real checkpoints. The small $\epsilon$ (1e-5 in GPT-2) prevents division by zero.
def layernorm(x, gamma, beta, eps=1e-5): # x [..., D]
mean = x.mean(-1, keepdim=True)
var = x.var(-1, keepdim=True, unbiased=False) # divide by D, not D-1
return gamma * (x - mean) / torch.sqrt(var + eps) + beta
inline std::vector<float> layernorm(const std::vector<float>& x, const std::vector<float>& g, const std::vector<float>& b, float eps) {
float mean = std::accumulate(x.begin(), x.end(), 0.f) / x.size(), var = 0;
for (float v : x) var += (v - mean) * (v - mean);
float inv = 1.f / std::sqrt(var / x.size() + eps);
std::vector<float> y(x.size());
for (size_t i = 0; i < x.size(); ++i) y[i] = (x[i] - mean) * inv * g[i] + b[i];
return y;
}
inline std::vector<float> rmsnorm(const std::vector<float>& x, const std::vector<float>& w, float eps) {
float ms = 0;
for (float v : x) ms += v * v;
float inv = 1.f / std::sqrt(ms / x.size() + eps);
std::vector<float> y(x.size());
for (size_t i = 0; i < x.size(); ++i) y[i] = x[i] * inv * w[i];
return y;
}
#![allow(unused)]
fn main() {
/// LayerNorm over one token vector: subtract the mean, divide by the standard deviation.
pub fn layernorm(x: &[f32], gamma: &[f32], beta: &[f32], eps: f32) -> Vec<f32> {
let n = x.len() as f32;
let mean = x.iter().sum::<f32>() / n;
let var = x.iter().map(|v| (v - mean) * (v - mean)).sum::<f32>() / n;
let inv = 1.0 / (var + eps).sqrt();
x.iter().zip(gamma).zip(beta).map(|((v, g), b)| (v - mean) * inv * g + b).collect()
}
/// RMSNorm: no mean subtraction, divide by the root mean square.
pub fn rmsnorm(x: &[f32], weight: &[f32], eps: f32) -> Vec<f32> {
let ms = x.iter().map(|v| v * v).sum::<f32>() / x.len() as f32;
let inv = 1.0 / (ms + eps).sqrt();
x.iter().zip(weight).map(|(v, w)| v * inv * w).collect()
}
}
(The C++ and Rust tabs also include RMSNorm, the simpler variant Qwen uses. It’s Chapter 17’s topic.)
The MLP: where most parameters live
Attention moves information between positions. The MLP (also called the feed-forward network) then processes each position on its own: expand to $4D$, apply a nonlinearity, contract back to $D$:
$$ \operatorname{MLP}(x) = W_{\text{down}},\operatorname{GELU}(W_{\text{up}},x + b_{\text{up}}) + b_{\text{down}} . $$
GPT-2 uses GELU, a smooth relative of ReLU. It’s close to 0 for very negative inputs, close to $x$ for positive ones, and smoothly curved in between. GPT-2 uses the tanh approximation:
$$ \operatorname{GELU}(x) \approx \tfrac12 x\left(1 + \tanh!\left(\sqrt{2/\pi},(x + 0.044715x^3)\right)\right). $$
The exact error-function version differs slightly. Using the wrong one is a small but real source of mismatch when loading checkpoints, so match the checkpoint’s convention (approximate="tanh" for GPT-2).
The MLP holds two thirds of each block’s weights ($8D^2$ versus attention’s $4D^2$). A useful mental model, explored in interpretability research, treats it as a key-value memory: the first matrix detects patterns, and the second writes associated information into the stream.
The block in code
Here is the reference block. Attention packs Q, K and V into one linear layer (qkv, width $3D$), which is a storage choice; GPT-2’s checkpoint stores it that way too. Then come the residual additions you just saw:
class GPTAttention(nn.Module):
"""Multi-head causal self-attention with one packed QKV projection."""
def __init__(self, cfg):
"""Create qkv: width -> 3*width and proj: width -> width, both with bias. (Your engine: Chapter 6)"""
super().__init__()
if cfg.width % cfg.heads:
raise ValueError("width must be divisible by heads")
self.heads = cfg.heads
self.qkv = nn.Linear(cfg.width, 3 * cfg.width)
self.proj = nn.Linear(cfg.width, cfg.width)
def forward(self, x, positions, cache=None, layer=0, rows=None):
"""Project, split heads, (cache), attend, merge heads, project. (Your engine: Chapter 6)"""
q, k, v = self.qkv(x).chunk(3, dim=-1)
q, k, v = (split_heads(t, self.heads) for t in (q, k, v))
key_positions = None
if cache is not None:
k, v, key_positions = cache.update(layer, k, v, positions, rows)
y = causal_attention(q, k, v, positions, key_positions)
return self.proj(merge_heads(y))
class GPTBlock(nn.Module):
def __init__(self, cfg):
"""Two LayerNorms, attention, and a 4x-wide GELU MLP. (Your engine: Chapter 6)"""
super().__init__()
self.ln1 = nn.LayerNorm(cfg.width, eps=cfg.eps)
self.attn = GPTAttention(cfg)
self.ln2 = nn.LayerNorm(cfg.width, eps=cfg.eps)
self.up = nn.Linear(cfg.width, 4 * cfg.width)
self.down = nn.Linear(4 * cfg.width, cfg.width)
self.drop = nn.Dropout(cfg.dropout)
def forward(self, x, positions, cache=None, layer=0, rows=None):
"""Pre-norm residual block: x + attn(ln1(x)), then + mlp(ln2(.)). (Your engine: Chapter 6)"""
x = x + self.drop(self.attn(self.ln1(x), positions, cache, layer, rows))
hidden = nn.functional.gelu(self.up(self.ln2(x)), approximate="tanh")
return x + self.drop(self.down(hidden))
The positions, cache and rows arguments let the same block serve generation with a KV cache later (Chapter 16) and batched serving (Chapter 24). For now they’re just passed through. cache is None and positions are $0..T-1$.
Assembling the model
class GPT(nn.Module):
def __init__(self, cfg):
"""Token and position embeddings, a stack of blocks, a final LayerNorm and a vocabulary head. (Your engine: Chapter 6)"""
super().__init__()
self.cfg = cfg
self.token = nn.Embedding(cfg.vocab, cfg.width)
self.position = nn.Embedding(cfg.context, cfg.width)
self.drop = nn.Dropout(cfg.dropout)
self.blocks = nn.ModuleList(GPTBlock(cfg) for _ in range(cfg.layers))
self.norm = nn.LayerNorm(cfg.width, eps=cfg.eps)
self.head = nn.Linear(cfg.width, cfg.vocab, bias=False)
if cfg.tie_weights:
self.head.weight = self.token.weight
self.apply(self._initialize)
@staticmethod
def _initialize(module):
if isinstance(module, (nn.Linear, nn.Embedding)):
nn.init.normal_(module.weight, std=0.02)
if isinstance(module, nn.Linear) and module.bias is not None:
nn.init.zeros_(module.bias)
def forward(self, ids, cache=None, positions=None, rows=None):
"""ids [B, T] -> logits [B, T, vocab]. (Your engine: Chapter 6)
Positions default to 0..T-1, or continue after the cache's current length.
"""
if positions is None:
start = cache.length if cache is not None else 0
if start + ids.shape[1] > self.cfg.context:
raise ValueError(f"Sequence exceeds the model's context of {self.cfg.context}")
positions = torch.arange(start, start + ids.shape[1], device=ids.device)
# Explicit position tensors are trusted: checking them would force a GPU->CPU sync.
x = self.drop(self.token(ids) + self.position(positions))
for i, block in enumerate(self.blocks):
x = block(x, positions, cache, i, rows)
return self.head(self.norm(x))
@property
def context_limit(self):
return self.cfg.context
def cache_spec(self):
"""(layers, kv_heads, head_dim): what every cache needs to know about this model."""
return self.cfg.layers, self.cfg.heads, self.cfg.width // self.cfg.heads
def new_cache(self, batch, capacity):
if capacity > self.context_limit:
raise ValueError("Requested cache exceeds the context limit")
p = self.token.weight
layers, kv_heads, head_dim = self.cache_spec()
return KVCache(layers, batch, kv_heads, capacity, head_dim, p.device, p.dtype)
Four details deserve attention:
- Two embeddings are added: token identity plus learned position. The position table has
contextrows, which hard-limits the sequence length. - The head produces logits, one score per vocabulary entry at every position: shape
[B, T, V]. Position $t$’s logits predict token $t+1$. - Weight tying:
self.head.weight = self.token.weightmakes the output head the same tensor as the embedding table, read in reverse. It saves $V \cdot D$ parameters (38.6M for GPT-2) and was used by GPT-2. The parameter count must count it once. - Initialization: weights start as small random normals (std 0.02, as in GPT-2) and biases at zero. Getting this wrong produces a model that trains slowly or not at all, but it doesn’t matter for loading pretrained weights, which overwrite everything.
Counting parameters exactly
For a block with biases and two LayerNorms:
| part | parameters |
|---|---|
attention: qkv ($D \times 3D$ + bias) and proj ($D\times D$ + bias) | $4D^2 + 4D$ |
MLP: up ($D \times 4D$ + bias) and down ($4D \times D$ + bias) | $8D^2 + 5D$ |
| two LayerNorms ($\gamma, \beta$ each) | $4D$ |
| block total | $12D^2 + 13D$ |
The whole tied model, with vocabulary $V$ and context $C$:
$$ N = VD + CD + L(12D^2 + 13D) + 2D . $$
For GPT-2 small ($V = 50{,}257$, $C = 1{,}024$, $D = 768$, $L = 12$), that’s 124,439,808, the “124M” in its name. python run.py gpt prints this alongside the count from your model, and your milestone tests that the two agree.
Notice how much of a small model is embedding: 31% of GPT-2 small. As models grow, the $12LD^2$ term dominates, which is why “parameters” and “FLOPs per token” are both close to $12LD^2$ for large models.
A random model already “works”
Before training, your GPT runs end to end and produces logits of the right shape, and the output is noise. That’s useful, not useless: plumbing correctness and learned behavior are separate claims, and you can test the first without the second. The milestone tests check shapes, causality (changing the last token never changes earlier logits), the parameter formula, weight tying, and the identity behavior of zeroed branches. None of that needs training.
Build it
Engine milestone 6: a complete GPT. In engine/gpt.py, implement __init__ and forward for GPTAttention, GPTBlock and GPT. Use your split_heads, merge_heads and causal_attention from Chapter 5. The config dataclass, initialization, parameter-counting helpers and cache factory are provided.
pytest tests/test_ch06_gpt.py
python run.py gpt --impl engine
Tip
Keep the exact attribute names of the reference (
token,position,blocks,norm,head,qkv,proj,ln1,ln2,up,down). Chapter 9’s GPT-2 loader maps checkpoint tensors onto them. If you name things differently, update the loader’s mapping.
Stretch exercises
- ★ Compute by hand the parameter count of GPT-2 medium ($D = 1024$, $L = 24$, 16 heads) and check it with
gpt_parameter_formula. Is it really “355M”? Where: paper, thenexperiments/ch06.py(create it) usingengine.gpt.gpt_parameter_formula. - ★★ Add a
return_hiddenoption that returns the residual stream after every block. Plot each block’s vector norm for a random input. How does it grow with depth? Where:GPT.forwardinengine/gpt.py. - ★★ Convert the block to post-norm (normalize after each addition) and verify that zeroed branches no longer give an identity block. Why? Where:
GPTBlock.forwardinengine/gpt.py; keep a separate post-norm variant for comparison. - ★★★ Implement the exact GELU and measure the maximum difference from the tanh version over [−6, 6]. Then estimate how much it would change GPT-2’s logits. Where:
experiments/ch06.py(create it) for the numerical comparison; change the GELU inGPTBlock.forwardinengine/gpt.pyfor the logit comparison.
Check your understanding
- Which axis does the MLP mix, and which does attention mix?
- Why does a pre-norm block with both branch outputs at zero leave its input unchanged?
- What changes, in storage and in training, when the head is tied to the token embedding?
- Why is LayerNorm’s output not always mean-zero and unit-variance?
- Where do the $12D^2$ parameters of a block come from?
Going deeper
- BALLM Chapter 4 (pp. 92-127): LayerNorm, GELU, shortcut connections and the GPT model, built up in the same order. Its parameter-counting exercise matches this chapter’s formula.
- PMPP §20.1 (pp. 478-482): the decoder block as a sequence of GEMMs and elementwise operations.
- Radford et al., Language Models are Unsupervised Multitask Learners (GPT-2, 2019); He et al., Deep Residual Learning (2015); Ba et al., Layer Normalization (2016).
- Anthropic, A Mathematical Framework for Transformer Circuits (2021): the residual stream view of transformers.
7. Train and evaluate your GPT
In this chapter
- The language-modeling objective: cross-entropy on the next token, and what its numbers mean.
- A complete training loop with AdamW, warmup, cosine decay, gradient clipping and periodic evaluation.
- Sanity checks that catch most training bugs in minutes.
- Overfitting, seen live on a real text, and what it does and doesn't tell you.
You will build
engine/train.py: lm_loss, evaluate and train. You'll train your Chapter 6 GPT on The Verdict and sample from it.
Time: 4-6 hours. GPU: helpful but not needed (the default experiment takes under a minute on a laptop CPU).
The objective: predict the next token
Chapter 4 turned text into windows of inputs and targets shifted by one. For every position $t$ in every window, the model outputs logits $z$, a score per vocabulary entry, and we want the probability of the true next token $y$ to be high. Softmax turns logits into probabilities, and the cross-entropy loss is the negative log of the probability assigned to the correct answer:
$$ p(y) = \frac{e^{z_y}}{\sum_j e^{z_j}}, \qquad L = -\log p(y). $$
The numbers are worth internalizing:
| probability of the right token | loss |
|---|---|
| 1.00 | 0.000 |
| 0.50 | 0.693 |
| 0.10 | 2.303 |
| 0.01 | 4.605 |
| 1/V (uniform guessing, V = 512) | 6.238 |
The last row is the most useful number in a training log. An untrained model should start at about $\log V$. If your first loss is 30, the initialization or the loss computation is broken. If it’s 0.5, the model can see its targets.
The gradient has a beautiful form
The derivative of cross-entropy with respect to each logit is
$$ \frac{\partial L}{\partial z_j} = p_j - \mathbb 1[j = y]. $$
If the target has probability 0.2, its logit’s gradient is −0.8, so gradient descent pushes it up. A wrong token with probability 0.3 gets +0.3 and is pushed down. Every logit is nudged in proportion to how wrong its probability is. This gradient then flows back through the head, every block, and the embeddings, via the chain rule from Chapter 3.
Use logits, not probabilities
PyTorch’s F.cross_entropy takes logits and computes a numerically stable log-softmax internally. Passing it probabilities (applying softmax yourself first) silently computes a different, wrong loss. It also wants 2-D logits and 1-D targets, so flatten batch and time together:
def lm_loss(model, x, y):
"""Mean next-token cross-entropy over every position. (Your engine: Chapter 7)
logits [B, T, V] and targets [B, T] are flattened to [B*T, V] and [B*T]. cross_entropy
applies a stable log-softmax itself, so it takes raw logits, never probabilities.
"""
logits = model(x)
return F.cross_entropy(logits.reshape(-1, logits.shape[-1]), y.reshape(-1))
For a validation set, average by tokens, not by batches. If batches have different sizes, a plain average of batch means over-weights the small ones:
@torch.no_grad()
def evaluate(model, x, y, batch_size=32):
"""Token-weighted mean loss over a whole split, in eval mode, without gradients. (Your engine: Chapter 7)"""
was_training = model.training
model.eval()
total, count = 0.0, 0
for start in range(0, len(x), batch_size):
xb, yb = x[start:start + batch_size], y[start:start + batch_size]
total += lm_loss(model, xb, yb).item() * yb.numel()
count += yb.numel()
model.train(was_training)
return total / count
The training loop
Here’s the full loop you’ll write. It’s the five-line core from Chapter 3, plus the details that make training stable:
def train(model, train_xy, valid_xy, steps, batch_size=16, lr=3e-3, weight_decay=0.1,
warmup=20, clip=1.0, eval_every=50, seed=0, log=print):
"""AdamW with linear warmup then cosine decay, gradient clipping and periodic evaluation. (Your engine: Chapter 7)
Returns a list of {"step", "train_loss", "valid_loss", "lr"} records.
"""
device = next(model.parameters()).device
tx, ty = (t.to(device) for t in train_xy)
vx, vy = (t.to(device) for t in valid_xy)
generator = torch.Generator().manual_seed(seed)
optimizer = torch.optim.AdamW(model.parameters(), lr=lr, weight_decay=weight_decay, betas=(0.9, 0.95))
def schedule(step):
if step < warmup:
return (step + 1) / warmup
progress = (step - warmup) / max(1, steps - warmup)
return 0.1 + 0.9 * 0.5 * (1 + math.cos(math.pi * progress))
history = []
for step in range(steps):
for group in optimizer.param_groups:
group["lr"] = lr * schedule(step)
model.train()
index = torch.randint(len(tx), (batch_size,), generator=generator).to(device)
loss = lm_loss(model, tx[index], ty[index])
optimizer.zero_grad(set_to_none=True)
loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), clip)
optimizer.step()
if step % eval_every == 0 or step + 1 == steps:
record = {"step": step + 1, "train_loss": round(loss.item(), 4),
"valid_loss": round(evaluate(model, vx, vy), 4),
"lr": round(optimizer.param_groups[0]["lr"], 6)}
history.append(record)
if log:
log(json.dumps(record))
return history
- AdamW with $\beta_2 = 0.95$ (rather than the default 0.999) is a common choice for transformers. It adapts faster when gradient statistics change.
- Warmup ramps the learning rate up over the first steps. Adam’s early estimates of gradient variance are noisy, so the first updates should be small.
- Cosine decay lowers the learning rate smoothly to 10% of its peak, so the model settles into a minimum instead of bouncing around it.
- Gradient clipping caps the global gradient norm at 1.0, so one unusual batch can’t throw the weights far off course.
- Evaluation runs on held-out data, in
eval()mode, without gradients, everyeval_everysteps.
First, prove the loop can learn
Before training on real data, run two checks that catch most bugs in minutes:
- Initial loss ≈ log V. It confirms the model and loss are wired correctly.
- Overfit one batch. Train repeatedly on a single fixed batch. A correct model and loop drive the loss near zero within a couple of hundred steps. If they don’t, there’s a bug: targets not shifted (or shifted twice), the optimizer created before parameters were replaced,
zero_gradmissing, the head disconnected, or a learning rate far too large.
Both are in your milestone tests. A model that can’t memorize one batch isn’t ready for a dataset.
Experiment 1: a task the model can master
The synthetic colors corpus repeats two phrases, “red green blue” and “blue green red”, 160 times each. With a character tokenizer and a 2-layer model, validation loss falls right along with training loss (python run.py train --corpus colors --steps 200 --context 32 --width 64 --layers 2):
{"step": 121, "train_loss": 0.0928, "valid_loss": 0.0676, ...}
{"step": 161, "train_loss": 0.0688, "valid_loss": 0.0636, ...}
{"step": 200, "train_loss": 0.0687, "valid_loss": 0.053, ...}
The validation text is new, but its patterns are the same as in training, so what the model learned transfers. This is what successful generalization looks like, on a deliberately easy task.
Experiment 2: a real text, and overfitting
Now train on The Verdict with a 512-token BPE vocabulary, 64-token windows and a 4-layer, 128-wide model of 867,072 parameters (python run.py train --steps 600). The 90/10 split gives 260 training windows and 28 validation windows. On a laptop CPU this takes about 40 seconds:
| step | train loss | valid loss |
|---|---|---|
| 1 | 6.248 | 6.201 |
| 61 | 4.833 | 5.016 |
| 121 | 3.849 | 4.424 |
| 181 | 3.100 | 4.434 |
| 241 | 2.628 | 4.704 |
| 361 | 1.599 | 5.534 |
| 481 | 0.717 | 6.099 |
| 600 | 0.437 | 6.410 |
The run starts at the uniform loss, 6.24 = log 512, as it should. Validation loss is best around step 120-180 and then gets worse, while training loss keeps falling toward zero. This is overfitting. The story is about 7,000 tokens long, and the model has 867k parameters, more than a hundred per training token. After learning the general statistics of English, its cheapest way to keep lowering the training loss is to memorize the training text verbatim. That memorized text is useless, even harmful, on the unseen 10%.
Samples make this concrete. At step 600, prompted with the story’s first words, the model reproduces the opening sentence and then degenerates into fragments of memorized phrases:
I HAD always thought Jack Gisburn rather a cheap genies be sply oweagal note that Emperors of thereerly a
Stopping at step 150, near the validation minimum, gives text that is less memorized but still clearly not fluent English. 7,000 tokens is far too little to learn a language from:
I HAD always thought, and he was the coree to doree of the coree to dorethethetheting, and the coreting
BALLM Chapter 5 shows the same effect on the same text. There’s nothing wrong with the code. The remedy is data, not cleverness: real pretraining uses trillions of tokens, so every token is seen once or a handful of times and memorization is rare. Regularization (dropout, used here at 0.1, and weight decay) and early stopping (keeping the checkpoint with the best validation loss) help at the margin.
Note
Perplexity is $e^{\text{loss}}$: here 512 at the start (uniform over 512 tokens), 83 at the validation minimum, and 1.5 on the memorized training set. Read it as “the model is as uncertain as if it were choosing uniformly among this many tokens”. Perplexities are only comparable between models with the same tokenizer: a 512-token vocabulary and a 150,000-token one measure uncertainty per very different units.
Teacher forcing and free-running generation
During training, every position sees the true previous tokens, even where the model would have predicted something else. This is teacher forcing, and it’s why one forward pass gives $T$ training examples. During generation, the model sees its own previous outputs. One early mistake changes every later input, and the model was never trained on contexts containing its own errors. So low validation loss is necessary but not sufficient for good generations. Always look at samples as well as loss curves.
Saving and loading
A checkpoint for inference needs the weights and the configuration that defines the architecture, plus whatever reconstructs the tokenizer:
def save_checkpoint(path, model, extra=None):
"""Weights + config + anything needed to rebuild the tokenizer. Not an exact-resume record:
optimizer and RNG state are deliberately omitted."""
path = Path(path)
path.parent.mkdir(parents=True, exist_ok=True)
torch.save({"config": vars(model.cfg), "state": model.state_dict(), "extra": extra or {}}, path)
def load_checkpoint(path, model_class, config_class, device="cpu"):
"""Only load checkpoints you created: torch.load of untrusted pickles can run code."""
blob = torch.load(path, map_location=device, weights_only=True)
model = model_class(config_class(**blob["config"])).to(device)
model.load_state_dict(blob["state"])
return model.eval(), blob["extra"]
Exactly resuming training also needs the optimizer state (Adam’s $m$ and $v$), the learning-rate schedule position, and the random-number-generator states. Our checkpoints deliberately omit those.
Warning
torch.loadon.pt/.binfiles unpickles Python objects, and a malicious pickle can run arbitrary code. Load only files you created, withweights_only=True. Real model distribution has moved to safetensors (Chapter 9), which is plain bytes plus a JSON header.
How big is real pretraining?
A useful rule of thumb: training costs about $6ND$ floating-point operations for $N$ parameters and $D$ training tokens. That’s 2 per parameter per token for the forward pass, and 4 for the backward. For our toy model, $6 \times 867\text{k} \times 600 \text{ steps} \times 16 \times 64$ tokens is about 3.2 TFLOP, a few seconds of a laptop. Qwen3-0.6B was trained on around 36 trillion tokens: about $1.3 \times 10^{23}$ FLOPs, or tens of thousands of GPU-days. This is why you’ll load pretrained weights from Chapter 9 on, and why fine-tuning (Part V) changes a trained model rather than starting over.
Build it
Engine milestone 7: train a GPT. In engine/train.py, implement lm_loss, evaluate and train (save/load helpers are provided).
pytest tests/test_ch07_training.py
python run.py train --impl engine --steps 600
python run.py generate --impl engine
The tests check the loss against a hand-computed cross-entropy, check that the initial loss is near log V, overfit one batch, and train on the colors task until validation loss halves.
Stretch exercises
- ★ Add early stopping to
train: keep a copy of the weights with the lowest validation loss and restore them at the end. How much better do the samples look? Where:traininengine/train.py. - ★★ Change one variable at a time (dropout 0.0 vs 0.3, width 64 vs 256, learning rate 1e-3 vs 1e-2) and record the best validation loss of each run. Predict each effect before you run it. Where:
experiments/ch07.py(create it), adaptingrun.py’scmd_trainand itsGPTConfig. - ★★ Train on a larger public-domain text (any Project Gutenberg novel, or several) with the same model. How does the gap between train and validation loss change with 10x more data? Where: put the new corpus in
data/; load it inexperiments/ch07.py(create it) usingengine.dataandengine.train. - ★★★ Implement gradient accumulation: compute gradients over $k$ small batches before each optimizer step, weighting by tokens. Verify that $k$ micro-batches of size 8 give the same update as one batch of 8k. Where:
traininengine/train.py.
Check your understanding
- Why does cross-entropy take logits rather than probabilities?
- What initial loss do you expect for a vocabulary of 50,257, and what does a much larger value suggest?
- Why can validation loss rise while training loss falls?
- Why can a model with low teacher-forced loss still produce poor generations?
- What extra state, beyond the weights, do you need to resume training exactly?
Going deeper
- BALLM Chapter 5 §§5.1-5.2 (pp. 128-150): evaluating generative models, training on The Verdict and observing overfitting. Appendix D (pp. 313-321): warmup, cosine decay, clipping.
- Kaplan et al., Scaling Laws for Neural Language Models (2020) and Hoffmann et al., Training Compute-Optimal Large Language Models (“Chinchilla”, 2022): how parameters, data and compute trade off.
- Andrej Karpathy, Let’s reproduce GPT-2 (124M) (video): a full pretraining run of the model your Chapter 9 loader will import.
8. Generating text: decoding and sampling
In this chapter
- The decoding policy, a separate decision from the model: how logits become one chosen token.
- Greedy decoding, temperature, top-k, top-p (nucleus) and min-p, with the corner cases that trip up implementations.
- Repetition penalties, stop conditions and reproducible randomness.
- The generation loop, and the first hint of why it should not recompute the past.
You will build
engine/sampling.py: sample and generate_stream, the decoding layer of your engine.
Time: 3-4 hours. GPU: not needed.
The model proposes, the policy decides
The model’s job ends at logits: one score per vocabulary entry for the next position. Turning those scores into one token is a separate decoding policy, and the same model behaves very differently under different policies. Keeping them separate in code is good engineering: you can test the policy on fixed logits, and change it per request without touching the model.
Greedy decoding
Take the highest-scoring token: logits.argmax(-1). No softmax is needed, since the largest logit always has the largest probability. Greedy is deterministic, which makes it ideal for testing (two implementations of the same model must produce identical greedy text). It’s also the best choice for short factual answers and for code.
Its weakness shows on open-ended text. Always taking the locally most likely token tends to fall into loops (“the coree to doree of the coree to doree…”, as the early-stopped model in Chapter 7 did), because the model’s own repeated output makes the repetition ever more likely.
Temperature
To sample instead, convert logits to probabilities and draw from them. Temperature $\tau$ reshapes the distribution first:
$$ p_i = \operatorname{softmax}(z / \tau)_i . $$
For logits $[2, 1, 0]$:
| τ | probabilities |
|---|---|
| 0.25 | [0.982, 0.018, 0.000] |
| 0.5 | [0.867, 0.117, 0.016] |
| 1.0 | [0.665, 0.245, 0.090] |
| 2.0 | [0.507, 0.307, 0.186] |
| 10 | [0.367, 0.332, 0.301] |
Low temperature sharpens toward greedy, and high temperature flattens toward uniform. $\tau = 1$ is the model’s own distribution. $\tau = 0$ would divide by zero, so implementations treat it as a separate greedy branch.
Truncation: top-k, top-p and min-p
Even at $\tau = 1$, the long tail of a 150,000-token vocabulary holds real probability mass. Thousands of individually unlikely tokens add up, and drawing one of them occasionally derails a generation. Truncation removes the tail before sampling:
- Top-k keeps the $k$ highest-scoring tokens and renormalizes. Simple, but fixed: $k = 40$ is too many when the model is certain and too few when many continuations are reasonable.
- Top-p (nucleus sampling, Holtzman et al. 2020) keeps the smallest set of most-likely tokens whose total probability reaches $p$. For probabilities [0.6, 0.25, 0.1, 0.05] and $p = 0.8$, it keeps the first two tokens (0.6 + 0.25 = 0.85). The set adapts: it’s small when the model is confident and large when it isn’t. Note the boundary rule: keep the token that crosses $p$, not just the tokens before it. Otherwise $p = 0.5$ would keep nothing here.
- Min-p keeps tokens whose probability is at least $p_{\min}$ times the top token’s. For [0.5, 0.3, 0.15, 0.04, 0.01] and $p_{\min} = 0.1$, the threshold is 0.05, so it keeps the first three. It scales its cutoff with the model’s confidence.
The order of operations matters: temperature, then top-k, then top-p, then min-p, then renormalize and draw. Production engines apply them in this order (sometimes with options to reorder).
def sample(logits, temperature=0.0, top_k=None, top_p=None, min_p=None, generator=None):
"""Choose one token per row of logits [B, V]; returns [B, 1]. (Your engine: Chapter 8)
temperature == 0 is greedy argmax. Otherwise: divide by temperature, keep the top_k
scores, keep the smallest prefix of the sorted distribution whose mass reaches top_p
(always keeping the first token that crosses it), drop tokens whose probability is below
min_p times the top probability, then draw from what is left.
"""
if temperature < 0:
raise ValueError("temperature must be >= 0")
if temperature == 0:
return logits.argmax(dim=-1, keepdim=True)
scores = logits.float() / temperature
if top_k is not None:
if not 1 <= top_k <= scores.shape[-1]:
raise ValueError("top_k must be in [1, vocab]")
threshold = scores.topk(top_k, dim=-1).values[..., -1:]
scores = scores.masked_fill(scores < threshold, float("-inf"))
if top_p is not None:
if not 0 < top_p <= 1:
raise ValueError("top_p must be in (0, 1]")
ordered, order = scores.sort(dim=-1, descending=True)
probs = ordered.softmax(-1)
# Remove a token when the mass BEFORE it already reaches top_p.
remove = (probs.cumsum(-1) - probs) >= top_p
ordered = ordered.masked_fill(remove, float("-inf"))
scores = torch.full_like(scores, float("-inf")).scatter(-1, order, ordered)
if min_p is not None:
probs = scores.softmax(-1)
scores = scores.masked_fill(probs < min_p * probs.amax(-1, keepdim=True), float("-inf"))
return torch.multinomial(scores.softmax(-1), 1, generator=generator)
struct Rng { // xorshift64*: tiny, seeded, reproducible
uint64_t s;
explicit Rng(uint64_t seed) : s(seed * 0x9E3779B97F4A7C15ull | 1) {}
float uniform() { s ^= s >> 12; s ^= s << 25; s ^= s >> 27; return float((s * 0x2545F4914F6CDD1Dull) >> 40) / float(1ull << 24); }
};
inline size_t sample_top_p(std::vector<float> logits, float temperature, float top_p, Rng& rng) {
std::vector<size_t> order(logits.size());
std::iota(order.begin(), order.end(), 0);
if (temperature == 0) return size_t(std::max_element(logits.begin(), logits.end()) - logits.begin());
std::sort(order.begin(), order.end(), [&](size_t a, size_t b) { return logits[a] > logits[b]; });
std::vector<float> p(order.size());
for (size_t r = 0; r < order.size(); ++r) p[r] = logits[order[r]] / temperature;
softmax(p.data(), p.size());
float mass = 0;
size_t keep = 0;
while (keep < p.size() && mass < top_p) mass += p[keep++]; // keep the token that crosses top_p
float r = rng.uniform() * mass;
for (size_t i = 0; i < keep; ++i) { if (r < p[i]) return order[i]; r -= p[i]; }
return order[keep - 1];
}
#![allow(unused)]
fn main() {
pub struct Policy {
pub temperature: f32,
pub top_k: Option<usize>,
pub top_p: Option<f32>,
}
pub fn sample(logits: &[f32], policy: &Policy, rng: &mut Rng) -> usize {
let argmax = || logits.iter().enumerate().fold(0, |best, (i, &x)| if x > logits[best] { i } else { best });
if policy.temperature == 0.0 {
return argmax();
}
// Sort token ids by score, highest first, and keep only the candidates the policy allows.
let mut order: Vec<usize> = (0..logits.len()).collect();
order.sort_by(|&a, &b| logits[b].partial_cmp(&logits[a]).unwrap());
let keep = policy.top_k.unwrap_or(logits.len()).min(logits.len());
let mut probs: Vec<f32> = order[..keep].iter().map(|&i| logits[i] / policy.temperature).collect();
softmax(&mut probs);
if let Some(p) = policy.top_p {
let (mut mass, mut cut) = (0.0, probs.len());
for (rank, q) in probs.iter().enumerate() {
if mass >= p {
cut = rank; // the token that crossed p (at rank - 1) stays
break;
}
mass += q;
}
probs.truncate(cut);
let total: f32 = probs.iter().sum();
probs.iter_mut().for_each(|q| *q /= total);
}
let mut r = rng.uniform();
for (rank, q) in probs.iter().enumerate() {
if r < *q {
return order[rank];
}
r -= q;
}
order[probs.len() - 1]
}
}
The Python version handles a batch of rows at once and implements every filter with masking (-inf for removed tokens), so it runs on the GPU without sorting through Python lists. The C++ and Rust versions sort one row explicitly, which makes the top-p boundary rule easy to see.
Repetition penalties
The repetition penalty (from the CTRL paper) discourages tokens already present in the context: positive logits are divided by the penalty, and negative ones multiplied by it. Penalties of 1.05-1.2 reduce loops, but they also discourage legitimately repeated words like names and code identifiers. Related variants subtract a fixed amount per occurrence (the frequency penalty) or once per token seen (the presence penalty). They’re all heuristics layered on the model’s distribution, so use them sparingly.
def repetition_penalty(logits, previous_ids, penalty=1.1):
"""CTRL-style penalty: shrink positive logits and grow negative logits of tokens already seen."""
if penalty == 1.0:
return logits
seen = torch.zeros_like(logits, dtype=torch.bool).scatter(-1, previous_ids, True)
adjusted = torch.where(logits > 0, logits / penalty, logits * penalty)
return torch.where(seen, adjusted, logits)
Seeds and reproducibility
Pass an explicit torch.Generator seeded per request. Then the same seed, model, policy and prompt give the same text. In a server, every request needs its own generator: if requests share one, the tokens drawn for one user depend on which other users happened to be in the batch. Reproducibility also depends on the device and kernel versions, so record them (Chapter 18’s manifest does).
The generation loop
Generation repeats: run the model, take the last position’s logits, choose a token, append it, stop if it’s a stop token or the budget is spent:
@torch.inference_mode()
def generate_stream(model, ids, max_new_tokens, cached=True, eos_ids=(), **policy):
"""Yield one new token ID at a time for a single prompt ids [1, P]. (Your engine: Chapter 8)
Prefill the whole prompt once; its last logits choose the first new token. Each later step
feeds only the newest token (cached) or the whole sequence again (uncached).
"""
if ids.ndim != 2 or ids.shape[0] != 1 or ids.shape[1] == 0:
raise ValueError("generate expects one non-empty prompt of shape [1, P]")
if ids.shape[1] + max_new_tokens > model.context_limit:
raise ValueError("Prompt plus new tokens exceed the model's context")
model.eval()
if max_new_tokens == 0:
return
cache = model.new_cache(1, ids.shape[1] + max_new_tokens) if cached else None
sequence = ids
logits = model(ids, cache)
for step in range(max_new_tokens):
token = sample(logits[:, -1], **policy)
yield int(token) # .item(): one host sync per token (Chapter 19 removes it)
if int(token) in eos_ids or step + 1 == max_new_tokens:
return
sequence = torch.cat((sequence, token), dim=1)
logits = model(token, cache) if cached else model(sequence)
def generate(model, ids, max_new_tokens, cached=True, eos_ids=(), **policy):
"""The prompt followed by the generated tokens, as a [1, P+N] tensor."""
new = list(generate_stream(model, ids, max_new_tokens, cached, eos_ids, **policy))
return torch.cat((ids, torch.tensor([new], dtype=ids.dtype, device=ids.device)), dim=1)
A few rules, each of which fixes a bug that shows up in real systems:
- The prompt’s last logits choose the first new token. No extra forward pass is needed before the first sample.
- Stop conditions: a stop token (EOS, or
<|im_end|>for chat models), a maximum number of new tokens, or the model’s context limit. Decide whether the stop token is included in the output. Here it is; a chat UI would hide it. - Reject, don’t silently crop. If prompt plus budget exceed the context, raise an error. Silently dropping the start of the prompt changes what the model sees without anyone noticing.
- One host synchronization per token.
int(token)copies the token to the CPU, which waits for the GPU to finish. To stream text you have to do that eventually, but Chapter 19 shows how to avoid doing it every step.
The cached flag switches between feeding only the new token (with a KV cache, Chapter 16) and re-running the whole sequence. Both must give identical tokens, and your milestone test checks that. The cache is what makes generation affordable; for now, model.new_cache comes provided.
Watching policies on your trained model
Prompting the model you trained in Chapter 7 (python run.py generate) shows each policy’s character, even on a tiny overfit model:
--- greedy
I HAD always thought Jack Gisburn rather a cheap genies be sply oweagal note that Emperors of thereerly a
--- t=0.8 top_k=20
I HAD always thought Jack Gisburn rat Gisburn wife's biggartw through the hush, why Jvert is, and he said, my enough the mre up the
--- t=1.0 top_p=0.9
I HAD always thought Jack Gisburn rather tap get not feltt."
Beyond sampling
Two other ideas are worth knowing by name. Beam search keeps the $b$ most likely partial sequences instead of one. It’s useful for translation, but it produces bland, repetitive text from LLMs and is rarely used for chat. Constrained decoding masks logits so the output must follow a grammar (valid JSON, a regex, a function signature). It’s a masked_fill before sampling, the same as top-k, just with a smarter mask.
Build it
Engine milestone 8: decoding. In engine/sampling.py, implement sample (greedy, temperature, top-k, top-p, min-p) and generate_stream. generate and repetition_penalty are provided.
pytest tests/test_ch08_generation.py
python run.py generate --impl engine --checkpoint runs/model.pt
The tests check that top-k=1 equals greedy, that top-p keeps exactly the crossing token, that seeded sampling is reproducible, and that cached and uncached generation produce identical tokens.
Tip
Implement top-p on sorted probabilities:
remove = (cumsum - probs) >= top_pmarks every token whose preceding mass already reached $p$, so the crossing token survives. Then scatter the masked scores back to vocabulary order withscatter.
Stretch exercises
- ★ Add
stop_stringssupport togenerate: stop when the decoded text ends with any given string. Why is this harder than stopping on token IDs? Where:generate/generate_streaminengine/sampling.py; supply a tokenizer for decoded-text matching. - ★★ Implement presence and frequency penalties and compare them with the repetition penalty on your Chapter 7 model. Where: add penalty helpers in
engine/sampling.pyand call them fromgenerate_stream. - ★★ Implement beam search with width 4 and compare its output with greedy decoding. Which has the higher total log-probability? Which reads better? Where: add a beam-search helper in
engine/sampling.py. - ★★★ Implement constrained decoding that forces outputs to be a valid decimal number: precompute, for each vocabulary entry, whether it can continue a partial number, and mask the rest. Where: add a numeric-output mask helper in
engine/sampling.pyand apply it beforesample.
Check your understanding
- Why is temperature 0 implemented as a separate branch?
- For probabilities [0.4, 0.3, 0.2, 0.1] and top-p 0.5, which tokens remain?
- Why should every request in a server have its own random generator?
- Why does generation not need an extra forward pass before choosing the first new token?
Going deeper
- BALLM §5.3 (pp. 151-158): temperature scaling and top-k sampling.
- Holtzman et al., The Curious Case of Neural Text Degeneration (2020): why greedy and pure sampling fail, and nucleus sampling.
- Keskar et al., CTRL (2019) for the repetition penalty; Nguyen et al., Turning Up the Heat: Min-p Sampling (2024).
- The vLLM and llama.cpp sampler implementations, to see the full set of production options (logit bias, penalties, grammars).
9. Real weights: safetensors and GPT-2
In this chapter
- What a model checkpoint contains: configuration, tokenizer, and tensors.
- The safetensors file format, byte by byte, and why it replaced pickle.
- Mapping someone else's tensor names and layouts onto your model, including the transpose trap.
- Proving a loaded model is correct, layer by layer, against an independent implementation.
You will build
engine/safetensors_io.py (a from-scratch reader) and load_gpt2 in engine/loaders.py. OpenAI's GPT-2 then runs in the GPT class you wrote in Chapter 6.
Time: 3-5 hours. GPU: not needed.
Why load, and what you’re loading
Chapter 7’s estimate put Qwen3-0.6B’s pretraining at about $10^{23}$ FLOPs. Nobody reproduces that to run a model. Open-weight models publish the result of training as files, and an inference engine’s first job is to read those files into exactly the computation that produced them.
A Hugging Face model snapshot is a folder:
models/gpt2/
├── config.json architecture hyperparameters: layers, widths, vocab size, epsilon, ...
├── tokenizer.json the tokenizer: vocabulary, merges, special tokens, pre-tokenizer
├── tokenizer_config.json chat template, special-token roles
├── generation_config.json default sampling settings and stop tokens
└── model.safetensors the weights (large models: model-00001-of-00004.safetensors ... + an index)
Download GPT-2 small (548 MB) to follow along:
uv pip install -r optional-requirements.txt
hf download openai-community/gpt2 --local-dir models/gpt2
The safetensors format
Older checkpoints were Python pickles (pytorch_model.bin). Unpickling can execute arbitrary code, so loading an untrusted model was a security hole. They also had to be fully deserialized before you could look at them. Safetensors fixes both with a deliberately simple layout:
┌───────────────┬──────────────────────────────┬────────────────────────────────────────┐
│ 8 bytes: N │ N bytes: JSON header (UTF-8) │ tensor data: raw little-endian bytes │
│ (u64, LE) │ │ at the offsets the header lists │
└───────────────┴──────────────────────────────┴────────────────────────────────────────┘
The header maps each tensor name to its dtype, shape and byte range within the data section:
{
"__metadata__": {"format": "pt"},
"wte.weight": {"dtype": "F32", "shape": [50257, 768], "data_offsets": [0, 154389504]},
"h.0.attn.c_attn.weight": {"dtype": "F32", "shape": [768, 2304], "data_offsets": [...]}
}
That’s the entire format. Because data is raw bytes at known offsets, a reader can memory-map the file and read any tensor without touching the others. A loader can inspect a 300 GB checkpoint’s header in milliseconds, which you’ll do in Chapter 30. You can write a valid file by hand:
import json, struct
header = json.dumps({"w": {"dtype": "F32", "shape": [2], "data_offsets": [0, 8]}}).encode()
blob = struct.pack("<Q", len(header)) + header + struct.pack("<2f", 1.5, -2.0)
open("hand.safetensors", "wb").write(blob) # a real, loadable checkpoint with one tensor
The reader is just as short:
def read_header(path):
"""Return (header dict without __metadata__, metadata dict, byte offset of the data section). (Your engine: Chapter 9)"""
with open(path, "rb") as file:
(length,) = struct.unpack("<Q", file.read(8))
if length > 100_000_000:
raise ValueError("Implausibly large safetensors header")
header = json.loads(file.read(length))
metadata = header.pop("__metadata__", {}) or {}
return header, metadata, 8 + length
def load_file(path, device="cpu", names=None):
"""Read tensors into a dict. Memory-maps the file and copies one tensor at a time. (Your engine: Chapter 9)"""
header, _, data_start = read_header(path)
selected = header if names is None else {n: header[n] for n in names}
tensors = {}
with open(path, "rb") as file, mmap.mmap(file.fileno(), 0, access=mmap.ACCESS_READ) as view:
for name, info in selected.items():
dtype = DTYPES[info["dtype"]]
begin, end = info["data_offsets"]
count = (end - begin) // torch.empty((), dtype=dtype).element_size()
if count != _numel(info["shape"]):
raise ValueError(f"{name}: byte range does not match its shape")
if count == 0: # frombuffer rejects empty reads
tensors[name] = torch.empty(info["shape"], dtype=dtype, device=device)
continue
with warnings.catch_warnings():
warnings.simplefilter("ignore") # frombuffer warns that mmap is read-only
flat = torch.frombuffer(view, dtype=dtype, count=count, offset=data_start + begin)
tensors[name] = flat.reshape(info["shape"]).to(device, copy=True)
return tensors
torch.frombuffer creates a tensor that views the mapped bytes with the right dtype, and .to(device, copy=True) makes an owned copy, on the GPU if requested. One tensor at a time keeps peak host memory low. Large models split their weights across shards plus a model.safetensors.index.json that maps each tensor to its file; snapshot_tensors walks them and checks that the index and files agree.
Dtypes, and BF16 for free
Header dtypes are short codes: F32, F16, BF16, I64, U8, F8_E4M3 and so on. BF16 is the most common for modern checkpoints, and it’s worth knowing why it’s so convenient. A BF16 number is exactly the top 16 bits of an FP32: same sign, same 8-bit exponent, mantissa cut from 23 bits to 7. Widening BF16 to FP32 is a 16-bit shift. FP16 has a different exponent width and needs real conversion logic:
# PyTorch does this for you: torch.frombuffer(raw, dtype=torch.bfloat16).float()
# By hand, on the raw 16-bit integers:
def bf16_to_f32(bits16): # int16 tensor of raw BF16 bit patterns
return (bits16.to(torch.int32) << 16).view(torch.float32)
// Read just the header: 8-byte little-endian length, then that many bytes of JSON.
inline std::string safetensors_header(const std::string& path) {
std::ifstream f(path, std::ios::binary);
if (!f) throw std::runtime_error("cannot open " + path);
uint64_t n = 0;
f.read(reinterpret_cast<char*>(&n), 8); // assumes a little-endian host, like x86 and ARM
std::string json(n, '\0');
f.read(json.data(), std::streamsize(n));
return json;
}
inline float bf16_to_float(uint16_t bits) {
uint32_t widened = uint32_t(bits) << 16; // BF16 is the top half of an FP32
float out;
std::memcpy(&out, &widened, 4);
return out;
}
#![allow(unused)]
fn main() {
/// BF16 is the top 16 bits of an f32: widening is a shift.
#[inline]
pub fn bf16_to_f32(bits: u16) -> f32 {
f32::from_bits((bits as u32) << 16)
}
/// IEEE half precision: 1 sign, 5 exponent, 10 mantissa bits.
pub fn f16_to_f32(bits: u16) -> f32 {
let sign = if bits >> 15 == 1 { -1.0 } else { 1.0 };
let exponent = ((bits >> 10) & 0x1f) as i32;
let mantissa = (bits & 0x3ff) as f32;
match exponent {
0 => sign * mantissa * 2f32.powi(-24), // subnormal
31 => if mantissa == 0.0 { sign * f32::INFINITY } else { f32::NAN },
e => sign * (1.0 + mantissa / 1024.0) * 2f32.powi(e - 15),
}
}
}
The Rust engine keeps weights in BF16 in memory and widens each value inside the matrix-vector product. That’s half the bytes of FP32, so decode runs about twice as fast, for exactly the reason Chapter 1 gave.
Mapping a checkpoint onto your model
Loading means: for every tensor in the file, find the parameter in your model it belongs to, check its shape, and copy it. Every tensor must be used, and every parameter must be filled. Reject the load if either set has leftovers, because a silently uninitialized parameter produces plausible-looking garbage.
GPT-2’s names map onto your Chapter 6 attribute names like this:
| checkpoint tensor | shape in file | your parameter | note |
|---|---|---|---|
wte.weight | [50257, 768] | token.weight | also the tied head |
wpe.weight | [1024, 768] | position.weight | |
h.{i}.ln_1.weight/.bias | [768] | blocks[i].ln1 | |
h.{i}.attn.c_attn.weight | [768, 2304] | blocks[i].attn.qkv.weight [2304, 768] | transpose |
h.{i}.attn.c_proj.weight | [768, 768] | blocks[i].attn.proj.weight | transpose |
h.{i}.mlp.c_fc.weight | [768, 3072] | blocks[i].up.weight [3072, 768] | transpose |
h.{i}.mlp.c_proj.weight | [3072, 768] | blocks[i].down.weight [768, 3072] | transpose |
h.{i}.attn.bias | [1, 1, 1024, 1024] | none | an old causal-mask buffer: skip |
ln_f.weight/.bias | [768] | norm |
The transpose trap
GPT-2 was written with a Conv1D layer that stores weights as [in, out], while nn.Linear stores [out, in] (Chapter 2). Three of the four matrices change shape when transposed, so forgetting the transpose fails loudly. But c_proj in attention is [768, 768]: square. Copying it untransposed passes the shape check and computes $x W$ instead of $x W^\top$. The model then produces fluent-looking nonsense. This is the canonical example of why shape checks are necessary but not sufficient.
@torch.no_grad()
def load_gpt2(directory, device="cpu", dtype=torch.float32):
"""GPT-2 from Hugging Face's safetensors layout. (Your engine: Chapter 9)
GPT-2 stores its projections in a Conv1D layout [in, out]; nn.Linear wants [out, in],
so those four weights are transposed. The head is tied to the token embedding.
"""
raw = read_config(directory)
if raw.get("model_type") != "gpt2" or raw.get("activation_function", "gelu_new") != "gelu_new":
raise ValueError("Expected a standard GPT-2 checkpoint")
model = GPT(GPTConfig(vocab=raw["vocab_size"], context=raw["n_positions"], width=raw["n_embd"],
layers=raw["n_layer"], heads=raw["n_head"],
eps=raw.get("layer_norm_epsilon", 1e-5))).to(device=device, dtype=dtype)
targets = {"wte.weight": (model.token.weight, False), "wpe.weight": (model.position.weight, False),
"ln_f.weight": (model.norm.weight, False), "ln_f.bias": (model.norm.bias, False)}
for i, block in enumerate(model.blocks):
for source, module, transpose in [("ln_1", block.ln1, False), ("ln_2", block.ln2, False),
("attn.c_attn", block.attn.qkv, True),
("attn.c_proj", block.attn.proj, True),
("mlp.c_fc", block.up, True), ("mlp.c_proj", block.down, True)]:
targets[f"h.{i}.{source}.weight"] = (module.weight, transpose)
targets[f"h.{i}.{source}.bias"] = (module.bias, False)
loaded, head_alias = set(), None
for name, value in snapshot_tensors(directory):
name = name.removeprefix("transformer.")
if name.endswith((".attn.bias", ".attn.masked_bias")): # old causal-mask buffers, not weights
continue
if name == "lm_head.weight": # an alias of wte when present
head_alias = value
continue
if name not in targets:
raise ValueError(f"Unexpected tensor {name}")
parameter, transpose = targets[name]
assign(parameter, value.T if transpose else value, name)
loaded.add(name)
if set(targets) - loaded:
raise ValueError(f"Missing tensors: {sorted(set(targets) - loaded)[:5]}")
if head_alias is not None and not torch.equal(head_alias.to(model.token.weight), model.token.weight):
raise ValueError("lm_head.weight disagrees with the tied token embedding")
return model.eval()
The packed c_attn also fixes the order of Q, K and V within its 2,304 outputs (Q first, then K, then V), which matches the chunk(3) in your attention. A checkpoint that packed them as K, Q, V would need reordering, and nothing in the shapes would tell you.
Proving the load is correct
“It loaded without errors” proves little. Correctness has levels, and you should climb them in order:
- Shapes and names: every tensor used, every parameter filled. The loader enforces this.
- Logits match an independent implementation on the same input IDs. Hugging Face Transformers is the usual reference. With FP32 on both sides, expect agreement to about 1e-4 or better. Your milestone test builds a small random GPT-2 with Transformers, saves it, loads it with your code, and compares logits to 1e-4.
- Greedy generations match for several prompts.
When logits disagree, don’t stare at the final output. Find the first intermediate that differs. Run both models with hooks that record each stage, and compare in order: token embeddings, position embeddings, first LayerNorm, Q/K/V of block 0, attention output, MLP output, the stream after block 0, and so on. The first mismatch localizes the bug to one operation, almost always an orientation, an epsilon, an activation variant, or an axis.
Tip
Compare fixed inputs, not generated text. Once two models choose different tokens, everything after is different input, and the comparison is meaningless.
Run GPT-2 in your engine
python run.py gpt2 --model-dir models/gpt2 --prompt "The GPU is a" --new-tokens 30
The tokenizer (from transformers) handles text at the boundary; everything else runs in your code: your safetensors reader, your loader, your GPT, your sampler. GPT-2 small is a 2019 model, so expect grammatical but meandering continuations, which is a big step up from Chapter 7’s model.
Build it
Engine milestone 9: load real weights. Implement read_header and load_file in engine/safetensors_io.py, and load_gpt2 in engine/loaders.py. save_file, snapshot_tensors and the assign helper are provided.
pytest tests/test_ch09_weights.py
python run.py gpt2 --impl engine --model-dir models/gpt2
The tests write files with every dtype (including zero-element tensors), read files written by the official safetensors library and vice versa, decode a hand-built file, and compare your GPT-2 against Transformers.
Tip
torch.frombuffer(view, dtype=..., count=..., offset=...)reads from a memory map without copying. It rejects zero-length reads, so handle empty tensors separately. Silence its “non-writable buffer” warning, then copy (.to(device, copy=True)) so the tensor outlives the mapping.
Stretch exercises
- ★ Write
inventory(dir)output for GPT-2: tensor count and total bytes. Does the total match the parameter formula from Chapter 6 times 4 bytes? Why are there extra tensors? Where:experiments/ch09.py(create it), calling the providedengine.safetensors_io.inventory. - ★★ Load GPT-2 in BF16 instead of FP32 and measure the maximum logit difference against the FP32 load. Is greedy text the same for a 50-token generation? Where:
experiments/ch09.py(create it), callingengine.loaders.load_gpt2with each dtype. - ★★ Deliberately skip the transpose for
c_projonly. Which tests or checks catch it? What does the generated text look like? Where: temporarily change thec_projassignments inload_gpt2inengine/loaders.py, then restore them. - ★★★ Write a converter that saves your Chapter 7 model as a Hugging Face–compatible GPT-2 snapshot (config, transposed weights, names), then load it with
transformersand compare logits. Where: newexperiments/export_gpt2.py, usingengine.train.load_checkpointandengine.safetensors_io.save_file.
Check your understanding
- Why can a safetensors header be read without loading any tensor data?
- Why is loading a pickle-based checkpoint from an untrusted source dangerous?
- Why can a missing transpose on a square matrix pass a shape check?
- A tied checkpoint omits
lm_head.weight. Why is that valid, and what must the loader do? - Your logits differ from the reference. What’s the first thing to compare?
Going deeper
- BALLM §§5.4-5.5 (pp. 159-168): saving and loading weights, and loading OpenAI’s GPT-2 into a from-scratch model (from TensorFlow checkpoints, with the same transposes).
- The safetensors specification and its security audit.
- Hugging Face’s
modeling_gpt2.py, especiallyConv1D, the origin of the[in, out]layout.
10. How GPUs run your code, and how to measure it
In this chapter
- Why GPUs are built so differently from CPUs, and what's inside one: SMs, warps, and a memory hierarchy.
- Asynchronous execution: why
time.time()around GPU code usually measures nothing. - Timing correctly, and profiling to find out why something is slow.
- The roofline model: predicting whether an operation is limited by arithmetic or by memory, before you write a kernel.
You will build
engine/measure.py: a synchronized timer and the roofline arithmetic you'll use to judge every optimization in the rest of the book.
Time: 4-6 hours. GPU: recommended (the arithmetic parts don't need one).
Why this part of the book exists
Your GPT works. Chapter 1 promised a fast engine, and “fast” means knowing what your hardware can do and making your code do it. The next six chapters teach GPU programming from the ground up: the execution model (this chapter), CUDA (11), matrix multiplication (12), reductions and floating point (13), Triton (14) and FlashAttention (15). Each technique gets used in the engine in Part IV.
Two philosophies of processor design
A CPU core is built to finish one stream of instructions as quickly as possible. Large caches, branch prediction and out-of-order execution all minimize the latency of each instruction. A desktop CPU has 8-32 of these sophisticated cores.
A GPU is built to finish an enormous amount of independent work, without caring how long any single piece takes. It has thousands of simple arithmetic units, small caches per unit, and no speculation. When one group of threads waits for memory, which takes hundreds of cycles, the hardware instantly switches to another group that’s ready. As long as there’s enough independent work, the arithmetic units stay busy. This is throughput-oriented design, and latency hiding through parallelism is the core idea. PMPP Chapter 1 frames the whole field this way.
Matrix multiplication across thousands of tokens and features is about as independent as work gets, which is why deep learning runs on GPUs.
Inside a GPU
An NVIDIA GPU is a collection of streaming multiprocessors (SMs): 132 on an H100, 128 on an RTX 4090. Each SM contains:
- Arithmetic units: FP32/INT32 cores, special-function units for
expandsin, and tensor cores, which perform small matrix multiplications (for example, 16×8×16 in BF16) in a single instruction. - A large register file (256 KB per SM): the fastest storage, private to each thread.
- Shared memory / L1 cache (up to about 228 KB per SM on Hopper): fast, explicitly managed, shared among the threads of a block.
- Warp schedulers that pick, every cycle, which group of threads issues next.
All SMs share an L2 cache (tens of MB) and the main device memory (HBM or GDDR, tens of GB). Approximate figures for three machines used as examples in this book:
| H100 SXM | RTX 4090 | DGX Spark (GB10) | |
|---|---|---|---|
| SMs | 132 | 128 | 48 |
| device memory | 80 GB HBM3 | 24 GB GDDR6X | 128 GB LPDDR5x (shared with CPU) |
| memory bandwidth | 3,350 GB/s | 1,008 GB/s | 273 GB/s |
| dense BF16 tensor throughput | ~990 TFLOP/s | ~165 TFLOP/s | (see NVIDIA’s datasheet) |
(Vendor datasheet figures; the dense BF16 numbers exclude structured sparsity.)
Threads execute in groups of 32 called warps. All 32 threads of a warp execute the same instruction at the same time on different data, a model NVIDIA calls SIMT (single instruction, multiple threads). Chapter 11 shows how you organize threads into blocks and grids. For now, the takeaway is the memory hierarchy: registers and shared memory are tiny but fast, device memory is huge but relatively slow, and making sure that data loaded from device memory is reused many times before being evicted is most of what “optimizing a kernel” means.
Your Python program submits work; the GPU does it later
When PyTorch runs y = x @ w on a CUDA tensor, the CPU doesn’t multiply anything. It enqueues a kernel on a stream (an ordered queue of GPU work) and returns immediately, usually within 5-20 µs. The GPU executes queued kernels in order, while the CPU races ahead submitting more.
This asynchrony is essential for performance: the GPU never waits for Python if the CPU keeps the queue full. But it has two consequences you must internalize:
- Anything that needs a GPU value on the CPU waits for the queue to drain.
print(y),y.item(),y.tolist(),if y > 0:,.cpu()and.numpy()all synchronize. In a decode loop, one.item()per token means the CPU stops, waits for the GPU, and only then submits the next step’s kernels, leaving the GPU idle while Python runs. Chapter 19 removes these. - Naive timing measures submission, not work.
Measuring time correctly
This looks reasonable and is wrong:
start = time.perf_counter()
y = x @ x # only *enqueued*
elapsed = time.perf_counter() - start # ~10 µs, regardless of the matrix size
Two correct options:
- Synchronized wall time. Call
torch.cuda.synchronize()before starting the clock (so earlier queued work isn’t counted) and after the operation (so this work is). This measures what a user experiences, including CPU overheads. - CUDA events. Record an event before and after the work on the GPU’s own timeline and ask for the elapsed time between them. This excludes CPU-side gaps, so it measures kernel time.
def measure_wall(function, samples=10, warmup=3, cuda=None):
"""Median wall-clock milliseconds of function(), synchronizing the GPU before starting and
after finishing each sample, so queued work is neither excluded nor borrowed. (Your engine: Chapter 10)"""
cuda = torch.cuda.is_available() if cuda is None else cuda
for _ in range(warmup):
function()
times = []
for _ in range(samples):
if cuda:
torch.cuda.synchronize()
start = time.perf_counter()
function()
if cuda:
torch.cuda.synchronize()
times.append((time.perf_counter() - start) * 1000)
return {"median_ms": statistics.median(times), "min_ms": min(times), "max_ms": max(times), "samples": samples}
def measure_cuda(function, warmup=5, samples=20, repeats=1):
"""GPU-side milliseconds between two CUDA events recorded on the current stream."""
for _ in range(warmup):
function()
torch.cuda.synchronize()
times = []
for _ in range(samples):
start, end = torch.cuda.Event(enable_timing=True), torch.cuda.Event(enable_timing=True)
start.record()
for _ in range(repeats):
function()
end.record()
end.synchronize()
times.append(start.elapsed_time(end) / repeats)
return {"median_ms": statistics.median(times), "min_ms": min(times), "max_ms": max(times), "samples": samples}
Three more rules make measurements trustworthy:
- Warm up. The first call pays one-time costs: CUDA context creation, library loading, kernel compilation (Triton,
torch.compile) and memory-pool growth. Run a few iterations untimed. - Report a distribution. Take the median of many samples and keep the min and max. GPU clocks, other processes and thermal throttling all add noise.
- Fix everything except the variable you’re testing. Shapes, dtype, device and input values (some kernels are faster on zeros!) must stay the same across the comparison.
Profiling: why is it slow?
A timer tells you how long; a profiler tells you why. PyTorch’s profiler records every operator and kernel, CPU side and GPU side:
python run.py profile # writes runs/trace.json and prints the top operators
Open the trace in Perfetto (or chrome://tracing). You’ll see two timelines: the CPU thread submitting operators, and the GPU stream executing kernels. Look for:
- Gaps on the GPU timeline: the GPU is idle waiting for the CPU. The cause is Python overhead, synchronization, or too many tiny kernels. Fixes: fusion, CUDA graphs, batching (Chapters 14, 19, 24).
- Unexpected kernels: a
copy_orcontiguousyou didn’t mean to cause, or a dtype conversion. - One kernel dominating: optimize that, and check with the roofline whether it’s already near its limit.
For kernel-level detail, NVIDIA’s Nsight Systems (nsys) shows the whole-program timeline, and Nsight Compute (ncu) shows one kernel’s achieved bandwidth, occupancy and stall reasons. GPU Mode L1 walks through all three tools.
Tip
Profilers add overhead. Use a profile to explain, and a separate un-profiled run to measure.
The roofline model
Before optimizing anything, ask what limits it. Every kernel does some arithmetic (FLOPs) and moves some bytes to and from device memory. Its arithmetic intensity is the ratio:
$$ I = \frac{\text{FLOPs}}{\text{bytes moved}} . $$
A device can do at most $F$ FLOP/s and move at most $W$ bytes/s. So the best achievable throughput is
$$ \text{attainable FLOP/s} = \min(F,; I \cdot W). $$
Plot that against $I$ and you get a “roofline”: a slanted line (memory-bound) that meets a flat line (compute-bound) at the ridge point $I^* = F/W$. For an H100, $I^* \approx 990{,}000/3{,}350 \approx 295$ FLOPs per byte. An operation with lower intensity can’t use the tensor cores fully, no matter how good the kernel is.
Now place real operations on it:
- Vector add in FP32 reads 8 bytes and writes 4 per element, and does 1 FLOP: $I = 1/12$. Hopelessly memory-bound. All you can do is approach the bandwidth limit.
- Matrix × vector (one token of decode through a
[4096, 4096]BF16 weight) does $2 \cdot 4096^2$ FLOPs and reads the $2 \cdot 4096^2$-byte matrix: $I \approx 1$. Memory-bound by a factor of about 300 on an H100. This is the Chapter 1 result again, now with a name. - Batched decode with $B$ tokens per weight read: $I \approx B$. That’s 1 at batch 1, 8 at batch 8, about 120 at batch 128, and about 410 at batch 512 (the input and output reads start to count). Batching is how decode climbs the roofline (Chapter 24).
- Prefill of a 4,096-token prompt through the same weight is a
[4096, 4096] @ [4096, 4096]GEMM: $I \approx 1{,}365$, comfortably compute-bound.
This is the most important analytical tool in the book. Before writing any kernel, compute its intensity and its ceiling, so you know when to stop optimizing.
The milestone turns these estimates into functions you can call from any experiment:
def matmul_intensity(m, n, k, bytes_per_element=2):
"""FLOPs per byte of [M,K] @ [K,N] if every operand is read once and C written once. (Your engine: Chapter 10)"""
flops = 2 * m * n * k
traffic = (m * k + k * n + m * n) * bytes_per_element
return flops / traffic
def attainable_flops(intensity, peak_flops, bandwidth):
"""The roofline: you cannot exceed the compute peak, nor bandwidth x intensity. (Your engine: Chapter 10)"""
return min(peak_flops, intensity * bandwidth)
def decode_ceiling(weight_bytes, bandwidth, batch=1, kv_bytes_per_sequence=0):
"""Upper bound on decode tokens/s when every step must read all weights once (shared by the
batch) plus each sequence's own KV cache. (Your engine: Chapter 10)"""
step_bytes = weight_bytes + batch * kv_bytes_per_sequence
return batch * bandwidth / step_bytes
Amdahl’s law: optimize what matters
If an operation takes a fraction $f$ of the total time and you make it $s$ times faster, the whole program gets
$$ \text{speedup} = \frac{1}{(1-f) + f/s} $$
times faster. Make a 10% operation infinitely fast and the program speeds up by only 1.11x. That’s why you profile the whole engine first (Chapter 19), and then optimize the biggest bar, not the most interesting kernel.
Important
DGX Spark. CPU and GPU share 128 GB of LPDDR5x. Two consequences: the 273 GB/s bandwidth is shared too (a CPU process streaming memory slows your GPU decode), and Linux’s “free” memory figure is misleading. Watch “available” in
free -hand usetorch.cuda.mem_get_info(). Moving data betweencpuandcudatensors still copies, even though the bytes sit in the same physical memory.
Build it
Engine milestone 10: measure and predict. In engine/measure.py, implement measure_wall (synchronized, with warmup and a median), matmul_intensity, attainable_flops and decode_ceiling. measure_cuda and amdahl_speedup are provided.
pytest tests/test_ch10_performance.py
python run.py profile
Then use them: predict the decode ceiling for GPT-2 small on your device, measure generate from Chapter 8 (python run.py cache does a version of this), and compute the fraction of the ceiling you achieve.
Stretch exercises
- ★ Time
x @ xfor a 4096² FP32 matrix three ways: unsynchronized wall clock, synchronized wall clock, and CUDA events. Explain each number. Where:experiments/ch10.py(create it), usingengine.measure.measure_wallandmeasure_cuda. - ★★ Measure the achieved bandwidth of
x + yfor vector sizes from $2^{10}$ to $2^{28}$ elements. Plot GB/s against size. Where does launch overhead dominate, and what fraction of datasheet bandwidth do you reach at the top? Where:experiments/ch10.py(create it). - ★★ Measure the TFLOP/s of
a @ bin BF16 for square sizes 256 to 8,192 and plot them on your device’s roofline. At what size do you reach 70% of peak? Where:experiments/ch10.py(create it), usingengine.measurefor timing and roofline arithmetic. - ★★★ Profile one decode step of your GPT-2 with
torch.profiler. Count the kernels, and estimate the fraction of step time spent in launch gaps rather than kernels. Where:experiments/ch10.py(create it), wrapping a decode step intorch.profiler.profile.
Check your understanding
- Why must you synchronize before starting a timer as well as after the operation?
- Why should profiling and benchmarking be separate runs?
- What is the arithmetic intensity of decode for a batch of 16 sequences, and is it memory-bound on an H100?
- A kernel taking 30% of runtime becomes 3x faster. What’s the overall speedup?
- On DGX Spark, why can a CPU-heavy job slow down GPU decoding?
Going deeper
- PMPP Chapter 1 (heterogeneous computing, latency vs throughput), Chapter 4 §§4.1-4.7 (GPU architecture, warps, scheduling, occupancy), and §22.5 (batching: latency versus throughput).
- GPU Mode L1 (Mark Saroufim, profiling and integrating kernels in PyTorch, with
nsys/ncuexamples) and L8 (the CUDA performance checklist). - Williams, Waterman and Patterson, Roofline: An Insightful Visual Performance Model (2009).
- NVIDIA Hopper and Ada architecture whitepapers for exact SM resources.
11. CUDA programming
In this chapter
- The CUDA execution model: kernels, threads, blocks and grids, and how each thread finds its data.
- Writing, launching and error-checking kernels, standalone and as PyTorch extensions.
- Warps, divergence and coalescing: the memory-access rule that matters most.
- A first matrix-multiplication kernel, and why it's slow.
You will build
Your first CUDA kernels in engine/kernels/cuda_ops.cu: vector addition and a naive matrix multiplication, compiled into PyTorch and tested against it.
Time: 6-8 hours. GPU: an NVIDIA GPU to run the kernels; without one you can still write them and compile-check them (see the end of the chapter).
One function, many threads
A CUDA kernel is a function that runs once per thread, on thousands of threads at the same time. Every thread executes the same code; what differs is its index, which it uses to decide which data to work on. This is the SPMD style (single program, multiple data), and the hardest part of learning it is to stop thinking “loop over elements” and start thinking “I am one element”.
Threads are organized in two levels:
- A block is a group of up to 1,024 threads that run on the same SM. They can cooperate through shared memory and synchronize with barriers (
__syncthreads()). - A grid is all the blocks of one launch. Blocks are independent: the hardware may run them in any order, in parallel or one after another, depending on how many fit on the GPU. A kernel must never assume an order between blocks.
That independence is what makes CUDA programs scale: the same kernel runs on a 48-SM laptop GPU and a 132-SM H100, just with more blocks in flight on the bigger one. PMPP calls this transparent scalability.
Vector addition
The “hello world” of CUDA adds two vectors. Each thread computes its global index from three built-in variables: blockIdx.x (which block), blockDim.x (threads per block) and threadIdx.x (which thread in the block):
$$ i = \texttt{blockIdx.x} \times \texttt{blockDim.x} + \texttt{threadIdx.x} $$
For $N = 1{,}000$ elements and 256 threads per block, we need $\lceil 1000 / 256 \rceil = 4$ blocks, which is 1,024 threads. The last 24 threads have no element to process, so every kernel includes a bounds check. Here’s the same kernel in three languages:
__global__ void add_kernel(const float* a, const float* b, float* out, int64_t n) {
int64_t i = int64_t(blockIdx.x) * blockDim.x + threadIdx.x; // this thread's element
if (i < n) out[i] = a[i] + b[i]; // the grid may overshoot n
}
@triton.jit
def add_kernel(x_ptr, y_ptr, out_ptr, n, BLOCK: tl.constexpr):
"""(Your engine: Chapter 14)"""
pid = tl.program_id(0) # which block of BLOCK elements is mine
offsets = pid * BLOCK + tl.arange(0, BLOCK) # a vector of BLOCK indices
mask = offsets < n # the last block may run past the end
x = tl.load(x_ptr + offsets, mask=mask)
y = tl.load(y_ptr + offsets, mask=mask)
tl.store(out_ptr + offsets, x + y, mask=mask)
def vector_add(x, y, block=1024):
check_device(x, y)
if x.shape != y.shape or not x.is_contiguous() or not y.is_contiguous():
raise ValueError("vector_add needs equal-shape contiguous tensors")
out = torch.empty_like(x)
n = x.numel()
add_kernel[(triton.cdiv(n, block),)](x, y, out, n, BLOCK=block)
return out
#![allow(unused)]
fn main() {
#[cuda_module]
mod kernels {
use super::*;
#[kernel]
pub fn twice(input: &[f32], mut output: DisjointSlice<f32>) {
let index = thread::index_1d();
if let Some(value) = output.get_mut(index) {
*value = input[index.get()] * 2.0;
}
}
}
}
The CUDA C++ version is one thread’s view of the work. Triton (Chapter 14) is one block’s view: a program handles BLOCK elements as a vector, and the mask plays the role of the bounds check. NVIDIA’s experimental cuda-oxide compiles Rust to GPU code; its DisjointSlice turns the bounds check into a type-system guarantee that no two threads write the same element.
Launching a kernel
From host code (the CPU side), a launch names the grid and block sizes in triple angle brackets:
int threads = 256;
int blocks = (n + threads - 1) / threads; // ceiling division
add_kernel<<<blocks, threads, 0, stream>>>(a, b, out, n);
Before that, the data must be in GPU memory. A standalone program allocates device buffers and copies to and from them explicitly (cudaMalloc, cudaMemcpy). The book’s standalone example does this with error checking and CUDA-event timing. Inside PyTorch, tensors already live on the device. You write a small C++ wrapper that checks its inputs, launches the kernel on PyTorch’s current stream, and returns a tensor:
torch::Tensor vector_add(torch::Tensor a, torch::Tensor b) {
check(a); check(b); // CUDA, float32, contiguous
TORCH_CHECK(a.sizes() == b.sizes(), "shapes differ");
c10::cuda::CUDAGuard guard(a.device()); // launch on a's GPU
auto out = torch::empty_like(a);
int64_t n = a.numel();
if (n) { // zero-size launches are invalid
add_kernel<<<(n + 255) / 256, 256, 0, stream()>>>(
a.data_ptr<float>(), b.data_ptr<float>(), out.data_ptr<float>(), n);
C10_CUDA_KERNEL_LAUNCH_CHECK(); // surface launch errors immediately
}
return out;
}
torch.utils.cpp_extension.load compiles the .cu file the first time you call it and imports the result as a Python module (izh/kernels/cuda.py). The wrapper’s checks are part of correctness. A kernel receives raw pointers and knows nothing about shapes, dtypes or strides; pass it a transposed tensor and it silently reads the wrong elements (Chapter 2).
Warning
Kernel launches are asynchronous, and so are their errors. An out-of-bounds write may only be reported at the next synchronizing call, far from the bug. When something crashes mysteriously, rerun with
CUDA_LAUNCH_BLOCKING=1(launches become synchronous) or undercompute-sanitizer(NVIDIA’s memory checker), which pinpoints the offending kernel and thread.
Two-dimensional grids
Grids and blocks can be 2-D or 3-D, which maps naturally onto matrices. For an M × N output with 16×16 blocks:
dim3 block(16, 16);
dim3 grid((N + 15) / 16, (M + 15) / 16);
int row = blockIdx.y * blockDim.y + threadIdx.y;
int col = blockIdx.x * blockDim.x + threadIdx.x;
A row-major matrix element (row, col) lives at row * N + col. For a width-5 matrix, (2, 3) is at offset 13.
Warps, divergence and occupancy
The hardware runs each block as warps of 32 consecutive threads that execute one instruction at a time, in lockstep. Two performance consequences follow:
- Divergence. If threads in a warp take different branches, the warp runs both paths one after the other, with some threads masked off each time. The bounds check in vector add diverges only in the last warp, which is harmless. Data-dependent branches in a hot loop can halve throughput or worse.
- Occupancy. An SM keeps many warps resident and switches between them to hide memory latency (Chapter 10). How many fit depends on each thread’s registers and each block’s shared memory: use more, and fewer warps fit. Occupancy is resident warps divided by the hardware maximum. It’s a means, not a goal. A kernel with low occupancy but lots of data reuse can beat one with high occupancy and none.
Coalescing: the rule that matters most
When the 32 threads of a warp load from global memory, the hardware combines their requests into as few memory transactions as possible. If thread $t$ reads element $\text{base} + t$, the 32 four-byte reads cover one contiguous 128-byte segment, served by a single transaction. That’s a coalesced access, and it uses the full bandwidth. If thread $t$ reads element $\text{base} + 32t$ (a stride), the warp touches 32 different segments and wastes most of each. Strided access can be 10-30x slower for the same number of useful bytes.
The practical rule: map threadIdx.x, which varies fastest within a warp, to the index that varies fastest in memory (the last axis of a row-major tensor). GPU Mode L8 measures this effect directly (coalesce.cu in the lecture repository).
A naive matrix multiplication
The direct translation of $C_{ij} = \sum_k A_{ik} B_{kj}$ gives each thread one output element:
// One thread per output element. Every thread re-reads a whole row of A and column of B
// from global memory: 2K loads for 2K flops, an arithmetic intensity of ~0.25 flop/byte.
__global__ void naive_matmul_kernel(const float* a, const float* b, float* c,
int M, int K, int N) {
int row = blockIdx.y * blockDim.y + threadIdx.y;
int col = blockIdx.x * blockDim.x + threadIdx.x; // x varies fastest -> coalesced B and C
if (row < M && col < N) {
float total = 0.f;
for (int k = 0; k < K; ++k) total += a[row * K + k] * b[k * N + col];
c[row * N + col] = total;
}
}
Check the address formulas against the shapes: A is [M, K] so its row stride is K, B is [K, N] so its row stride is N. A common bug uses one width for both, which works only for square matrices. Test with non-square shapes like [3, 5] @ [5, 7].
Is it coalesced? Within a warp, col varies (x is the fastest thread index) and row is fixed. Reads of B[k * N + col] are consecutive, which is coalesced. Reads of A[row * K + k] are the same address for every thread, a broadcast, which is fine. Writes to C[row * N + col] are consecutive, also good.
So why is it slow? Apply Chapter 10’s roofline. Each thread loads $2K$ values to do $2K$ FLOPs: an arithmetic intensity of about 0.25 FLOP per byte in FP32. Every element of A is re-read by each of the N threads that use it, and every element of B by M threads. The kernel is starved for memory bandwidth, and its speed is a small fraction of what the arithmetic units can do. Fixing that, by loading tiles into shared memory and reusing them many times, is the subject of the next chapter.
Build it
Engine milestone 11: your first kernels. In engine/kernels/cuda_ops.cu, write add_kernel and naive_matmul_kernel. The PyTorch bindings and input checks are already there. Leave the tiled matmul and reductions for Chapters 12 and 13.
uv pip install ninja setuptools # build tools for PyTorch extensions
export CUDA_HOME=/usr/local/cuda # where your CUDA Toolkit lives
pytest tests/test_ch11_cuda.py -k "vector_add or matmuls"
python run.py kernels --backend cuda --impl engine
The tests cover lengths 0, 1, 31, 32, 33 and 1,000,003 (around warp and block boundaries) and non-square matrices.
No GPU? You can still write and compile the kernels. NVIDIA’s compiler is pip-installable, so a compile check catches syntax and type errors:
uv pip install "nvidia-cuda-nvcc==13.3.*" "nvidia-cuda-runtime==13.3.*" "nvidia-cuda-cccl==13.3.*"
NV=$(python -c "import nvidia, os; print(os.path.dirname(nvidia.__path__[0]))")/nvidia/cu13
$NV/bin/nvcc -std=c++17 -arch=sm_90 -I$NV/include -I$NV/include/cccl cpp/kernels.cu -o build/kernels
Important
DGX Spark compiles for
-arch=sm_121with CUDA 13. PyTorch’s extension builder picks the architecture from the GPU it finds, or fromTORCH_CUDA_ARCH_LIST="12.1".
Stretch exercises
- ★ Make vector add process 4 elements per thread with
float4loads (reinterpret_cast<const float4*>). Measure the bandwidth before and after on 2²⁸ elements. Where:add_kernelinengine/kernels/cuda_ops.cu; handle the tail when the length is not divisible by four. - ★★ Write a deliberately uncoalesced copy kernel (thread $t$ reads element $32t \bmod N$) and compare its bandwidth with a coalesced copy. Where: add copy kernels and host launchers in
engine/kernels/cuda_ops.cu, then expose them in itsPYBIND11_MODULE. - ★★ Swap the roles of x and y in the naive matmul (map
threadIdx.xto rows). Measure the slowdown and explain it with coalescing. Where:naive_matmul_kerneland its launch geometry inengine/kernels/cuda_ops.cu. - ★★★ Write
transpose_kernelfor a[M, N]FP32 matrix. A naive version must choose between coalesced reads and coalesced writes. Then fix it with a 32×32 shared-memory tile (you’ll need the bank-conflict padding from Chapter 12). Where: addtranspose_kernel, a host launcher and a binding inengine/kernels/cuda_ops.cu.
Check your understanding
- Why launch more threads than elements and then bounds-check, instead of launching exactly N threads?
- Why can’t one block wait for another block to finish?
- A warp reads
x[threadIdx.x * 16](FP32). How many 128-byte segments does it touch, and how much of the transferred data is useful? - What is the naive matmul’s arithmetic intensity, and what does the roofline predict about its speed?
Going deeper
- PMPP Chapters 2-3 (pp. 23-73): data parallelism, CUDA program structure, multidimensional grids and the naive matmul. Chapter 4 (warps, divergence, occupancy) and §6.1 (coalescing).
- GPU Mode L2 (a recap of PMPP Chapters 1-3), L3 (Jeremy Howard, CUDA for Python programmers, writing kernels inline from a notebook), L4 (compute and memory architecture), and L8 (the performance checklist, with runnable
coalesce.cu,divergence.cuandoccupancy.cu). - NVIDIA’s CUDA C++ Programming Guide, chapters on the programming model and memory hierarchy;
compute-sanitizerdocumentation.
12. Fast matrix multiplication
In this chapter
- Data reuse: why the naive kernel is slow, and how tiling into shared memory fixes it.
- Barriers, boundary tiles and shared-memory bank conflicts.
- Register tiling (thread coarsening) and tensor cores: how libraries reach most of peak.
- Why decode's matrix-vector products are a different problem from prefill's GEMMs.
You will build
engine/tiling.py: a tiled matmul that counts its memory traffic (runs anywhere), and the tiled CUDA kernel in engine/kernels/cuda_ops.cu (needs a GPU).
Time: 6-8 hours. GPU: optional.
Why matrix multiplication deserves a chapter
In a transformer, nearly all arithmetic is matrix multiplication: the QKV, output and MLP projections of every layer, and the head. Prefill is a sequence of large GEMMs (general matrix-matrix multiplications), and their efficiency sets time-to-first-token. You’ll rarely beat NVIDIA’s cuBLAS or CUTLASS for large GEMMs, and you shouldn’t try. But every kernel you’ll write later (FlashAttention, quantized matmuls, paged attention) is built from the same tiling ideas. This chapter is where you learn them.
Counting the waste
Chapter 11’s naive kernel computes each output with $2K$ loads from global memory. For $M = N = K = 4096$, that’s $2 \times 4096^3 \approx 137$ billion loads to compute 16.8 million outputs, while the matrices themselves contain only $3 \times 4096^2 \approx 50$ million distinct values. Every element of A and B is fetched 4,096 times. The L2 cache catches some of the repeats, but the arithmetic units still mostly wait.
Tiling: load once, use many times
Divide the output into $T \times T$ tiles and give each tile to one thread block. To compute its tile, the block needs a $T$-row strip of A and a $T$-column strip of B. Walk along $K$ in phases. In each phase, the block cooperatively copies one $T\times T$ tile of A and one of B into shared memory, one element per thread, and then every thread accumulates $T$ products from those tiles:
Each loaded value is now used $T$ times instead of once, so global traffic drops by a factor of $T$:
$$ \text{loads}{\text{tiled}} = \frac{2MNK}{T}, \qquad I{\text{phase}} = \frac{2T^3 \text{ FLOPs}}{2T^2 \times 4 \text{ bytes}} = \frac{T}{4} \text{ FLOP/byte (FP32)}. $$
$T = 16$ gives an intensity of 4 instead of 0.25. That’s a 16x higher ceiling on the roofline.
You can verify the traffic arithmetic without a GPU. tiled_matmul performs the same tile-by-tile computation on the CPU and counts every element it “loads into shared memory”:
def tiled_matmul(a, b, tile=16):
"""C = A @ B computed one TILE x TILE output tile at a time, the way a CUDA block does it.
Returns (C, elements_loaded): every tile of A and B copied into "shared memory" is counted
as a global-memory load, so you can compare traffic with the naive kernel's 2*M*N*K. (Your engine: Chapter 12)
"""
m, k = a.shape
n = b.shape[1]
c = torch.zeros(m, n, dtype=torch.float32)
loads = 0
for row in range(0, m, tile): # one "block" per output tile
for col in range(0, n, tile):
acc = torch.zeros(min(tile, m - row), min(tile, n - col))
for phase in range(0, k, tile): # march along K one tile at a time
a_tile = a[row:row + tile, phase:phase + tile].float() # load into shared memory
b_tile = b[phase:phase + tile, col:col + tile].float()
loads += a_tile.numel() + b_tile.numel()
acc += a_tile @ b_tile # reuse each value `tile` times
c[row:row + tile, col:col + tile] = acc
return c, loads
For 64×64 matrices, it loads exactly $1/T$ as many elements as the naive kernel, for every tile size. That’s what the milestone tests check.
The CUDA kernel
constexpr int TILE = 16;
// Each block computes a TILE x TILE output tile. In phase p every thread loads one element of
// A's tile and one of B's into shared memory; then all threads reuse those 2*TILE^2 values
// TILE times each. Global traffic drops by a factor of TILE.
__global__ void tiled_matmul_kernel(const float* a, const float* b, float* c, int M, int K, int N) {
__shared__ float As[TILE][TILE];
__shared__ float Bs[TILE][TILE];
int row = blockIdx.y * TILE + threadIdx.y;
int col = blockIdx.x * TILE + threadIdx.x;
float total = 0.f;
for (int phase = 0; phase < K; phase += TILE) {
int ak = phase + threadIdx.x, bk = phase + threadIdx.y;
As[threadIdx.y][threadIdx.x] = (row < M && ak < K) ? a[row * K + ak] : 0.f; // zero-pad edges
Bs[threadIdx.y][threadIdx.x] = (bk < K && col < N) ? b[bk * N + col] : 0.f;
__syncthreads(); // tile fully loaded before anyone reads it
for (int k = 0; k < TILE; ++k) total += As[threadIdx.y][k] * Bs[k][threadIdx.x];
__syncthreads(); // everyone done reading before the next overwrite
}
if (row < M && col < N) c[row * N + col] = total;
}
@triton.jit
def matmul_kernel(a_ptr, b_ptr, c_ptr, M, N, K,
stride_am, stride_ak, stride_bk, stride_bn, stride_cm, stride_cn,
BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr,
GROUP_M: tl.constexpr):
"""(Your engine: Chapter 14)"""
# Map the 1-D program id to an output tile. Walking GROUP_M tile-rows at a time keeps the
# B tiles those rows need hot in L2 cache ("grouped ordering", Triton tutorial 03).
pid = tl.program_id(0)
tiles_m, tiles_n = tl.cdiv(M, BLOCK_M), tl.cdiv(N, BLOCK_N)
group = pid // (GROUP_M * tiles_n)
first_m = group * GROUP_M
group_size = min(tiles_m - first_m, GROUP_M)
pid_m = first_m + (pid % group_size)
pid_n = (pid % (GROUP_M * tiles_n)) // group_size
rm = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
rn = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
rk = tl.arange(0, BLOCK_K)
acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32) # accumulate in FP32 whatever the inputs
for k0 in range(0, K, BLOCK_K):
a = tl.load(a_ptr + rm[:, None] * stride_am + (k0 + rk)[None, :] * stride_ak,
mask=(rm[:, None] < M) & ((k0 + rk)[None, :] < K), other=0.0)
b = tl.load(b_ptr + (k0 + rk)[:, None] * stride_bk + rn[None, :] * stride_bn,
mask=((k0 + rk)[:, None] < K) & (rn[None, :] < N), other=0.0)
acc += tl.dot(a, b) # tensor cores when dtypes allow
c = acc.to(c_ptr.dtype.element_ty)
tl.store(c_ptr + rm[:, None] * stride_cm + rn[None, :] * stride_cn, c,
mask=(rm[:, None] < M) & (rn[None, :] < N))
def matmul(a, b, block_m=64, block_n=64, block_k=32, group_m=8):
check_device(a, b)
if a.ndim != 2 or b.ndim != 2 or a.shape[1] != b.shape[0]:
raise ValueError("matmul expects [M, K] @ [K, N]")
m, k = a.shape
n = b.shape[1]
c = torch.empty((m, n), device=a.device, dtype=a.dtype)
grid = (triton.cdiv(m, block_m) * triton.cdiv(n, block_n),)
matmul_kernel[grid](a, b, c, m, n, k, a.stride(0), a.stride(1), b.stride(0), b.stride(1),
c.stride(0), c.stride(1), BLOCK_M=block_m, BLOCK_N=block_n,
BLOCK_K=block_k, GROUP_M=group_m)
return c
Three details make the CUDA version correct:
- Two barriers per phase. The first
__syncthreads()ensures the whole tile is loaded before anyone reads it. The second ensures everyone has finished reading before the next phase overwrites the tile. Remove the second and fast threads corrupt the tile under slow ones. The result is wrong only sometimes, which makes it a nightmare to debug. - Zero-padded edges. When $M$, $N$ or $K$ isn’t a multiple of $T$, threads that would load out of bounds load 0 instead. Zeros don’t change the dot product, so edge tiles need no special arithmetic.
- Every thread reaches every barrier, including threads whose output is outside the matrix. They still help load tiles; they just skip the final write. A thread that returns early would make the barrier wait forever, or worse.
The Triton version (Chapter 14 explains the language) expresses the same algorithm at the level of tiles: tl.load a [BLOCK_M, BLOCK_K] block, tl.dot it with a [BLOCK_K, BLOCK_N] block, and accumulate in FP32. Triton handles shared memory and synchronization itself. It also maps tl.dot onto tensor cores, which our CUDA kernel doesn’t use.
Shared-memory bank conflicts
Shared memory is divided into 32 banks, and consecutive 4-byte words go to consecutive banks. A warp’s 32 accesses complete in one cycle if they hit 32 different banks (or the same word, which is broadcast). If several threads hit different words in the same bank, the accesses are serialized.
Reading a row of a float tile[32][32] across a warp touches 32 consecutive words: 32 banks, no conflict. Reading a column touches words 32 apart, all in the same bank: a 32-way conflict. The classic fix is padding: declare float tile[32][33]. Now advancing one row moves by 33 words, which is one bank over, so a column spans all 32 banks. Our tiled matmul reads As[ty][k] (a broadcast within a warp) and Bs[k][tx] (a row), so it’s conflict-free without padding. The transpose exercise in Chapter 11 needs the padding.
Register tiling and thread coarsening
Tiling into shared memory raised the intensity of global memory traffic. But each FMA in the inner loop still reads two values from shared memory, which has its own bandwidth limit. The next level of reuse moves into registers: each thread computes a small block of outputs, say $4\times4$ or $8\times8$, instead of one. Each value it reads from shared memory then feeds 4 or 8 FMAs held in registers:
for k in tile:
a_frag[0..7] = As[ty*8 .. ty*8+7][k] # 8 values of A
b_frag[0..7] = Bs[k][tx*8 .. tx*8+7] # 8 values of B
acc[i][j] += a_frag[i] * b_frag[j] # 64 FMAs from 16 loads
This is thread coarsening (PMPP §6.5 and Chapter 15): fewer threads, each doing more work, with higher reuse. The costs are more registers per thread (lower occupancy) and more complex code. Well-tuned FP32 SIMT kernels reach roughly 50-70% of FP32 peak this way.
Tensor cores
Modern NVIDIA GPUs have tensor cores: units that compute a small matrix product, such as a 16×8×16 BF16 tile with FP32 accumulation, as a single warp-wide instruction. Their throughput is roughly 8-16x that of the ordinary FP32 cores, which is where the “~990 TFLOP/s” of an H100 comes from. Using them requires BF16, FP16, FP8 or INT8 inputs, data arranged in specific fragment layouts, and enough reuse to feed them, so tiling matters even more.
You can program tensor cores with CUDA’s wmma/mma instructions (GPU Mode L23), with CUTLASS/CuTe templates (L15, L36, L57), or let Triton’s tl.dot do it. For the engine, the right choice is clear: use cuBLAS (torch.matmul) for the big GEMMs, and write kernels only where you can fuse something it can’t. Typical results for a 4096³ BF16 GEMM look like this:
| kernel | fraction of tensor-core peak |
|---|---|
| naive (Chapter 11), FP32 cores | under 1% |
| shared-memory tiled, FP32 cores | 2-5% |
| register-tiled, FP32 cores | 5-10% |
Triton tl.dot, BF16, autotuned | 60-85% |
| cuBLAS, BF16 | 70-90% |
(The fractions are relative to the tensor-core peak, which is why even a well-tuned FP32 SIMT kernel looks small here. Measure your own device in the stretch exercises.)
Decode is a different problem: GEMV
During decode, each projection multiplies a single vector (or a few, with batching) by a weight matrix. That’s a matrix-vector product (GEMV), with an arithmetic intensity of about 1 (Chapter 10). Tiling can’t help, because there’s nothing to reuse: each weight is used exactly once per token. A good GEMV kernel just streams the matrix at full bandwidth:
- coalesced, vectorized loads (each thread reads 16 bytes at a time, consecutive threads read consecutive memory);
- enough parallelism: one or more warps per output row, combining partial sums with warp reductions (Chapter 13);
- no wasted bytes: weights stored compactly, ideally in 4-bit form and dequantized in registers (Chapter 20).
That’s why decode optimization in practice is quantization plus bandwidth-efficient GEMV, and prefill optimization is GEMM. Your Rust engine’s matvec (rows split across CPU threads, BF16 widened in the inner loop) is a CPU GEMV in exactly this spirit.
Build it
Engine milestone 12: tiling. Implement tiled_matmul in engine/tiling.py. On a GPU, also write tiled_matmul_kernel in engine/kernels/cuda_ops.cu.
pytest tests/test_ch12_matmul.py # runs anywhere
pytest tests/test_ch11_cuda.py -k matmuls # GPU: naive and tiled vs torch
Then measure. Time naive_matmul, tiled_matmul and torch.matmul for 1024³ and 4096³ FP32 matrices (set torch.backends.cuda.matmul.allow_tf32 = False for a fair FP32 comparison) and convert the times to TFLOP/s.
Stretch exercises
- ★ Run
tiled_matmulwith tiles 4, 8, 16 and 32 on 256×256 matrices and plot loads against tile size. Then estimate the shared-memory capacity a 64×64 FP32 tile pair would need. Would it fit on your GPU? Where:experiments/ch12.py(create it), callingengine.tiling.tiled_matmulwith each tile size. - ★★ Add 2×2 register tiling to the CUDA kernel (each thread computes four outputs) and measure the speedup. Where:
tiled_matmul_kerneland its launch geometry inengine/kernels/cuda_ops.cu. - ★★ Write a GEMV kernel (
y = W x, W[N, K]BF16) where each warp computes one output row with a warp-shuffle reduction. Measure its achieved bandwidth against the datasheet. Where: add a GEMV kernel, host launcher and binding inengine/kernels/cuda_ops.cu. - ★★★ Use
wmmato write a tensor-core BF16 GEMM for multiples of 16, and compare it with cuBLAS. Where: add a WMMA kernel, host launcher and binding inengine/kernels/cuda_ops.cu.
Check your understanding
- Why does a tiled matmul need two barriers per phase, and what goes wrong without the second?
- Why can larger tiles be slower, even though they increase reuse?
- A warp reads column 5 of
float s[32][32]. How many bank conflicts occur, and how doess[32][33]fix it? - Why doesn’t tiling help a decode-time matrix-vector product?
Going deeper
- PMPP Chapter 5 (pp. 103-130): memory types, tiling, the tiled matmul kernel, boundary checks. §6.4-6.5: bank conflicts and coarsening. Chapter 15 (GEMM): register tiling, software pipelining and tensor-core considerations.
- GPU Mode L5 (Jeremy Howard, tiled matmul in CUDA and Numba from Python), L23 (tensor cores), L15/L36/L57 (CUTLASS and CuTe).
- Simon Boehm, How to Optimize a CUDA Matmul Kernel for cuBLAS-like Performance (blog, 2022): ten kernels from naive to 90% of cuBLAS, each measured.
13. Reductions, softmax and floating point
In this chapter
- Parallel reductions: trees, shared memory and warp shuffles, the pattern behind softmax, norms and losses.
- Floating-point formats from FP32 to FP4: what they can represent, and what they silently can't.
- Numerically stable softmax and log-sum-exp, and why accumulation precision matters.
- How to compare two implementations honestly: tolerances, aggregate errors, and the right reference.
You will build
stable_softmax in engine/numerics.py and, on a GPU, warp- and block-level reduction kernels, ending with an RMSNorm kernel.
Time: 5-7 hours. GPU: optional.
One answer from many values
A reduction combines many values into one: a sum, a maximum, a mean. Transformers are full of them. Softmax needs a maximum and a sum per row, LayerNorm and RMSNorm need a mean of squares per token, and the loss averages over tokens. Unlike elementwise operations, reductions need threads to combine their results, which raises two questions: how to do it in parallel, and how much numerical precision the combination loses.
Parallel reduction: a tree
A serial sum does $N - 1$ additions one after another. A tree adds independent pairs in parallel: $N/2$ additions, then $N/4$, and so on. It’s still $N-1$ additions in total, but only $\log_2 N$ sequential steps. For 8 values, that’s 3 steps instead of 7.
On a GPU, a block reduces its share of the data in two levels:
- Within a warp, threads exchange register values directly with shuffle instructions.
__shfl_down_sync(mask, v, offset)gives each thread the value held by the threadoffsetlanes above it. Five rounds (offsets 16, 8, 4, 2, 1) leave the warp’s total in lane 0, with no shared memory and no barriers. - Across warps, each warp’s lane 0 writes its partial sum to shared memory, the block synchronizes, and the first warp reduces those partials the same way.
// Sum within a warp: each step halves the number of live partial sums, using register
// exchanges instead of shared memory. After 5 steps lane 0 holds the warp's total.
__device__ float warp_sum(float v) {
for (int offset = 16; offset > 0; offset >>= 1) v += __shfl_down_sync(0xffffffff, v, offset);
return v;
}
// Block-wide sum: every warp reduces, warp leaders publish to shared memory, warp 0 finishes.
__device__ float block_sum(float v) {
__shared__ float partial[32];
int lane = threadIdx.x % 32, warp = threadIdx.x / 32;
v = warp_sum(v);
if (lane == 0) partial[warp] = v;
__syncthreads();
v = (threadIdx.x < blockDim.x / 32) ? partial[lane] : 0.f;
if (warp == 0) v = warp_sum(v);
__syncthreads(); // partial[] may be reused by the caller
return v; // valid in thread 0
}
// One block per row; each thread strides over the row (thread coarsening), then reduces.
__global__ void row_sum_kernel(const float* x, float* out, int width) {
const float* row = x + int64_t(blockIdx.x) * width;
float v = 0.f;
for (int i = threadIdx.x; i < width; i += blockDim.x) v += row[i];
v = block_sum(v);
if (threadIdx.x == 0) out[blockIdx.x] = v;
}
row_sum_kernel also uses thread coarsening: each thread first sums many elements serially in a register (for (i = tid; i < width; i += blockDim.x)), and only then does the block reduce 256 partials. Reading consecutive elements per iteration keeps the loads coalesced. Two details are worth remembering:
- Identity values. Threads with no data contribute the operation’s identity: 0 for a sum, but $-\infty$ for a max. Contributing 0 to a max of all-negative values would be wrong.
- Barriers across blocks don’t exist. If a reduction spans several blocks, write per-block partial results to global memory and reduce them in a second kernel, or use atomics.
RMSNorm: a reduction plus a broadcast
RMSNorm (Chapter 17) needs one reduction per token (the mean of squares), then the result must reach every thread so each can scale its elements. Thread 0 writes the statistic to shared memory, the block synchronizes, and everyone reads it:
__global__ void rmsnorm_kernel(const float* x, const float* w, float* y, int width, float eps) {
const float* row = x + int64_t(blockIdx.x) * width;
float* out = y + int64_t(blockIdx.x) * width;
__shared__ float inv_rms;
float sq = 0.f;
for (int i = threadIdx.x; i < width; i += blockDim.x) sq += row[i] * row[i];
sq = block_sum(sq);
if (threadIdx.x == 0) inv_rms = rsqrtf(sq / width + eps);
__syncthreads(); // broadcast the statistic to every thread
for (int i = threadIdx.x; i < width; i += blockDim.x) out[i] = row[i] * inv_rms * w[i];
}
One block per token row, reading the row twice: once for the statistic, once to normalize. A fused Triton version (Chapter 14) also adds the residual connection in the same pass.
Floating point: what your numbers can hold
A floating-point number is $\pm, 1.m \times 2^{e}$: a sign bit, an exponent field, and a mantissa (fraction) field. The exponent bits set the range, and the mantissa bits set the precision, meaning the gap between neighbouring representable numbers. Every format makes a different trade:
| format | sign / exponent / mantissa bits | largest value | gap just above 1.0 | typical use |
|---|---|---|---|---|
| FP32 | 1 / 8 / 23 | 3.4 × 10³⁸ | 1.2 × 10⁻⁷ | accumulation, optimizer state |
| BF16 | 1 / 8 / 7 | 3.4 × 10³⁸ | 0.0078 | weights and activations |
| FP16 | 1 / 5 / 10 | 65,504 | 0.00098 | older GPUs, some kernels |
| FP8 E4M3 | 1 / 4 / 3 | 448 | 0.125 | weights, activations (scaled) |
| FP8 E5M2 | 1 / 5 / 2 | 57,344 | 0.25 | gradients (scaled) |
| FP4 E2M1 | 1 / 2 / 1 | 6 | 0.5 | weights, with per-block scales |
BF16 keeps FP32’s 8-bit exponent, so it has the same range and never overflows where FP32 wouldn’t. It pays with precision: only about 2-3 significant decimal digits. That’s why it dominates training and inference. FP16 has more precision but overflows above 65,504, a real problem for activations. FP8 and FP4 are so coarse that they only work with scale factors stored alongside groups of values (Chapter 20).
Rounding makes order matter
In exact arithmetic, $(10^8 + 1) - 10^8 = 1$. In FP32, the gap between representable numbers near $10^8$ is 8, so $10^8 + 1$ rounds back to $10^8$ and the result is 0. Computing $(10^8 - 10^8) + 1$ instead gives 1. Floating-point addition is not associative: a tree reduction, a fused kernel and a serial loop can all give different answers, all legitimately.
Precision of the accumulator matters even more. Summing 0.1 ten thousand times should give 1,000:
BF16 accumulator, sequential: 32.0
FP32 accumulator: 1000.98 (0.1 itself isn't exact in BF16)
The BF16 sum is stuck at 32, where the gap between representable numbers is 0.25: adding 0.1 rounds back down every time. This is why every serious kernel stores in low precision but accumulates in FP32: matmuls on tensor cores, reductions in norms, the softmax denominator. Your attention and norm code converts to FP32 (.float()) for its reductions for exactly this reason.
Stable softmax
$e^{1000}$ overflows every format, so a naive softmax of [1000, 999, 998] returns [nan, nan, nan]. Since softmax is unchanged by subtracting a constant from every input,
$$ \operatorname{softmax}(x)_i = \frac{e^{x_i - m}}{\sum_j e^{x_j - m}}, \qquad m = \max_j x_j , $$
subtracting the row maximum makes the largest exponent $e^0 = 1$ and everything else smaller. The result is $[0.6652, 0.2447, 0.0900]$, exactly as if the inputs had been $[2, 1, 0]$.
def stable_softmax(scores, dim=-1):
"""Softmax that subtracts the row maximum first. (Your engine: Chapter 13)
Computes in FP32 and returns the input dtype. Each row needs at least one finite score.
"""
x = scores.float()
shifted = x - x.amax(dim=dim, keepdim=True)
numerators = shifted.exp()
return (numerators / numerators.sum(dim=dim, keepdim=True)).to(scores.dtype)
inline void softmax(float* x, size_t n) {
float m = *std::max_element(x, x + n), sum = 0;
for (size_t i = 0; i < n; ++i) sum += (x[i] = std::exp(x[i] - m));
for (size_t i = 0; i < n; ++i) x[i] /= sum;
}
#![allow(unused)]
fn main() {
/// Stable softmax in place: subtract the maximum so the largest exponent is exp(0) = 1.
pub fn softmax(x: &mut [f32]) {
let max = x.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
let mut sum = 0.0;
for v in x.iter_mut() {
*v = (*v - max).exp();
sum += *v;
}
for v in x.iter_mut() {
*v /= sum;
}
}
}
The same trick gives a stable log-sum-exp, $\log\sum_j e^{x_j} = m + \log\sum_j e^{x_j - m}$, which is how cross-entropy is computed without ever forming probabilities (Chapter 7). Keep the softmax’s edge case in mind: if every entry of a row is $-\infty$ (a fully masked attention row), the max is $-\infty$ and the result is NaN. Your mask logic must guarantee at least one visible key, or define what an empty row means.
In Chapter 15, you’ll extend this to compute softmax in tiles, carrying the running maximum forward. That’s the online softmax at the heart of FlashAttention.
Comparing implementations honestly
Because different correct implementations round differently, “equal” means “close enough”, and you must decide how close before you look at the results. The standard test combines absolute and relative tolerances:
$$ |a - r| \le \text{atol} + \text{rtol}\cdot|r| . $$
The absolute term handles values near zero; the relative term handles scale. Good error reports also include aggregate measures, because a maximum alone can hide a systematic drift:
def compare(actual, expected):
"""An error report, not a pass/fail verdict. Choose tolerances per dtype and operation."""
difference = (actual.float() - expected.float()).abs()
reference = expected.float()
return {
"max_abs": difference.max().item() if difference.numel() else 0.0,
"mean_abs": difference.mean().item() if difference.numel() else 0.0,
"relative_l2": (
torch.linalg.vector_norm(difference)
/ torch.linalg.vector_norm(reference).clamp_min(1e-12)
).item(),
}
Which tolerances? They depend on dtype, operation and magnitude. A BF16 result can’t be closer than about 0.4% relative to anything, so atol=1e-5 on BF16 logits is meaningless. A useful discipline is to make three comparisons:
- Trusted FP32 vs your FP32: tests your implementation. Expect ~1e-5 or better.
- Trusted BF16 vs trusted FP32: measures what the precision itself costs.
- Your BF16 vs trusted BF16: should be about as close as (2), not much worse.
If (3) is much worse than (2), you have a bug or a bad accumulation choice. And whatever the tolerance, rounding never excuses a semantic error: a wrong mask or transpose produces differences far larger than any rounding.
Warning
On NVIDIA GPUs, PyTorch may run “FP32” matmuls with TF32 tensor cores (10-bit mantissa) when
torch.backends.cuda.matmul.allow_tf32is on, giving about 1e-3 relative error. Set the flag explicitly whenever you compare FP32 results.
Note
Some GPU reductions are nondeterministic: using atomic additions or split-K matmuls, the order of additions changes run to run, and so do the last bits of the result. GPU Mode L9 (
nondeterminism.py) demonstrates it. Serving engines that promise bitwise-reproducible outputs have to avoid these kernels.
Build it
Engine milestone 13: stable numerics. Implement stable_softmax in engine/numerics.py (logsumexp, compare and float_bits are provided). On a GPU, implement warp_sum, block_sum, row_sum_kernel and rmsnorm_kernel in engine/kernels/cuda_ops.cu.
pytest tests/test_ch13_numerics.py
pytest tests/test_ch11_cuda.py -k reductions # GPU
Stretch exercises
- ★ Reproduce the BF16 summation failure, then implement Kahan summation in BF16 (carry a compensation term) and see how close it gets to 1,000. Where:
experiments/ch13.py(create it); add a Kahan helper toengine/numerics.py. - ★★ Write a block reduction for the maximum and check it on rows of all-negative numbers. What breaks if invalid lanes contribute 0? Where: add max-reduction helpers and a row-max kernel, launcher and binding in
engine/kernels/cuda_ops.cu. - ★★ Run Qwen3 or your GPT in FP32 and in BF16 on the same input, and report max-abs and relative-L2 error of the logits, and whether greedy tokens agree over 50 steps. Where:
experiments/ch13.py(create it), importingengine.numerics.compare. - ★★★ Write a single-pass softmax kernel in CUDA: one block per row, a max reduction, a sum reduction, then normalize, with the row held in registers. Compare its bandwidth with
torch.softmax. Where: add a row-softmax kernel, launcher and binding inengine/kernels/cuda_ops.cu.
Check your understanding
- Why does subtracting the row maximum leave softmax unchanged?
- Why does BF16 have FP32’s range but far less precision?
- Why should a BF16 kernel accumulate sums in FP32?
- Why report aggregate error in addition to the maximum error?
- What identity value should inactive threads contribute to a max reduction?
Going deeper
- PMPP Chapter 10 (reductions: trees, divergence, warp shuffles, coarsening, pp. 225-248) and Appendix A (floating-point considerations).
- GPU Mode L9 (Mark Saroufim, reductions, with runnable kernels and the nondeterminism demo) and L84 (numerics and AI).
- David Goldberg, What Every Computer Scientist Should Know About Floating-Point Arithmetic (1991); Micikevicius et al., FP8 Formats for Deep Learning (2022).
14. Triton and kernel fusion
In this chapter
- Triton's programming model: you write the work of one block as vector code, and the compiler handles threads, shared memory and tensor cores.
- Kernel fusion: why combining operations into one kernel is the most effective optimization for memory-bound code.
- Fused softmax, fused residual-add + RMSNorm, and a tiled matmul with tensor cores, all in Python.
- How
torch.compilegenerates Triton for you, and when to write your own.
You will build
engine/kernels/triton_basics.py and triton_matmul.py: four kernels, tested on GPU, or on CPU with Triton's interpreter.
Time: 5-7 hours. GPU: optional (the interpreter runs everything, slowly).
Why another GPU language
CUDA C++ gives you full control, and you pay for it in detail: thread indices, shared-memory layouts, barriers, bank conflicts, vector loads, tensor-core fragment layouts. Triton (Tillet et al., 2019, now part of PyTorch) raises the level of abstraction by one step. You write the program for one block of data, using operations on whole tiles (tl.load a tile, tl.dot two tiles, tl.sum along an axis), and the compiler decides how threads map onto them, stages data through shared memory, coalesces loads and uses tensor cores. You still choose the tiling, which is the decision that matters most, but the bookkeeping disappears.
The trade-off: Triton reaches 80-100% of hand-written CUDA for most kernels an inference engine needs, in a fraction of the code. PyTorch’s compiler emits Triton, and most new kernels in vLLM and SGLang start as Triton. When it falls short (warp-specialized FlashAttention 3, the very latest hardware features), engines drop to CUDA, CUTLASS or newer tile languages.
The programming model
A Triton kernel is a Python function decorated with @triton.jit. It’s launched over a grid of program instances, and each instance asks which one it is with tl.program_id(axis):
@triton.jit
def add_kernel(x_ptr, y_ptr, out_ptr, n, BLOCK: tl.constexpr):
"""(Your engine: Chapter 14)"""
pid = tl.program_id(0) # which block of BLOCK elements is mine
offsets = pid * BLOCK + tl.arange(0, BLOCK) # a vector of BLOCK indices
mask = offsets < n # the last block may run past the end
x = tl.load(x_ptr + offsets, mask=mask)
y = tl.load(y_ptr + offsets, mask=mask)
tl.store(out_ptr + offsets, x + y, mask=mask)
def vector_add(x, y, block=1024):
check_device(x, y)
if x.shape != y.shape or not x.is_contiguous() or not y.is_contiguous():
raise ValueError("vector_add needs equal-shape contiguous tensors")
out = torch.empty_like(x)
n = x.numel()
add_kernel[(triton.cdiv(n, block),)](x, y, out, n, BLOCK=block)
return out
Read it as “program pid handles elements [pid*BLOCK, (pid+1)*BLOCK)”:
tl.arange(0, BLOCK)is a vector ofBLOCKindices, not a loop.BLOCKmust be a compile-time constant (tl.constexpr) and a power of two.- Pointers support arithmetic:
x_ptr + offsetsis a vector of addresses. maskdisables lanes past the end. Every load and store of a partial block needs it, and masked loads take another=fill value.- The launch
add_kernel[grid](...)takes a grid tuple, computed withtriton.cdiv.
The first call compiles the kernel for the given constants and dtypes (seconds); later calls reuse the compiled binary. Time steady state, not the first call.
Tip
No GPU? Run Triton on the CPU. Set
TRITON_INTERPRET=1before importing triton and pass CPU tensors. The interpreter executes the kernel with NumPy, about 100-1,000x slower but numerically faithful, and you can put
Fusion: the optimization that matters most for memory-bound code
Consider softmax over the rows of a [4096, 32000] FP32 matrix (512 MB), written as separate PyTorch operations:
m = x.max(dim=-1, keepdim=True) # read x (512 MB) write m (small)
e = torch.exp(x - m) # read x, write x - m, read it back, write e (4 passes over 512 MB)
s = e.sum(dim=-1, keepdim=True) # read e
y = e / s # read e, write y
That’s around 8 full passes over memory for an operation whose minimum is 2 (read x once, write y once). Since softmax is memory-bound (Chapter 10), runtime is proportional to passes, so a fused kernel that keeps each row in registers is about 4x faster:
@triton.jit
def softmax_kernel(x_ptr, out_ptr, row_stride, width, BLOCK: tl.constexpr):
"""(Your engine: Chapter 14)"""
row = tl.program_id(0)
cols = tl.arange(0, BLOCK)
mask = cols < width
x = tl.load(x_ptr + row * row_stride + cols, mask=mask, other=-float("inf")).to(tl.float32)
x = x - tl.max(x, axis=0) # stable: largest exponent is exp(0) = 1
num = tl.exp(x)
tl.store(out_ptr + row * row_stride + cols, num / tl.sum(num, axis=0), mask=mask)
def softmax(x):
"""Row softmax over the last dim of a contiguous 2-D tensor, one program per row.
Reads x once and writes once: the unfused version (max, sub, exp, sum, div) makes ~5 passes."""
check_device(x)
if x.ndim != 2 or not x.is_contiguous():
raise ValueError("softmax expects a contiguous matrix")
out = torch.empty_like(x)
block = triton.next_power_of_2(x.shape[1])
if block > 65536:
raise ValueError("Rows wider than 65536 need a multi-pass (online) kernel")
softmax_kernel[(x.shape[0],)](x, out, x.stride(0), x.shape[1], BLOCK=block)
return out
One program per row: load the row (padded with $-\infty$, the identity for max), subtract the max, exponentiate, sum, divide, and store. The intermediates never touch memory. Fusion is how almost every non-GEMM operation in a production engine is implemented: normalization, activations, rotary embeddings, residual additions, sampling.
Fused residual add + RMSNorm
Every transformer layer does x = x + branch_output followed by h = rmsnorm(x). Unfused, that’s a read of both inputs, a write of the sum, a read of the sum for the statistic, and another read to normalize, then a write. Fused, it’s two reads and two writes, all in one kernel:
@triton.jit
def add_rmsnorm_kernel(x_ptr, res_ptr, w_ptr, out_ptr, res_out_ptr, width, eps,
HAS_RESIDUAL: tl.constexpr, BLOCK: tl.constexpr):
"""(Your engine: Chapter 14)"""
row = tl.program_id(0)
cols = tl.arange(0, BLOCK)
mask = cols < width
x = tl.load(x_ptr + row * width + cols, mask=mask, other=0.0).to(tl.float32)
if HAS_RESIDUAL: # fuse "x = x + residual" into the same pass
x += tl.load(res_ptr + row * width + cols, mask=mask, other=0.0).to(tl.float32)
tl.store(res_out_ptr + row * width + cols, x, mask=mask)
inv_rms = 1.0 / tl.sqrt(tl.sum(x * x, axis=0) / width + eps)
w = tl.load(w_ptr + cols, mask=mask, other=0.0).to(tl.float32)
tl.store(out_ptr + row * width + cols, x * inv_rms * w, mask=mask)
def rmsnorm(x, weight, eps=1e-6, residual=None):
"""RMSNorm(x [+ residual]) * weight over the last dim. With a residual, also returns x + residual,
so a transformer layer's 'add then normalize' costs one memory pass instead of three."""
check_device(x, weight)
shape = x.shape
x2 = x.reshape(-1, shape[-1]).contiguous()
out = torch.empty_like(x2)
res_out = torch.empty_like(x2) if residual is not None else out
res = residual.reshape(-1, shape[-1]).contiguous() if residual is not None else x2
add_rmsnorm_kernel[(x2.shape[0],)](x2, res, weight, out, res_out, shape[-1], eps,
HAS_RESIDUAL=residual is not None,
BLOCK=triton.next_power_of_2(shape[-1]))
if residual is not None:
return out.view(shape), res_out.view(shape)
return out.view(shape)
HAS_RESIDUAL is a constexpr, so Triton compiles two specialized versions and the if costs nothing at runtime. You’ll swap this kernel into the engine in Chapter 19.
Matmul in Triton
The tiled algorithm from Chapter 12, at tile level. Each program computes one [BLOCK_M, BLOCK_N] output tile by looping over K in BLOCK_K steps, with tl.dot doing the tile product, on tensor cores when the dtypes allow:
@triton.jit
def matmul_kernel(a_ptr, b_ptr, c_ptr, M, N, K,
stride_am, stride_ak, stride_bk, stride_bn, stride_cm, stride_cn,
BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr,
GROUP_M: tl.constexpr):
"""(Your engine: Chapter 14)"""
# Map the 1-D program id to an output tile. Walking GROUP_M tile-rows at a time keeps the
# B tiles those rows need hot in L2 cache ("grouped ordering", Triton tutorial 03).
pid = tl.program_id(0)
tiles_m, tiles_n = tl.cdiv(M, BLOCK_M), tl.cdiv(N, BLOCK_N)
group = pid // (GROUP_M * tiles_n)
first_m = group * GROUP_M
group_size = min(tiles_m - first_m, GROUP_M)
pid_m = first_m + (pid % group_size)
pid_n = (pid % (GROUP_M * tiles_n)) // group_size
rm = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
rn = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
rk = tl.arange(0, BLOCK_K)
acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32) # accumulate in FP32 whatever the inputs
for k0 in range(0, K, BLOCK_K):
a = tl.load(a_ptr + rm[:, None] * stride_am + (k0 + rk)[None, :] * stride_ak,
mask=(rm[:, None] < M) & ((k0 + rk)[None, :] < K), other=0.0)
b = tl.load(b_ptr + (k0 + rk)[:, None] * stride_bk + rn[None, :] * stride_bn,
mask=((k0 + rk)[:, None] < K) & (rn[None, :] < N), other=0.0)
acc += tl.dot(a, b) # tensor cores when dtypes allow
c = acc.to(c_ptr.dtype.element_ty)
tl.store(c_ptr + rm[:, None] * stride_cm + rn[None, :] * stride_cn, c,
mask=(rm[:, None] < M) & (rn[None, :] < N))
def matmul(a, b, block_m=64, block_n=64, block_k=32, group_m=8):
check_device(a, b)
if a.ndim != 2 or b.ndim != 2 or a.shape[1] != b.shape[0]:
raise ValueError("matmul expects [M, K] @ [K, N]")
m, k = a.shape
n = b.shape[1]
c = torch.empty((m, n), device=a.device, dtype=a.dtype)
grid = (triton.cdiv(m, block_m) * triton.cdiv(n, block_n),)
matmul_kernel[grid](a, b, c, m, n, k, a.stride(0), a.stride(1), b.stride(0), b.stride(1),
c.stride(0), c.stride(1), BLOCK_M=block_m, BLOCK_N=block_n,
BLOCK_K=block_k, GROUP_M=group_m)
return c
The grouped ordering at the top deserves a word. The GPU runs programs roughly in pid order. If consecutive programs walked along one row of output tiles, each would need a different column strip of B, and B’s tiles would be evicted from L2 before neighbouring rows reused them. Grouping GROUP_M tile-rows together makes nearby programs share B tiles while they’re still cached. This often gives a 10-20% speedup for free.
Autotuning
The best BLOCK_M, BLOCK_N, BLOCK_K, number of warps and number of pipeline stages depend on the GPU and the matrix shape. Triton can try several and remember the fastest per shape:
@triton.autotune(configs=[
triton.Config({"BLOCK_M": 128, "BLOCK_N": 128, "BLOCK_K": 32}, num_warps=8, num_stages=3),
triton.Config({"BLOCK_M": 64, "BLOCK_N": 128, "BLOCK_K": 64}, num_warps=4, num_stages=4),
triton.Config({"BLOCK_M": 32, "BLOCK_N": 64, "BLOCK_K": 64}, num_warps=4, num_stages=4),
], key=["M", "N", "K"])
@triton.jit
def matmul_kernel(...): ...
num_stages controls software pipelining: loading the next K-tile while computing on the current one, which hides memory latency behind arithmetic (PMPP §15.8).
Let the compiler write it: torch.compile
PyTorch’s compiler traces your model, fuses chains of elementwise and reduction operations, and generates Triton kernels. Wrapping a function in torch.compile often gives much of the benefit of hand fusion with no kernel code:
fast_norm = torch.compile(lambda x, r, w: rmsnorm_reference(x + r, w))
# TORCH_LOGS=output_code python script.py prints the Triton it generated
Reading the generated code is one of the best ways to learn Triton idioms. When does writing your own still pay? When the fusion crosses boundaries the compiler won’t cross (attention, quantized matmuls, paged memory access), when you need a specific algorithm (online softmax), or when you need control over memory layout.
Integrating kernels into an engine
Keep a readable PyTorch reference next to every custom kernel, and select between them in one place:
def rmsnorm(x, w, eps, backend="torch"):
if backend == "triton" and x.is_cuda and x.shape[-1] <= 65536:
return triton_basics.rmsnorm(x, w, eps)
return reference_rmsnorm(x, w, eps) # always-correct fallback
Then test at two levels. First the kernel against the reference, on awkward shapes (odd widths, one row, the largest supported width). Then the whole model’s logits with the kernel swapped in. A correct kernel can still be called with the wrong axis or a non-contiguous input. Record which backend produced every benchmark number.
| kernel | shape | reference ms | custom ms | max error | whole-engine effect |
|---|---|---|---|---|---|
| softmax | 32 × 32000 | ||||
| add+rmsnorm | 1 × 1024 (decode) | ||||
| add+rmsnorm | 4096 × 1024 (prefill) |
Fill a table like this with measurements, not expectations. Decode-sized inputs often benefit less than you’d hope: launch overhead, not memory traffic, dominates tiny kernels (Chapter 19).
Build it
Engine milestone 14: Triton kernels. Write the four @triton.jit kernels in engine/kernels/triton_basics.py (add_kernel, softmax_kernel, add_rmsnorm_kernel) and engine/kernels/triton_matmul.py (matmul_kernel). The Python wrappers that launch them are provided.
pytest tests/test_ch14_triton.py # on CPU this uses TRITON_INTERPRET=1 automatically
python run.py kernels --impl engine
On a GPU, benchmark your fused softmax and add+RMSNorm against the unfused PyTorch versions with measure_cuda, and report the achieved GB/s.
Stretch exercises
- ★ Add a
GELUepilogue to the matmul kernel (apply GELU toaccbefore storing) and test it againstF.gelu(a @ b). That’s a fused linear + activation. Where:matmul_kernelinengine/kernels/triton_matmul.py. - ★★ Add
@triton.autotuneto the matmul and plot TFLOP/s againsttorch.matmulfor square BF16 sizes 512-8,192. Where:matmul_kerneland its launcher inengine/kernels/triton_matmul.py. - ★★ Write a Triton RoPE kernel that rotates Q and K in place for given positions (Chapter 17 explains RoPE), and test it against the reference
apply_rope. Where: newengine/kernels/triton_rope.py; compare with the providedizh.qwen3.apply_rope. - ★★★ Write a softmax for rows too long for one block: two passes, or an online version that carries a running max and sum (preview of Chapter 15). Where:
softmax_kernelandsoftmaxinengine/kernels/triton_basics.py.
Check your understanding
- What does one Triton program instance correspond to, and how does it differ from one CUDA thread?
- Why does fusing a softmax make it several times faster, even though it does the same arithmetic?
- Why must masked loads in a max-reduction use
other=-inf? - Why should you keep a PyTorch reference path alongside every custom kernel?
Going deeper
- GPU Mode L14 (Umer Adil, A Practitioner’s Guide to Triton: notebook with debugging tips and the interpreter), L18 (fused kernels), L28 (Liger kernels: fused RMSNorm, RoPE and cross-entropy in production), L29 (Triton internals).
- The official Triton tutorials: vector add, fused softmax, matmul with grouped ordering and autotuning, layer norm.
- Tillet, Kung and Cox, Triton: an intermediate language and compiler for tiled neural network computations (2019).
- PMPP §15.8 (software pipelining) for what
num_stagesdoes underneath.
15. FlashAttention
In this chapter
- Why standard attention is limited by memory traffic, not arithmetic.
- The online-softmax recurrence, derived step by step: how to normalize scores you haven't all seen yet.
- FlashAttention as tiling plus the online recurrence, in PyTorch, C++, Rust and Triton.
- Flash-decoding: splitting a long KV cache across many blocks and merging the partial results.
You will build
online_attention in engine/attention.py and a FlashAttention forward kernel in engine/kernels/triton_flash.py, with causal masking, query offsets and grouped-query heads.
Time: 6-8 hours. GPU: optional.
The score matrix is the problem
Standard attention (Chapter 5) computes $S = QK^\top/\sqrt d$, then $P = \operatorname{softmax}(S)$, then $O = PV$. For a sequence of $T$ tokens, $S$ and $P$ are $T \times T$ per head. At $T = 8{,}192$, each is 67 million entries: 256 MiB per head in FP32. With 32 heads and a batch of 8, the intermediate scores alone would need 64 GiB.
Even when the memory fits, the traffic kills performance. Unfused, $S$ is written to HBM, read back for the softmax, $P$ is written, then read again for $PV$: four passes over $T^2$ numbers. All that just to produce an output of size $T \times d$, with $d$ typically 64-256. By Chapter 10’s roofline, unfused attention is memory-bound at long sequence lengths even though its arithmetic is large.
FlashAttention (Dao et al., 2022) computes exactly the same output without ever storing $S$ or $P$ in global memory. It tiles $Q$, $K$ and $V$ into on-chip memory and uses a clever recurrence for the softmax. The output matches standard attention up to rounding; it isn’t an approximation. Sparse and linear attention (Chapters 28-29) are approximations, or different models altogether. This isn’t.
The obstacle: softmax needs the whole row
Tiling a matmul works because a dot product is a sum: process $K$ in chunks and add up partial sums. Softmax is harder. Each weight is $e^{s_j}/\sum_k e^{s_k}$, and the denominator needs every score in the row. The stable version also subtracts the row’s maximum, which you don’t know until you’ve seen every score. It looks like you need the whole row first.
Online softmax: normalize as you go
The trick (Milakov and Gimelshein, 2018) is to keep a running answer that’s correct for the scores seen so far, and to fix it up when new scores arrive. For one query row, keep three running values:
- $m$: the largest score seen so far,
- $\ell$: the sum of $e^{s_j - m}$ over scores seen so far,
- $a$: the sum of $e^{s_j - m}, v_j$ over scores seen so far (a vector of length $d$, not yet normalized).
When a new tile of scores arrives with a larger maximum $m_{\text{new}}$, every old term $e^{s_j - m_{\text{old}}}$ must become $e^{s_j - m_{\text{new}}}$. Since
$$ e^{s_j - m_{\text{new}}} = e^{s_j - m_{\text{old}}} \cdot e^{m_{\text{old}} - m_{\text{new}}}, $$
every old term is fixed by multiplying by the same factor $\alpha = e^{m_{\text{old}} - m_{\text{new}}}$. So the update for a tile with scores $s$ and values $V_{\text{tile}}$ is:
$$ \begin{aligned} m_{\text{new}} &= \max!\big(m_{\text{old}},, \max(s)\big) \ p &= e^{,s - m_{\text{new}}} \ \ell_{\text{new}} &= \alpha,\ell_{\text{old}} + \textstyle\sum p \ a_{\text{new}} &= \alpha, a_{\text{old}} + p, V_{\text{tile}} \end{aligned} \qquad\text{and at the end}\qquad O = a / \ell . $$
Both $\ell$ and $a$ must be rescaled. Forgetting one mixes terms measured against different references, a classic bug.
Worked example
One query, scores $[0, 1, 2, -1]$ against scalar values $[2, 4, 1, 3]$, processed in two tiles of two.
Tile 1, scores $[0, 1]$: $m = 1$, $p = [e^{-1}, e^0] = [0.3679, 1]$, $\ell = 1.3679$, $a = 0.3679\cdot 2 + 1\cdot 4 = 4.7358$.
Tile 2, scores $[2, -1]$: the maximum rises to $m = 2$, so $\alpha = e^{1-2} = 0.3679$. $p = [e^0, e^{-3}] = [1, 0.0498]$.
$$ \ell = 0.3679 \times 1.3679 + 1 + 0.0498 = 1.5530, \qquad a = 0.3679 \times 4.7358 + 1\cdot 1 + 0.0498 \cdot 3 = 2.8915 . $$
$O = 2.8915 / 1.5530 = 1.8619$, exactly what softmax([0,1,2,-1]) @ [2,4,1,3] gives. Your milestone tests check this example.
The algorithm in plain PyTorch
Here’s the recurrence over key tiles, for all queries and heads at once, with causal masking by position:
def online_attention(q, k, v, query_positions=None, key_positions=None, tile=16):
"""The same result as causal_attention, computed one key tile at a time without ever
holding the full [T, S] score matrix. This is FlashAttention's recurrence in plain
PyTorch. (Your engine: Chapter 15)
Running state per query row: m (max score so far), l (sum of exp(score - m)),
acc (sum of exp(score - m) * value). A new tile with a larger maximum rescales the old
l and acc by exp(m_old - m_new) so every term shares one exponent reference.
"""
if tile < 1:
raise ValueError("tile must be positive")
batch, q_heads, t, d = q.shape
kv_heads, s = k.shape[1], k.shape[2]
if q_heads % kv_heads:
raise ValueError("Query heads must be a multiple of KV heads")
group = q_heads // kv_heads
k = k.repeat_interleave(group, dim=1) if group > 1 else k
v = v.repeat_interleave(group, dim=1) if group > 1 else v
qp = _positions(query_positions, batch, t, s - t, q.device)
kp = _positions(key_positions, batch, s, 0, q.device)
q32 = q.float() / math.sqrt(d)
m = torch.full((batch, q_heads, t, 1), float("-inf"), device=q.device)
l = torch.zeros((batch, q_heads, t, 1), device=q.device)
acc = torch.zeros((batch, q_heads, t, v.shape[-1]), device=q.device)
for start in range(0, s, tile):
stop = min(start + tile, s)
scores = q32 @ k[:, :, start:stop].float().transpose(-2, -1)
visible = kp[:, None, None, start:stop] <= qp[:, None, :, None]
scores = scores.masked_fill(~visible, float("-inf"))
m_new = torch.maximum(m, scores.amax(-1, keepdim=True))
# A row that has seen no visible key yet keeps m = -inf; use 0 as a safe reference.
reference = torch.where(torch.isfinite(m_new), m_new, torch.zeros_like(m_new))
alpha = torch.exp(m - reference) # exp(-inf) = 0 for the first visible tile
p = torch.exp(scores - reference)
l = alpha * l + p.sum(-1, keepdim=True)
acc = alpha * acc + p @ v[:, :, start:stop].float()
m = m_new
if bool((l == 0).any()):
raise ValueError("A query row has no visible key")
return (acc / l).to(q.dtype)
Two edge cases need care. A query row may see no visible key in an early tile (with causal masking, the first tiles of a long query block). Then $m = -\infty$, and $e^{-\infty - (-\infty)}$ is NaN. The code uses a safe reference of 0 for rows that haven’t seen any score yet. And if a row never sees a key at all, the result is undefined, so the function raises.
The same recurrence for one query and one head, in C++ and Rust:
inline std::vector<float> online_attention(const std::vector<float>& q, const std::vector<float>& k,
const std::vector<float>& v, size_t S, size_t d, size_t tile) {
float m = -INFINITY, l = 0, scale = 1.f / std::sqrt(float(d));
std::vector<float> acc(d, 0.f);
for (size_t start = 0; start < S; start += tile) {
size_t end = std::min(S, start + tile);
std::vector<float> s(end - start);
float m_new = m;
for (size_t j = start; j < end; ++j) {
s[j - start] = std::inner_product(q.begin(), q.end(), k.begin() + j * d, 0.f) * scale;
m_new = std::max(m_new, s[j - start]);
}
float alpha = std::exp(m - m_new); // rescale what was accumulated under the old max
l *= alpha;
for (float& a : acc) a *= alpha;
for (size_t j = start; j < end; ++j) {
float p = std::exp(s[j - start] - m_new);
l += p;
for (size_t e = 0; e < d; ++e) acc[e] += p * v[j * d + e];
}
m = m_new;
}
for (float& a : acc) a /= l;
return acc;
}
#![allow(unused)]
fn main() {
/// One query, one head, keys processed in tiles of `tile` with a running max m, running
/// denominator l and an unnormalized accumulator acc (Chapter 15). Equals softmax(qK^T)V.
pub fn online_attention(q: &[f32], k: &[f32], v: &[f32], s: usize, d: usize, tile: usize) -> Vec<f32> {
let scale = 1.0 / (d as f32).sqrt();
let (mut m, mut l) = (f32::NEG_INFINITY, 0.0f32);
let mut acc = vec![0.0f32; d];
for start in (0..s).step_by(tile) {
let end = (start + tile).min(s);
let scores: Vec<f32> = (start..end)
.map(|j| q.iter().zip(&k[j * d..][..d]).map(|(a, b)| a * b).sum::<f32>() * scale)
.collect();
let m_new = scores.iter().cloned().fold(m, f32::max);
let alpha = (m - m_new).exp(); // rescales everything accumulated so far
l *= alpha;
acc.iter_mut().for_each(|a| *a *= alpha);
for (j, sc) in (start..end).zip(&scores) {
let p = (sc - m_new).exp();
l += p;
for (a, vv) in acc.iter_mut().zip(&v[j * d..][..d]) {
*a += p * vv;
}
}
m = m_new;
}
acc.iter().map(|a| a / l).collect()
}
}
This PyTorch version saves memory but not time: each tile is still several separate kernels launched from a Python loop. The speed comes from doing a whole tile’s work inside one kernel, with the tile in on-chip memory.
The kernel
A FlashAttention forward kernel assigns each program a block of BLOCK_M queries for one (batch, head). The program loads its queries once, then streams BLOCK_N keys and values at a time through on-chip memory, applying the recurrence:
@triton.jit
def flash_fwd_kernel(q_ptr, k_ptr, v_ptr, o_ptr,
sqb, sqh, sqt, skb, skh, sks, svb, svh, svs, sob, soh, sot,
T, S, offset, heads, group, scale,
D: tl.constexpr, BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr):
"""(Your engine: Chapter 15)"""
pid_m = tl.program_id(0) # which tile of queries
bh = tl.program_id(1) # which (batch, query head)
b = bh // heads
h = bh % heads
kvh = h // group # GQA: several query heads read one KV head
rm = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
rd = tl.arange(0, D)
q = tl.load(q_ptr + b * sqb + h * sqh + rm[:, None] * sqt + rd[None, :], mask=rm[:, None] < T, other=0.0)
q_pos = rm + offset # absolute position of each query row
m = tl.full((BLOCK_M,), -float("inf"), tl.float32)
l = tl.zeros((BLOCK_M,), tl.float32)
acc = tl.zeros((BLOCK_M, D), tl.float32)
# Causality: no key beyond the last query's position is ever needed.
last_key = tl.minimum(S, offset + (pid_m + 1) * BLOCK_M)
for start in range(0, last_key, BLOCK_N):
rn = start + tl.arange(0, BLOCK_N)
k = tl.load(k_ptr + b * skb + kvh * skh + rn[:, None] * sks + rd[None, :], mask=rn[:, None] < S, other=0.0)
v = tl.load(v_ptr + b * svb + kvh * svh + rn[:, None] * svs + rd[None, :], mask=rn[:, None] < S, other=0.0)
# input_precision="ieee": FP32 inputs get true FP32 math, not TF32 (BF16/FP16 are unaffected).
s = tl.dot(q, tl.trans(k), input_precision="ieee").to(tl.float32) * scale # [BLOCK_M, BLOCK_N]
visible = (rn[None, :] <= q_pos[:, None]) & (rn[None, :] < S)
s = tl.where(visible, s, -float("inf"))
m_new = tl.maximum(m, tl.max(s, axis=1))
m_safe = tl.where(m_new == -float("inf"), 0.0, m_new) # rows with nothing visible yet
alpha = tl.exp(m - m_safe) # rescale the old partial sums
p = tl.exp(s - m_safe[:, None])
l = alpha * l + tl.sum(p, axis=1)
acc = acc * alpha[:, None] + tl.dot(p.to(v.dtype), v, input_precision="ieee").to(tl.float32)
m = m_new
out = acc / l[:, None]
tl.store(o_ptr + b * sob + h * soh + rm[:, None] * sot + rd[None, :], out.to(o_ptr.dtype.element_ty),
mask=rm[:, None] < T)
def flash_attention(q, k, v, query_offset=None, block_m=32, block_n=32):
"""q [B, Hq, T, D], k/v [B, Hkv, S, D]; query i sits at absolute position offset + i
(default offset = S - T, i.e. the queries are the newest tokens). D must be a power of two >= 16."""
check_device(q, k, v)
batch, heads, t, d = q.shape
kv_heads, s = k.shape[1], k.shape[2]
if heads % kv_heads or d & (d - 1) or d < 16:
raise ValueError("Need Hq % Hkv == 0 and a power-of-two head dim >= 16")
offset = s - t if query_offset is None else query_offset
q, k, v = (x.contiguous() for x in (q, k, v))
o = torch.empty_like(q)
grid = (triton.cdiv(t, block_m), batch * heads)
flash_fwd_kernel[grid](q, k, v, o, *q.stride()[:3], *k.stride()[:3], *v.stride()[:3], *o.stride()[:3],
t, s, offset, heads, heads // kv_heads, 1.0 / math.sqrt(d),
D=d, BLOCK_M=block_m, BLOCK_N=block_n)
return o
Details worth noticing:
- Causal block skipping. No key past the last query’s position is needed, so the loop stops at
offset + (pid_m + 1) * BLOCK_M. That halves the work for causal attention. - Query offset. Query row
isits at absolute positionoffset + i. Withoffset = S - T, one kernel serves full prefill (T = S), chunked prefill (T < S) and decode (T = 1), exactly like your Chapter 5 position rule. - Grouped-query attention for free. Query head
hreads KV headh // group. No repeated K/V tensors are materialized. - FP32 statistics, low-precision inputs.
m,landaccstay in FP32 registers.pis cast to the value dtype for the secondtl.dot, so both matmuls run on tensor cores. input_precision="ieee". On NVIDIA GPUstl.dotruns FP32 inputs on TF32 tensor cores by default (Chapter 13’s 10-bit mantissa, about 1e-3 error), so an FP32 test againstcausal_attentionfails at a 1e-4 tolerance even though the interpreter passes."ieee"asks for true FP32 math. It changes only FP32 inputs: BF16 and FP16 still run on tensor cores.
How much memory traffic does this save? Each query block streams the visible K and V once, so the kernel reads about $2T^2 d/\text{BLOCK_M}$ elements of K and V. With $d = 128$ and BLOCK_M = 128, that’s about $2T^2$, against roughly $4T^2$ score elements moved by the unfused version (plus its own K/V reads). The bigger wins are elsewhere: nothing of size $T^2$ is ever allocated or written, and the time moves out of memory-bound elementwise kernels into two compute-bound matmuls per tile.
Note
Production FlashAttention adds much more. FlashAttention-2 reorders loops to parallelize over the sequence and cut non-matmul work. FlashAttention-3 uses Hopper’s asynchronous copies (TMA) and warp specialization (some warps load, others compute) and supports FP8. FlashInfer and vLLM’s kernels add paged KV caches (Chapter 25) and many variants. GPU Mode L12 (Thomas Viehmann’s FlashAttention lecture, with a from-scratch CUDA version) and L36 (FlashAttention-3 with CUTLASS) cover them.
Decode is different: flash-decoding
During decode, there’s one query per sequence and a long KV cache. One program per (batch, head) would leave most of the GPU idle: with batch 1 and 8 KV heads, that’s 8 programs for 132 SMs. Flash-decoding splits the KV sequence into chunks processed by different programs in parallel. Each produces a partial result $(m_i, \ell_i, a_i)$ for its chunk, and a second small kernel merges them with the same rescaling rule:
$$ m = \max_i m_i,\qquad \ell = \sum_i e^{m_i - m},\ell_i,\qquad a = \sum_i e^{m_i - m},a_i,\qquad O = a/\ell . $$
This merge rule is associative: any number of partial softmaxes can be combined in any grouping. That same property powers ring attention across multiple GPUs (GPU Mode L13) and the paged decode kernel in Chapter 25.
What FlashAttention does and doesn’t change
- Changes: memory for intermediates goes from $O(T^2)$ to $O(T)$, and HBM traffic drops several-fold, so long-sequence attention becomes compute-bound and fast.
- Doesn’t change: the result, up to rounding, and the arithmetic. Attention still costs $O(T^2 d)$ FLOPs. At very long contexts that quadratic arithmetic itself becomes the bottleneck, which is the motivation for Chapters 28 and 29.
Build it
Engine milestone 15: FlashAttention. Implement online_attention in engine/attention.py and flash_fwd_kernel in engine/kernels/triton_flash.py (the launcher flash_attention is provided).
pytest tests/test_ch15_flash.py
The tests try tile sizes from 1 to 64 (the answer must not depend on the tile size), rectangular queries with GQA, the worked example above, and the Triton kernel on prefill, chunked-prefill and decode shapes. On a GPU, benchmark your kernel against F.scaled_dot_product_attention for T = 512 to 8,192.
Stretch exercises
- ★ Instrument
causal_attentionandonline_attentionwithtorch.cuda.max_memory_allocated()for T = 4,096. How much memory does the online version save? Where:experiments/ch15.py(create it), importing both functions fromengine.attention. - ★★ Implement the flash-decoding merge: split
online_attention’s key range into 4 chunks, compute each chunk’s $(m, \ell, a)$ separately, and merge. Verify against the unsplit result. Where: add a split/merge attention helper besideonline_attentioninengine/attention.py. - ★★ Add a sliding-window option to the Triton kernel: skip tiles entirely before
q_pos - window, and mask within the boundary tile. Where:flash_fwd_kernelandflash_attentioninengine/kernels/triton_flash.py. - ★★★ Write the backward pass of attention in PyTorch, using FlashAttention’s trick: recompute $P$ from the saved log-sum-exp per row instead of storing it. Check it against autograd. Where: add a backward helper in
engine/attention.py; compare with autograd intests/test_ch15_stretch.py.
Check your understanding
- Why is the accumulator $a$ kept unnormalized until the very end?
- Why is applying an ordinary softmax to each tile independently incorrect?
- Which complexity does FlashAttention reduce: arithmetic, intermediate memory, or both?
- Why does decode need a different parallelization (flash-decoding) than prefill?
Going deeper
- PMPP §20.5 (pp. 492-503): FlashAttention, derived and implemented in CUDA; §20.6 (KV-cache arithmetic intensity) previews Chapter 16.
- GPU Mode L12 (FlashAttention, with
flash_attention.cuand a notebook), L13 (ring attention and the log-sum-exp merge, withhowto_log_sum_exp.ipynb), L36 (FlashAttention-3). - Dao et al., FlashAttention (2022) and FlashAttention-2 (2023); Shah et al., FlashAttention-3 (2024); Milakov and Gimelshein, Online normalizer calculation for softmax (2018).
- The Triton tutorial Fused Attention, which this chapter’s kernel simplifies.
16. Prefill, decode and the KV cache
In this chapter
- Why the generation loop from Chapter 8 does quadratic work, and the observation that removes it.
- The KV cache: what to store, how to lay it out, and how positions and masks continue across calls.
- Prefill versus decode as two different workloads, and chunked prefill.
- KV-cache memory arithmetic, and why it shapes modern architectures.
You will build
engine/kv_cache.py: a preallocated per-layer cache that both your GPT and (next chapter) your Qwen3 use, with exact full / chunked / token-by-token equivalence.
Time: 4-6 hours. GPU: not needed.
Generation without a cache repeats itself
Chapter 8’s loop with cached=False re-runs the entire sequence for every new token. With a prompt of $P$ tokens, generating $N$ tokens processes $P$, then $P+1$, …, up to $P+N-1$ positions. Each of those passes recomputes the same projections, MLPs and attention for every earlier position, getting exactly the same numbers as last time. The total work grows like $N \cdot P + N^2/2$, so it’s quadratic in the output length.
Measured on the reference engine (a tiny 4-layer Qwen3 on a laptop CPU, python run.py cache):
| new tokens | uncached | cached |
|---|---|---|
| 16 | 0.10 s | 0.05 s |
| 64 | 0.50 s | 0.19 s |
| 128 | 1.90 s | 0.33 s |
The uncached time grows quadratically and the cached time roughly linearly. Real models with thousand-token prompts make the gap enormous.
The observation: the past doesn’t change
In a causal model, position $t$’s hidden states depend only on tokens $\le t$. Appending a token can’t change anything computed for earlier positions. So their keys and values, the only things later tokens read from them, are fixed once computed.
A new token needs, at every layer: its own query, and the keys and values of all positions so far. So store each layer’s K and V as they’re computed, and on the next step process only the new token: compute its Q, K and V, append K and V to the cache, and attend over the whole cache. That’s the KV cache. Two things it deliberately does not store:
- Old queries. A query is used only once, by its own position.
- Old hidden states or MLP results. They’re never read again; only K and V are.
Prefill and decode
Generation now has two phases, as Chapter 1 promised:
prefill: input IDs at positions 0,1,2 write K/V at 0,1,2 last logits -> token 3
decode 1: input token 3 at position 3 write K/V at 3 logits -> token 4
decode 2: input token 4 at position 4 write K/V at 4 logits -> token 5
Two off-by-one facts trip people up:
- The first new token comes from prefill. The prompt’s last position already predicts it. No separate decode step is needed before the first sample.
- The last generated token is never fed back unless you want another prediction. After $N$ new tokens, decode has run $N-1$ times.
For a prompt of 3 tokens and 3 generated tokens, the uncached loop projects $3 + 4 + 5 = 12$ token positions; the cached one projects $3 + 1 + 1 = 5$.
Layout and validity
For each layer, store keys and values as [B, H_kv, capacity, D_h], preallocated to the maximum length you’ll need, plus a count of valid positions:
class KVCache:
"""Contiguous preallocated cache; every row in the batch has the same length. (Your engine: Chapter 16)"""
def __init__(self, layers, batch, kv_heads, capacity, head_dim, device="cpu", dtype=torch.float32):
self.capacity = capacity
shape = (batch, kv_heads, capacity, head_dim)
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.lengths = [0] * layers
@property
def length(self):
"""Tokens already processed into every layer."""
if len(set(self.lengths)) != 1:
raise RuntimeError("A forward pass stopped part-way; reset or truncate the cache")
return self.lengths[0]
def update(self, layer, k, v, positions=None, rows=None):
"""Write k/v after the existing entries and return views of the whole valid history. (Your engine: Chapter 16)"""
if rows is not None:
raise ValueError("KVCache rows always advance together; use StaticKVCache for slots")
start = self.lengths[layer]
end = start + k.shape[2]
if end > self.capacity:
raise ValueError(f"KV capacity {self.capacity} exceeded")
if k.shape[:2] != self.keys[layer].shape[:2] or k.shape[3] != self.keys[layer].shape[3]:
raise ValueError("Cache batch/head/width shapes differ")
self.keys[layer][:, :, start:end].copy_(k)
self.values[layer][:, :, start:end].copy_(v)
self.lengths[layer] = end
return self.keys[layer][:, :, :end], self.values[layer][:, :, :end], None
def truncate(self, length):
"""Forget everything after `length` tokens (speculative-decoding rollback). (Your engine: Chapter 16)"""
if not 0 <= length <= self.length:
raise ValueError("Invalid truncation length")
self.lengths = [length] * len(self.lengths)
def reset(self):
self.truncate(0)
@property
def bytes(self):
return sum(t.numel() * t.element_size() for t in self.keys + self.values)
struct KvCache {
std::vector<std::vector<float>> keys, values; // per layer: [positions, kv_heads*head_dim]
size_t width, len = 0;
KvCache(size_t layers, size_t width_, size_t capacity) : keys(layers), values(layers), width(width_) {
for (size_t l = 0; l < layers; ++l) { keys[l].reserve(capacity * width); values[l].reserve(capacity * width); }
}
void append(size_t layer, const std::vector<float>& k, const std::vector<float>& v) {
keys[layer].insert(keys[layer].end(), k.begin(), k.end());
values[layer].insert(values[layer].end(), v.begin(), v.end());
}
void truncate(size_t n) {
for (size_t l = 0; l < keys.size(); ++l) { keys[l].resize(n * width); values[l].resize(n * width); }
len = n;
}
};
#![allow(unused)]
fn main() {
/// Rows are positions; each row holds kv_heads * head_dim values. Capacity is reserved up
/// front so appending never reallocates (and never copies the history).
pub struct KvCache {
pub keys: Vec<Vec<f32>>, // one buffer per layer
pub values: Vec<Vec<f32>>,
pub width: usize, // kv_heads * head_dim
pub len: usize, // positions already processed
}
impl KvCache {
pub fn new(layers: usize, width: usize, capacity: usize) -> Self {
KvCache {
keys: (0..layers).map(|_| Vec::with_capacity(capacity * width)).collect(),
values: (0..layers).map(|_| Vec::with_capacity(capacity * width)).collect(),
width,
len: 0,
}
}
/// Append one position's K and V for one layer; return the whole history for that layer.
pub fn append(&mut self, layer: usize, k: &[f32], v: &[f32]) -> (&[f32], &[f32]) {
self.keys[layer].extend_from_slice(k);
self.values[layer].extend_from_slice(v);
(&self.keys[layer], &self.values[layer])
}
/// Roll back to `len` positions (speculative decoding).
pub fn truncate(&mut self, len: usize) {
for layer in 0..self.keys.len() {
self.keys[layer].truncate(len * self.width);
self.values[layer].truncate(len * self.width);
}
self.len = len;
}
pub fn bytes(&self) -> usize {
2 * self.keys.len() * self.len * self.width * 4
}
}
}
Design choices worth noticing:
- Preallocate; don’t concatenate.
torch.caton every step reallocates and copies the whole history, which is $O(T)$ work per token and $O(T^2)$ overall, the very cost the cache was meant to remove. Writing into a preallocated slice is $O(1)$ per token. (Chapter 25’s paged cache gets the same benefit without reserving the maximum up front.) - Capacity versus valid length. Slots beyond
lengthhold zeros or stale data.updatereturns a view of just the valid prefix, so attention never sees them. - Per-layer lengths. Every layer appends during a forward pass. If an exception interrupts a pass halfway, some layers are one token ahead of others. The
lengthproperty detects this and refuses to continue, rather than silently misaligning positions. - Truncate for rollback.
truncate(n)forgets everything after $n$ tokens without copying. Speculative decoding (Chapter 26) depends on this.
The cache follows a one-method protocol: update(layer, k, v, positions, rows) writes new entries and returns everything attention may read. Your models never touch cache internals. Later caches (static for CUDA graphs in Chapter 19, slot-based for batching in Chapter 24, paged in Chapter 25) implement the same method and plug into the same model code unchanged.
Positions continue, and masks become rectangles
When a decode step feeds token 57, it must be treated as position 57, not position 0. GPT adds position_embedding[57], and Qwen rotates its query and key by angle 57 (Chapter 17). Your models compute positions as cache.length + arange(T).
The mask changes too. Without a cache, queries and keys are the same $T$ positions, and the mask is a square lower triangle. With a cache, $T$ new queries attend to $S$ keys ($S > T$): a rectangle. For a cached prefix of 3 tokens and a chunk of 2 new tokens:
key 0 1 2 3 4
query pos 3 1 1 1 1 0
query pos 4 1 1 1 1 1
A square-triangle shortcut applied to this 2-row chunk would let query 3 see only key 0 and query 4 only keys 0-1, hiding the most recent history. This is exactly why your causal_attention takes positions (Chapter 5): the rule key_pos <= query_pos produces the right rectangle automatically. (Watch out for PyTorch’s scaled_dot_product_attention(is_causal=True): with $T \ne S$ it uses a top-left-aligned triangle, which is the wrong rectangle for cached decoding.)
Prove equivalence
The cache is an optimization, so it must change nothing except speed. The strongest test feeds one fixed sequence three ways:
- the whole sequence at once, with no cache;
- a prefix, then a multi-token chunk, then single tokens, all through one cache;
- compare all logits position by position.
They must agree to floating-point tolerance (about 1e-5 in FP32). Use fixed token IDs, not generated ones: once two runs pick different tokens, they’re processing different inputs and the comparison is meaningless.
How big is the cache?
Each layer stores K and V for every position, so:
$$ \text{KV bytes} = 2 \times \text{layers} \times \text{batch} \times \text{tokens} \times H_{kv} \times D_h \times \text{bytes per value}. $$
For Qwen3-0.6B (28 layers, 8 KV heads of dimension 128, BF16):
- 112 KiB per token;
- 224 MiB for a 2,048-token conversation;
- 3.5 GiB for a 32k-token context, about three times the model’s own weights.
def kv_cache_bytes(layers, batch, tokens, kv_heads, head_dim, bytes_per_value=2):
"""2 (K and V) x layers x batch x tokens x kv_heads x head_dim x bytes."""
return 2 * layers * batch * tokens * kv_heads * head_dim * bytes_per_value
Why the cache shapes model design
Two consequences follow from this formula, and they explain several architecture choices you’ll meet:
- Capacity. In serving, the cache limits how many conversations fit on a GPU at once (Chapters 24-25). Note that the formula uses KV heads, not query heads. Grouped-query attention with 8 KV heads instead of 32 shrinks the cache 4x: for an 8B-class model (32 layers, head dimension 128) at 8k tokens, that’s 1 GiB per sequence instead of 4 GiB. Multi-query attention (one KV head) and DeepSeek’s multi-head latent attention push further.
- Bandwidth. Each decode step reads the whole cache of each sequence once, in addition to the weights. Attention during decode has an arithmetic intensity of about 1 FLOP per byte per query head (PMPP §20.6), as memory-bound as the weight reads. At long contexts, the KV reads dominate the step time. That’s what motivates KV-cache quantization (Chapter 20), sparse attention that reads only part of the cache (Chapter 29), and linear attention with a fixed-size state instead of a cache (Chapter 28).
Build it
Engine milestone 16: a KV cache. Implement KVCache.update and KVCache.truncate in engine/kv_cache.py (the constructor, length, reset and bytes are provided). Your GPT from Chapter 6 already calls cache.update(...).
pytest tests/test_ch16_kv_cache.py
python run.py cache --impl engine
The tests run full, chunked and token-by-token forwards through GPT and Qwen3 models and compare all logits, check rollback with truncate, check that capacity is enforced, and check the memory formula against cache.bytes.
Stretch exercises
- ★ Replace the preallocated write with
torch.catand measure generation time for 64, 256 and 1,024 new tokens. Where does the quadratic copy cost become visible? Where: a separate concatenating-cache variant ofKVCache.updateinengine/kv_cache.py; compare it inexperiments/ch16.py(create it). - ★★ Add a
layer_bytes()report and print how much of a 32k-token Qwen3-0.6B decode step’s traffic is KV cache versus weights. Where: addKVCache.layer_bytesinengine/kv_cache.py. - ★★ Implement a sliding-window cache that keeps only the last $w$ positions in a ring buffer. What must change in the positions passed to attention? Where: add a ring-buffer cache class in
engine/kv_cache.py; wire its absolute positions intoQwen3Attention.forwardinengine/qwen3.py. - ★★★ Store the cache in FP8 (E4M3, with one scale per head) and dequantize on read. Measure the logit error against a BF16 cache for 1,000 decode steps. Where: add an FP8 cache variant in
engine/kv_cache.py; integrate its dequantized reads inengine/qwen3.py.
Check your understanding
- Why cache keys and values, but not queries?
- Which forward pass produces the logits for the first generated token?
- Why can a lower-triangular mask be wrong for a 2-token chunk after a 3-token prefix?
- Why does grouped-query attention reduce cache size, but not the number of attention score computations?
- Why does decode at long context become dominated by cache reads?
Going deeper
- PMPP §20.4 (pp. 488-492): KV caching; §20.6 (pp. 504-508): the arithmetic intensity and memory requirement of the KV cache; §20.7: MQA and GQA.
- Pope et al., Efficiently Scaling Transformer Inference (2022): the analysis of cache memory and bandwidth that much of modern serving builds on.
- Shazeer, Fast Transformer Decoding: One Write-Head is All You Need (MQA, 2019); DeepSeek-V2 (2024) for multi-head latent attention.
17. A modern architecture: Qwen3
In this chapter
- What changed between GPT-2 (2019) and today's open models, and why.
- RMSNorm, rotary position embeddings (RoPE), grouped-query attention with Q/K normalization, and SwiGLU, each derived and worked with numbers.
- How to treat an architecture as an exact contract with its checkpoint.
You will build
engine/qwen3.py: dense Qwen3 from scratch, verified against the official Hugging Face implementation, ready to load the real 0.6B checkpoint in Chapter 18.
Time: 5-7 hours. GPU: not needed.
From GPT-2 to a 2025 model
The transformer block hasn’t changed shape since GPT-2: embed, then $L$ blocks of attention and MLP on a residual stream, then norm and head. What changed is the ingredients. Almost every open model since Llama (2023), including Qwen, Mistral and DeepSeek, made the same set of substitutions:
| GPT-2 small | Qwen3-0.6B | why it changed | |
|---|---|---|---|
| normalization | LayerNorm (mean, variance, γ, β) | RMSNorm (RMS, γ only) | cheaper, just as good |
| positions | learned table added to embeddings | RoPE: rotate Q and K | relative positions, longer contexts |
| attention | 12 heads, MHA | 16 query heads, 8 KV heads (GQA) | half the KV cache |
| Q/K | used directly | RMSNorm per head on Q and K | training stability |
| head width | D / heads = 64 | head_dim = 128, independent of D | more expressive heads |
| MLP | 4D, GELU | SwiGLU, 3D, gated | better quality per parameter |
| biases | everywhere | none | no measurable benefit |
| context | 1,024 (table size) | 40,960 configured | RoPE has no table |
Dense Qwen3 is the first real model your engine will run. Here’s the contract for the 0.6B checkpoint, from its official config.json:
| field | value |
|---|---|
| hidden size $D$ | 1,024 |
| layers | 28 |
| query heads / KV heads | 16 / 8 |
| head dimension | 128 |
| MLP intermediate size | 3,072 |
| vocabulary | 151,936 |
| RMSNorm ε | 1e-6 |
| RoPE base θ | 1,000,000 |
| tied embeddings | yes |
| max positions | 40,960 |
Notice that the query projection has width $16 \times 128 = 2{,}048$, which is twice the hidden size. Code that assumes head_dim = hidden_size / heads reshapes the projections wrongly and crashes (or worse, doesn’t). The book’s tiny default config deliberately has the same mismatch, so the bug surfaces in the tests.
RMSNorm
LayerNorm subtracts the mean and divides by the standard deviation. RMSNorm (Zhang and Sennrich, 2019) skips the mean and divides by the root mean square:
$$ \operatorname{RMSNorm}(x)_i = \gamma_i, \frac{x_i}{\sqrt{\frac1D\sum_j x_j^2 + \epsilon}} . $$
For $x = [3, 4]$: the mean square is $(9 + 16)/2 = 12.5$, the RMS is 3.5355, and the output is $[0.8485, 1.1314]$ before $\gamma$. The vector’s direction is unchanged, and its RMS becomes 1. It’s one reduction instead of two, with no bias, and models train just as well.
class RMSNorm(nn.Module):
def __init__(self, width, eps=1e-6):
super().__init__()
self.weight = nn.Parameter(torch.ones(width))
self.eps = eps
def forward(self, x):
"""x / sqrt(mean(x^2) + eps) * weight, reduced in FP32. (Your engine: Chapter 17)"""
x32 = x.float()
normalized = x32 * torch.rsqrt(x32.square().mean(-1, keepdim=True) + self.eps)
return self.weight * normalized.to(x.dtype)
Two parity details matter when loading real weights:
- Compute the mean of squares in FP32 even for BF16 inputs (Chapter 13).
- The cast order is part of the contract. Qwen3 normalizes in FP32, casts back to the input dtype, then multiplies by $\gamma$. Gemma and the Qwen3.5 family use $(1+\gamma)$ with $\gamma$ initialized at zero, multiplying before the cast (you’ll meet that “zero-centered” variant in Chapter 29). In BF16, the difference changes the last bits of every layer’s output.
Rotary position embeddings
GPT-2 added a learned position vector to each token. RoPE (Su et al., 2021) instead rotates the query and key vectors by an angle proportional to their position, right before the attention scores are computed.
Take one pair of coordinates and rotate it by angle $\phi$:
$$ R(\phi) = \begin{pmatrix} \cos\phi & -\sin\phi \ \sin\phi & \cos\phi \end{pmatrix}. $$
Rotate a query at position $p$ by $p,\omega$ and a key at position $s$ by $s,\omega$. Their dot product is
$$ \big(R(p\omega),q\big)\cdot\big(R(s\omega),k\big) = q^{\top} R\big((s-p),\omega\big), k , $$
because rotations compose and $R(a)^\top = R(-a)$. The score depends only on the distance $s - p$, not on absolute positions. Shift both tokens by 10 and nothing changes. For example, with $q = k = [1, 0]$ and $\omega = 0.35$, two tokens 2 apart score $\cos(0.7) = 0.7648$, at positions (0, 2) or (10, 12) alike.
A head of dimension $d$ has $d/2$ coordinate pairs, each rotating at its own frequency:
$$ \omega_i = \theta^{-2i/d}, \qquad i = 0, \ldots, d/2 - 1 . $$
With $\theta = 10^6$ and $d = 128$, the fastest pair turns 1 radian per token and the slowest about $1.2 \times 10^{-6}$ radians, a wavelength of about 5 million tokens. Fast pairs resolve nearby order; slow pairs let distant tokens still be distinguished. Raising $\theta$ (Qwen3 uses $10^6$, the original paper $10^4$) stretches every wavelength, which helps long contexts.
The pairing convention is part of the contract
Which coordinates form a pair? The original paper paired adjacent coordinates $(0,1), (2,3), \ldots$ Llama, Qwen and most Hugging Face models pair coordinate $i$ with $i + d/2$, the split-half (“rotate half”) convention. Both are valid RoPEs, and a model is trained with exactly one. Using the other produces correct shapes and fluent-looking garbage. Our implementation uses split-half:
def rope_cos_sin(positions, rotary_dim, theta, dtype=torch.float32):
"""cos/sin tables for absolute positions [T] or [B, T] -> [B or 1, 1, T, rotary_dim]. (Your engine: Chapter 17)
Frequencies theta^(-2i/rotary_dim) for i in [0, rotary_dim/2), repeated for the two halves.
"""
if rotary_dim % 2:
raise ValueError("RoPE needs an even number of rotated dimensions")
inv_freq = theta ** (-torch.arange(0, rotary_dim, 2, device=positions.device, dtype=torch.float32) / rotary_dim)
pos = positions.float() if positions.ndim == 2 else positions.float()[None]
angles = pos[..., None] * inv_freq # [B, T, rotary_dim/2]
angles = torch.cat((angles, angles), dim=-1)[:, None] # [B, 1, T, rotary_dim]
return angles.cos().to(dtype), angles.sin().to(dtype)
def rotate_half(x):
a, b = x.chunk(2, dim=-1)
return torch.cat((-b, a), dim=-1)
def apply_rope(x, cos, sin):
"""Rotate the first cos.shape[-1] features of x (split-half pairing); pass the rest through. (Your engine: Chapter 17)
Pairs are (i, i + rotary_dim/2), NOT adjacent features (2i, 2i+1): that is the checkpoint's
convention and getting it wrong still produces plausible-looking shapes.
"""
rotary = cos.shape[-1]
x_rot, x_pass = x[..., :rotary], x[..., rotary:]
rotated = x_rot * cos + rotate_half(x_rot) * sin
return torch.cat((rotated, x_pass), dim=-1) if x_pass.shape[-1] else rotated
// Split-half pairing: feature i rotates with feature i + d/2.
inline void rope(float* x, size_t d, size_t pos, float theta) {
for (size_t i = 0; i < d / 2; ++i) {
float freq = std::pow(theta, -2.f * i / d), s = std::sin(pos * freq), c = std::cos(pos * freq);
float a = x[i], b = x[i + d / 2];
x[i] = a * c - b * s;
x[i + d / 2] = a * s + b * c;
}
}
#![allow(unused)]
fn main() {
/// Rotate one head vector in place for absolute position `pos`, pairing feature i with
/// feature i + d/2 (the split-half convention used by Qwen, Llama and the book's Python code).
pub fn rope(x: &mut [f32], pos: usize, theta: f32) {
let half = x.len() / 2;
for i in 0..half {
let freq = theta.powf(-2.0 * i as f32 / x.len() as f32);
let (sin, cos) = (pos as f32 * freq).sin_cos();
let (a, b) = (x[i], x[i + half]);
x[i] = a * cos - b * sin;
x[i + half] = a * sin + b * cos;
}
}
}
apply_rope rotates only the first cos.shape[-1] features and passes the rest through. Qwen3 rotates all 128, but Flash-Next’s attention rotates only the first 64 of 256 (partial RoPE, Chapter 29). The same function handles both.
Where RoPE sits in the computation matters for caching: Q and K are rotated before K is cached. A cached key already contains its position. Never rotate cached keys again; each new token rotates only its own Q and K, at its own position.
Attention with grouped KV heads and Q/K normalization
Qwen3’s attention is your Chapter 5 attention with three additions:
- Separate Q, K, V projections, no biases, with output widths $H_q D_h$, $H_{kv} D_h$, $H_{kv} D_h$.
- Q/K norm: an RMSNorm over each head’s 128 features, with a learned $\gamma$ shared across heads, applied to Q and K before RoPE. This keeps attention logits from growing without bound during training, a known source of instability.
- GQA: 16 query heads share 8 KV heads; your
causal_attentionhandles the mapping.
class Qwen3Attention(nn.Module):
def __init__(self, cfg):
super().__init__()
if cfg.num_attention_heads % cfg.num_key_value_heads:
raise ValueError("Query heads must be a multiple of KV heads")
self.cfg = cfg
d, hq, hkv = cfg.head_dim, cfg.num_attention_heads, cfg.num_key_value_heads
self.q_proj = nn.Linear(cfg.hidden_size, hq * d, bias=False)
self.k_proj = nn.Linear(cfg.hidden_size, hkv * d, bias=False)
self.v_proj = nn.Linear(cfg.hidden_size, hkv * d, bias=False)
self.o_proj = nn.Linear(hq * d, cfg.hidden_size, bias=False)
self.q_norm = RMSNorm(d, cfg.rms_norm_eps)
self.k_norm = RMSNorm(d, cfg.rms_norm_eps)
def forward(self, x, positions, rope, cache=None, layer=0, rows=None):
"""Project, per-head RMSNorm on q and k, RoPE, cache, GQA attention, output projection. (Your engine: Chapter 17)"""
b, t, _ = x.shape
c = self.cfg
q = self.q_proj(x).view(b, t, c.num_attention_heads, c.head_dim)
k = self.k_proj(x).view(b, t, c.num_key_value_heads, c.head_dim)
v = self.v_proj(x).view(b, t, c.num_key_value_heads, c.head_dim)
q = self.q_norm(q).transpose(1, 2) # norm over head_dim, before RoPE
k = self.k_norm(k).transpose(1, 2)
v = v.transpose(1, 2)
cos, sin = rope
q, k = apply_rope(q, cos, sin), apply_rope(k, cos, sin)
key_positions = None
if cache is not None:
k, v, key_positions = cache.update(layer, k, v, positions, rows) # keys are cached already rotated
y = causal_attention(q, k, v, positions, key_positions)
return self.o_proj(y.transpose(1, 2).reshape(b, t, -1))
SwiGLU
The GPT-2 MLP was expand → GELU → contract. Qwen3’s MLP has two expansions, one of which gates the other (Shazeer, 2020):
$$ \operatorname{SwiGLU}(x) = W_{\text{down}}\big(\operatorname{SiLU}(W_{\text{gate}},x) \odot W_{\text{up}},x\big), \qquad \operatorname{SiLU}(z) = z,\sigma(z). $$
The gate decides, feature by feature, how much of the “up” signal passes. If a gate pre-activation is 0, $\operatorname{SiLU}(0) = 0$ and the feature is shut off. If it’s 2, $\operatorname{SiLU}(2) = 1.7616$, and with an up value of 3 the feature passes $5.2848$. Gated MLPs consistently outperform ungated ones at the same parameter count. To keep that count comparable with a $4D$ GELU MLP, the intermediate size is usually about $\frac{8}{3}D$ (Qwen3-0.6B uses $3D$).
class SwiGLU(nn.Module):
def __init__(self, width, hidden):
super().__init__()
self.gate_proj = nn.Linear(width, hidden, bias=False)
self.up_proj = nn.Linear(width, hidden, bias=False)
self.down_proj = nn.Linear(hidden, width, bias=False)
def forward(self, x):
"""down(silu(gate(x)) * up(x)). (Your engine: Chapter 17)"""
return self.down_proj(nn.functional.silu(self.gate_proj(x)) * self.up_proj(x))
The whole model
class Qwen3Layer(nn.Module):
def __init__(self, cfg):
super().__init__()
self.input_layernorm = RMSNorm(cfg.hidden_size, cfg.rms_norm_eps)
self.self_attn = Qwen3Attention(cfg)
self.post_attention_layernorm = RMSNorm(cfg.hidden_size, cfg.rms_norm_eps)
self.mlp = SwiGLU(cfg.hidden_size, cfg.intermediate_size)
def forward(self, x, positions, rope, cache=None, layer=0, rows=None):
x = x + self.self_attn(self.input_layernorm(x), positions, rope, cache, layer, rows)
return x + self.mlp(self.post_attention_layernorm(x))
class Qwen3(nn.Module):
def __init__(self, cfg):
super().__init__()
self.cfg = cfg
self.model = Qwen3Backbone(cfg)
self.lm_head = nn.Linear(cfg.hidden_size, cfg.vocab_size, bias=False)
if cfg.tie_word_embeddings:
self.lm_head.weight = self.model.embed_tokens.weight
@property
def context_limit(self):
return self.cfg.max_position_embeddings
def forward(self, ids, cache=None, positions=None, rows=None, return_hidden=False):
"""ids [B, T] -> logits [B, T, vocab]. Same cache/position contract as GPT.forward. (Your engine: Chapter 17)"""
if positions is None:
start = cache.length if cache is not None else 0
if start + ids.shape[1] > self.context_limit:
raise ValueError("Sequence exceeds max_position_embeddings")
positions = torch.arange(start, start + ids.shape[1], device=ids.device)
x = self.model.embed_tokens(ids)
rope = rope_cos_sin(positions, self.cfg.head_dim, self.cfg.rope_theta, x.dtype)
for i, layer in enumerate(self.model.layers):
x = layer(x, positions, rope, cache, i, rows)
hidden = self.model.norm(x)
return hidden if return_hidden else self.lm_head(hidden)
def cache_spec(self):
c = self.cfg
return c.num_hidden_layers, c.num_key_value_heads, c.head_dim
def new_cache(self, batch, capacity):
if capacity > self.context_limit:
raise ValueError("Requested cache exceeds the context limit")
c, p = self.cfg, self.model.embed_tokens.weight
return KVCache(c.num_hidden_layers, batch, c.num_key_value_heads, capacity, c.head_dim, p.device, p.dtype)
The parameter budget of Qwen3-0.6B, counted from this code:
| part | parameters |
|---|---|
| embedding (tied with the head) | 155,582,464 |
| per layer: attention (q, k, v, o + q/k norms) | 6,291,712 |
| per layer: SwiGLU MLP | 9,437,184 |
| per layer total (with 2 RMSNorms) | 15,730,944 |
| whole model (28 layers + final norm) | 596,049,920 |
A quarter of this small model is the embedding table. In bigger models, the 28 × 15.7M of layers dominates.
Why parity tests matter now
From now on, your models must match someone else’s trained weights exactly, and small discrepancies don’t announce themselves. A wrong RoPE pairing, an epsilon of 1e-5 instead of 1e-6, LayerNorm instead of RMSNorm, or norm-then-RoPE swapped to RoPE-then-norm all produce a model that runs, generates English-like text after loading, and is subtly or badly wrong. The milestone tests therefore compare your Qwen3 against Hugging Face’s on small random configurations. Random weights exercise every code path, and a 1e-5 match on random weights is strong evidence of a correct implementation.
Build it
Engine milestone 17: Qwen3. In engine/qwen3.py, implement RMSNorm.forward, rope_cos_sin, apply_rope, Qwen3Attention.forward, SwiGLU.forward and Qwen3.forward. Module construction, configuration parsing and cache creation are provided, and the parameter names already match the official checkpoint.
uv pip install -r optional-requirements.txt # for the comparison with Transformers
pytest tests/test_ch17_qwen3.py
python run.py cache --impl engine # your Qwen3, cached vs uncached
The tests check RMSNorm on the worked example, RoPE’s relative-position property and norm preservation, partial rotation, the query-width-versus-hidden-size trap, and logits against Hugging Face’s Qwen3 with both tied and untied heads.
Stretch exercises
- ★ Verify numerically that RoPE preserves vector norms, and that the score at distance 0 is just the unrotated dot product. Where:
experiments/ch17.py(create it), importingengine.qwen3.rope_cos_sinandapply_rope. - ★★ Implement the adjacent-pair RoPE convention and show that loading Qwen3 weights with it changes the logits (by how much?). Where: add an adjacent-pair variant of
apply_ropeinengine/qwen3.pyand select it inQwen3Attention.forward. - ★★ Implement YaRN or linear RoPE scaling (multiply positions by a factor, or interpolate frequencies) and explain what it does to the wavelength table. Where:
rope_cos_sininengine/qwen3.py. - ★★★ Write a fused Triton kernel that applies Q/K-norm and RoPE to the Q and K projections in one pass, and plug it into
Qwen3Attentionbehind a backend switch. Where: newengine/kernels/triton_rope.py, selected byQwen3Attention.forwardinengine/qwen3.py.
Check your understanding
- Why can a correct-looking reshape still be wrong for Qwen3’s attention projections?
- What happens to a key, in order, before it enters the cache?
- What property of rotations makes a shared position shift cancel out in a RoPE dot product?
- Why does SwiGLU have three weight matrices where GPT-2’s MLP has two?
- Why are random-weight parity tests enough to establish the architecture, before you ever load the real checkpoint?
Going deeper
- BALLM has no Qwen3 chapter, but Raschka’s companion repository does:
ch05/11_qwen3builds Qwen3 from scratch in the book’s style, andch05/07_gpt_to_llamawalks the GPT-2 → Llama changes one by one. - Su et al., RoFormer: Enhanced Transformer with Rotary Position Embedding (2021); Zhang and Sennrich, Root Mean Square Layer Normalization (2019); Shazeer, GLU Variants Improve Transformer (2020).
- Qwen Team, Qwen3 Technical Report (2025); the official
modeling_qwen3.pyin Transformers.
18. Engine v1: load Qwen3 and chat
In this chapter
- What an engine owns, beyond the model: the request contract, state lifetime, stopping, measurements.
- Loading a real checkpoint efficiently and safely, and the text boundary: tokenizer, chat template, stop tokens.
- A ladder of evidence that your engine computes exactly what the checkpoint expects.
- Measuring time-to-first-token and decode speed against the hardware ceiling.
You will build
load_qwen3 in engine/loaders.py and LLM.stream in engine/engine.py. At the end of the chapter you chat with Qwen3-0.6B through code you wrote, in Python, and on the CPU in Rust.
Time: 5-7 hours. GPU: recommended (Qwen3-0.6B also runs acceptably on a modern CPU).
From model to engine
You have a model (Chapter 17), a sampler (8) and a cache (16). An engine joins them under a contract:
- Input: token IDs plus generation settings (temperature, top-p, max new tokens, stop tokens, seed).
- Output: generated token IDs, a finish reason (
stoporlength), and measurements. - Owned state: the KV cache for the request, created when the request starts and released when it ends, even if it ends with an error.
Text appears only at the boundary. Turning a chat into token IDs and IDs back into text is the tokenizer’s job, and keeping it outside the engine core has practical benefits. The engine is testable with integer inputs alone, the same engine serves different tokenizers, and token-level outputs make parity checks exact.
Loading Qwen3-0.6B
Download the checkpoint (if you didn’t in Chapter 1):
uv pip install -r optional-requirements.txt
hf download Qwen/Qwen3-0.6B --local-dir models/Qwen3-0.6B
Because your parameter names match the checkpoint’s exactly (Chapter 17), loading is a checked copy:
@torch.no_grad()
def load_qwen3(directory, device="cpu", dtype=torch.bfloat16):
"""Dense Qwen3: our parameter names equal the checkpoint's, so this is a checked copy. (Your engine: Chapter 18)"""
cfg = Qwen3Config.from_hf(read_config(directory))
with torch.device("meta"):
model = Qwen3(cfg) # no memory yet: parameters are shapes only
model = model.to_empty(device=device).to(dtype)
if cfg.tie_word_embeddings: # to_empty breaks the tie; restore it
model.lm_head.weight = model.model.embed_tokens.weight
destinations = dict(model.named_parameters())
loaded, head_alias = set(), None
for name, value in snapshot_tensors(directory):
if name == "lm_head.weight" and cfg.tie_word_embeddings:
head_alias = value
continue
if name not in destinations:
raise ValueError(f"Unexpected tensor {name}")
assign(destinations[name], value, name)
loaded.add(name)
missing = set(destinations) - loaded
if missing:
raise ValueError(f"Missing tensors: {sorted(missing)[:5]}")
if head_alias is not None and not torch.equal(head_alias.to(model.lm_head.weight), model.lm_head.weight):
raise ValueError("Tied lm_head and embed_tokens disagree")
return model.eval()
Three engineering points:
- Build on the
metadevice first.Qwen3(cfg)on the CPU would allocate and randomly initialize 596M parameters, only to overwrite them. Built undertorch.device("meta"), parameters are shapes without storage.to_empty(device=...)then allocates uninitialized memory directly on the target device, and the loader fills it shard by shard. Peak memory is about one model plus one shard, not two models. - Restore weight tying.
to_emptycreates fresh storage for every parameter, which silently unties the head from the embedding. Re-tie them, and if the checkpoint also containslm_head.weight, verify it equals the embedding rather than letting file order decide which copy wins. - Refuse what you don’t implement.
Qwen3Config.from_hfrejects other model types, scaled RoPE and sliding windows. A loader that “mostly works” on an unsupported variant is worse than one that refuses.
The text boundary
The tokenizer comes from the checkpoint’s own files. Your engine uses Hugging Face’s implementation through AutoTokenizer, loaded with local_files_only=True so nothing is fetched or executed remotely. (Chapter 4’s BPE explains what it does; reading Qwen’s tokenizer.json with your own BPE is a stretch exercise.)
Three details decide whether chat works at all:
- The chat template wraps messages in
<|im_start|>role ... <|im_end|>markers and ends with the generation prompt<|im_start|>assistant\n(Chapter 4). - Thinking mode. Qwen3 can “think” in a
<think>...</think>block before answering.enable_thinking=Falsein the template turns it off for short, direct answers;Truegives better reasoning at the cost of many more tokens. - Stop tokens. The assistant ends its turn with
<|im_end|>(ID 151645);<|endoftext|>(151643) ends a document. Stop on both, or the model will continue with an imaginary next user turn.
The engine
class LLM:
def __init__(self, model, tokenizer=None, manifest=None):
self.model = model.eval()
self.tokenizer = tokenizer
self.device = next(model.parameters()).device
self.manifest = manifest or {}
@classmethod
def from_pretrained(cls, directory, device=None, dtype="bf16", **load_options):
device = device or default_device()
model = load_model(directory, device, parse_dtype(dtype), **load_options)
tokenizer = None
try: # The tokenizer is the only third-party piece, used at the text boundary.
from transformers import AutoTokenizer
tokenizer = AutoTokenizer.from_pretrained(directory, local_files_only=True)
except Exception:
pass
config = read_config(directory)
manifest = {"model_dir": str(directory), "model_type": config.get("model_type"),
"dtype": str(parse_dtype(dtype)), "device": str(device),
"torch": torch.__version__, "python": platform.python_version(),
"gpu": torch.cuda.get_device_name() if str(device).startswith("cuda") else None}
return cls(model, tokenizer, manifest)
def encode(self, text, chat=False, thinking=False):
if self.tokenizer is None:
raise RuntimeError("No tokenizer: pass token IDs directly")
if chat:
text = self.tokenizer.apply_chat_template([{"role": "user", "content": text}], tokenize=False,
add_generation_prompt=True, enable_thinking=thinking)
return self.tokenizer(text, add_special_tokens=False).input_ids
def default_stop_ids(self):
if self.tokenizer is None or self.tokenizer.eos_token_id is None:
return ()
ids = {self.tokenizer.eos_token_id}
for token in ("<|im_end|>", "<|endoftext|>"):
index = self.tokenizer.convert_tokens_to_ids(token)
if isinstance(index, int) and index >= 0 and index != self.tokenizer.unk_token_id:
ids.add(index)
return tuple(ids)
@torch.inference_mode()
def stream(self, prompt_ids, config=GenerationConfig()):
"""Yield token IDs as they are produced; afterwards self.last_result holds the summary. (Your engine: Chapter 18)
Measures time to first token (prefill + first sample) and the mean time per later token.
"""
ids = torch.tensor([list(prompt_ids)], device=self.device)
total = ids.shape[1] + config.max_new_tokens
if total > self.model.context_limit:
raise ValueError("Prompt plus max_new_tokens exceed the model's context")
generator = torch.Generator(device=self.device).manual_seed(config.seed)
policy = dict(temperature=config.temperature, top_k=config.top_k, top_p=config.top_p,
min_p=config.min_p, generator=generator)
cache = self.model.new_cache(1, total)
synchronize(self.device)
start = time.perf_counter()
logits = self.model(ids, cache)
produced, finish, first_token_time = [], "length", None
try:
for step in range(config.max_new_tokens):
token = sample(logits[:, -1], **policy)
value = int(token) # host sync: we must know it to stream it
if first_token_time is None:
first_token_time = time.perf_counter()
produced.append(value)
yield value
if value in config.stop_ids:
finish = "stop"
break
if step + 1 < config.max_new_tokens:
logits = self.model(token, cache)
finally:
end = time.perf_counter()
decode_steps = max(len(produced) - 1, 0)
self.last_result = GenerationResult(produced, finish, metrics={
"prompt_tokens": ids.shape[1], "new_tokens": len(produced),
"ttft_ms": round(1000 * ((first_token_time or end) - start), 3),
"ms_per_output_token": round(1000 * (end - (first_token_time or end)) / decode_steps, 3) if decode_steps else None,
"tokens_per_s": round(len(produced) / (end - start), 2) if produced else 0.0})
def generate(self, prompt, config=None, chat=False):
"""prompt is text (needs a tokenizer) or a list of token IDs."""
config = config or GenerationConfig(stop_ids=self.default_stop_ids())
prompt_ids = self.encode(prompt, chat=chat) if isinstance(prompt, str) else list(prompt)
for _ in self.stream(prompt_ids, config):
pass
result = self.last_result
if self.tokenizer is not None:
result.text = self.tokenizer.decode(result.token_ids, skip_special_tokens=True)
return result
def run_manifest(self, prompt_ids, config):
"""Everything needed to reproduce (and compare) a run."""
return {**self.manifest, "prompt_token_ids": list(prompt_ids), "generation": asdict(config)}
stream is the heart: prefill the prompt, sample, then decode one token at a time, yielding each token as soon as it exists so a UI can display it. The try/finally records the result even if the consumer stops reading early (a user pressing “stop” is a normal event, not an error). It measures two numbers users feel:
- TTFT (time to first token): from the start of the request to the first sampled token. That’s prefill plus one sample. It grows with prompt length.
- Time per output token (TPOT): the average time between later tokens. That’s decode. It’s nearly constant, rising slowly as the KV cache grows.
run_manifest records everything needed to reproduce a run: model directory and type, dtype, device, library versions, prompt token IDs and generation settings. Save one next to every result you want to keep or compare. Later, “it got faster” or “it got worse” questions become answerable.
Chat with it
python run.py chat --model-dir models/Qwen3-0.6B --prompt "Why is the sky blue? Answer in two sentences."
The text streams token by token, followed by the measurements, like this (exact text and timings will vary):
The sky appears blue because molecules in the atmosphere scatter shorter (blue) wavelengths of sunlight
much more strongly than longer (red) ones. ...
{"prompt_tokens": 23, "new_tokens": 41, "ttft_ms": ..., "ms_per_output_token": ..., "tokens_per_s": ..., "finish": "stop"}
Tip
If the model rambles, repeats the question or starts a fake “user:” turn, check the boundary before suspecting the model: was the generation prompt added, and are both stop IDs in
stop_ids?llm.encode(text, chat=True)andtokenizer.convert_ids_to_tokens(ids)show exactly what the model sees.
A ladder of evidence
“It produces English” is weak evidence. A model with a subtly wrong RoPE still writes fluent sentences. Climb this ladder instead, and stop at the first rung that fails:
- Unit operations against small hand-checked cases: Chapter 17’s tests.
- Internal consistency: full, chunked and token-by-token forwards agree (Chapter 16’s tests). This catches position and cache bugs.
- Independent parity on the real checkpoint: load the same snapshot into Hugging Face Transformers and into your engine, in FP32, and compare logits on fixed token IDs. Expect a maximum absolute difference around 1e-4 to 1e-3.
- Layer-by-layer localization when rung 3 fails: register forward hooks on both models’ decoder layers, and compare the residual stream after each layer. The first layer that disagrees contains the bug.
- Behavior: identical greedy text over several prompts, in FP32.
import torch
from transformers import AutoModelForCausalLM
from engine.loaders import load_qwen3
path = "models/Qwen3-0.6B"
ref = AutoModelForCausalLM.from_pretrained(path, dtype=torch.float32, attn_implementation="eager").eval()
mine = load_qwen3(path, device="cpu", dtype=torch.float32)
ids = torch.randint(0, 151936, (1, 33))
with torch.no_grad():
diff = (mine(ids) - ref(ids).logits).abs()
print("max abs logit difference:", diff.max().item())
For BF16, compare like with like: your BF16 against Transformers’ BF16, both judged against FP32 (the three-comparison rule from Chapter 13). Greedy tokens can legitimately diverge after a few dozen steps in BF16, because a near-tie between two tokens flips on rounding. Compare logits, not long generations.
Memory budget
For Qwen3-0.6B in BF16 on a single request:
| item | size |
|---|---|
| weights (596M × 2 bytes, tied head counted once) | 1.19 GB |
| KV cache at 4,096 tokens (112 KiB per token) | 0.46 GB |
| activations during a 4,096-token prefill (largest: logits 4,096 × 151,936 in FP32 if materialized) | up to 2.5 GB |
| CUDA context, allocator caching, kernels’ workspaces | 0.3-1 GB |
The logits row is a surprise to most people. Prefill only needs the last position’s logits, so computing all 4,096 × 151,936 of them wastes memory and a large matmul. Production engines slice the hidden states to the last position before the head. That’s a one-line stretch exercise below.
Measure against the ceiling
Use your Chapter 10 tools. Predict, then measure:
- Decode ceiling:
decode_ceiling(1.19e9, bandwidth), about 230 tokens/s on DGX Spark and about 840 on an RTX 4090. - Achieved:
ms_per_output_tokenfrom the engine.
Engine v1 typically reaches a small fraction of the ceiling on a GPU. A 0.6B model’s decode step is hundreds of small kernels, each waiting on Python, plus one int(token) synchronization per step, so the GPU idles between kernels. A profile (torch.profiler) shows the gaps clearly. Chapter 19 removes them.
The Rust engine does the same, on the CPU
The Rust track loads the same snapshot and generates on the CPU. Its weights stay in BF16 and are widened inside a multi-threaded mat-vec, which is a direct application of the memory-bound decode model:
#![allow(unused)]
fn main() {
/// Process one token at position cache.len and return the next-token logits.
pub fn forward_token(&self, token: u32, cache: &mut KvCache) -> Vec<f32> {
let c = &self.cfg;
let pos = cache.len;
let d = c.head_dim;
let mut x = self.embed.row(token as usize);
for (li, layer) in self.layers.iter().enumerate() {
let h = rmsnorm(&x, &layer.input_norm, c.eps);
let (mut q, mut k, v) = (layer.q.matvec(&h), layer.k.matvec(&h), layer.v.matvec(&h));
for head in q.chunks_mut(d) { // per-head RMSNorm, then rotate
head.copy_from_slice(&rmsnorm(head, &layer.q_norm, c.eps));
rope(head, pos, c.theta);
}
for head in k.chunks_mut(d) {
head.copy_from_slice(&rmsnorm(head, &layer.k_norm, c.eps));
rope(head, pos, c.theta);
}
let (keys, values) = cache.append(li, &k, &v);
let attended = causal_attention(&q, keys, values, 1, pos + 1, c.heads, c.kv_heads, d);
for (xi, oi) in x.iter_mut().zip(layer.o.matvec(&attended)) {
*xi += oi;
}
let h = rmsnorm(&x, &layer.post_norm, c.eps);
let gated: Vec<f32> = layer.gate.matvec(&h).iter().zip(layer.up.matvec(&h)).map(|(g, u)| silu(*g) * u).collect();
for (xi, mi) in x.iter_mut().zip(layer.down.matvec(&gated)) {
*xi += mi;
}
}
cache.len += 1;
let h = rmsnorm(&x, &self.norm, c.eps);
self.head.as_ref().unwrap_or(&self.embed).matvec(&h)
}
}
#![allow(unused)]
fn main() {
/// Prefill the prompt token by token (simple, not fast), then decode `new_tokens` tokens.
pub fn generate(model: &Qwen3, prompt: &[u32], new_tokens: usize, policy: &crate::sampling::Policy,
rng: &mut crate::rng::Rng, stop: &[u32]) -> Vec<u32> {
let mut cache = model.new_cache(prompt.len() + new_tokens);
let mut logits = vec![];
for &t in prompt {
logits = model.forward_token(t, &mut cache);
}
let mut out = vec![];
for _ in 0..new_tokens {
let next = crate::sampling::sample(&logits, policy, rng) as u32;
out.push(next);
if stop.contains(&next) || out.len() == new_tokens {
break;
}
logits = model.forward_token(next, &mut cache);
}
out
}
}
cd rust
IDS=$(python tokenize_ids.py encode --model-dir ../models/Qwen3-0.6B "Why is the sky blue?")
OUT=$(cargo run --release -- generate --model-dir ../models/Qwen3-0.6B --ids "$IDS" --new-tokens 48)
python tokenize_ids.py decode --model-dir ../models/Qwen3-0.6B "$OUT"
Its cargo test reproduces the Python reference’s logits at every position to 1e-4, for both F32 and BF16 checkpoints. Compare its CPU tokens/s with your machine’s memory bandwidth divided by 1.19 GB.
Build it
Engine milestone 18: engine v1. Implement load_qwen3 in engine/loaders.py and LLM.stream in engine/engine.py (from_pretrained, encode, generate and the manifest are provided).
pytest tests/test_ch18_engine.py
python run.py chat --impl engine --model-dir models/Qwen3-0.6B
The tests check that your engine’s greedy output equals sampling.generate, that stop tokens and streaming behave, that seeded sampling reproduces, and that the manifest records the settings. Then run rung 3 of the ladder on the real checkpoint and record the maximum logit difference in your notes.
Stretch exercises
- ★ Compute logits only for the last position during prefill (
return_hidden=True, slice, thenlm_head). Measure the memory saved on a 4,096-token prompt. Where: the prefill model call inLLM.streaminengine/engine.py. - ★★ Implement rung 4: a function that registers hooks on every decoder layer of both models and prints the first layer whose output differs by more than a tolerance. Where:
experiments/ch18.py(create it), usingengine.qwen3and the corresponding Transformers model. - ★★ Add a
--thinkingflag torun.py chatand compare the answer quality and token count of thinking and non-thinking modes on five reasoning questions. Where: the argument parser andcmd_chatinrun.py; pass it asthinking=...toLLM.encode(which forwards it to the chat template). - ★★★ Load Qwen’s
tokenizer.jsonwith your own BPE implementation: parse the vocabulary and merges, implement the byte-to-unicode mapping and the pre-tokenizer regex (use the third-partyregexpackage), and check that you reproduceAutoTokenizer’s IDs on 1,000 lines of text. Where: add a checkpoint-tokenizer loader inengine/tokenizer.py; compare it inexperiments/ch18.py(create it).
Check your understanding
- Why should you compare intermediate layer outputs before generated text?
- Why can a model that fits in memory still fail to allocate for a long prompt?
- Why are “tokens emitted” and “tokens processed into the cache” different counts?
- Why does
to_emptyrequire you to re-tie the embedding and head? - What two stop tokens does a Qwen3 chat need, and what happens with only one?
Going deeper
- BALLM Chapter 5 §§5.4-5.5 for the GPT-2 version of loading pretrained weights; Raschka’s
ch05/11_qwen3notebook loads real Qwen3 weights into a from-scratch model. - PMPP §20.4 and §20.6 for cache sizing and decode intensity.
- The Qwen3 model card (chat template, thinking mode, recommended sampling settings: temperature 0.7, top-p 0.8 for non-thinking mode).
- vLLM’s
LLMclass and SGLang’sEngine: the production versions of this chapter’s interface.
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.
20. Quantization
In this chapter
- Why fewer bits per weight means faster decode, with the arithmetic to predict how much.
- Symmetric integer quantization derived from scratch: scales, codes, the error bound and granularity (per tensor, per row, per group).
- Packing two 4-bit values per byte, and the metadata cost that "4-bit" leaves out.
- Calibration: choosing clipping by output error, not weight error.
- A W4A16 Triton kernel that dequantizes in registers, and the landscape beyond round-to-nearest: GPTQ, AWQ, SmoothQuant, FP8, NF4 and KV-cache quantization.
You will build
engine/quant.py (INT8/INT4 groupwise quantization, packing, QuantLinear) and a W4A16 matmul kernel in engine/kernels/triton_quant.py. You'll quantize your Qwen3 engine and measure size, quality and speed.
Time: 5-7 hours. GPU: optional (needed for the speed measurements and the bitsandbytes workflow).
Fewer bytes, faster tokens
Decode reads every weight once per token, so its ceiling is bandwidth ÷ bytes per token (Chapter 10). Halve the bytes and you double the ceiling. For an 8B-parameter model:
| format | bits per weight (incl. scales) | weights | ceiling at 273 GB/s (DGX Spark) | at 1,008 GB/s (RTX 4090) |
|---|---|---|---|---|
| BF16 | 16 | 16.4 GB | 17 tokens/s | 61 tokens/s |
| INT8, per row | ~8 | 8.2 GB | 33 tokens/s | 123 tokens/s |
| INT4, groups of 128 | ~4.1 | 4.3 GB | 63 tokens/s | 234 tokens/s |
Quantization also decides what fits. A 32B model in BF16 needs 64 GB, more than an RTX 4090’s 24 GB; in INT4 it needs about 17 GB. For the sparse models of Part VII, with hundreds of billions of parameters, low-bit weights are the difference between running and not running.
The price is error: every weight becomes the nearest value on a coarse grid. The whole subject is about spending a fixed bit budget where it matters most.
Symmetric quantization, from scratch
Take a group of weights $w_1, \dots, w_n$ and a bit width $b$. A signed $b$-bit integer can hold $-2^{b-1} \dots 2^{b-1}-1$; we use the symmetric range $[-Q, Q]$ with $Q = 2^{b-1}-1$ (127 for INT8, 7 for INT4) so that zero is exact and the grid is symmetric. Map the largest magnitude $a = \max_i |w_i|$ to $Q$:
$$ s = \frac{a}{Q}, \qquad q_i = \operatorname{clamp}!\Big(\operatorname{round}\big(w_i / s\big), -Q, Q\Big), \qquad \hat w_i = s, q_i . $$
The integers $q_i$ (the codes) are stored, along with one floating-point scale $s$ per group. Without clamping (the range covers every weight), rounding moves each value by at most half a step:
$$ |\hat w_i - w_i| \le \tfrac{s}{2} = \frac{a}{2Q}. $$
Worked example. $w = [-2, -1, 0, 1, 2]$ at 4 bits: $Q = 7$, $s = 2/7 \approx 0.2857$. Dividing gives $[-7, -3.5, 0, 3.5, 7]$. PyTorch rounds halves to even, so the codes are $[-7, -4, 0, 4, 7]$ and $\hat w = [-2, -1.143, 0, 1.143, 2]$. The error on $\pm 1$ is $0.143 = s/2$, exactly at the bound. (Rounding conventions differ between libraries, which is one reason to test against the specific runtime you’ll use.)
Granularity: who shares a scale
The bound says the error is proportional to the largest magnitude in the group. One large weight coarsens the grid for everything that shares its scale. So the choice of groups matters as much as the bit width:
- Per tensor: one scale for the whole matrix. A single outlier anywhere ruins it. Fine for INT8 activations with care, poor for low-bit weights.
- Per row (per output channel): one scale per row of the
[out, in]weight. Standard for INT8 weights. - Per group: one scale per
gconsecutive inputs within a row, typically $g = 32$-$128$. Standard for INT4.
def quantize_int8_rows(weight):
"""One scale per output row: scale = max|w| / 127, q = round(w / scale). (Your engine: Chapter 20)"""
maximum = weight.float().abs().amax(dim=-1, keepdim=True)
scale = torch.where(maximum > 0, maximum / 127.0, torch.ones_like(maximum))
return (weight.float() / scale).round().clamp(-127, 127).to(torch.int8), scale
def dequantize_int8_rows(q, scale):
return q.float() * scale
def quantize_groupwise(weight, bits=4, group_size=128, clip=1.0):
"""Symmetric signed codes with one scale per (row, group of group_size inputs). (Your engine: Chapter 20)
weight [O, I] -> codes int8 [O, I] in [-(2^(b-1)-1), 2^(b-1)-1], scales [O, ceil(I/g)].
clip < 1 shrinks the range: large outliers saturate, everything else gets finer steps.
"""
if weight.ndim != 2 or bits not in (2, 3, 4, 8) or group_size < 1 or not 0 < clip <= 1:
raise ValueError("Expected a matrix, 2/3/4/8 bits, positive group size and 0 < clip <= 1")
out_features, in_features = weight.shape
groups = -(-in_features // group_size)
padded = F.pad(weight.float(), (0, groups * group_size - in_features))
blocks = padded.view(out_features, groups, group_size)
limit = 2 ** (bits - 1) - 1
maximum = blocks.abs().amax(-1, keepdim=True) * clip
scales = torch.where(maximum > 0, maximum / limit, torch.ones_like(maximum))
codes = (blocks / scales).round().clamp(-limit, limit).to(torch.int8)
return codes.view(out_features, -1)[:, :in_features], scales.squeeze(-1)
def dequantize_groupwise(codes, scales, group_size):
out_features, in_features = codes.shape
expanded = scales.repeat_interleave(group_size, dim=1)[:, :in_features]
return codes.float() * expanded
Two edge cases the code handles: a group of all zeros gets scale 1 (any scale works, and 0/0 doesn’t); the last group of a row may be short when in_features isn’t a multiple of g, so the row is padded for the computation and the padding is dropped.
The metadata the name leaves out
“4-bit” counts the codes only. With groups of $g$ inputs and a 16-bit scale per group, the true cost is $4 + 16/g$ bits per weight: 4.125 for $g = 128$, 4.5 for $g = 32$. Asymmetric schemes add a zero point per group. Layers usually left in high precision (embeddings, the LM head, MoE routers, norms) add more. For Qwen3-0.6B, whose tied embedding is 26% of the parameters, “4-bit” weights come to well over 4 bits on average. Always report measured bytes.
Packing
PyTorch has no 4-bit dtype, so an int8 tensor holding values in $[-7, 7]$ still costs 8 bits each. Pack two codes per byte: the first in the low nibble, the second in the high nibble, each in 4-bit two’s complement (values 8-15 of a nibble mean -8 to -1):
def pack_int4(codes):
"""Two signed 4-bit values per byte, two's complement, first value in the low nibble."""
values = codes.to(torch.int16)
if bool(((values < -8) | (values > 7)).any()):
raise ValueError("INT4 values must lie in [-8, 7]")
flat = values.flatten()
if flat.numel() % 2:
flat = torch.cat((flat, flat.new_zeros(1)))
nibbles = flat & 0xF
return (nibbles[0::2] | (nibbles[1::2] << 4)).to(torch.uint8)
def unpack_int4(packed, shape):
"""Inverse of pack_int4: nibbles 8..15 decode to -8..-1."""
nibbles = torch.stack((packed & 0xF, packed >> 4), dim=-1).flatten().to(torch.int16)
signed = torch.where(nibbles >= 8, nibbles - 16, nibbles)
count = 1
for size in shape:
count *= size
return signed[:count].to(torch.int8).reshape(shape)
inline std::pair<std::vector<int8_t>, std::vector<float>> quantize_row(const std::vector<float>& w, int bits, size_t group) {
float limit = float((1 << (bits - 1)) - 1);
std::vector<int8_t> codes;
std::vector<float> scales;
for (size_t g = 0; g < w.size(); g += group) {
size_t end = std::min(w.size(), g + group);
float m = 0;
for (size_t i = g; i < end; ++i) m = std::max(m, std::fabs(w[i]));
float scale = m > 0 ? m / limit : 1.f;
scales.push_back(scale);
for (size_t i = g; i < end; ++i) codes.push_back(int8_t(std::clamp(std::round(w[i] / scale), -limit, limit)));
}
return {codes, scales};
}
inline std::vector<uint8_t> pack_int4(const std::vector<int8_t>& c) { // first value in the low nibble
std::vector<uint8_t> out((c.size() + 1) / 2, 0);
for (size_t i = 0; i < c.size(); ++i) out[i / 2] |= uint8_t((c[i] & 0x0f) << (4 * (i % 2)));
return out;
}
inline int8_t unpack_nibble(uint8_t byte, int high) {
int n = (byte >> (4 * high)) & 0x0f;
return int8_t(n >= 8 ? n - 16 : n); // two's complement
}
#![allow(unused)]
fn main() {
/// Quantize one row in groups of `group` values: scale = max|w| / (2^(bits-1) - 1).
pub fn quantize_row(w: &[f32], bits: u32, group: usize) -> (Vec<i8>, Vec<f32>) {
let limit = ((1 << (bits - 1)) - 1) as f32;
let mut codes = Vec::with_capacity(w.len());
let mut scales = vec![];
for chunk in w.chunks(group) {
let max = chunk.iter().fold(0.0f32, |m, v| m.max(v.abs()));
let scale = if max > 0.0 { max / limit } else { 1.0 };
scales.push(scale);
codes.extend(chunk.iter().map(|v| (v / scale).round().clamp(-limit, limit) as i8));
}
(codes, scales)
}
/// Two signed 4-bit codes per byte, first code in the low nibble (two's complement).
pub fn pack_int4(codes: &[i8]) -> Vec<u8> {
codes.chunks(2).map(|p| {
let lo = (p[0] as u8) & 0x0f;
let hi = (*p.get(1).unwrap_or(&0) as u8) & 0x0f;
lo | (hi << 4)
}).collect()
}
pub fn unpack_int4(packed: &[u8], count: usize) -> Vec<i8> {
let nibble = |n: u8| if n >= 8 { n as i8 - 16 } else { n as i8 };
packed.iter().flat_map(|b| [nibble(b & 0x0f), nibble(b >> 4)]).take(count).collect()
}
}
Packed layouts are a contract between the quantizer and the kernel, and every library chooses differently: nibble order, interleaving for vectorized loads, and transposition for tensor-core fragment layouts. A GPTQ checkpoint packs eight 4-bit values into an int32; Marlin kernels reorder them again for fast loading. A checkpoint is only usable by a runtime that implements its exact format.
Calibration: optimize what you care about
Clipping (shrinking $a$ by a factor clip < 1) saturates the largest weights but gives everything else a finer grid. Whether that helps depends on how the weights are used. A large weight attached to an input feature that’s almost always near zero barely affects the output, so it’s worth clipping; a large weight on an active feature isn’t.
So choose quantization parameters by output error on representative inputs (calibration data), not by weight error:
def choose_clip(weight, calibration_inputs, bits=4, group_size=128, candidates=(1.0, 0.95, 0.9, 0.85, 0.8, 0.7)):
"""Pick the clip that minimizes OUTPUT error on calibration activations, not weight error."""
reference = calibration_inputs.float() @ weight.float().T
best = None
for clip in candidates:
codes, scales = quantize_groupwise(weight, bits, group_size, clip)
approx = calibration_inputs.float() @ dequantize_groupwise(codes, scales, group_size).T
error = ((approx - reference).norm() / reference.norm().clamp_min(1e-12)).item()
if best is None or error < best[0]:
best = (error, clip)
return best[1], best[0]
The lab plants exactly that situation: a 16 × 128 weight whose input feature 0 has weights 12× larger than the rest, while that feature is nearly silent in the activations. It picks the clip on 256 calibration rows and reports error on 256 different held-out rows:
python model_workflows.py quantize
{"bits": 8, "group": 128, "clip": 0.95, "weight_error": 0.03627, "held_out_output_error": 0.02119, ...}
{"bits": 8, "group": 32, "clip": 1.0, "weight_error": 0.00874, "held_out_output_error": 0.01161, ...}
{"bits": 4, "group": 128, "clip": 0.7, "weight_error": 0.28655, "held_out_output_error": 0.28023, ...}
{"bits": 4, "group": 64, "clip": 0.7, "weight_error": 0.25714, "held_out_output_error": 0.22791, ...}
{"bits": 4, "group": 32, "clip": 0.85, "weight_error": 0.17493, "held_out_output_error": 0.19074, ...}
Three lessons: 4-bit chose aggressive clipping because the outlier is harmless to clip; smaller groups help most at low bit widths; and at group 32 the output error is larger than the weight error, while at group 64 it’s smaller. Weight error is not the objective. Keep calibration and evaluation data separate, or you’ll fool yourself.
A quantized linear layer
class QuantLinear(nn.Module):
"""Weight-only quantized linear layer: integer weights in memory, higher-precision math. (Your engine: Chapter 20)
The reference forward dequantizes the whole matrix each call, which saves memory but not
time. A real W4A16 kernel (kernels/triton_quant.py) dequantizes tiles in registers so the
weights cross the memory bus as 4-bit values.
"""
def __init__(self, linear, bits=4, group_size=128, clip=1.0):
super().__init__()
self.in_features, self.out_features = linear.in_features, linear.out_features
self.bits, self.group_size = bits, group_size
codes, scales = quantize_groupwise(linear.weight.data, bits, group_size, clip)
if bits == 4:
self.register_buffer("packed", pack_int4(codes))
else:
self.register_buffer("packed", codes)
self.register_buffer("scales", scales.to(linear.weight.dtype))
self.bias = None if linear.bias is None else nn.Parameter(linear.bias.data.clone(), requires_grad=False)
self.compute_dtype = linear.weight.dtype
def codes(self):
shape = (self.out_features, self.in_features)
return unpack_int4(self.packed, shape) if self.bits == 4 else self.packed
def dequantized_weight(self):
"""(Your engine: Chapter 20)"""
return dequantize_groupwise(self.codes(), self.scales.float(), self.group_size).to(self.compute_dtype)
def forward(self, x):
"""(Your engine: Chapter 20)"""
return F.linear(x, self.dequantized_weight(), self.bias)
def storage_bytes(self):
return sum(t.numel() * t.element_size() for t in (self.packed, self.scales))
def quantize_model(model, bits=4, group_size=128, skip=("lm_head", "head", "gate", "router", "shared_expert_gate")):
"""Replace every nn.Linear with a QuantLinear, in place, unless a component of its name is in
`skip` (output heads and MoE routers are small and sensitive, so they stay in high precision)."""
replaced = 0
for name, module in list(model.named_modules()):
for child_name, child in list(module.named_children()):
full = f"{name}.{child_name}" if name else child_name
if isinstance(child, nn.Linear) and not set(full.split(".")) & set(skip):
setattr(module, child_name, QuantLinear(child, bits, group_size))
replaced += 1
return replaced
quantize_model replaces every nn.Linear in place, except layers whose names mark them as sensitive: the LM head (its errors land directly on the logits) and MoE routers (small, and a flipped expert choice changes the computation discretely). On the tiny Qwen3 test model in FP32:
python run.py quant --bits 4
{"bits": 4, "bytes_before": 15739904, "bytes_after": 3124224,
"logit_error": {"max_abs": 0.1275, "mean_abs": 0.0178, "relative_l2": 0.0396}, "same_top1": 0.953}
{"bits": 8, "bytes_before": 15739904, "bytes_after": 4959232,
"logit_error": {"max_abs": 0.0075, "mean_abs": 0.00098, "relative_l2": 0.0022}, "same_top1": 1.0}
(The 8-bit line comes from --bits 8; both runs are from a laptop CPU.) The 8-bit model’s size isn’t a quarter of FP32’s because the embedding stays unquantized. At 4 bits, 95% of positions keep the same top token, with 4% relative logit error. Real models are more robust than this tiny random one, but measure, don’t assume.
Speed needs a kernel
QuantLinear.forward dequantizes the whole matrix to BF16 and calls an ordinary matmul. That saves memory but not time: the full-precision weights are materialized and read again every call, so decode is slower than BF16. The speedup only appears when a kernel reads the packed bytes and dequantizes them in registers, so weights cross the memory bus at 4 bits. That’s a W4A16 kernel: 4-bit weights, 16-bit activations, 16/32-bit math.
@triton.jit
def w4a16_kernel(x_ptr, packed_ptr, scales_ptr, y_ptr, M, N, K, groups,
GROUP: tl.constexpr, BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr):
"""(Your engine: Chapter 20)"""
pid_m, pid_n = tl.program_id(0), tl.program_id(1)
rm = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
rn = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
half = tl.arange(0, BLOCK_K // 2) # byte index inside the K tile
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)
x = tl.load(x_ptr + rm[:, None] * K + rk[None, :], mask=(rm[:, None] < M) & (rk[None, :] < K), other=0.0)
# Packed row n holds K/2 bytes; byte j carries code 2j (low nibble) and 2j+1 (high nibble).
byte_cols = k0 // 2 + half
packed = tl.load(packed_ptr + rn[:, None] * (K // 2) + byte_cols[None, :],
mask=(rn[:, None] < N) & (byte_cols[None, :] < K // 2), other=0).to(tl.int32)
low, high = packed & 0xF, (packed >> 4) & 0xF
codes = tl.reshape(tl.join(low, high), (BLOCK_N, BLOCK_K)) # interleave back to k order
codes = tl.where(codes >= 8, codes - 16, codes) # two's complement nibble
scale = tl.load(scales_ptr + rn * groups + k0 // GROUP, mask=rn < N, other=0.0).to(tl.float32)
w = codes.to(tl.float32) * scale[:, None] # [BLOCK_N, BLOCK_K]
acc += tl.dot(x.to(tl.float32), tl.trans(w), input_precision="ieee") # FP32, not TF32
tl.store(y_ptr + rm[:, None] * N + rn[None, :], acc.to(y_ptr.dtype.element_ty),
mask=(rm[:, None] < M) & (rn[None, :] < N))
def w4a16_matmul(x, packed, scales, group_size, out_features, block_m=16, block_n=32, block_k=32):
"""x [M, K]; packed uint8 [N * K / 2] (row-major, from quant.pack_int4); scales [N, K / group]."""
check_device(x, packed, scales)
m, k = x.shape
if k % 2 or group_size % block_k or k % block_k:
raise ValueError("K must be a multiple of BLOCK_K, and BLOCK_K must divide the group size")
y = torch.empty((m, out_features), device=x.device, dtype=x.dtype)
grid = (triton.cdiv(m, block_m), triton.cdiv(out_features, block_n))
w4a16_kernel[grid](x.contiguous(), packed, scales.contiguous(), y, m, out_features, k, scales.shape[1],
GROUP=group_size, BLOCK_M=block_m, BLOCK_N=block_n, BLOCK_K=block_k)
return y
It’s the Chapter 14 matmul with the B-tile loaded differently. Each program loads BLOCK_K / 2 bytes per weight row, splits each byte into its two nibbles, interleaves them back into $k$ order with tl.join and a reshape, sign-extends, and multiplies by the group’s scale. Requiring BLOCK_K to divide the group size means one scale per row per K-tile. As in Chapter 15, input_precision="ieee" keeps the FP32 tl.dot exact on the GPU instead of TF32.
What production W4A16 kernels (Marlin, Machete, AWQ’s GEMV, ExLlama) add: weight layouts pre-shuffled so a 128-bit load fills a tensor-core fragment directly, dequantization via bit tricks instead of shifts and multiplies, and split-K for the batch-1 decode shape. At batch 1 a good INT4 kernel approaches 3.5-3.9× BF16 decode speed; at large batches the kernel becomes compute-bound and the advantage shrinks, because the dequantization arithmetic is extra work.
Beyond round-to-nearest
Round-to-nearest with good groups is a strong baseline at 8 bits and a reasonable one at 4. The research methods address different parts of the problem:
| method | idea | needs |
|---|---|---|
| LLM.int8() (Dettmers et al., 2022) | A few activation features have huge outliers; compute those columns in FP16 and the rest in INT8 | mixed-precision kernels |
| SmoothQuant (Xiao et al., 2022) | Move activation outliers into the weights with a per-channel scale: $XW^\top = (X S^{-1})(W S)^\top$, exact before rounding | calibration; enables W8A8 |
| GPTQ (Frantar et al., 2022) | Quantize weights one column at a time and adjust the remaining columns to compensate the error, using second-order statistics of calibration inputs | calibration; a few minutes per model |
| AWQ (Lin et al., 2023) | Protect the ~1% of weight channels attached to large activations by scaling them up before quantization | calibration; very fast |
| NF4 (QLoRA, Dettmers et al., 2023) | A 16-value non-uniform grid matched to normally distributed weights; a table lookup to dequantize | used for frozen fine-tuning bases |
| FP8 (E4M3, E5M2) | 8-bit floating point with per-tensor or per-block scales; native tensor-core support on Hopper and later | W8A8 with almost no accuracy loss |
| MXFP4 / NVFP4 | 4-bit floats with a shared scale per block of 16-32 values, native on Blackwell | hardware support |
| QAT | Simulate quantization during training so the model adapts to it | training compute |
The scaling identity behind SmoothQuant and AWQ is worth remembering: for a diagonal $S$, $XW^\top = (XS^{-1})(WS)^\top$ exactly. It moves magnitude between activations and weights without changing the function, so you can choose where the outliers live before you round.
Quantizing the KV cache
At long contexts, the KV cache dominates memory traffic (Chapter 16). Quantizing it to FP8 or INT8, with a scale per head or per token, halves the traffic of attention during decode. Keys are more sensitive than values (errors pass through the softmax’s exponential), and keys often have outlier channels, so KV quantization usually uses per-channel scales for K and per-token scales for V (KIVI, 2024). Chapter 16’s last stretch exercise implements a simple version.
Quantize a real checkpoint
For Qwen3-0.6B, the companion script compares BF16 against bitsandbytes NF4 or LLM.int8() on held-out instruction data, then saves and reloads the quantized model to confirm the output survives the round trip:
uv pip install -r workflow-requirements.txt # transformers, safetensors, peft
uv pip install 'bitsandbytes>=0.48' # check platform support first
python quantize_checkpoint.py --model-dir models/Qwen3-0.6B --bits 4 \
--valid-file data/sft-valid.jsonl --output runs/qwen3-nf4
It reports the response loss of both models, the KL divergence between their next-token distributions, and parameter bytes. These are what to look at when judging a quantization: loss and KL on your own data first, then task scores and some generations read by a human.
Warning
bitsandbytes ships compiled CUDA binaries. If its CUDA version doesn’t match PyTorch’s, import fails or 4-bit layers silently fall back. The bitsandbytes installation guide explains
BNB_CUDA_VERSION; Chapter 22 shows how to diagnose it. The exported checkpoint runs in Transformers + bitsandbytes, not in your engine’s loader, which expects ordinary floating-point tensors.
An evaluation protocol
- Fix the evaluation set, token IDs and decoding settings before quantizing anything.
- Report held-out loss (or perplexity) and the KL from the original model, then task metrics.
- Measure memory (weights, KV cache, peak) and speed separately for prefill and decode, with batch size, lengths, GPU, backend and warmup recorded.
- Re-evaluate after every transformation that follows (merging adapters, editing, exporting): quantization doesn’t commute with them.
Build it
Engine milestone 20: quantization. Implement quantize_int8_rows, quantize_groupwise, QuantLinear.dequantized_weight and QuantLinear.forward in engine/quant.py (packing, clip search and quantize_model are provided), and w4a16_kernel in engine/kernels/triton_quant.py.
pytest tests/test_ch20_quant.py
python run.py quant --impl engine --bits 4
The tests check the half-step error bound for every bit width, short tail groups, zero groups, INT4 packing at both ends of the range, that a quantized Qwen3 shrinks and stays close, and your Triton kernel against dequantize-then-matmul. Then quantize your Qwen3-0.6B engine to 8 and 4 bits and record held-out loss and size in your notes.
Stretch exercises
- ★ Plot held-out output error against bits per weight (including scales) for group sizes 16-256 on a real Qwen3-0.6B layer. Where:
experiments/ch20.py(create it), usingengine.quant.quantize_groupwiseanddequantize_groupwise. - ★★ Make
QuantLinearcall your W4A16 kernel on CUDA, and measure decode tokens/s for Qwen3-0.6B against BF16 in graph mode (Chapter 19). Where:QuantLinear.forwardinengine/quant.py, callingengine.kernels.triton_quant.w4a16_matmul. - ★★ Implement AWQ’s core: for one linear layer, search a per-input-channel scale $s_j = \bar{|x_j|}^{\alpha}$ over $\alpha \in [0, 1]$ that minimizes calibration output error after 4-bit quantization of $W \operatorname{diag}(s)$. Where: add a calibration/search helper in
engine/quant.py; Chapter 39’sengine/formats/awq.pyprovides the later integration. - ★★★ Implement GPTQ for one layer: accumulate $H = 2X^\top X$ from calibration inputs, then quantize column by column, spreading each column’s error over the remaining columns with the Cholesky factor of $H^{-1}$. Compare with round-to-nearest at 3 and 4 bits. Where: add a one-layer GPTQ helper in
engine/quant.py; Chapter 39’sengine/formats/gptq.pyprovides the later integration.
Check your understanding
- Why does halving the bytes per weight roughly double the decode ceiling, but barely change large-batch prefill speed?
- Why do smaller groups reduce error, and what do they cost?
- Why can a clip below 1.0 reduce output error while increasing weight error?
- Why doesn’t dequantize-then-matmul make decode faster?
- Why are the LM head and MoE routers often left unquantized?
Going deeper
- GPU Mode L7 (Advanced quantization: weight-only and dynamic quantization in torchao, with Triton kernels), L33 (BitBLAS: mixed-precision kernels for arbitrary bit widths), L30 (Quantized training).
- PMPP Appendix A (Numerical considerations): the low-precision floating-point formats (FP16, BF16, FP8, block-scaled formats) and how GPUs implement them.
- Papers: LLM.int8() (2208.07339), SmoothQuant (2211.10438), GPTQ (2210.17323), AWQ (2306.00978), QLoRA/NF4 (2305.14314), KIVI (2402.02750); Micikevicius et al., FP8 Formats for Deep Learning (2022).
gpt-fast’squantize.py(int8 and int4 weight-only with GPTQ in a few hundred lines); the Marlin kernel write-up; Maarten Grootendorst’s A Visual Guide to Quantization.
21. Fine-tuning: classifiers and instruction following
In this chapter
- What fine-tuning changes in a pretrained model, and what it can't add.
- Classification fine-tuning: replacing the vocabulary head and reading the last real token.
- Instruction fine-tuning: prompt templates, response-only loss, label shifting and padding, done exactly right.
- Fine-tuning Qwen3-0.6B, the memory full fine-tuning really needs, and how to evaluate the result.
You will build
instruction_batch in engine/lora.py and last_token_logits in engine/sft.py. You'll train a classifier, instruction-tune a GPT on 1,100 examples, and fine-tune a real Qwen3 checkpoint.
Time: 5-7 hours. GPU: recommended for GPT-2 and Qwen3 (the from-scratch demos run on a CPU in about a minute).
What fine-tuning changes
A pretrained model has learned language, facts and patterns from trillions of tokens, but its only skill is continuing text. Ask a base model “What is the capital of France?” and it may continue with three more quiz questions. Fine-tuning keeps training the same weights on a smaller, targeted dataset so the model does a specific job.
| goal | data | objective | what changes |
|---|---|---|---|
| continued pretraining | domain text | next-token loss on every token | knowledge and style of a domain |
| classification | labelled texts | cross-entropy over classes, from one position | a new output head, top layers |
| instruction tuning (SFT) | prompt-response pairs | next-token loss on response tokens only | format and behavior: answering, stopping |
| preference tuning (DPO, RLHF) | prompts with better and worse responses | a preference objective | which of several plausible answers is preferred |
The important intuition: fine-tuning mostly teaches format and behavior, and draws on knowledge already in the weights. A thousand examples can teach a model to answer in one sentence and stop. They can’t teach it chemistry. You’ll see this directly below, when a model trained from scratch learns the answer format perfectly and the facts not at all.
BALLM Chapters 6 and 7 build both classification and instruction tuning on GPT-2; this chapter follows them closely, then moves the same ideas to Qwen3.
Classification: replace the head
A language model ends in a head that maps the final hidden state to vocabulary scores. For a classifier, replace it with a head that maps to class scores: for spam detection, nn.Linear(width, 2) instead of nn.Linear(width, 50257).
Which position’s hidden state should be classified? Under a causal mask, position $t$ has seen tokens $0 \dots t$ only. The last token is the only one that has seen the whole message. With right-padded batches, the last real token is at length - 1, not at -1:
def replace_head(model, num_classes, train_last_blocks=1):
"""Freeze the model, swap the vocabulary head for a num_classes head, and unfreeze the
last few blocks and the final norm (BALLM Chapter 6's recipe)."""
for parameter in model.parameters():
parameter.requires_grad_(False)
width = model.head.in_features
p = model.head.weight
model.head = nn.Linear(width, num_classes, device=p.device, dtype=p.dtype) # new, trainable
for block in list(model.blocks)[len(model.blocks) - train_last_blocks:]:
block.requires_grad_(True)
model.norm.requires_grad_(True)
return model
def last_token_logits(model, ids, lengths):
"""Class scores read at each sequence's last REAL token. (Your engine: Chapter 21)
Under a causal mask only the last position has seen the whole text, and with right padding
that position is lengths - 1, not -1.
"""
logits = model(ids)
return logits[torch.arange(ids.shape[0], device=ids.device), lengths - 1]
def pad_sequences(sequences, pad_id, device="cpu"):
width = max(len(s) for s in sequences)
ids = torch.full((len(sequences), width), pad_id, dtype=torch.long)
for row, seq in enumerate(sequences):
ids[row, :len(seq)] = torch.tensor(seq)
return ids.to(device), torch.tensor([len(s) for s in sequences], device=device)
@torch.no_grad()
def accuracy(model, sequences, labels, pad_id, batch_size=32):
model.eval()
device = next(model.parameters()).device
correct = 0
for i in range(0, len(sequences), batch_size):
ids, lengths = pad_sequences(sequences[i:i + batch_size], pad_id, device)
predicted = last_token_logits(model, ids, lengths).argmax(-1).cpu()
correct += int((predicted == torch.tensor(labels[i:i + batch_size])).sum())
return correct / len(sequences)
def train_classifier(model, sequences, labels, pad_id, steps, batch_size=16, lr=5e-4, seed=0):
device = next(model.parameters()).device
rng = random.Random(seed)
trainable = [p for p in model.parameters() if p.requires_grad]
optimizer = torch.optim.AdamW(trainable, lr=lr, weight_decay=0.1)
for _ in range(steps):
model.train()
index = rng.sample(range(len(sequences)), min(batch_size, len(sequences)))
ids, lengths = pad_sequences([sequences[i] for i in index], pad_id, device)
target = torch.tensor([labels[i] for i in index], device=device)
loss = F.cross_entropy(last_token_logits(model, ids, lengths), target)
optimizer.zero_grad(set_to_none=True)
loss.backward()
optimizer.step()
return loss.item()
replace_head follows BALLM’s recipe: freeze everything, add a fresh trainable head, and unfreeze the last transformer block and the final norm. Training only the top of the network is fast, needs little data, and keeps the pretrained features intact. Train more blocks when you have more data and the task differs more from pretraining.
Run it
Without any download, the demo trains a small GPT on a synthetic task that requires reading the whole sequence (does the token 7 appear anywhere?):
python run.py classify --steps 300
{"before": 0.49166666666666664}
{"train_accuracy": 1.0, "test_accuracy": 1.0}
That took 12 seconds on a laptop CPU. For a real task, download the SMS Spam Collection (5,574 messages, tab-separated label<TAB>text) and pass --data SMSSpamCollection. With pretrained GPT-2 small, frozen except the last block, BALLM reports 95.7% test accuracy after five epochs (BALLM §6.7); balance the classes first, as the book does, so that “always say ham” doesn’t score 87%.
Note
The demo trains its GPT from scratch with all blocks unfrozen (
train_last_blocks=layers), because there are no pretrained features to preserve. With GPT-2 weights, keep the default of one block.
Instruction tuning
Instruction tuning is still next-token prediction. What changes is the data: each example is a prompt and a desired response, joined by a template, and the loss covers only the response.
A template
The template marks where the instruction ends and the answer begins. BALLM uses the Alpaca format:
Below is an instruction that describes a task. Write a response that appropriately completes the request.
### Instruction:
Rewrite the sentence using a simile.
### Input:
The car is very fast.
### Response:
The car is as fast as lightning.<|endoftext|>
Chat models use their own chat template with special tokens instead (Qwen’s <|im_start|> and <|im_end|>, Chapter 4). Two rules matter more than the template you choose: training and inference must use exactly the same format, and the response must end with the stop token the engine stops on. A model never trained to emit <|endoftext|> never stops.
def format_prompt(entry):
"""The Alpaca-style template from BALLM Chapter 7. The response follows the final header."""
text = ("Below is an instruction that describes a task. "
"Write a response that appropriately completes the request."
f"\n\n### Instruction:\n{entry['instruction']}")
if entry.get("input"):
text += f"\n\n### Input:\n{entry['input']}"
return text + "\n\n### Response:\n"
def load_instructions(path):
with open(path, encoding="utf-8") as f:
return json.load(f)
def split_entries(entries, valid_fraction=0.1, test_fraction=0.1, seed=0):
"""Shuffle once, then split by item: every version of an item stays in one split."""
order = list(range(len(entries)))
random.Random(seed).shuffle(order)
n_test, n_valid = int(len(entries) * test_fraction), int(len(entries) * valid_fraction)
pick = lambda idx: [entries[i] for i in idx]
return pick(order[n_test + n_valid:]), pick(order[n_test:n_test + n_valid]), pick(order[:n_test])
def encode_examples(entries, encode, eos_id, max_length):
"""(prompt_ids, response_ids + [eos]) pairs. Over-long examples are dropped and counted, never
truncated: truncation would silently cut the response, the only part that is learned."""
examples, dropped = [], 0
for entry in entries:
prompt, response = encode(format_prompt(entry)), encode(entry["output"]) + [eos_id]
if len(prompt) + len(response) > max_length + 1: # inputs are one shorter than the pair
dropped += 1
continue
examples.append((prompt, response))
return examples, dropped
The dataset is BALLM’s 1,100 instruction-response pairs, in data/instruction-data.json. split_entries splits it into 935 training, 55 validation and 110 test examples before any processing, so no item leaks between splits. encode_examples drops over-long examples instead of truncating them; truncation would cut off the response, which is the only part that’s learned.
Response-only loss and the label shift
Should the model learn to predict the prompt tokens too? Usually not: the prompt is given at inference time, and learning to generate instructions wastes capacity. So prompt positions get the target $-100$, which F.cross_entropy ignores by default. Padding positions get $-100$ too.
The subtle part is the shift. For a prompt [10, 11] and a response [20, 21, EOS], the complete sequence is [10, 11, 20, 21, EOS]. Inputs are all tokens but the last; targets are all tokens but the first:
inputs 10 11 20 21
targets -100 20 21 EOS
The input 11, the last prompt token, must predict 20, the first response token. That’s the most important prediction in the example, and the one most often lost by masking on the wrong side of the shift.
def instruction_batch(examples, pad_id, device="cpu"):
"""examples: list of (prompt_ids, response_ids_ending_in_eos). (Your engine: Chapter 21)
Returns inputs x and targets y, both [B, L-1], already shifted for next-token loss.
Prompt and padding targets are -100 so cross_entropy ignores them; only response tokens
(including the final EOS) are learned.
"""
if not examples or any(not p or not r for p, r in examples):
raise ValueError("Each example needs a non-empty prompt and response")
width = max(len(p) + len(r) for p, r in examples)
ids = torch.full((len(examples), width), pad_id, dtype=torch.long)
labels = torch.full_like(ids, -100)
for row, (prompt, response) in enumerate(examples):
complete = list(prompt) + list(response)
ids[row, :len(complete)] = torch.tensor(complete)
labels[row, len(prompt):len(complete)] = torch.tensor(list(response))
return ids[:, :-1].to(device), labels[:, 1:].to(device)
Warning
Hugging Face models shift internally: you pass unshifted
input_idsandlabelsof the same length, and the model drops the first label. If you shift in your data pipeline and pass the result to a Hugging Face model, every target moves one token too far and the model learns to predict two tokens ahead. Your own GPT expects shifted targets;finetune.pyfor Qwen3 passes unshifted ones. Know which convention each model uses.
When the pad token equals the EOS token (common: GPT-2 has no pad token), mask padding by position, not by ID. Masking every occurrence of the EOS ID would also mask the real end of the response, and the model would never learn to stop.
The training loop
def response_loss(model, x, y):
"""Mean cross-entropy over response targets only; -100 marks prompt and padding."""
logits = model(x)
return F.cross_entropy(logits.reshape(-1, logits.shape[-1]), y.reshape(-1), ignore_index=-100)
@torch.no_grad()
def evaluate_responses(model, examples, pad_id, batch_size=8):
"""Average loss per response token (not per batch), so long and short answers weigh fairly."""
model.eval()
device = next(model.parameters()).device
total, count = 0.0, 0
for i in range(0, len(examples), batch_size):
x, y = instruction_batch(examples[i:i + batch_size], pad_id, device)
logits = model(x)
total += F.cross_entropy(logits.reshape(-1, logits.shape[-1]), y.reshape(-1),
ignore_index=-100, reduction="sum").item()
count += int((y != -100).sum())
return total / max(count, 1)
def finetune(model, train_examples, valid_examples, pad_id, steps, batch_size=8, lr=5e-5,
weight_decay=0.1, warmup=10, eval_every=25, seed=0, log=print):
"""AdamW on response-only loss, each batch padded only to its own longest example."""
device = next(model.parameters()).device
rng = random.Random(seed)
trainable = [p for p in model.parameters() if p.requires_grad]
optimizer = torch.optim.AdamW(trainable, lr=lr, weight_decay=weight_decay)
history = []
for step in range(steps):
scale = (step + 1) / warmup if step < warmup else 0.5 * (1 + math.cos(math.pi * (step - warmup) / max(1, steps - warmup)))
for group in optimizer.param_groups:
group["lr"] = lr * scale
model.train()
x, y = instruction_batch(rng.sample(train_examples, min(batch_size, len(train_examples))), pad_id, device)
loss = response_loss(model, x, y)
optimizer.zero_grad(set_to_none=True)
loss.backward()
torch.nn.utils.clip_grad_norm_(trainable, 1.0)
optimizer.step()
if step % eval_every == 0 or step + 1 == steps:
record = {"step": step + 1, "train_loss": round(loss.item(), 4),
"valid_loss": round(evaluate_responses(model, valid_examples, pad_id), 4)}
history.append(record)
if log:
log(json.dumps(record))
return history
Two details differ from Chapter 7’s loop:
- Dynamic padding. Each batch is padded to its own longest example, not to the dataset’s. Instruction lengths vary a lot, and padding to the maximum wastes compute.
- Token-weighted evaluation. Validation loss is the sum over response tokens divided by their count, not an average of batch averages, which would weight a three-word answer the same as a three-sentence one.
Run it
Without pretrained weights, run.py sft trains its own BPE and a 4-layer GPT from scratch on the 935 training examples only:
python run.py sft --steps 600 --width 192 --layers 4
On a laptop CPU, in 55 seconds:
{"train": 935, "valid": 55, "test": 110, "dropped_overlong": 0}
{"step": 1, "train_loss": 6.9229, "valid_loss": 6.8699}
{"step": 301, "train_loss": 3.9705, "valid_loss": 4.2244}
{"step": 600, "train_loss": 2.8521, "valid_loss": 3.8971}
### What is the contraction for 'it is'?
expected: The contraction for 'it is' is 'it's.'
model: The gas' is 'bua'.
### Convert the number 5 from decimal to binary.
expected: The binary equivalent of the decimal number 5 is 101.
model: The approximately 14 is 24.
### Rewrite the sentence using a simile. The baby is very cute.
expected: The baby is as cute as a button.
model: The ma is a pa.
The model has learned the shape of an answer (“The X is Y.”, then stop) and nothing else. It has never read enough text to know words, let alone facts. That’s fine-tuning without pretraining.
Now the real thing. Load GPT-2 (Chapter 9) and run the same command with --model-dir models/gpt2-medium (a GPU makes this a couple of minutes; BALLM reports 15.8 minutes for GPT-2 medium on an M3 MacBook Air CPU). BALLM’s GPT-2 medium (355M), after two epochs on this dataset, answered (BALLM §7.7):
### Rewrite the sentence using a simile. The car is very fast.
expected: The car is as fast as lightning.
model: The car is as fast as a bullet.
### What type of cloud is typically associated with thunderstorms?
expected: The type of cloud typically associated with thunderstorms is cumulonimbus.
model: The type of cloud associated with thunderstorms is a cumulus cloud.
Same data, same loop; the difference is entirely pretraining. Note the second answer: fluent, confident and wrong. Evaluating instruction-tuned models is hard.
Fine-tuning Qwen3-0.6B
For the real checkpoint, finetune.py uses Hugging Face Transformers for the model, so you can focus on the data and the measurements. The same instruction data is in chat format in data/instruction-{train,valid,test}.jsonl:
uv pip install -r optional-requirements.txt
python finetune.py train --mode full --model-dir models/Qwen3-0.6B \
--train-file data/instruction-train.jsonl --valid-file data/instruction-valid.jsonl \
--output runs/qwen3-sft --steps 200 --lr 1e-5 --batch-size 4 --accumulation 4
python finetune.py evaluate --model-dir models/Qwen3-0.6B \
--valid-file data/instruction-valid.jsonl --output runs/qwen3-sft
What the script does, each of which you’ve now seen from first principles:
- Formats every example with Qwen’s own chat template (
enable_thinking=False), and checks that the prompt-with-generation-prompt is an exact prefix of the full conversation, so the mask boundary is right. - Masks the prompt and padding, keeps the
<|im_end|>that ends the response, and refuses over-long examples. - Accumulates gradients over
--accumulationmicro-batches, dividing by the total number of response tokens in the update (not by the number of micro-batches). - Saves the model with a
manifest.jsonrecording the base config’s checksum, dataset checksums, arguments and library versions, andevaluatereloads it in a fresh process.
Qwen3-0.6B is already instruction-tuned, so this run teaches it the dataset’s terse answer style. To see the full effect of SFT, start from a base checkpoint such as Qwen/Qwen3-0.6B-Base.
The memory of full fine-tuning
Inference needs the weights. Full fine-tuning with AdamW in mixed precision needs, per parameter:
| item | bytes |
|---|---|
| BF16 weights used in the forward pass | 2 |
| FP32 master copy of the weights | 4 |
| gradients (FP32) | 4 |
| AdamW first and second moments (FP32) | 8 |
| total | 18 |
For 0.6B parameters, that’s about 11 GB before activations; for an 8B model, 144 GB. Activations add more, growing with batch size × sequence length × layers (gradient checkpointing trades recomputation for most of it). That’s why the next chapter’s methods exist.
Evaluating a fine-tuned model
A falling training loss proves only that the model memorizes the training set. In increasing order of cost and value:
- Held-out loss on response tokens, computed identically before and after.
- Exact checks where answers are checkable: classification accuracy, arithmetic, format compliance (“did it stop?”, “is it valid JSON?”).
- Regression checks on prompts unrelated to the fine-tuning data: did general ability survive? Small datasets and high learning rates cause catastrophic forgetting.
- Model-as-judge: a stronger model scores responses against references. BALLM §7.8 does this with Llama 3 8B through Ollama, scoring each response 0-100. Useful, but biased toward long and confident answers; spot-check it.
- Human review of a sample, especially of failures.
Build it
Engine milestone 21: fine-tuning. Implement instruction_batch in engine/lora.py and last_token_logits in engine/sft.py (formatting, the training loops and replace_head are provided).
pytest tests/test_ch21_finetune.py
python run.py classify --impl engine
python run.py sft --impl engine --steps 600 --width 192
The tests check the shifted targets and masks of a padded batch, that masked positions contribute nothing to the loss, that classification reads each sequence’s last real token, that replace_head freezes the right parameters, and that over-long examples are dropped rather than truncated. Then fine-tune GPT-2 (or Qwen3-0.6B) and record held-out loss before and after, and three test responses.
Stretch exercises
- ★ Train on all tokens (prompt included) instead of response-only, with the same budget. Compare held-out response loss and the test answers. Where: the response mask built by
instruction_batchinengine/lora.py, consumed byengine.sft.response_loss. - ★★ Classify SMS spam with pretrained GPT-2, comparing three settings: head only, head + last block, all layers. Plot test accuracy against trainable parameters. Where:
experiments/ch21.py(create it), adaptingrun.py’scmd_classifyand callingengine.sft.train_classifier. - ★★ Implement sequence packing: concatenate several examples into one row, with a block-diagonal causal mask and positions restarting at 0 for each example. Your
causal_attentionalready takes anallowedmask (Chapter 5). Measure the speedup over dynamic padding. Where: add a packed batch builder inengine/lora.py; pass its masks/positions throughengine/sft.pyandengine/gpt.pytoengine.attention.causal_attention. - ★★★ Implement DPO (Rafailov et al., 2023) on BALLM’s preference dataset (
ch07/04_preference-tuning-with-dpoin Raschka’s repository): the loss compares the policy’s and a frozen reference model’s log-probabilities of the chosen and rejected responses. Where: add a DPO loss/training helper inengine/sft.py; load chosen/rejected pairs inexperiments/ch21.py(create it).
Check your understanding
- Why does a classifier read the last token’s hidden state, not the first?
- In the shifted inputs and targets, which input position predicts the first response token?
- Why must padding be masked by position when the pad ID equals the EOS ID?
- Why does full fine-tuning need many times the memory of inference?
- Why did the from-scratch model learn the answer format but none of the facts?
Going deeper
- BALLM Chapter 6 (classification fine-tuning: §6.5 adding the head, §6.6 the last token, §6.7 training) and Chapter 7 (instruction fine-tuning: §7.3 batching and masking, §7.6 training, §7.8 evaluation with a judge model). Raschka’s repository has bonus material on DPO, LoRA variants and larger models.
- Ouyang et al., Training language models to follow instructions with human feedback (InstructGPT, 2022); Taori et al., Alpaca (2023); Zhou et al., LIMA: Less Is More for Alignment (2023), on how little data format-teaching needs.
- Rafailov et al., Direct Preference Optimization (2023).
- Hugging Face TRL’s
SFTTrainerdocumentation, the production version of this chapter’s loop.
22. LoRA and QLoRA
In this chapter
- Why fine-tuning updates can be low-rank, and how LoRA trains a model through two small matrices per layer.
- Initialization, the α/r scale, and which layers to adapt.
- Merging an adapter into the weights for zero inference cost, or keeping it separate to serve many adapters on one base.
- QLoRA: training adapters on a 4-bit frozen base, and the memory arithmetic of all three approaches.
You will build
LoRALinear in engine/lora.py: the adapted forward pass and the merge. You'll teach your instruction-tuned model a new style with 5% of its parameters, merge the adapter and verify that nothing changed.
Time: 3-5 hours. GPU: recommended for Qwen3 (the demo runs on a CPU in under a minute).
The cost of full fine-tuning
Chapter 21 ended with an uncomfortable table: full fine-tuning with AdamW needs about 18 bytes per parameter, 144 GB for an 8B model before activations. And every fine-tuned variant is a complete copy of the model: ten customers with ten fine-tunes of an 8B model means 160 GB of checkpoints.
LoRA (Low-Rank Adaptation, Hu et al., 2021) fixes both. It freezes every pretrained weight and learns a small correction for selected matrices. Only the corrections get gradients and optimizer state, and only they are saved.
The low-rank idea
Fine-tuning changes a weight matrix $W \in \mathbb{R}^{d_\text{out} \times d_\text{in}}$ to $W + \Delta W$. Empirically, the useful $\Delta W$ for adapting a pretrained model has low intrinsic rank (Aghajanyan et al., 2020): it can be well approximated by a product of two thin matrices. So LoRA parameterizes it that way:
$$ W’ = W + \frac{\alpha}{r} B A, \qquad A \in \mathbb{R}^{r \times d_\text{in}},\ B \in \mathbb{R}^{d_\text{out} \times r},\ r \ll \min(d_\text{in}, d_\text{out}). $$
For a 4,096 × 4,096 projection, full fine-tuning trains 16.8M numbers; LoRA with $r = 8$ trains $8 \times (4096 + 4096) = 65{,}536$, which is 0.39%.
The forward pass never forms $BA$. It computes the frozen path and the low-rank path separately and adds them:
$$ y = W x + \frac{\alpha}{r}, B (A x). $$
$Ax$ is an $r$-vector, cheap to compute. The extra FLOPs are $2r(d_\text{in} + d_\text{out})$ per token, under 1% of the layer.
class LoRALinear(nn.Module):
"""y = base(x) + (alpha / r) * B(A(x)), with the base frozen. (Your engine: Chapter 22)
A [r, in] starts random and B [out, r] starts at zero, so the adapted layer initially
equals the base exactly; B receives a gradient on the first step, A on the next.
"""
def __init__(self, base, rank=8, alpha=16):
super().__init__()
if rank < 1:
raise ValueError("rank must be positive")
self.base = base
for parameter in base.parameters():
parameter.requires_grad_(False)
self.scale = alpha / rank
self.A = nn.Parameter(base.weight.new_empty(rank, base.in_features))
self.B = nn.Parameter(base.weight.new_zeros(base.out_features, rank))
nn.init.kaiming_uniform_(self.A, a=math.sqrt(5))
def forward(self, x):
"""(Your engine: Chapter 22)"""
return self.base(x) + F.linear(F.linear(x, self.A), self.B) * self.scale
@torch.no_grad()
def merged(self):
"""A plain nn.Linear with W + scale * B @ A folded in. Do not also keep the adapter active. (Your engine: Chapter 22)"""
out = nn.Linear(self.base.in_features, self.base.out_features, bias=self.base.bias is not None,
device=self.base.weight.device, dtype=self.base.weight.dtype)
out.weight.copy_(self.base.weight + self.scale * (self.B @ self.A))
if self.base.bias is not None:
out.bias.copy_(self.base.bias)
return out
Initialization and the scale
- $B = 0$, $A$ random. The adapted model starts exactly equal to the base, so training begins from the pretrained behavior. On the first step, $B$ gets a gradient ($\partial L / \partial B = \frac{\alpha}{r} g, (Ax)^\top$, which is non-zero) while $A$’s gradient is zero ($\partial L/\partial A = \frac{\alpha}{r} B^\top g, x^\top$, and $B = 0$). After one update, $B \ne 0$ and both train. The milestone test checks this.
- $\alpha/r$. Scaling the update by $\alpha/r$ keeps its size roughly stable when you change $r$, so a learning rate tuned at one rank still works at another. A common choice is $\alpha = 2r$. (rsLoRA, Kalajdzievski 2023, argues for $\alpha/\sqrt{r}$ at large ranks.)
Which layers to adapt
def add_lora(model, targets=("q_proj", "v_proj"), rank=8, alpha=16):
"""Freeze the model, then wrap every Linear whose attribute name is in targets."""
for parameter in model.parameters():
parameter.requires_grad_(False)
wrapped = []
for name, module in list(model.named_modules()):
for child_name, child in list(module.named_children()):
if child_name in targets and isinstance(child, nn.Linear):
setattr(module, child_name, LoRALinear(child, rank, alpha))
wrapped.append(f"{name}.{child_name}")
return wrapped
def merge_lora(model):
"""Replace every LoRALinear with its merged Linear, in place."""
for module in list(model.modules()):
for child_name, child in list(module.named_children()):
if isinstance(child, LoRALinear):
setattr(module, child_name, child.merged())
return model
The original paper adapted only the attention query and value projections. The QLoRA paper found that adapting all linear layers (attention and MLP) matters more than the rank: with all layers adapted, ranks from 8 to 64 performed similarly. Embeddings and the LM head are usually left alone; for a model with tied embeddings, adapting the head would also change the input embedding through the shared tensor.
Run it: a new style with 5% of the parameters
Chapter 21’s run.py sft saved its instruction-tuned GPT to runs/sft.pt. Now teach it a new behavior, answering in capital letters, by training LoRA adapters on every linear layer with the same instructions and upper-cased responses:
python run.py lora --steps 300
On a laptop CPU, in 33 seconds:
{"before_lora": "The gas' is 'bua'.", "valid_loss_on_caps": 6.5372, "dropped_overlong": 0}
{"wrapped_layers": 16, "trainable": 98304, "total": 2025600, "trainable_fraction": 0.0485}
{"step": 1, "train_loss": 6.5902, "valid_loss": 6.4817}
{"step": 151, "train_loss": 3.6265, "valid_loss": 3.6524}
{"step": 300, "train_loss": 3.2451, "valid_loss": 3.5377}
{"instruction": "What is the contraction for 'it is'?", "after_lora": "THE OMES 'TES 'S 'S 'TETES 'TETETETETESTETESSTESTESTE"}
{"instruction": "Convert the number 5 from decimal to binary.", "after_lora": "15 MES: 14"}
{"merged_max_abs_difference": 4.410743713378906e-06, "parameters_after_merge": 2025600}
The style moved completely; the content was nonsense before and remains so (this model never pretrained), and the first answer degenerates into repetition, Chapter 8’s classic greedy failure. The fraction is high here only because the model is tiny. For Qwen3-0.6B, rank 8 on q_proj and v_proj trains 1.15M parameters, 0.19% of the model. On pretrained GPT-2 small, BALLM Appendix E’s LoRA (rank 16 on every linear layer, 2.7M trainable parameters, 2% of the model) reaches 98.0% test accuracy on spam classification, slightly above full fine-tuning of the last block in Chapter 6.
The last line is the merge check, discussed next.
Merge, or keep adapters separate
After training, an adapter can be merged: compute $W + \frac{\alpha}{r}BA$ once and replace $W$. The merged model is an ordinary model with zero extra inference cost. The demo’s merge changed the logits by $4 \times 10^{-6}$, floating-point rounding from adding in a different order.
Or keep adapters separate. One base model in GPU memory can then serve hundreds of fine-tunes, each a few megabytes, applying each request’s adapter in its forward pass. Batching requests with different adapters needs a gathered low-rank matmul (S-LoRA and Punica’s BGMV kernels; vLLM and SGLang support this). That’s how serving providers host many customer fine-tunes cheaply.
Rules that prevent subtle bugs:
- Never apply an adapter to a model it has already been merged into: the update would count twice.
- A merged model that you then quantize is a different model from the adapter-on-quantized-base you trained with QLoRA. Evaluate the artifact you’ll ship.
- An adapter is only valid for the exact base weights it was trained on. Record the base’s revision with the adapter (
finetune.pywrites it to its manifest).
QLoRA: a 4-bit frozen base
LoRA removes gradients and optimizer state for the base, but the base weights themselves still sit in BF16. QLoRA (Dettmers et al., 2023) stores the frozen base in 4 bits and dequantizes on the fly in the forward and backward passes, while the adapters train in BF16. It combines three ideas:
- NF4, a 4-bit data type whose 16 levels are quantiles of a normal distribution, matching how pretrained weights are distributed (Chapter 20). Unlike Chapter 20’s uniform INT4 grid, dequantizing is a table lookup.
- Double quantization: the per-block scales (one per 64 weights) are themselves quantized to 8 bits, saving about 0.37 bits per parameter.
- Paged optimizers: optimizer state that can spill to CPU memory during memory spikes.
Memory per parameter of the base, with activations and the small adapters aside:
| method | base weights | gradients + optimizer | total per base parameter | 8B model |
|---|---|---|---|---|
| full fine-tuning (mixed precision) | 2 + 4 (master) | 4 + 8 | 18 bytes | ~144 GB |
| LoRA, BF16 base | 2 | ~0 | ~2 bytes | ~16 GB |
| QLoRA, NF4 base | ~0.52 | ~0 | ~0.52 bytes | ~4.2 GB |
That’s why QLoRA fine-tunes an 8B model on a single consumer GPU, and why the QLoRA paper could fine-tune a 65B model on one 48 GB GPU. The cost is speed: dequantizing every weight in every forward and backward pass makes each step slower than BF16 LoRA.
LoRA and QLoRA on Qwen3
finetune.py uses Hugging Face PEFT for the adapters, with the same data and loss handling as Chapter 21:
uv pip install -r workflow-requirements.txt
python finetune.py train --mode lora --rank 8 --model-dir models/Qwen3-0.6B \
--train-file data/instruction-train.jsonl --valid-file data/instruction-valid.jsonl \
--output runs/qwen3-lora --steps 200 --lr 2e-4 --batch-size 4 --accumulation 4
python finetune.py evaluate --model-dir models/Qwen3-0.6B \
--valid-file data/instruction-valid.jsonl --output runs/qwen3-lora
The saved directory holds only the adapter (a few MB) and the tokenizer. Note the learning rate: LoRA typically needs 10-20× the learning rate of full fine-tuning, because the update starts at zero and lives in a small subspace.
For QLoRA, install bitsandbytes and pass --mode qlora. The script loads the base with NF4, double quantization and BF16 compute, prepares it for k-bit training (casting norms to FP32, enabling gradient checkpointing), and adds the same adapters.
When bitsandbytes’ CUDA build doesn’t match PyTorch’s
bitsandbytes ships compiled CUDA libraries for specific CUDA versions. If PyTorch uses a newer CUDA than any bundled library, bitsandbytes looks for a file that doesn’t exist (for example libbitsandbytes_cuda132.so with PyTorch on CUDA 13.2) and fails or falls back. Diagnose first:
python -m bitsandbytes # prints the CUDA version it detected and the library it loaded
If a library for an older compatible CUDA is bundled and that toolkit is installed, the documented override selects it, as in this example for a CUDA 13.0 build:
export BNB_CUDA_VERSION=130
export LD_LIBRARY_PATH="/usr/local/cuda-13.0/lib64${LD_LIBRARY_PATH:+:$LD_LIBRARY_PATH}"
python -m bitsandbytes # must now report success before you train
This is an environment-specific workaround, not a fix for an unsupported GPU or platform (check the installation guide for ARM64 and newer architectures). If QLoRA can’t run, use plain LoRA; don’t call the result QLoRA.
Build it
Engine milestone 22: LoRA. Implement LoRALinear.forward and LoRALinear.merged in engine/lora.py (the constructor, add_lora and merge_lora are provided).
pytest tests/test_ch22_lora.py
python run.py lora --impl engine --steps 300
The tests check that an adapted layer starts equal to its base, that only $A$ and $B$ receive gradients (and that $A$’s is zero on the first step), that the merged layer matches the live adapter, and that adding and merging adapters on a whole Qwen3 model preserves its outputs.
Stretch exercises
- ★ Run
run.py lorawith ranks 1, 2, 8 and 32, and plot final validation loss against trainable parameters. Where does a larger rank stop helping? Where: terminal:python run.py lora --impl engine --rank Rfor each rank; record results inexperiments/ch22.py(create it). - ★★ Compare full SFT (Chapter 21) and LoRA on GPT-2 with the same data and number of steps: held-out loss, peak memory (
torch.cuda.max_memory_allocated) and seconds per step. Where:experiments/ch22.py(create it), adaptingrun.py’scmd_sftandcmd_lora. - ★★ Serve two adapters at once: for a batch where row 0 uses adapter 1 and row 1 uses adapter 2, compute $B_{i}(A_{i} x)$ per row with gathered weights, and verify against running each row separately. Where: add a gathered-adapter helper in
engine/lora.py; Chapter 43 integrates it inengine/serve/adapters.py. - ★★★ Implement DoRA (Liu et al., 2024): decompose $W$ into magnitude and direction, $W’ = m \cdot \frac{W + BA}{\lVert W + BA \rVert_c}$ with a trainable per-column magnitude $m$, and compare with LoRA at equal parameter count. Where: add a DoRA layer beside
LoRALinearinengine/lora.py.
Check your understanding
- Why does initializing $B$ to zero make the adapted model equal to the base at the start?
- Why is $A$’s gradient zero on the first step, and why doesn’t that stop training?
- What does merging cost at inference time, and what does keeping adapters separate buy you?
- Why does QLoRA reduce memory but slow each step down?
- Why can’t an adapter trained on one base model be applied to another?
Going deeper
- BALLM Appendix E (Parameter-efficient fine-tuning with LoRA): LoRA from scratch on GPT-2 for spam classification, the source of this chapter’s comparison.
- Hu et al., LoRA (2021); Dettmers et al., QLoRA (2023); Aghajanyan et al., Intrinsic Dimensionality Explains the Effectiveness of Language Model Fine-Tuning (2020); Liu et al., DoRA (2024); Sheng et al., S-LoRA (2023) and Chen et al., Punica (2023) for multi-adapter serving.
- GPU Mode L32 (Unsloth: LLM systems engineering): fused Triton kernels for LoRA and QLoRA training.
- The Hugging Face PEFT documentation (LoRA configuration, merging, quantized bases) and Sebastian Raschka’s article Practical Tips for Finetuning LLMs Using LoRA.
23. Inside the model: hooks, the logit lens and steering
In this chapter
- The residual stream as a shared workspace that every layer reads and writes, and how to observe it with forward hooks.
- The logit lens: decoding what the model "would say" at every layer.
- Directions in activation space: finding one from contrasting examples, adding it to steer behavior, projecting it out to remove behavior.
- Making an edit permanent by orthogonalizing weights, and evaluating edits honestly.
You will build
direction_from_means and project_out in engine/steering.py. You'll read your model's layers with the logit lens, double the probability of an all-caps answer with one added vector, and remove a direction from a real Qwen3.
Time: 4-5 hours. GPU: optional.
The residual stream
Look at Chapter 6’s block again: x = x + attention(norm(x)), then x = x + mlp(norm(x)). Every layer reads the current $x$ and adds its result back. Nothing ever overwrites $x$. So the hidden state after layer $L$ is a sum:
$$ x_L = \underbrace{x_0}{\text{embedding}} + \sum{\ell=1}^{L} \big(\text{attn}\ell + \text{mlp}\ell\big). $$
This residual stream is a shared workspace: a $d$-dimensional vector per position that each component reads from and writes to, and that the final norm and head decode into the next-token prediction. Two consequences make this chapter possible:
- Because the head reads the stream through one linear map, you can apply it to the stream at any layer and see what’s there. That’s the logit lens.
- Because components communicate by adding vectors, a feature the model uses is often represented as a direction in this space (the linear representation hypothesis). Adding or removing a direction can change behavior directly.
Hooks: observing without changing code
PyTorch’s forward hooks run a function after a module’s forward pass, with access to its inputs and output. If the function returns a value, it replaces the output. That’s all you need to read and edit activations in any model, yours or a library’s, without changing its source:
@contextmanager
def capture(module, store, name):
"""Record a module's output into store[name] for every forward call inside the block."""
def hook(_module, _inputs, output):
store.setdefault(name, []).append((output[0] if isinstance(output, tuple) else output).detach())
handle = module.register_forward_hook(hook)
try:
yield store
finally:
handle.remove() # always detach, even when the forward raises
@torch.no_grad()
def residuals(model, ids, layers=None):
"""Hidden state after each requested decoder layer, as {layer_index: [B, T, D]}."""
blocks = list(_blocks(model))
layers = range(len(blocks)) if layers is None else layers
store, handles = {}, []
for i in layers:
handles.append(blocks[i].register_forward_hook(
lambda _m, _inp, out, i=i: store.__setitem__(i, (out[0] if isinstance(out, tuple) else out).detach())))
try:
model(ids)
finally:
for handle in handles:
handle.remove()
return store
Two habits matter. Always remove hooks in a finally block: a forgotten hook keeps editing every later forward pass, and an exception inside a forward is exactly when you’d forget. And detach() anything you store, or you’ll keep the whole autograd graph alive.
The logit lens
Apply the final norm and the head to the residual stream after each layer, and look at the top tokens:
@torch.no_grad()
def logit_lens(model, ids, top=5):
"""Decode every layer's residual through the final norm and head: what would the model
predict if it stopped here? Returns {layer: top token IDs at the last position}."""
norm = model.norm if hasattr(model, "norm") else model.model.norm
head = model.head if hasattr(model, "head") else model.lm_head
return {layer: head(norm(h[:, -1])).topk(top, dim=-1).indices[0].tolist()
for layer, h in residuals(model, ids).items()}
On a real model the picture is striking (nostalgebraist, 2020). Early layers predict tokens related to the current input; middle layers converge toward plausible continuations; the last few layers sharpen the final answer. For GPT-2 on “The Eiffel Tower is in the city of”, “ Paris“ typically rises to the top only in the later layers. Different models use their early layers differently, and the plain logit lens can be misleading for some; the tuned lens (Belrose et al., 2023) learns a small linear translator per layer to fix this.
Your 4-layer instruction-tuned GPT from Chapter 21, teacher-forced through “The contraction for”:
python run.py steer
{"text_end": " Response:\nThe contraction for", "layer_0": [" '", " \"", "'"], "layer_1": [" '", " \"", "'"],
"layer_2": [" '", " \"", " the"], "layer_3": [" '", " \"", " the"]}
Every layer already expects a quote, the dataset’s style for this kind of answer (“The contraction for ‘it is’ is …”). A small model decides early. The last layer’s lens equals the real prediction exactly, which the milestone test checks.
Directions from contrasting examples
How do you find the direction for a concept? The simplest method that works surprisingly well is a difference of means. Collect activations at one layer for examples with the property (positive) and without it (negative), and take the normalized difference of their averages:
$$ d = \frac{\bar{h}\text{pos} - \bar{h}\text{neg}}{\lVert \bar{h}\text{pos} - \bar{h}\text{neg} \rVert}. $$
The pairs should differ only in the property, so everything else averages out: “Explain recursion in a cheerful conversational tone” versus “… in a formal technical tone”, for 32 topics (data/style-positive.jsonl and style-negative.jsonl).
Two edits use a direction:
- Steering (adding): $h \leftarrow h + \lambda d$ at one layer, at every position. Positive $\lambda$ pushes the model toward the property.
- Directional ablation (projecting out): $h \leftarrow h - (h \cdot d), d$. The component along $d$ becomes exactly zero, so downstream layers can no longer read the property from it.
def direction_from_means(positive, negative):
"""Unit vector from the mean of negative examples toward the mean of positive ones. (Your engine: Chapter 23)"""
delta = positive.float().mean(0) - negative.float().mean(0)
norm = delta.norm()
if not torch.isfinite(norm) or norm < 1e-8:
raise ValueError("The two groups have (almost) the same mean: no usable direction")
return delta / norm
def project_out(hidden, direction, strength=1.0):
"""h - strength * (h . d) d. With strength 1 the d-component becomes exactly zero. (Your engine: Chapter 23)"""
h = hidden.float()
return (h - strength * (h @ direction.float()).unsqueeze(-1) * direction.float()).to(hidden.dtype)
@contextmanager
def steer(model, layer, direction, mode="add", strength=4.0):
"""Temporarily edit one block's output: 'add' pushes along the direction, 'ablate' removes it."""
block = _blocks(model)[layer]
def hook(_module, _inputs, output):
hidden = output[0] if isinstance(output, tuple) else output
if mode == "add":
edited = hidden + strength * direction.to(hidden)
elif mode == "ablate":
edited = project_out(hidden, direction.to(hidden.device), strength)
else:
raise ValueError("mode must be 'add' or 'ablate'")
return (edited, *output[1:]) if isinstance(output, tuple) else edited
handle = block.register_forward_hook(hook)
try:
yield model
finally:
handle.remove()
Run it
run.py steer builds a direction at layer 2 of your Chapter 21 model from 300 training answers, written normally (negative) and in capitals (positive). Then it adds the direction, scaled by multiples of the gap between the two means, and measures the probability that the first response token is in capitals, over 60 test prompts:
{"layer": 2, "added": "0.0 x gap (0.0)", "first_token_all_caps_probability": 0.2227}
{"layer": 2, "added": "0.5 x gap (5.4)", "first_token_all_caps_probability": 0.2877}
{"layer": 2, "added": "1.0 x gap (10.8)", "first_token_all_caps_probability": 0.3699}
{"layer": 2, "added": "1.5 x gap (16.1)", "first_token_all_caps_probability": 0.4431}
One vector, no training, and the probability doubles. (The baseline is 22% because single capital-letter tokens like “A” count.) In a model this small, though, the direction is crude: push harder, or generate long texts, and fluency collapses before the style changes cleanly. Steering works far better in large models, whose representations are more linear and more robust. Turner et al. (2023) steered GPT-2 XL’s topic and sentiment with a single prompt-pair difference, and Arditi et al. (2024) found that refusal in 13 open chat models is mediated by one direction.
Making an edit permanent
A hook works at runtime. To ship an edit, change the weights. Every write into the residual stream comes from a matrix: the embedding, each attention output projection and each MLP down projection. If each writer’s output has no component along $d$, the stream never gains one. For a linear layer $y = Wx + b$, replace
$$ W \leftarrow (I - d d^\top) W, \qquad b \leftarrow (I - d d^\top), b , $$
so that $d^\top y = 0$ for every input:
@torch.no_grad()
def orthogonalize_output(linear, direction):
"""Make a linear layer unable to write along `direction`: W <- (I - d d^T) W, b <- (I - d d^T) b.
Applied to every module that writes into the residual stream, this makes ablation permanent."""
d = direction.to(linear.weight)
linear.weight -= torch.outer(d, d @ linear.weight)
if linear.bias is not None:
linear.bias -= d * (d @ linear.bias)
Applied to all writers (and the embedding’s rows), this is weight orthogonalization, popularized as “abliteration” after Arditi et al.’s refusal paper. The milestone test checks that the edited layer equals the runtime projection exactly.
Warning
Directional ablation of a refusal direction removes a model’s safety behavior. It’s a well-documented interpretability result, and it’s why open-weight safety can’t rely on refusals alone. This book uses tone and capitalization as its examples; if you study refusal, do it on models and in settings where removing safeguards is appropriate, and don’t distribute such edits.
Editing a real Qwen3
edit_checkpoint.py runs the whole procedure on a Hugging Face Qwen3 checkpoint: it fits a direction at one block’s output from the last prompt token of the positive and negative prompts, ablates it with a hook on every position, and compares the edited and original models on held-out data:
python edit_checkpoint.py --model-dir models/Qwen3-0.6B --layer 14 \
--positive-file data/style-positive.jsonl --negative-file data/style-negative.jsonl \
--valid-file data/instruction-valid.jsonl --output runs/tone-direction.pt
It reports response loss before and after, the KL divergence from the original model on held-out responses, one greedy generation each way, and saves the direction with a manifest. No weights are changed; apply orthogonalize_output to make it permanent once you’re convinced.
Evaluate edits like any other change
- Did it do what you wanted? Measure the targeted behavior on held-out prompts, not the ones used to fit the direction.
- What else did it change? Report KL from the original on unrelated data, and held-out loss. A direction that changes everything is not a “tone” direction.
- Which layer? Directions work best in the middle layers; sweep a few and pick by held-out measurements, not by eye.
- Combine carefully. Ablation, LoRA (Chapter 22) and quantization (Chapter 20) don’t commute. Evaluate the artifact you ship.
Build it
Engine milestone 23: activation editing. Implement direction_from_means and project_out in engine/steering.py (hooks, the logit lens and weight orthogonalization are provided).
pytest tests/test_ch23_steering.py
python run.py steer --impl engine
The tests check that projecting out removes the component exactly and is idempotent, that hooks are removed even when the forward raises, that orthogonalized weights equal the runtime projection, and that the logit lens at the last layer equals the model’s real prediction.
Stretch exercises
- ★ Run
run.py steer --layer Lfor every layer, and plot the all-caps probability at 1.0 × gap against the layer. Where: terminal:python run.py steer --impl engine --layer L; plot results inexperiments/ch23.py(create it). - ★★ Load GPT-2 (Chapter 9) and print the logit lens for “The Eiffel Tower is in the city of”. At which layer does “ Paris“ enter the top 5? Where:
experiments/ch23.py(create it), callingengine.loaders.load_gpt2andengine.steering.logit_lens. - ★★ Steer Qwen3-0.6B’s tone with the style prompts: add $\lambda d$ at the best layer during generation and compare five answers at $\lambda = 0$, 4 and 8. Where:
experiments/ch23.py(create it), usingengine.steering.steerwithengine.engine.LLM. - ★★★ Orthogonalize every residual writer of your Qwen3 engine against a direction, save the edited checkpoint, and verify its logits equal the hooked model’s. Where: extend
orthogonalize_outputinengine/steering.pyfor Qwen3’s residual writers; save withengine.safetensors_io.save_file.
Check your understanding
- Why can the final norm and head be applied to the residual stream after any layer?
- Why should positive and negative examples differ only in the property you want?
- What’s the difference between adding a direction and projecting it out?
- Why does orthogonalizing every residual writer make ablation permanent?
- How would you show that an edit changed only what you intended?
Going deeper
- nostalgebraist, interpreting GPT: the logit lens (2020); Belrose et al., Eliciting Latent Predictions from Transformers with the Tuned Lens (2023).
- Elhage et al., A Mathematical Framework for Transformer Circuits (Anthropic, 2021): the residual-stream view this chapter uses.
- Turner et al., Activation Addition: Steering Language Models Without Optimization (2023); Zou et al., Representation Engineering (2023); Arditi et al., Refusal in Language Models Is Mediated by a Single Direction (2024).
- Sparse autoencoders, which find thousands of interpretable directions at once: Bricken et al., Towards Monosemanticity (2023) and Templeton et al., Scaling Monosemanticity (2024); Gemma Scope and the
TransformerLensandnnsightlibraries for doing this on real models.
24. Serving many users: continuous batching
In this chapter
- Why batching makes decode almost free per extra user, and where that stops: weights amortize, KV-cache reads don't.
- Static batching and its waste; iteration-level (continuous) batching, which fixes it.
- A scheduler with a token budget, decode-first priority and chunked prefill.
- Per-request state, and the test that matters most: every request gets exactly its solo output.
- Serving metrics: time to first token, time per output token, throughput and tail latency.
You will build
Scheduler.plan and ContinuousBatchingEngine.step in engine/scheduler.py: a serving loop that runs many requests through your Qwen3 at once, each in its own cache slot.
Time: 5-7 hours. GPU: recommended for the measurements (everything runs on a CPU).
One user leaves the GPU idle
Decode at batch 1 is memory-bound: each step reads every weight once to produce one token (Chapter 10). A 1,024 × 4,096 BF16 linear layer does 2 FLOPs per weight read, for an arithmetic intensity of about 1 FLOP per byte. An H100 can do around 300 FLOPs in the time it reads one byte. At batch 1, more than 99% of its arithmetic capability idles.
Now run 64 users’ decode steps in one forward pass. The weights are read once and multiplied by 64 vectors: 64 times the work for the same weight traffic. PMPP §20.6 works the numbers for that layer: arithmetic intensity grows from 1 FLOP/B at batch 1 to 315 FLOP/B at batch 512. Until the batch is large enough to become compute-bound, each extra user’s decode is nearly free.
There’s an important exception. Each user has their own KV cache, so attention reads grow linearly with the batch:
$$ \text{bytes per step} \approx \underbrace{\text{weight bytes}}{\text{shared}} + \sum{\text{requests}} \underbrace{\text{KV bytes}(\text{context}i)}{\text{per request}} . $$
For Qwen3-0.6B (1.19 GB of weights, 112 KiB of KV per token), 64 users at 2,000 tokens each read about 15 GB of cache per step, twelve times the weights. Batching helps until the KV cache dominates, and the KV cache also limits how many users fit in memory. Those two facts drive the next two chapters: Chapter 25 packs caches without waste, and Chapter 26 gets more tokens out of each cache read.
Static batching and its waste
The simplest batching collects $B$ requests, pads their prompts to the same length, and generates until all of them finish:
step: 1 2 3 4 5 6 7 8 9 10
request A: ■ ■ ■ ✓ · · · · · · finished at step 4, slot idles
request B: ■ ■ ■ ■ ■ ■ ■ ■ ■ ✓
request C: ■ ■ ✓ · · · · · · · new request D waits for the whole batch
Outputs vary from a few tokens to thousands, so most slots idle most of the time, and newly arrived requests wait for the slowest member of the current batch (head-of-line blocking).
Continuous batching (Orca, Yu et al., 2022) makes the decision at every iteration instead: after each step, finished requests leave and waiting requests join. The batch is a set of independent requests that happen to share a forward pass.
Prefill and decode in one step
Joining requests need their prompts prefilled, which is a very different workload from decode (Chapter 16): hundreds of tokens at once, compute-bound. If a 4,000-token prompt is prefilled in one step, every running user’s next token waits for it, and their time per output token jumps.
The fix is a token budget per step and chunked prefill (Sarathi-Serve, Agrawal et al., 2024):
- Admit waiting requests while cache slots are free.
- Decode first: every request that has finished its prefill gets one token in this step.
- Spend the rest of the budget on prefill chunks of at most
prefill_chunktokens, oldest request first.
A long prompt is then spread over several steps, and decode latency stays steady. The budget is the knob: a large budget finishes prefills sooner (lower time to first token) but makes each step slower (higher time per output token).
class Scheduler:
"""Decode-first, token-budgeted, chunked-prefill planner. (Your engine: Chapter 24)
Each step: admit waiting requests while slots are free; give every running request that
has finished prefill one decode token; spend what is left of the budget on prompt chunks
of at most prefill_chunk tokens, oldest request first.
"""
def __init__(self, max_slots=8, token_budget=64, prefill_chunk=32):
if min(max_slots, token_budget, prefill_chunk) < 1:
raise ValueError("Scheduler limits must be positive")
self.max_slots, self.token_budget, self.prefill_chunk = max_slots, token_budget, prefill_chunk
self.free_slots = deque(range(max_slots))
self.waiting, self.running = deque(), []
def submit(self, request):
if not request.prompt or request.max_new_tokens < 1:
raise ValueError("A request needs a prompt and at least one new token")
self.waiting.append(request)
def plan(self):
"""(Your engine: Chapter 24)"""
while self.waiting and self.free_slots:
request = self.waiting.popleft()
request.slot = self.free_slots.popleft()
self.running.append(request)
budget = self.token_budget
decode = [r for r in self.running if r.prefilled == len(r.prompt)][:budget]
budget -= len(decode)
prefill = []
for request in self.running:
remaining = len(request.prompt) - request.prefilled
if remaining and budget:
length = min(remaining, self.prefill_chunk, budget)
prefill.append((request, request.prefilled, length))
budget -= length
return Plan(prefill, decode)
def release(self, request):
"""Return a finished or cancelled request's slot."""
self.running.remove(request)
self.free_slots.append(request.slot)
@property
def idle(self):
return not self.waiting and not self.running
The scheduler is pure policy: it moves no tensors. Keeping it separate from execution makes it testable with plain Python lists, and it’s where production engines differ most (priorities, preemption, fairness, prefix-aware ordering).
The engine loop
class ContinuousBatchingEngine:
"""Runs Scheduler plans on a model that follows the book's cache protocol. (Your engine: Chapter 24)"""
def __init__(self, model, max_slots=8, capacity=512, token_budget=64, prefill_chunk=32):
self.model = model.eval()
p = next(model.parameters())
layers, kv_heads, head_dim = model.cache_spec()
self.cache = StaticKVCache(layers, max_slots, kv_heads, capacity, head_dim, p.device, p.dtype)
self.capacity, self.device = capacity, p.device
self.scheduler = Scheduler(max_slots, token_budget, prefill_chunk)
self.generators = {}
self.steps = 0
def submit(self, request):
if len(request.prompt) + request.max_new_tokens > self.capacity:
raise ValueError(f"{request.rid}: prompt plus output exceed slot capacity")
request.arrival = time.perf_counter()
self.generators[request.rid] = torch.Generator(device=self.device).manual_seed(request.seed)
self.scheduler.submit(request)
def _emit(self, request, logits_row):
token = int(sample(logits_row[None], request.temperature, generator=self.generators[request.rid]))
if request.first_token_time is None:
request.first_token_time = time.perf_counter()
request.output.append(token)
if token in request.stop_ids:
request.finish_reason = "stop"
elif len(request.output) >= request.max_new_tokens:
request.finish_reason = "length"
@torch.inference_mode()
def step(self):
"""One engine iteration. Returns the requests that finished during it. (Your engine: Chapter 24)"""
plan = self.scheduler.plan()
for request, start, length in plan.prefill: # prompt chunks, one request at a time
ids = torch.tensor([request.prompt[start:start + length]], device=self.device)
positions = torch.arange(start, start + length, device=self.device)[None]
rows = torch.tensor([request.slot], device=self.device)
logits = self.model(ids, self.cache, positions=positions, rows=rows)
request.prefilled += length
if request.prefilled == len(request.prompt): # the last prompt position predicts token 1
self._emit(request, logits[0, -1])
if plan.decode: # every decoding request, one batched forward
ids = torch.tensor([[r.output[-1]] for r in plan.decode], device=self.device)
positions = torch.tensor([[len(r.prompt) + len(r.output) - 1] for r in plan.decode], device=self.device)
rows = torch.tensor([r.slot for r in plan.decode], device=self.device)
logits = self.model(ids, self.cache, positions=positions, rows=rows)
for i, request in enumerate(plan.decode):
self._emit(request, logits[i, -1])
finished = [r for r in self.scheduler.running if r.done]
for request in finished:
request.finish_time = time.perf_counter()
self.scheduler.release(request)
self.generators.pop(request.rid, None)
self.steps += 1
return finished
def run(self, requests):
"""Submit everything, step until idle, return {rid: request}."""
for request in requests:
self.submit(request)
results = {}
while not self.scheduler.idle:
for request in self.step():
results[request.rid] = request
return results
Each request owns one row of a StaticKVCache (Chapter 19), its slot. Rows hold different lengths; writes go to (row, position) pairs and the causal rule hides everything a row hasn’t written. That’s why one batched forward can decode requests at positions 7, 58 and 1,203 at once: the model receives positions of shape [B, 1] and rows of shape [B], and the cache protocol does the rest. No padding of prompts, no attention mask beyond key_pos <= query_pos.
Per-request state lives on the Request: prompt, output, how much has been prefilled, stop tokens, finish reason, and its own random generator, seeded per request so that a request’s samples don’t depend on who else shares the batch.
The correctness test
Batching is an optimization, so it must not change any request’s output. The milestone test runs four requests with different prompt and output lengths through 3 slots, with chunked prefill and a small budget, and checks that each one’s greedy output equals the output of generating it alone. This catches wrong positions, wrong rows, a stale cache slot from a previous request and off-by-one errors in the prefill-to-decode handoff. (With sampling, outputs match only if each request’s random stream is independent of the batch, which is why generators are per request.)
Note
Even with correct code, batched and solo outputs can differ in the last bits in BF16 on a GPU: a matmul over a batch of 8 may use a different kernel or reduction order than a batch of 1. Greedy tokens can then diverge at a near-tie. Production engines that need batch-invariant results use kernels designed for it (Thinking Machines, Defeating Nondeterminism in LLM Inference, 2025).
Run it
python run.py batch --requests 12 --slots 8
Twelve requests with prompts of 8-64 tokens and outputs of 8-48 tokens, through the 2-layer test model on a laptop CPU, first one at a time, then continuously batched:
{"requests": 12, "output_tokens": 322, "serial_tok_s": 398.4, "batched_tok_s": 1037.6, "engine_steps": 55}
{"rid": "5", "ttft_ms": 25.06, "tpot_ms": 8.15, "tokens": 9, "finish": "length"}
{"rid": "6", "ttft_ms": 29.77, "tpot_ms": 7.69, "tokens": 10, "finish": "length"}
{"rid": "2", "ttft_ms": 9.11, "tpot_ms": 7.53, "tokens": 21, "finish": "length"}
2.6× the throughput, in 55 engine steps instead of 322 sequential ones. With --slots 1 the engine does no batching and reached 549 tokens/s against the serial loop’s 640 on the same run: the scheduler’s bookkeeping costs a little, and batching is where the win comes from. On a GPU, where a batch of 8 costs about the same as a batch of 1, the gain approaches the slot count.
Serving metrics
| metric | definition | what users feel |
|---|---|---|
| TTFT (time to first token) | arrival → first token | responsiveness; queueing + prefill |
| TPOT (time per output token) | average gap between later tokens | streaming speed; should beat reading speed (~5-10 tokens/s) |
| throughput | output tokens per second, all requests | cost per token |
| goodput | throughput of requests that met their latency targets | what you can actually sell |
Report percentiles (p50, p90, p99), not averages: a scheduler that’s great on average and terrible for 1% of users is a bad scheduler. And benchmark with a realistic arrival process (requests arriving over time, for example Poisson), not everything submitted at once, which hides queueing.
Production engines add more scheduling tools: preemption (pause a request and free its cache when memory runs out, then recompute or swap it back in), priorities, prefix-aware ordering (Chapter 25), and disaggregation, running prefill and decode on separate GPUs so the two workloads stop interfering (DistServe, Splitwise, 2024).
Batched prefill in production
This engine prefills one request’s chunk per forward call, which is simple and correct but launches many small forward passes. Production engines flatten all of a step’s tokens, prefill chunks and decode tokens of every request, into one long sequence of $N$ tokens, with per-token positions and request IDs. The linear layers see one [N, d] matrix, ideal for tensor cores. Attention uses a “varlen” kernel that receives the boundaries between requests (cu_seqlens in FlashAttention) and each request’s cache location. Your model’s positions-based interface already supports this; what’s missing is an attention kernel that takes ragged batches (Chapter 25’s paged decode kernel is a step in that direction).
Build it
Engine milestone 24: continuous batching. Implement Scheduler.plan and ContinuousBatchingEngine.step in engine/scheduler.py (request bookkeeping, submit, release, run and the latency report are provided).
pytest tests/test_ch24_batching.py
python run.py batch --impl engine
The tests check that the scheduler respects the slot limit and the token budget and serves decode before prefill, and that four requests sharing three slots with chunked prefill each produce exactly their solo greedy output and return every slot when done.
Stretch exercises
- ★ Sweep
--slotsfrom 1 to 32 on a GPU with Qwen3-0.6B, and plot throughput and mean TPOT. Where does throughput stop growing? Where: createexperiments/ch24.pyfromrun.py’scmd_batch, replacing its tiny model withengine.loaders.load_qwen3and collecting per-request timings. Sweepmax_slotsin that script;run.py batch --slots Nis the tiny-model baseline. - ★★ Add Poisson arrivals: submit requests at random times while the engine runs, and report p50 and p99 TTFT for token budgets 64, 256 and 1,024. Where:
experiments/ch24.py(create it), adaptingrun.py’scmd_batchand submitting requests toengine.scheduler.ContinuousBatchingEngine. - ★★ Serve your engine over HTTP with an OpenAI-compatible
/v1/completionsendpoint (FastAPI and server-sent events), running the engine loop in a background thread and streaming each request’s tokens as they’re emitted. Where: newexperiments/batching_server.py, wrappingengine.scheduler.ContinuousBatchingEngine; Chapter 36 supplies the later unified server. - ★★★ Implement preemption: when a new high-priority request arrives and no slot is free, evict the request with the most remaining work, and recompute its cache from prompt + output when it’s readmitted. Verify outputs are still identical to solo runs. Where:
SchedulerandContinuousBatchingEngineinengine/scheduler.py.
Check your understanding
- Why does batching speed up the weight reads of decode but not its KV-cache reads?
- What does continuous batching fix that static batching can’t?
- Why does the scheduler give decode tokens priority over prefill chunks?
- How does one batched forward serve requests at different positions without padding?
- Why should each request have its own random generator?
Going deeper
- PMPP §20.6 (KV cache arithmetic intensity and memory requirement, pp. 504-508): batching’s effect on intensity and the batch-size versus context-length trade-off.
- Yu et al., Orca: A Distributed Serving System for Transformer-Based Generative Models (OSDI 2022); Agrawal et al., Taming Throughput-Latency Tradeoff in LLM Inference with Sarathi-Serve (OSDI 2024); Zhong et al., DistServe (2024).
- GPU Mode L35 (SGLang performance optimization), and the vLLM and SGLang schedulers (
vllm/v1/core/sched/scheduler.py,sglang/srt/managers/scheduler.py), which follow the plan-then-execute structure used here. - Anyscale’s blog post How continuous batching enables 23x throughput in LLM inference for measurements on real GPUs.
25. Paged attention and prefix caching
In this chapter
- Why reserving a maximum-length cache per request wastes most of the GPU's memory.
- Paging, borrowed from operating systems: fixed-size blocks, a shared pool and per-sequence block tables.
- Sharing blocks between sequences with reference counts and copy-on-write: parallel sampling and beam search for free.
- Prefix caching: reusing the KV cache of a shared system prompt across requests.
- A decode attention kernel that reads keys and values straight through the block table.
You will build
BlockAllocator and PagedKVCache in engine/paged.py, and a paged decode attention kernel in engine/kernels/triton_paged.py.
Time: 5-7 hours. GPU: optional (the Triton kernel runs in the interpreter).
Reserved memory is wasted memory
Chapter 24’s engine gives each request a slot of fixed capacity, say 1,024 tokens. A request that uses 50 tokens still reserves 1,024. With requests of varied lengths, most of the reserved cache is empty:
python run.py paged
{"requests": 32, "tokens_stored": 9924, "static_slots_reserved": 32768, "paged_slots_reserved": 10144,
"static_utilization": 0.303, "paged_utilization": 0.978}
Thirty-two live requests of 20-500 tokens fill 30% of their reserved slots. The vLLM paper (Kwon et al., 2023) measured existing systems using only 20-38% of their KV memory for actual tokens. Since the cache limits how many requests fit (Chapter 24), that waste directly cuts throughput.
Allocating exactly the right amount isn’t possible either: you don’t know how long an answer will be until it’s done. And growing a contiguous buffer means copying it, or fragmenting memory into holes that no request fits.
Paging
Operating systems solved the same problem for process memory decades ago. PagedAttention applies the solution to the KV cache:
- Split every sequence’s cache into fixed-size blocks of $b$ tokens (16 is typical).
- Keep all blocks of all sequences in one pool: per layer, tensors of shape
[num_blocks, kv_heads, b, head_dim]. - Give each sequence a block table mapping its logical blocks to physical blocks in the pool, which need not be contiguous or in order.
Logical position $p$ lives in physical block table[p // b] at offset p % b. For $b = 4$ and table [5, 1, 7], position 9 is in logical block 2, physical block 7, offset 1.
print(table[9 // 4], 9 % 4) # (7, 1)
inline std::pair<size_t, size_t> physical_slot(const std::vector<size_t>& table, size_t pos, size_t block) {
return {table[pos / block], pos % block};
}
#![allow(unused)]
fn main() {
pub fn physical_slot(table: &[usize], position: usize, block_size: usize) -> (usize, usize) {
(table[position / block_size], position % block_size)
}
}
A sequence allocates a new block only when it fills its last one, so it wastes at most $b - 1$ slots, under 4% in the demo above. Blocks come from a free list and go back when the request finishes.
The allocator
class BlockAllocator:
"""Free list plus reference counts. (Your engine: Chapter 25)"""
def __init__(self, num_blocks):
self.free = deque(range(num_blocks))
self.refs = [0] * num_blocks
def allocate(self):
"""(Your engine: Chapter 25)"""
if not self.free:
raise MemoryError("KV block pool exhausted")
block = self.free.popleft()
self.refs[block] = 1
return block
def share(self, block):
self.refs[block] += 1
def release(self, block):
"""(Your engine: Chapter 25)"""
self.refs[block] -= 1
if self.refs[block] == 0:
self.free.append(block)
elif self.refs[block] < 0:
raise RuntimeError(f"Block {block} released more times than it was referenced")
@property
def num_free(self):
return len(self.free)
Every block carries a reference count: the number of block tables (and caches, below) that point to it. release decrements it, and only a count of zero returns the block to the free list.
The paged cache
class PagedKVCache:
"""Implements the book's cache protocol over a block pool. (Your engine: Chapter 25)
Before a forward pass the caller names the sequences in the batch (set_batch) and
reserves room for the new tokens (reserve); allocation failures therefore happen
before any GPU work is launched.
"""
def __init__(self, layers, num_blocks, block_size, kv_heads, head_dim, device="cpu", dtype=torch.float32):
shape = (num_blocks, kv_heads, block_size, head_dim)
self.k_pool = [torch.zeros(shape, device=device, dtype=dtype) for _ in range(layers)]
self.v_pool = [torch.zeros(shape, device=device, dtype=dtype) for _ in range(layers)]
self.block_size, self.device = block_size, device
self.allocator = BlockAllocator(num_blocks)
self.tables, self.lengths = {}, {}
self.batch = []
def blocks_needed(self, tokens):
return -(-tokens // self.block_size)
def create(self, seq):
if seq in self.tables:
raise ValueError(f"Sequence {seq} already exists")
self.tables[seq], self.lengths[seq] = [], 0
def reserve(self, seq, new_tokens):
"""Make blocks exist for positions [length, length + new_tokens), copying any shared
block that is about to be written (copy-on-write). (Your engine: Chapter 25)"""
table, start = self.tables[seq], self.lengths[seq]
end = start + new_tokens
first_written = start // self.block_size
for index in range(first_written, min(len(table), self.blocks_needed(end))):
block = table[index]
if self.allocator.refs[block] > 1: # shared: copy before writing
fresh = self.allocator.allocate()
for pool in self.k_pool + self.v_pool:
pool[fresh].copy_(pool[block])
self.allocator.release(block)
table[index] = fresh
while len(table) < self.blocks_needed(end):
table.append(self.allocator.allocate())
def set_batch(self, seqs):
self.batch = list(seqs)
width = max(len(self.tables[s]) for s in self.batch)
rows = [self.tables[s] + [0] * (width - len(self.tables[s])) for s in self.batch]
self.table_tensor = torch.tensor(rows, device=self.device, dtype=torch.long)
def update(self, layer, k, v, positions, rows=None):
"""Scatter new K/V into their blocks, then gather each sequence's history (reference path). (Your engine: Chapter 25)"""
if positions.ndim == 1:
positions = positions.expand(k.shape[0], -1)
batch_index = torch.arange(k.shape[0], device=k.device)[:, None]
blocks = self.table_tensor[batch_index, positions // self.block_size] # [B, T]
offsets = positions % self.block_size
self.k_pool[layer][blocks, :, offsets] = k.transpose(1, 2)
self.v_pool[layer][blocks, :, offsets] = v.transpose(1, 2)
if layer == len(self.k_pool) - 1: # the last layer commits the lengths
for b, seq in enumerate(self.batch):
self.lengths[seq] = max(self.lengths[seq], int(positions[b].max()) + 1)
span = int(positions.max()) + 1
logical = torch.arange(span, device=k.device)
gather_blocks = self.table_tensor[:, logical // self.block_size] # [B, span]
keys = self.k_pool[layer][gather_blocks, :, logical % self.block_size] # [B, span, H, D]
values = self.v_pool[layer][gather_blocks, :, logical % self.block_size]
return keys.transpose(1, 2), values.transpose(1, 2), logical
def fork(self, parent, child):
"""child starts as an alias of parent's blocks (beam search, n>1 sampling, shared prompts). (Your engine: Chapter 25)"""
self.create(child)
self.tables[child] = list(self.tables[parent])
self.lengths[child] = self.lengths[parent]
for block in self.tables[child]:
self.allocator.share(block)
def free(self, seq):
for block in self.tables.pop(seq):
self.allocator.release(block)
self.lengths.pop(seq)
def truncate(self, seq, length):
"""Roll back to `length` tokens, returning now-unused blocks."""
keep = self.blocks_needed(length)
for block in self.tables[seq][keep:]:
self.allocator.release(block)
del self.tables[seq][keep:]
self.lengths[seq] = length
The cache implements the same update protocol as every cache since Chapter 16, so your Qwen3 runs on it unchanged. The flow per step:
reserve(seq, n)before the forward pass makes sure blocks exist for the $n$ new positions. Allocation failures happen here, before any GPU work is launched, where the scheduler can handle them (by waiting, or preempting another request).set_batch(seqs)builds the block-table tensor for this batch, padded to the longest table.updatescatters the new keys and values to(table[p // b], p % b)and returns each sequence’s history. This reference version gathers the history into a contiguous tensor for the existing attention code, which costs a copy; the kernel at the end of the chapter removes it.
Sharing blocks: forks and copy-on-write
Many requests share content. Parallel sampling (n = 4 answers to one prompt) and beam search need several sequences with the same prompt. With block tables, a fork copies the table, not the data, and increments each block’s reference count.
The sequences can then diverge. When one of them is about to write into a block whose count is above 1, reserve first copies that block to a fresh one and points this sequence’s table at the copy: copy-on-write. Full blocks of the prompt stay shared forever; only the partially filled last block gets copied.
{"parent_blocks": 3, "new_blocks_after_4_forks": 0, "new_blocks_after_each_sample_writes_a_token": 4}
A 40-token prompt takes 3 blocks of 16 (the last one half full). Four forks cost no new blocks. When each sample writes its first token, each copies the shared partial block: 4 new blocks in total, instead of 12 for four independent copies of the prompt.
The milestone test checks the subtle part: after a fork, the parent and the child each continue with a different token, and both must produce exactly the logits of running their full sequence from scratch. If copy-on-write is missing, one sequence’s write corrupts the other’s history.
Prefix caching
Most production traffic shares long prefixes: the same system prompt, the same few-shot examples, the same document with different questions, the growing history of a multi-turn chat. Their KV cache is identical, so it can be computed once.
class PrefixCache:
"""Reuse full blocks whose token contents (and everything before them) match exactly.
Keys are the whole token prefix up to the end of a block, so two blocks with the same
text but different histories never collide. The cache holds its own reference on each
block; least-recently-used entries are evicted when the pool runs low.
"""
def __init__(self, cache):
self.cache = cache
self.entries = OrderedDict() # tuple(prefix tokens) -> block
def attach(self, seq, prompt):
"""Create seq with the longest cached prefix of prompt. Returns tokens reused."""
self.cache.create(seq)
size, reused = self.cache.block_size, 0
# Keep at least one prompt token to prefill: its logits predict the first output.
while (reused + size) <= len(prompt) - 1:
key = tuple(prompt[:reused + size])
block = self.entries.get(key)
if block is None:
break
self.entries.move_to_end(key)
self.cache.allocator.share(block)
self.cache.tables[seq].append(block)
reused += size
self.cache.lengths[seq] = reused
return reused
def remember(self, seq, tokens):
"""After prefill, publish seq's full blocks for later requests."""
size = self.cache.block_size
for index in range(len(tokens) // size):
key = tuple(tokens[:(index + 1) * size])
if key not in self.entries:
block = self.cache.tables[seq][index]
self.cache.allocator.share(block)
self.entries[key] = block
def evict(self, blocks_wanted):
"""Drop LRU entries that nobody else is using until enough blocks are free."""
for key in list(self.entries):
if self.cache.allocator.num_free >= blocks_wanted:
break
block = self.entries[key]
if self.cache.allocator.refs[block] == 1:
del self.entries[key]
self.cache.allocator.release(block)
The key for a cached block is the entire token prefix up to the end of that block, not just the block’s own tokens. A block’s keys and values depend on every earlier token, so two blocks with identical text but different histories must never match. (Production engines hash the chain of block contents instead of storing the tuples, and the key must also include anything else that changes activations: the model revision, LoRA adapter, and multimodal inputs.)
attach walks the prompt one full block at a time while the prefix is cached, shares those blocks, and leaves at least one token to prefill, because the prompt’s last position must run to produce the first output’s logits. remember publishes a sequence’s full blocks after prefill. The cache holds its own reference to each block, so cached blocks survive their requests, and evict drops least-recently-used entries nobody else is using when the pool runs low.
Eight requests sharing a 96-token system prompt, each with 8 tokens of its own:
{"prompt_tokens": 832, "tokens_prefilled": 160, "blocks_in_use": 14, "blocks_without_sharing": 56}
The first request prefills all 104 tokens; the other seven prefill only their own 8. That’s 81% less prefill compute, which lowers time to first token directly, and a quarter of the memory. SGLang’s RadixAttention generalizes this to a radix tree over all cached prefixes and schedules requests to maximize hits.
Attention that reads through the table
The gather in update builds a contiguous copy of every sequence’s history at every layer, every step. A paged attention kernel reads K and V directly from the pool instead. Here’s a decode kernel: one program per (sequence, query head), walking the sequence’s block table and applying Chapter 15’s online softmax one block at a time:
@triton.jit
def paged_decode_kernel(q_ptr, k_pool, v_pool, tables_ptr, lengths_ptr, out_ptr,
max_blocks, heads, group, scale,
s_pool_block, s_pool_head, s_pool_slot,
D: tl.constexpr, BLOCK: tl.constexpr):
"""(Your engine: Chapter 25)"""
seq, h = tl.program_id(0), tl.program_id(1)
kvh = h // group
rd = tl.arange(0, D)
slots = tl.arange(0, BLOCK)
q = tl.load(q_ptr + (seq * heads + h) * D + rd).to(tl.float32)
length = tl.load(lengths_ptr + seq)
m = tl.full((), -float("inf"), tl.float32)
l = tl.zeros((), tl.float32)
acc = tl.zeros((D,), tl.float32)
for i in range(0, tl.cdiv(length, BLOCK)):
block = tl.load(tables_ptr + seq * max_blocks + i) # logical block i -> physical block
base = block * s_pool_block + kvh * s_pool_head
valid = i * BLOCK + slots < length
k = tl.load(k_pool + base + slots[:, None] * s_pool_slot + rd[None, :], mask=valid[:, None], other=0.0)
v = tl.load(v_pool + base + slots[:, None] * s_pool_slot + rd[None, :], mask=valid[:, None], other=0.0)
s = tl.sum(k.to(tl.float32) * q[None, :], axis=1) * scale
s = tl.where(valid, s, -float("inf"))
m_new = tl.maximum(m, tl.max(s, axis=0))
alpha = tl.exp(m - m_new)
p = tl.exp(s - m_new)
l = l * alpha + tl.sum(p, axis=0)
acc = acc * alpha + tl.sum(p[:, None] * v.to(tl.float32), axis=0)
m = m_new
tl.store(out_ptr + (seq * heads + h) * D + rd, (acc / l).to(out_ptr.dtype.element_ty))
def paged_decode_attention(q, k_pool, v_pool, block_tables, lengths, block_size):
"""q [B, Hq, D]; pools [N, Hkv, block_size, D]; block_tables int [B, max_blocks]; lengths int [B] (>= 1)."""
check_device(q, k_pool, v_pool, block_tables, lengths)
batch, heads, d = q.shape
out = torch.empty_like(q)
k_pool, v_pool = k_pool.contiguous(), v_pool.contiguous()
paged_decode_kernel[(batch, heads)](q.contiguous(), k_pool, v_pool, block_tables.contiguous(),
lengths.contiguous(), out, block_tables.shape[1], heads,
heads // k_pool.shape[1], 1.0 / math.sqrt(d),
k_pool.stride(0), k_pool.stride(1), k_pool.stride(2),
D=d, BLOCK=block_size)
return out
Compared with the flash kernel, the only new thing is the indirection: tl.load(tables_ptr + seq * max_blocks + i) turns logical block i into a physical block number, and the K/V addresses come from that. The lengths tensor masks the tail of the last block. Grouped-query attention works the same way as before: query head h reads KV head h // group.
The reference paged_decode_attention in paged.py gathers and computes the same result; the test checks the kernel against it with scrambled block tables.
What production paged kernels add: splitting long sequences across programs with a final merge (flash-decoding, Chapter 15), processing all query heads of a KV group together so K/V is loaded once, and FP8 KV caches. FlashInfer and vLLM’s attention backends implement these for prefill, decode and mixed batches.
Choosing the block size
| smaller blocks (8) | larger blocks (32-128) |
|---|---|
| less waste in each sequence’s last block | longer contiguous reads, better kernel efficiency |
| finer-grained prefix sharing | smaller block tables, fewer allocations |
16 is vLLM’s default. Some FlashAttention-3-based paged kernels prefer larger blocks.
Build it
Engine milestone 25: a paged KV cache. Implement BlockAllocator.allocate and release, and PagedKVCache.reserve, update and fork in engine/paged.py, and paged_decode_kernel in engine/kernels/triton_paged.py (set_batch, free, truncate, the prefix cache and the launcher are provided).
pytest tests/test_ch25_paged.py
python run.py paged --impl engine
The tests check reference counts and pool exhaustion, that a Qwen3 running on the paged cache (prefill, then token-by-token) matches the full forward, that a forked parent and child diverge correctly through copy-on-write and return every block when freed, that a second request reuses three cached prefix blocks and still computes the right logits, and your Triton kernel against the reference with non-contiguous block tables.
Stretch exercises
- ★ Plot
paged_utilizationfromrun.py pagedagainst block sizes 1-128 for the same request lengths. Where:experiments/ch25.py(create it), adaptingrun.py’scmd_pagedwith different block sizes. - ★★ Replace
StaticKVCachein Chapter 24’s engine withPagedKVCacheand aPrefixCache, callingreservefor every planned chunk. WhenreserveraisesMemoryError, tryevict, then preempt the newest request. Verify outputs still match solo runs. Where:ContinuousBatchingEngine.__init__/stepinengine/scheduler.py, usingengine.paged. - ★★ Implement $n$-way parallel sampling in the batching engine with
fork: one prefill, then $n$ independently seeded samples. Where: request/slot handling inengine/scheduler.pyandPagedKVCache.forkinengine/paged.py. - ★★★ Write a paged prefill kernel: a block of queries from one sequence attends causally to that sequence’s paged history plus the new tokens. Where: add a prefill kernel and launcher beside
paged_decode_kernelinengine/kernels/triton_paged.py.
Check your understanding
- Why can’t an engine simply allocate exactly the cache each request will need?
- Where does the logical position $p$ of a sequence live in the pool?
- Why is the key of a cached prefix block the whole prefix, not just the block’s tokens?
- When does copy-on-write copy a block, and which block is it usually?
- Why does
attachalways leave at least one prompt token to prefill?
Going deeper
- Kwon et al., Efficient Memory Management for Large Language Model Serving with PagedAttention (SOSP 2023), the vLLM paper; Zheng et al., SGLang: Efficient Execution of Structured Language Model Programs (2024) for RadixAttention.
- GPU Mode L40 (FlashInfer): paged and ragged attention kernels in production; L35 (SGLang).
- PMPP §20.7 (Alleviating the memory requirements of the attention mechanism).
- vLLM’s
vllm/v1/core/block_pool.pyandkv_cache_manager.py, and its paged attention kernels, which follow the structure of this chapter.
26. Speculative decoding
In this chapter
- Why checking several tokens costs a memory-bound model about as much as generating one.
- Draft-then-verify: greedy speculation, and rejection sampling that keeps the output distribution exactly the target model's.
- The expected speedup as a function of acceptance rate, draft length and draft cost.
- Rolling back the KV cache, and the zoo of draft methods: small models, early exits, n-gram lookup, Medusa, EAGLE and multi-token prediction.
You will build
accept_or_correct and speculative_generate in engine/speculative.py: lossless speculative decoding for any two models that share a tokenizer.
Time: 4-6 hours. GPU: recommended for speed measurements.
Verification is almost free
At batch 1, a decode step reads every weight to produce one token’s logits (Chapter 10). Feeding the model five tokens instead of one reads the same weights and returns five positions’ logits: the step is still memory-bound and takes nearly the same time. Chapter 24 used this by batching different users. Speculative decoding uses it along the sequence of one user.
The catch: to feed five tokens, you need to know them, and you only know the next token after computing it. Unless you guess. If a cheap draft model guesses the next $k$ tokens, the expensive target model can check all of them in one forward pass. Where its own predictions agree, those tokens are done. At the first disagreement, the target’s own prediction replaces the guess, so every round produces at least one token, and up to $k+1$.
Greedy speculation
With greedy decoding, the rule is simple. Feed the target [last, d1, d2, d3, d4]. Its logits at each position give its own choice for the next token:
target input: last d1 d2 d3 d4
target's argmax: t1 t2 t3 t4 t5
draft proposed: d1 d2 d3 d4
Accept $d_1$ if $t_1 = d_1$, then $d_2$ if $t_2 = d_2$, and so on. At the first mismatch, take the target’s token $t_i$ instead and stop. If all four match, $t_5$ is a free bonus token. Every committed token is exactly what greedy decoding of the target alone would have produced, so the output is identical; only the number of target calls changes.
Sampling: rejection sampling keeps the distribution exact
With temperature sampling, “agree with the argmax” is too strict and would also change the output distribution. Leviathan et al. (2023) and Chen et al. (2023) showed how to accept draft tokens so that every committed token is distributed exactly as if sampled from the target.
Let $q$ be the draft’s distribution at a position, $p$ the target’s, and $x \sim q$ the draft’s token:
- Accept $x$ with probability $\min!\big(1, p(x)/q(x)\big)$.
- Otherwise sample a replacement from the residual distribution $r(y) = \dfrac{\max(0,\ p(y) - q(y))}{\sum_z \max(0,\ p(z) - q(z))}$, and stop the round.
Why this produces exactly $p$: the probability of outputting $y$ through acceptance is $q(y) \cdot \min(1, p(y)/q(y)) = \min(q(y), p(y))$. The probability of rejecting is $1 - \sum_z \min(q(z), p(z)) = \sum_z \max(0, p(z) - q(z))$, which is exactly the residual’s normalizer. So rejection contributes $\max(0, p(y) - q(y))$. Adding the two:
$$ \min(p(y), q(y)) + \max(0,\ p(y) - q(y)) = p(y). \quad\checkmark $$
Worked example. Vocabulary of three tokens, $p = [0.6, 0.3, 0.1]$, $q = [0.2, 0.3, 0.5]$. The draft proposes token 2 half the time, but the target accepts it with probability $0.1/0.5 = 0.2$. Token 0 is accepted always ($0.6/0.2 > 1$). The residual is $\max(0, p - q) = [0.4, 0, 0]$: every rejection becomes token 0. Total probability of token 0: $0.2$ (proposed and accepted) $+ (0.5 \times 0.8)$ (token 2 rejected) $= 0.6$. The milestone test draws 20,000 samples and checks all three frequencies.
def accept_or_correct(p, q, token, generator=None):
"""One verification step for draft token `token` drawn from q. (Your engine: Chapter 26)
Accept with probability min(1, p[token] / q[token]). On rejection, draw from the
residual max(p - q, 0) / sum(max(p - q, 0)). Returns (token, accepted).
The two cases together produce exactly p: accepted mass min(p, q) plus residual mass
max(p - q, 0) add up to p at every token.
"""
ratio = (p[token] / q[token].clamp_min(1e-20)).clamp(max=1.0)
if torch.rand((), generator=generator, device=p.device) < ratio:
return token, True
residual = (p - q).clamp_min(0)
if residual.sum() <= 0: # p == q numerically: nothing to correct
residual = p
return int(torch.multinomial(residual / residual.sum(), 1, generator=generator)), False
inline std::pair<size_t, bool> accept_or_correct(const std::vector<double>& p, const std::vector<double>& q, size_t x, Rng& rng) {
if (rng.uniform() < std::min(1.0, p[x] / q[x])) return {x, true};
std::vector<double> residual(p.size());
double total = 0;
for (size_t i = 0; i < p.size(); ++i) total += (residual[i] = std::max(0.0, p[i] - q[i]));
double r = rng.uniform() * total;
for (size_t i = 0; i < p.size(); ++i) { if (r < residual[i]) return {i, false}; r -= residual[i]; }
return {x, false};
}
#![allow(unused)]
fn main() {
/// Draft token `x` came from q. Accept with probability min(1, p[x]/q[x]); otherwise draw
/// from the leftover distribution max(p - q, 0), renormalized. The result is distributed as p.
pub fn accept_or_correct(p: &[f64], q: &[f64], x: usize, rng: &mut Rng) -> (usize, bool) {
if (rng.uniform() as f64) < (p[x] / q[x]).min(1.0) {
return (x, true);
}
let residual: Vec<f64> = p.iter().zip(q).map(|(a, b)| (a - b).max(0.0)).collect();
let total: f64 = residual.iter().sum();
let mut r = rng.uniform() as f64 * total;
for (i, mass) in residual.iter().enumerate() {
if r < *mass {
return (i, false);
}
r -= mass;
}
(residual.iter().rposition(|m| *m > 0.0).unwrap_or(x), false)
}
}
How much faster?
Suppose each draft token is accepted independently with probability $\alpha$ (a simplification, but a useful one). A round accepts $i$ tokens with probability $\alpha^i(1-\alpha)$ and always adds one target token, so the expected tokens per round is
$$ E[\text{tokens}] = 1 + \alpha + \alpha^2 + \dots + \alpha^k = \frac{1 - \alpha^{k+1}}{1 - \alpha}. $$
A round costs one target forward plus $k$ draft forwards. If a draft step costs a fraction $c$ of a target step, the speedup over plain decoding is
$$ \text{speedup} = \frac{1 - \alpha^{k+1}}{(1 - \alpha)(1 + k c)} . $$
With $\alpha = 0.8$, $k = 4$ and $c = 0.05$ (a draft about 20× smaller), that’s $3.36 / 1.2 = 2.8\times$. With $\alpha = 0.5$, it’s $1.94 / 1.2 = 1.6\times$. Larger $k$ helps only while $\alpha$ is high: the chance of reaching the $k$-th guess shrinks geometrically, while its cost doesn’t.
def expected_tokens_per_round(alpha, k):
"""With per-token acceptance probability alpha, a round yields (1 - alpha^(k+1)) / (1 - alpha) tokens."""
if alpha >= 1:
return k + 1
return (1 - alpha ** (k + 1)) / (1 - alpha)
The algorithm, with cache rollback
@torch.inference_mode()
def speculative_generate(target, draft, prompt_ids, max_new_tokens, k=4, temperature=0.0,
stop_ids=(), generator=None):
"""Generate with a draft model. Both models must share the tokenizer. (Your engine: Chapter 26)
Invariant at the top of each round: both caches hold every committed token except the
newest one, `last`, which has been chosen but not yet fed to either model.
Returns (new_token_ids, stats).
"""
device = next(target.parameters()).device
ids = torch.as_tensor(prompt_ids, device=device).view(1, -1)
capacity = ids.shape[1] + max_new_tokens + k + 1
t_cache, d_cache = target.new_cache(1, capacity), draft.new_cache(1, capacity)
target(ids[:, :-1], t_cache) if ids.shape[1] > 1 else None
draft(ids[:, :-1], d_cache) if ids.shape[1] > 1 else None
last = ids[:, -1:]
output, rounds, accepted_total, proposed_total = [], 0, 0, 0
while len(output) < max_new_tokens:
rounds += 1
base = t_cache.length
# 1. Draft proposes k tokens, one cheap forward each.
proposals, q_dists, token = [], [], last
for _ in range(k):
logits = draft(token, d_cache)[:, -1]
if temperature == 0:
token = logits.argmax(-1, keepdim=True)
else:
q = _probs(logits, temperature)[0]
q_dists.append(q)
token = torch.multinomial(q, 1, generator=generator)[None]
proposals.append(token)
# 2. Target scores last + all proposals in ONE forward: k+1 next-token distributions.
block = torch.cat([last] + proposals, dim=1)
target_logits = target(block, t_cache)[0] # [k+1, V]
# 3. Accept a prefix of the proposals, then add one corrected or bonus token.
new_tokens, all_accepted = [], True
for i, proposal in enumerate(proposals):
token_id = int(proposal)
if temperature == 0:
best = int(target_logits[i].argmax())
ok, chosen = best == token_id, best
else:
chosen, ok = accept_or_correct(_probs(target_logits[i], temperature), q_dists[i], token_id, generator)
new_tokens.append(chosen)
if not ok:
all_accepted = False
break
accepted = len(new_tokens) if all_accepted else len(new_tokens) - 1
if all_accepted: # all accepted: free bonus token
bonus = target_logits[k]
chosen = int(bonus.argmax()) if temperature == 0 else int(
torch.multinomial(_probs(bonus, temperature), 1, generator=generator))
new_tokens.append(chosen)
accepted_total += accepted
proposed_total += k
# 4. Roll both caches back to "everything committed except the newest token".
t_cache.truncate(base + accepted + 1)
if accepted == k:
draft(proposals[-1], d_cache) # the draft never fed its own last proposal
d_cache.truncate(base + accepted + 1)
for token_id in new_tokens:
output.append(token_id)
if token_id in stop_ids or len(output) == max_new_tokens:
return output, {"rounds": rounds, "acceptance_rate": accepted_total / proposed_total,
"tokens_per_target_call": len(output) / rounds}
last = torch.tensor([[new_tokens[-1]]], device=device)
return output, {"rounds": rounds, "acceptance_rate": accepted_total / max(proposed_total, 1),
"tokens_per_target_call": len(output) / max(rounds, 1)}
The bookkeeping that makes it correct is one invariant: at the top of each round, both caches hold every committed token except the newest one, last, which has been chosen but not yet fed to either model. Then:
- The draft feeds
lastand its own proposals one by one, writing $k$ positions to its cache. - The target feeds
lastand all $k$ proposals at once, writing $k+1$ positions to its cache and returning $k+1$ distributions. - Some prefix of the proposals is accepted, plus one corrected or bonus token.
- Both caches are truncated back to the base length plus the accepted tokens plus
last, the new invariant. Rejected proposals’ keys and values are simply forgotten. That’s why Chapter 16’s cache hastruncate: rollback costs nothing.
One edge case: if all $k$ proposals are accepted, the draft never fed its own last proposal, so it’s fed once more before truncating.
Run it
python run.py speculate --k 4
Both models are random here (a 6-layer tiny Qwen3 as target), so the numbers show mechanics, not real-world speed:
{"draft": "target itself (perfect draft)", "rounds": 7, "acceptance_rate": 1.0, "tokens_per_target_call": 4.57}
{"draft": "first 2 of 6 layers", "identical_to_target_greedy": true, "rounds": 14, "acceptance_rate": 0.339, "tokens_per_target_call": 2.29}
With a perfect draft, every round yields $k+1 = 5$ tokens (the last round is cut short at 32 tokens). The “early exit” draft, the target’s own first two layers followed by its final norm and head, agrees with the full model a third of the time, still halving the target calls. And the output is identical to the target’s greedy output, as it must be. With --k 2 the early-exit draft’s acceptance rose to 44% but tokens per call fell to 1.78; with --k 8 acceptance fell to 21%.
On real models, a well-matched pair does much better: drafts from the same family and training data agree on most “easy” tokens (punctuation, the rest of a word, boilerplate code). Typical acceptance rates are 0.6-0.8 for chat and higher for code.
Where drafts come from
| draft | how it guesses | trade-off |
|---|---|---|
| small model of the same family (Qwen3-0.6B for Qwen3-8B) | its own forward passes | needs the same tokenizer; extra memory and its own KV cache |
| early exit / self-speculation (LayerSkip, 2024) | the target’s first layers plus its head | no extra model; works best if trained for early exit |
| n-gram / prompt lookup | copies the continuation of the last few tokens from earlier in the prompt | free; excellent for editing, summarization, code with repetition |
| Medusa (Cai et al., 2024) | extra heads on the target predict tokens $t+2, t+3, \dots$ | small trained heads; verify a tree of candidates |
| EAGLE 1-3 (Li et al., 2024-25) | a one-layer model predicting the target’s next hidden state | among the best acceptance rates; trained per target |
| multi-token prediction (MTP) | heads trained with the model (DeepSeek-V3, Qwen3-Next, Qwen3.8-Flash-Next) | free at inference if the checkpoint ships them |
Tree verification checks several candidate continuations at once: the drafts form a tree, and one target forward pass evaluates all branches using an attention mask where each node sees only its ancestors. Your position-based attention (Chapter 5) supports this directly: give each node its depth as its position and an allowed mask from the tree.
When it doesn’t help
Speculation turns spare memory bandwidth into tokens. At large batch sizes, decode is already compute-bound (Chapter 24), so verifying $k+1$ tokens per request costs $k+1$ times as much, and rejected tokens are pure waste. Production engines reduce $k$ or turn speculation off as the batch grows. It also helps less when acceptance is low: high-temperature creative writing, or a draft trained on different data.
Build it
Engine milestone 26: speculative decoding. Implement accept_or_correct and speculative_generate in engine/speculative.py (expected_tokens_per_round is provided).
pytest tests/test_ch26_speculative.py
python run.py speculate --impl engine --k 4
The tests check that greedy speculation with an imperfect draft reproduces the target’s greedy output exactly, that a perfect draft accepts every proposal, that rejection sampling reproduces $p$ over 20,000 draws, and the expected-tokens formula.
Stretch exercises
- ★ Measure acceptance rate and wall-clock speedup for Qwen3-1.7B (or 8B) as target and Qwen3-0.6B as draft on a GPU, for $k = 2, 4, 6$. Compare with the formula. Where:
experiments/ch26.py(create it), callingengine.speculative.speculative_generate. - ★★ Implement prompt lookup decoding: find the most recent earlier occurrence of the last 3 tokens in the context and propose the $k$ tokens that followed. Measure acceptance on a “rewrite this paragraph” task. Where: add a prompt-lookup draft helper in
engine/speculative.py. - ★★ Make $k$ adaptive: track a running acceptance rate and choose the $k$ that maximizes the expected speedup each round. Where: the round loop in
speculative_generateinengine/speculative.py. - ★★★ Implement tree verification for greedy decoding: the draft proposes its top-2 tokens at each of 3 depths (a tree of 14 nodes), verify all branches in one target call with a tree attention mask, and commit the longest accepted path. Where: add tree drafting/verification in
engine/speculative.py; pass tree masks through the target model toengine.attention.causal_attention.
Check your understanding
- Why does verifying five tokens take about as long as generating one, at batch 1?
- Show that the accept-or-resample rule outputs token $y$ with probability exactly $p(y)$.
- Why does every round produce at least one token?
- Why must both caches be truncated after a round, and to what length?
- Why can speculation make a large-batch server slower?
Going deeper
- Leviathan, Kalman and Matias, Fast Inference from Transformers via Speculative Decoding (2023); Chen et al., Accelerating Large Language Model Decoding with Speculative Sampling (DeepMind, 2023).
- GPU Mode L22 (Hacker’s guide to speculative decoding in vLLM, Cade Daniels).
- PMPP §20.6 names batching and speculative decoding as the two main ways to raise the arithmetic intensity of the generation phase.
- Cai et al., Medusa (2024); Li et al., EAGLE, EAGLE-2, EAGLE-3 (2024-25); Elhoushi et al., LayerSkip (2024); Gloeckle et al., Better & Faster Large Language Models via Multi-token Prediction (2024); DeepSeek-V3 technical report (MTP).
27. Mixture of experts
In this chapter
- How a mixture-of-experts layer decouples a model's size from the compute per token.
- The router: softmax, top-k, renormalization, and a gated shared expert.
- Dispatch: from a readable per-expert loop to the sort-by-expert layout that grouped matrix multiplications consume.
- The economics of MoE inference: why batch-1 decode is fast, why memory is the price, and what happens as the batch grows.
- Loading a real Qwen3-MoE checkpoint and matching Hugging Face's output.
You will build
TopKRouter.forward and Experts.forward_grouped in engine/moe.py. Your Qwen3 grows into Qwen3-MoE, and the same block becomes Flash-Next's MoE in Chapter 30.
Time: 5-6 hours. GPU: optional.
Size without compute
In a dense transformer, every parameter takes part in every token. Doubling the parameters doubles the compute per token and the bytes read per decode step. A mixture of experts (MoE) breaks that link. Each MLP is replaced by $E$ smaller MLPs, the experts, and a router that sends each token to only $k$ of them.
dense layer: x ──► MLP ──► y every token uses all of it
MoE layer: x ──► router ──► top-k of E experts ──► weighted sum ──► y
Qwen3-30B-A3B has 128 experts per layer and uses 8 per token. It stores 30.5B parameters but each token is processed by about 3.3B active ones. The model knows roughly what a 30B model knows, and computes like a 3B one.
python run.py moe
{"model": "Qwen3-30B-A3B", "expert_params_stored_B": 29.0, "expert_params_active_B": 1.82, "active_fraction": 0.063}
{"model": "Qwen3-235B-A22B", "expert_params_stored_B": 227.15, "expert_params_active_B": 14.24, "active_fraction": 0.063}
{"model": "Qwen3.8-Flash-Next", "expert_params_stored_B": 121.09, "expert_params_active_B": 2.66, "active_fraction": 0.022}
(Expert parameters only; attention and embeddings add the rest.) Flash-Next, the capstone’s target, routes each token to 10 of 512 experts plus a shared expert: it stores 121B parameters of experts and uses 2.2% of them per token.
The router
The router is one linear layer from the hidden state to $E$ scores. Qwen-family models then use:
- Softmax over all experts, in FP32: $\pi = \operatorname{softmax}(W_r x)$.
- Top-k: keep the $k$ largest probabilities and their expert indices.
- Renormalize (when
norm_topk_probis set): divide the kept weights by their sum so they add to 1.
The layer’s output is the weighted sum of the chosen experts’ outputs: $y = \sum_{i \in \text{top-}k} w_i, \text{expert}_i(x)$.
Example. Four experts, $k = 2$, router probabilities $[0.1, 0.5, 0.3, 0.1]$. Experts 1 and 2 are chosen with weights $0.5$ and $0.3$; renormalized, they become $0.625$ and $0.375$.
class TopKRouter(nn.Module):
def __init__(self, hidden, experts, top_k, normalize=True):
super().__init__()
self.weight = nn.Parameter(torch.zeros(experts, hidden))
self.top_k, self.normalize = top_k, normalize
def forward(self, x):
"""x [N, D] -> (logits [N, E], weights [N, k], experts [N, k]). (Your engine: Chapter 27)
Softmax over all experts in FP32, keep the k largest, and (if normalize) rescale the
kept probabilities to sum to 1.
"""
logits = F.linear(x, self.weight)
probs = torch.softmax(logits, dim=-1, dtype=torch.float32)
weights, experts = probs.topk(self.top_k, dim=-1)
if self.normalize:
weights = weights / weights.sum(-1, keepdim=True)
return logits, weights.to(x.dtype), experts
Two details matter for exact parity. The softmax must be in FP32: the router’s decision is discrete, so a BF16 rounding difference can swap two near-tied experts and change the output completely. And ties in topk are broken by index, which matters when comparing against another implementation (Chapter 29 meets this problem in earnest).
Other routers exist: DeepSeek-V3 uses a sigmoid score per expert with a learned bias for load balancing, and some models route groups of experts first. Each checkpoint documents its own; from_hf refuses configurations it doesn’t implement.
The shared expert
Some models add one shared expert that every token uses, alongside the routed ones. In Qwen3-Next and Flash-Next, its output is scaled by a learned gate, $\sigma(w_s^\top x)$, so each token decides how much of the shared knowledge to mix in:
$$ y = \sum_{i \in \text{top-}k} w_i, \text{expert}_i(x) + \sigma(w_s^\top x), \text{shared}(x). $$
class SparseMoeBlock(nn.Module):
"""Router + routed experts (+ an optional shared expert whose output is scaled by
sigmoid(shared_expert_gate(x)), as in Qwen3-Next and Flash-Next)."""
def __init__(self, hidden, experts, top_k, intermediate, shared_intermediate=0, normalize=True):
super().__init__()
self.gate = TopKRouter(hidden, experts, top_k, normalize)
self.experts = Experts(experts, hidden, intermediate)
self.shared_expert = SwiGLU(hidden, shared_intermediate) if shared_intermediate else None
self.shared_expert_gate = nn.Linear(hidden, 1, bias=False) if shared_intermediate else None
self.dispatch = "grouped"
def forward(self, x):
shape = x.shape
flat = x.reshape(-1, shape[-1])
_, weights, experts = self.gate(flat)
run = self.experts.forward_grouped if self.dispatch == "grouped" else self.experts.forward_loop
out = run(flat, weights, experts)
if self.shared_expert is not None:
out = out + torch.sigmoid(self.shared_expert_gate(flat)) * self.shared_expert(flat)
return out.reshape(shape)
Dispatch: getting tokens to their experts
Experts are stored stacked: one tensor gate_up_proj [E, 2I, D] holding every expert’s gate and up projections, and one down_proj [E, D, I]. Checkpoints usually store each expert’s matrices separately (mlp.experts.17.gate_proj.weight), and the loader stacks them.
class Experts(nn.Module):
"""E SwiGLU experts in two stacked tensors: gate_up_proj [E, 2I, D] (gate rows, then up rows)
and down_proj [E, D, I]. One tensor per kind keeps the layout friendly to grouped kernels."""
def __init__(self, experts, hidden, intermediate):
super().__init__()
self.gate_up_proj = nn.Parameter(torch.empty(experts, 2 * intermediate, hidden))
self.down_proj = nn.Parameter(torch.empty(experts, hidden, intermediate))
nn.init.normal_(self.gate_up_proj, std=0.02)
nn.init.normal_(self.down_proj, std=0.02)
@property
def num_experts(self):
return self.gate_up_proj.shape[0]
def expert(self, e, x):
gate, up = F.linear(x, self.gate_up_proj[e]).chunk(2, dim=-1)
return F.linear(F.silu(gate) * up, self.down_proj[e])
def forward_loop(self, x, weights, experts):
"""Readable reference: for each expert, gather its tokens, run it, scatter-add back."""
out = torch.zeros_like(x)
for e in experts.unique().tolist():
token, slot = torch.where(experts == e)
out.index_add_(0, token, self.expert(e, x[token]) * weights[token, slot, None])
return out
def forward_grouped(self, x, weights, experts):
"""Sort the N*k assignments by expert so each expert's rows are contiguous: this is the
layout a grouped GEMM consumes (one launch, many independent matmuls). (Your engine: Chapter 27)"""
n, k = experts.shape
flat_expert = experts.reshape(-1)
order = flat_expert.argsort(stable=True)
token_of = order // k # which token each sorted assignment came from
counts = torch.bincount(flat_expert, minlength=self.num_experts).tolist()
rows = x[token_of] # [N*k, D] grouped by expert
results, start = [], 0
for e, count in enumerate(counts): # a grouped GEMM does these as one kernel
if count:
results.append(self.expert(e, rows[start:start + count]))
start += count
expert_out = torch.cat(results) * weights.reshape(-1)[order, None]
return torch.zeros_like(x).index_add_(0, token_of, expert_out)
The loop dispatch is the readable reference: for each expert used in the batch, find its tokens, run the expert on them, and add the weighted result back to those tokens’ rows. Each expert is a separate small matmul, and the loop body runs once per active expert, with a host sync from unique().
The grouped dispatch prepares the layout a fast kernel needs:
- Flatten the $N \times k$ assignments and sort them by expert (
argsort, stable). Now each expert’s rows are contiguous. - Permute the token rows into that order (
x[token_of]), duplicating each token $k$ times. - Run each expert on its contiguous slice. A grouped GEMM kernel does all of these as one launch: many independent matmuls of different sizes.
- Multiply by the routing weights and unpermute with
index_add_, summing each token’s $k$ contributions.
On a CPU, with 512 tokens, 32 experts and $k = 4$:
{"tokens": 512, "top_k": 4, "assignments": 2048, "busiest_expert": 80, "idlest_expert": 52, "ideal_per_expert": 64}
{"dispatch": "loop", "ms": 16.1, "checksum": 1027.975}
{"dispatch": "grouped", "ms": 9.3, "checksum": 1027.975}
Same result, and even the Python-level grouped version is faster. On a GPU, production MoE layers fuse the permutation, the grouped GEMMs and the unpermutation into a few kernels (vLLM’s fused_moe Triton kernel, MegaBlocks, DeepGEMM).
The first line shows load imbalance: with a random router, one expert got 80 tokens and another 52, against an ideal of 64. Trained routers are pushed toward balance during training, with an auxiliary loss (Switch Transformer: $E \sum_i f_i P_i$, the fraction of tokens routed to expert $i$ times its mean probability) or a bias adjusted on the fly (DeepSeek-V3’s auxiliary-loss-free balancing). At inference, imbalance means the busiest expert sets the step time.
The economics of MoE inference
Batch-1 decode is fast. A decode step reads only the experts its token uses. For Qwen3-30B-A3B in BF16, that’s about 6.6 GB per token (3.3B active parameters × 2 bytes), not 61 GB: decode runs at the speed of a 3B dense model.
Memory is the price. All experts must be resident, because the next token may need any of them. A 30B MoE needs 30B parameters of memory, the same as a 30B dense model. That’s why MoE models pair so well with quantization (experts are most of the bytes, Chapter 20) and with large-memory machines: the 128 GB of unified memory in a DGX Spark holds Flash-Next with 4-bit experts.
Larger batches touch more experts. With $B$ tokens each choosing $k$ of $E$ experts roughly uniformly, a layer touches about $E,\big(1 - (1 - k/E)^B\big)$ distinct experts:
| batch | Qwen3-30B-A3B ($E = 128$, $k = 8$) | Flash-Next ($E = 512$, $k = 10$) |
|---|---|---|
| 1 | 8 | 10 |
| 8 | 52 | 75 |
| 32 | 112 | 240 |
| 128 | 128 (all) | 471 |
By batch 32, a decode step of Qwen3-30B-A3B reads almost every expert, the full 61 GB, while each expert does only a little work (about $Bk/E = 2$ tokens each). MoE decode at moderate batch is therefore more memory-bound per token than a dense model with the same active parameters. Throughput still grows with batch until each expert gets enough tokens to be compute-bound, which needs batches in the hundreds or thousands. That’s why large MoE deployments use expert parallelism: experts spread over many GPUs, tokens exchanged with all-to-all communication, so that each GPU holds few experts and gets many tokens for each (Chapter 41).
Prefill is easy. Thousands of prompt tokens give every expert plenty of rows; grouped GEMMs run near peak.
Qwen3-MoE
Qwen3-MoE is Qwen3 with every MLP replaced by a SparseMoeBlock, and nothing else changed. Your Chapter 17 code gives you everything but the block:
@dataclass
class Qwen3MoeConfig(Qwen3Config):
num_experts: int = 8
num_experts_per_tok: int = 2
moe_intermediate_size: int = 64
norm_topk_prob: bool = True
@classmethod
def from_hf(cls, raw):
if raw.get("model_type") != "qwen3_moe":
raise ValueError("Expected model_type 'qwen3_moe'")
if raw.get("mlp_only_layers") or raw.get("decoder_sparse_step", 1) != 1:
raise ValueError("Dense layers inside a Qwen3-MoE stack are not implemented")
raw = dict(raw)
if "num_experts" not in raw and "num_local_experts" in raw: # the name Transformers 5 writes
raw["num_experts"] = raw["num_local_experts"]
dense = Qwen3Config.from_hf({**raw, "model_type": "qwen3"})
own = [f.name for f in fields(cls) if not hasattr(Qwen3Config, f.name)]
missing = [name for name in own if name not in raw]
if missing: # never fall back to defaults for the model's shape
raise ValueError(f"config.json lacks {missing}")
return cls(**vars(dense), **{name: raw[name] for name in own})
class Qwen3MoeLayer(Qwen3Layer):
def __init__(self, cfg):
super().__init__(cfg)
self.mlp = SparseMoeBlock(cfg.hidden_size, cfg.num_experts, cfg.num_experts_per_tok,
cfg.moe_intermediate_size, 0, cfg.norm_topk_prob)
class Qwen3Moe(Qwen3):
"""Qwen3 with every MLP replaced by a routed-expert block."""
def __init__(self, cfg):
super().__init__(cfg)
self.model.layers = nn.ModuleList(Qwen3MoeLayer(cfg) for _ in range(cfg.num_hidden_layers))
@torch.no_grad()
def load_qwen3_moe(directory, device="cpu", dtype=torch.bfloat16):
cfg = Qwen3MoeConfig.from_hf(read_config(directory))
model = Qwen3Moe(cfg).to(device=device, dtype=dtype)
params = dict(model.named_parameters())
pending = {}
for name, value in snapshot_tensors(directory):
parts = name.split(".")
if ".mlp.experts." in name: # model.layers.L.mlp.experts.E.kind.weight
layer, e, kind = int(parts[2]), int(parts[5]), parts[6]
pending.setdefault(layer, {})[(e, kind)] = value
continue
if name == "lm_head.weight" and cfg.tie_word_embeddings:
continue
if name not in params:
raise ValueError(f"Unexpected tensor {name}")
assign(params[name], value, name)
for layer, tensors in pending.items():
gate_up, down = stack_expert_tensors(tensors, cfg.num_experts, cfg.moe_intermediate_size)
block = model.model.layers[layer].mlp.experts
assign(block.gate_up_proj, gate_up, f"layer {layer} gate_up_proj")
assign(block.down_proj, down, f"layer {layer} down_proj")
return model.eval()
The loader collects the per-expert tensors as it streams the safetensors shards, then stacks them per layer. Qwen3-30B-A3B is 61 GB in BF16, so this chapter’s tests build tiny random Qwen3-MoE checkpoints, save them in Hugging Face’s format, and compare your model’s logits with transformers’ Qwen3MoeForCausalLM, with and without renormalization. The largest logit difference is about $10^{-7}$ in FP32, on logits of magnitude 0.5.
Note
Writing that test exposed a real bug. Transformers 5 saves the expert count as
num_local_experts, while Qwen’s published configs saynum_experts. The first version offrom_hfsilently fell back to its default of 8 experts and passed every test that happened to use 8. Now it accepts both names and refuses a config that lacks the field. Two lessons: never let a model’s shape come from a default, and make test models use non-default sizes everywhere.
def moe_parameter_split(hidden, layers, experts, top_k, intermediate, shared_intermediate=0):
"""(stored expert parameters, active expert parameters per token) for the MLP part only."""
per_expert = 3 * hidden * intermediate
shared = 3 * hidden * shared_intermediate
stored = layers * (experts * per_expert + shared + experts * hidden)
active = layers * (top_k * per_expert + shared + experts * hidden)
return stored, active
Build it
Engine milestone 27: mixture of experts. Implement TopKRouter.forward and Experts.forward_grouped in engine/moe.py (the loop dispatch, the block, the model and the loader are provided).
pytest tests/test_ch27_moe.py
python run.py moe --impl engine
The tests check the router’s selection and renormalization, that grouped dispatch equals the loop for random routings (including experts that receive no tokens), the parameter split, and logit parity with Hugging Face’s Qwen3-MoE on saved tiny checkpoints. If you have the memory, load Qwen/Qwen3-30B-A3B with load_model and chat with it through your Chapter 18 engine.
Stretch exercises
- ★ Run a tiny Qwen3-MoE on 1,000 tokens of text and record every layer’s expert counts. Is the load more balanced in early or late layers? (Use a real checkpoint if you can; random routers are uninformative.) Where:
experiments/ch27.py(create it), registering hooks onTopKRouterinengine/moe.py. - ★★ Write a Triton grouped GEMM: one program per (expert, output tile), with a prefix-sum array of row offsets telling each program where its expert’s rows start. Where: new
engine/kernels/triton_grouped.py, called byExperts.forward_groupedinengine/moe.py. - ★★ Quantize only the experts to 4 bits with Chapter 20’s
QuantLinearlogic adapted to stacked tensors, leaving attention and the router in BF16. Measure memory and logit error. Where: add quantized stacked-expert storage/dispatch inengine/moe.py, usingengine.quant. - ★★★ Expert offloading: keep experts in CPU memory, keep a GPU cache of the most recently used ones, and copy missing experts on demand. Measure cache hit rates over a conversation. Where: add an expert-cache variant in
engine/moe.py; Chapter 40’sengine/offload.pyprovides the later integration.
Check your understanding
- Why does an MoE’s decode speed at batch 1 depend on active parameters, while its memory depends on total parameters?
- Why must the router’s softmax run in FP32?
- What does sorting the assignments by expert buy you?
- Why does a moderate batch make MoE decode read nearly all expert weights?
- What problem does a load-balancing loss solve, and why does imbalance matter at inference too?
Going deeper
- Shazeer et al., Outrageously Large Neural Networks: The Sparsely-Gated Mixture-of-Experts Layer (2017); Fedus, Zoph and Shazeer, Switch Transformers (2021); Jiang et al., Mixtral of Experts (2024); DeepSeek-AI, DeepSeekMoE (2024) and the DeepSeek-V3 technical report (fine-grained and shared experts, auxiliary-loss-free balancing).
- Gale et al., MegaBlocks: Efficient Sparse Training with Mixture-of-Experts (2022), on grouped and block-sparse kernels.
- The Qwen3 technical report (2025) for Qwen3-MoE, and the Qwen3-Next model card for the gated shared expert.
- vLLM’s
fused_moeTriton kernel and SGLang’s MoE runner; GPU Mode L11 (Sparsity) for the broader picture of sparse computation on GPUs.
28. Linear attention and Gated DeltaNet
In this chapter
- Why softmax attention's cost grows with context, and how removing the softmax turns attention into a recurrence with a fixed-size state.
- The state as an associative memory, why additive writes interfere, and how decay and the delta rule fix it.
- Gated DeltaNet: the recurrent form for decode and the chunked form, all matrix multiplications, for prefill.
- The full Qwen3.5 / Flash-Next linear-attention layer: short convolution, gates, and the state each request carries.
You will build
engine/gdn.py: the recurrent and chunked Gated DeltaNet, the causal convolution with carried state, and the complete layer, matching Hugging Face's Flash-Next layer.
Time: 6-8 hours. GPU: not needed.
The cost of remembering everything
Softmax attention keeps every past key and value. At decode, each new token reads the whole KV cache: time and memory per token grow linearly with context, and prefill grows quadratically. Chapter 16’s formula put Qwen3-0.6B’s cache at 3.5 GiB for a 32k context, three times the model. At the million-token contexts that agents and long documents want, the cache, not the weights, is the model.
Could a layer instead keep a fixed-size summary of the past, updated once per token, like an RNN? That’s linear attention. Flash-Next uses it in 36 of its 48 layers.
Removing the softmax
Attention for query $t$ is $o_t = \sum_{i \le t} \frac{\exp(q_t \cdot k_i)}{Z}, v_i$. The exponential couples $q_t$ and $k_i$ inside each term, so nothing can be precomputed. Replace $\exp(q \cdot k)$ by a plain dot product (or by $\phi(q) \cdot \phi(k)$ for some feature map $\phi$), and the sum factorizes (Katharopoulos et al., 2020, Transformers are RNNs):
$$ o_t = \sum_{i \le t} (q_t^\top k_i), v_i = q_t^\top \underbrace{\Big(\sum_{i \le t} k_i v_i^\top\Big)}_{S_t ,\in, \mathbb{R}^{d_k \times d_v}} . $$
$S_t$ is a $d_k \times d_v$ matrix, the state, and it updates with one outer product per token:
$$ S_t = S_{t-1} + k_t v_t^\top, \qquad o_t = S_t^\top q_t . $$
Per-token cost and memory are now constant, whatever the context. Prefill is linear in $T$ instead of quadratic.
The state is an associative memory
Think of $S$ as a memory that stores value $v$ under key $k$. Reading with a query equal to a stored key returns $S^\top k_j = \sum_i (k_i \cdot k_j), v_i$: the right value, if keys are orthonormal, plus interference from every other stored value in proportion to how similar its key is. A $d_k$-dimensional state can hold at most $d_k$ orthogonal keys. Writing more than that, and real sequences are thousands of tokens, the memories blur together.
python run.py linear
The first part stores $n$ random key-value pairs in a 64 × 64 state and reads each back with its key (unit-norm keys, relative error of the retrieved value):
{"pairs": 16, "key_dim": 64, "additive_error_all": 0.45, "delta_error_all": 0.303, "delta_error_last_8": 0.224}
{"pairs": 64, "key_dim": 64, "additive_error_all": 0.967, "delta_error_all": 0.714, "delta_error_last_8": 0.23}
{"pairs": 256, "key_dim": 64, "additive_error_all": 2.013, "delta_error_all": 1.199, "delta_error_last_8": 0.282}
Plain additive writes degrade steadily, and past 64 pairs the “retrieved” value is mostly noise. Two ideas fix this, and Gated DeltaNet uses both.
Forget: decay
Multiply the state by a decay $\alpha_t \in (0, 1)$ before each write: $S_t = \alpha_t S_{t-1} + k_t v_t^\top$. Old memories fade, making room for new ones. If $\alpha_t$ is computed from the input (a gate), the model can decide per token how much to forget: keep everything inside a sentence, wipe the slate at a document boundary. This is the idea behind RetNet, GLA and Mamba-2.
Correct: the delta rule
Instead of adding $v_t$ blindly, first ask what the memory currently returns for $k_t$, and write only the error:
$$ S_t = S_{t-1} + \beta_t, k_t \big(v_t - S_{t-1}^\top k_t\big)^\top . $$
This is the delta rule (Widrow-Hoff, 1960; DeltaNet, Schlag et al., 2021): one step of gradient descent on $\tfrac12 \lVert S^\top k_t - v_t \rVert^2$ with learning rate $\beta_t$. With $\beta = 1$ and a unit key, reading $k_t$ afterwards returns exactly $v_t$: the old association along $k_t$ is replaced, not piled on. That’s the “delta” columns above: the most recent pairs are retrieved well even at 256 pairs, and overall error is much lower.
Worked example (the milestone test). $k = q = [1, 0]$, $v = [2, 4]$, empty state. With $\beta = 0.5$: the prediction is $[0, 0]$, the error is $[2, 4]$, and $S$ gains $0.5 \cdot [1, 0]^\top [2, 4]$, so its first row is $[1, 2]$ and the output is $[1, 2]$. Repeat with $\beta = 1$: the prediction is now $[1, 2]$, the error $[1, 2]$, and the first row becomes $[2, 4]$, exactly $v$.
Gated DeltaNet
Combine both (Yang, Kautz and Hatamizadeh, 2024):
$$ S_t = \alpha_t S_{t-1} + \beta_t, k_t \big(v_t - \alpha_t S_{t-1}^\top k_t\big)^\top, \qquad o_t = S_t^\top q_t , $$
with $\alpha_t = e^{g_t}$ for a learned log-decay $g_t \le 0$, and $q, k$ L2-normalized (so that $\beta \le 1$ keeps updates stable).
The recurrent form
Decode processes one token at a time, so the recurrence is exactly what it needs:
def recurrent_gated_delta_rule(q, k, v, g, beta, state=None):
"""One token at a time. (Your engine: Chapter 28)
q, k [B, H, T, Dk] (already L2-normalized; q also scaled by Dk^-0.5); v [B, H, T, Dv];
g, beta [B, H, T]; state [B, H, Dk, Dv] or None. Returns (o [B, H, T, Dv], final state).
"""
b, h, t, dk = k.shape
dv = v.shape[-1]
S = torch.zeros(b, h, dk, dv, device=q.device, dtype=torch.float32) if state is None else state.float()
q, k, v, g, beta = (x.float() for x in (q, k, v, g, beta))
out = torch.empty(b, h, t, dv, device=q.device, dtype=torch.float32)
for i in range(t):
S = S * g[:, :, i, None, None].exp()
prediction = torch.einsum("bhk,bhkv->bhv", k[:, :, i], S)
delta = (v[:, :, i] - prediction) * beta[:, :, i, None]
S = S + torch.einsum("bhk,bhv->bhkv", k[:, :, i], delta)
out[:, :, i] = torch.einsum("bhk,bhkv->bhv", q[:, :, i], S)
return out, S
// S is [dk][dv] row-major: decay, predict v from k, write beta * error along k, read with q.
inline std::vector<float> delta_step(std::vector<float>& S, const std::vector<float>& q, const std::vector<float>& k,
const std::vector<float>& v, float decay, float beta) {
size_t dv = v.size();
for (float& s : S) s *= decay;
std::vector<float> pred(dv, 0.f), out(dv, 0.f);
for (size_t i = 0; i < k.size(); ++i) for (size_t j = 0; j < dv; ++j) pred[j] += k[i] * S[i * dv + j];
for (size_t i = 0; i < k.size(); ++i) for (size_t j = 0; j < dv; ++j) S[i * dv + j] += beta * k[i] * (v[j] - pred[j]);
for (size_t i = 0; i < q.size(); ++i) for (size_t j = 0; j < dv; ++j) out[j] += q[i] * S[i * dv + j];
return out;
}
#![allow(unused)]
fn main() {
/// S is [dk][dv] stored row-major. Decay, predict v from k, write the scaled error, read with q.
pub fn delta_step(s: &mut [f32], q: &[f32], k: &[f32], v: &[f32], decay: f32, beta: f32) -> Vec<f32> {
let dv = v.len();
s.iter_mut().for_each(|x| *x *= decay);
let mut prediction = vec![0.0; dv];
for (i, ki) in k.iter().enumerate() {
for j in 0..dv {
prediction[j] += ki * s[i * dv + j];
}
}
for (i, ki) in k.iter().enumerate() {
for j in 0..dv {
s[i * dv + j] += beta * ki * (v[j] - prediction[j]);
}
}
let mut out = vec![0.0; dv];
for (i, qi) in q.iter().enumerate() {
for j in 0..dv {
out[j] += qi * s[i * dv + j];
}
}
out
}
}
Each step is a few $d_k \times d_v$ elementwise operations and mat-vecs per head: memory-bound on the state, but the state is small (16 KiB per head in FP32 for $d_k = d_v = 64$) and doesn’t grow. The state stays in FP32 even when the model runs in BF16, because it accumulates over thousands of steps (Chapter 13).
The chunked form: recurrence as matrix multiplication
For prefill, a loop of $T$ sequential steps wastes a GPU. The chunked algorithm processes the sequence in chunks of $C$ tokens (64 is typical). Within a chunk, it unrolls the recurrence algebraically into matrix products; across chunks, it carries the state.
Let $G_i = \sum_{j \le i} g_j$ be the cumulative log-decay within the chunk, and $D_{ij} = e^{G_i - G_j}$ for $j \le i$ (the decay from position $j$ to $i$). Unrolling shows that the values actually written, $U$ (each $v_i$ minus what the memory predicted for $k_i$, scaled by $\beta_i$), satisfy a unit lower-triangular linear system: each position’s write depends on the writes before it in the chunk. Solving that system (the “UT transform”) turns the sequential dependency into one triangular solve per chunk. Then the outputs and the next chunk’s state are matrix products:
$$ \begin{aligned} \big(I + \operatorname{strict_tril}(\beta_i, k_i \cdot k_j, D_{ij})\big), U &= \beta \odot V - \beta \odot e^{G} \odot (K S_0) \ O &= e^{G} \odot (Q S_0) + \operatorname{tril}!\big(QK^\top \odot D\big), U \ S_\text{end} &= e^{G_C} S_0 + \textstyle\sum_j e^{G_C - G_j}, k_j U_j^\top \end{aligned} $$
def chunk_gated_delta_rule(q, k, v, g, beta, state=None, chunk=64):
"""Same contract and result as the recurrent form, computed chunk by chunk. (Your engine: Chapter 28)
Inside a chunk, let G_i be the cumulative log-decay up to position i and
D_ij = exp(G_i - G_j) for j <= i (0 above the diagonal). Unrolling the recurrence shows the
values actually written, U, solve a unit-lower-triangular system:
(I + strict_tril(beta_i k_i.k_j D_ij)) U = beta * v - beta * exp(G) * (K S_0)
(the "UT transform"); the chunk's outputs and final state then follow from matmuls:
O = exp(G) * (Q S_0) + (tril(Q K^T * D)) U
S_end = exp(G_last) S_0 + sum_j exp(G_last - G_j) k_j U_j^T
"""
b, h, t, dk = k.shape
dv = v.shape[-1]
q, k, v, g, beta = (x.float() for x in (q, k, v, g, beta))
pad = (-t) % chunk
if pad:
q, k, v = (F.pad(x, (0, 0, 0, pad)) for x in (q, k, v))
g, beta = (F.pad(x, (0, pad)) for x in (g, beta))
n = q.shape[2] // chunk
q, k, v = (x.view(b, h, n, chunk, -1) for x in (q, k, v))
g, beta = g.view(b, h, n, chunk), beta.view(b, h, n, chunk)
G = g.cumsum(-1) # [b,h,n,C]
lower = torch.ones(chunk, chunk, dtype=torch.bool, device=q.device).tril()
decay = (G[..., :, None] - G[..., None, :]).masked_fill(~lower, float("-inf")).exp() # D_ij
k_beta, v_beta = k * beta[..., None], v * beta[..., None]
system = (k_beta @ k.transpose(-1, -2)) * decay # beta_i k_i.k_j D_ij
system = system.tril(-1) + torch.eye(chunk, device=q.device) # I + strictly-lower part
# Solve once for both right-hand sides: the part that depends on S_0 and the part that does not.
w = torch.linalg.solve_triangular(system, k_beta * G[..., None].exp(), upper=False, unitriangular=True)
u = torch.linalg.solve_triangular(system, v_beta, upper=False, unitriangular=True)
attn = (q @ k.transpose(-1, -2)) * decay # causal, decayed Q K^T
S = torch.zeros(b, h, dk, dv, device=q.device) if state is None else state.float()
out = torch.empty_like(v)
for c in range(n):
U = u[:, :, c] - w[:, :, c] @ S # values actually written
out[:, :, c] = (q[:, :, c] * G[:, :, c, :, None].exp()) @ S + attn[:, :, c] @ U
tail = (G[:, :, c, -1:] - G[:, :, c]).exp() # exp(G_last - G_j)
S = S * G[:, :, c, -1, None, None].exp() + (k[:, :, c] * tail[..., None]).transpose(-1, -2) @ U
return out.view(b, h, n * chunk, dv)[:, :, :t], S
The milestone test runs chunk sizes 1, 5, 8 and 64 against the recurrent form on a sequence of 37 tokens (so the last chunk is partial) with a non-zero initial state. They all agree to $10^{-5}$.
Note
The first version of this function kept the wrong triangle of the system matrix, a transposition that’s easy to make when deriving it on paper. Outputs looked plausible, and the chunk-size-1 test even passed, because a 1 × 1 chunk has no off-diagonal part. Testing several chunk sizes, including ones that don’t divide the length, caught it. Test a fast algorithm against the slow one over every boundary case you can construct.
On a laptop CPU, for 4 heads and 1,024 tokens:
{"form": "recurrent", "tokens": 1024, "ms": 132.0}
{"form": "chunked", "tokens": 1024, "ms": 12.2}
{"max_difference": 1.043081283569336e-07, "state_bytes_per_head": 16384, "kv_bytes_per_head_at_1024_tokens": 262144}
Eleven times faster, numerically identical, and the state is 16× smaller than one head’s KV cache at only 1,024 tokens (in BF16). The production kernels (the flash-linear-attention library’s Triton kernels, which Qwen’s models use) fuse these steps per chunk and run at near-matmul speed.
The full layer
The Gated DeltaNet layer in Qwen3.5 and Flash-Next wraps the rule with projections, gates and a short convolution. Parameter names match the checkpoints:
class GatedDeltaNet(nn.Module):
"""A Qwen3.5 / Flash-Next linear-attention layer. (Your engine: Chapter 28)"""
def __init__(self, hidden, key_heads, value_heads, key_dim, value_dim, conv_kernel=4, eps=1e-6,
gate_activation="silu"):
super().__init__()
if value_heads % key_heads:
raise ValueError("value heads must be a multiple of key heads")
self.kh, self.vh, self.dk, self.dv = key_heads, value_heads, key_dim, value_dim
self.key_size, self.value_size = key_heads * key_dim, value_heads * value_dim
conv_dim = 2 * self.key_size + self.value_size
self.in_proj_qkv = nn.Linear(hidden, conv_dim, bias=False)
self.in_proj_z = nn.Linear(hidden, self.value_size, bias=False)
self.in_proj_b = nn.Linear(hidden, value_heads, bias=False)
self.in_proj_a = nn.Linear(hidden, value_heads, bias=False)
self.conv1d = nn.Conv1d(conv_dim, conv_dim, conv_kernel, groups=conv_dim, bias=False)
self.dt_bias = nn.Parameter(torch.ones(value_heads))
self.A_log = nn.Parameter(torch.log(torch.empty(value_heads).uniform_(0.01, 16)))
self.norm = RMSNormGated(value_dim, eps, gate_activation)
self.out_proj = nn.Linear(self.value_size, hidden, bias=False)
def forward(self, x, state=None, chunk=64):
"""x [B, T, D] -> (y [B, T, D], new LinearAttentionState). (Your engine: Chapter 28)"""
b, t, _ = x.shape
state = state or LinearAttentionState()
mixed, conv_state = causal_conv1d(self.in_proj_qkv(x).transpose(1, 2), self.conv1d.weight[:, 0], state.conv)
q, k, v = mixed.transpose(1, 2).split([self.key_size, self.key_size, self.value_size], dim=-1)
q = q.reshape(b, t, self.kh, self.dk).transpose(1, 2)
k = k.reshape(b, t, self.kh, self.dk).transpose(1, 2)
v = v.reshape(b, t, self.vh, self.dv).transpose(1, 2)
repeat = self.vh // self.kh # several value heads share a key head
q, k = q.repeat_interleave(repeat, 1), k.repeat_interleave(repeat, 1)
q = l2norm(q.float()) * self.dk ** -0.5
k = l2norm(k.float())
beta = torch.sigmoid(self.in_proj_b(x)).transpose(1, 2) # write strength in (0,1)
g = (-self.A_log.float().exp() * F.softplus(self.in_proj_a(x).float() + self.dt_bias)).transpose(1, 2)
rule = recurrent_gated_delta_rule if t == 1 else chunk_gated_delta_rule
kwargs = {} if t == 1 else {"chunk": chunk}
o, recurrent = rule(q, k, v, g, beta, state.recurrent, **kwargs)
z = self.in_proj_z(x).reshape(b, t, self.vh, self.dv)
o = self.norm(o.transpose(1, 2).to(x.dtype), z)
return self.out_proj(o.reshape(b, t, -1)), LinearAttentionState(conv_state, recurrent)
Step by step:
- Project $x$ to $q, k, v$ in one matrix (
in_proj_qkv), plus an output gate $z$ (in_proj_z), a write strength $b$ and a decay input $a$ (one scalar per value head each). - Short causal convolution over time on $q, k, v$ (kernel 4, one filter per channel, then SiLU). It lets each token mix in its three predecessors before the state sees it: cheap local context. Its own state is the last 3 inputs.
- Heads: there are more value heads than key heads (Flash-Next: 16 key heads, 48 value heads), and each key head is shared by several value heads, like grouped-query attention in reverse.
- Normalize: $q$ and $k$ L2-normalized, $q$ scaled by $d_k^{-1/2}$.
- Gates: $\beta = \sigma(b)$, and the log-decay $g = -e^{A_\text{log}} \cdot \operatorname{softplus}(a + \text{dt_bias})$, the same parameterization as Mamba’s $\Delta$: always negative, so $\alpha = e^{g} \in (0, 1)$.
- The rule: recurrent for a single token, chunked otherwise.
- Gated RMSNorm: normalize each head’s output, multiply by $\operatorname{SiLU}(z)$, then project back to the hidden size.
def causal_conv1d(x, weight, state=None, activation=True):
"""Depthwise causal convolution over time with carried history. (Your engine: Chapter 28)
x [B, C, T]; weight [C, K]; state [B, C, K-1] = the previous K-1 inputs (zeros at start).
Returns (y [B, C, T], new_state). Output t sees inputs t-K+1 .. t only.
"""
kernel = weight.shape[-1]
if state is None:
state = x.new_zeros(x.shape[0], x.shape[1], kernel - 1)
joined = torch.cat((state.to(x.dtype), x), dim=-1)
y = F.conv1d(joined, weight[:, None, :].to(x.dtype), groups=x.shape[1])
return (F.silu(y) if activation else y), joined[..., -(kernel - 1):]
class RMSNormGated(nn.Module):
"""RMSNorm(x) * weight * silu(z): normalize the read-out, then let a gate decide how much passes."""
def __init__(self, width, eps=1e-6, activation="silu"):
super().__init__()
self.weight = nn.Parameter(torch.ones(width))
self.eps, self.activation = eps, activation
def forward(self, x, z):
x32 = x.float()
normed = x32 * torch.rsqrt(x32.square().mean(-1, keepdim=True) + self.eps)
gate = F.silu(z.float()) if self.activation == "silu" else torch.sigmoid(z.float())
return (self.weight * normed.to(x.dtype) * gate).to(x.dtype)
The state a request carries
A Gated DeltaNet layer’s per-request state is the convolution history plus the recurrent matrices, fixed size whatever the context. For Flash-Next: 48 value heads × 128 × 128 FP32 values per layer, 3 MiB, times 36 linear layers, 108 MiB per sequence. Its 12 attention layers’ KV cache at 32k tokens is 0.75 GiB; at 262k tokens, 6 GiB. An all-attention model of the same shape would need four times as much.
The fixed state changes the engine in three ways:
- No truncation. A KV cache can be rolled back by forgetting positions (Chapter 26’s speculative decoding). A recurrent state can’t: the rejected tokens have been mixed into it. Engines snapshot the state before speculating and restore it on rejection (Chapter 30’s
HybridState). - Prefix caching needs snapshots too, stored at block boundaries: the state after a shared prefix, not a list of per-position entries.
- Batching is simple: every request’s state has the same shape.
Why hybrid?
Linear attention’s fixed memory is also its weakness. Tasks that need exact recall of an arbitrary earlier token (copy this ID, find the needle in the haystack) are hard for any fixed-size state once the context is long, as the recall demo showed. A few full-attention layers restore precise lookup, and models with a 3:1 ratio of linear to attention layers (Qwen3-Next, Qwen3.5, Flash-Next) match or beat pure attention models on long-context benchmarks at a fraction of the cache. Flash-Next goes one step further: its attention layers are themselves sparse (Chapter 29).
Build it
Engine milestone 28: Gated DeltaNet. Implement recurrent_gated_delta_rule, chunk_gated_delta_rule, causal_conv1d and GatedDeltaNet.forward in engine/gdn.py (the gated norm, the state class and the constructor are provided).
pytest tests/test_ch28_gdn.py
python run.py linear --impl engine
The tests check the worked example, chunked against recurrent for four chunk sizes with an initial state, that the convolution’s carried state makes split calls equal one call, that a layer prefilling 13 tokens and then decoding 8 matches one full forward, and parity with Hugging Face’s Flash-Next Qwen4ExpTextGatedDeltaNet layer with the same weights.
Stretch exercises
- ★ Extend the recall demo with decay: at $\alpha = 0.95$, how does error depend on how long ago a pair was written? Where:
experiments/ch28.py(create it), adapting the recall example inrun.py’scmd_linear. - ★★ Measure the chunked form’s time for chunk sizes 16, 32, 64, 128 and $T$ = 4,096. Explain the optimum in terms of the triangular solve’s cost ($O(C^2)$ per position) against the number of sequential chunk steps ($T/C$). Where:
experiments/ch28.py(create it), callingengine.gdn.chunk_gated_delta_rule. - ★★ Write a Triton kernel for the recurrent step at decode: one program per (sequence, head), the state tile in registers, all of decay, predict, correct and read fused. Where: new
engine/kernels/triton_gdn.py, selected byGatedDeltaNet.forwardinengine/gdn.py. - ★★★ Implement Mamba-2’s selective state-space recurrence ($S_t = \alpha_t S_{t-1} + k_t v_t^\top$ with scalar $\alpha_t$ per head) in the same chunked style, and show it’s Gated DeltaNet with $\beta$’s correction term removed. Where: add a selective-state-space helper beside
chunk_gated_delta_ruleinengine/gdn.py.
Check your understanding
- Why does replacing $\exp(q \cdot k)$ by $q \cdot k$ make attention a recurrence?
- Why do additive writes interfere, and what limits how many associations a $d_k \times d_v$ state can hold?
- What does the delta rule write instead of $v_t$, and why does $\beta = 1$ replace an association exactly?
- Why does the chunked form need a triangular solve?
- Why can’t a recurrent state be truncated like a KV cache, and what do engines do instead?
Going deeper
- Katharopoulos et al., Transformers are RNNs (2020); Schlag, Irie and Schmidhuber, Linear Transformers Are Secretly Fast Weight Programmers (2021); Yang et al., Gated Linear Attention (2023), Parallelizing Linear Transformers with the Delta Rule over Sequence Length (2024) and Gated Delta Networks (2024).
- Gu and Dao, Mamba (2023) and Dao and Gu, Transformers are SSMs (Mamba-2, 2024), for the state-space view of the same family.
- Songlin Yang’s blog series DeltaNet Explained (three parts), the clearest derivation of the chunked algorithm; the
flash-linear-attentionrepository for production kernels. - GPU Mode L20, L21 and L24 (Scan), the parallel-prefix algorithm behind Mamba’s selective scan; PMPP Chapter 11 (Scan: Kogge-Stone and Brent-Kung parallel prefix sums).
- The Qwen3-Next and Qwen3.5 model cards for the hybrid 3:1 layout.
29. Sparse attention: windows, top-k and learned indexers
In this chapter
- When attention can safely read only some keys, and when it can't.
- Fixed patterns: sliding windows and attention sinks.
- Content-based selection: top-k attention, and why exact top-k saves nothing by itself.
- Qwen Sparse Attention: a cheap block indexer that chooses what full attention reads, plus gated attention with partial RoPE, as in Flash-Next.
- Cache consistency: why selection must give identical results for full, chunked and token-by-token forwards.
You will build
QSAIndexer.forward and GatedAttention.forward in engine/sparse.py: Flash-Next's attention layer.
Time: 5-7 hours. GPU: not needed.
Most attention weights are tiny
A trained model’s attention is usually peaked: for a given query, a handful of keys carry most of the weight and thousands share the rest. If the weight on a key is $10^{-6}$, skipping it changes the output by about $10^{-6}$ of that value. Sparse attention bets on this: read only the keys that matter, and save the memory traffic of the rest.
The bet doesn’t always pay. run.py sparse compares exact top-$k$ attention (keep each query’s $k$ highest-scoring keys, renormalize) with dense attention over 1,024 keys, for flat and for peaked score distributions:
python run.py sparse
{"context": 1024, "score_scale": 1.0, "relative_error": {"budget_16": 4.4, "budget_64": 1.892, "budget_256": 0.595}}
{"context": 1024, "score_scale": 4.0, "relative_error": {"budget_16": 0.123, "budget_64": 0.026, "budget_256": 0.002}}
With flat scores (random queries and keys), the output is an average over many values, and keeping 256 of 1,024 keys still leaves a 60% error. With peaked scores, 64 keys give 2.6% error and 256 give 0.2%. Real models sit closer to the second case for most heads and most layers, which is why sparse attention works in practice, and why it’s trained into the model rather than applied after the fact.
Fixed patterns: windows and sinks
The simplest selection ignores content. A sliding window lets each query see only the last $w$ keys: memory is $O(w)$ per sequence, and the KV cache becomes a ring buffer. Mistral 7B used $w = 4{,}096$; Gemma alternates windowed and global layers.
Xiao et al. (2023) found a twist: windowed attention collapses when the first tokens leave the window, because models learn to dump excess attention weight on the very first positions (“attention sinks”). Keeping the first few tokens visible forever fixes it.
def sliding_window_allowed(query_positions, key_positions, window, sinks=0):
"""Keys within `window` positions of the query, plus the first `sinks` tokens of the sequence."""
qp = query_positions[..., :, None]
kp = key_positions[..., None, :]
return ((qp - kp) < window) | (kp < sinks)
Your causal_attention takes an allowed mask (Chapter 5), so a window is one more boolean term. The milestone test checks a window of 2 with 1 sink: position 5 sees keys 0, 4 and 5.
Content-based selection: top-k
Windows can’t recall something important from 50,000 tokens ago. Top-k attention selects by score instead:
def topk_attention(q, k, v, budget, query_positions=None):
"""Exact scores, but each query keeps only its `budget` highest-scoring visible keys.
Teaches selection; it still computes every score, so it saves nothing by itself."""
t, s = q.shape[-2], k.shape[-2]
qp = torch.arange(s - t, s, device=q.device) if query_positions is None else query_positions
scores = (q.float() @ k.float().transpose(-2, -1)) / math.sqrt(q.shape[-1])
visible = torch.arange(s, device=q.device)[None, :] <= qp[:, None]
scores = scores.masked_fill(~visible, float("-inf"))
keep = scores.topk(min(budget, s), dim=-1).indices
allowed = torch.zeros_like(scores, dtype=torch.bool).scatter(-1, keep, True) & visible
return causal_attention(q, k, v, qp, None, allowed)
This version is useful for understanding and for measuring error, but it saves nothing: to find the top $k$ scores it computes all of them, reading every key. The savings come only when something cheaper than attention decides which keys to read. That’s an indexer.
Qwen Sparse Attention
Flash-Next’s 12 attention layers use QSA: a small learned indexer picks, for each query, a budget of 2,048 keys out of the whole context, and full attention reads only those. Three ideas make the indexer cheap:
- Blocks, not tokens. The indexer’s keys are averaged over consecutive blocks of 4 tokens (the compression ratio). Selection happens per block: one decision per 4 tokens, and contiguous memory reads.
- Small and shared. The indexer has its own tiny projection: 4 query heads of dimension 128, and one key per block shared by all of them, against attention’s 2 KV heads of dimension 256 per token.
- A simple score. For query $q$ (with heads $h$) and pooled block key $\bar k_b$: $$ \text{score}(q, b) = \frac{1}{\sqrt{d}} \sum_h \operatorname{ReLU}\big(q_h \cdot \bar k_b\big). $$ The ReLU lets each indexer head vote only for blocks, never against.
Then each query keeps its top $\text{budget}/\text{ratio} = 512$ complete visible blocks, plus the open block, the partially filled block it’s in, which is always readable so recent context is never lost. Both the indexer’s queries and the pooled block keys get RoPE (the block key at the block’s first position), so the indexer is position-aware.
class QSAIndexer(nn.Module):
"""Chooses, for every query, which keys attention may read. (Your engine: Chapter 29)
Keys are averaged over consecutive blocks of `ratio` tokens; one shared indexer key per
block (RMSNorm, then RoPE at the block's first position) is scored against several
small indexer query heads: score = sum over heads of ReLU(q . k) / sqrt(d). Each query
keeps its top budget/ratio complete visible blocks plus the incomplete trailing block.
"""
def __init__(self, hidden, heads, head_dim, budget, ratio, eps=1e-6):
super().__init__()
if budget % ratio:
raise ValueError("budget must be a multiple of the compression ratio")
self.heads, self.head_dim, self.ratio, self.block_topk = heads, head_dim, ratio, budget // ratio
self.index_qk_proj = nn.Linear(hidden, (heads + 1) * head_dim, bias=False)
self.q_layernorm = ZeroCenteredRMSNorm(head_dim, eps)
self.k_layernorm = ZeroCenteredRMSNorm(head_dim, eps)
def forward(self, x, positions, rotary_dim, theta, cached_keys=None):
"""x [B, T, D] at absolute positions [T] -> (allowed [B, 1, T, S], all raw keys [B, S, d]). (Your engine: Chapter 29)"""
b, t, _ = x.shape
q, raw = self.index_qk_proj(x).split([self.heads * self.head_dim, self.head_dim], dim=-1)
q = self.q_layernorm(q.view(b, t, self.heads, self.head_dim))
cos, sin = rope_cos_sin(positions, rotary_dim, theta, q.dtype) # [1, 1, T, r]
q = apply_rope(q, cos.transpose(1, 2), sin.transpose(1, 2)) # broadcast over heads
keys = raw if cached_keys is None else torch.cat((cached_keys, raw), dim=1)
s = keys.shape[1]
n_blocks = s // self.ratio
key_pos = torch.arange(s, device=x.device)
visible_tail_start = ((positions + 1) // self.ratio) * self.ratio # first token of the open block
tail = (key_pos[None, :] >= visible_tail_start[:, None]) & (key_pos[None, :] <= positions[:, None])
if n_blocks == 0:
return tail[None, None].expand(b, 1, t, s), keys
pooled = keys[:, :n_blocks * self.ratio].view(b, n_blocks, self.ratio, -1).float().mean(2).to(keys.dtype)
pooled = self.k_layernorm(pooled)
starts = torch.arange(n_blocks, device=x.device) * self.ratio
bcos, bsin = rope_cos_sin(starts, rotary_dim, theta, pooled.dtype) # [1, 1, n, r]
pooled = apply_rope(pooled, bcos[:, 0], bsin[:, 0])
scores = torch.relu(torch.einsum("bthd,bnd->btnh", q.float(), pooled.float())).sum(-1) / math.sqrt(self.head_dim)
# Queries that see the same number of complete blocks run topk over exactly those blocks.
# (topk breaks ties differently for different vector lengths; ReLU makes ties at 0 common,
# so slicing rather than masking with -inf is what keeps us identical to the reference.)
visible_blocks = (positions + 1) // self.ratio # [T]
chosen = torch.zeros(b, t, n_blocks, dtype=torch.bool, device=x.device)
for count in visible_blocks.unique().tolist():
if count == 0:
continue
rows = visible_blocks == count
top = scores[:, rows, :count].topk(min(self.block_topk, count), dim=-1).indices
chosen[:, rows] = chosen[:, rows].scatter(-1, top, True)
block_of_key = (key_pos // self.ratio).clamp(max=n_blocks - 1) # tail keys are masked below
from_blocks = chosen.gather(-1, block_of_key.expand(b, t, s))
from_blocks = from_blocks & (key_pos < n_blocks * self.ratio)
return (from_blocks | tail[None])[:, None], keys
The indexer’s output is just an allowed mask [B, 1, T, S], and attention uses it exactly like a causal mask. A real kernel instead gathers the selected blocks’ K and V (Chapter 25’s block tables are the natural fit) and never reads the others.
What it saves
The arithmetic for one Flash-Next attention layer at decode, in BF16 (the second half of run.py sparse):
{"context": 8192, "dense_kv_MB_per_layer_step": 16.8, "qsa_MB_per_layer_step": 4.7, "reduction": 3.5}
{"context": 32768, "dense_kv_MB_per_layer_step": 67.1, "qsa_MB_per_layer_step": 6.3, "reduction": 10.7}
{"context": 262144, "dense_kv_MB_per_layer_step": 536.9, "qsa_MB_per_layer_step": 21.0, "reduction": 25.6}
QSA reads the K and V of 2,048 selected tokens plus every block’s 128-dimensional indexer key. The indexer’s reads still grow with context, but 16 times more slowly than the full KV cache (one 256-byte key per 4 tokens, against 4 KiB of K and V per token). The cache itself still stores everything: sparse attention saves bandwidth and compute, not memory. That’s why Flash-Next combines it with linear attention (Chapter 28), which saves memory.
A subtlety: a query may not see itself
A query at position 11, with blocks of 4, completes block 2. Its open block is now empty, and block 2 is a complete block that must compete in the top-$k$ with blocks 0 and 1. With a budget of two blocks and a low score for block 2, the query attends to blocks 0 and 1 and not to itself. The milestone test asserts that this can happen, rather than assuming that every query sees its own token. It always sees some keys (at least one complete block wins the top-$k$), just not necessarily its own.
Gated attention with partial RoPE
Flash-Next’s attention layer (inherited from Qwen3-Next and Qwen3.5) changes three more things from Chapter 17’s Qwen3:
- An output gate.
q_projproduces twice the query width: a query and a gate per head. The attention output is multiplied by $\sigma(\text{gate})$ beforeo_proj. The gate lets a head output nothing when it has nothing useful to add, which also removes the need for attention sinks (Qiu et al., 2025, Gated Attention for Large Language Models). - Partial RoPE. Only the first 25% of each head’s 256 dimensions are rotated (
partial_rotary_factor); the remaining 192 carry content without position. Yourapply_ropefrom Chapter 17 rotates the firstrotary_dimfeatures. - Zero-centered RMSNorm for the Q and K norms: the scale is $1 + w$ with $w$ initialized at 0, so weight decay pulls toward the identity instead of toward zero.
class ZeroCenteredRMSNorm(nn.Module):
"""RMSNorm with scale (1 + weight), weight initialized at 0. Optionally normalizes each
group of `group_size` features separately (used on the 4 residual streams in Chapter 30)."""
def __init__(self, width, eps=1e-6, group_size=None):
super().__init__()
self.weight = nn.Parameter(torch.zeros(width))
self.eps, self.group_size = eps, group_size
def forward(self, x):
x32 = x.float()
if self.group_size:
x32 = x32.unflatten(-1, (-1, self.group_size))
x32 = x32 * torch.rsqrt(x32.square().mean(-1, keepdim=True) + self.eps)
if self.group_size:
x32 = x32.flatten(-2)
return (x32 * (1.0 + self.weight.float())).type_as(x)
class GatedAttention(nn.Module):
"""Qwen3.5/Flash-Next attention: zero-centered Q/K norms, partial RoPE, GQA, an optional QSA
indexer, and a sigmoid output gate produced by the same q_proj. (Your engine: Chapter 29)"""
def __init__(self, hidden, heads, kv_heads, head_dim, rotary_dim, theta, eps=1e-6, indexer=None):
super().__init__()
self.heads, self.kv_heads, self.head_dim = heads, kv_heads, head_dim
self.rotary_dim, self.theta = rotary_dim, theta
self.q_proj = nn.Linear(hidden, heads * head_dim * 2, bias=False) # [query | gate] per head
self.k_proj = nn.Linear(hidden, kv_heads * head_dim, bias=False)
self.v_proj = nn.Linear(hidden, kv_heads * head_dim, bias=False)
self.o_proj = nn.Linear(heads * head_dim, hidden, bias=False)
self.q_norm = ZeroCenteredRMSNorm(head_dim, eps)
self.k_norm = ZeroCenteredRMSNorm(head_dim, eps)
self.indexer = indexer
def forward(self, x, positions, state=None):
"""x [B, T, D], positions [T] -> (y [B, T, D], new AttentionState). (Your engine: Chapter 29)"""
b, t, _ = x.shape
state = state or AttentionState()
query, gate = self.q_proj(x).view(b, t, self.heads, 2 * self.head_dim).chunk(2, dim=-1)
q = self.q_norm(query).transpose(1, 2)
k = self.k_norm(self.k_proj(x).view(b, t, self.kv_heads, self.head_dim)).transpose(1, 2)
v = self.v_proj(x).view(b, t, self.kv_heads, self.head_dim).transpose(1, 2)
cos, sin = rope_cos_sin(positions, self.rotary_dim, self.theta, q.dtype)
q, k = apply_rope(q, cos, sin), apply_rope(k, cos, sin) # only the first rotary_dim features
allowed, index_keys = None, None
if self.indexer is not None:
allowed, index_keys = self.indexer(x, positions, self.rotary_dim, self.theta, state.index_keys)
state = state.append(k, v)
state.index_keys = index_keys # all raw indexer keys so far
y = causal_attention(q, state.k, state.v, positions, None, allowed)
y = y.transpose(1, 2).reshape(b, t, -1) * torch.sigmoid(gate.reshape(b, t, -1))
return self.o_proj(y), state
AttentionState carries the layer’s per-request cache: K, V and all raw indexer keys (pooling happens on the fly, so a decode step that completes a block sees it immediately). This reference concatenates; in Chapter 30’s engine this is the state that a paged or static cache would replace.
Consistency: the hard requirement
Selection is discrete: a block is in or out. If a full forward and a token-by-token decode compute scores with tiny differences, they can select different blocks and produce very different outputs, and the cache-equivalence test from Chapter 16 fails in a way that looks random. Two rules keep them identical:
- Score each query against exactly the blocks it can see, the same way in both paths. A decode step at position $p$ must see the same pooled keys a full forward does for row $p$.
- Break ties the same way. ReLU makes exact ties at 0 common (many blocks score 0). PyTorch’s
topkbreaks ties differently depending on the length of the vector it searches. Masking unseen blocks with $-\infty$ in a long vector and slicing a short vector to just the visible blocks can choose different zero-scored blocks. The reference therefore groups queries by their number of visible blocks and runstopkover exactly that slice, for every path.
The second rule was found the hard way: the first version of the whole Flash-Next model differed from Hugging Face’s on a few positions, and the cause was exactly this tie-breaking. The milestone test runs prefill of 11 tokens followed by 14 single-token steps against one full forward, with a budget small enough that selection matters.
Build it
Engine milestone 29: sparse attention. Implement QSAIndexer.forward and GatedAttention.forward in engine/sparse.py (windows, the norm, top-k attention and the state are provided).
pytest tests/test_ch29_sparse.py
python run.py sparse --impl engine
The tests check windows with sinks, that top-k with a full budget equals dense attention, the grouped zero-centered norm, that the indexer always keeps the open block, respects the budget and never selects a future key, and that prefill plus token-by-token decode with selection equals one full forward. Chapter 30’s parity test checks the layer against Hugging Face’s implementation inside the full model.
Stretch exercises
- ★ Add an attention-sink option to
topk_attention(always keep the first 4 keys) and measure the error change with flat and peaked scores. Where:topk_attentioninengine/sparse.py. - ★★ Implement a ring-buffer KV cache for sliding-window layers, storing only the last $w$ positions, and verify a windowed model’s logits against the full-cache version. Where:
AttentionStateinengine/sparse.py, with position handling inGatedAttention.forward. - ★★ Make the indexer’s selection explicit: return the selected block indices
[B, T, budget/ratio]instead of a mask, and write attention that gathers only those blocks’ K and V. Check it against the mask version. Where:QSAIndexer.forwardandGatedAttention.forwardinengine/sparse.py. - ★★★ Write a Triton decode kernel for QSA: one program per (sequence, head) that reads the selected block list and streams those blocks’ K/V with the online softmax (Chapter 25’s paged kernel is most of it). Where: new
engine/kernels/triton_sparse.py, called byGatedAttention.forwardinengine/sparse.py.
Check your understanding
- Why does top-k attention with exact scores save no computation?
- Why does Flash-Next’s indexer select blocks rather than tokens?
- What is the open block, and why is it always readable?
- Why does sparse attention save bandwidth but not cache memory?
- Why must full and incremental forwards break top-k ties identically?
Going deeper
- Beltagy et al., Longformer (2020) and Zaheer et al., BigBird (2020) for fixed sparse patterns; Xiao et al., Efficient Streaming Language Models with Attention Sinks (2023).
- DeepSeek-AI, Native Sparse Attention (2025) and the DeepSeek-V3.2 report (DeepSeek Sparse Attention, a lightning indexer with ReLU scores): the closest published relatives of QSA.
- Qiu et al., Gated Attention for Large Language Models: Non-linearity, Sparsity, and Attention-Sink-Free (2025), the gate used here.
- The Transformers
qwen4_expmodeling file, the reference implementation this chapter’s code is checked against.
30. Capstone: Qwen3.8-Flash-Next from scratch
In this chapter
- Reading a frontier model's configuration and mapping every component to something you've built.
- The two genuinely new pieces: hyper-connections (four residual streams) and the hashed n-gram memory.
- Assembling the hybrid model, its per-request state, and an exact parity check against the official implementation.
- Loading the real checkpoint on a single machine: experts quantized while loading, a 95 GiB n-gram table left on disk, and a byte budget for each choice.
You will build
engine/flashnext.py: the complete Qwen3.8-Flash-Next text model, its loader and its memory plan, served through your Chapter 18 engine.
Time: 2-3 weeks. GPU: the tests run on a CPU; the real model needs about 70 GB of GPU or unified memory with 4-bit experts.
The target
Qwen3.8-Flash-Next is a multimodal model whose language part combines nearly every idea in Part VII. Its model card describes a 125B-parameter backbone with 6B parameters active per token, plus 51B parameters of n-gram memory, 4B parameters of multi-token-prediction heads, and a vision encoder. The text configuration (qwen4_exp_text in Transformers 5.18 and later):
| field | value | chapter |
|---|---|---|
| layers / hidden size | 48 / 2,560 | 6, 17 |
| layer pattern | 3 linear-attention layers, then 1 attention layer, ×12 | 28, 29 |
| attention: query / KV heads, head dim | 24 / 2, 256 | 17 |
| partial RoPE | 25% of each head (64 of 256 dims), θ = 10⁷ | 17, 29 |
| QSA indexer: heads × dim, block ratio, budget | 4 × 128, 4, 2,048 tokens | 29 |
| Gated DeltaNet: key / value heads, dims | 16 / 48, 128 / 128 | 28 |
| routed experts / per token, expert width | 512 / 10, 640 | 27 |
| shared expert width (sigmoid-gated) | 640 | 27 |
| residual streams / low-rank width | 4 / 320 | this chapter |
| n-gram memory: orders, hash heads per order, layer | 2-3, 8, layer 2 | this chapter |
| vocabulary / context | 248,320 / 262,144 | 4 |
Twenty-eight chapters of preparation cover all but two rows. Before writing any code, the professional habit is an architecture ledger: for every component, the equations, tensor names and shapes, and the reference you’ll check against. The official implementation (transformers/models/qwen4_exp/modeling_qwen4_exp.py) is the ground truth; this chapter’s code was written from it and is tested against it.
What’s already built
- Gated DeltaNet layers with their convolution and recurrent state, exactly Chapter 28’s
GatedDeltaNet, with the same parameter names. - Gated attention with zero-centered Q/K norms, partial RoPE, two KV heads and the QSA indexer: Chapter 29’s
GatedAttention. - MoE with 512 experts, top-10, renormalized softmax router and a sigmoid-gated shared expert: Chapter 27’s
SparseMoeBlock. - Loading sharded safetensors, meta-device construction and expert stacking: Chapters 9, 18 and 27.
Two things are new.
Hyper-connections: four residual streams
Every model so far had one residual stream: x = x + sublayer(norm(x)) (Chapter 23). Flash-Next keeps four parallel streams of width 2,560, and each sublayer learns how to read from them and how to write back (hyper-connections, Zhu et al., 2024). The embedding initializes all four streams to the same vector, and the streams drift apart as layers write to them differently.
Around each sublayer (attention or linear attention, then the MoE), a GatedResidual module:
- Normalizes each stream separately: a zero-centered RMSNorm over groups of 2,560 features. This replaces the usual pre-norm; there’s no other norm before the sublayer.
- Reads: computes per-feature weights for each stream through a low-rank bottleneck, $w = \sigma\big(W_\text{up}, \operatorname{SiLU}(W_\text{down}, \hat s / 4)\big)$, with rank 320. The sublayer’s input is the mean over streams of $w \odot \hat s$.
- Writes: computes one scalar per stream, $\gamma = 2\sigma(W_\text{inject}, \hat s / 4) \in (0, 2)$, and adds $\gamma_i \cdot \text{output}$ to stream $i$.
At the end, a final GatedResidual without the write part (combine=False) reads one vector from the four streams for the LM head. There’s no final norm.
class GatedResidual(nn.Module):
"""Hyper-connection around one sublayer. (Your engine: Chapter 30)
streams [B, T, hc*D] -> normalize each stream (grouped zero-centered RMSNorm);
read weights w = sigmoid(up(silu(down(normed) / hc))) [B, T, hc, D]
sublayer input = mean over streams of w * normed [B, T, D]
write weights = 2 * sigmoid(inject(normed) / hc) [B, T, hc]
After the sublayer: streams + write_weights[..., None] * output[..., None, :].
"""
def __init__(self, hidden, hc, rank, eps=1e-6, combine=True):
super().__init__()
self.hc, self.hidden = hc, hidden
self.hc_norm = ZeroCenteredRMSNorm(hc * hidden, eps, group_size=hidden)
self.input_mix_weight_down = nn.Linear(hc * hidden, rank, bias=False)
self.input_mix_weight_up = nn.Linear(rank, hc * hidden, bias=False)
self.block_inject_weight = nn.Linear(hc * hidden, hc, bias=False) if combine else None
def read(self, streams):
"""(Your engine: Chapter 30)"""
normed = self.hc_norm(streams)
weights = torch.sigmoid(self.input_mix_weight_up(F.silu(self.input_mix_weight_down(normed) / self.hc)))
mixed = (weights.unflatten(-1, (self.hc, self.hidden)) * normed.unflatten(-1, (self.hc, self.hidden))).mean(-2)
if self.block_inject_weight is None:
return mixed, None
return mixed, 2 * torch.sigmoid(self.block_inject_weight(normed) / self.hc)
@staticmethod
def write(streams, output, inject):
return streams + (output.unsqueeze(-2) * inject.unsqueeze(-1)).flatten(-2)
Why bother? With one stream, each layer’s output is added with weight 1, and depth works through a single shared channel. Multiple streams with learned read and write weights let the network keep some information untouched by later layers and vary how strongly each layer contributes, which in the hyper-connection papers improves training stability and quality at a cost of a few small matrices per layer. For inference, it means the hidden state is four times wider between layers: 20 KiB per token instead of 5 KiB in BF16, which matters for activation memory during prefill and not at all for the KV cache.
The n-gram memory
The second new piece is a huge lookup table indexed by the last few tokens. The idea: many next-token facts depend only on the immediately preceding tokens (“New York” → “City”), and a model shouldn’t spend attention and MLP compute rediscovering them. A hashed table can store an embedding for every frequent bigram and trigram, and the model looks it up in $O(1)$.
Hashing n-grams
For position $t$, the bigram is $(x_{t-1}, x_t)$ and the trigram $(x_{t-2}, x_{t-1}, x_t)$. There are $248{,}320^3$ possible trigrams, far too many to store, so each n-gram is hashed into a table of about 20 million rows:
$$ h = \Big(\bigoplus_{i=0}^{n-1} x_{t-i} \cdot m_i\Big) \bmod p , $$
with odd multipliers $m_i$ derived from a seed with splitmix64, XOR ($\oplus$) to combine, and a prime table size $p$. Hash collisions are inevitable, so each n-gram order uses 8 independent hash heads, each with its own prime size (the 16 smallest primes above 20 million), each returning a 160-dimensional row. Concatenated, the 16 rows form one 2,560-dimensional embedding. Two n-grams colliding in one head are very unlikely to collide in all eight.
MASK64 = (1 << 64) - 1
def splitmix64(value):
value = (value + 0x9E3779B97F4A7C15) & MASK64
value = ((value ^ (value >> 30)) * 0xBF58476D1CE4E5B9) & MASK64
value = ((value ^ (value >> 27)) * 0x94D049BB133111EB) & MASK64
return (value ^ (value >> 31)) & MASK64
def layer_multipliers(vocab, ngram_size, ple_index, seed):
"""Odd multipliers, small enough that token_id * multiplier never overflows int64."""
half = max(1, (((1 << 63) - 1) // max(vocab, 1)) // 2)
base = seed + 10007 * ple_index
return [2 * (splitmix64((base + 0x9E3779B97F4A7C15 * (i + 1)) & MASK64) % half) + 1 for i in range(ngram_size)]
def next_primes(start, count):
"""The `count` smallest primes greater than `start` (one distinct table size per hash head)."""
def is_prime(n):
if n < 2 or (n % 2 == 0 and n != 2):
return n == 2
return all(n % d for d in range(3, math.isqrt(n) + 1, 2))
primes, n = [], start
while len(primes) < count:
n += 1
if is_prime(n):
primes.append(n)
return primes
Two details that a careless implementation gets wrong:
- Document boundaries. An n-gram must not span an end-of-sequence token: after an EOS, the “previous token” positions read as EOS.
shiftedcomputes, for every position, how far back the current document starts, and substitutes EOS beyond it. - Carried history. In decode, the n-gram of the new token needs the previous two tokens, which belong to earlier calls. The layer’s state carries them.
class NGramEmbedding(nn.Module):
"""Hashed bigram and trigram embeddings, heads_per_ngram independent hashes per order. (Your engine: Chapter 30)
The n-gram ending at token t is (x[t-n+1], ..., x[t]); its id for one head is
(XOR over i of x[t-i] * m_i) mod p_head, offset into one big table. N-grams never cross an
EOS: positions before the most recent EOS read as EOS.
"""
def __init__(self, cfg, ple_index):
super().__init__()
self.n, self.heads_per = cfg.ngram_size, cfg.heads_per_ngram
self.heads = (cfg.ngram_size - 1) * cfg.heads_per_ngram
self.eos = cfg.eos_token_id
sizes = next_primes(cfg.ngram_vocab_size_base - 1, (ple_index + 1) * self.heads)[ple_index * self.heads:]
self._sizes = sizes
offsets = [sum(sizes[:i]) for i in range(len(sizes))]
padded = -(-sum(sizes) // cfg.make_ngram_vocab_size_divisible_by) * cfg.make_ngram_vocab_size_divisible_by
self.register_buffer("layer_multipliers", torch.tensor(layer_multipliers(cfg.vocab_size, self.n, ple_index, cfg.seed)))
self.register_buffer("ngram_heads_vocab_sizes", torch.tensor(sizes))
self.register_buffer("ngram_heads_offsets", torch.tensor(offsets))
self.ngram_embedding = nn.Embedding(padded, cfg.ple_embed_dim // self.heads)
def shifted(self, history, shift):
"""history[t - shift], or EOS when that position is before the current document. (Your engine: Chapter 30)"""
if shift == 0:
return history
b, length = history.shape
pos = torch.arange(length, device=history.device)
eos_at = torch.where(history == self.eos, pos, -1)
last_eos_before = torch.cat((eos_at.new_full((b, 1), -1), eos_at.cummax(1).values[:, :-1]), dim=1)
in_document = pos - (last_eos_before + 1)
source = (pos - shift).clamp_min(0).expand(b, -1)
valid = (in_document >= shift) & ((pos - shift) >= 0)
return torch.where(valid, history.gather(1, source), torch.full_like(history, self.eos))
def forward(self, ids, context=None):
"""ids [B, T]; context = the previous n-1 token ids (EOS at the start). Returns (emb [B, T, E], new context). (Your engine: Chapter 30)"""
if context is None:
context = ids.new_full((ids.shape[0], self.n - 1), self.eos)
history = torch.cat((context, ids.long()), dim=1)
shifted = [self.shifted(history, s) for s in range(self.n)]
blocks = []
for order in range(2, self.n + 1):
mixed = shifted[0] * self.layer_multipliers[0]
for position in range(1, order):
mixed = mixed ^ (shifted[position] * self.layer_multipliers[position])
first = (order - 2) * self.heads_per
sizes = self.ngram_heads_vocab_sizes[first:first + self.heads_per]
blocks.append(mixed[..., None] % sizes + self.ngram_heads_offsets[first:first + self.heads_per])
index = torch.cat(blocks, dim=-1)[:, -ids.shape[1]:]
table = self.ngram_embedding.weight
rows = table[index.to(table.device)].to(ids.device) # the table may live on the host
return rows.flatten(-2), history[:, -(self.n - 1):]
Injecting it: per-layer embeddings
The lookup result is injected into the residual streams at layer 2 by a PLELayer (“per-layer embedding”):
- A key projection of the n-gram embedding (one per stream) and a query from the normalized streams give a gate per stream: their dot product over $\sqrt{d}$, passed through a signed square root (to tame large values while keeping the sign) and a sigmoid.
- A value projection of the n-gram embedding, scaled by each stream’s gate, is the addition to that stream.
- A dilated causal convolution (kernel 4, dilation 3, depthwise) over recent gated values adds local context, with its own carried state of the last 9 positions.
class PLELayer(nn.Module):
"""Per-layer embedding: gate the n-gram value into each residual stream, then add a dilated
causal depthwise convolution over recent gated values. (Your engine: Chapter 30)"""
def __init__(self, cfg, ple_index):
super().__init__()
d, hc, e = cfg.hidden_size, cfg.hc_count, cfg.ple_embed_dim
self.hidden, self.hc, self.dilation = d, hc, cfg.ngram_size
self.state_len = (cfg.ple_conv_kernel_size - 1) * cfg.ngram_size
self.ple_embedding = NGramEmbedding(cfg, ple_index)
self.key_proj = nn.Linear(e, hc * d, bias=False)
self.value_proj = nn.Linear(e, d, bias=False)
self.norm_key = ZeroCenteredRMSNorm(hc * d, cfg.rms_norm_eps, group_size=d)
self.norm_query = ZeroCenteredRMSNorm(hc * d, cfg.rms_norm_eps, group_size=d)
self.norm_conv = ZeroCenteredRMSNorm(hc * d, cfg.rms_norm_eps, group_size=d)
self.conv1d = nn.Conv1d(hc * d, hc * d, cfg.ple_conv_kernel_size, groups=hc * d,
dilation=cfg.ngram_size, bias=False)
def forward(self, streams, ids, state=None):
"""state = (ngram context ids, conv history). Returns (addition to streams, new state). (Your engine: Chapter 30)"""
context, conv_state = state if state is not None else (None, None)
embedding, context = self.ple_embedding(ids, context)
key = self.norm_key(self.key_proj(embedding)).unflatten(-1, (self.hc, self.hidden))
value = self.value_proj(embedding)
query = self.norm_query(streams).unflatten(-1, (self.hc, self.hidden))
gate = (key * query).sum(-1, keepdim=True) / math.sqrt(self.hidden)
gate = gate.abs().clamp_min(1e-6).sqrt() * gate.sign() # signed square root
gated = (torch.sigmoid(gate) * value.unsqueeze(-2)).flatten(-2) # [B, T, hc*D]
x = self.norm_conv(gated).transpose(1, 2)
if conv_state is None:
conv_state = x.new_zeros(x.shape[0], x.shape[1], self.state_len)
joined = torch.cat((conv_state, x), dim=-1)
conv = F.silu(self.conv1d(joined)).transpose(1, 2)
return gated + conv, (context, joined[..., -self.state_len:])
Why it’s cheap at inference: each token reads 16 rows of 160 BF16 values, about 5 KB, from a table of 51 billion parameters. The table’s size is a storage problem, not a bandwidth problem, so it can live in host memory or even on an SSD. Closely related published work: DeepSeek’s Engram conditional memory (2026).
The layer and the model
class FlashNextLayer(nn.Module):
def __init__(self, cfg, index):
super().__init__()
self.kind = cfg.layer_types[index]
if self.kind == "linear_attention":
self.linear_attn = GatedDeltaNet(cfg.hidden_size, cfg.linear_num_key_heads, cfg.linear_num_value_heads,
cfg.linear_key_head_dim, cfg.linear_value_head_dim,
cfg.linear_conv_kernel_dim, cfg.rms_norm_eps, cfg.output_gate_type)
else:
indexer = None
if cfg.indexer_n_heads:
indexer = QSAIndexer(cfg.hidden_size, cfg.indexer_n_heads, cfg.indexer_head_dim,
cfg.indexer_budget, cfg.indexer_compress_ratio, cfg.rms_norm_eps)
self.self_attn = GatedAttention(cfg.hidden_size, cfg.num_attention_heads, cfg.num_key_value_heads,
cfg.head_dim, cfg.rotary_dim, cfg.rope_theta, cfg.rms_norm_eps, indexer)
self.mlp = SparseMoeBlock(cfg.hidden_size, cfg.num_experts, cfg.num_experts_per_tok, cfg.moe_intermediate_size,
cfg.shared_expert_intermediate_size, cfg.norm_topk_prob)
ple_index = cfg.ple_layer_ids.index(index + 1) if (index + 1) in cfg.ple_layer_ids else None
self.ple = PLELayer(cfg, ple_index) if ple_index is not None else None
self.attn_hyper_connection = GatedResidual(cfg.hidden_size, cfg.hc_count, cfg.hc_lowrank, cfg.rms_norm_eps)
self.mlp_hyper_connection = GatedResidual(cfg.hidden_size, cfg.hc_count, cfg.hc_lowrank, cfg.rms_norm_eps)
def forward(self, streams, ids, positions, state):
"""state = {"mixer": layer state, "ple": PLE state}. Returns (streams, new state). (Your engine: Chapter 30)"""
new_state = {}
if self.ple is not None:
addition, new_state["ple"] = self.ple(streams, ids, state.get("ple"))
streams = streams + addition
x, inject = self.attn_hyper_connection.read(streams)
if self.kind == "linear_attention":
y, new_state["mixer"] = self.linear_attn(x, state.get("mixer"))
else:
y, new_state["mixer"] = self.self_attn(x, positions, state.get("mixer"))
streams = GatedResidual.write(streams, y, inject)
x, inject = self.mlp_hyper_connection.read(streams)
streams = GatedResidual.write(streams, self.mlp(x), inject)
return streams, new_state
Each layer: inject the n-gram memory (layer 2 only), read from the streams, run the mixer (Gated DeltaNet or gated sparse attention), write, read again, run the MoE, write.
class HybridState:
"""Everything one sequence carries between calls. Updates return a new object, so keeping the
old one is a free snapshot: speculative decoding rolls back by simply not adopting the new state."""
def __init__(self, layers=None, length=0):
self.layers = layers or []
self.length = length
class HybridCache:
"""Adapts the functional HybridState to the engines' mutable-cache convention (Chapters 18-19):
model(ids, cache) updates cache.state in place and returns logits only."""
def __init__(self):
self.state = None
@property
def length(self):
return 0 if self.state is None else self.state.length
def snapshot(self):
return self.state # states are never modified in place: this is a free copy
def restore(self, state):
self.state = state
def truncate(self, length):
raise NotImplementedError("A recurrent state cannot be truncated: restore a snapshot instead")
class FlashNextBackbone(nn.Module):
def __init__(self, cfg):
super().__init__()
self.embed_tokens = nn.Embedding(cfg.vocab_size, cfg.hidden_size)
self.layers = nn.ModuleList(FlashNextLayer(cfg, i) for i in range(cfg.num_hidden_layers))
self.hyper_connection_mixer = GatedResidual(cfg.hidden_size, cfg.hc_count, cfg.hc_lowrank,
cfg.rms_norm_eps, combine=False)
class FlashNext(nn.Module):
def __init__(self, cfg):
super().__init__()
self.cfg = cfg
self.model = FlashNextBackbone(cfg)
self.lm_head = nn.Linear(cfg.hidden_size, cfg.vocab_size, bias=False)
@property
def context_limit(self):
return self.cfg.max_position_embeddings
def new_cache(self, batch, capacity):
if capacity > self.context_limit:
raise ValueError("Requested cache exceeds the context limit")
return HybridCache()
def forward(self, ids, state=None):
"""ids [B, T] continuing `state` -> (logits, new HybridState); with a HybridCache, logits only."""
if isinstance(state, HybridCache):
logits, state.state = self.step(ids, state.state)
return logits
return self.step(ids, state)
def step(self, ids, state=None):
"""ids [B, T] continuing `state` -> (logits [B, T, V], new HybridState). (Your engine: Chapter 30)"""
state = state or HybridState([{} for _ in self.model.layers])
positions = torch.arange(state.length, state.length + ids.shape[1], device=ids.device)
streams = self.model.embed_tokens(ids).repeat(1, 1, self.cfg.hc_count) # every stream starts equal
new_layers = []
for layer, layer_state in zip(self.model.layers, state.layers):
streams, layer_state = layer(streams, ids, positions, layer_state)
new_layers.append(layer_state)
hidden, _ = self.model.hyper_connection_mixer.read(streams) # collapse streams
return self.lm_head(hidden), HybridState(new_layers, state.length + ids.shape[1])
The state of a request
One Flash-Next request carries four kinds of state, all owned by HybridState:
| state | where | size per sequence |
|---|---|---|
| recurrent matrices + conv history | 36 linear-attention layers | 108 MiB, fixed |
| K, V and indexer keys | 12 attention layers | 27 KiB per token (KV 24 KiB, indexer 3 KiB): 0.84 GiB at 32k |
| last 2 token IDs + PLE conv history | layer 2 | a few KB, fixed |
| position | model | one integer |
step never modifies a state in place: it returns a new HybridState. Keeping the old one is therefore a free snapshot, which is exactly what speculative decoding needs, since the recurrent state can’t be truncated (Chapter 28). The milestone test runs a branch from a state, discards it, runs another from the same state, and checks they’re identical. For the engines of Chapters 18-19, new_cache returns a HybridCache, a small mutable wrapper with snapshot and restore, so LLM.stream serves Flash-Next unchanged.
Exactly equal to the official implementation
The tests build a tiny random Flash-Next with Transformers’ Qwen4ExpForCausalLM: 8 layers, 2 of them attention with a QSA budget small enough to matter, 8 experts, two PLE layers with sharded n-gram tables. Every norm and convolution weight is randomized, since zero-initialized ones would hide bugs. The tests save it as a checkpoint and load it with your loader:
max abs logit difference: 1.49e-08 (logits of magnitude 0.12; FP32)
and check that:
- a full forward over 2 × 41 tokens with an EOS in the middle matches the reference;
- prefill of 13 tokens, a chunk of 16, then 11 single-token steps equals one full forward;
- the state is a free snapshot;
- the real configuration’s memory plan has the published sizes;
- the n-gram tables memory-mapped from disk give bit-identical logits, and 4-bit experts stay close;
- your Chapter 18
LLMengine generates the same tokens as the model’s own greedy loop.
Getting there surfaced real bugs, each now a lesson in an earlier chapter: a transposed triangle in the chunked delta rule (28) and topk tie-breaking under ReLU zeros (29). A third, a config field renamed between library versions (27), would have broken any loader that trusted defaults.
Loading the real checkpoint
@torch.no_grad()
def load_flashnext(directory, device="cpu", dtype=torch.bfloat16, ngram_device="cpu", expert_bits=None,
ngram_mmap=False, group_size=128):
"""Load the text model from a local snapshot of the official checkpoint. (Your engine: Chapter 30)
Renames multimodal prefixes, skips the vision tower and MTP head, and stacks per-expert
tensors. expert_bits=4 or 8 quantizes each layer's experts as soon as they are read;
ngram_mmap=True leaves the n-gram tables on disk (MappedRows); otherwise they are loaded to
`ngram_device` (host memory by default: lookups touch a few rows per token).
"""
raw = json.loads((Path(directory) / "config.json").read_text())
cfg = FlashNextConfig.from_hf(raw)
with torch.device("meta"):
model = FlashNext(cfg)
for layer in model.model.layers: # never allocate what will be stored differently
if expert_bits:
layer.mlp.experts = nn.Module()
if ngram_mmap and layer.ple is not None:
del layer.ple.ple_embedding.ngram_embedding
model = model.to_empty(device=device).to(dtype)
for index, layer in enumerate(model.model.layers):
if layer.ple is not None: # buffers are derived, not learned: recompute them
fresh = _ngram_buffers(cfg, cfg.ple_layer_ids.index(index + 1))
for name, value in fresh.items():
setattr(layer.ple.ple_embedding, name, value.to(device))
if not ngram_mmap:
layer.ple.ple_embedding.ngram_embedding.to(ngram_device)
params = dict(model.named_parameters())
rename = lambda name: name.replace("model.language_model.", "model.")
is_ngram = lambda name: ".ngram_embedding." in name
experts, ngram_shards = {}, {}
skip = (lambda name: is_ngram(name)) if ngram_mmap else None
for name, value in snapshot_tensors(directory, skip=skip):
name = rename(name)
if name.startswith(("model.visual.", "mtp.", "model.mtp")) or name.endswith(
("layer_multipliers", "ngram_heads_vocab_sizes", "ngram_heads_offsets")):
continue
if ".mlp.experts." in name:
parts = name.split(".")
layer = int(parts[2])
experts.setdefault(layer, {})[(int(parts[5]), parts[6])] = value
if len(experts[layer]) == 3 * cfg.num_experts: # a whole layer has arrived
_install_experts(model, cfg, layer, experts.pop(layer), expert_bits, group_size, device, dtype)
continue
if is_ngram(name):
prefix, _, shard = name.partition(".shard_")
ngram_shards.setdefault(prefix.removesuffix(".weight"), {})[int(shard.split(".")[0]) if shard else 0] = value
continue
if name not in params:
raise ValueError(f"Unexpected tensor {name}")
assign(params[name], value, name)
if experts:
raise ValueError(f"Incomplete experts for layers {sorted(experts)}")
if ngram_mmap:
for original, path in tensor_files(directory).items():
name = rename(original)
if is_ngram(name):
prefix, _, shard = name.partition(".shard_")
ngram_shards.setdefault(prefix.removesuffix(".weight"), {})[
int(shard.split(".")[0]) if shard else 0] = map_tensor(path, original)
for prefix, shards in ngram_shards.items():
ordered = [shards[i] for i in sorted(shards)]
owner = model.get_submodule(prefix.rsplit(".", 1)[0])
if ngram_mmap:
owner.ngram_embedding = MappedRows(ordered, dtype)
else:
assign(params[prefix + ".weight"], torch.cat(ordered, dim=0), prefix)
return model.eval()
def _ngram_buffers(cfg, ple_index):
with torch.device("meta"):
probe = NGramEmbedding(cfg, ple_index) # meta: the big table is not allocated
return {"layer_multipliers": torch.tensor(layer_multipliers(cfg.vocab_size, cfg.ngram_size, ple_index, cfg.seed)),
"ngram_heads_vocab_sizes": torch.tensor(probe._sizes),
"ngram_heads_offsets": torch.tensor([sum(probe._sizes[:i]) for i in range(len(probe._sizes))])}
def _install_experts(model, cfg, layer, tensors, bits, group_size, device, dtype):
gate_up, down = stack_expert_tensors(tensors, cfg.num_experts, cfg.moe_intermediate_size)
block = model.model.layers[layer].mlp
if bits:
block.experts = QuantizedExperts(gate_up.to(device, dtype), down.to(device, dtype), bits, group_size)
else:
assign(block.experts.gate_up_proj, gate_up, "gate_up_proj")
assign(block.experts.down_proj, down, "down_proj")
The loader streams the safetensors shards one at a time and handles the checkpoint’s quirks: the multimodal prefix model.language_model. is renamed, the vision tower and MTP weights are skipped, per-expert tensors are stacked, and n-gram tables stored as shards are concatenated. Two options make it fit on one machine:
class QuantizedExperts(nn.Module):
"""Routed experts stored as groupwise INT4/INT8 codes (Chapter 20). An expert is dequantized
only when a token is routed to it, so a BF16 copy of all experts never exists."""
def __init__(self, gate_up, down, bits=4, group_size=128):
super().__init__()
self.bits, self.dtype, self.shapes, self.groups = bits, gate_up.dtype, {}, {}
for name, w in (("gate_up", gate_up), ("down", down)):
experts, rows, cols = w.shape
codes, scales = quantize_groupwise(w.reshape(experts * rows, cols), bits, group_size)
codes = pack_int4(codes).view(experts, -1) if bits == 4 else codes.view(experts, rows, cols)
self.register_buffer(f"{name}_codes", codes)
self.register_buffer(f"{name}_scales", scales.view(experts, rows, -1).to(torch.float16))
self.shapes[name], self.groups[name] = (rows, cols), group_size
self.num_experts = gate_up.shape[0]
def weight(self, name, e):
codes = getattr(self, f"{name}_codes")[e]
if self.bits == 4:
codes = unpack_int4(codes, self.shapes[name])
scales = getattr(self, f"{name}_scales")[e].float()
return dequantize_groupwise(codes, scales, self.groups[name]).to(self.dtype)
def expert(self, e, x):
gate, up = F.linear(x, self.weight("gate_up", e)).chunk(2, dim=-1)
return F.linear(F.silu(gate) * up, self.weight("down", e))
forward_loop = Experts.forward_loop
forward_grouped = Experts.forward_grouped
class MappedRows:
"""An embedding table left on disk: rows are read through memory maps when looked up, and the
operating system's page cache keeps the hot ones in RAM. Lookups touch a few rows per token,
so the 95 GiB n-gram table never has to fit in memory."""
def __init__(self, shards, dtype):
self.shards, self.dtype, self.device = shards, dtype, torch.device("cpu")
sizes = torch.tensor([0] + [t.shape[0] for t in shards])
self.starts = sizes.cumsum(0)
@property
def weight(self):
return self
def __getitem__(self, index):
flat = index.reshape(-1).cpu()
shard = torch.searchsorted(self.starts, flat, right=True) - 1
out = torch.empty(flat.numel(), self.shards[0].shape[1], dtype=self.dtype)
for s in shard.unique().tolist():
hit = shard == s
out[hit] = self.shards[s][flat[hit] - self.starts[s]].to(self.dtype)
return out.view(*index.shape, -1)
expert_bits=4replaces each layer’s experts withQuantizedExpertsas soon as that layer’s 1,536 expert tensors have arrived: the BF16 experts of more than one layer never exist at once, and the model is built on the meta device so their full-size placeholders are never allocated. An expert is dequantized only when a token is routed to it (a W4A16 grouped kernel, Chapters 20 and 27, is the fast version).ngram_mmap=Truenever reads the n-gram shards into memory.MappedRowswraps memory maps of the safetensors files; a lookup reads a few pages and the operating system caches the hot ones.
The memory plan
def memory_plan(cfg, weight_bits=16, expert_bits=4, ngram_bits=16, context=32768, kv_bits=16):
"""Rough byte budget for the text model: what must be resident, and where."""
d, l = cfg.hidden_size, cfg.num_hidden_layers
n_attn = sum(t != "linear_attention" for t in cfg.layer_types)
n_lin = l - n_attn
expert = cfg.num_experts * 3 * d * cfg.moe_intermediate_size * l
shared = 3 * d * cfg.shared_expert_intermediate_size * l + l * (cfg.num_experts + 1) * d
kd, vd = cfg.linear_num_key_heads * cfg.linear_key_head_dim, cfg.linear_num_value_heads * cfg.linear_value_head_dim
linear = n_lin * (d * (2 * kd + vd) + d * vd + vd * d + 2 * d * cfg.linear_num_value_heads)
attn = n_attn * (d * cfg.num_attention_heads * cfg.head_dim * 3 + 2 * d * cfg.num_key_value_heads * cfg.head_dim)
hc = (2 * l + 1) * 2 * cfg.hc_count * d * cfg.hc_lowrank
embed = 2 * cfg.vocab_size * d
ngram_rows = sum(next_primes(cfg.ngram_vocab_size_base - 1, (cfg.ngram_size - 1) * cfg.heads_per_ngram * len(cfg.ple_layer_ids)))
ngram = ngram_rows * cfg.ple_embed_dim // ((cfg.ngram_size - 1) * cfg.heads_per_ngram)
kv = 2 * n_attn * context * cfg.num_key_value_heads * cfg.head_dim * kv_bits // 8
recurrent = n_lin * cfg.linear_num_value_heads * cfg.linear_key_head_dim * cfg.linear_value_head_dim * 4
gib = 1024 ** 3
return {
"routed_experts_GiB": expert * expert_bits / 8 / gib,
"dense_weights_GiB": (shared + linear + attn + hc + embed) * weight_bits / 8 / gib,
"ngram_tables_GiB": ngram * ngram_bits / 8 / gib,
"kv_cache_GiB_per_sequence": kv / gib,
"recurrent_state_MiB_per_sequence": recurrent / 1024 ** 2,
"active_params_per_token_B": (cfg.num_experts_per_tok * 3 * d * cfg.moe_intermediate_size * l
+ shared + linear + attn + hc) / 1e9,
}
python run.py flashnext
{"expert_bits": 16, "context": 32768, "routed_experts_GiB": 225.0, "dense_weights_GiB": 9.11, "ngram_tables_GiB": 95.37, "kv_cache_GiB_per_sequence": 0.75, "recurrent_state_MiB_per_sequence": 108.0, "active_params_per_token_B": 5.98}
{"expert_bits": 8, "context": 32768, "routed_experts_GiB": 112.5, ...}
{"expert_bits": 4, "context": 32768, "routed_experts_GiB": 56.25, ...}
{"expert_bits": 4, "context": 262144, "routed_experts_GiB": 56.25, ..., "kv_cache_GiB_per_sequence": 6.0, ...}
Putting it together for three machines (estimates from the plan, not measurements):
| machine | experts | dense weights | n-gram table | resident total | decode ceiling |
|---|---|---|---|---|---|
| DGX Spark (128 GB unified, 273 GB/s) | 4-bit, in memory | BF16 | memory-mapped from NVMe | ~68 GiB | ~28 tokens/s |
| same, dense weights in INT8 | 4-bit | INT8 | memory-mapped | ~63 GiB | ~50 tokens/s |
| H100 80 GB (3,350 GB/s) | 4-bit | BF16 | host RAM | ~68 GiB | ~350 tokens/s |
| RTX 4090 24 GB + 128 GB host RAM | 4-bit, offloaded to host (stretch) | BF16 | host RAM | 11 GiB on GPU | PCIe-bound |
The ceilings use Chapter 10’s rule with the bytes read per token: 2.36B active expert parameters at 4 bits (1.2 GB) plus 4.25B other active parameters including the LM head at 2 bytes (8.5 GB), about 9.7 GB per token in BF16, or 5.4 GB with INT8 dense weights. The dense part, not the experts, dominates decode traffic once the experts are 4-bit: quantize it next.
Run it
hf download Qwen/Qwen3.8-Flash-Next --local-dir models/Qwen3.8-Flash-Next # ~360 GB on disk
python run.py chat --model-dir models/Qwen3.8-Flash-Next --expert-bits 4 --ngram-mmap \
--prompt "Explain hyper-connections in two sentences." --new-tokens 128
Important
This book’s code was validated against the official implementation on tiny random checkpoints with the official architecture (Appendix F); the full checkpoint was not run in the validation environment. On a real machine, climb Chapter 18’s ladder of evidence again: compare your logits with Transformers’ on a few hundred tokens in BF16, layer by layer if they differ, before trusting generations.
Multi-token prediction
The checkpoint also ships about 4B parameters of multi-token prediction (MTP) heads, which both your loader and Transformers skip. In the DeepSeek-V3 style, an MTP module takes the final hidden state at position $t$ and the embedding of token $t+1$ and predicts token $t+2$, reusing the model’s embedding and head. At inference, it’s a built-in draft model for speculative decoding (Chapter 26): one cheap extra module proposes the next token, and the main model verifies. With HybridState snapshots, your speculative loop needs only one change: instead of truncating caches after a rejection, restore the snapshot and re-run the accepted tokens. That’s the first stretch exercise.
Build it
Engine milestone 30: the capstone. Implement in engine/flashnext.py: GatedResidual.read, NGramEmbedding.shifted and forward, PLELayer.forward, FlashNextLayer.forward, FlashNext.step and load_flashnext (the configuration, hashing constants, quantized experts, mapped tables, state classes and memory plan are provided).
pytest tests/test_ch30_flashnext.py
python run.py flashnext --impl engine
The tests check the hashing constants, logit parity with Transformers’ Qwen4ExpForCausalLM (including an EOS mid-sequence), that chunked prefill plus decode equals the full forward, free snapshots, the real model’s memory plan, memory-mapped and quantized loading, and generation through your LLM engine. When they pass, your engine runs Qwen3.8-Flash-Next.
Stretch exercises
- ★★ Speculative decoding for Flash-Next: change
speculative_generateto snapshot and restoreHybridCacheinstead of truncating, and use a smaller model with the same tokenizer as the draft. Verify greedy outputs are unchanged. Where:speculative_generateinengine/speculative.py, usingHybridCache.snapshot/restoreinengine/flashnext.py. - ★★ Quantize the dense weights to INT8 (attention, linear attention, shared experts, LM head) and measure the logit error and decode speed against the plan’s prediction. Where: weight installation in
load_flashnextinengine/flashnext.py, usingengine.quant. - ★★★ Load and use the MTP head as a draft: read its weights (prefix
mtp.) and implement its forward from the configuration and the DeepSeek-V3 report’s description, then measure acceptance rates. Where: add an MTP module and load its weights inengine/flashnext.py; call it fromengine/speculative.py. - ★★★ Expert offloading for a 24 GB GPU: keep
QuantizedExpertscodes in pinned host memory and copy only the routed experts per layer, overlapping the copy for layer $\ell+1$ with compute for layer $\ell$ using CUDA streams. Report tokens/s against the PCIe bound. Where:QuantizedExpertsand expert installation inengine/flashnext.py. - ★★★ Continuous batching for a hybrid model: give each request a
HybridState, batch the linear-attention layers (every state has the same shape) and the attention layers (ragged KV), and verify solo-equivalence as in Chapter 24. Where: newexperiments/hybrid_batching.py, adaptingengine.scheduler.ContinuousBatchingEngineforengine.flashnext.HybridState.
Check your understanding
- Which Flash-Next components come from which earlier chapters, and which two are new?
- How do the read and write weights of a hyper-connection differ in shape and range?
- Why does the n-gram memory use 8 hash heads with different prime sizes?
- Why can a 51B-parameter table live on disk without slowing decode much?
- Why is a functional (never modified in place) state convenient for speculative decoding?
- After quantizing the experts to 4 bits, what dominates the bytes read per decode token?
Going deeper
- The Qwen3.8-Flash-Next model card and configuration, and
modeling_qwen4_exp.pyin Transformers 5.18+. - Zhu et al., Hyper-Connections (2024), and the follow-ups on manifold-constrained hyper-connections; DeepSeek-AI’s Engram (2026) for hashed n-gram memory; the DeepSeek-V3 technical report for MTP.
- The Qwen3-Next and Qwen3.5 model cards, the direct ancestors of this architecture (3:1 Gated DeltaNet / gated attention, sigmoid-gated shared expert).
- vLLM and SGLang’s model files for Qwen3-Next, to see how production engines batch a hybrid model’s states.
31. Engine v2: one core for every request
In this chapter
- Why the serving pieces of Parts IV-VI don't compose as they are, and the one idea that makes them compose: every request is "compute its next n tokens".
- Flattened ragged batches: prefill chunks and decode tokens of many requests in a single
[N, D]forward pass, with a slot mapping and block tables for attention. - Automatic prefix caching with hash chains: blocks that stay findable after their request finishes, LRU eviction, and why the hash must be cryptographic.
- A unified scheduler with a token budget, chunked prefill, priorities and preemption by recompute or by swapping to host memory.
- Sizing the KV pool from free GPU memory, and the test that keeps all of it honest.
You will build
The core of the production engine in engine/serve/: BlockManager (blocks.py), Scheduler (scheduler.py), build_batch (batch.py), the reference paged attention backend (attention.py), FlatModel (model.py) and EngineCore.step (engine.py).
Time: 8-10 hours. GPU: not needed (everything runs on a CPU; Chapter 32 makes it fast).
Five good pieces that don’t fit together
Look at what Parts IV-VI gave you. Chapter 19’s FastDecoder removes launch overhead, but serves one request. Chapter 24’s ContinuousBatchingEngine serves many, but gives each request a fixed-size row of a static cache and prefills one request’s chunk per forward call. Chapter 25’s PagedKVCache packs caches without waste and shares prefixes, but nothing schedules requests onto it. Chapter 26’s speculative decoding handles one request with a contiguous cache. Each piece passed its tests. None of them can be combined with the others without rewriting the seams between them.
That’s the situation every serving engine reaches. vLLM’s first version grew one feature at a time and was rewritten in 2025 as “V1” around a simpler core; SGLang, TensorRT-LLM and LMDeploy converged on similar designs. This chapter builds that core. The next twelve chapters plug the rest of a production server into it: fast kernels, CUDA graphs, a full sampler, constrained decoding, a tokenizer, an HTTP API, speculative decoding, real checkpoint formats, other hardware, several GPUs and more model families.
The design rests on three decisions:
- One kind of work. A request’s state is its token list (prompt, then output) and one number,
num_computed_tokens: how many of those tokens already have keys and values in the cache. The scheduler’s only output is a list of(request, n): compute the nextntokens of this request. - One batch layout. All scheduled tokens of all requests are concatenated into one flat sequence of $N$ tokens. Matmuls see
[N, D]; attention gets a small metadata object that says where each request’s rows start and where its blocks live. - One memory pool. Every layer’s K and V live in a pool of fixed-size blocks shared by all requests, with reference counts and a prefix cache keyed by hash chains.
One kind of work
With num_computed_tokens as the only state, every situation an engine meets is a value of n:
| situation | tokens in the list | computed | scheduled n |
|---|---|---|---|
| a new 1,000-token prompt, budget 512 | 1,000 | 0 | 512 (a prefill chunk) |
| the same request, next step | 1,000 | 512 | 488 (the rest; its last row produces the first output token) |
| decoding, 37 tokens generated | 1,037 | 1,036 | 1 |
| new request whose first 768 tokens are cached | 900 | 768 (set at admission) | 132 |
| resumed after preemption by recompute | 1,037 | 0 | 512, then 512, then 13 |
| verifying 4 draft tokens (Chapter 37) | 1,037 + 4 drafts | 1,036 | 5 |
The invariant that makes this work: the newest token is always uncomputed. After a decode step samples token $t$, it’s appended to the list, so num_tokens = num_computed_tokens + 1 and the next step feeds exactly that token. A request needs a sample from the model exactly when a step’s tokens reach the end of its list, which is how the batch builder knows which rows to send to the LM head.
@dataclass(eq=False)
class Request:
request_id: str
prompt_token_ids: list
params: SamplingParams = field(default_factory=SamplingParams)
eos_token_id: int | None = None
priority: int = 0 # lower runs first under the "priority" policy (vLLM's convention)
arrival_time: float = field(default_factory=time.monotonic)
cache_key: tuple = () # anything else that changes activations: LoRA, images, tenant salt
status: Status = Status.WAITING
num_computed_tokens: int = 0 # tokens whose K/V are in the cache
num_cached_tokens: int = 0 # of those, how many came from the prefix cache at admission
spec_token_ids: list = field(default_factory=list) # draft tokens to verify this step (Chapter 37)
swapped_out: bool = False # KV lives in host memory (swap preemption)
num_preemptions: int = 0
first_token_time: float | None = None
finish_time: float | None = None
stop_reason: int | str | None = None
extra: dict = field(default_factory=dict) # per-feature state: guided FSM, LoRA slot, images
def __post_init__(self):
if not self.prompt_token_ids:
raise ValueError(f"{self.request_id}: empty prompt")
self.token_ids = list(self.prompt_token_ids) # prompt then output, one list
self.num_prompt_tokens = len(self.token_ids)
self.block_hashes = [] # filled by the block manager
self.num_placeholders = 0 # trailing PLACEHOLDER tokens (async scheduling)
@property
def output_token_ids(self):
return self.token_ids[self.num_prompt_tokens:]
@property
def num_tokens(self):
return len(self.token_ids)
@property
def num_output_tokens(self):
return len(self.token_ids) - self.num_prompt_tokens
@property
def num_tokens_with_spec(self):
return len(self.token_ids) + len(self.spec_token_ids)
def append(self, token_id):
self.token_ids.append(int(token_id))
num_cached_tokens and num_preemptions are statistics; cache_key and extra are hooks for later chapters (LoRA adapters, images and per-tenant salts change activations, so they must be part of the prefix-cache key). Status distinguishes waiting, running, preempted and the three ways to finish.
Flattened ragged batches
Chapter 24’s engine ran one forward call per prefill chunk, plus one for all decodes. With twenty requests in flight that’s up to twenty small forward calls per step, each reading every weight from memory. The fix is to give the model all of a step’s tokens at once, without padding:
request: a (chunk of 3) b (decode) c (decode) d (new, 4 tokens)
input_ids: a6 a7 a8 b41 c9 d0 d1 d2 d3
positions: 6 7 8 41 9 0 1 2 3
query_start_loc: 0 3 4 5 9
seq_lens: 9 42 10 4
The linear layers, norms and MLPs don’t care which request a row belongs to: they process [9, D] as one matrix, so the weights are read once for the whole step and the matmul is as large as the batch allows. Only attention needs the structure, and it gets it from the metadata:
@dataclass
class BatchMeta:
query_start_loc: torch.Tensor
seq_lens: torch.Tensor
block_table: torch.Tensor
slot_mapping: torch.Tensor
block_size: int
query_start_loc_cpu: list # host copies: kernels launch without a device sync
seq_lens_cpu: list
max_query_len: int
max_seq_len: int
@property
def num_reqs(self):
return len(self.seq_lens_cpu)
@property
def decode_only(self):
return self.max_query_len == 1
Two pieces connect tokens to the pool. The block table of request $r$ lists its physical blocks in logical order, exactly as in Chapter 25. The slot mapping gives, for each of the $N$ new tokens, the flat pool slot where its K and V go: position $p$ of a request lives in block table[p // block_size] at offset p % block_size, so its slot is table[p // block_size] * block_size + p % block_size. Writing a step’s keys is then one scatter over $N$ slots, whatever mix of requests the step contains.
def build_batch(scheduled, block_tables, block_size, device="cpu", pad_to=None):
"""Lay out [(request, n), ...] as one flattened batch. (Your engine: Chapter 31)
block_tables[request_id] lists the request's physical blocks in logical order. Request r's
n new tokens are token_ids[c : c + n] (plus its draft tokens, if any), at positions
c .. c + n - 1, where c is its num_computed_tokens. Position p lives in slot
table[p // block_size] * block_size + p % block_size.
pad_to (CUDA graphs, Chapter 33) appends dummy decode rows with slot -1 and length 0.
"""
ids, positions, slots, starts, lengths, tables, logits_rows, counts = [], [], [], [0], [], [], [], []
prompt_rows, prompt_spans = [], []
for request, n in scheduled:
c = request.num_computed_tokens
tokens = (request.token_ids + request.spec_token_ids)[c:c + n]
if len(tokens) != n:
raise ValueError(f"{request.request_id}: scheduled {n} tokens but only {len(tokens)} exist")
table = block_tables[request.request_id]
ids.extend(tokens)
positions.extend(range(c, c + n))
slots.extend(table[p // block_size] * block_size + p % block_size for p in range(c, c + n))
starts.append(starts[-1] + n)
lengths.append(c + n)
tables.append(table)
# Rows from the last real token onwards produce samples: one for a finished prefill or a
# decode, 1 + (draft tokens scheduled) when verifying. A mid-prompt chunk produces none.
k = min(n, max(0, c + n - (request.num_tokens - 1)))
logits_rows.extend(range(starts[-1] - k, starts[-1]))
counts.append(k)
if request.params.prompt_logprobs is not None: # rows whose next token is a prompt token
last = min(c + n, request.num_prompt_tokens - 1)
if last > c:
prompt_rows.extend(range(starts[-2], starts[-2] + last - c))
prompt_spans.append((request, c, request.token_ids[c + 1:last + 1]))
if pad_to is not None:
for _ in range(pad_to - len(lengths)):
ids.append(0), positions.append(0), slots.append(-1)
starts.append(starts[-1] + 1)
lengths.append(0)
tables.append([])
width = max(1, max(len(t) for t in tables))
table_tensor = torch.zeros((len(tables), width), dtype=torch.int32)
for r, table in enumerate(tables):
table_tensor[r, :len(table)] = torch.tensor(table, dtype=torch.int32)
meta = BatchMeta(
query_start_loc=torch.tensor(starts, dtype=torch.int32, device=device),
seq_lens=torch.tensor(lengths, dtype=torch.int32, device=device),
block_table=table_tensor.to(device),
slot_mapping=torch.tensor(slots, dtype=torch.int64, device=device),
block_size=block_size, query_start_loc_cpu=starts, seq_lens_cpu=lengths,
max_query_len=max(b - a for a, b in zip(starts, starts[1:])), max_seq_len=max(lengths))
return Batch(torch.tensor(ids, dtype=torch.long, device=device),
torch.tensor(positions, dtype=torch.long, device=device), meta,
torch.tensor(logits_rows + prompt_rows, dtype=torch.long, device=device), counts,
[r.request_id for r, _ in scheduled], {}, prompt_spans)
logits_indices selects the rows whose logits the sampler needs. A decode contributes its one row; a prefill chunk that ends at the prompt’s last token contributes that row; a mid-prompt chunk contributes nothing. This is a large saving: Qwen3-8B’s LM head is a 4,096 × 151,936 matmul, the single largest in the model. A 2,000-token prefill that sent every row through it would spend more time there than in the 36 layers’ attention. Production engines all compute logits only where they’re sampled.
The pad_to argument, used in Chapter 33, adds dummy rows with slot -1 and length 0, so a decode batch of 13 requests can run in a CUDA graph captured for 16.
Attention over the pool
The pool stores, per layer, k_cache and v_cache of shape [num_blocks, block_size, kv_heads, head_dim]. That’s a different axis order from Chapter 25’s [num_blocks, kv_heads, block_size, head_dim]. Slot-major rows make writing one index_copy_ over flat slot numbers:
def write_kv(k_cache, v_cache, k, v, slot_mapping):
"""Scatter k, v [N, Hkv, D] into the pools at flat slots. (Your engine: Chapter 31)
Slot -1 marks a padding row (Chapter 33's CUDA graphs pad batches). Its write goes to the
pool's last slot, a scratch block the runner allocates but the block manager never hands
out, instead of being filtered out: filtering would need the host to know how many rows
are real, a sync that a CUDA graph can't contain.
"""
scratch = k_cache.shape[0] * k_cache.shape[1] - 1
slots = torch.where(slot_mapping >= 0, slot_mapping, scratch)
k_cache.view(-1, *k_cache.shape[2:]).index_copy_(0, slots, k.to(k_cache.dtype))
v_cache.view(-1, *v_cache.shape[2:]).index_copy_(0, slots, v.to(v_cache.dtype))
The reference backend gathers each request’s context through its block table and calls the causal_attention you wrote in Chapter 5. The positions do the work, as they have since Chapter 16: request $r$’s $t$ queries are the last $t$ positions of its seq_len-token context, so query $i$ is at position seq_len - t + i.
class ReferenceBackend:
"""Gathers each request's context and calls Chapter 5's causal_attention. Slow and obviously
correct: every faster backend is tested against it."""
name = "reference"
def allocate(self, shape, dtype, device):
return torch.zeros(shape, dtype=dtype, device=device), torch.zeros(shape, dtype=dtype, device=device)
def write(self, cache, k, v, meta):
write_kv(cache[0], cache[1], k, v, meta.slot_mapping)
def forward(self, q, cache, meta, scale=None, window=0):
"""(Your engine: Chapter 31)
Request r's queries are rows query_start_loc[r]:query_start_loc[r+1]; they are the
LAST t positions of its seq_len-token context, so query i sits at position
seq_len - t + i and may see keys 0 .. that position (and, with a sliding window of
w > 0, only the last w of them: Chapter 42).
"""
k_cache, v_cache = cache[0], cache[1]
out = torch.zeros_like(q)
starts, lengths = meta.query_start_loc_cpu, meta.seq_lens_cpu
for r, length in enumerate(lengths):
begin, end = starts[r], starts[r + 1]
if length == 0 or begin == end: # padding row
continue
t = end - begin
k = gather_context(k_cache, meta.block_table[r], length)
v = gather_context(v_cache, meta.block_table[r], length)
q_pos, k_pos = torch.arange(length - t, length, device=q.device), torch.arange(length, device=q.device)
allowed = (k_pos[None, :] > q_pos[:, None] - window) if window > 0 else None
y = causal_attention(q[begin:end].transpose(0, 1)[None], self.dequantize(k, cache, 0, meta, r, length),
self.dequantize(v, cache, 1, meta, r, length), q_pos, k_pos, allowed, scale)
out[begin:end] = y[0].transpose(0, 1)
return out
@staticmethod
def dequantize(x, cache, which, meta, r, length):
"""[length, Hkv, D] -> [1, Hkv, length, D] in float; quantized pools carry scales at cache[2:]."""
x = x.float()
if len(cache) > 2:
x = x * gather_context(cache[2 + which], meta.block_table[r], length)[..., None]
return x.transpose(0, 1)[None]
It loops over requests in Python, which is slow, and it’s meant to be: every faster backend in Chapter 32 is tested against it. Chapter 32’s Triton kernel does the same computation in one launch, reading K/V tiles straight from the blocks.
The flat model
The model needs one change: attention reads and writes the pool through a backend, and every tensor is [N, ...] instead of [B, T, ...]. Rather than writing a new model, FlatModel reuses the modules and weights of the Qwen3 (or Qwen3Moe) you already loaded and tested. Nothing is copied or renamed, so every parity test of Chapters 17 and 27 still covers the weights.
class FlatModel(nn.Module):
def __init__(self, model):
super().__init__()
self.model = model.eval()
self.cfg = model.cfg
self.backbone = model.model
self.layers = model.model.layers
self.rope = getattr(model, "rope_tables", None) or self.default_rope
self.fused = False
def default_rope(self, positions, dtype):
cos, sin = rope_cos_sin(positions, self.cfg.head_dim, self.cfg.rope_theta, dtype)
return cos[0, 0][:, None], sin[0, 0][:, None] # [N, 1, rotary_dim]: broadcast over heads
def kv_spec(self):
"""(layers, kv_heads, head_dim) of the pool every attention layer writes."""
custom = getattr(self.model, "kv_spec", None) # latent caches (Chapter 42)
if custom is not None:
return custom()
return self.cfg.num_hidden_layers, self.cfg.num_key_value_heads, self.cfg.head_dim
def attention(self, layer, attn, h, positions, rope, kv, meta, backend):
"""One attention sublayer on [N, D] rows. (Your engine: Chapter 31)
The same steps as Qwen3Attention.forward, but heads are the middle axis of [N, H, D]
and the cache is the pool: write this step's keys (after RoPE), then attend.
"""
paged = getattr(attn, "paged_forward", None) # attention of another shape: MLA (Chapter 42)
if paged is not None:
return paged(h, positions, rope, kv, meta, backend)
n = h.shape[0]
c = self.cfg
q, k, v = self.project_qkv(attn, h)
q = q.view(n, c.num_attention_heads, c.head_dim)
k = k.view(n, c.num_key_value_heads, c.head_dim)
v = v.view(n, c.num_key_value_heads, c.head_dim)
if hasattr(attn, "q_norm"):
q, k = attn.q_norm(q), attn.k_norm(k)
cos, sin = rope
q, k = apply_rope(q, cos, sin), apply_rope(k, cos, sin)
backend.write(kv, k, v, meta)
out = backend.forward(q, kv, meta, window=getattr(attn, "sliding_window", 0))
return attn.o_proj(out.reshape(n, -1))
def forward(self, input_ids, positions, kv_caches, meta, backend, embeds=None, **features):
"""input_ids [N] -> final hidden states [N, D] (before the LM head). (Your engine: Chapter 31)
add_norm(norm, x, delta) returns (norm(x + delta), x + delta): the residual add and the
next norm, which Chapter 33 fuses into one kernel. mlp(module, h) runs the MLP or MoE.
"""
bank = getattr(self, "adapters", None)
with bank.activate(features.get("lora_ids"), meta.decode_only) if bank is not None else nullcontext():
return self._forward(input_ids, positions, kv_caches, meta, backend, embeds, features)
def _forward(self, input_ids, positions, kv_caches, meta, backend, embeds, features):
x = self.backbone.embed_tokens(input_ids) if embeds is None else embeds
if "mrope_positions" in features:
from ..multimodal import mrope_cos_sin
rope = mrope_cos_sin(features["mrope_positions"], self.cfg.head_dim, self.cfg.rope_theta,
features["mrope_sections"], x.dtype)
else:
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)
h, _ = self.add_norm(self.backbone.norm, x, delta)
return h
def compute_logits(self, hidden):
"""Only the rows the sampler needs reach the LM head: for a 2,000-token prefill
that is 1 row instead of 2,000, the largest matmul in the model skipped."""
return self.model.lm_head(hidden).float()
RoPE tables are computed once per step for all $N$ positions and broadcast over heads. The MoE block of Chapter 27 already flattens its input to [N, D], so Qwen3Moe runs through the same code unchanged. Chapter 42 turns this class into a registry that covers Llama, Mistral, Qwen2, DeepSeek’s latent attention and more.
Prefix caching with hash chains
Chapter 25’s PrefixCache keyed each cached block by the entire token prefix up to its end. That’s correct but costly: the key for block 1,000 of a long document is a 16,000-token tuple, and comparing keys is proportional to their length. Production engines use a hash chain instead:
$$ h_0 = H(\text{root},\ \text{tokens}0,\ \text{extra}), \qquad h_i = H(h{i-1},\ \text{tokens}_i,\ \text{extra}) $$
Each key has a fixed size, depends on the whole history through $h_{i-1}$, and is computed once per block as the request grows. extra holds anything else that changes the keys and values: the LoRA adapter (Chapter 43), hashes of images in the prompt (Chapter 43), and a per-tenant salt if users must never share cache entries.
def hash_block(parent, tokens, extra=()):
"""Key for one full block: depends on the previous block's key, this block's tokens and
any extra keys (LoRA adapter, image hashes, tenant salt). (Your engine: Chapter 31)
A cryptographic hash, not Python's hash(): a crafted collision would let one user's
request read another user's cached keys and values.
"""
h = hashlib.blake2b(digest_size=16)
h.update(parent if parent is not None else b"\x00root")
h.update(array("q", tokens).tobytes())
if extra:
h.update(repr(tuple(extra)).encode())
return h.digest()
Warning
Use a cryptographic hash. With Python’s
hash()or a weak 64-bit hash, an attacker who can submit prompts can search for a token sequence whose block hash collides with someone else’s, and their request would then attend to the victim’s cached keys and values, which leaks information about the victim’s prompt. vLLM offers SHA-256 block hashing (--prefix-caching-hash-algo) for this reason. BLAKE2b with a 128-bit digest is as safe for this purpose and fast in Python.
Blocks that outlive their requests
The second improvement over Chapter 25 is what happens when a request finishes. Chapter 25’s prefix cache held an extra reference on every published block, so cached blocks were never free, and you had to evict them explicitly. vLLM V1’s design is better: a freed block keeps its contents and its hash, and goes onto the free queue. It’s still findable by hash and can be reused for free, but it’s also the next candidate for allocation. Only when the allocator actually hands it out is it evicted from the cache. A block is in one of three states:
| state | ref count | in free queue | findable by hash |
|---|---|---|---|
| in use | ≥ 1 | no | yes, if full and published |
| cached free | 0 | yes | yes |
| empty free | 0 | yes | no |
The free queue is an LRU list: allocation takes from the front, freeing appends to the back. One detail matters for hit rates: a finished request’s blocks are appended tail first, so its last blocks (unique to that request) are evicted before its first blocks (the system prompt and shared history, most likely to be reused).
class BlockManager:
"""Reference-counted block pool with an LRU free queue and a hash -> block map."""
def __init__(self, num_blocks, block_size, enable_prefix_caching=True):
if num_blocks < 1 or block_size < 1:
raise ValueError("Need at least one block of at least one token")
self.num_blocks, self.block_size = num_blocks, block_size
self.enable_prefix_caching = enable_prefix_caching
self.ref_count = [0] * num_blocks
self.block_hash = [None] * num_blocks
self.free_queue = OrderedDict((b, None) for b in range(num_blocks)) # front = evicted first
self.cached = {} # hash -> block id
self.req_blocks = {} # request id -> [block ids], logical order
self.num_registered = {} # request id -> how many of its blocks are published by hash
self.stats = {"queries": 0, "hits": 0}
@property
def num_free(self):
return len(self.free_queue)
@property
def usage(self):
return 1.0 - self.num_free / self.num_blocks
def blocks_for(self, num_tokens):
return -(-num_tokens // self.block_size)
def update_hashes(self, request):
"""Extend request.block_hashes to cover every full block of its known tokens (not the
placeholders of tokens still being computed, Chapter 33)."""
size, hashes = self.block_size, request.block_hashes
known = request.num_tokens - request.num_placeholders
while (len(hashes) + 1) * size <= known:
start = len(hashes) * size
parent = hashes[-1] if hashes else None
hashes.append(hash_block(parent, request.token_ids[start:start + size], request.cache_key))
return hashes
def find_cached_prefix(self, request):
"""Longest run of cached full blocks at the start of the request. (Your engine: Chapter 31)
Stops at the first miss. Leaves at least one token to compute, because the last
prompt position's logits are what choose the first output token.
"""
if not self.enable_prefix_caching or request.params.prompt_logprobs is not None:
return [] # prompt logprobs need every prompt position computed
hashes = self.update_hashes(request)
limit = (request.num_tokens - 1) // self.block_size
found = []
for h in hashes[:limit]:
block = self.cached.get(h)
if block is None:
break
found.append(block)
self.stats["queries"] += request.num_tokens
self.stats["hits"] += len(found) * self.block_size
return found
def allocate_slots(self, request, num_new_tokens, cached_blocks=()):
"""Make room for num_new_tokens more tokens of request, after its computed tokens and
any cached_blocks it is about to adopt. Returns False, changing nothing, if the pool is
too small; the scheduler then preempts someone or waits. (Your engine: Chapter 31)
"""
blocks = self.req_blocks.setdefault(request.request_id, [])
total = request.num_computed_tokens + len(cached_blocks) * self.block_size + num_new_tokens
needed = self.blocks_for(total) - len(blocks) - len(cached_blocks)
# Adopting a cached-but-free block takes it out of the free queue too.
reclaimed = sum(1 for b in cached_blocks if self.ref_count[b] == 0)
if needed > self.num_free - reclaimed:
return False
for b in cached_blocks:
if self.ref_count[b] == 0:
del self.free_queue[b]
self.ref_count[b] += 1
blocks.append(b)
for _ in range(max(needed, 0)):
blocks.append(self._pop_free())
if cached_blocks:
self.num_registered[request.request_id] = len(cached_blocks)
return True
def _pop_free(self):
block, _ = self.free_queue.popitem(last=False) # least recently freed
h = self.block_hash[block]
if h is not None: # evict it from the prefix cache
if self.cached.get(h) == block:
del self.cached[h]
self.block_hash[block] = None
self.ref_count[block] = 1
return block
def cache_full_blocks(self, request, num_tokens=None):
"""Publish request's full blocks among its first num_tokens tokens (default: the computed
ones) under their hashes. (Your engine: Chapter 31)
The scheduler publishes blocks as soon as they are allocated for this step's tokens,
before they are computed. A request admitted later in the same step may then adopt
them: in a flattened batch every layer writes all K/V before any attention reads, so
it reads them after they are written.
"""
if not self.enable_prefix_caching:
return
blocks = self.req_blocks.get(request.request_id, [])
hashes = self.update_hashes(request)
num_tokens = request.num_computed_tokens if num_tokens is None else num_tokens
full = min(num_tokens // self.block_size, len(hashes), len(blocks))
start = self.num_registered.get(request.request_id, 0)
for i in range(start, full):
block, h = blocks[i], hashes[i]
if self.block_hash[block] is None:
self.block_hash[block] = h
self.cached.setdefault(h, block) # first writer wins; duplicates stay private
self.num_registered[request.request_id] = max(start, full)
def free(self, request):
"""Drop request's references. Blocks are queued tail-first, so a sequence's last blocks
are evicted before its prefix, which is the part most likely to be shared. (Your engine: Chapter 31)"""
for block in reversed(self.req_blocks.pop(request.request_id, [])):
self.ref_count[block] -= 1
if self.ref_count[block] == 0:
self.free_queue[block] = None
elif self.ref_count[block] < 0:
raise RuntimeError(f"Block {block} freed more often than referenced")
self.num_registered.pop(request.request_id, None)
def trim(self, request, num_tokens):
"""Release blocks past num_tokens (rejected speculative tokens, Chapter 37)."""
blocks = self.req_blocks.get(request.request_id, [])
keep = self.blocks_for(num_tokens)
for block in reversed(blocks[keep:]):
self.ref_count[block] -= 1
if self.ref_count[block] == 0:
h = self.block_hash[block]
if h is not None and self.cached.get(h) == block:
del self.cached[h]
self.block_hash[block] = None
self.free_queue[block] = None
del blocks[keep:]
def reset_prefix_cache(self):
"""Forget every cached-free block (for example after loading new weights)."""
for block in self.free_queue:
self.block_hash[block] = None
self.cached = {h: b for h, b in self.cached.items() if self.ref_count[b] > 0}
@property
def hit_rate(self):
return self.stats["hits"] / self.stats["queries"] if self.stats["queries"] else 0.0
Three rules in this code are easy to get wrong:
- Leave one token to compute.
find_cached_prefixnever matches the block that contains the prompt’s last token, even if it’s cached. The first output token is sampled from the logits at the last prompt position, so that position must run through the model. (A request with a 32-token prompt and 16-token blocks reuses at most the first block.) - Adopting a cached-free block costs a free slot.
allocate_slotscounts how many of the adopted blocks are currently in the free queue, because taking them out of it shrinks what’s left for new allocations. Getting this wrong over-commits the pool. - Allocation is all or nothing. If the request doesn’t fit,
allocate_slotsreturnsFalsebefore changing anything. The scheduler can then preempt someone and retry, without undoing a half-finished allocation.
Only full blocks are published, and they’re published as soon as they’re allocated for this step’s tokens, before the forward pass computes them. That looks premature, but it’s safe and it matters: four requests that share a long prompt and arrive together (or the four samples of an n=4 request, Chapter 34) are admitted in the same step, and only the first would compute the prompt if the others could adopt its blocks right away. In a flattened batch every layer writes all $N$ tokens’ keys and values before any attention reads them, so a request that adopted a block in this step reads it after it was written. The engine also publishes after each step, which covers generated tokens as they fill blocks. That’s what makes multi-turn chat cheap: turn 2’s prompt is turn 1’s prompt plus turn 1’s answer plus the new message, and most of it is already cached.
The scheduler
With one kind of work, the scheduler is short:
class Scheduler:
def __init__(self, config, block_manager):
self.config, self.blocks = config, block_manager
self.waiting, self.running = deque(), []
self.requests = {}
self.on_finish = None # called before a finished request's blocks are freed (Chapter 41)
def add(self, request):
capacity = self.blocks.num_blocks * self.blocks.block_size
if request.num_tokens + request.params.max_tokens > capacity:
raise ValueError(f"{request.request_id}: needs more KV slots than the whole pool holds ({capacity})")
if not self.config.enable_chunked_prefill and request.num_tokens > self.config.max_num_batched_tokens:
raise ValueError(f"{request.request_id}: prompt exceeds the token budget and chunked prefill is off")
self.requests[request.request_id] = request
self._enqueue(request)
def _enqueue(self, request, front=False):
if self.config.policy == "priority":
key = (request.priority, request.arrival_time)
index = next((i for i, r in enumerate(self.waiting) if (r.priority, r.arrival_time) > key), len(self.waiting))
self.waiting.insert(index, request)
elif front:
self.waiting.appendleft(request)
else:
self.waiting.append(request)
def _chunk(self, n, budget):
limit = self.config.long_prefill_token_threshold
return min(n, budget, limit) if limit else min(n, budget)
def _preempt(self, request, out):
"""Take request out of the running set and give its blocks back. (Your engine: Chapter 31)
recompute: forget its KV; when readmitted it prefills prompt + output again (and
probably hits its own blocks in the prefix cache). swap: its blocks are copied to host
memory first, and num_computed_tokens is kept.
"""
self.running.remove(request)
request.spec_token_ids = [] # drafts are only valid for the very next step
if self.config.preemption_mode == "swap":
out.swap_out.append((request, list(self.blocks.req_blocks[request.request_id])))
request.swapped_out = True
else:
request.num_computed_tokens = 0
self.blocks.free(request)
request.status = Status.PREEMPTED
request.num_preemptions += 1
out.preempted.append(request)
self._enqueue(request, front=True)
def schedule(self):
"""Choose (request, num_new_tokens) pairs for one step. (Your engine: Chapter 31)"""
out = SchedulerOutput()
budget = self.config.max_num_batched_tokens
if self.config.policy == "priority":
self.running.sort(key=lambda r: (r.priority, r.arrival_time))
# 1. Running requests: decodes, unfinished prefill chunks, verification blocks.
index = 0
while index < len(self.running) and budget > 0:
request = self.running[index]
n = self._chunk(request.num_tokens_with_spec - request.num_computed_tokens, budget)
if request.num_output_tokens >= request.params.max_tokens:
n = 0 # its last token is in flight (Chapter 33)
if n <= 0:
index += 1
continue
while not self.blocks.allocate_slots(request, n):
victim = self.running[-1] # lowest priority: last in order
self._preempt(victim, out)
if victim is request:
break
if request.status is Status.PREEMPTED:
break # it was the last one; nothing left to try
out.scheduled.append((request, n))
self.blocks.cache_full_blocks(request, request.num_computed_tokens + n)
budget -= n
index += 1
# 2. Waiting requests, unless memory was just so tight that we preempted.
while self.waiting and budget > 0 and len(self.running) < self.config.max_num_seqs and not out.preempted:
request = self.waiting[0]
cached = [] if request.num_computed_tokens else self.blocks.find_cached_prefix(request)
remaining = request.num_tokens - request.num_computed_tokens - len(cached) * self.blocks.block_size
if not self.config.enable_chunked_prefill and remaining > budget:
break
n = self._chunk(remaining, budget)
if not self.blocks.allocate_slots(request, n, cached):
break
self.waiting.popleft()
if request.swapped_out:
out.swap_in.append((request, list(self.blocks.req_blocks[request.request_id])))
request.swapped_out = False
request.num_computed_tokens += len(cached) * self.blocks.block_size
if request.status is Status.WAITING:
request.num_cached_tokens = len(cached) * self.blocks.block_size
request.status = Status.RUNNING
self.running.append(request)
out.scheduled.append((request, n))
self.blocks.cache_full_blocks(request, request.num_computed_tokens + n) # visible to later admissions
budget -= n
return out
Running requests come first. They hold blocks and users are watching their streams; a new request can wait one more step. Within the running list, order is arrival (FCFS) or (priority, arrival). Each running request asks for everything it still needs: one token if it’s decoding, the rest of its prompt if it’s mid-prefill (capped by long_prefill_token_threshold, the chunk size), or its history if it was just resumed.
Admission spends what’s left. A waiting request starts after its longest cached prefix, so a request whose 2,000-token system prompt is cached costs only its own question. If the pool can’t hold it, admission stops: the scheduler never preempts a running request to admit a new one.
Chunked prefill falls out of the budget: a 5,000-token prompt with a 2,048-token budget runs as 2,048 + 2,048 + 904 over three steps, sharing each step with every running request’s decode token. Chapter 24 explained why: decode latency stays steady while long prompts arrive. With enable_chunked_prefill=False a prompt must fit in one step’s budget, which some engines prefer for simplicity at low load.
When memory runs out, the scheduler preempts the lowest-priority running request (the last in order) and retries. If the victim is the request it was trying to schedule, it stops: nothing with lower priority is left. After any preemption it also stops admitting, because the pool is evidently full.
Preemption: recompute or swap
A preempted request gives back its blocks. There are two ways to bring it back later:
- Recompute (
preemption_mode="recompute"): forget its K/V and setnum_computed_tokens = 0. When readmitted, it prefills prompt plus output again, in chunks like any prompt. Its blocks were just freed into the queue with their hashes, so if they haven’t been reused yet, the prefix cache gives most of them back for free. - Swap (
preemption_mode="swap"): copy its blocks to host memory before giving them back, and keepnum_computed_tokens. When readmitted, allocate fresh blocks and copy the K/V back. The runner does both copies before the step’s forward pass, so the freed blocks can be reused by other requests in the same step.
Which costs less? For Qwen3-8B in BF16, one token’s K/V across all 36 layers is $2 \times 36 \times 8 \times 128 \times 2 = 147{,}456$ bytes. Swapping it out and back over PCIe 5.0 at about 50 GB/s takes $2 \times 147{,}456 / 50 \times 10^9 \approx 5.9$ µs. Recomputing it costs about $2 \times 8 \times 10^9 = 16$ GFLOPs, roughly 27 µs at a realistic 600 TFLOP/s on an H100. So swapping is cheaper per token, and the copy can overlap compute. vLLM V1 nevertheless supports only recompute: preemption should be rare in a well-sized engine, recompute needs no host-memory management, and with prefix caching the “recompute” is often a cache hit. Both modes are here so you can measure the trade-off on your own workload (stretch exercise 2). For models with a tiny KV cache per token, such as DeepSeek’s latent attention (Chapter 42), swapping is cheaper still.
The engine loop
EngineCore.step ties the parts together: schedule, perform swaps, build one batch, run it, sample, update.
class EngineCore:
def __init__(self, model, config=None, eos_token_id=None, sampler=None, vocab_bytes=None):
self.config = config = config or EngineConfig()
p = next(model.parameters())
self.device = p.device
flat = model if isinstance(model, FlatModel) else FlatModel(model)
options = {"kv_cache_dtype": config.kv_cache_dtype} if config.kv_cache_dtype != "auto" else {}
backend = get_backend(config.attention_backend, **options)
layers, kv_heads, head_dim = flat.kv_spec()
num_blocks = config.num_blocks
if num_blocks is None:
if self.device.type == "cuda":
num_blocks = profile_num_blocks(flat, backend, config.block_size, p.dtype, self.device,
config.max_num_batched_tokens, config.gpu_memory_utilization)
else:
num_blocks = num_blocks_for_memory(config.kv_cache_bytes, layers, kv_heads, head_dim,
config.block_size, p.dtype, getattr(flat.model, "kv_pools", 2))
self.runner = ModelRunner(flat, num_blocks, config.block_size, backend)
self.blocks = BlockManager(num_blocks, config.block_size, config.enable_prefix_caching)
self.scheduler = Scheduler(config.scheduler_config(), self.blocks)
self.sampler = sampler or (Sampler() if config.sampler == "full" else SimpleSampler())
vocab_size = flat.model.cfg.vocab_size
self.guides = GuideFactory(vocab_bytes, eos_token_id, vocab_size) if vocab_bytes is not None else None
self.eos_token_id = eos_token_id
self.max_model_len = config.max_model_len or getattr(flat.model, "context_limit", 1 << 30)
self.generators = {}
self.engine_generator = torch.Generator(device=self.device).manual_seed(config.seed)
self.steps = self.preemptions = 0
self.drafter = None # speculative decoding (Chapter 37)
self.inflight = None # the launched, unread step (async scheduling)
self.mrope_sections = getattr(flat.cfg, "mrope_sections", None)
if config.cuda_graphs != "off" and (config.cuda_graphs != "auto" or self.device.type == "cuda"):
mode = {"static": "off"}.get(config.cuda_graphs, config.cuda_graphs)
self.runner.enable_graphs(config.max_num_seqs, self.max_model_len, mode)
def add_request(self, request_id, prompt_token_ids, params=None, priority=0, cache_key=(), arrival_time=None,
lora=None, features=None):
params = params or SamplingParams()
if params.n > 1: # n samples = n requests; the prefix cache shares the prompt
return [self.add_request(f"{request_id}:{i}", prompt_token_ids, params.child(i), priority, cache_key,
arrival_time, lora, features) for i in range(params.n)]
if request_id in self.scheduler.requests:
raise ValueError(f"Duplicate request id {request_id!r}")
if (lora is not None or features) and self.drafter is not None:
raise ValueError("This drafter does not support image or adapter requests")
guided = params.guided_regex is not None or params.guided_json is not None
if guided and (self.guides is None or self.config.async_scheduling or self.config.sampler != "full"):
raise ValueError("Guided decoding needs vocab_bytes, sampler='full' and synchronous scheduling")
if (params.logprobs is not None or params.prompt_logprobs is not None or params.needs_penalties
or params.logit_bias or params.allowed_token_ids is not None or params.min_tokens) \
and self.config.sampler != "full":
raise ValueError("This request needs EngineConfig(sampler='full')")
if len(prompt_token_ids) + params.max_tokens > self.max_model_len:
raise ValueError(f"{request_id}: prompt + max_tokens exceed the model's {self.max_model_len} positions")
request = Request(request_id, list(prompt_token_ids), params, self.eos_token_id, priority,
cache_key=tuple(cache_key))
if arrival_time is not None:
request.arrival_time = arrival_time
if guided:
request.extra["guide"] = self.guides(params)
if features:
from ..multimodal import validate_features
extra, key = validate_features(features, len(prompt_token_ids), self.runner.flat.cfg.hidden_size,
self.runner.flat.cfg.head_dim)
if "mrope_sections" in extra:
if getattr(self.runner.flat.cfg, "rope_scaling", None):
raise ValueError("This M-RoPE path does not implement scaled rotary frequencies")
if self.mrope_sections is not None and self.mrope_sections != extra["mrope_sections"]:
raise ValueError("M-RoPE sections must match the model")
self.mrope_sections = extra["mrope_sections"]
request.extra.update(extra)
request.cache_key += (key,) if key else ()
bank = getattr(self.runner.flat, "adapters", None)
if lora is not None:
if bank is None:
raise ValueError("No adapter bank installed")
slot, key = bank.acquire(lora)
request.extra["lora_slot"] = slot
request.cache_key += (key,)
try:
self.scheduler.add(request)
except Exception:
if lora is not None:
bank.release(slot)
raise
if params.seed is not None:
self.generators[request_id] = torch.Generator(device=self.device).manual_seed(params.seed)
return request
def abort_request(self, request_id):
request = self.scheduler.abort(request_id)
if request is not None:
self.release_adapter(request)
self.generators.pop(request_id, None)
self.runner.swapped.pop(request_id, None)
return request
@property
def has_unfinished(self):
return self.scheduler.has_unfinished or self.inflight is not None
def check_stop(self, request, index=-1):
"""Finish reason after the token at `index` (the newest by default), or None. (Your engine: Chapter 31)"""
p, token = request.params, request.token_ids[index]
produced = (index % request.num_tokens) + 1 - request.num_prompt_tokens # output tokens up to it
if produced >= p.min_tokens:
if token in p.stop_token_ids:
request.stop_reason = token
return Status.FINISHED_STOPPED
if not p.ignore_eos and token == request.eos_token_id:
return Status.FINISHED_STOPPED
if produced >= p.max_tokens or produced + request.num_prompt_tokens >= self.max_model_len:
return Status.FINISHED_LENGTH
guide = request.extra.get("guide")
if guide is not None and guide.finished: # the pattern admits no further byte
request.stop_reason = "guided"
return Status.FINISHED_STOPPED
return None
def step(self):
"""One iteration: schedule, run, sample, update. Returns a RequestOutput per request that
produced tokens or finished. (Your engine: Chapter 31)"""
if self.config.async_scheduling:
return self.step_async()
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)
if plan.empty:
return []
if self.drafter is not None:
return self.drafter.step(self, plan) # Chapter 37 replaces sample-and-update
batch = self.runner.prepare(plan.scheduled, self.blocks.req_blocks)
logits = self.runner.execute(batch)
if batch.prompt_spans:
self.record_prompt_logprobs(batch, logits[batch.num_sample_rows:])
logits = logits[:batch.num_sample_rows]
sampling = [r for (r, _), k in zip(plan.scheduled, batch.sample_counts) if k]
tokens, logprobs = self.sampler(logits, sampling, self.generators) if sampling else ([], None)
sampled = dict(zip((r.request_id for r in sampling), tokens))
lp = dict(zip((r.request_id for r in sampling), logprobs)) if logprobs else {}
outputs, now = [], time.monotonic()
for request, n in plan.scheduled:
request.num_computed_tokens += n
new = sampled.get(request.request_id, [])
for token in new:
request.append(token)
self.blocks.cache_full_blocks(request)
if not new:
continue # mid-prompt chunk: nothing to report
request.first_token_time = request.first_token_time or now
status = self.check_stop(request)
if status is not None:
self.scheduler.finish(request, status)
self.generators.pop(request.request_id, None)
outputs.append(self.make_output(request, new, lp.get(request.request_id)))
self.steps += 1
return outputs
After the forward pass, every scheduled request advances num_computed_tokens by its n. Requests whose rows produced logits get a new token, which is appended to the list (it becomes the next step’s input) and checked against the stop conditions: a stop token, EOS unless ignore_eos, max_tokens, or the model’s context limit. min_tokens suppresses the stop checks until enough tokens exist; Chapter 34’s sampler also masks those tokens so they can’t be chosen too early. Finished requests free their blocks at once, which is why the scheduler can admit a new request in the same step that an old one finishes.
Notice what the loop doesn’t do: there’s no padding, no per-request forward, no slot bookkeeping, and no copying of caches. The sampler here is SimpleSampler, which takes greedy rows with one argmax and calls Chapter 8’s sample for each sampled row. Chapter 34 replaces it.
Sizing the pool
On a GPU, the engine gives the KV pool everything that’s left after the weights and the activations of the largest possible step:
def kv_bytes_per_block(layers, kv_heads, head_dim, block_size, dtype, pools=2):
"""K and V (pools=2; a latent cache has one), every layer, one block."""
return pools * layers * block_size * kv_heads * head_dim * torch.empty((), dtype=dtype).element_size()
def num_blocks_for_memory(free_bytes, layers, kv_heads, head_dim, block_size, dtype, pools=2):
return int(free_bytes // kv_bytes_per_block(layers, kv_heads, head_dim, block_size, dtype, pools))
@torch.inference_mode()
def profile_num_blocks(flat, backend, block_size, dtype, device, max_tokens, utilization=0.9):
"""CUDA: weights are loaded; run the largest batch the scheduler may build with a dummy
single-block pool, measure peak memory, and give everything else (up to `utilization` of
the GPU) to the KV pool. That is how an engine fills a GPU without running out mid-traffic."""
from .batch import build_batch as _build
from .request import Request
layers, kv_heads, head_dim = flat.kv_spec()
torch.cuda.empty_cache()
torch.cuda.reset_peak_memory_stats()
dummy = Request("profile", [0] * max_tokens)
probe = backend.allocate((-(-max_tokens // block_size), block_size, kv_heads, head_dim), dtype, device)
batch = _build([(dummy, max_tokens)], {"profile": list(range(probe[0].shape[0]))}, block_size, device)
per_block = sum(t[0].numel() * t.element_size() for t in probe) * layers
hidden = flat(batch.input_ids, batch.positions, [probe] * layers, batch.meta, backend)
flat.compute_logits(hidden[batch.logits_indices])
torch.cuda.synchronize()
peak = torch.cuda.max_memory_allocated()
total = torch.cuda.get_device_properties(device).total_memory
del probe, hidden
return int((total * utilization - peak) // per_block)
profile_num_blocks runs one forward pass of max_num_batched_tokens tokens with a dummy one-request pool, records the peak, and divides what remains of gpu_memory_utilization × total by the bytes per block. For Qwen3-8B in BF16 on an 80 GB H100 at 90% utilization: 72 GB available, minus 16.4 GB of weights and about 2 GB of activations for a 2,048-token step, leaves about 53 GB, or $53 \times 10^9 / 147{,}456 \approx 360{,}000$ tokens of cache. At 2,000 tokens per request that’s 180 concurrent requests’ worth of memory, against the 22 that a static 16,384-token slot per request would allow.
The runner owns the pools and does the swaps:
class ModelRunner:
def __init__(self, model, num_blocks, block_size, backend="reference", device=None, dtype=None):
self.flat = model if isinstance(model, FlatModel) else FlatModel(model)
p = next(self.flat.parameters())
self.device, self.dtype = torch.device(device or p.device), dtype or p.dtype
self.backend = get_backend(backend) if isinstance(backend, str) else backend
self.block_size, self.num_blocks = block_size, num_blocks
layers, kv_heads, head_dim = self.flat.kv_spec()
shape = (num_blocks + 1, block_size, kv_heads, head_dim) # + 1: the scratch block for padding rows
where = getattr(self.flat, "layer_device", lambda i: self.device) # Chapter 40: layers on several devices
allocate = getattr(self.flat.model, "allocate_kv", None) or self.backend.allocate # MLA: one latent pool
self.kv_caches = [allocate(shape, self.dtype, where(i)) for i in range(layers)]
self.swapped = {} # request id -> [(k, v) per layer] in host memory
self.graphs = None # GraphRunner for decode batches (Chapter 33)
self.last_hidden = None
def swap_out(self, request, blocks):
"""Copy a preempted request's blocks to (pinned, on CUDA) host memory."""
index = torch.tensor(blocks, device=self.device)
host = (lambda t: t.cpu().pin_memory()) if self.device.type == "cuda" else (lambda t: t.cpu())
self.swapped[request.request_id] = [tuple(host(_bytes(t)[index]) for t in layer) for layer in self.kv_caches]
def swap_in(self, request, blocks):
"""Copy a resumed request's saved KV into its newly allocated blocks."""
saved = self.swapped.pop(request.request_id)
count = saved[0][0].shape[0]
index = torch.tensor(blocks[:count], device=self.device)
for layer, host in zip(self.kv_caches, saved): # K, V and, when quantized, their scales
for tensor, copy in zip(layer, host):
_bytes(tensor).index_copy_(0, index, copy.to(self.device, non_blocking=True))
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
@torch.inference_mode()
def prepare_features(self, batch, scheduled):
"""Map absolute feature positions and adapter slots to this step's ragged rows. (Your engine: Chapter 43)
A prefill chunk can cut through an image; recomputation uses the same saved rows.
"""
extras = [r.extra for r, _ in scheduled]
if any("lora_slot" in e for e in extras):
slots = torch.full_like(batch.input_ids, -1)
for (request, n), start in zip(scheduled, batch.meta.query_start_loc_cpu):
slots[start:start + n] = request.extra.get("lora_slot", -1)
batch.extra["lora_ids"] = slots
mrope = [e for e in extras if "mrope_positions" in e]
if mrope:
sections = mrope[0]["mrope_sections"]
if any(e["mrope_sections"] != sections for e in mrope):
raise ValueError("All requests must use the model's M-RoPE sections")
axes = batch.positions[None].expand(3, -1).clone()
for (request, n), start in zip(scheduled, batch.meta.query_start_loc_cpu):
if "mrope_positions" not in request.extra:
continue
for offset in range(n):
pos = request.num_computed_tokens + offset
if pos < request.num_prompt_tokens:
axes[:, start + offset] = request.extra["mrope_positions"][:, pos].to(self.device)
else:
axes[:, start + offset] = pos + request.extra["mrope_delta"]
batch.extra.update(mrope_positions=axes, mrope_sections=sections)
replacements = []
for (request, n), start in zip(scheduled, batch.meta.query_start_loc_cpu):
for pos, row in request.extra.get("image_rows", {}).items():
offset = pos - request.num_computed_tokens
if 0 <= offset < n:
replacements.append((start + offset, row))
if replacements:
embeds = self.flat.backbone.embed_tokens(batch.input_ids)
for row, value in replacements:
embeds[row] = value.to(embeds)
batch.extra["embeds"] = embeds
def enable_graphs(self, max_batch, max_model_len, mode="auto"):
from .graphs import GraphRunner
blocks = -(-max_model_len // self.block_size)
self.graphs = GraphRunner(self, max_batch, blocks, use_graphs=None if mode == "auto" else mode == "on")
self.graphs.capture()
@torch.inference_mode()
def execute(self, batch):
"""Run one flattened batch; returns logits [L, vocab] for batch.logits_indices."""
if self.graphs is not None and not batch.extra and self.graphs.can_run(batch):
return self.graphs.run(batch)
hidden = self.flat(batch.input_ids, batch.positions, self.kv_caches, batch.meta, self.backend,
**(batch.extra or {}))
if hidden is None: # a pipeline stage other than the first (Chapter 41)
return None
if batch.logits_indices.numel() == 0:
return torch.empty((0, 0), device=self.device)
self.last_hidden = hidden[batch.logits_indices] # what Medusa / MTP heads read (Chapter 37)
return self.flat.compute_logits(self.last_hidden)
The correctness test
As in Chapters 16, 24 and 25, the optimizations must not change any output. The milestone runs five requests of different lengths through the engine in four configurations: small chunks and a tight budget; a 12-block pool that forces preemption by recompute; the same pool with swapping; and prefix caching off with only two sequences at a time. One request shares its first two blocks with another, so the prefix cache is exercised too. In every configuration, each request’s greedy output must equal sampling.generate run on that prompt alone, and when the engine is idle every block must be free again.
The test catches the bugs that this design invites: a slot mapping off by one at a block boundary, a request resumed with a stale num_computed_tokens, a cached block adopted without leaving a token to compute, a swapped request restored into the wrong blocks, logits taken from a mid-prompt row, and blocks leaked by aborted or preempted requests.
Run it
python run.py core --requests 16
Sixteen requests that share a 192-token system prompt, each with its own 8-48-token question and 8-40 output tokens, through the 2-layer test model on a laptop CPU. First Chapter 24’s engine with 8 slots of 320 tokens, then the engine core with an ample pool and with a 24-block (384-token) pool:
{"engine": "ch24 slots", "kv_slots": 2560, "output_tokens": 322, "tok_s": 516.5, "steps": 55}
{"engine": "ch31 core, ample pool", "kv_slots": 65536, "output_tokens": 322, "tok_s": 1108.1, "steps": 52, "prefix_hit_rate": 0.814, "preemptions": 0}
{"engine": "ch31 core, 24-block pool", "kv_slots": 384, "output_tokens": 322, "tok_s": 804.9, "steps": 96, "prefix_hit_rate": 0.866, "preemptions": 8}
The core is about twice as fast as Chapter 24’s engine, even on a CPU (timings vary by 10-20% from run to run on a laptop), for two reasons: 81% of prompt tokens come from the prefix cache, and a step’s prefill chunks run in one forward pass instead of one per request. With a pool of only 384 token slots, 15% of what the slot engine reserved, it still serves every request correctly: it preempts 8 times and takes more steps, and the prefix hit rate rises because resumed requests find their own blocks still cached.
Build it
Engine milestone 31: the engine core. Implement, in engine/serve/:
blocks.py:hash_block,BlockManager.find_cached_prefix,allocate_slots,cache_full_blocks,free;scheduler.py:Scheduler._preempt,Scheduler.schedule;batch.py:build_batch;attention.py:write_kv,ReferenceBackend.forward;model.py:FlatModel.attention,FlatModel.forward;engine.py:EngineCore.check_stop,EngineCore.step.
The request and parameter classes, swap copies, pool sizing, generate and the stats are provided.
pytest tests/test_ch31_engine_v2.py
python run.py core --impl engine
The tests check hash chains and extra keys, prefix reuse with reference counts and LRU eviction of freed blocks, the one-token rule, all-or-nothing allocation, the budget, chunk and sequence limits, running-before-waiting order, preemption of the newest request, the flattened layout of a mixed batch, paged attention against dense attention, and that every request’s output equals its solo output under chunking, recompute and swap preemption, and without prefix caching, for dense and MoE models, with stop tokens and aborts.
Stretch exercises
- ★ Add a
prefix_cache_statsendpoint-style method that reports, per step, how many tokens were cached, computed and evicted. Run a multi-turn chat workload (each turn’s prompt is the previous prompt + answer + a new message) and plot the hit rate against the pool size. Where: add a reporting method toEngineCoreinengine/serve/engine.py, readingBlockManagercounters inengine/serve/blocks.py. - ★★ Measure recompute against swap on a GPU with Qwen3-0.6B: force preemption with a small pool and long outputs, and compare total time and p99 time per output token. Then repeat with prefix caching off. When does swapping win? Where:
experiments/ch31.py(create it), configuringengine.serve.engine.EngineCoreand adaptingrun.py’scmd_core. - ★★ Implement SGLang’s cache-aware scheduling: when several requests are waiting, admit first the one with the longest cached prefix (longest-prefix-match), instead of FCFS. Show the hit-rate improvement on a workload with several different system prompts, and the starvation risk it introduces. Where: waiting-request selection in
Scheduler.scheduleinengine/serve/scheduler.py. - ★★★ Replace the flat hash map with a radix tree over token blocks (SGLang’s RadixAttention). Support partial-block matches for the last block by copying it, and evict leaves in LRU order. Compare hit rates with the hash chain on a tree-shaped workload (many branches from shared prefixes). Where: prefix lookup/storage/eviction in
BlockManagerinengine/serve/blocks.py.
Check your understanding
- Why is a prefill chunk, a decode step and a resumed request the same kind of work for this scheduler?
- In a flattened batch, which operations need the metadata and which don’t? Why does that make the matmuls faster than Chapter 24’s per-request prefill calls?
- Why can’t the prefix cache match the block containing the prompt’s last token?
- A finished request’s blocks go to the back of the free queue tail-first. What would happen to hit rates if they went head-first?
- Why must the block hash include the LoRA adapter and image hashes, and why must it be cryptographic?
- Why does the scheduler stop admitting new requests in a step where it preempted one?
- For a model with 4 KiB of K/V per token and 3 B parameters, which preemption mode would you expect to cost less, and why?
Going deeper
- vLLM V1:
vllm/v1/core/sched/scheduler.py(the unified token-budget scheduler),vllm/v1/core/kv_cache_manager.pyandvllm/v1/core/block_pool.py(hash-chained prefix caching with a free-queue LRU), and the design notes vLLM V1: A Major Upgrade to vLLM’s Core Architecture (vLLM blog, 2025) and Automatic Prefix Caching in the vLLM docs. - Zheng et al., SGLang: Efficient Execution of Structured Language Model Programs (2024), for RadixAttention and cache-aware scheduling;
python/sglang/srt/mem_cache/radix_cache.py. - Agrawal et al., Sarathi-Serve (OSDI 2024), for chunked prefill and the token budget; Kwon et al., PagedAttention (SOSP 2023), §4.5 on swapping versus recomputation.
- GPU Mode L35 (SGLang performance optimization) and L40 (FlashInfer) for how the flattened batch’s metadata reaches the kernels.
32. Production attention kernels
In this chapter
- Why the reference backend is slow, and what a fast one must do: one launch per layer for the whole mixed batch, K/V read straight from the blocks, each byte read once.
- A unified Triton kernel for flattened batches: prefill chunks, decode tokens and verification blocks together, with grouped-query heads packed into one tile.
- Split-KV ("flash-decoding") for long-context decode, and the log-sum-exp merge that makes it exact.
- INT8 and FP8 KV caches with a scale per token and head: twice the context in the same memory.
- When to use vendor kernels (FlashAttention, FlashInfer) and how to plug them in behind the same test.
You will build
unified_attention_kernel, splitkv_decode_kernel and merge_kernel in engine/kernels/triton_unified.py, and quantize_kv in engine/serve/triton_backend.py.
Time: 6-8 hours. GPU: recommended for timing (every test runs in Triton's interpreter on a CPU).
What the reference backend costs
Chapter 31’s ReferenceBackend is correct and slow, in three separate ways:
- It copies the context. For each request it gathers the request’s blocks into a contiguous tensor, then attention reads that copy. Every byte of K and V crosses the memory bus three times (read from the pool, written to the copy, read again) instead of once.
- It launches per request. A Python loop over requests issues a handful of kernels each. With 64 decoding requests that’s hundreds of launches per layer, and Chapter 19 showed what launch overhead does to decode.
- It reads K/V once per query head.
causal_attentionrepeats each KV head for its group of query heads (repeat_interleave). Qwen3-8B has 4 query heads per KV head, so the cache is read four times.
How fast could it be? Decode attention is memory-bound: each step must read every cached key and value once, and does about 4 FLOPs per element read (two for $q \cdot k$, two for $p \cdot v$), far below an H100’s ~300 FLOPs per byte. So its time is bytes divided by bandwidth. Take 64 requests at 4,096 tokens of context on Qwen3-8B. One token of one layer holds $2 \times 8 \times 128 = 2{,}048$ values, 4 KiB in BF16, so one layer’s cache for the batch is $64 \times 4{,}096 \times 4$ KiB = 1.07 GB, and all 36 layers hold 38.7 GB. At 3.35 TB/s that’s 11.5 ms per decode step for attention alone, a floor no kernel can beat, against about 4.9 ms to read the 16.4 GB of weights. At this batch size and context length attention dominates the step, and a kernel that reads the cache more than once makes it proportionally slower.
A production attention kernel therefore has three goals: read each cached byte once, in one launch per layer for the whole batch, and keep enough programs running to use every SM.
One kernel for the whole batch
The flattened batch of Chapter 31 mixes requests in different states: a 512-token prefill chunk, thirty decode tokens, a 5-token verification block. A kernel written for one of these shapes is bad at the others. Prefill kernels tile over many query rows; decode kernels have one query row and must parallelize over something else. The unified kernel handles both by choosing its tiles from the metadata:
- one program per (request, tile of
BLOCK_Qquery tokens, KV head); - each program loads the query rows of all
GROUPquery heads that share its KV head, so one tile of K and V serves all of them; - it walks the request’s block table one KV block per iteration, applying the causal rule by position and the online softmax of Chapter 15.
program (r, qt, g): rows = BLOCK_Q tokens × GROUP_PAD heads
row m -> token qt*BLOCK_Q + m // GROUP_PAD, head g*GROUP + m % GROUP_PAD
q tile [BLOCK_Q·GROUP_PAD, D] K block [BLOCK_SIZE, D] (block = table[r, j])
┌──────────┐ ┌──────┐
tok 0 │ h0 h1 h2 h3 │ · Kᵀ → scores [rows, BLOCK_SIZE] → online softmax → · V
tok 1 │ h0 h1 h2 h3 │
... └──────────┘
Packing the group into the tile’s rows is the GQA trick that FlashInfer and vLLM’s kernels use. It fixes problem 3 (K/V are loaded once per group), and it also makes decode tensor-core friendly: a single decode token has only one query row, but with 4 heads per group and BLOCK_Q = 4 the tile has 16 rows, the minimum tl.dot shape.
@triton.jit
def unified_attention_kernel(q_ptr, k_ptr, v_ptr, o_ptr, ks_ptr, vs_ptr,
qsl_ptr, seq_lens_ptr, table_ptr,
s_qt, s_qh, s_kb, s_ks, s_kh, s_sb, s_ss, s_table,
scale, window,
GROUP: tl.constexpr, GROUP_PAD: tl.constexpr, BLOCK_Q: tl.constexpr,
BLOCK_SIZE: tl.constexpr, D: tl.constexpr, QUANT: tl.constexpr):
"""(Your engine: Chapter 32)"""
r, qt, g = tl.program_id(0), tl.program_id(1), tl.program_id(2)
q_start = tl.load(qsl_ptr + r)
q_len = tl.load(qsl_ptr + r + 1) - q_start
length = tl.load(seq_lens_ptr + r)
if qt * BLOCK_Q < q_len:
# Row m of the tile is (query token m // GROUP_PAD, query head m % GROUP_PAD of group g).
rows = tl.arange(0, BLOCK_Q * GROUP_PAD)
token = qt * BLOCK_Q + rows // GROUP_PAD
head = g * GROUP + rows % GROUP_PAD
row_ok = (token < q_len) & (rows % GROUP_PAD < GROUP)
d = tl.arange(0, D)
q = tl.load(q_ptr + (q_start + token)[:, None] * s_qt + head[:, None] * s_qh + d[None, :],
mask=row_ok[:, None], other=0.0)
if QUANT:
q = q.to(tl.float32)
q_pos = length - q_len + token # queries are the newest positions
m_i = tl.full((BLOCK_Q * GROUP_PAD,), -float("inf"), tl.float32)
l_i = tl.zeros((BLOCK_Q * GROUP_PAD,), tl.float32)
acc = tl.zeros((BLOCK_Q * GROUP_PAD, D), tl.float32)
# Keys needed: up to the tile's last query; with a sliding window, from its first query's window.
last_key = tl.minimum(length, length - q_len + tl.minimum(q_len, (qt + 1) * BLOCK_Q))
first_key = tl.where(window > 0, tl.maximum(length - q_len + qt * BLOCK_Q - window + 1, 0), 0)
slots = tl.arange(0, BLOCK_SIZE)
for j in range(first_key // BLOCK_SIZE, tl.cdiv(last_key, BLOCK_SIZE)):
block = tl.load(table_ptr + r * s_table + j).to(tl.int64)
base = block * s_kb + g * s_kh
k = tl.load(k_ptr + base + slots[:, None] * s_ks + d[None, :])
v = tl.load(v_ptr + base + slots[:, None] * s_ks + d[None, :])
if QUANT: # dequantize the tile in registers
k = k.to(tl.float32) * tl.load(ks_ptr + block * s_sb + slots * s_ss + g)[:, None]
v = v.to(tl.float32) * tl.load(vs_ptr + block * s_sb + slots * s_ss + g)[:, None]
s = tl.dot(q, tl.trans(k), input_precision="ieee").to(tl.float32) * scale
key = j * BLOCK_SIZE + slots
visible = (key[None, :] <= q_pos[:, None]) & (key[None, :] < length)
visible = visible & ((window <= 0) | (key[None, :] > q_pos[:, None] - window))
s = tl.where(visible, s, -float("inf"))
m_new = tl.maximum(m_i, tl.max(s, axis=1))
m_safe = tl.where(m_new == -float("inf"), 0.0, m_new)
alpha = tl.exp(m_i - m_safe)
p = tl.exp(s - m_safe[:, None])
l_i = alpha * l_i + tl.sum(p, axis=1)
acc = acc * alpha[:, None] + tl.dot(p.to(v.dtype), v, input_precision="ieee").to(tl.float32)
m_i = m_new
out = acc / tl.where(l_i > 0, l_i, 1.0)[:, None]
tl.store(o_ptr + (q_start + token)[:, None] * s_qt + head[:, None] * s_qh + d[None, :],
out.to(o_ptr.dtype.element_ty), mask=row_ok[:, None])
Points to notice:
- Positions come from lengths, not from the input. Request $r$’s $t$ queries are the last $t$ of its
seq_lenpositions, the same rule as the reference backend, soq_pos = length - q_len + token. - The key loop stops at the tile’s last query. For a prefill chunk, a tile of early queries never loads keys after its last row’s position. For a decode token, the loop covers exactly the context.
- Programs with nothing to do exit at once. The grid’s second dimension is sized for the longest query in the batch (
max_query_len), so a decode request’s programs for tiles 1, 2, … seeqt * BLOCK_Q >= q_lenand return. That wastes a few launches’ worth of scheduling, not memory traffic. - Sliding windows (Chapter 42’s Mistral and Gemma layers) need two changes: the loop starts at the first block inside the oldest query’s window, and the mask also drops keys older than
windowpositions. - No contiguous copy. K and V tiles are loaded straight from the block with
block * stride + slot * stride + head * stride; the block table is the only indirection.
The launcher computes the grid and the packing:
def unified_attention(q, k_cache, v_cache, meta, scale=None, window=0, k_scale=None, v_scale=None, block_q=None):
"""q [N, Hq, D] -> out [N, Hq, D] for a flattened batch (BatchMeta from serve/batch.py)."""
check_device(q, k_cache, v_cache)
n, heads, d = q.shape
kv_heads, block_size = k_cache.shape[2], k_cache.shape[1]
if heads % kv_heads or d & (d - 1) or block_size & (block_size - 1):
raise ValueError("Need Hq % Hkv == 0 and power-of-two head_dim and block_size")
group = heads // kv_heads
pad = _group_pad(group)
block_q = block_q or max(1, 16 // pad) # at least 16 rows per tile: tensor-core friendly
out = torch.empty_like(q)
quant = k_scale is not None
ks = k_scale if quant else q.new_empty(1)
vs = v_scale if quant else q.new_empty(1)
grid = (meta.num_reqs, triton.cdiv(meta.max_query_len, block_q), kv_heads)
unified_attention_kernel[grid](
q, k_cache, v_cache, out, ks, vs, meta.query_start_loc, meta.seq_lens, meta.block_table,
q.stride(0), q.stride(1), k_cache.stride(0), k_cache.stride(1), k_cache.stride(2),
ks.stride(0) if quant else 0, ks.stride(1) if quant else 0, meta.block_table.stride(0),
scale or 1.0 / math.sqrt(d), window,
GROUP=group, GROUP_PAD=pad, BLOCK_Q=block_q, BLOCK_SIZE=block_size, D=d, QUANT=quant)
return out
GROUP_PAD rounds the group up to a power of two, because Triton’s tl.arange needs one. A model with 6 query heads per KV head pads to 8, wasting a quarter of each tile’s rows, a cost that vendor kernels avoid with specialized code paths. For grouped heads, block sizes and head dimensions that are powers of two, nothing is wasted.
Long contexts: split the keys
The unified kernel’s parallelism for a decode batch is requests × KV heads. Four requests with 32,768-token contexts on Qwen3-8B give $4 \times 8 = 32$ programs, on a GPU with 132 SMs. Three quarters of the GPU idles while each program walks 2,048 blocks one after another. This is the long-context, small-batch case: a single user summarizing a book.
Split-KV, published as Flash-Decoding (Dao et al., 2023), adds a third grid dimension: each request’s context is cut into $S$ splits, and each program attends over its split only. Each produces a partial output $o_s$, normalized within its split, and its log-sum-exp $\ell_s = m_s + \log \sum_{j \in s} e^{x_j - m_s}$. The exact result is a weighted average:
$$ o = \sum_s w_s, o_s, \qquad w_s = \frac{e^{\ell_s}}{\sum_{s’} e^{\ell_{s’}}} = e^{\ell_s - \ell}, \qquad \ell = \log \sum_s e^{\ell_s}. $$
That’s the online-softmax recurrence of Chapter 15 applied to whole splits instead of tiles: each partial result carries the normalizer it was computed with, so they can be combined in any order. The same merge appears in cascade attention (attend to a shared prefix once for many requests, then merge with each request’s own suffix) and in ring attention (Chapter 41), where the splits live on different GPUs.
@triton.jit
def splitkv_decode_kernel(q_ptr, k_ptr, v_ptr, ks_ptr, vs_ptr, seq_lens_ptr, table_ptr, po_ptr, pl_ptr,
s_qt, s_qh, s_kb, s_ks, s_kh, s_sb, s_ss, s_table,
s_por, s_pos, s_poh, s_plr, s_pls,
scale, blocks_per_split,
GROUP: tl.constexpr, GROUP_PAD: tl.constexpr, BLOCK_SIZE: tl.constexpr,
D: tl.constexpr, QUANT: tl.constexpr):
"""One (request, KV head, split) of a decode batch: partial output and log-sum-exp. (Your engine: Chapter 32)"""
r, g, split = tl.program_id(0), tl.program_id(1), tl.program_id(2)
length = tl.load(seq_lens_ptr + r)
rows = tl.arange(0, GROUP_PAD)
head = g * GROUP + rows
row_ok = rows < GROUP
d = tl.arange(0, D)
q = tl.load(q_ptr + r * s_qt + head[:, None] * s_qh + d[None, :], mask=row_ok[:, None], other=0.0).to(tl.float32)
m_i = tl.full((GROUP_PAD,), -float("inf"), tl.float32)
l_i = tl.zeros((GROUP_PAD,), tl.float32)
acc = tl.zeros((GROUP_PAD, D), tl.float32)
first = split * blocks_per_split
last = tl.minimum(first + blocks_per_split, tl.cdiv(length, BLOCK_SIZE))
slots = tl.arange(0, BLOCK_SIZE)
for j in range(first, last):
block = tl.load(table_ptr + r * s_table + j).to(tl.int64)
base = block * s_kb + g * s_kh
k = tl.load(k_ptr + base + slots[:, None] * s_ks + d[None, :]).to(tl.float32)
v = tl.load(v_ptr + base + slots[:, None] * s_ks + d[None, :]).to(tl.float32)
if QUANT:
k = k * tl.load(ks_ptr + block * s_sb + slots * s_ss + g)[:, None]
v = v * tl.load(vs_ptr + block * s_sb + slots * s_ss + g)[:, None]
s = tl.dot(q, tl.trans(k), input_precision="ieee") * scale # [GROUP_PAD, BLOCK_SIZE]
s = tl.where((j * BLOCK_SIZE + slots)[None, :] < length, s, -float("inf"))
m_new = tl.maximum(m_i, tl.max(s, axis=1))
m_safe = tl.where(m_new == -float("inf"), 0.0, m_new)
alpha = tl.exp(m_i - m_safe)
p = tl.exp(s - m_safe[:, None])
l_i = alpha * l_i + tl.sum(p, axis=1)
acc = acc * alpha[:, None] + tl.dot(p, v, input_precision="ieee")
m_i = m_new
out = acc / tl.where(l_i > 0, l_i, 1.0)[:, None]
lse = tl.where(l_i > 0, m_i + tl.log(tl.where(l_i > 0, l_i, 1.0)), -float("inf")) # empty split: -inf
tl.store(po_ptr + r * s_por + split * s_pos + head[:, None] * s_poh + d[None, :], out, mask=row_ok[:, None])
tl.store(pl_ptr + r * s_plr + split * s_pls + head, lse, mask=row_ok)
@triton.jit
def merge_kernel(po_ptr, pl_ptr, o_ptr, splits, s_por, s_pos, s_poh, s_plr, s_pls, s_ot, s_oh,
SPLITS_PAD: tl.constexpr, D: tl.constexpr):
"""out = sum_s exp(lse_s - lse) * out_s, with lse = logsumexp_s(lse_s). (Your engine: Chapter 32)"""
r, h = tl.program_id(0), tl.program_id(1)
s = tl.arange(0, SPLITS_PAD)
d = tl.arange(0, D)
lse = tl.load(pl_ptr + r * s_plr + s * s_pls + h, mask=s < splits, other=-float("inf"))
top = tl.max(lse, axis=0)
top = tl.where(top == -float("inf"), 0.0, top)
w = tl.exp(lse - top) # 0 for empty splits
parts = tl.load(po_ptr + r * s_por + s[:, None] * s_pos + h * s_poh + d[None, :], mask=(s < splits)[:, None], other=0.0)
total = tl.sum(w, axis=0)
out = tl.sum(w[:, None] * parts, axis=0) / tl.where(total > 0, total, 1.0)
tl.store(o_ptr + r * s_ot + h * s_oh + d, out.to(o_ptr.dtype.element_ty))
An empty split (a short request in a batch with long ones, or a padding row) writes $\ell_s = -\infty$, whose weight is $e^{-\infty} = 0$. The merge program guards the all-empty case, so padding rows produce zeros rather than NaNs, which matters when a CUDA graph (Chapter 33) runs padded batches.
How many splits? Enough programs to fill the GPU, but not so many that each split is a few blocks and the merge costs more than it saves:
def choose_splits(num_reqs, kv_heads, max_blocks, target_programs=264, min_blocks_per_split=4, max_splits=32):
"""Enough programs to fill the GPU (about two per SM on a 132-SM H100), without splits so
small that the merge costs more than it saves."""
want = triton.cdiv(target_programs, max(num_reqs * kv_heads, 1))
return max(1, min(want, max_splits, triton.cdiv(max_blocks, min_blocks_per_split)))
def split_kv_decode(q, k_cache, v_cache, meta, num_splits=None, scale=None, k_scale=None, v_scale=None):
"""Decode-only batch (one query row per request): q [R, Hq, D] -> out [R, Hq, D]."""
check_device(q, k_cache, v_cache)
reqs, heads, d = q.shape
kv_heads, block_size = k_cache.shape[2], k_cache.shape[1]
group = heads // kv_heads
max_blocks = triton.cdiv(max(meta.max_seq_len, 1), block_size)
splits = num_splits or choose_splits(reqs, kv_heads, max_blocks)
per_split = triton.cdiv(max_blocks, splits)
parts = torch.empty((reqs, splits, heads, d), device=q.device, dtype=torch.float32)
lse = torch.empty((reqs, splits, heads), device=q.device, dtype=torch.float32)
quant = k_scale is not None
ks = k_scale if quant else parts
vs = v_scale if quant else parts
splitkv_decode_kernel[(reqs, kv_heads, splits)](
q, k_cache, v_cache, ks, vs, meta.seq_lens, meta.block_table, parts, lse,
q.stride(0), q.stride(1), k_cache.stride(0), k_cache.stride(1), k_cache.stride(2),
ks.stride(0) if quant else 0, ks.stride(1) if quant else 0, meta.block_table.stride(0),
parts.stride(0), parts.stride(1), parts.stride(2), lse.stride(0), lse.stride(1),
scale or 1.0 / math.sqrt(d), per_split,
GROUP=group, GROUP_PAD=max(_group_pad(group), 16), BLOCK_SIZE=block_size, D=d, QUANT=quant)
out = torch.empty_like(q)
merge_kernel[(reqs, heads)](parts, lse, out, splits, parts.stride(0), parts.stride(1), parts.stride(2),
lse.stride(0), lse.stride(1), out.stride(0), out.stride(1),
SPLITS_PAD=max(_group_pad(splits), 2), D=d)
return out
The backend uses split-KV only for decode-only batches with few (request, KV head) pairs and contexts of at least 2,048 tokens. For a batch of 64 decoding requests, the unified kernel already has $64 \times 8 = 512$ programs, and splitting would only add the merge.
Quantized KV caches
At long contexts the KV cache dominates both memory and decode time (Chapter 16), so storing it in 8 bits instead of 16 doubles the context that fits and nearly halves attention’s memory traffic. Two 8-bit formats are common:
| format | values | strengths |
|---|---|---|
| INT8 with a scale | 255 evenly spaced levels in $[-127s, 127s]$ | uniform precision; works on every GPU |
| FP8 E4M3 with a scale | 4 exponent bits, 3 mantissa bits, max 448 | relative precision across a wide range; native in Hopper and later tensor cores |
The scale’s granularity matters more than the format. One scale per layer (vLLM’s FP8 KV cache, which uses 1.0 unless the checkpoint ships calibrated scales) is cheapest, but a single token or channel with a large value then wastes most of the code range for every other entry. This chapter uses one scale per (token, KV head): 4 extra bytes per 128-value head vector, a 3% overhead, which follows outliers that vary from token to token.
def quantize_kv(x, dtype):
"""x [N, Hkv, D] -> (codes [N, Hkv, D], scales [N, Hkv]) with one scale per token and head. (Your engine: Chapter 32)
Symmetric: scale = max|x| / q_max, so the largest element maps to the code range's edge
(127 for int8, 448 for fp8 E4M3). Per-token scales follow outliers that differ token
to token; per-head scales keep one head's large channels from crushing another head.
"""
code_dtype, q_max = KV_DTYPES[dtype]
amax = x.float().abs().amax(-1).clamp_min(1e-8)
scale = amax / q_max
scaled = x.float() / scale[..., None]
codes = scaled.round().clamp(-127, 127).to(code_dtype) if dtype == "int8" else scaled.to(code_dtype)
return codes, scale
Quantization happens on write, once per token. Dequantization happens in the kernel’s registers: the tile of codes is loaded (half the bytes), converted and multiplied by its per-slot scale, then used exactly like a BF16 tile. Both kernels take a QUANT constant that switches this on, so the same code serves both pools.
class TritonBackend(ReferenceBackend):
name = "triton"
def __init__(self, kv_cache_dtype="auto", split_kv="auto"):
if kv_cache_dtype not in ("auto", *KV_DTYPES):
raise ValueError(f"kv_cache_dtype must be auto, int8 or fp8, not {kv_cache_dtype!r}")
self.kv_cache_dtype, self.split_kv = kv_cache_dtype, split_kv
@property
def quantized(self):
return self.kv_cache_dtype != "auto"
def allocate(self, shape, dtype, device):
if not self.quantized:
return super().allocate(shape, dtype, device)
code_dtype = KV_DTYPES[self.kv_cache_dtype][0]
k, v = (torch.zeros(shape, dtype=code_dtype, device=device) for _ in range(2))
k_scale, v_scale = (torch.zeros(shape[:3], dtype=torch.float32, device=device) for _ in range(2))
return k, v, k_scale, v_scale
def write(self, cache, k, v, meta):
if not self.quantized:
return write_kv(cache[0], cache[1], k, v, meta.slot_mapping)
(kc, ks), (vc, vs) = quantize_kv(k, self.kv_cache_dtype), quantize_kv(v, self.kv_cache_dtype)
bits = (lambda t: t.view(torch.uint8)) if self.kv_cache_dtype == "fp8" else (lambda t: t)
write_kv(bits(cache[0]), bits(cache[1]), bits(kc), bits(vc), meta.slot_mapping) # same bytes, any dtype
write_kv(cache[2][..., None], cache[3][..., None], ks[..., None], vs[..., None], meta.slot_mapping)
def use_split_kv(self, meta, kv_heads):
"""Split long decode contexts when one program per (request, KV head) can't fill the GPU."""
if self.split_kv != "auto":
return bool(self.split_kv) and meta.decode_only
return meta.decode_only and meta.num_reqs * kv_heads < 128 and meta.max_seq_len >= 2048
def forward(self, q, cache, meta, scale=None, window=0):
from ..kernels.triton_unified import split_kv_decode, unified_attention
scales = (cache[2], cache[3]) if self.quantized else (None, None)
if window == 0 and self.use_split_kv(meta, cache[0].shape[2]):
return split_kv_decode(q, cache[0], cache[1], meta, None, scale, *scales)
return unified_attention(q, cache[0], cache[1], meta, scale, window, *scales)
FP8 has a catch on the CPU: PyTorch’s CPU kernels don’t implement index_copy_ for float8_e4m3fn. The write therefore copies the same bytes through a uint8 view. A dtype is just an interpretation of bytes; for a copy, the interpretation doesn’t matter.
How much accuracy does it cost? Keys are more sensitive than values, because errors in $q \cdot k$ pass through the exponential of the softmax, and real models’ keys have a few large outlier channels. KIVI (Liu et al., 2024) quantizes keys per channel and values per token for this reason; KVQuant (Hooper et al., 2024) goes further with non-uniform codes. Per-token-and-head scales at 8 bits are a robust default; measure with the model’s own perplexity before going lower (stretch exercise 3).
Vendor kernels
Should a production engine write its own attention kernels at all? Every major engine uses vendor kernels on its main path: vLLM uses FlashAttention 2/3 and FlashInfer on NVIDIA GPUs, SGLang uses FlashInfer and its own Triton kernels, TensorRT-LLM uses NVIDIA’s hand-written kernels. On Hopper, FlashAttention-3 (Shah et al., 2024) reaches 75% of peak by using features that Triton exposes only partly: asynchronous TMA loads, warp-specialized producer and consumer warps, wgmma instructions, and FP8 compute with incoherent processing. For prefill on an H100, expect FA3 to be well ahead of a straightforward Triton kernel; for decode, which is memory-bound, a good Triton kernel gets much closer, because the limit is bandwidth rather than instruction scheduling.
Your kernels remain useful in three ways: they run where vendor kernels don’t (AMD GPUs through Triton, Chapter 39; new architectures before the vendor library supports them; the CPU interpreter for testing), they’re the reference you understand completely, and they’re the starting point for operations no library has yet (the stretch exercises). The engine’s backend interface makes vendor kernels a drop-in:
class FlashAttnBackend(ReferenceBackend):
"""Delegates to flash_attn_varlen_func(..., block_table=...). FlashAttention-2 requires the
page (block) size to be a multiple of 256; vLLM's fork of it accepts 16."""
name = "flash_attn"
def __init__(self):
from flash_attn import flash_attn_varlen_func # noqa: F401 (fail early if missing)
self.fn = flash_attn_varlen_func
def forward(self, q, cache, meta, scale=None, window=0):
k_cache, v_cache = cache[0], cache[1]
if k_cache.shape[1] % 256:
raise ValueError("flash-attn 2 paged attention needs block_size % 256 == 0")
cu_k = torch.zeros(meta.num_reqs + 1, dtype=torch.int32, device=q.device)
cu_k[1:] = torch.cumsum(meta.seq_lens, 0)
return self.fn(q, k_cache, v_cache, meta.query_start_loc, cu_k, meta.max_query_len, meta.max_seq_len,
softmax_scale=scale or 1.0 / math.sqrt(q.shape[-1]), causal=True,
window_size=(window - 1, 0) if window else (-1, -1), block_table=meta.block_table)
The adapter passes the same metadata the Triton kernel uses. FlashAttention-2’s paged interface requires a page size divisible by 256, so you’d run the engine with block_size=256, which lowers prefix-cache granularity (vLLM ships a fork of FlashAttention that accepts 16). FlashInfer’s interface is different again: a plan call on the CPU partitions the batch’s work once per step, then run launches the kernel for each layer, which moves the load-balancing decisions out of the kernel. Whatever you plug in, it must pass the same test as your own kernel: equal to the reference backend on a mixed batch.
The tests and the measurements
The milestone tests compare the unified kernel with the reference backend on a batch with requests in every state (mid-prefill chunk, a full 16-token prompt, a 21-token prompt, decodes at 70 and 130 tokens), for grouped heads with group sizes 4 and 3 (padded to 4) and for plain multi-head attention; the sliding window; split-KV with 1, 3 and 8 splits; padding rows producing exact zeros; INT8 and FP8 caches matching the reference on the dequantized values and staying close to full precision; and the engine core producing every request’s solo output with the Triton backend.
python run.py backends
On a CPU, with Qwen3-8B’s attention shape (32 query heads, 8 KV heads, head dim 128) and contexts shrunk 8× so that the interpreter finishes in a few minutes, it prints the errors against the reference:
{"case": "mixed (2 prefill chunks + 30 decodes)", "backend": "triton unified", "max_abs_error": 7.152557373046875e-07, "ms": null}
{"case": "mixed (2 prefill chunks + 30 decodes)", "backend": "triton, int8 KV", "max_abs_error": 0.009663641452789307, "ms": null}
{"case": "mixed (2 prefill chunks + 30 decodes)", "backend": "triton, fp8 KV", "max_abs_error": 0.03983457386493683, "ms": null}
{"case": "decode, 4 long contexts", "backend": "triton unified", "max_abs_error": 1.564621925354004e-07, "ms": null}
{"case": "decode, 4 long contexts", "backend": "triton split-kv", "max_abs_error": 1.4156103134155273e-07, "ms": null}
{"case": "decode, 4 long contexts", "backend": "triton, int8 KV", "max_abs_error": 0.0009696260094642639, "ms": null}
{"case": "decode, 4 long contexts", "backend": "triton, fp8 KV", "max_abs_error": 0.005297387018799782, "ms": null}
{"kv_cache": "bf16", "qwen3_8b_bytes_per_token": 147456, "tokens_per_GB": 6781}
{"kv_cache": "int8", "qwen3_8b_bytes_per_token": 76032, "tokens_per_GB": 13152}
The exact kernels agree with the reference to FP32 rounding. INT8 is four to five times more accurate than FP8 at the same size, because E4M3 has only 3 mantissa bits, and both errors shrink with longer contexts as the softmax averages over more values. On a GPU, the ms column times each backend; compare each against the bandwidth floor you computed at the start of the chapter (bytes of K and V read, divided by your GPU’s measured bandwidth from Chapter 10).
Build it
Engine milestone 32: production attention. Implement unified_attention_kernel, splitkv_decode_kernel and merge_kernel in engine/kernels/triton_unified.py, and quantize_kv in engine/serve/triton_backend.py (the launchers, the split heuristic, the backend class and the FlashAttention adapter are provided).
pytest tests/test_ch32_attention_kernels.py
python run.py backends --impl engine
Then serve with it: EngineConfig(attention_backend="triton", block_size=16), and kv_cache_dtype="int8" or "fp8" for a quantized cache.
Stretch exercises
- ★ Autotune the unified kernel with
@triton.autotuneoverBLOCK_Q,num_warpsandnum_stages, keyed onmax_query_lenand the group size. Plot achieved bandwidth for decode batches of 1-256 requests at 4,096 tokens against the floor. Where:unified_attention_kerneland its launcher inengine/kernels/triton_unified.py. - ★★ Implement cascade attention: for a batch whose requests share a long cached prefix, attend all of their queries to the shared blocks in one tile-friendly pass (the prefix’s K/V read once for the whole batch), attend each to its own suffix, and merge with
merge_kernel. Measure the saving with 64 requests sharing a 4,000-token system prompt. Where: new cascade kernels inengine/kernels/triton_unified.py, dispatched byTritonBackend.forwardinengine/serve/triton_backend.py. - ★★ Measure KV quantization’s effect on quality: perplexity of Qwen3-0.6B on a held-out text with BF16, INT8 and FP8 caches, then per-channel keys (KIVI) at 4 bits. Where:
experiments/ch32.py(create it) for BF16/INT8/FP8; extend cache layout/read/write inengine/serve/triton_backend.pyandengine/kernels/triton_unified.pyfor KIVI. - ★★★ Add attention sinks (gpt-oss): one learned logit per head that joins every softmax’s denominator but contributes no value. It’s two lines in the online softmax: start
m_iandl_ifrom the sink instead of $-\infty$ and 0. Test against a dense implementation. Where: online-softmax initialization inengine/kernels/triton_unified.py; add the dense comparison inengine/serve/attention.py.
Check your understanding
- Why is decode attention memory-bound, and what’s the minimum time for a step, given the context lengths and the bandwidth?
- What does packing the GQA group into a tile’s rows save, and why does it also help decode use tensor cores?
- Why does a decode-only batch with 4 long requests need split-KV, while one with 64 requests doesn’t?
- Derive the merge weights $w_s$ from the definition of softmax. Why can the splits be merged in any order?
- Why must an empty split write $-\infty$ as its log-sum-exp, and what would a padding row produce if the merge didn’t guard the all-empty case?
- Why are per-token-and-head scales more accurate than a single scale for the whole KV cache?
Going deeper
- Dao, Haziza, Massa, Sizov, Flash-Decoding for long-context inference (Stanford CRFM, 2023); Shah et al., FlashAttention-3: Fast and Accurate Attention with Asynchrony and Low-precision (2024).
- Ye et al., FlashInfer: Efficient and Customizable Attention Engine for LLM Inference Serving (MLSys 2025), for load-balanced scheduling of ragged batches and cascade attention; GPU Mode L40 (FlashInfer).
- vLLM’s Triton attention backend (
vllm/attention/ops/andvllm/v1/attention/backends/), which follows the unified design of this chapter and is one of the backends vLLM runs on AMD GPUs. - Liu et al., KIVI: A Tuning-Free Asymmetric 2bit Quantization for KV Cache (2024); Hooper et al., KVQuant (2024).
- PMPP §20.5 (FlashAttention, pp. 492-503) and GPU Mode L12 (Flash Attention), for the tiling ideas this kernel builds on.
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 overhead | where it comes from | size, Qwen3-0.6B, batch 16, H100 |
|---|---|---|
| GPU work | weights + KV cache read once per step | 1.2 GB / 3.35 TB/s ≈ 0.36 ms |
| kernel launches | 28 layers × ~16 kernels, ~5 µs of CPU each | ≈ 2.2 ms |
| host bookkeeping | scheduling, building the batch, Python | 0.3 ms measured below, growing with batch size |
| host syncs | .tolist() of sampled tokens; MoE’s expert counts | GPU 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, whichwrite_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:
-
Launch step $s$; sample on the device (no
.tolist()); for each request that samples, append aPLACEHOLDERand advancenum_computed_tokensnow. -
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 -
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_hashesstops 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
- ★ Profile one decode step on a GPU with
torch.profilerbefore and afterfuse()and graphs. Count kernels per step and plot the GPU’s idle gaps. Where:experiments/ch33.py(create it), adaptingrun.py’scmd_overhead. - ★★ 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 againstbuild_batch. Where: persistent metadata inengine/serve/batch.py, updated fromModelRunner.prepareinengine/serve/runner.py. - ★★★ Piecewise graphs: split
FlatModel.forwardat 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._forwardinengine/serve/model.pyand capture/replay inengine/serve/graphs.py. - ★★ 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 inengine/kernels/triton_unified.py; integrate it inengine/serve/triton_backend.pyandengine/serve/model.py.
Check your understanding
- Why is Qwen3-0.6B’s decode launch-bound at batch 16 on an H100, while Qwen3-32B’s isn’t?
- What three things does fusing Q/K/V into one matmul save?
- Why can’t
forward_groupedrun inside a CUDA graph, and which two changes inmoe_alignfix that? - 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?
- Why must the graph’s
BatchMetacarry the bucket’snum_reqsrather than the step’s? - 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?
- 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.pyand 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_moeTriton 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).
34. Sampling at scale and structured output
In this chapter
- Everything an inference API lets a request ask for, and how one batched sampler serves a batch where every request asked for something different.
- Penalties, logit bias,
min_tokens, logprobs and prompt logprobs, and seeds that reproduce whatever else is in the batch. - Parallel samples (
n> 1) and beam search, built on the prefix cache instead of special cache plumbing. - Constrained decoding: compiling regular expressions to a byte-level DFA, turning the DFA into per-token masks for a 150,000-token vocabulary, and JSON Schema to regex.
You will build
apply_processors, apply_penalties, top_k_top_p_min_p and Sampler.sample_device (engine/serve/sampler.py); beam_search (engine/serve/beam.py); utf8_sequences, TokenIndex.next_states, Guide.advance and json_schema_to_regex (engine/serve/structured.py).
Time: 7-9 hours. GPU: not needed.
What a request can ask for
Chapter 8’s sample chose one token from one row of logits with four settings. An API request can set many more, and in a serving engine every row of the batch belongs to a different request with different settings. SamplingParams collects them, with the names the OpenAI and vLLM APIs use:
@dataclass
class SamplingParams:
max_tokens: int = 16
temperature: float = 0.0 # 0 = greedy
top_p: float = 1.0
top_k: int = 0 # 0 = off
min_p: float = 0.0
seed: int | None = None # None: draw from the engine's generator
n: int = 1 # parallel samples of one prompt (Chapter 34)
stop_token_ids: tuple = ()
stop: tuple = () # stop strings, checked on detokenized text (Chapter 35)
include_stop_str_in_output: bool = False
ignore_eos: bool = False
min_tokens: int = 0 # EOS and stop tokens are masked until this many tokens exist
presence_penalty: float = 0.0 # OpenAI: subtract once per distinct generated token
frequency_penalty: float = 0.0 # OpenAI: subtract once per occurrence
repetition_penalty: float = 1.0 # CTRL/HF: divide positive / multiply negative logits of seen tokens
logit_bias: dict = field(default_factory=dict) # {token_id: bias}
allowed_token_ids: tuple | None = None
logprobs: int | None = None # return the chosen token's logprob and the top-n alternatives
prompt_logprobs: int | None = None
guided_regex: str | None = None # constrained decoding (Chapter 34)
guided_json: dict | str | None = None
skip_special_tokens: bool = True
def __post_init__(self):
if self.max_tokens < 1:
raise ValueError("max_tokens must be at least 1")
if self.temperature < 0 or not 0 < self.top_p <= 1 or self.top_k < 0 or not 0 <= self.min_p <= 1:
raise ValueError("Need temperature >= 0, 0 < top_p <= 1, top_k >= 0 and 0 <= min_p <= 1")
if self.n < 1 or self.repetition_penalty <= 0 or self.min_tokens < 0:
raise ValueError("Need n >= 1, repetition_penalty > 0 and min_tokens >= 0")
if self.guided_regex is not None and self.guided_json is not None:
raise ValueError("Choose one of guided_regex and guided_json")
self.stop_token_ids = tuple(self.stop_token_ids)
self.stop = (self.stop,) if isinstance(self.stop, str) else tuple(self.stop)
def child(self, i):
"""Settings for sample i of an n > 1 request: one sample each, distinct seeds."""
from dataclasses import replace
return replace(self, n=1, seed=None if self.seed is None else self.seed + i)
@property
def greedy(self):
return self.temperature == 0
@property
def needs_penalties(self):
return bool(self.presence_penalty or self.frequency_penalty or self.repetition_penalty != 1.0)
| setting | what it does | where it acts |
|---|---|---|
temperature, top_k, top_p, min_p | shape and truncate the distribution (Chapter 8) | filters |
seed | this request’s draws are reproducible, whatever else is in the batch | the draw |
presence_penalty, frequency_penalty | subtract a constant per distinct / per repeated output token (OpenAI) | penalties |
repetition_penalty | shrink the logits of every token seen in prompt or output (CTRL, Hugging Face) | penalties |
logit_bias, allowed_token_ids | add to chosen logits; forbid everything else | processors |
min_tokens | EOS and stop tokens are impossible until this many tokens exist | processors |
guided_regex, guided_json | the output must match a pattern or a JSON Schema | processors |
logprobs, prompt_logprobs | report log-probabilities of chosen, alternative and prompt tokens | reported, not applied |
n | several independent samples of one prompt | the engine |
stop, include_stop_str_in_output, skip_special_tokens | stop on generated text; how to detokenize | the detokenizer (Chapter 35) |
Chapter 31’s SimpleSampler handled the first two rows of this table, with a Python loop over sampled rows. This chapter’s Sampler handles all of it with batched tensor operations, so its cost barely depends on how many requests asked for what.
One sampler for a batch
The pipeline, in the order vLLM uses:
logits [L, V] --> processors: allowed ids, logit bias, min_tokens, grammar masks
--> penalties: repetition, presence, frequency
--> greedy rows: argmax (done)
--> sampled rows: / temperature, top-k, top-p, min-p, draw
Order matters in two places. Processors come first because they express hard constraints: a grammar mask must not be undone by a penalty, and a forbidden token must stay forbidden whatever the temperature. Logprobs are taken from the raw logits, before anything else touches them. That’s the OpenAI convention and vLLM’s default: a reported logprob says what the model thought, so it’s comparable across requests with different settings, and evaluations that score answers by logprob aren’t disturbed by temperature or penalties.
Processors
def apply_processors(logits, requests):
"""Per-request logit edits that must happen before anything else. (Your engine: Chapter 34)
allowed_token_ids: everything else -> -inf. logit_bias: add bias[token]. min_tokens: EOS
and stop tokens -> -inf until the request has that many output tokens. Grammar guides
(structured.py): tokens that would break the pattern -> -inf.
"""
vocab = logits.shape[-1]
for i, request in enumerate(requests):
p = request.params
if p.allowed_token_ids is not None:
keep = torch.zeros(vocab, dtype=torch.bool, device=logits.device)
keep[list(p.allowed_token_ids)] = True
logits[i].masked_fill_(~keep, float("-inf"))
if p.logit_bias:
ids = torch.tensor([int(t) for t in p.logit_bias], device=logits.device)
logits[i, ids] += torch.tensor([float(b) for b in p.logit_bias.values()], device=logits.device)
if request.num_output_tokens < p.min_tokens:
banned = list(p.stop_token_ids) + ([request.eos_token_id] if request.eos_token_id is not None else [])
if banned:
logits[i, banned] = float("-inf")
guide = request.extra.get("guide")
if guide is not None:
logits[i].masked_fill_(~guide.allowed(logits.device), float("-inf"))
return logits
These are per-request edits with Python loops over the requests that set them. They’re cheap because most requests set none, and each edit is a single indexed write. min_tokens is worth a second look: setting EOS and every stop token to $-\infty$ makes stopping impossible rather than unlikely, so a model that wants to stop early must pick its next-best token. The engine’s check_stop also refuses to stop before min_tokens, which matters for stop strings (Chapter 35), which no logit mask can prevent.
Penalties
def apply_penalties(logits, requests):
"""Repetition, presence and frequency penalties, batched. (Your engine: Chapter 34)
counts[r, t] = occurrences of token t in request r's OUTPUT; seen[r, t] = t occurs in its
prompt or output. Then, as in vLLM:
repetition (HF/CTRL): seen positive logits / p, seen negative logits * p
frequency (OpenAI): logits -= frequency_penalty * counts
presence (OpenAI): logits -= presence_penalty * (counts > 0)
"""
rows = [i for i, r in enumerate(requests) if r.params.needs_penalties]
if not rows:
return logits
device, vocab = logits.device, logits.shape[-1]
sub = [requests[i] for i in rows]
width = max(max(r.num_output_tokens for r in sub), 1)
outputs = torch.full((len(sub), width), vocab, dtype=torch.long, device=device) # vocab = padding
for j, r in enumerate(sub):
if r.num_output_tokens:
outputs[j, :r.num_output_tokens] = torch.tensor(r.output_token_ids, device=device)
counts = torch.zeros((len(sub), vocab + 1), device=device).scatter_add_(
1, outputs, torch.ones_like(outputs, dtype=torch.float))[:, :vocab]
seen = counts > 0
for j, r in enumerate(sub):
seen[j, torch.tensor(r.prompt_token_ids, device=device)] = True
rep = torch.tensor([r.params.repetition_penalty for r in sub], device=device)[:, None]
freq = torch.tensor([r.params.frequency_penalty for r in sub], device=device)[:, None]
pres = torch.tensor([r.params.presence_penalty for r in sub], device=device)[:, None]
x = logits[rows]
x = torch.where(seen, torch.where(x > 0, x / rep, x * rep), x)
x = x - freq * counts - pres * (counts > 0).float()
logits[rows] = x
return logits
The three penalties do different things and are easy to confuse:
- Presence ($\alpha$) and frequency ($\beta$) come from the OpenAI API: $\ell_t \leftarrow \ell_t - \alpha \cdot [c_t > 0] - \beta \cdot c_t$, where $c_t$ counts token $t$ in the output so far. They’re additive and unbounded: a frequency penalty of 0.5 makes a token that already appeared 20 times 10 logits less likely.
- Repetition ($\rho$), from CTRL (Keskar et al., 2019) and Hugging Face, is multiplicative and applies to every token seen in the prompt or output: positive logits are divided by $\rho$, negative ones multiplied. It’s bounded, but it penalizes copying from the prompt, which hurts tasks that should quote it, like extraction and summarization.
The counts come from one scatter_add_ over the batch’s output tokens padded into a rectangle (the padding value is vocab, a column that’s dropped afterwards). Rebuilding them from Python lists every step costs host time proportional to the outputs’ total length; vLLM keeps persistent count tensors and updates them with each step’s tokens (stretch exercise 1).
Filters and the draw
def top_k_top_p_min_p(logits, top_k, top_p, min_p):
"""Per-row filters with one descending sort. (Your engine: Chapter 34)
top_k [L] (vocab = off), top_p [L] (1 = off), min_p [L] (0 = off), applied in Chapter 8's
order: keep the k best; renormalize; keep a token while the mass BEFORE it is < p (so the
token that crosses p stays); drop tokens below min_p times the top probability (a ratio,
so it doesn't care about renormalization). The top token always survives.
"""
ordered, order = logits.sort(dim=-1, descending=True)
rank = torch.arange(logits.shape[-1], device=logits.device)[None]
ordered = ordered.masked_fill(rank >= top_k[:, None], float("-inf"))
probs = ordered.softmax(-1)
keep = (probs.cumsum(-1) - probs) < top_p[:, None]
keep &= probs >= min_p[:, None] * probs[:, :1]
keep[:, 0] = True
filtered = ordered.masked_fill(~keep, float("-inf"))
return torch.full_like(logits, float("-inf")).scatter(-1, order, filtered)
One descending sort serves all three filters for every row. Each row’s top_k becomes a rank mask, its top_p a cumulative-mass mask on the renormalized survivors, and its min_p a ratio to the top probability. A row that asked for none of them has top_k = vocab, top_p = 1, min_p = 0, and keeps everything. The sort is skipped entirely when no row needs it.
class Sampler:
def __init__(self, max_logprobs=20):
self.max_logprobs = max_logprobs
def sample_device(self, logits, requests, generators):
"""[L, V] logits -> [L] token ids on the device, no host sync. (Your engine: Chapter 34)"""
logits = apply_penalties(apply_processors(logits.float().clone(), requests), requests)
device, vocab = logits.device, logits.shape[-1]
greedy = torch.tensor([r.params.greedy for r in requests], device=device)
temperature = torch.tensor([r.params.temperature or 1.0 for r in requests], device=device)
x = logits / temperature[:, None]
if any(r.params.top_k or r.params.top_p < 1 or r.params.min_p for r in requests):
top_k = torch.tensor([r.params.top_k or vocab for r in requests], device=device)
top_p = torch.tensor([r.params.top_p for r in requests], device=device)
min_p = torch.tensor([r.params.min_p for r in requests], device=device)
x = top_k_top_p_min_p(x, top_k, top_p, min_p)
probs = x.softmax(-1)
# Exponential race: argmax(p_i / E_i) with E_i ~ Exp(1) draws i with probability p_i.
noise = torch.empty_like(probs).exponential_()
for i, r in enumerate(requests):
if r.request_id in generators: # seeded: the same draws whatever the batch
noise[i] = torch.empty(vocab, device=device).exponential_(generator=generators[r.request_id])
sampled = (probs / noise).argmax(-1)
self.last_logits = logits
return torch.where(greedy, logits.argmax(-1), sampled)
def __call__(self, logits, requests, generators):
tokens = self.sample_device(logits, requests, generators)
token_list = tokens.tolist()
logprobs = None
wanted = [r.params.logprobs for r in requests]
if any(n is not None for n in wanted):
logprobs = self.logprobs(logits, tokens, wanted)
for request, token in zip(requests, token_list):
guide = request.extra.get("guide")
if guide is not None:
guide.advance(token)
return [[t] for t in token_list], logprobs
def logprobs(self, raw_logits, tokens, wanted):
"""Chosen token's logprob and rank, and the top-n alternatives, from the raw logits."""
lp = raw_logits.float().log_softmax(-1)
chosen = lp.gather(-1, tokens[:, None])[:, 0]
rank = (lp > chosen[:, None]).sum(-1) + 1
n = min(max(w or 0 for w in wanted), self.max_logprobs)
top_v, top_i = lp.topk(max(n, 1), dim=-1)
out = []
for i, w in enumerate(wanted):
if w is None:
out.append(None)
continue
out.append([{"token_id": int(tokens[i]), "logprob": float(chosen[i]), "rank": int(rank[i]),
"top": dict(zip(top_i[i, :w].tolist(), top_v[i, :w].tolist()))}])
return out
The draw uses the exponential race: with $E_i \sim \text{Exp}(1)$ independent, $\arg\max_i p_i / E_i$ is token $i$ with probability exactly $p_i$. It’s the Gumbel-max trick of Chapter 19 in another form ($-\log E_i$ is a Gumbel variable), so it needs no sort and no cumulative sum, only elementwise operations and an argmax, and it never syncs with the host.
Seeded requests draw their noise from their own torch.Generator. That’s what makes seed meaningful in a server: the row’s random numbers depend only on its generator, not on its position in the batch or on how many other rows drew before it. The test checks exactly that: the same seeded request draws the same tokens alone and in a batch.
Logprobs, and why prompt logprobs disable the prefix cache
logprobs=n returns, for each generated token, its log-probability, its rank, and the $n$ most likely alternatives. prompt_logprobs=n returns the log-probability of each prompt token given the ones before it, which is how evaluation harnesses score multiple-choice answers and compute perplexity through an API.
Prompt logprobs need the logits of every prompt position, not just the last. build_batch therefore adds, for such requests, the rows whose next token is a prompt token after the sampling rows, and the engine turns them into log-probabilities as they’re computed, chunk by chunk:
def build_batch(scheduled, block_tables, block_size, device="cpu", pad_to=None):
"""Lay out [(request, n), ...] as one flattened batch. (Your engine: Chapter 31)
block_tables[request_id] lists the request's physical blocks in logical order. Request r's
n new tokens are token_ids[c : c + n] (plus its draft tokens, if any), at positions
c .. c + n - 1, where c is its num_computed_tokens. Position p lives in slot
table[p // block_size] * block_size + p % block_size.
pad_to (CUDA graphs, Chapter 33) appends dummy decode rows with slot -1 and length 0.
"""
ids, positions, slots, starts, lengths, tables, logits_rows, counts = [], [], [], [0], [], [], [], []
prompt_rows, prompt_spans = [], []
for request, n in scheduled:
c = request.num_computed_tokens
tokens = (request.token_ids + request.spec_token_ids)[c:c + n]
if len(tokens) != n:
raise ValueError(f"{request.request_id}: scheduled {n} tokens but only {len(tokens)} exist")
table = block_tables[request.request_id]
ids.extend(tokens)
positions.extend(range(c, c + n))
slots.extend(table[p // block_size] * block_size + p % block_size for p in range(c, c + n))
starts.append(starts[-1] + n)
lengths.append(c + n)
tables.append(table)
# Rows from the last real token onwards produce samples: one for a finished prefill or a
# decode, 1 + (draft tokens scheduled) when verifying. A mid-prompt chunk produces none.
k = min(n, max(0, c + n - (request.num_tokens - 1)))
logits_rows.extend(range(starts[-1] - k, starts[-1]))
counts.append(k)
if request.params.prompt_logprobs is not None: # rows whose next token is a prompt token
last = min(c + n, request.num_prompt_tokens - 1)
if last > c:
prompt_rows.extend(range(starts[-2], starts[-2] + last - c))
prompt_spans.append((request, c, request.token_ids[c + 1:last + 1]))
if pad_to is not None:
for _ in range(pad_to - len(lengths)):
ids.append(0), positions.append(0), slots.append(-1)
starts.append(starts[-1] + 1)
lengths.append(0)
tables.append([])
width = max(1, max(len(t) for t in tables))
table_tensor = torch.zeros((len(tables), width), dtype=torch.int32)
for r, table in enumerate(tables):
table_tensor[r, :len(table)] = torch.tensor(table, dtype=torch.int32)
meta = BatchMeta(
query_start_loc=torch.tensor(starts, dtype=torch.int32, device=device),
seq_lens=torch.tensor(lengths, dtype=torch.int32, device=device),
block_table=table_tensor.to(device),
slot_mapping=torch.tensor(slots, dtype=torch.int64, device=device),
block_size=block_size, query_start_loc_cpu=starts, seq_lens_cpu=lengths,
max_query_len=max(b - a for a, b in zip(starts, starts[1:])), max_seq_len=max(lengths))
return Batch(torch.tensor(ids, dtype=torch.long, device=device),
torch.tensor(positions, dtype=torch.long, device=device), meta,
torch.tensor(logits_rows + prompt_rows, dtype=torch.long, device=device), counts,
[r.request_id for r, _ in scheduled], {}, prompt_spans)
A cached prefix is never computed, so its logits don’t exist. The block manager therefore gives no prefix-cache hits to requests that ask for prompt logprobs, the same rule vLLM uses. Such requests still publish their blocks, so they help later requests.
Many samples: n and beam search
Parallel sampling (n=4: four independent answers to one prompt) looks like it needs forked caches, as in Chapter 25. With the prefix cache, it doesn’t: the engine turns it into four ordinary requests, rid:0 to rid:3, with seeds seed + i. They’re admitted in the same step, and since the scheduler publishes blocks as soon as it allocates them (Chapter 31), children 1-3 adopt child 0’s prompt blocks immediately. Only the last, partial block is computed four times. The frontend (Chapter 36) gathers the children back into one response.
Beam search keeps the $W$ most probable sequences, extending each by its most probable next tokens. It’s deterministic and favors high-likelihood outputs, which suits translation and some structured tasks, though it’s known to produce bland, repetitive text in open-ended generation (Holtzman et al., 2020). Inside an engine it’s awkward: beams fork and die every step, and the scheduler would need to know about them. vLLM V1 moved beam search out of the engine, and so does this book:
def beam_search(engine, prompt_ids, beam_width, max_tokens, eos_token_id=None, length_penalty=1.0):
"""Returns [(tokens, score)] for the best beam_width sequences, best first. (Your engine: Chapter 34)
score = sum of token logprobs / (generated length ** length_penalty). A beam that emits EOS
is finished and stops growing; the search ends when W beams have finished or max_tokens
is reached.
"""
beams, finished = [([], 0.0)], []
params = SamplingParams(max_tokens=1, logprobs=2 * beam_width, ignore_eos=True)
for step in range(max_tokens):
outputs = []
prompts = {f"beam-{step}-{i}": list(prompt_ids) + tokens for i, (tokens, _) in enumerate(beams)}
engine.generate(prompts, params, outputs=outputs)
top = {o.request_id: o.logprobs[0]["top"] for o in outputs}
candidates = []
for i, (tokens, logprob) in enumerate(beams):
for token, lp in top[f"beam-{step}-{i}"].items():
candidates.append((tokens + [token], logprob + lp))
candidates.sort(key=lambda c: -c[1])
beams = []
for tokens, logprob in candidates:
if eos_token_id is not None and tokens[-1] == eos_token_id:
finished.append((tokens, logprob))
else:
beams.append((tokens, logprob))
if len(beams) == beam_width:
break
if len(finished) >= beam_width or not beams:
break
pool = finished + beams
scored = [(tokens, logprob / len(tokens) ** length_penalty) for tokens, logprob in pool]
return sorted(scored, key=lambda c: -c[1])[:beam_width]
Each step submits one single-token request per beam, asking for the top $2W$ logprobs (enough that $W$ survive even if some beams end). A beam’s request shares everything but its newest token with its parent’s request of the previous step, so the prefix cache recomputes about one block per beam per step. The cost is one engine round-trip per generated token rather than a fused loop, which is the right trade-off for a feature that’s rarely used in serving.
Constrained decoding
A model asked for JSON usually produces JSON. A program that parses the output needs always: one missing quote in a million requests is a production incident. Constrained decoding guarantees it by masking, at every step, every token that would make the output impossible to complete validly. The model still chooses, among the tokens that keep the output valid.
The pieces: compile the pattern to an automaton, and at each step allow exactly the tokens that keep the automaton alive.
Why the automaton runs over bytes
Tokens are byte strings. In a byte-level BPE vocabulary (Chapter 4), é (UTF-8 C3 A9) may be one token, or two tokens C3 and A9, and many tokens end in the middle of a multi-byte character, especially for Chinese, Japanese and emoji. A character-level automaton can’t say whether the token C3 is allowed. A byte-level automaton can: after C3 it’s in a state that expects one continuation byte 80-BF.
So the regex compiler turns every character class into byte sequences. A range of code points becomes a few sequences of byte ranges:
U+0000-U+007F [00-7F]
U+0080-U+07FF [C2-DF][80-BF]
U+0800-U+0FFF [E0][A0-BF][80-BF]
U+1000-U+CFFF [E1-EC][80-BF][80-BF]
...
def utf8_sequences(lo, hi):
"""Byte-range sequences whose concatenations are exactly the UTF-8 encodings of the code
points lo..hi. (Your engine: Chapter 34)
Example: U+0080..U+07FF is [C2-DF][80-BF]. A range is first split where the encoded length
changes (and around the surrogates, which have no encoding), then wherever a prefix
byte doesn't cover a whole block of continuation bytes, until each piece is a product of
byte ranges.
"""
for a, b in ((0, 0x7F), (0x80, 0x7FF), (0x800, 0xD7FF), (0xE000, 0xFFFF), (0x10000, MAX_CODEPOINT)):
if lo <= b and hi >= a:
yield from _split(max(lo, a), min(hi, b))
def _split(lo, hi):
n = len(chr(lo).encode())
for i in range(1, n):
mask = (1 << (6 * i)) - 1 # the low i continuation bytes
if lo & ~mask != hi & ~mask:
if lo & mask:
yield from _split(lo, lo | mask)
yield from _split((lo | mask) + 1, hi)
return
if hi & mask != mask:
yield from _split(lo, (hi & ~mask) - 1)
yield from _split(hi & ~mask, hi)
return
yield list(zip(chr(lo).encode(), chr(hi).encode()))
The splitting rule is the one from RE2 and Rust’s regex-syntax: a range is cut where the encoded length changes and around the surrogates (U+D800-U+DFFF, which UTF-8 can’t encode), then wherever a prefix byte doesn’t cover a whole block of continuation bytes, until each piece is a product of byte ranges. The test checks it by brute force: the byte strings a set of sequences generates must be exactly the encodings of the code points in the range.
From regex to DFA
The parser is ordinary recursive descent: alternation of sequences of quantified atoms. Its output is a small tree:
class Parser:
"""Recursive descent over the regex text. AST nodes are tuples:
("chars", [(lo, hi), ...]) one code point from these ranges
("cat", [nodes]), ("alt", [nodes]), ("repeat", node, min, max or None)
"""
CLASSES = {"d": [(48, 57)], "w": [(48, 57), (65, 90), (95, 95), (97, 122)],
"s": [(9, 13), (32, 32)]}
ESCAPES = {"n": 10, "t": 9, "r": 13, "f": 12, "v": 11, "0": 0}
def __init__(self, text):
self.text, self.i = text, 0
def parse(self):
node = self.alternation()
if self.i != len(self.text):
raise ValueError(f"Unexpected {self.text[self.i]!r} at {self.i} in {self.text!r}")
return node
def peek(self):
return self.text[self.i] if self.i < len(self.text) else None
def take(self):
ch = self.text[self.i]
self.i += 1
return ch
def alternation(self):
branches = [self.sequence()]
while self.peek() == "|":
self.take()
branches.append(self.sequence())
return branches[0] if len(branches) == 1 else ("alt", branches)
def sequence(self):
items = []
while self.peek() not in (None, "|", ")"):
items.append(self.quantified())
return ("cat", items)
def quantified(self):
node = self.atom()
while self.peek() in ("*", "+", "?", "{"):
ch = self.peek()
if ch == "{":
match = re.match(r"\{(\d+)(,(\d*))?\}", self.text[self.i:])
if not match:
break # a literal brace
self.i += match.end()
lo = int(match.group(1))
hi = lo if match.group(2) is None else (int(match.group(3)) if match.group(3) else None)
node = ("repeat", node, lo, hi)
continue
self.take()
node = ("repeat", node, {"*": 0, "+": 1, "?": 0}[ch], {"*": None, "+": None, "?": 1}[ch])
return node
def atom(self):
ch = self.take()
if ch == "(":
if self.text.startswith("?:", self.i):
self.i += 2
node = self.alternation()
if self.peek() != ")":
raise ValueError(f"Unclosed group in {self.text!r}")
self.take()
return node
if ch == "[":
return ("chars", self.char_class())
if ch == ".":
return ("chars", [(0, 9), (11, MAX_CODEPOINT)])
if ch == "\\":
return ("chars", self.escape())
if ch in "*+?":
raise ValueError(f"Nothing to repeat at {self.i - 1} in {self.text!r}")
return ("chars", [(ord(ch), ord(ch))])
def escape(self):
ch = self.take()
if ch in self.CLASSES:
return self.CLASSES[ch]
if ch.lower() in self.CLASSES:
return complement(self.CLASSES[ch.lower()])
if ch == "u":
code = int(self.text[self.i:self.i + 4], 16)
self.i += 4
return [(code, code)]
if ch == "x":
code = int(self.text[self.i:self.i + 2], 16)
self.i += 2
return [(code, code)]
code = self.ESCAPES.get(ch, ord(ch))
return [(code, code)]
def char_class(self):
negate = self.peek() == "^"
if negate:
self.take()
ranges, first = [], True
while first or self.peek() != "]":
first = False
if self.peek() is None:
raise ValueError(f"Unclosed class in {self.text!r}")
ch = self.take()
lo = self.escape() if ch == "\\" else [(ord(ch), ord(ch))]
if len(lo) == 1 and lo[0][0] == lo[0][1] and self.peek() == "-" and self.text[self.i + 1:self.i + 2] not in ("]", ""):
self.take()
ch = self.take()
hi = self.escape() if ch == "\\" else [(ord(ch), ord(ch))]
ranges.append((lo[0][0], hi[0][0]))
else:
ranges.extend(lo)
self.take()
return complement(ranges) if negate else normalize(ranges)
Thompson’s construction turns the tree into an NFA with byte-range edges and epsilon edges, one small fragment per node; {m,n} becomes $m$ mandatory copies and $n - m$ optional ones. Subset construction then turns the NFA into a DFA whose states are sets of NFA states, with a full 256-entry transition row per state:
class DFA:
"""table[s, byte] = next state; state 0 is dead (absorbing), state 1 is the start."""
def __init__(self, pattern, max_states=20000):
nfa = NFA()
start, end = nfa.state(), nfa.state()
nfa.build(Parser(pattern).parse(), start, end)
first = nfa.closure([start])
ids = {first: 1}
rows, accepting, todo = [[0] * 256, [0] * 256], [False, end in first], [first]
while todo: # subset construction
current = todo.pop()
row = rows[ids[current]]
moves = {}
for s in current:
for a, b, nxt in nfa.edges[s]:
for byte in range(a, b + 1):
moves.setdefault(byte, set()).add(nxt)
for byte, targets in moves.items():
target = nfa.closure(targets)
if target not in ids:
if len(ids) + 1 > max_states:
raise ValueError("Pattern needs too many DFA states; simplify it")
ids[target] = len(rows)
rows.append([0] * 256)
accepting.append(end in target)
todo.append(target)
row[byte] = ids[target]
self.table = torch.tensor(rows, dtype=torch.int32)
self.accepting = torch.tensor(accepting)
# A state is "final" when it accepts and no byte leads anywhere: only EOS may follow.
self.final = self.accepting & (self.table == 0).all(-1)
def match(self, data):
state = 1
for byte in data:
state = int(self.table[state, byte])
return bool(self.accepting[state])
State 0 is dead: every byte leads back to it, and any token that reaches it is forbidden. A state is final when it accepts and no byte leads anywhere; the output is complete, and only EOS may follow. Russ Cox’s Regular Expression Matching Can Be Simple And Fast (2007) explains why this construction is linear in the input, unlike the backtracking matchers in most languages’ standard libraries.
From DFA to token masks
For a DFA state $s$ and a vocabulary of $V$ tokens, the mask says which tokens’ bytes, fed one by one from $s$, never reach the dead state. Walking 150,000 tokens one at a time in Python would take a second per state. Walking them all at once takes milliseconds: put every token’s bytes in a padded [V, max_len] tensor, start every token at $s$, and advance all of them one byte position at a time with a single table lookup:
class TokenIndex:
"""Which tokens each DFA state allows, computed by walking ALL tokens' bytes at once."""
def __init__(self, dfa, vocab_bytes, eos_token_id, vocab_size=None):
self.dfa, self.eos = dfa, eos_token_id
self.vocab_size = vocab_size or len(vocab_bytes)
width = max(1, max(len(b) for b in vocab_bytes))
self.lengths = torch.tensor([len(b) for b in vocab_bytes])
self.bytes = torch.zeros((len(vocab_bytes), width), dtype=torch.long)
for i, b in enumerate(vocab_bytes):
if b:
self.bytes[i, :len(b)] = torch.tensor(list(b))
self.cache = {}
def next_states(self, state):
"""[vocab] next DFA state after each token from `state` (0 = forbidden). (Your engine: Chapter 34)"""
if state not in self.cache:
current = torch.full((self.bytes.shape[0],), state, dtype=torch.int32)
for position in range(self.bytes.shape[1]):
step = self.dfa.table[current.long(), self.bytes[:, position]]
current = torch.where(position < self.lengths, step, current)
current[self.lengths == 0] = 0 # empty tokens (specials) never advance a pattern
self.cache[state] = current
return self.cache[state]
def allowed(self, state, device="cpu"):
"""[vocab_size] bool: tokens that keep the pattern alive, plus EOS in accepting states."""
mask = torch.zeros(self.vocab_size, dtype=torch.bool)
nxt = self.next_states(state)
mask[:nxt.shape[0]] = nxt != 0
if self.eos is not None:
mask[self.eos] = bool(self.dfa.accepting[state])
return mask.to(device)
The result also gives each token’s next state, so advancing a request after sampling is one lookup. Results are cached per state, and states are computed lazily, only when some request reaches them. A pattern with 440 DFA states rarely visits more than a few dozen. EOS is allowed exactly in accepting states.
class Guide:
"""One request's position in its pattern."""
def __init__(self, index):
self.index, self.state = index, 1
def allowed(self, device="cpu"):
return self.index.allowed(self.state, device)
def advance(self, token):
"""(Your engine: Chapter 34)"""
if token == self.index.eos:
self.state = -1 # done
return
nxt = int(self.index.next_states(self.state)[token])
if nxt == 0:
raise ValueError(f"Token {token} violates the pattern")
self.state = nxt
@property
def finished(self):
"""No byte can follow: the output is complete even without an EOS token."""
return self.state == -1 or bool(self.index.dfa.final[self.state])
Each request holds a Guide with its current state. The sampler masks with guide.allowed() before anything else and calls guide.advance(token) after the draw; the engine stops the request when the guide is final, even without an EOS token.
JSON Schema to regex
Most structured output in practice is “JSON matching this schema”, from OpenAI’s response_format and from tool calling (Chapter 35). A practical subset of JSON Schema is regular, so it can be translated into a regex, the approach of Outlines (Willard and Louf, 2023):
def json_schema_to_regex(schema, defs=None, depth=0):
"""A regex whose matches are exactly the (compact, optionally single-spaced) JSON documents
valid under a practical subset of JSON Schema. (Your engine: Chapter 34)
Supported: type (one or a list), properties + required (in declared order), items with
minItems/maxItems, enum, const, anyOf/oneOf, $ref into $defs, string minLength/maxLength/
pattern, and the primitives. Objects are closed (no extra keys).
"""
if depth > 16:
raise ValueError("Schema nests too deeply (recursive $ref?)")
defs = defs if defs is not None else schema.get("$defs", schema.get("definitions", {}))
recurse = lambda s: json_schema_to_regex(s, defs, depth + 1) # noqa: E731
if "$ref" in schema:
return recurse(defs[schema["$ref"].split("/")[-1]])
if "const" in schema:
return regex_escape(json.dumps(schema["const"]))
if "enum" in schema:
return "(?:" + "|".join(regex_escape(json.dumps(v)) for v in schema["enum"]) + ")"
for key in ("anyOf", "oneOf"):
if key in schema:
return "(?:" + "|".join(recurse(s) for s in schema[key]) + ")"
kind = schema.get("type")
if isinstance(kind, list):
return "(?:" + "|".join(recurse({**schema, "type": k}) for k in kind) + ")"
if kind == "object" or (kind is None and "properties" in schema):
props = schema.get("properties", {})
required = set(schema.get("required", []))
items = [(f'"{regex_escape(name)}"{WS}:{WS}{recurse(sub)}', name in required) for name, sub in props.items()]
if not items:
return r"\{" + WS + r"\}"
branches = []
for first, (pattern, _) in enumerate(items): # which property comes first decides the commas
if any(req for _, req in items[:first]):
break # a required property can't be skipped
body = pattern
for later, req in items[first + 1:]:
part = f"{WS},{WS}{later}"
body += part if req else f"(?:{part})?"
branches.append(body)
inner = "(?:" + "|".join(branches) + ")"
if not required:
inner += "?"
return r"\{" + WS + inner + WS + r"\}"
if kind == "array":
item = recurse(schema.get("items", {"type": "string"}))
lo, hi = schema.get("minItems", 0), schema.get("maxItems")
more = f"(?:{WS},{WS}{item})"
if hi is not None and hi == 0:
body = ""
elif lo == 0:
body = f"(?:{item}{more}{{0,{hi - 1}}})?" if hi is not None else f"(?:{item}{more}*)?"
else:
body = item + (more + (f"{{{lo - 1},{hi - 1}}}" if hi is not None else f"{{{lo - 1},}}"))
return r"\[" + WS + body + WS + r"\]"
if kind == "string":
if "pattern" in schema:
return f'"{schema["pattern"].lstrip("^").rstrip("$")}"'
lo, hi = schema.get("minLength"), schema.get("maxLength")
if lo is not None or hi is not None:
return f'"{STRING_CHAR}{{{lo or 0},{"" if hi is None else hi}}}"'
return PRIMITIVES["string"]
if kind in PRIMITIVES:
return PRIMITIVES[kind]
if kind is None: # {} = any primitive (bounded: no nesting)
return "(?:" + "|".join(PRIMITIVES.values()) + ")"
raise ValueError(f"Unsupported schema type {kind!r}")
The only subtle part is objects with optional properties: a comma separates two properties only if both are present. The translation chooses which property comes first (the alternation), after which every later property is , "key": value, mandatory if required and optional otherwise. Properties stay in declared order, which keeps the regex, and the DFA, small. Objects are closed: no keys the schema doesn’t list.
What doesn’t fit: recursive schemas (a tree whose nodes contain trees) and arbitrary JSON ({"type": "json_object"} with unbounded nesting) aren’t regular, because a finite automaton can’t count brackets. They need a pushdown automaton, or a context-free grammar engine such as XGrammar (Dong et al., 2024) or llguidance, which keep a stack, precompute masks for the tokens whose validity doesn’t depend on it, and check the rest at runtime. The interface stays the same: a mask per step and an advance per token (stretch exercise 3).
What constrained decoding costs
Three costs, all addressed by production engines:
- Compile time: building the DFA and its first masks. Cache compiled patterns (
compile_patternand the factory’s indexes do); tool schemas repeat across requests. - Mask time per step: a few milliseconds for a new state, under a millisecond for a cached one, on the CPU. Engines compute the next step’s masks while the GPU runs the forward pass, and apply them as bitmasks (32 tokens per
int32). - Synchronous scheduling: the next mask depends on the token just sampled, so a guided request can’t use Chapter 33’s async scheduling, whose next step is planned before the token is known. This engine refuses the combination; production engines advance the grammar on the device or delay the mask by a step.
SGLang adds a speedup worth knowing: when the DFA has only one path forward (inside a fixed key like "name": "), the engine can append all its tokens at once without asking the model, called jump-forward decoding (stretch exercise 4).
Run it
python run.py guided
{"regex_chars": 295, "dfa_states": 440, "compile_ms": 22.4}
{"state": "start", "allowed_tokens": 195, "first_mask_ms": 17.3, "cached_mask_ms": 0.67}
{"state": "after '{\"name\": \"'", "allowed_tokens": 141111, "first_mask_ms": 15.9, "cached_mask_ms": 0.45}
{"seed": 0, "output": "{\"name\" : \"age\",\"age\": -243,\"admin\":true }", "valid_json": true}
{"seed": 1, "output": "{\"name\":\"name\" ,\"age\":-6638951103811364110 }", "valid_json": true}
{"seed": 2, "output": "{\"name\" : \"c++\", \"age\" :3 , \"admin\":true}", "valid_json": true}
{"greedy": [244, 158, 145, 235, 484, 225], "beam_best": [244, 158, 145, 235, 484, 225], "beam_best_mean_logprob": -4.52}
A schema with four properties, an enum array and optional fields compiles to a 440-state DFA in 22 ms. With a synthetic vocabulary the size of Qwen3’s (151,936 byte strings), the first mask for a state takes about 17 ms and a cached one under a millisecond. At the start, only the 195 tokens that begin with { and continue validly are allowed; inside a string value, 141,111 are. Then the random 2-layer test model, which knows nothing about JSON, produces valid, schema-conforming JSON with three different seeds. (Its choices are silly; the guarantee isn’t.) Finally, beam search with four beams finds the greedy sequence here, as it often does when one continuation dominates.
Build it
Engine milestone 34: the sampler and structured output. Implement apply_processors, apply_penalties, top_k_top_p_min_p and Sampler.sample_device in engine/serve/sampler.py; beam_search in engine/serve/beam.py; and utf8_sequences, TokenIndex.next_states, Guide.advance and json_schema_to_regex in engine/serve/structured.py (the regex parser, NFA, DFA, logprob formatting and the engine wiring are provided). Use it with EngineConfig(sampler="full") and EngineCore(..., vocab_bytes=...).
pytest tests/test_ch34_sampling.py
python run.py guided --impl engine
The tests check UTF-8 sequences by brute force, the DFA against Python’s re on ten patterns and twenty-five strings, token masks against a byte-by-byte walk (including half of a two-byte character), EOS only in accepting states, the JSON Schema translation on valid and invalid documents with required, optional and $ref properties, each processor and penalty against its formula, the filters against Chapter 8’s order, the sampling distribution and seed independence from the batch, logprobs and prompt logprobs against a full forward pass, n=4 sharing the prompt’s blocks, beam search against a brute-force implementation, and guided generation producing pattern- and schema-valid output.
Stretch exercises
- ★ Keep penalty counts in a persistent
[max_num_seqs, vocab]tensor indexed by request slot, updated with each step’s sampled tokens. Measure host time per step againstapply_penaltiesat 256 requests with 1,000-token outputs. Where: penalty state inSamplerinengine/serve/sampler.py, with request-slot lifecycle inengine/serve/engine.py. - ★★ Pack masks as
int32bitmasks (32 tokens per word) and apply them with a small Triton kernel that expands bits to $-\infty$. Compare memory and time against boolean masks for 256 guided requests. Where: mask storage inengine/serve/structured.py; newengine/kernels/triton_masks.py, called fromengine/serve/sampler.py. - ★★★ Support arbitrary JSON (
{"type": "json_object"}) with a pushdown guide: a byte-level automaton for tokens plus a stack of open{and[. Precompute, per automaton state, the tokens that don’t touch the stack, and check the rest at runtime (the core idea of XGrammar). Where: guide state andGuideFactoryinengine/serve/structured.py. - ★★ Implement jump-forward decoding: when the guide’s state has exactly one path for the next several bytes, tokenize that text and append its tokens directly, skipping the model for those positions. Measure steps saved on the schema of
run.py guided. Where: deterministic-path detection inengine/serve/structured.py, with token/cache progression inengine/serve/engine.py.
Check your understanding
- Why are logprobs reported from the raw logits rather than after temperature and filters?
- What’s the difference between
repetition_penaltyandfrequency_penalty, and which one would you avoid for summarization? - Why does the exponential race draw token $i$ with probability $p_i$? Why is it graph-friendly?
- Why does a seeded request use its own generator, and what would go wrong with one shared generator?
- Why can’t a request with
prompt_logprobsuse the prefix cache? - How does
n=4share the prompt’s KV cache without forking, and which block is computed four times? - Why must the constraint automaton run over bytes rather than characters?
- Why can a regex describe “an object with these properties” but not “any JSON value”?
Going deeper
- Willard and Louf, Efficient Guided Generation for Large Language Models (2023), the Outlines paper; Dong et al., XGrammar: Flexible and Efficient Structured Generation Engine for Large Language Models (2024); Microsoft’s llguidance; LMSYS, Fast JSON Decoding for Local LLMs with Compressed Finite State Machine (2024), for jump-forward decoding.
- Russ Cox, Regular Expression Matching Can Be Simple And Fast (2007), and the
utf8module of Rust’sregex-syntaxcrate, for Thompson NFAs and UTF-8 range compilation. - Holtzman et al., The Curious Case of Neural Text Degeneration (ICLR 2020), for nucleus sampling and beam search’s failure modes; Nguyen et al., Turning Up the Heat: Min-p Sampling for Creative and Coherent LLM Outputs (2024); Keskar et al., CTRL (2019), for the repetition penalty.
- vLLM’s
vllm/v1/sample/(sampler, penalties, logprobs) andvllm/v1/structured_output/, which follow the structure of this chapter.
35. Text in, text out: tokenizers, templates and tool calls
In this chapter
- Loading any byte-level or SentencePiece-style BPE tokenizer from
tokenizer.json, matching Hugging Face'stokenizerstoken for token, with no dependency on it. - Streaming text from a stream of tokens without ever printing half a character.
- Stop strings that span tokens, and holding back text that might be the start of one.
- Chat templates rendered exactly as
transformersrenders them, safely. - Turning a model's tool calls and reasoning into the API's structured fields, while streaming, and forcing a tool call with Chapter 34's constrained decoding.
You will build
Tokenizer.encode and Tokenizer.bpe, IncrementalDetokenizer and StopChecker (engine/serve/tokenizer.py); render_chat and TagStreamer.feed (engine/serve/chat.py).
Time: 5-7 hours. GPU: not needed.
The text boundary is part of the engine
Chapter 18’s engine used Hugging Face’s tokenizer “at the text boundary” and stopped there. A server can’t. Every request crosses the boundary twice: its messages become a prompt through a chat template and a tokenizer, and its tokens become streamed text, stop-string checks, tool calls and reasoning. Mistakes here look like model bugs: a chat template missing one newline costs measurable accuracy, a detokenizer that prints each token’s bytes shows � in every emoji, a stop string split across two tokens is never noticed, and a tool call printed as text breaks every agent that called the API.
Owning this code also matters for speed and deployment. The engine of Part VIII has no dependency on transformers at serving time, and Chapter 36 runs tokenization in a separate process from the engine loop, so it must be code you control.
Loading tokenizer.json
A tokenizer.json file describes a pipeline: a normalizer (often Unicode NFC), a pre-tokenizer that splits text into words, a model (BPE: a vocabulary and a ranked list of merges), added tokens that are matched before anything else, and a decoder. Two families of BPE cover nearly all open models:
| byte-level BPE | SentencePiece-style BPE | |
|---|---|---|
| models | GPT-2, Llama 3, Qwen, DeepSeek, Mistral (Tekken) | Llama 2, Mistral v0.1-v0.3, Gemma |
| pre-tokenizer | a regex splits words; each word’s UTF-8 bytes are mapped to printable characters | spaces become ▁; words start at ▁ |
| unknown text | impossible: all 256 bytes are in the vocabulary | characters outside the vocabulary fall back to <0x00>-<0xFF> byte tokens |
| decoder | map characters back to bytes, UTF-8 decode | ▁ → space, byte tokens → bytes, drop the leading space |
Byte-level vocabularies store tokens as strings over a 256-character alphabet that GPT-2 introduced, so that every byte has a printable stand-in (Ġ is a space, Ċ a newline):
@lru_cache(maxsize=1)
def bytes_to_unicode():
"""GPT-2's reversible map from the 256 byte values to printable characters: printable
Latin-1 bytes map to themselves, the rest to code points 256 and up. Byte-level vocabularies
store tokens as strings over this alphabet ("Ġ" is a space, byte 0x20)."""
keep = list(range(ord("!"), ord("~") + 1)) + list(range(ord("¡"), ord("¬") + 1)) + list(range(ord("®"), ord("ÿ") + 1))
chars, extra = {}, 0
for b in range(256):
if b in keep:
chars[b] = chr(b)
else:
chars[b] = chr(256 + extra)
extra += 1
return chars
Encoding runs the pipeline. Added tokens such as <|im_start|> are found first, longest first, so a chat’s control tokens are never split or merged with text. The rest is normalized, split into words by the pre-tokenizer’s regex, and each word is merged independently:
class Tokenizer:
def __init__(self, spec, config=None):
model = spec["model"]
if model.get("type", "BPE") != "BPE":
raise ValueError(f"Only BPE tokenizers are implemented, not {model.get('type')}")
self.vocab = dict(model["vocab"])
self.ranks = {pair: i for i, pair in enumerate(_merges(model.get("merges", [])))}
self.byte_fallback = model.get("byte_fallback", False)
self.ignore_merges = model.get("ignore_merges", False)
self.unk = model.get("unk_token")
self.added = {t["content"]: t for t in spec.get("added_tokens", [])}
for content, token in self.added.items():
self.vocab.setdefault(content, token["id"])
self.id_to_token = {i: t for t, i in self.vocab.items()}
self.special_ids = {t["id"] for t in self.added.values() if t.get("special")}
pre = _flatten(spec.get("pre_tokenizer"), "pre")
self.splits = [regex.compile(p["pattern"].get("Regex") or regex.escape(p["pattern"]["String"]))
for p in pre if p["type"] == "Split"]
byte_level = [p for p in pre if p["type"] == "ByteLevel"]
self.byte_level = bool(byte_level) or any(d["type"] == "ByteLevel" for d in _flatten(spec.get("decoder"), "dec"))
if byte_level and byte_level[0].get("use_regex", True) and not self.splits:
self.splits = [regex.compile(r"'s|'t|'re|'ve|'m|'ll|'d| ?\p{L}+| ?\p{N}+| ?[^\s\p{L}\p{N}]+|\s+(?!\S)|\s+")]
self.add_prefix_space = bool(byte_level and byte_level[0].get("add_prefix_space"))
metaspace = [p for p in pre if p["type"] == "Metaspace"]
if metaspace and metaspace[0].get("split", True):
self.splits = self.splits + [regex.compile("▁?[^▁]+|▁+(?=▁)|▁+$")] # words start at ▁
norms = _flatten(spec.get("normalizer"), "norm")
self.nfc = any(n["type"] == "NFC" for n in norms)
self.space_to_meta = bool(metaspace) or any(n["type"] == "Replace" and n["content"] == "▁" for n in norms)
scheme = metaspace[0].get("prepend_scheme", "always") if metaspace else "always"
self.prepend_meta = self.space_to_meta and (scheme != "never" if metaspace else
any(n["type"] == "Prepend" for n in norms))
decoders = _flatten(spec.get("decoder"), "dec")
self.strip_leading_space = any(d["type"] == "Strip" and d.get("start", 0) for d in decoders) or \
any(d["type"] == "Metaspace" and d.get("prepend_scheme", "always") != "never" for d in decoders)
contents = sorted(self.added, key=len, reverse=True)
self.added_pattern = regex.compile("|".join(regex.escape(c) for c in contents)) if contents else None
self.byte_map = bytes_to_unicode()
self.byte_unmap = {c: b for b, c in self.byte_map.items()}
config = config or {}
self.chat_template = config.get("chat_template")
self.bos_token = _content(config.get("bos_token"))
self.eos_token = _content(config.get("eos_token"))
self.cache = {}
@classmethod
def from_dir(cls, directory):
directory = Path(directory)
config_path = directory / "tokenizer_config.json"
config = json.loads(config_path.read_text()) if config_path.exists() else {}
return cls(json.loads((directory / "tokenizer.json").read_text()), config)
@property
def vocab_size(self):
return max(self.id_to_token) + 1
def token_to_id(self, token):
return self.vocab.get(token)
@property
def eos_token_id(self):
return self.vocab.get(self.eos_token) if self.eos_token else None
# ------------------------------------------------------------ encoding
def encode(self, text):
"""Text -> token ids: added tokens first, then normalize, pre-tokenize and merge. (Your engine: Chapter 35)"""
ids, start = [], 0
for match in (self.added_pattern.finditer(text) if self.added_pattern else ()):
ids.extend(self._encode_ordinary(text[start:match.start()]))
ids.append(self.vocab[match.group()])
start = match.end()
ids.extend(self._encode_ordinary(text[start:]))
return ids
def _encode_ordinary(self, text):
if not text:
return []
if self.nfc:
text = unicodedata.normalize("NFC", text)
if self.byte_level:
if self.add_prefix_space and not text.startswith(" "):
text = " " + text
ids = []
for word in self.pre_tokenize(text):
ids.extend(self.bpe("".join(self.byte_map[b] for b in word.encode("utf-8"))))
return ids
if self.space_to_meta:
text = text.replace(" ", "▁")
if self.prepend_meta and not text.startswith("▁"):
text = "▁" + text
ids = []
for word in (self.pre_tokenize(text) if self.splits else [text]):
ids.extend(self.bpe(word))
return ids
def pre_tokenize(self, text):
"""Split into pieces with each Split regex, keeping both matches and gaps ("Isolated")."""
pieces = [text]
for pattern in self.splits:
out = []
for piece in pieces:
pos = 0
for m in pattern.finditer(piece):
if m.start() > pos:
out.append(piece[pos:m.start()])
if m.end() > m.start():
out.append(m.group())
pos = m.end()
if pos < len(piece):
out.append(piece[pos:])
pieces = out
return pieces
def bpe(self, word):
"""Merge the lowest-ranked adjacent pair until none is in the merge table. (Your engine: Chapter 35)"""
if word in self.cache:
return self.cache[word]
if self.ignore_merges and word in self.vocab:
return [self.vocab[word]]
parts = list(word)
while len(parts) > 1:
best, where = None, -1
for i, pair in enumerate(zip(parts, parts[1:])):
rank = self.ranks.get(pair)
if rank is not None and (best is None or rank < best):
best, where = rank, i
if best is None:
break
parts[where:where + 2] = [parts[where] + parts[where + 1]]
ids = []
for part in parts:
if part in self.vocab:
ids.append(self.vocab[part])
elif self.byte_fallback:
ids.extend(self.vocab[f"<0x{b:02X}>"] for b in part.encode("utf-8"))
elif self.unk is not None:
ids.append(self.vocab[self.unk])
else:
raise ValueError(f"Cannot encode {part!r}: not in the vocabulary and no byte fallback")
if len(self.cache) < 200_000:
self.cache[word] = ids
return ids
# ------------------------------------------------------------ decoding
def token_bytes(self, token_id, skip_special_tokens=False):
"""The bytes one token contributes to the text (b"" for skipped specials)."""
token = self.id_to_token.get(token_id)
if token is None or (skip_special_tokens and token_id in self.special_ids):
return b""
if token in self.added:
return token.encode("utf-8")
if self.byte_level:
return bytes(self.byte_unmap[c] for c in token)
if self.byte_fallback and len(token) == 6 and token.startswith("<0x") and token.endswith(">"):
return bytes([int(token[3:5], 16)])
return token.replace("▁", " ").encode("utf-8")
def decode(self, ids, skip_special_tokens=False):
data = b"".join(self.token_bytes(i, skip_special_tokens) for i in ids)
text = data.decode("utf-8", errors="replace")
if self.strip_leading_space and text.startswith(" "):
text = text[1:]
return text
def vocab_bytes(self):
"""bytes per token id for constrained decoding (Chapter 34); special tokens are b"" (never text)."""
return [b"" if i in self.special_ids else self.token_bytes(i) for i in range(self.vocab_size)]
The pre-tokenizer regex decides which merges are even possible, so it’s part of the model’s definition, not an implementation detail. Qwen’s splits letters from digits and every digit from the next (\p{N}); Llama 3’s allows digit groups of up to three (\p{N}{1,3}); GPT-2’s attaches a leading space to words. These patterns use Unicode property classes (\p{L} is any letter in any script), which Python’s built-in re lacks, so the loader uses the regex package.
bpe is Chapter 4’s algorithm with a learned merge table: repeatedly merge the adjacent pair with the lowest rank. Two production details are worth noticing. Results are cached per word, because natural text repeats words constantly; that cache is why this pure-Python implementation keeps up with Rust on ordinary text. And Llama 3’s ignore_merges flag says that a word that’s already in the vocabulary is one token, without running the merges.
The test trains three tokenizers with Hugging Face’s tokenizers library: a Qwen-style byte-level one with its split regex and NFC, a GPT-2-style one, and a SentencePiece-style one with byte fallback. It then checks that this loader produces exactly the same token IDs and the same decoded text on sixteen strings: control tokens, whitespace runs, contractions, code, Chinese, Japanese, emoji with zero-width joiners, and characters the tokenizer never saw.
vocab_bytes() gives each token’s bytes, which is what Chapter 34’s constrained decoding walks. Special tokens map to b"", so a pattern can never be satisfied by <|im_end|>. Tokens that were added but aren’t special, like Qwen3’s <think>, are ordinary text.
Streaming text without half characters
The engine produces tokens one at a time, and the server must stream text. Decoding each token’s bytes separately fails whenever a token ends inside a UTF-8 character, which byte-level BPE does routinely for rare characters:
tokens: a | ␠ | F0 9F | A7 91 | E2 80 8D | ...
per-token decode: "a ��������� on ����������������"
incremental: "a", " ", "", "🧑", "", "", "", "", "🚀", " on", ...
The fix is an incremental UTF-8 decoder: it returns every complete character and holds back an incomplete one until the bytes that finish it arrive. Python’s codecs module has one:
class IncrementalDetokenizer:
"""Turns a stream of token ids into a stream of text deltas. (Your engine: Chapter 35)
A token can end in the middle of a UTF-8 character (byte-level BPE splits "é" or an emoji
across tokens), so printing each token's bytes as they arrive would print replacement
characters. Bytes go through an incremental UTF-8 decoder, which holds back an incomplete
trailing character until the token that completes it arrives.
"""
def __init__(self, tokenizer, skip_special_tokens=True):
self.tokenizer, self.skip = tokenizer, skip_special_tokens
self.decoder = codecs.getincrementaldecoder("utf-8")(errors="replace")
self.text = ""
def add(self, token_ids):
data = b"".join(self.tokenizer.token_bytes(t, self.skip) for t in token_ids)
delta = self.decoder.decode(data, final=False)
self.text += delta
return delta
def flush(self):
delta = self.decoder.decode(b"", final=True)
self.text += delta
return delta
Engines that call a string-level decode (vLLM, Text Generation Inference) get the same effect with a more complex trick: they decode a window of recent tokens twice, with and without the newest one, and stream the difference only if it doesn’t end in a replacement character. Working at the byte level makes the problem disappear, because bytes are what tokens really are.
Stop strings
stop=["</answer>"] must end generation when the text contains </answer>, even if it arrives as </ans + wer>, and the client must never see the stop string, or any part of it, unless include_stop_str_in_output is set. Two rules handle that:
- Search the text the new delta could have completed: the last
len(longest stop) - 1characters before it, plus the delta. - Hold back any tail of the text that is a prefix of a stop string, until the next delta proves it isn’t.
"2.</ans"streams"2."and keeps"</ans". If the next delta is"wer>", the stop is complete and nothing more is sent; if it’s"ia", the held text is released as"</ansia".
class StopChecker:
"""Stop strings on streamed text. (Your engine: Chapter 35)
A stop string can span several deltas, so the checker searches the text that the new delta
could have completed, and it holds back (doesn't stream yet) any tail of the text that is
the start of a stop string, so that a client never sees "</ans" before the stop is found.
"""
def __init__(self, stops, include_stop=False):
self.stops, self.include = [s for s in stops if s], include_stop
self.longest = max((len(s) for s in self.stops), default=0)
self.text, self.sent = "", 0
self.stopped = None
def add(self, delta):
"""Append delta; returns the text that may be streamed now. self.stopped is set on a match."""
old = len(self.text)
self.text += delta
if self.stops:
window = max(0, old - self.longest + 1)
hits = [(self.text.find(s, window), s) for s in self.stops]
hits = [(i, s) for i, s in hits if i >= 0]
if hits:
index, stop = min(hits)
self.text = self.text[:index + (len(stop) if self.include else 0)]
self.stopped = stop
return self._release(len(self.text))
hold = 0
for s in self.stops: # the longest tail that begins some stop string
for k in range(min(len(s) - 1, len(self.text)), 0, -1):
if self.text.endswith(s[:k]):
hold = max(hold, k)
break
return self._release(len(self.text) - hold)
def finish(self):
return self._release(len(self.text))
def _release(self, upto):
out = self.text[self.sent:upto]
self.sent = max(self.sent, upto)
return out
Stop strings are checked on detokenized text, so they belong to the frontend (Chapter 36), not to the engine core. When the frontend finds one, it aborts the request in the core, which frees its blocks.
Chat templates
A chat model saw every training conversation rendered in one exact format. Qwen uses ChatML:
<|im_start|>system
You are terse.<|im_end|>
<|im_start|>user
Weather in Paris?<|im_end|>
<|im_start|>assistant
The format is defined by a Jinja template shipped in the model’s tokenizer_config.json, and the model’s quality depends on rendering it exactly: the newlines, the tools block, whether a <think> block is pre-filled. The only reliable way to match is to render the template the way transformers does:
def render_chat(template, messages, tools=None, add_generation_prompt=True, **variables):
"""Render a Hugging Face chat template exactly as transformers' apply_chat_template does. (Your engine: Chapter 35)
Same environment: a sandbox (templates come from downloaded files and must not run code),
trim_blocks and lstrip_blocks, the loopcontrols extension, and the helpers templates use:
raise_exception, strftime_now and a tojson filter that keeps non-ASCII text.
"""
from jinja2.ext import loopcontrols
from jinja2.sandbox import ImmutableSandboxedEnvironment
def raise_exception(message):
raise ValueError(message)
def tojson(value, ensure_ascii=False, indent=None, separators=None, sort_keys=False):
return json.dumps(value, ensure_ascii=ensure_ascii, indent=indent, separators=separators, sort_keys=sort_keys)
env = ImmutableSandboxedEnvironment(trim_blocks=True, lstrip_blocks=True, extensions=[loopcontrols])
env.filters["tojson"] = tojson
env.globals["raise_exception"] = raise_exception
env.globals["strftime_now"] = lambda fmt: datetime.now().strftime(fmt)
return env.from_string(template).render(messages=messages, tools=tools,
add_generation_prompt=add_generation_prompt, **variables)
Three details make it match. Whitespace control: trim_blocks and lstrip_blocks remove the newlines and indentation around {% %} tags, which templates are written to expect. Helpers: templates call raise_exception, strftime_now and a tojson filter that, unlike Jinja’s built-in one, keeps non-ASCII characters and accepts indent. Loop controls: templates use {% continue %}, which needs Jinja’s loopcontrols extension. And one detail makes it safe: the sandbox. A template is code that came with a downloaded model; the sandboxed environment refuses attribute access that could reach Python internals, so a malicious template can’t run code on the server.
The book ships CHATML_TEMPLATE, a ChatML template with tools and switchable thinking in the style of Qwen3’s, for models without one (the test models). The test renders a conversation with a system prompt, a tool call, a tool result and non-ASCII text, with and without tools and with thinking on and off, and checks that the result equals transformers’ apply_chat_template character for character.
Tool calls and reasoning
Tool calling is a convention on top of text. The template lists the tools in the system prompt; a model trained for it answers with
<tool_call>
{"name": "get_weather", "arguments": {"city": "Paris"}}
</tool_call>
and the API must return it as structured data: "tool_calls": [{"id": ..., "type": "function", "function": {"name": "get_weather", "arguments": "{\"city\": \"Paris\"}"}}] (the arguments are a JSON string in OpenAI’s format). Reasoning models similarly write <think>...</think> before the answer, which the API returns as reasoning_content.
Parsing a complete answer is a regex. Parsing a stream is harder: the text arrives a few characters at a time, <tool_call> can be split across deltas, and the client must never see half a tag as content. TagStreamer splits a stream into “outside” and “inside” segments for one pair of tags, holding back any tail that might start a tag:
TOOL_OPEN, TOOL_CLOSE = "<tool_call>", "</tool_call>"
def parse_tool_calls(text):
"""Complete output -> (content, [tool calls in OpenAI's format]). Hermes/Qwen style:
<tool_call>{"name": ..., "arguments": {...}}</tool_call>, possibly several."""
calls = []
for body in re.findall(re.escape(TOOL_OPEN) + r"(.*?)" + re.escape(TOOL_CLOSE), text, re.S):
calls.append(openai_tool_call(json.loads(body)))
content = re.sub(re.escape(TOOL_OPEN) + r".*?" + re.escape(TOOL_CLOSE), "", text, flags=re.S).strip()
return content, calls
def openai_tool_call(call):
arguments = call.get("arguments", call.get("parameters", {}))
return {"id": f"call_{uuid.uuid4().hex[:24]}", "type": "function",
"function": {"name": call["name"],
"arguments": arguments if isinstance(arguments, str) else json.dumps(arguments, ensure_ascii=False)}}
class TagStreamer:
"""Splits a text stream into "outside" and "inside" segments for one pair of tags, holding back
any tail that might be the start of a tag. (Your engine: Chapter 35)
feed(delta) returns [(inside, text), ...] for text that is certain; an inside segment is
reported once, whole, when its closing tag arrives.
"""
def __init__(self, open_tag, close_tag, start_inside=False):
self.open, self.close = open_tag, close_tag
self.inside, self.buffer = start_inside, ""
def feed(self, delta, final=False):
self.buffer += delta
out = []
while True:
tag = self.close if self.inside else self.open
index = self.buffer.find(tag)
if index >= 0:
if index or self.inside:
out.append((self.inside, self.buffer[:index]))
self.buffer = self.buffer[index + len(tag):]
self.inside = not self.inside
continue
if self.inside and not final:
return out # wait for the closing tag
hold = 0 if final else max((k for k in range(1, len(tag)) if self.buffer.endswith(tag[:k])), default=0)
if len(self.buffer) > hold:
out.append((self.inside, self.buffer[:len(self.buffer) - hold]))
self.buffer = self.buffer[len(self.buffer) - hold:]
return out
OutputParser chains two of them, one for <think> and one for <tool_call>, and emits OpenAI-style deltas:
class OutputParser:
"""Streams a Qwen3-style answer into OpenAI deltas: reasoning_content, content and tool_calls.
Newlines around the reasoning are formatting, not content: "<think>\n...\n</think>\n\n"
yields the reasoning without them and content that starts at the first real character.
A tool call that isn't valid JSON (or is cut off) is returned as content, not dropped.
"""
def __init__(self, reasoning=False, tools=False, starts_in_reasoning=False):
self.reasoning = TagStreamer("<think>", "</think>", starts_in_reasoning) if reasoning else None
self.tools = TagStreamer(TOOL_OPEN, TOOL_CLOSE) if tools else None
self.num_calls, self.content_started = 0, False
def feed(self, text, final=False):
deltas = []
segments = self.reasoning.feed(text, final) if self.reasoning else [(False, text)]
for in_think, part in segments:
if in_think:
if part.strip("\n"):
deltas.append({"reasoning_content": part.strip("\n")})
continue
for in_call, piece in (self.tools.feed(part, final) if self.tools else [(False, part)]):
if in_call:
try:
call = openai_tool_call(json.loads(piece))
except (ValueError, KeyError, TypeError):
deltas.append({"content": TOOL_OPEN + piece})
continue
deltas.append({"tool_calls": [{"index": self.num_calls, **call}]})
self.num_calls += 1
continue
if not self.content_started:
piece = piece.lstrip("\n")
self.content_started = bool(piece)
if piece:
deltas.append({"content": piece})
return deltas
A tool call is emitted whole when its closing tag arrives. vLLM streams the arguments’ JSON incrementally instead, which a few clients rely on; that’s stretch exercise 2. A malformed or unterminated call is returned as content rather than dropped, so a client can at least see what the model said.
Forcing a tool call
tool_choice="required" (call some tool) and tool_choice={"function": {"name": "get_weather"}} (call this tool) must be guarantees, not hopes. Chapter 34’s constrained decoding makes them so: the wrapper is literal text and the arguments follow the tool’s own JSON Schema, so the regex for “call get_weather” is the escaped prefix, the schema’s regex, and the escaped suffix:
def tool_call_pattern(tools, name=None):
"""A regex that forces the model to call a tool (tool_choice="required") or one named tool:
the call's wrapper is literal text and its arguments follow the tool's JSON Schema
(Chapter 34)."""
from .structured import json_schema_to_regex, regex_escape
options = []
for tool in tools:
fn = tool["function"]
if name is not None and fn["name"] != name:
continue
arguments = json_schema_to_regex(fn.get("parameters") or {"type": "object", "properties": {}})
options.append(regex_escape(f'{TOOL_OPEN}\n{{"name": "{fn["name"]}", "arguments": ') + arguments
+ regex_escape(f"}}\n{TOOL_CLOSE}"))
if not options:
raise ValueError(f"No tool named {name!r}")
return "(?:" + "|".join(options) + ")"
With it, even a model that would rather chat produces a parseable call with every required argument present and of the right type.
Run it
python run.py text # or --model-dir to use a real checkpoint's tokenizer
Without a checkpoint, the command trains a 4,000-token Qwen-style tokenizer on The Verdict with Hugging Face’s tokenizers, then loads its tokenizer.json with the book’s loader:
{"text": "The Verdict x4 (repetitive)", "tokenizer": "izh (Python)", "tokens": 18152, "tokens_per_s": 586506}
{"text": "The Verdict x4 (repetitive)", "tokenizer": "tokenizers (Rust)", "tokens": 18152, "tokens_per_s": 480221}
{"text": "The Verdict x4 (repetitive)", "identical_ids": true}
{"text": "20,000 random words (no cache hits)", "tokenizer": "izh (Python)", "tokens": 121540, "tokens_per_s": 1009295}
{"text": "20,000 random words (no cache hits)", "tokenizer": "tokenizers (Rust)", "tokens": 121540, "tokens_per_s": 1338480}
{"text": "20,000 random words (no cache hits)", "identical_ids": true}
{"tokens": 27, "per_token_decode": "a ��������� on ����������������", "incremental_deltas": ["a", " ", "", "🧑", "", "", "", "", "🚀", " on", " ", "", "", "", "𝔐", "", "", "", "𝔞", "", "", "", "𝔯", "", "", "", "𝔰"]}
{"stop_streamed": ["The answer is 4", "2", ""], "final": "The answer is 42"}
{"chat_prompt_tokens": 270, "chat_prompt_tail": "n Paris?<|im_end|>\n<|im_start|>assistant\n<think>\n\n</think>\n\n"}
{"parsed": [{"reasoning_content": "Need the tool."}, {"tool_calls": {"name": "get_weather", "arguments": "{\"city\": \"Paris\"}"}}]}
The IDs are identical to the Rust library’s. The speeds are close on this small vocabulary, and the word cache even puts Python ahead on repetitive text; with a real 150,000-token vocabulary each uncached word takes more merge rounds, and the Rust library also offers parallel batch encoding. The number that matters more for serving is that tokenization runs on the CPU, in Python, competing with the engine loop for the interpreter, which is why Chapter 36 moves it to its own process. The astronaut emoji, a zero-width-joiner sequence the tokenizer never saw, arrives as 27 tokens, and the incremental detokenizer turns them into whole characters with empty deltas in between, where per-token decoding prints replacement characters.
Build it
Engine milestone 35: the text boundary. Implement Tokenizer.encode, Tokenizer.bpe, IncrementalDetokenizer.add, IncrementalDetokenizer.flush and StopChecker.add in engine/serve/tokenizer.py, and render_chat and TagStreamer.feed in engine/serve/chat.py (the tokenizer.json parsing, byte map, decoding, the output parser and the forced-tool-call pattern are provided).
pytest tests/test_ch35_text.py
python run.py text --impl engine
The tests check encoding and decoding against Hugging Face’s tokenizers for Qwen-style, GPT-2-style and SentencePiece-style tokenizers on sixteen hard strings; that streaming never prints a replacement character while some tokens split characters; stop strings across deltas, with and without the stop text, and false alarms; vocab_bytes for constrained decoding; chat templates against transformers with tools, tool results and thinking on and off; tool-call and reasoning parsing when streamed one character at a time, with no partial tag leaking; and the forced-tool-call pattern.
Stretch exercises
- ★ Load a real checkpoint’s tokenizer (Qwen3, Llama 3, Mistral) with
run.py text --model-dir, and compare IDs withtokenizerson a few megabytes of mixed text. Find and fix any difference. Where:experiments/ch35.py(create it) for parity; fix differences inengine/serve/tokenizer.py. - ★★ Stream tool-call arguments incrementally: emit the name as soon as
"name": "..."is complete, then the arguments’ JSON text as it arrives, in OpenAI’s delta format. Where:OutputParserinengine/serve/chat.pyand delta formatting inServer.chat_chunksinengine/serve/api.py. - ★★ Make
bpefaster for long words: a linked list of symbols and a heap of candidate merges keyed by rank and position. Measure on a 50,000-character word without spaces (some languages and code produce these). Where:Tokenizer.bpeinengine/serve/tokenizer.py. - ★★★ Implement the Unigram model (T5, ALBERT, some multilingual models): Viterbi segmentation over the vocabulary’s log-probabilities. Match
tokenizerson a trained Unigram tokenizer. Where: model loading inTokenizer.from_dirand segmentation inengine/serve/tokenizer.py.
Check your understanding
- Why do byte-level vocabularies map bytes to printable characters rather than storing raw bytes?
- Why are added tokens matched before pre-tokenization, and what would go wrong otherwise?
- Why does the pre-tokenizer regex belong to the model’s definition?
- Why does per-token decoding print replacement characters, and how does an incremental UTF-8 decoder avoid them?
- A stop string is
"STOP"and the text so far ends with"ST". What does the checker stream, and when does it release the"ST"? - Why must chat templates be rendered in a sandbox?
- How does constrained decoding turn
tool_choicefrom a hint into a guarantee?
Going deeper
- Hugging Face tokenizers documentation (the pipeline: normalizers, pre-tokenizers, models, post-processors, decoders) and its
tokenizer.jsonformat; Sennrich et al., Neural Machine Translation of Rare Words with Subword Units (2016) for BPE; Kudo and Richardson, SentencePiece (2018). - Andrej Karpathy, Let’s build the GPT Tokenizer (2024) and his
minbperepository, for byte-level BPE and its regex pre-tokenizers. - Chat templates in the
transformersdocumentation, and vLLM’svllm/entrypoints/chat_utils.py,vllm/entrypoints/openai/tool_parsers/andreasoning/. - OpenAI’s API reference for chat completions, tool calls and streaming deltas, which defines the format the parser emits.
36. An OpenAI-compatible server
In this chapter
- The architecture of a serving process: an engine core that only schedules and runs the model, and an asyncio frontend that owns text, HTTP and streaming.
- Why production servers run the engine core in its own process, and how the two halves talk.
- The OpenAI API in detail: completions and chat completions, streaming with server-sent events, tools, JSON output, logprobs,
n, errors and usage. - The behaviors that separate a demo from a service: cancellation when clients leave, backpressure, error isolation, metrics and health.
You will build
core_loop, AsyncLLM._dispatch and AsyncLLM.generate (engine/serve/async_engine.py), and Server.sampling_params and Server.chat_prompt (engine/serve/api.py). Then python serve.py --model-dir ... serves a checkpoint to any OpenAI client.
Time: 6-8 hours. GPU: not needed.
One request’s journey
A chat request to the finished server passes through these stages:
client ──HTTP──> FastAPI route ──> Server.chat
│ validate, apply defaults, render the chat template,
│ tokenize, build SamplingParams (+ grammar for tools / JSON)
▼
AsyncLLM.generate ──("add", rid, ids, params)──> engine core (own process)
▲ │ Scheduler, ModelRunner,
│ per-request asyncio.Queue │ Sampler: Chapters 31-34
│ ▼
AsyncLLM._dispatch <──("outputs", [RequestOutput...])──
│ detokenize, stop strings, metrics
▼
OutputParser (reasoning, tool calls) ──SSE chunks──> client
The rest of Part VIII’s machinery sits inside the engine-core box. This chapter builds everything around it.
Two halves
The engine core is a tight loop: plan, run, sample, update, repeat. Anything that delays it delays every request’s next token. The frontend does work that’s unpredictable in size: parsing JSON bodies, rendering templates, tokenizing a 100,000-token document, detokenizing, serializing streamed chunks for hundreds of connections. In one Python process, the two compete for a single interpreter. CPython runs one thread’s bytecode at a time (the GIL), and Chapter 33 measured a millisecond of host work per step at 256 requests, so a frontend that tokenizes a long prompt can easily stall the decode loop for several steps.
So production servers split them. vLLM V1 runs an API-server process and an EngineCore process connected by ZeroMQ sockets; SGLang runs a tokenizer manager, a scheduler and a detokenizer manager as separate processes. This book’s AsyncLLM does the same with two queues, and can run the core either in a thread (simple to debug, and what the tests use) or in its own process (mode="process", the default for serve.py):
def core_loop(factory, inputs, outputs, stats_every=8):
"""Run an EngineCore until told to stop. (Your engine: Chapter 36)
Block for input only when idle; otherwise take whatever messages are waiting and step.
An invalid request fails alone ("error"); an exception inside step() fails every request
("fatal"), because the engine's state can no longer be trusted.
"""
try:
core = factory()
except Exception as error: # noqa: BLE001
outputs.put(("fatal", f"engine failed to start: {error!r}"))
return
outputs.put(("ready", {"max_model_len": core.max_model_len, "num_blocks": core.blocks.num_blocks}))
busy = False
while True:
if busy and not core.has_unfinished: # just went idle (finished or aborted): report it
outputs.put(("stats", core.stats()))
busy = core.has_unfinished
block = not busy
while True:
try:
kind, payload = inputs.get(block=block)
except queue.Empty:
break
block = False
if kind == "shutdown":
return
if kind == "abort":
core.abort_request(payload)
elif kind == "add":
rid, prompt, params, priority, options = payload
try:
core.add_request(rid, prompt, params, priority, **options)
except ValueError as error:
outputs.put(("error", (rid, str(error))))
if core.has_unfinished:
try:
step_outputs = core.step()
except Exception as error: # noqa: BLE001
outputs.put(("fatal", f"engine step failed: {error!r}"))
return
if step_outputs:
outputs.put(("outputs", step_outputs))
if core.steps % stats_every == 0:
outputs.put(("stats", core.stats()))
busy = True
The loop blocks only when idle. While requests are running it takes whatever messages are waiting and steps, so a burst of new requests is admitted at the next step without waiting for the batch to drain.
It also distinguishes two kinds of failure. A request the engine rejects (too long, a guided-decoding pattern it can’t compile) fails alone: the core reports ("error", rid, message) and keeps serving everyone else. An exception inside step() is different: the engine’s state, its block tables and its num_computed_tokens, may now be inconsistent, so the core reports "fatal", every in-flight request fails, and /health turns red so that an orchestrator (Kubernetes, systemd) restarts the process. Trying to limp on after a failed step risks serving one request’s cache to another.
In process mode, the factory that builds the engine is picklable (functools.partial of a module-level function with a checkpoint path), and the model is loaded in the child. Weights are never copied between processes, and the child process is started with spawn, because a forked process inherits CUDA state that it can’t use.
The frontend
AsyncLLM runs in the server’s event loop. Each request gets an asyncio.Queue; a background thread receives the core’s messages and hands them to the loop with call_soon_threadsafe, where _dispatch turns token IDs into text:
class AsyncLLM:
def __init__(self, factory, tokenizer, mode="thread", max_pending=1024, metrics=None,
request_timeout=120, max_buffered_updates=256):
if mode not in ("thread", "process") or max_pending < 1 or request_timeout <= 0 or max_buffered_updates < 1:
raise ValueError("Need thread|process mode and positive admission/deadline/buffer limits")
self.tokenizer, self.max_pending, self.mode = tokenizer, max_pending, mode
context = mp.get_context("spawn")
self.inputs = context.Queue() if mode == "process" else queue.Queue()
self.outputs = context.Queue() if mode == "process" else queue.Queue()
target = (context.Process if mode == "process" else threading.Thread)
self.worker = target(target=run_core, args=(factory, self.inputs, self.outputs), daemon=True)
self.requests, self.core_stats, self.info = {}, {}, None
self.metrics = metrics
self.request_timeout, self.max_buffered_updates = request_timeout, max_buffered_updates
self.closing = False
self.dead = None
self.ids = itertools.count()
async def start(self):
"""Start the core, and a daemon thread that hands the core's messages to the event loop;
wait until the model is loaded. (A daemon thread, not a loop executor: a thread blocked
on queue.get() must never keep the process from exiting.)"""
self.loop = asyncio.get_running_loop()
self.ready = self.loop.create_future()
self.worker.start()
threading.Thread(target=self._pump, daemon=True).start()
kind, payload = await self.ready
if kind != "ready":
self.dead = payload
raise EngineDead(payload)
self.info = payload
async def stop(self, drain_timeout=30):
"""Reject new work, give existing streams a deadline, then abort and join the worker."""
self.closing = True
deadline = time.monotonic() + drain_timeout
while self.requests and time.monotonic() < deadline:
await asyncio.sleep(0.01)
for rid, state in list(self.requests.items()):
self.abort(rid)
state.queue.put_nowait(EngineDead("Server drain deadline exceeded"))
self.inputs.put(("shutdown", None))
await asyncio.to_thread(self.worker.join, 2)
if self.mode == "process" and self.worker.is_alive():
self.worker.terminate()
await asyncio.to_thread(self.worker.join, 2)
self.outputs.put(("closed", None)) # ends the pump
await asyncio.sleep(0)
def _pump(self):
while True:
kind, payload = self.outputs.get()
if kind == "closed":
return
self.loop.call_soon_threadsafe(self._handle, kind, payload)
if kind == "fatal":
return
def _handle(self, kind, payload):
if not self.ready.done():
self.ready.set_result((kind, payload))
if kind == "ready":
return
if kind == "outputs":
for out in payload:
self._dispatch(out)
elif kind == "stats":
self.core_stats = payload
elif kind == "error":
rid, message = payload
state = self.requests.get(rid)
if state is not None:
state.queue.put_nowait(ValueError(message))
elif kind == "fatal":
self.dead = payload
for state in self.requests.values():
state.queue.put_nowait(EngineDead(payload))
def _dispatch(self, out: RequestOutput):
"""Token ids -> text for one request: detokenize, check stop strings, time it. (Your engine: Chapter 36)"""
state = self.requests.get(out.request_id)
if state is None or state.finished:
return # aborted, or already stopped by a stop string
if state.queue.qsize() >= self.max_buffered_updates:
self.abort(out.request_id)
state.finished = True
while not state.queue.empty():
state.queue.get_nowait()
state.queue.put_nowait(Overloaded("Client is not consuming streamed updates"))
return
now = time.monotonic()
if out.new_token_ids:
if state.first_token is None:
state.first_token = now
if self.metrics:
self.metrics.observe("time_to_first_token_seconds", now - state.arrival)
elif self.metrics:
self.metrics.observe("inter_token_latency_seconds", (now - state.last_token) / len(out.new_token_ids))
state.last_token = now
text = state.detok.add(out.new_token_ids)
if out.finished:
text += state.detok.flush()
visible = state.stop.add(text)
finished, reason, stop_reason = out.finished, out.finish_reason, out.stop_reason
if state.stop.stopped is not None:
if not out.finished:
self.abort(out.request_id) # the core doesn't know about stop strings
finished, reason, stop_reason = True, "stop", state.stop.stopped
elif finished:
visible += state.stop.finish()
state.finished = finished
if finished and self.metrics:
self.metrics.inc("request_success_total", labels={"finished_reason": reason})
self.metrics.observe("e2e_request_latency_seconds", now - state.arrival)
self.metrics.inc("prompt_tokens_total", out.num_prompt_tokens)
self.metrics.inc("generation_tokens_total", out.num_output_tokens)
state.queue.put_nowait(TextOutput(out.request_id, visible, out.new_token_ids, finished, reason, stop_reason,
out.logprobs, out.prompt_logprobs, out.num_prompt_tokens,
out.num_output_tokens, out.num_cached_tokens))
async def generate(self, prompt_ids, params, request_id=None, priority=0, lora=None, features=None):
"""Async iterator of TextOutputs for one request (n must be 1; the API layer fans out). (Your engine: Chapter 36)
Leaving the loop early (a client disconnect cancels the task) aborts the request in the
core, so its KV blocks are freed at once.
"""
if self.dead:
raise EngineDead(self.dead)
if self.closing:
raise Overloaded("Server is draining")
if len(self.requests) >= self.max_pending:
raise Overloaded(f"{len(self.requests)} requests in flight")
rid = request_id or f"req-{next(self.ids)}"
if rid in self.requests:
raise ValueError(f"Duplicate request id {rid!r}")
state = RequestState(asyncio.Queue(), IncrementalDetokenizer(self.tokenizer, params.skip_special_tokens),
StopChecker(params.stop, params.include_stop_str_in_output))
self.requests[rid] = state
self.inputs.put(("add", (rid, list(prompt_ids), params, priority, {"lora": lora, "features": features})))
deadline = time.monotonic() + self.request_timeout
try:
while True:
item = await asyncio.wait_for(state.queue.get(), max(0, deadline - time.monotonic()))
if isinstance(item, Exception):
raise item
yield item
if item.finished:
return
finally:
if not state.finished:
state.finished = True
self.abort(rid)
self.requests.pop(rid, None)
def abort(self, request_id):
self.inputs.put(("abort", request_id))
@property
def num_in_flight(self):
return len(self.requests)
Three responsibilities live here because they need text:
- Detokenization with Chapter 35’s incremental UTF-8 decoder, one per request.
- Stop strings, which the core can’t see. When the frontend finds one, it truncates the text, marks the request finished with
finish_reason="stop"and the matched string asstop_reason, and sends("abort", rid)to the core so that its blocks are freed at once. - Latency metrics: time to first token (arrival to first output), inter-token latency, end-to-end latency, measured where the client would measure them.
The reader is a daemon thread, not a task that calls queue.get() through the event loop’s executor. That’s a detail that bit this chapter’s first draft: a test failed an assertion before stopping the engine, and the interpreter then hung forever, because asyncio.run waits for its executor’s threads at exit, and one of them was blocked in queue.get(). A server whose shutdown can hang can’t be restarted cleanly; a daemon thread never holds the process.
The API
OpenAI’s API is the de facto standard: every client library, agent framework, evaluation harness and benchmark tool speaks it. Matching it closely is worth more than any feature you could add.
Requests and defaults
The request schemas list OpenAI’s fields plus vLLM’s extensions (top_k, min_p, repetition_penalty, min_tokens, ignore_eos, guided_json, guided_regex, priority). Unknown fields are ignored, because clients send fields like user and parallel_tool_calls that a server may not use:
class SamplingFields(BaseModel):
"""Fields shared by both endpoints. None means "the server's default" (generation_config.json)."""
model_config = ConfigDict(extra="ignore")
model: str | None = None
max_tokens: int | None = None
temperature: float | None = None
top_p: float | None = None
top_k: int | None = None
min_p: float | None = None
repetition_penalty: float | None = None
presence_penalty: float = 0.0
frequency_penalty: float = 0.0
n: int = 1
seed: int | None = None
stop: str | list[str] | None = None
stop_token_ids: list[int] | None = None
include_stop_str_in_output: bool = False
ignore_eos: bool = False
min_tokens: int = 0
logit_bias: dict[str, float] | None = None
skip_special_tokens: bool = True
stream: bool = False
stream_options: dict | None = None
guided_json: dict | str | None = None
guided_regex: str | None = None
priority: int = 0
lora: str | None = None
class CompletionRequest(SamplingFields):
prompt: str | list[str] | list[int] | list[list[int]]
logprobs: int | None = None
prompt_logprobs: int | None = None
echo: bool = False
class ChatRequest(SamplingFields):
messages: list[dict]
tools: list[dict] | None = None
tool_choice: str | dict | None = None
response_format: dict | None = None
logprobs: bool = False
top_logprobs: int | None = None
max_completion_tokens: int | None = None
chat_template_kwargs: dict | None = None
Sampling fields default to None, meaning “the server’s default”, and the server’s defaults come from the checkpoint’s generation_config.json. Qwen3’s instruct models, for example, ship temperature=0.6, top_p=0.95, top_k=20, and its authors recommend against greedy decoding for thinking mode. A request that sets nothing should get what the model’s authors intended, not OpenAI’s default of temperature=1.0.
def sampling_params(self, req, prompt_len, default_max, logprobs=None, prompt_logprobs=None,
guided_regex=None, guided_json=None):
"""Request fields -> SamplingParams, with server defaults and validation. (Your engine: Chapter 36)"""
max_tokens = req.max_tokens if req.max_tokens is not None else default_max
if max_tokens is None:
max_tokens = self.max_model_len - prompt_len
if prompt_len + max_tokens > self.max_model_len:
raise APIError(400, f"This model's maximum context length is {self.max_model_len} tokens; the request has "
f"{prompt_len} prompt tokens and asks for {max_tokens} more.", param="max_tokens")
pick = lambda name: getattr(req, name) if getattr(req, name) is not None else self.defaults[name] # noqa: E731
try:
return SamplingParams(
max_tokens=max_tokens, temperature=pick("temperature"), top_p=pick("top_p"), top_k=pick("top_k"),
min_p=pick("min_p"), repetition_penalty=pick("repetition_penalty"), seed=req.seed, n=req.n,
presence_penalty=req.presence_penalty, frequency_penalty=req.frequency_penalty,
stop=req.stop or (), include_stop_str_in_output=req.include_stop_str_in_output,
stop_token_ids=tuple(req.stop_token_ids or ()) + self.default_stop_token_ids,
ignore_eos=req.ignore_eos, min_tokens=req.min_tokens,
logit_bias={int(k): v for k, v in (req.logit_bias or {}).items()},
logprobs=logprobs, prompt_logprobs=prompt_logprobs, skip_special_tokens=req.skip_special_tokens,
guided_regex=guided_regex or req.guided_regex, guided_json=guided_json or req.guided_json)
except ValueError as error:
raise APIError(400, str(error)) from None
max_tokens defaults to 16 for completions (OpenAI’s legacy default) and to the rest of the context for chat. A request that can’t fit is rejected with a 400 that says why, in OpenAI’s error format: {"error": {"message", "type", "param", "code"}}. Clients parse that format and some retry on it, so every error path uses it, including validation errors from SamplingParams (a negative temperature is the client’s mistake, not a server crash).
The generation config’s eos_token_id list is added to every request’s stop tokens. Qwen3 lists both <|im_end|> (end of turn) and <|endoftext|>; a server that stops only on the tokenizer’s EOS lets the model ramble past the end of its answer.
Completions
async def completions(self, req, http_request):
self.check(http_request, req.model)
self.check_lora(req.lora)
prompt = req.prompt
if isinstance(prompt, str) or (prompt and isinstance(prompt[0], int)):
prompt = [prompt]
prompts = [self.tokenizer.encode(p) if isinstance(p, str) else list(p) for p in prompt]
if not prompts or any(not p for p in prompts):
raise APIError(400, "The prompt is empty", param="prompt")
params = [self.sampling_params(req, len(p), 16, req.logprobs, req.prompt_logprobs) for p in prompts]
rid, created = f"cmpl-{uuid.uuid4().hex}", int(time.time())
self.admit(len(prompts) * params[0].n)
streams = []
for i, (ids, p) in enumerate(zip(prompts, params)):
streams += [self.llm.generate(ids, p.child(j) if p.n > 1 else p, f"{rid}-{i}-{j}", req.priority,
req.lora) for j in range(p.n)]
n = params[0].n
base = {"id": rid, "object": "text_completion", "created": created, "model": self.model_name}
if req.stream:
return self.sse(self.completion_chunks(streams, base, n, prompts, req))
texts, finish, logprobs, usage = [""] * len(streams), [None] * len(streams), [[] for _ in streams], [0, 0, 0]
async for k, out in merge(streams):
texts[k] += out.text
logprobs[k] += out.logprobs or []
if out.finished:
finish[k] = (out.finish_reason, out.stop_reason)
usage[1] += out.num_output_tokens
usage[2] += out.num_cached_tokens
usage[0] = sum(len(p) for p in prompts)
choices = [{"index": k, "text": (req.prompt if req.echo and isinstance(req.prompt, str) else "") + texts[k],
"logprobs": self.completion_logprobs(logprobs[k]) if req.logprobs is not None else None,
"finish_reason": finish[k][0], "stop_reason": finish[k][1]} for k in range(len(streams))]
return {**base, "choices": choices, "usage": self.usage(*usage)}
async def completion_chunks(self, streams, base, n, prompts, req):
completion_tokens = cached = 0
async for k, out in merge(streams):
lp = self.completion_logprobs(out.logprobs) if out.logprobs and req.logprobs is not None else None
choice = {"index": k, "text": out.text, "logprobs": lp,
"finish_reason": out.finish_reason if out.finished else None}
yield {**base, "choices": [choice]}
if out.finished:
completion_tokens += out.num_output_tokens
cached += out.num_cached_tokens
if (req.stream_options or {}).get("include_usage"):
yield {**base, "choices": [], "usage": self.usage(sum(len(p) for p in prompts), completion_tokens, cached)}
prompt can be text, token IDs, or a list of either: a batch. With n samples per prompt, one HTTP request becomes len(prompts) × n engine requests, with choice index prompt × n + sample. merge interleaves their streams as items arrive, so a streaming response sends each choice’s chunks as soon as they exist rather than one choice after another.
Chat completions
def chat_prompt(self, req):
"""Messages -> token ids, plus the constraints the request implies. (Your engine: Chapter 36)"""
messages = []
for m in req.messages:
content = m.get("content")
if isinstance(content, list): # [{"type": "text", "text": ...}, ...]
parts = []
for part in content:
if part.get("type") == "text":
parts.append(part.get("text", ""))
elif part.get("type") == "image_url" and self.multimodal_processor is not None:
parts.append(self.multimodal_processor.marker)
else:
raise APIError(400, "This server does not support that content part", param="messages")
content = "".join(parts)
messages.append({**m, "content": content if content is not None else ""})
tools = req.tools if req.tools and req.tool_choice != "none" else None
guided_regex = guided_json = None
if tools and req.tool_choice == "required":
guided_regex = tool_call_pattern(tools)
elif tools and isinstance(req.tool_choice, dict):
guided_regex = tool_call_pattern(tools, req.tool_choice["function"]["name"])
fmt = req.response_format or {}
if fmt.get("type") == "json_schema":
guided_json = fmt["json_schema"].get("schema", {})
elif fmt.get("type") == "json_object":
guided_regex = json_object_regex()
try:
text = render_chat(self.chat_template, messages, tools=tools, add_generation_prompt=True,
**(req.chat_template_kwargs or {}))
except Exception as error: # noqa: BLE001 (template errors are the client's)
raise APIError(400, f"Chat template error: {error}") from None
return self.tokenizer.encode(text), tools is not None, guided_regex, guided_json
async def chat(self, req, http_request):
self.check(http_request, req.model)
self.check_lora(req.lora)
ids, use_tools, guided_regex, guided_json = self.chat_prompt(req)
images = [part for m in req.messages if isinstance(m.get("content"), list)
for part in m["content"] if part.get("type") == "image_url"]
features = None
if images:
ids, features = await asyncio.to_thread(self.multimodal_processor, ids, images)
if req.max_completion_tokens is not None:
req.max_tokens = req.max_completion_tokens
logprobs = (req.top_logprobs or 0) if req.logprobs else None
params = self.sampling_params(req, len(ids), None, logprobs, None, guided_regex, guided_json)
rid, created = f"chatcmpl-{uuid.uuid4().hex}", int(time.time())
streams = self.streams([ids], params, rid, req.priority, req.lora, features)
parsers = [OutputParser(self.reasoning, use_tools) for _ in streams]
base = {"id": rid, "created": created, "model": self.model_name}
if req.stream:
return self.sse(self.chat_chunks(streams, parsers, base, len(ids), req))
messages = [{"role": "assistant", "content": ""} for _ in streams]
finish, entries, completion, cached = [None] * len(streams), [[] for _ in streams], 0, 0
async for k, out in merge(streams):
for delta in parsers[k].feed(out.text, final=out.finished):
self.accumulate(messages[k], delta)
entries[k] += out.logprobs or []
if out.finished:
finish[k] = "tool_calls" if messages[k].get("tool_calls") else out.finish_reason
completion += out.num_output_tokens
cached += out.num_cached_tokens
for m in messages:
if m.get("tool_calls") and not m["content"]:
m["content"] = None
choices = [{"index": k, "message": messages[k], "finish_reason": finish[k],
"logprobs": self.chat_logprobs(entries[k]) if req.logprobs else None} for k in range(len(streams))]
return {**base, "object": "chat.completion", "choices": choices,
"usage": self.usage(len(ids), completion, cached)}
async def chat_chunks(self, streams, parsers, base, prompt_len, req):
base = {**base, "object": "chat.completion.chunk"}
for k in range(len(streams)):
yield {**base, "choices": [{"index": k, "delta": {"role": "assistant", "content": ""}, "finish_reason": None}]}
completion, cached, called = 0, 0, [False] * len(streams)
async for k, out in merge(streams):
lp = self.chat_logprobs(out.logprobs) if out.logprobs and req.logprobs else None
deltas = parsers[k].feed(out.text, final=out.finished) or ([{}] if lp else [])
for delta in deltas:
called[k] |= "tool_calls" in delta
yield {**base, "choices": [{"index": k, "delta": delta, "logprobs": lp, "finish_reason": None}]}
lp = None
if out.finished:
reason = "tool_calls" if called[k] else out.finish_reason
yield {**base, "choices": [{"index": k, "delta": {}, "finish_reason": reason}]}
completion += out.num_output_tokens
cached += out.num_cached_tokens
if (req.stream_options or {}).get("include_usage"):
yield {**base, "choices": [], "usage": self.usage(prompt_len, completion, cached)}
chat_prompt normalizes messages (OpenAI allows content to be a list of typed parts; Chapter 43 adds images), renders the chat template with the tools, and turns two request fields into grammars with Chapter 34’s machinery:
| request | constraint |
|---|---|
tool_choice: "required" | the output is one of the tools’ call formats, arguments matching each tool’s schema |
tool_choice: {"function": {"name": "get_weather"}} | that tool’s call, with valid arguments |
response_format: {"type": "json_schema", ...} | the schema’s regex |
response_format: {"type": "json_object"} | any JSON object, nested up to three levels |
Without a constraint (tool_choice: "auto"), the model decides whether to call a tool, and the OutputParser of Chapter 35 separates calls, reasoning and content as the text streams in. A choice that contains a tool call finishes with finish_reason: "tool_calls", and its content is null if there was no other text, as OpenAI clients expect.
Streaming
A streamed response is a sequence of server-sent events: data: {json} lines separated by blank lines, ending with data: [DONE]. For chat, the first chunk of each choice carries {"role": "assistant", "content": ""}, later chunks carry content, reasoning_content or tool_calls deltas, a final chunk carries the finish_reason, and if the client asked for stream_options: {"include_usage": true}, one more chunk with empty choices carries the token counts, including prompt_tokens_details.cached_tokens, the prefix-cache hits that some providers bill at a discount.
An error after streaming has started can’t change the HTTP status, which was 200 and is long gone. It’s sent as an error event before [DONE]. That’s why the server does its admission check before returning the streaming response.
Behaviors that make it a service
Cancellation
Users close tabs; agents time out; load balancers drop connections. Each abandoned request still holds KV blocks and still costs a slot in every batch until it reaches max_tokens. A busy server can spend a large fraction of its GPU on output nobody will read.
When a client disconnects during a streamed response, the server’s response task is cancelled. Cancellation propagates into merge, which cancels its pumps, which raises CancelledError inside each AsyncLLM.generate, whose finally block sends ("abort", rid). The core frees the request’s blocks at its next loop. The test checks the whole chain: start a 500-token request, read one chunk, close the stream, and within a few steps the core reports zero running requests and an empty KV pool.
Backpressure
Accepting every request isn’t kind to clients: past some load, queueing delay grows without bound and every request times out. max_pending caps the requests in flight; beyond it the server answers 503, and well-behaved clients back off and retry, or a load balancer sends the request to another replica (Chapter 41). The right cap comes from measurement (Chapter 44): the load at which time to first token breaks its target.
Observability
/metrics serves Prometheus’s text format, with names that mirror vLLM’s (izh: instead of vllm:), so existing dashboards translate directly:
"""Prometheus metrics for the server (Chapter 36), in the text exposition format, without a client
library. Names follow vLLM's (vllm:...) with an izh: prefix, so existing dashboards translate."""
import bisect
import threading
LATENCY_BUCKETS = (0.001, 0.005, 0.01, 0.02, 0.04, 0.06, 0.08, 0.1, 0.25, 0.5, 0.75, 1.0, 2.5, 5.0, 7.5, 10.0,
20.0, 40.0, 80.0)
HELP = {
"time_to_first_token_seconds": "Time from arrival to the first output token.",
"inter_token_latency_seconds": "Time between output tokens of one request.",
"e2e_request_latency_seconds": "Time from arrival to completion.",
"request_success_total": "Finished requests by finish reason.",
"prompt_tokens_total": "Prompt tokens of finished requests.",
"generation_tokens_total": "Output tokens of finished requests.",
"num_requests_running": "Requests in the running batch.",
"num_requests_waiting": "Requests waiting for admission.",
"kv_cache_usage_perc": "Fraction of KV blocks in use.",
"prefix_cache_hit_rate": "Fraction of prompt tokens served by the prefix cache.",
"requests_in_flight": "Requests the server is streaming.",
}
class Metrics:
def __init__(self, prefix="izh:"):
self.prefix, self.lock = prefix, threading.Lock()
self.counters, self.histograms = {}, {}
def inc(self, name, value=1, labels=None):
key = (name, tuple(sorted((labels or {}).items())))
with self.lock:
self.counters[key] = self.counters.get(key, 0) + value
def observe(self, name, value):
with self.lock:
counts, total = self.histograms.setdefault(name, ([0] * (len(LATENCY_BUCKETS) + 1), [0.0, 0]))
counts[bisect.bisect_left(LATENCY_BUCKETS, value)] += 1
total[0] += value
total[1] += 1
def render(self, gauges):
"""The /metrics page: counters, histograms (cumulative buckets) and current gauges."""
lines = []
def header(name, kind):
lines.append(f"# HELP {self.prefix}{name} {HELP.get(name, name)}")
lines.append(f"# TYPE {self.prefix}{name} {kind}")
with self.lock:
for name in sorted({n for n, _ in self.counters}):
header(name, "counter")
for (n, labels), value in sorted(self.counters.items()):
if n == name:
label = ",".join(f'{k}="{v}"' for k, v in labels)
lines.append(f"{self.prefix}{name}{{{label}}} {value}" if label else f"{self.prefix}{name} {value}")
for name, (counts, (total, count)) in sorted(self.histograms.items()):
header(name, "histogram")
running = 0
for bound, c in zip(list(LATENCY_BUCKETS) + ["+Inf"], counts):
running += c
lines.append(f'{self.prefix}{name}_bucket{{le="{bound}"}} {running}')
lines.append(f"{self.prefix}{name}_sum {total}")
lines.append(f"{self.prefix}{name}_count {count}")
for name, value in sorted(gauges.items()):
header(name, "gauge")
lines.append(f"{self.prefix}{name} {value}")
return "\n".join(lines) + "\n"
The ones to alert on: time_to_first_token_seconds p99 (users waiting), num_requests_waiting (load beyond capacity), kv_cache_usage_perc near 1 together with preemptions (memory-bound; Chapter 31), and request_success_total{finished_reason="abort"} (clients giving up).
Security
The API key check is a minimum. The other protections are already in place from earlier chapters: chat templates render in a sandbox (Chapter 35), prefix-cache keys are cryptographic and can carry a per-tenant salt (Chapter 31), and every request is bounded by max_model_len and the server’s limits. Chapter 44 adds request deadlines, raw-body limits, bounded stream buffers and drain/readiness; per-key rate limiting remains deployment work.
Serving a checkpoint
uv pip install -r serve-requirements.txt
python serve.py --model-dir models/Qwen3-0.6B --attention-backend triton --cuda-graphs auto --reasoning
serve.py --help lists every flag; each maps to a field of Chapter 31’s EngineConfig or to this chapter’s server settings. Any OpenAI client works:
from openai import OpenAI
client = OpenAI(base_url="http://localhost:8000/v1", api_key="unused")
reply = client.chat.completions.create(model="Qwen3-0.6B", messages=[{"role": "user", "content": "Hi!"}])
Run it
python run.py serve --requests 32 --slots 32 --new-tokens 32
Without a checkpoint, the demo trains a 1,000-token tokenizer, builds a random 2-layer model with the same vocabulary, starts the server with the engine core in its own process, and acts as a client:
{"models": [{"id": "izh-tiny", "object": "model", "owned_by": "izh", "max_model_len": 2048}]}
{"finish_reason": "tool_calls", "tool_calls": [{"id": "call_35b136318e8944f082ba71c1", "type": "function", "function": {"name": "get_weather", "arguments": "{\"city\": \" depreerဵs\"}"}}], "usage": {"prompt_tokens": 281, "completion_tokens": 68, "total_tokens": 349, "prompt_tokens_details": {"cached_tokens": 0}}}
{"sse_events": 13, "first": "data: {\"id\": \"cmpl-0c6f8af0647d4499850aa6a2112b09e1\", \"object\": \"text_completion\", \"create...", "last": "data: [DONE]"}
{"concurrent_requests": 32, "output_tokens": 1024, "seconds": 0.67, "tok_s": 1529.8, "ttft_p50_ms": 120.4, "ttft_max_ms": 124.1}
{"metrics": ["izh:generation_tokens_total 1104", "izh:request_success_total{finished_reason=\"length\"} 33", "izh:request_success_total{finished_reason=\"stop\"} 1", "izh:time_to_first_token_seconds_count 34"]}
tool_choice="required" forced a well-formed call to get_weather from a model with random weights (its choice of city is as random as its weights). A streamed completion of 12 tokens arrived as 13 events plus [DONE]. Thirty-two concurrent streaming requests through HTTP, the frontend, the queues and the engine process ran at about 1,500 tokens per second on a laptop CPU, every request receiving its first token within 125 ms of the others: they were admitted in the same step.
Build it
Engine milestone 36: the server. Implement core_loop, AsyncLLM._dispatch and AsyncLLM.generate in engine/serve/async_engine.py, and Server.sampling_params and Server.chat_prompt in engine/serve/api.py (the route handlers, streaming, merging, metrics and the launcher are provided).
pytest tests/test_ch36_server.py
python run.py serve --impl engine
python serve.py --model-dir <a Qwen3 checkpoint> # then point any OpenAI client at it
The tests check that greedy completions equal the offline engine’s output, streamed and not, with usage chunks; batch prompts with n=2 and seeded reproducibility; stop strings and logprobs, streamed and not; a forced tool call in chat, streamed and not; JSON-schema response format with chat logprobs; 401, 400 (too long, invalid temperature) and 404 errors in OpenAI’s format and the metrics page; cancellation freeing the core’s blocks and backpressure refusing excess requests; and the engine core in its own process producing the same answer.
Stretch exercises
- ★ Add request deadlines: an
X-Request-Timeoutheader after which the server aborts the request and returns what it has withfinish_reason: "length". Where: header handling inengine/serve/api.py, with timed cancellation viaAsyncLLM.abortinengine/serve/async_engine.py. - ★★ Rate-limit per API key with a token bucket on tokens, not requests (prompt +
max_tokens), and return 429 with aRetry-Afterheader. Where: API-key admission inengine/serve/api.py; add token-bucket state there or inengine/serve/limits.py. - ★★ Graceful shutdown: on SIGTERM, stop accepting requests (503), finish the running ones, then exit. Test it with a request in flight. Where: application lifespan/shutdown in
build_appinengine/serve/api.py, with draining inengine/serve/async_engine.py. - ★★★ Replace the queues with ZeroMQ sockets and msgpack serialization, and measure the frontend-to-core overhead per step at 256 concurrent streams against
multiprocessing.Queue. Where: queue creation and send/receive paths inengine/serve/async_engine.py.
Check your understanding
- Why does running the engine core in its own process improve the decode latency of every request?
- Which failures should fail one request, and which should fail the whole engine? Why?
- Why are stop strings handled by the frontend and not the engine core?
- Trace what happens, step by step, when a client closes a streaming connection.
- Why must the server refuse an overloaded request before it starts streaming the response?
- Why should a server’s default temperature come from the checkpoint rather than from OpenAI’s API defaults?
Going deeper
- vLLM:
vllm/entrypoints/openai/api_server.pyandserving_chat.py(routes and streaming),vllm/v1/engine/async_llm.py,core_client.py(the ZeroMQ link to the engine process) andoutput_processor.py(detokenization and stop strings in the frontend). - SGLang:
python/sglang/srt/managers/tokenizer_manager.py,scheduler.pyanddetokenizer_manager.py, a three-process design. - OpenAI’s API reference for Completions, Chat Completions and streaming; the WHATWG HTML standard’s section on server-sent events; Prometheus’s Exposition formats documentation.
- Beyer et al., Site Reliability Engineering (O’Reilly, 2016), chapters on handling overload and cascading failures, for backpressure and load shedding.
37. Speculative decoding in the server
In this chapter
- Moving speculation from one request (Chapter 26) into the batched engine: drafts become scheduled tokens, verification is one more kind of logits row, and rollback is a change of
num_computed_tokens. - Three drafters with different costs: n-gram prompt lookup, a draft model whose KV pool shares the target's block tables, and Medusa heads that read the target's hidden state.
- Rejection sampling with deterministic drafts, and a statistical test that the output distribution is exactly the target's.
- When speculation pays in a server and when it doesn't: batch size, acceptance and the memory-bound regime.
You will build
NgramDrafter.propose_one, DraftModelDrafter.propose, verify and SpeculativeStep.step in engine/serve/spec.py.
Time: 6-8 hours. GPU: recommended for the speedups (every test runs on a CPU).
From one request to a batch
Chapter 26 showed why speculation works: a decode step reads every weight to produce one token, so checking $k+1$ tokens in one forward pass costs about the same as checking one, and every accepted draft token is a step saved. Its implementation served one request with a contiguous cache, its own loop and its own truncate. A serving engine can’t run a separate loop per request. The engine core of Chapter 31 already has everything speculation needs:
| speculation needs | the engine core has |
|---|---|
| feed the newest token and $k$ drafts | request.spec_token_ids; the scheduler asks for num_tokens_with_spec - num_computed_tokens tokens |
| $k+1$ logits rows per request | build_batch emits rows from the last real token onward (Chapter 31’s logits_indices) |
| blocks for the drafts’ K/V | allocate_slots for those tokens, like any other |
| roll back rejected tokens | set num_computed_tokens to the accepted length, BlockManager.trim the rest |
So speculation in the server is a different second half of EngineCore.step: verify instead of sample, roll back, then draft for the next step.
class SpeculativeStep:
"""Replaces the sample-and-update half of EngineCore.step when speculation is on."""
def __init__(self, drafter):
self.drafter = drafter
self.stats = {"proposed": 0, "accepted": 0, "verify_steps": 0}
def eligible(self, request):
p = request.params
return not (p.needs_penalties or p.logit_bias or p.allowed_token_ids is not None or p.logprobs is not None
or p.prompt_logprobs is not None or "guide" in request.extra
or p.min_tokens > request.num_output_tokens)
def step(self, engine, plan):
"""Verify every request's scheduled drafts, roll back, then draft for the next step. (Your engine: Chapter 37)"""
batch = engine.runner.prepare(plan.scheduled, engine.blocks.req_blocks)
logits = engine.runner.execute(batch)
hidden = engine.runner.last_hidden
if batch.prompt_spans:
engine.record_prompt_logprobs(batch, logits[batch.num_sample_rows:])
outputs, now, row, last_rows = [], time.monotonic(), 0, {}
plain = [(r, k) for (r, _), k in zip(plan.scheduled, batch.sample_counts)
if k and not (r.spec_token_ids and r.num_computed_tokens == r.num_tokens - 1)]
plain_rows = []
for (request, n), k in zip(plan.scheduled, batch.sample_counts):
before = request.num_computed_tokens
request.num_computed_tokens += n
if not k:
engine.blocks.cache_full_blocks(request)
continue
if request.spec_token_ids and before == request.num_tokens - 1: # a decode step with drafts
drafts = request.spec_token_ids[:k - 1]
new, accepted = verify(logits[row:row + k], drafts, request.params, engine.generators.get(request.request_id))
self.stats["proposed"] += len(drafts)
self.stats["accepted"] += accepted
self.stats["verify_steps"] += 1
request.spec_token_ids = []
request.num_computed_tokens = before + 1 + accepted # newest + accepted drafts
engine.blocks.trim(request, request.num_computed_tokens + 1) # free blocks of rejected drafts
last_rows[request.request_id] = row + accepted
self.finish_tokens(engine, request, new, now, outputs)
else:
plain_rows.append(row)
last_rows[request.request_id] = row
row += k
if plain: # rows with nothing to verify: normal sampling
tokens, logprobs = engine.sampler(logits[plain_rows], [r for r, _ in plain], engine.generators)
for (request, _), new, lp in zip(plain, tokens, logprobs or [None] * len(plain)):
self.finish_tokens(engine, request, new, now, outputs, lp)
running = [r for r in engine.scheduler.running if r.request_id in last_rows and not r.status.finished]
engine.spec_hidden = hidden[[last_rows[r.request_id] for r in running]] if hidden is not None and running else None
candidates = [r for r in running if self.eligible(r)]
proposals = self.drafter.propose(engine, candidates) if candidates else {}
for request in candidates:
room = min(request.params.max_tokens - request.num_output_tokens, engine.max_model_len - request.num_tokens) - 1
request.spec_token_ids = proposals.get(request.request_id, [])[:max(room, 0)]
engine.steps += 1
return outputs
def finish_tokens(self, engine, request, new, now, outputs, logprobs=None):
"""Append tokens one at a time so that a stop condition in the middle cuts the rest."""
emitted = []
for token in new:
request.append(token)
emitted.append(token)
status = engine.check_stop(request)
if status is not None:
engine.scheduler.finish(request, status)
engine.generators.pop(request.request_id, None)
break
request.first_token_time = request.first_token_time or now
if not request.status.finished:
engine.blocks.cache_full_blocks(request)
outputs.append(engine.make_output(request, emitted, logprobs))
Requests that ask for features speculation can’t serve exactly are simply not drafted for: penalties, logit bias, allowed tokens and grammars change the target distribution at each position depending on the tokens before it, and logprobs would have to be reported for accepted drafts as well. They decode one token per step in the same batch, sampled by Chapter 34’s sampler, while their neighbors speculate.
Verification with deterministic drafts
Every drafter in this chapter proposes its single best guess. The proposal distribution $q$ is then a point mass on the draft $d$, and Chapter 26’s rule simplifies:
- accept $d$ with probability $\min(1, p(d)/q(d)) = p(d)$;
- on rejection, sample from the residual $\max(p - q, 0)$, which is $p$ with $d$’s entry set to zero, renormalized;
- if all $k$ drafts are accepted, sample one bonus token from the last row.
The total probability of emitting token $t$ at the first position is $p(d)$ for $t = d$, and $(1 - p(d)) \cdot p(t) / (1 - p(d)) = p(t)$ for every other $t$: exactly $p$. Greedy decoding is the special case “accept while the target’s argmax equals the draft”.
def verify(logits, drafts, params, generator=None):
"""logits [len(drafts) + 1, V] from the target; returns (accepted drafts + one new token, accepted). (Your engine: Chapter 37)
Greedy: accept while the target's argmax equals the draft. Sampled: with p the target's
filtered distribution at each row, accept draft d with probability p(d); on the first
rejection, draw from p with d removed. If all are accepted, the last row gives a bonus token.
"""
if params.greedy:
best = logits.argmax(-1).tolist()
out = []
for i, d in enumerate(drafts):
if best[i] != d:
return out + [best[i]], len(out)
out.append(d)
return out + [best[len(drafts)]], len(drafts)
x = logits.float() / params.temperature
if params.top_k or params.top_p < 1 or params.min_p:
n = x.shape[0]
x = top_k_top_p_min_p(x, torch.full((n,), params.top_k or x.shape[-1], device=x.device),
torch.full((n,), params.top_p, device=x.device), torch.full((n,), params.min_p, device=x.device))
p = x.softmax(-1)
out = []
for i, d in enumerate(drafts):
if torch.rand((), generator=generator, device=p.device) < p[i, d]:
out.append(d)
continue
residual = p[i].clone()
residual[d] = 0
if residual.sum() <= 0:
residual = p[i]
return out + [int(torch.multinomial(residual / residual.sum(), 1, generator=generator))], len(out)
return out + [int(torch.multinomial(p[len(drafts)], 1, generator=generator))], len(drafts)
The target’s distribution $p$ includes the request’s temperature and its top-k, top-p and min-p filters, computed with Chapter 34’s batched filter, so a request gets the same distribution with or without speculation. The test checks this statistically: with a perfect drafter at temperature 1, the empirical distribution of the second generated token over 3,000 samples must be within a total variation distance of 0.07 of its exact law, which the test computes by summing over every possible first token. The test also checks that drafts were both accepted and rejected, so both branches were exercised.
Rollback
After verification the request’s token list holds the accepted drafts and one new token. Its K/V cache is valid through the last accepted draft: the newest token (the correction or bonus) hasn’t been fed yet, which is exactly Chapter 31’s invariant “the newest token is uncomputed”. So:
$$ \texttt{num_computed_tokens} = (\text{before}) + 1 + \text{accepted}, $$
and BlockManager.trim releases blocks that only held rejected drafts. Nothing is copied, and nothing about the rejected positions needs to be erased: those slots will be overwritten before any query can see them, by the same causal rule that has protected stale slots since Chapter 19.
Two bookkeeping rules were found by this chapter’s tests. A request that is preempted while holding drafts must drop them (the scheduler clears spec_token_ids in _preempt); otherwise, when it resumes, recomputing its whole history would be mistaken for a verification step. And a request that stops in the middle of its accepted tokens, because one of them is a stop token, must discard the rest: finish_tokens appends one token at a time and checks check_stop after each.
Drafter 1: n-gram prompt lookup
The cheapest drafter needs no model at all. If the last few tokens of the context appeared earlier, guess that what followed them then follows them now:
class NgramDrafter:
"""Prompt-lookup decoding: find the most recent earlier occurrence of the context's last n
tokens (longest n first) and propose the k tokens that followed it. Free to run, and very
effective when the output copies the input: code edits, extraction, RAG answers, chat history.
"""
def __init__(self, k=4, max_n=4, min_n=1):
self.k, self.max_n, self.min_n = k, max_n, min_n
def propose_one(self, tokens, k):
"""(Your engine: Chapter 37)"""
for n in range(min(self.max_n, len(tokens) - 1), self.min_n - 1, -1):
suffix = tokens[-n:]
for start in range(len(tokens) - n - 1, -1, -1): # most recent match first
if tokens[start:start + n] == suffix:
follow = tokens[start + n:start + n + k]
if follow:
return follow
return []
def propose(self, engine, requests):
return {r.request_id: self.propose_one(r.token_ids, self.k) for r in requests}
This “prompt lookup decoding” (Saxena, 2023) costs microseconds and is remarkably effective whenever the output copies the input: editing code, extracting fields, answering questions about a document in the prompt (RAG), continuing a conversation that quotes itself. For free-form generation it rarely finds a match, proposes nothing, and costs nothing. Most engines enable it by default for that reason (vLLM’s ngram method).
Drafter 2: a draft model that shares block tables
A small model from the same family (Qwen3-0.6B drafting for Qwen3-32B) guesses well on all kinds of text. It needs its own K/V cache, and that’s where the paged design pays off again: give the draft model a pool with the same number of blocks as the target’s, and every request’s block table addresses both pools. The block manager allocates once, for two models.
class DraftModelDrafter:
"""A small model (same tokenizer) drafts greedily. Its KV pool has the target's shape in blocks,
so a request's block table addresses both pools: the block manager allocates once for two
models. The draft's own progress is request.extra["draft_computed"]."""
def __init__(self, draft_model, engine, k=4):
self.k = k
self.flat = draft_model if isinstance(draft_model, FlatModel) else FlatModel(draft_model)
runner = engine.runner
self.backend = get_backend(runner.backend.name) if hasattr(runner.backend, "name") else runner.backend
layers, kv_heads, head_dim = self.flat.kv_spec()
shape = (runner.num_blocks + 1, runner.block_size, kv_heads, head_dim)
p = next(self.flat.parameters())
self.caches = [self.backend.allocate(shape, p.dtype, runner.device) for _ in range(layers)]
@torch.inference_mode()
def propose(self, engine, requests):
"""Catch the draft up to each request's newest token, then draft k tokens greedily. (Your engine: Chapter 37)"""
live = []
for r in requests:
done = r.extra.get("draft_computed", 0)
if r.extra.get("draft_epoch") != r.num_preemptions: # its blocks were freed or swapped
done = 0
done = min(done, r.num_computed_tokens)
# Room for the drafts' K/V: positions up to num_tokens + k - 2.
if engine.blocks.allocate_slots(r, self.k):
live.append((r, done))
drafts = {r.request_id: [] for r in requests}
if not live:
return drafts
tables = engine.blocks.req_blocks
views = [_View(r, done) for r, done in live]
scheduled = [(v, v.num_tokens - v.num_computed_tokens) for v in views]
for step in range(self.k):
batch = build_batch(scheduled, tables, engine.runner.block_size, engine.runner.device)
hidden = self.flat(batch.input_ids, batch.positions, self.caches, batch.meta, self.backend)
tokens = self.flat.compute_logits(hidden[batch.logits_indices]).argmax(-1).tolist()
for (view, n), token in zip(scheduled, tokens):
view.num_computed_tokens += n
view.spec_token_ids.append(token)
scheduled = [(v, 1) for v in views]
for (r, _), view in zip(live, views):
drafts[r.request_id] = view.spec_token_ids
r.extra["draft_computed"] = r.num_tokens + self.k - 1 # draft K/V is valid this far, if accepted
r.extra["draft_epoch"] = r.num_preemptions
return drafts
Each step, the drafter first catches up: it runs the draft model over every token the draft hasn’t seen yet (the whole prompt the first time; afterwards, the newest token and any correction), as one flattened batch over all requests. Then it drafts $k$ tokens greedily, one batched forward per token. It tracks its own progress in request.extra["draft_computed"], which is valid only up to the tokens that were actually accepted, and resets when the request is preempted (its blocks were freed or swapped, and only the target’s pool is swapped).
Before drafting, it asks the block manager for room for the drafts’ K/V. If the pool is too full, that request simply gets no drafts this step. Speculation never causes a preemption on its own.
Drafter 3: Medusa heads
The draft model above repeats work the target already did: it re-reads the context and builds its own representation of it. Medusa (Cai et al., 2024) instead adds $k$ small heads to the target. Each reads the target’s final hidden state at the newest position and guesses the token $1, 2, \ldots, k$ places further on, through the target’s own LM head. The heads have no attention and no cache, so drafting is a few matrix-vector products:
class MedusaHeads(nn.Module):
"""k residual heads on the target's last hidden state; head i guesses the token i + 1 places
after the one the target just sampled. They share the target's LM head."""
def __init__(self, hidden, k):
super().__init__()
self.blocks = nn.ModuleList(nn.Linear(hidden, hidden) for _ in range(k))
for block in self.blocks:
nn.init.zeros_(block.weight)
nn.init.zeros_(block.bias)
def forward(self, h): # [R, D] -> [k, R, D]
return torch.stack([h + nn.functional.silu(block(h)) for block in self.blocks])
class MedusaDrafter:
def __init__(self, heads, k=None):
self.heads, self.k = heads, k or len(heads.blocks)
@torch.inference_mode()
def propose(self, engine, requests):
hidden = engine.spec_hidden # [R, D]: each request's newest accepted position
if hidden is None or not requests:
return {r.request_id: [] for r in requests}
logits = engine.runner.flat.compute_logits(self.heads(hidden)[:self.k]) # [k, R, V]
guesses = logits.argmax(-1).T.tolist()
return {r.request_id: g for r, g in zip(requests, guesses)}
The heads must be trained, which is cheap: freeze the target, run it over text (ideally its own outputs, so the heads learn to predict the target), and train each head with cross-entropy against the token $i + 1$ positions ahead. run.py drafters trains three heads for 300 steps on the CPU.
Medusa’s guesses for positions 2 and 3 ignore the tokens in between, which limits its acceptance. EAGLE (Li et al., 2024) and DeepSeek-V3’s multi-token prediction (MTP) modules fix that with one small transformer layer that takes the target’s hidden state and the embedding of the next token, and drafts autoregressively with a cache of its own. These are the most effective drafters in production today, and with this engine they’re a combination of the two drafters above: a draft model whose input embedding is fused with the target’s hidden state (stretch exercise 3).
When speculation pays in a server
Speculation trades compute for memory bandwidth. Verifying $k$ drafts multiplies the step’s tokens, and therefore its FLOPs, by up to $k + 1$, but barely changes its memory traffic. So:
- At small batch sizes, decode is memory-bound, the extra FLOPs are nearly free, and every accepted token is a step saved. Speedups of 2-3× are common for chat with a good drafter.
- At large batch sizes, decode is approaching compute-bound (Chapter 24). Extra verification tokens now cost real time, and rejected drafts are wasted compute that other requests could have used. Speculation can make throughput worse.
Production engines therefore turn speculation down as load rises: fewer drafts per request, or none, when the batch is large (stretch exercise 1). The expected tokens per step for acceptance rate $\alpha$ and $k$ drafts is Chapter 26’s $(1 - \alpha^{k+1})/(1 - \alpha)$; whether that’s worth the extra $k$ tokens of compute depends on where the step sits on the roofline.
Tree verification extends this: instead of one chain of $k$ guesses, verify a small tree of alternatives (Medusa’s and EAGLE-2’s top-2 at each position) in one forward pass, with an attention mask that lets each node see only its ancestors. It accepts more tokens per step at the price of more verification compute, so it shines at batch 1 and fades with load. With the reference backend’s allowed mask, it’s a stretch exercise.
Run it
python run.py drafters --requests 8 --slots 8 --new-tokens 64
Eight requests through the 6-layer test model (random weights) on a laptop CPU, with each drafter:
{"drafter": "none", "steps": 64, "tokens_per_step_per_request": 1.0, "acceptance": null, "tok_s": 448.7}
{"drafter": "n-gram, k=2", "steps": 61, "tokens_per_step_per_request": 1.05, "acceptance": 0.521, "tok_s": 437.3}
{"drafter": "n-gram, k=4", "steps": 61, "tokens_per_step_per_request": 1.05, "acceptance": 0.398, "tok_s": 435.0}
{"drafter": "draft model (first 2 of 6 layers), k=3", "steps": 32, "tokens_per_step_per_request": 2.0, "acceptance": 0.479, "tok_s": 443.6}
{"drafter": "draft model = target, k=3", "steps": 17, "tokens_per_step_per_request": 3.76, "acceptance": 1.0, "tok_s": 450.0}
{"drafter": "none (Medusa prompts)", "steps": 64, "acceptance": null, "mean_accepted_per_step": null}
{"drafter": "Medusa, 3 trained heads", "steps": 34, "acceptance": 0.441, "mean_accepted_per_step": 1.3}
Every configuration produced exactly the tokens of plain decoding (the command asserts it). The n-gram drafter rarely finds a repeat in a random model’s output, but when it does, half its guesses are right. A draft model made of the target’s first two layers, an “early exit” draft, halves the number of steps; the target drafting for itself reaches 3.76 tokens per step, the ceiling for $k = 3$ minus the requests’ final steps. Three Medusa heads, trained for 300 steps on the target’s own outputs, also halve the steps.
The last column is the honest one: on a CPU, tokens per second didn’t move. A CPU forward pass is compute-bound, so verifying four tokens costs about four times as much as decoding one, cancelling the saved steps. On a GPU, where the same decode step is memory-bound, the steps saved translate into time saved; measure it there with this command and compare it with the formula above.
Build it
Engine milestone 37: speculation in the engine. Implement NgramDrafter.propose_one, DraftModelDrafter.propose, verify and SpeculativeStep.step in engine/serve/spec.py (MedusaHeads, MedusaDrafter, finish_tokens and enable_speculation are provided). Turn it on with spec.enable_speculation(engine, drafter).
pytest tests/test_ch37_speculative.py
python run.py drafters --impl engine
The tests check n-gram lookup, greedy verification, the sampled verification law against the target distribution, n-gram speculation equal to plain decoding with and without preemption and prefix caching, a perfect draft model accepting every draft, an early-exit draft model under preemption, Medusa heads, a stop token inside accepted drafts, and the distribution of sampled speculative output against its exact law.
Stretch exercises
- ★★ Make $k$ adaptive: track each request’s recent acceptance rate and the batch size, and choose $k \in {0, 1, 2, 4}$ to maximize expected tokens per unit of step time, using the roofline model of Chapter 10 for the step’s cost. Where: draft-length selection in
SpeculativeStepinengine/serve/spec.py. - ★★ Capture CUDA graphs for verification batches: every decoding request schedules exactly $1 + k$ tokens, so the batch is uniform, and Chapter 33’s buckets apply with $(1 + k) \times$ the rows. Where: verification execution in
engine/serve/spec.py, with graph buckets inengine/serve/graphs.py. - ★★★ Implement an EAGLE-style drafter: one Qwen3 layer whose input is
fc(concat(embed(token), target_hidden)), with its own KV pool sharing the target’s block tables. Train it on the target’s outputs and compare acceptance with Medusa’s. Where: add a drafter inengine/serve/spec.py; expose target hidden rows inengine/serve/model.pyand train it inexperiments/ch37.py(create it). - ★★★ Tree verification: verify a tree of drafts (the top 2 at each of 3 positions) with the reference backend’s
allowedmask, and accept the longest path the target agrees with. Where: tree proposals/verification inengine/serve/spec.py, with new tree-mask metadata inengine/serve/batch.pyand mask handling inengine/serve/attention.py.
Check your understanding
- Which four mechanisms of the Chapter 31 engine make batched speculation a small change?
- Why does a deterministic drafter make the acceptance probability simply $p(d)$, and why does the output still follow $p$ exactly?
- After accepting 2 of 4 drafts, what is the request’s new
num_computed_tokens, and which blocks can be freed? - Why must a preempted request drop its drafts?
- How can a draft model use the target’s block tables without the block manager allocating twice?
- Why can speculation reduce throughput at large batch sizes even with a 70% acceptance rate?
Going deeper
- Leviathan et al., Fast Inference from Transformers via Speculative Decoding (ICML 2023) and Chen et al., Accelerating Large Language Model Decoding with Speculative Sampling (2023), for the rejection rule.
- Cai et al., Medusa: Simple LLM Inference Acceleration Framework with Multiple Decoding Heads (2024); Li et al., EAGLE: Speculative Sampling Requires Rethinking Feature Uncertainty (2024) and EAGLE-2 (2024); DeepSeek-AI, DeepSeek-V3 Technical Report (2024), §2.2 on multi-token prediction.
- Apoorv Saxena, Prompt Lookup Decoding (2023); Liu et al., Optimizing Speculative Decoding for Serving Large Language Models Using Goodput (2024), for adapting speculation to load.
- vLLM’s
vllm/v1/spec_decode/(n-gram, EAGLE, Medusa, MTP proposers) andvllm/v1/sample/rejection_sampler.py.
38. GGUF: running llama.cpp’s models
In this chapter
- The GGUF file format: typed metadata, tensor descriptors and aligned data, in one memory-mappable file.
- llama.cpp's quantization types, from Q8_0 to the K-quants' super-blocks, decoded byte for byte and checked against llama.cpp's own implementation.
- A tokenizer rebuilt from GGUF metadata, and a Qwen3 or Qwen3-MoE model loaded from a
.gguffile. - Serving quantized weights: repacking every format into one group-affine layout and one Triton kernel, and exporting your own models to GGUF.
You will build
GGUFReader.value and load_gguf (engine/formats/gguf.py), scale_min_k4, dequant_q4_k and dequant_q6_k (engine/formats/ggml_quants.py), repack_ggml (engine/formats/runtime.py) and affine_matmul_kernel (engine/kernels/triton_formats.py).
Time: 6-8 hours. GPU: not needed.
Why GGUF matters
Most people who run open models locally never touch a safetensors checkpoint. They download a single .gguf file, often a “Q4_K_M” quantization a quarter the size of the original, and run it with llama.cpp, Ollama or LM Studio. Hugging Face hosts tens of thousands of them, and many new models appear in GGUF within hours of their release. An engine that can’t read GGUF can’t run what most local users have on disk.
Chapter 20 explained quantization from first principles with one simple format of its own. GGUF is what quantization looks like after years of engineering for CPUs and consumer GPUs: a dozen block formats with different bit budgets, and a file format that bundles them with everything else a runtime needs. This chapter reads all of it.
The file
A GGUF file is one header followed by the tensor data:
"GGUF" version (u32 = 3) tensor count (u64) metadata count (u64)
metadata x count: key (u64 length + UTF-8) type (u32) value
tensor info x count: name n_dims (u32) dims (u64 each, innermost first) ggml type (u32) offset (u64)
padding to general.alignment (32 by default)
tensor data, each tensor at its offset, aligned
Metadata values are typed: integers of every width, floats, booleans, strings, and arrays of any of them. That’s where a GGUF file keeps what a Hugging Face checkpoint spreads over config.json, tokenizer.json, tokenizer_config.json and generation_config.json:
| keys | content |
|---|---|
general.architecture, general.name | "qwen3", the model’s name |
qwen3.block_count, qwen3.embedding_length, qwen3.attention.head_count_kv, qwen3.rope.freq_base, … | the architecture’s hyperparameters, prefixed by its name |
tokenizer.ggml.model, .tokens, .merges, .token_type, .eos_token_id, .pre | the whole tokenizer |
tokenizer.chat_template | the Jinja chat template |
class GGUFReader:
"""Parses the header and memory-maps the data; tensors are read lazily, one at a time."""
def __init__(self, path):
self.path = Path(path)
self.file = open(self.path, "rb")
self.data = mmap.mmap(self.file.fileno(), 0, access=mmap.ACCESS_READ)
self.pos = 0
if self.read(4) != MAGIC:
raise ValueError(f"{path} is not a GGUF file")
self.version = self.scalar("<I")
if self.version not in (2, 3):
raise ValueError(f"Unsupported GGUF version {self.version}")
tensor_count, kv_count = self.scalar("<Q"), self.scalar("<Q")
self.metadata = {}
for _ in range(kv_count):
key = self.string()
self.metadata[key] = self.value(self.scalar("<I"))
self.tensors = {}
for _ in range(tensor_count):
name = self.string()
dims = [self.scalar("<Q") for _ in range(self.scalar("<I"))]
kind, offset = self.scalar("<I"), self.scalar("<Q")
self.tensors[name] = (list(reversed(dims)), ggml_quants.NAMES.get(kind, kind), offset)
alignment = self.metadata.get("general.alignment", 32)
self.data_start = -(-self.pos // alignment) * alignment
def read(self, n):
out = self.data[self.pos:self.pos + n]
self.pos += n
return out
def scalar(self, fmt):
(value,) = struct.unpack(fmt, self.read(struct.calcsize(fmt)))
return value
def string(self):
return self.read(self.scalar("<Q")).decode("utf-8", errors="replace")
def value(self, kind):
"""One typed metadata value. (Your engine: Chapter 38)"""
if kind in SCALARS:
return self.scalar(SCALARS[kind])
if kind == STRING:
return self.string()
if kind == ARRAY:
element, count = self.scalar("<I"), self.scalar("<Q")
if element in SCALARS and element != 7: # numeric arrays: one bulk read
dtype = np.dtype(SCALARS[element])
values = np.frombuffer(self.read(count * dtype.itemsize), dtype=dtype)
return values.tolist()
return [self.value(element) for _ in range(count)]
raise ValueError(f"Unknown GGUF value type {kind}")
def raw(self, name):
"""The tensor's bytes, without copying (a view of the memory map)."""
shape, kind, offset = self.tensors[name]
weights, size = ggml_quants.BLOCK[kind]
count = int(np.prod(shape))
nbytes = count // weights * size
start = self.data_start + offset
with warnings.catch_warnings(): # the map is read-only, and so is our use of it
warnings.simplefilter("ignore")
return torch.frombuffer(memoryview(self.data)[start:start + nbytes], dtype=torch.uint8)
def tensor(self, name):
"""Dequantized float32 tensor in PyTorch's [rows, cols] order."""
shape, kind, _ = self.tensors[name]
return ggml_quants.dequantize(self.raw(name), kind, shape)
The reader parses the header and memory-maps the file. A tensor’s bytes are a view into the map, never copied until they’re converted, so opening a 40 GB file is instant and the operating system’s page cache serves repeated loads. Dimensions are stored innermost first (GGML’s ne[0] is the row length), so the reader reverses them into PyTorch’s [rows, cols].
Quantization types
Every quantized GGML type stores a tensor as a sequence of fixed-size blocks, each holding a few scales and the codes of 32 or 256 consecutive weights of a row. A row’s length must be a multiple of the block size, so a writer falls back to another type (usually Q8_0) for tensors whose rows don’t fit.
The legacy types
Q8_0 (34 bytes / 32 weights = 8.5 bits): d:f16 | q:i8 x 32 w = d * q
Q4_0 (18 bytes / 32 weights = 4.5 bits): d:f16 | q:u4 x 32 w = d * (q - 8)
Q4_1 (20 bytes / 32 weights = 5.0 bits): d:f16 | m:f16 | q:u4 x 32 w = d * q + m
Q5_0, Q5_1: a fifth bit per weight in a 32-bit mask
Two 4-bit codes share each byte, with llama.cpp’s own order: byte $j$’s low nibble is weight $j$ and its high nibble is weight $j + 16$, not weight $j + 1$. Getting that wrong produces a perfectly plausible-looking tensor with its columns shuffled, the same trap as RoPE’s pairing in Chapter 17.
def dequant_q8_0(b):
return f16(b[:, :2])[:, None] * b[:, 2:].view(torch.int8).float()
def dequant_q4_0(b):
return f16(b[:, :2])[:, None] * (nibbles(b[:, 2:]).float() - 8)
def dequant_q4_1(b):
return f16(b[:, :2])[:, None] * nibbles(b[:, 4:]).float() + f16(b[:, 2:4])[:, None]
def _high_bits(qh_bytes):
"""A little-endian u32 of 32 flags -> [n, 32] values 0/16: bit j belongs to weight j."""
bits = (qh_bytes[:, :, None] >> torch.arange(8, dtype=torch.uint8)) & 1 # [n, 4, 8]
return bits.reshape(-1, 32).float() * 16
def dequant_q5_0(b):
return f16(b[:, :2])[:, None] * (nibbles(b[:, 6:]).float() + _high_bits(b[:, 2:6]) - 16)
def dequant_q5_1(b):
return f16(b[:, :2])[:, None] * (nibbles(b[:, 8:]).float() + _high_bits(b[:, 4:8])) + f16(b[:, 2:4])[:, None]
K-quants
The K-quants (Q2_K to Q6_K, by Kawrakow, 2023) improved quality at a given size with a two-level scheme. A super-block of 256 weights is split into 8 sub-blocks of 32 (or 16 of 16). Each sub-block has its own scale and minimum, which adapt to its local range, and those scales are themselves quantized, to 6 bits in Q4_K, against one f16 d and dmin per super-block. Q4_K’s scale fields cost 128 bits per 256 weights, 0.5 bits per weight, exactly Q4_0’s overhead; for that price every 32-weight sub-block gets both a scale and a minimum, which Q4_1 needed 5 bits per weight to afford.
Q4_K’s 144-byte block:
d:f16 dmin:f16 scales:12 bytes (8 six-bit scales + 8 six-bit mins) qs:128 bytes (256 four-bit codes)
weight = (d * scale[s]) * q - (dmin * min[s]), s = sub-block 0..7
The twelve scale bytes pack sixteen 6-bit numbers. Sub-blocks 0-3 keep their 6 bits in bytes 0-3 (scales) and 4-7 (mins); sub-blocks 4-7 keep their low 4 bits in bytes 8-11 and their top 2 bits in the otherwise unused top bits of bytes 0-7:
def scale_min_k4(scales):
"""Q4_K / Q5_K: 12 bytes -> 8 six-bit scales and 8 six-bit mins. (Your engine: Chapter 38)
Sub-blocks 0-3 keep their low 6 bits in bytes 0-3 (scales) and 4-7 (mins); sub-blocks 4-7
keep their low 4 bits in bytes 8-11 and their top 2 bits in the spare top bits of bytes 0-7.
"""
s = scales.to(torch.int32)
low_sc, low_m = s[:, 0:4] & 63, s[:, 4:8] & 63
high_sc = (s[:, 8:12] & 0xF) | ((s[:, 0:4] >> 6) << 4)
high_m = (s[:, 8:12] >> 4) | ((s[:, 4:8] >> 6) << 4)
return torch.cat((low_sc, high_sc), 1).float(), torch.cat((low_m, high_m), 1).float()
def dequant_q4_k(b):
"""(Your engine: Chapter 38)"""
d, dmin = f16(b[:, 0:2])[:, None], f16(b[:, 2:4])[:, None]
sc, m = scale_min_k4(b[:, 4:16])
q = b[:, 16:].reshape(-1, 4, 32) # 4 chunks of 32 bytes = 8 sub-blocks of 32
codes = torch.stack((q & 0xF, q >> 4), dim=2).reshape(-1, 8, 32).float() # chunk i: low -> sub 2i, high -> 2i+1
return ((d * sc)[:, :, None] * codes - (dmin * m)[:, :, None]).reshape(-1, QK_K)
def dequant_q5_k(b):
d, dmin = f16(b[:, 0:2])[:, None], f16(b[:, 2:4])[:, None]
sc, m = scale_min_k4(b[:, 4:16])
qh = b[:, 16:48] # bit 2i / 2i+1 of qh[l]: the 5th bit for subs 2i, 2i+1
q = b[:, 48:].reshape(-1, 4, 32)
low = torch.stack((q & 0xF, q >> 4), dim=2).reshape(-1, 8, 32).to(torch.int32)
bit = (qh[:, None, :].to(torch.int32) >> torch.arange(8)[None, :, None]) & 1 # [n, 8, 32]
codes = (low + 16 * bit).float()
return ((d * sc)[:, :, None] * codes - (dmin * m)[:, :, None]).reshape(-1, QK_K)
def q6_k_codes(b):
"""Q6_K's 256 six-bit codes (0..63) in weight order: 4 low bits from ql, 2 high bits from qh."""
ql, qh = b[:, :128].to(torch.int32), b[:, 128:192].to(torch.int32)
out = []
for half in range(2): # two halves of 128 weights
l_lo, l_hi = ql[:, 64 * half:64 * half + 32], ql[:, 64 * half + 32:64 * half + 64]
h = qh[:, 32 * half:32 * half + 32]
out += [(l_lo & 0xF) | (((h >> 0) & 3) << 4), (l_hi & 0xF) | (((h >> 2) & 3) << 4),
(l_lo >> 4) | (((h >> 4) & 3) << 4), (l_hi >> 4) | (((h >> 6) & 3) << 4)]
return torch.cat(out, 1)
def dequant_q6_k(b):
"""(Your engine: Chapter 38)"""
sc, d = b[:, 192:208].view(torch.int8).float(), f16(b[:, 208:210])[:, None]
return d * sc.repeat_interleave(16, 1) * (q6_k_codes(b).float() - 32) # one scale per 16 weights
def dequant_q2_k(b):
sc, q = b[:, :16].to(torch.int32), b[:, 16:80].to(torch.int32)
d, dmin = f16(b[:, 80:82])[:, None], f16(b[:, 82:84])[:, None]
out = []
for half in range(2):
chunk = q[:, 32 * half:32 * half + 32]
for shift in range(4):
codes = ((chunk >> (2 * shift)) & 3).float() # 32 codes: two sub-blocks of 16
for part in range(2):
s = sc[:, 8 * half + 2 * shift + part]
out.append(d * (s & 0xF)[:, None].float() * codes[:, 16 * part:16 * part + 16]
- dmin * (s >> 4)[:, None].float())
return torch.cat(out, 1)
def dequant_q3_k(b):
hmask, q = b[:, :32].to(torch.int32), b[:, 32:96].to(torch.int32)
raw, d = b[:, 96:108].to(torch.int32), f16(b[:, 108:110])[:, None]
# 16 six-bit scales: low 4 bits in bytes 0-7 (two per byte), top 2 bits in bytes 8-11.
low = torch.cat((raw[:, 0:8] & 0xF, raw[:, 0:8] >> 4), 1)
high = torch.cat([(raw[:, 8:12] >> (2 * i)) & 3 for i in range(4)], 1)
scales = (low | (high << 4)).float() - 32
out, bit = [], 0
for half in range(2):
chunk = q[:, 32 * half:32 * half + 32]
for shift in range(4):
codes = ((chunk >> (2 * shift)) & 3) - 4 * (1 - ((hmask >> bit) & 1))
for part in range(2):
out.append(d * scales[:, 8 * half + 2 * shift + part][:, None] * codes[:, 16 * part:16 * part + 16].float())
bit += 1
return torch.cat(out, 1)
Q6_K, which the popular “_M” mixes use for the most sensitive tensors (the output head, some attention and FFN-down layers), splits each 6-bit code into 4 low bits and 2 high bits stored in separate arrays, with an 8-bit signed scale per 16 weights. Q2_K and Q3_K go further down, with 2-bit codes (Q3_K adds a third bit from a mask) and 4-bit or 6-bit sub-block scales.
A name like Q4_K_M isn’t a tensor type. It’s a recipe: mostly Q4_K, with Q6_K for the tensors that hurt most when quantized. A GGUF file can mix types freely, tensor by tensor.
Nothing in this code is checked against prose or memory. The test uses llama.cpp’s own Python package, gguf, as the reference: for each legacy type, gguf.quants.quantize produces blocks, and our decoder must reproduce gguf.quants.dequantize bit for bit; for the K-quants, random blocks (with sane f16 scale fields) must decode identically. Our files must open in llama.cpp’s reader, and its files in ours.
Beyond the K-quants
llama.cpp keeps adding formats. IQ types (“importance” quants, IQ1_S to IQ4_XS) use non-uniform codebooks and lattice codes chosen with an importance matrix from calibration data; IQ4_NL’s 16 values (included here) are a non-uniform table that fits the bell shape of weight distributions better than a uniform grid. MXFP4 is the OCP microscaling format that OpenAI’s gpt-oss models ship in: 32 FP4 (E2M1) values sharing a power-of-two scale (Chapter 39). Each is another block decoder of the same shape.
The tokenizer is in the file
tokenizer.ggml.model = "gpt2" means byte-level BPE (Chapter 35), with the vocabulary in tokens, the merge ranks in merges and each token’s kind in token_type (3 = control, which is special; 4 = user-defined, which is added but ordinary text). One thing is missing: the pre-tokenizer’s regex, which GGUF doesn’t store. Instead, tokenizer.ggml.pre names it ("qwen2", "llama-bpe", …), and every runtime keeps a table from names to patterns, which is why a llama.cpp release is needed for each new model family:
def tokenizer_from_gguf(metadata):
"""Rebuild a byte-level BPE tokenizer.json spec from GGUF metadata. (Your engine: Chapter 38)
token_type 3 (control) tokens are special; type 4 (user-defined) tokens are added but
ordinary text, like <think>.
"""
from ..serve.tokenizer import Tokenizer
model = metadata.get("tokenizer.ggml.model")
if model != "gpt2":
raise ValueError(f"Only byte-level BPE ('gpt2') GGUF tokenizers are implemented, not {model!r}")
pre = metadata.get("tokenizer.ggml.pre", "default")
if pre not in PRE_TOKENIZERS:
raise ValueError(f"Unknown pre-tokenizer {pre!r}; add its regex to PRE_TOKENIZERS")
tokens = metadata["tokenizer.ggml.tokens"]
kinds = metadata.get("tokenizer.ggml.token_type", [1] * len(tokens))
added = [{"id": i, "content": t, "special": k == 3} for i, (t, k) in enumerate(zip(tokens, kinds)) if k in (3, 4)]
spec = {"model": {"type": "BPE", "vocab": {t: i for i, t in enumerate(tokens)},
"merges": metadata.get("tokenizer.ggml.merges", [])},
"pre_tokenizer": {"type": "Sequence", "pretokenizers": [
{"type": "Split", "pattern": {"Regex": PRE_TOKENIZERS[pre]}, "behavior": "Isolated"},
{"type": "ByteLevel", "add_prefix_space": False, "use_regex": False}]},
"normalizer": {"type": "NFC"} if pre == "qwen2" else None,
"decoder": {"type": "ByteLevel"}, "added_tokens": added}
eos = metadata.get("tokenizer.ggml.eos_token_id")
config = {"chat_template": metadata.get("tokenizer.chat_template"),
"eos_token": tokens[eos] if eos is not None else None}
return Tokenizer(spec, config)
The result is the same Tokenizer class as Chapter 35’s, so the server, the detokenizer and constrained decoding all work unchanged on a GGUF model.
Loading the model
GGUF tensor names are short and architecture-neutral: token_embd, blk.{i}.attn_q, blk.{i}.ffn_down, output_norm, output. Loading Qwen3 is a name table plus the hyperparameters from metadata:
@torch.no_grad()
def load_gguf(path, device="cpu", dtype=torch.bfloat16, keep_quantized=True):
"""A Qwen3 or Qwen3-MoE model from a GGUF file. (Your engine: Chapter 38)
Quantized 2-D weights become QuantLinear-style modules that keep their codes (repacked once
into the engine's group-affine layout, formats/runtime.py) unless keep_quantized=False.
Returns (model, tokenizer).
"""
from ..qwen3 import Qwen3, Qwen3Config
from ..moe import Qwen3Moe, Qwen3MoeConfig
from .runtime import AffineQuantLinear, repack_ggml
reader = GGUFReader(path)
raw = config_from_gguf(reader.metadata)
arch = raw["model_type"]
if arch not in ("qwen3", "qwen3_moe", "llama", "qwen2"):
raise ValueError(f"GGUF architecture {arch!r} is not supported")
names = set(reader.tensors)
raw["tie_word_embeddings"] = "output.weight" not in names
if arch in ("llama", "qwen2"): # Chapter 42's registry
from ..models import build_model
with torch.device("meta"):
model = build_model(raw)
cfg = model.cfg
else:
cfg = (Qwen3MoeConfig if arch == "qwen3_moe" else Qwen3Config).from_hf(raw)
with torch.device("meta"):
model = (Qwen3Moe if arch == "qwen3_moe" else Qwen3)(cfg)
model = model.to_empty(device=device).to(dtype)
if cfg.tie_word_embeddings:
model.lm_head.weight = model.model.embed_tokens.weight
params = dict(model.named_parameters())
def put(target, name, heads=None):
"""heads: undo llama.cpp's Q/K row permutation (Llama-architecture files only)."""
shape, kind, _ = reader.tensors[name]
module_name, _, attr = target.rpartition(".")
module = model.get_submodule(module_name)
rows = None if heads is None else unpermute_qk(torch.arange(shape[0])[:, None], heads)[:, 0]
if keep_quantized and kind not in ("F32", "F16", "BF16") and isinstance(module, torch.nn.Linear) and attr == "weight":
codes, scales, offsets, group, packed = repack_ggml(reader.raw(name), kind, shape)
if rows is not None: # quantization is per row: permuting rows is exact
codes, scales, offsets = codes[rows], scales[rows], offsets[rows]
parent, _, child = module_name.rpartition(".")
layer = AffineQuantLinear.from_packed(codes, scales, offsets, group, packed, dtype=dtype)
if module.bias is not None:
layer.bias = module.bias
setattr(model.get_submodule(parent), child, layer.to(device))
return
value = reader.tensor(name).to(device=device, dtype=dtype).reshape(params[target].shape)
params[target].copy_(value if rows is None else value[rows.to(device)])
for name in names:
if name == "token_embd.weight":
put("model.embed_tokens.weight", name)
elif name == "output_norm.weight":
put("model.norm.weight", name)
elif name == "output.weight":
put("lm_head.weight", name)
elif name == "rope_freqs.weight": # Llama 3's frequency divisors
model.inv_freq = model.inv_freq / reader.tensor(name).float()
elif name.startswith("blk."):
_, layer, kind, part = name.split(".", 3)
prefix = f"model.layers.{layer}."
if kind.endswith("_exps"): # MoE experts, stacked [E, out, in]
experts = model.model.layers[int(layer)].mlp.experts
tensor = reader.tensor(name).to(device=device, dtype=dtype)
if kind == "ffn_down_exps":
experts.down_proj.copy_(tensor)
else:
inter = experts.gate_up_proj.shape[1] // 2
part = slice(0, inter) if kind == "ffn_gate_exps" else slice(inter, 2 * inter)
experts.gate_up_proj[:, part].copy_(tensor)
else:
heads = {"attn_q": cfg.num_attention_heads, "attn_k": cfg.num_key_value_heads}.get(kind)
put(prefix + GGUF_NAMES[kind] + "." + part, name, heads if arch == "llama" and part == "weight" else None)
else:
raise ValueError(f"Unexpected GGUF tensor {name}")
return model.eval(), tokenizer_from_gguf(reader.metadata)
MoE layers store their experts stacked: ffn_gate_exps is one 3-D tensor of shape [E, I, D], the same layout Chapter 27 chose for grouped kernels.
One more trap waits in Llama-family GGUFs. llama.cpp’s converter permutes the rows of Llama’s Q and K projections, turning Hugging Face’s rotate-half RoPE pairs $(i, i + d/2)$ into llama.cpp’s interleaved pairs $(2i, 2i+1)$. A Llama GGUF loaded into a rotate-half model must undo it (unpermute_qk; Chapter 42’s registry uses it). Qwen models use rotate-half RoPE in llama.cpp too, so their weights aren’t permuted.
Serving quantized weights
Dequantizing every weight to BF16 at load time works, but throws away the reason for quantizing: a Q4_K model would occupy 3.5 times its file size in GPU memory, and decode, being memory-bound, would read 3.5 times more bytes per token. The weights must stay quantized, and the matmul must decode them on the fly, like Chapter 20’s W4A16 kernel.
Writing one kernel per GGML type (llama.cpp has dozens) is a large job. Look at what the types decode to, though:
| type | group | weight = | code range |
|---|---|---|---|
| Q8_0 | 32 | $d \cdot q$ | −127..127 |
| Q4_0 | 32 | $d \cdot q - 8d$ | 0..15 |
| Q4_1 | 32 | $d \cdot q + m$ | 0..15 |
| Q4_K | 32 | $(d \cdot sc) \cdot q - (d_{min} \cdot m)$ | 0..15 |
| Q5_K | 32 | the same, 5-bit codes | 0..31 |
| Q6_K | 16 | $(d \cdot sc) \cdot q - 32 (d \cdot sc)$ | 0..63 |
Every one is code × scale + offset per group. So are GPTQ and AWQ (Chapter 39). The engine therefore has one quantized layer, AffineQuantLinear, and each format is a function that repacks its blocks into that layout once, at load time:
def repack_ggml(raw, kind, shape):
"""GGML blocks -> (codes, scales, offsets, group size, packed). Exactly the same weights, in
the engine's layout. (Your engine: Chapter 38)"""
rows, cols = shape
weights, size = ggml_quants.BLOCK[kind]
b = raw.reshape(-1, size)
f16 = ggml_quants.f16
if kind == "Q8_0":
codes = b[:, 2:].view(torch.int8).reshape(rows, cols)
scales = f16(b[:, :2]).reshape(rows, -1)
return codes.clone(), scales, torch.zeros_like(scales), 32, False
if kind == "Q4_0":
d = f16(b[:, :2]).reshape(rows, -1)
codes = ggml_quants.nibbles(b[:, 2:]).reshape(rows, cols)
return pack_nibbles(codes), d, -8 * d, 32, True
if kind == "Q4_1":
codes = ggml_quants.nibbles(b[:, 4:]).reshape(rows, cols)
return pack_nibbles(codes), f16(b[:, :2]).reshape(rows, -1), f16(b[:, 2:4]).reshape(rows, -1), 32, True
if kind in ("Q4_K", "Q5_K"):
d, dmin = f16(b[:, 0:2])[:, None], f16(b[:, 2:4])[:, None]
sc, m = ggml_quants.scale_min_k4(b[:, 4:16])
scales, offsets = (d * sc).reshape(rows, -1), (-dmin * m).reshape(rows, -1)
if kind == "Q4_K":
q = b[:, 16:].reshape(-1, 4, 32)
codes = torch.stack((q & 0xF, q >> 4), dim=2).reshape(rows, cols)
return pack_nibbles(codes), scales, offsets, 32, True
qh = b[:, 16:48]
q = b[:, 48:].reshape(-1, 4, 32)
low = torch.stack((q & 0xF, q >> 4), dim=2).reshape(-1, 8, 32).to(torch.int32)
bit = (qh[:, None, :].to(torch.int32) >> torch.arange(8)[None, :, None]) & 1
return (low + 16 * bit).to(torch.int8).reshape(rows, cols), scales, offsets, 32, False
if kind == "Q6_K": # w = d * sc * (q - 32), one scale per 16
scales = (f16(b[:, 208:210])[:, None] * b[:, 192:208].view(torch.int8).float()).reshape(rows, -1)
codes = ggml_quants.q6_k_codes(b).to(torch.int8).reshape(rows, cols)
return codes, scales, -32 * scales, 16, False
raise NotImplementedError(f"No repacking for {kind}; load it dense (keep_quantized=False)")
class AffineQuantLinear(nn.Module):
"""A linear layer whose weight stays as codes + per-group scale and offset."""
def __init__(self, codes, scales, offsets, group_size, packed, perm=None, bias=None, dtype=torch.bfloat16):
super().__init__()
self.register_buffer("codes", codes)
self.register_buffer("scales", scales.float())
self.register_buffer("offsets", offsets.float())
self.register_buffer("perm", perm) # GPTQ act-order: input columns grouped by g_idx
self.bias = None if bias is None else nn.Parameter(bias.to(dtype), requires_grad=False)
self.group_size, self.packed, self.compute_dtype = group_size, packed, dtype
self.out_features = codes.shape[0]
self.in_features = codes.shape[1] * (2 if packed else 1)
@classmethod
def from_packed(cls, codes, scales, offsets, group_size, packed, dtype=torch.bfloat16):
return cls(codes, scales, offsets, group_size, packed, dtype=dtype)
@property
def weight_codes(self):
return unpack_nibbles(self.codes) if self.packed else self.codes
def dequantized_weight(self):
g = self.group_size
return (self.weight_codes.float() * self.scales.repeat_interleave(g, 1)[:, :self.in_features]
+ self.offsets.repeat_interleave(g, 1)[:, :self.in_features])
def forward(self, x):
shape = x.shape
x2 = x.reshape(-1, shape[-1])
if self.perm is not None:
x2 = x2[:, self.perm]
if self.use_kernel(x2):
from ..kernels.triton_formats import affine_matmul
y = affine_matmul(x2, self.codes, self.scales, self.offsets, self.group_size, self.out_features, self.packed)
elif self.use_cpu_kernel(x2):
from ..kernels import cpu
y = cpu.affine_gemm(x2, self.codes, self.scales, self.offsets, self.group_size, self.packed).to(x.dtype)
else:
y = F.linear(x2, self.dequantized_weight().to(x.dtype))
if self.bias is not None:
y = y + self.bias
return y.reshape(*shape[:-1], self.out_features)
def use_kernel(self, x):
from ..kernels import interpreting
return x.is_cuda or (interpreting() and getattr(self, "force_kernel", False))
def use_cpu_kernel(self, x):
"""Decode-sized batches on the CPU, when the SIMD extension is built (Chapter 40)."""
if not CPU_KERNELS["enabled"] or x.device.type != "cpu" or x.shape[0] > 16:
return False
if self.in_features % 32 or self.group_size % 32:
return False
from ..kernels import cpu
return cpu.available()
@property
def storage_bytes(self):
return sum(t.numel() * t.element_size() for t in (self.codes, self.scales, self.offsets))
and one Triton kernel that dequantizes tiles in registers, the same structure as Chapter 20’s W4A16 kernel with an offset and an arbitrary group size:
@triton.jit
def affine_matmul_kernel(x_ptr, codes_ptr, scales_ptr, offsets_ptr, y_ptr, M, N, K, groups, G,
PACKED: tl.constexpr, BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr):
"""(Your engine: Chapter 38)"""
pid_m, pid_n = tl.program_id(0), tl.program_id(1)
rm = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
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)
k_ok = rk < K
x = tl.load(x_ptr + rm[:, None] * K + rk[None, :], mask=(rm[:, None] < M) & k_ok[None, :], other=0.0)
w_ok = (rn[:, None] < N) & k_ok[None, :]
if PACKED: # byte k // 2 holds code k in its low (even k) or high nibble
byte = tl.load(codes_ptr + rn[:, None] * (K // 2) + rk[None, :] // 2, mask=w_ok, other=0).to(tl.int32)
code = (byte >> ((rk[None, :] % 2) * 4)) & 0xF
else:
code = tl.load(codes_ptr + rn[:, None] * K + rk[None, :], mask=w_ok, other=0).to(tl.int32)
group = rn[:, None] * groups + rk[None, :] // G
scale = tl.load(scales_ptr + group, mask=w_ok, other=0.0)
offset = tl.load(offsets_ptr + group, mask=w_ok, other=0.0)
w = code.to(tl.float32) * scale + offset # [BLOCK_N, BLOCK_K], in registers
acc += tl.dot(x.to(tl.float32), tl.trans(w), input_precision="ieee")
tl.store(y_ptr + rm[:, None] * N + rn[None, :], acc.to(y_ptr.dtype.element_ty),
mask=(rm[:, None] < M) & (rn[None, :] < N))
def affine_matmul(x, codes, scales, offsets, group_size, out_features, packed, block_m=16, block_n=32, block_k=32):
check_device(x, codes, scales, offsets)
m, k = x.shape
y = torch.empty((m, out_features), device=x.device, dtype=x.dtype)
grid = (triton.cdiv(m, block_m), triton.cdiv(out_features, block_n))
affine_matmul_kernel[grid](x.contiguous(), codes, scales, offsets, y, m, out_features, k, scales.shape[1], group_size,
PACKED=packed, BLOCK_M=block_m, BLOCK_N=block_n, BLOCK_K=block_k)
return y
The repacking is exact. The test checks that the repacked layer’s weights equal llama.cpp’s dequantization for Q8_0, Q4_0, Q4_1, Q4_K, Q5_K and Q6_K, and that the kernel’s output matches a dense matmul. Exactness has a price in memory: the per-group scales and offsets are stored as FP32 (Q4_K’s d·sc products don’t round-trip through FP16), so Q4_K takes 6 bits per weight resident instead of 4.5, and Q8_0 takes 10 instead of 8.5. A kernel that decodes Q4_K’s super-blocks directly, reading the 12 scale bytes as llama.cpp’s CUDA kernels do, recovers the difference (stretch exercise 2). It’s the right next step for a production engine, and the layout above is the right first one, because it makes every format run, correctly, with one kernel.
Writing GGUF
The writer is the reader in reverse, and with it the engine can export: quantize a model it has loaded, fine-tuned (Chapter 22) or edited (Chapter 23), and hand the file to llama.cpp, Ollama or LM Studio.
def fallback_type(kind, row_length):
"""Rows must hold whole blocks: like llama.cpp, fall back to Q8_0 (or F16) when they can't."""
weights = ggml_quants.BLOCK[kind][0]
if row_length % weights == 0:
return kind
return "Q8_0" if row_length % 32 == 0 else "F16"
def write_gguf(path, metadata, tensors, alignment=32):
"""metadata: {key: python value}; tensors: {name: (float tensor [rows, cols], ggml type name)}."""
def pack_string(s):
data = s.encode("utf-8")
return struct.pack("<Q", len(data)) + data
def pack_value(v):
if isinstance(v, bool):
return struct.pack("<I?", 7, v)
if isinstance(v, int):
return struct.pack("<Iq", 11, v) if v < 0 else struct.pack("<IQ", 10, v) if v >= 2 ** 32 else struct.pack("<II", 4, v)
if isinstance(v, float):
return struct.pack("<If", 6, v)
if isinstance(v, str):
return struct.pack("<I", STRING) + pack_string(v)
if isinstance(v, (list, tuple)):
if all(isinstance(x, str) for x in v):
return struct.pack("<IIQ", ARRAY, STRING, len(v)) + b"".join(pack_string(x) for x in v)
if all(isinstance(x, float) for x in v):
return struct.pack("<IIQ", ARRAY, 6, len(v)) + np.asarray(v, dtype="<f4").tobytes()
return struct.pack("<IIQ", ARRAY, 5, len(v)) + np.asarray(v, dtype="<i4").tobytes()
raise TypeError(f"Can't store {type(v)} in GGUF metadata")
metadata = {"general.alignment": alignment, **metadata}
blobs, infos, offset = [], [], 0
for name, (tensor, kind) in tensors.items():
kind = fallback_type(kind, tensor.shape[-1])
data = ggml_quants.quantize(tensor, kind).numpy().tobytes()
dims = list(reversed(tensor.shape))
infos.append(pack_string(name) + struct.pack("<I", len(dims)) + struct.pack(f"<{len(dims)}Q", *dims)
+ struct.pack("<IQ", ggml_quants.TYPES[kind], offset))
padded = data + b"\0" * (-len(data) % alignment)
blobs.append(padded)
offset += len(padded)
header = MAGIC + struct.pack("<IQQ", 3, len(tensors), len(metadata))
header += b"".join(pack_string(k) + pack_value(v) for k, v in metadata.items()) + b"".join(infos)
header += b"\0" * (-len(header) % alignment)
Path(path).write_bytes(header + b"".join(blobs))
def quantize_q8_0(w):
blocks = w.float().reshape(-1, 32)
d = blocks.abs().amax(1) / 127
q = torch.where(d[:, None] > 0, blocks / d[:, None].clamp_min(1e-30), torch.zeros_like(blocks)).round().clamp(-127, 127)
return torch.cat((_f16_bytes(d), q.to(torch.int8).view(torch.uint8)), 1)
def quantize_q4_k(w):
"""A straightforward Q4_K encoder: per 32-weight sub-block, min/max -> scale and min; the
8 scales and 8 mins are then quantized to 6 bits against the super-block's d and dmin.
(llama.cpp's encoder also searches each sub-block's range to minimize error.)"""
x = w.float().reshape(-1, 8, 32)
lo = x.amin(-1).clamp(max=0) # Q4_K stores w = scale*q - min, min >= 0
hi = x.amax(-1)
sub_scale = (hi - lo) / 15
sub_min = -lo
d = sub_scale.amax(1) / 63
dmin = sub_min.amax(1) / 63
sc = torch.where(d[:, None] > 0, sub_scale / d[:, None].clamp_min(1e-30), torch.zeros_like(sub_scale)).round().clamp(0, 63)
m = torch.where(dmin[:, None] > 0, sub_min / dmin[:, None].clamp_min(1e-30), torch.zeros_like(sub_min)).round().clamp(0, 63)
d_eff = (d.to(torch.float16).float()[:, None] * sc)
m_eff = (dmin.to(torch.float16).float()[:, None] * m)
q = torch.where(d_eff[:, :, None] > 0, (x + m_eff[:, :, None]) / d_eff[:, :, None].clamp_min(1e-30),
torch.zeros_like(x)).round().clamp(0, 15).to(torch.uint8)
sc, m = sc.to(torch.uint8), m.to(torch.uint8)
packed = torch.zeros(x.shape[0], 12, dtype=torch.uint8)
packed[:, 0:4] = (sc[:, 0:4] & 63) | ((sc[:, 4:8] >> 4) << 6)
packed[:, 4:8] = (m[:, 0:4] & 63) | ((m[:, 4:8] >> 4) << 6)
packed[:, 8:12] = (sc[:, 4:8] & 0xF) | ((m[:, 4:8] & 0xF) << 4)
q = q.reshape(-1, 4, 2, 32)
qs = (q[:, :, 0] | (q[:, :, 1] << 4)).reshape(-1, 128)
return torch.cat((_f16_bytes(d), _f16_bytes(dmin), packed, qs), 1)
QUANTIZE = {"Q8_0": quantize_q8_0, "Q4_K": quantize_q4_k}
def quantize(w, type_name):
"""float tensor -> raw GGML bytes (rows must be multiples of the block size)."""
if type_name == "F32":
return w.float().contiguous().view(torch.uint8).reshape(-1)
if type_name in ("F16", "BF16"):
return w.to(torch.float16 if type_name == "F16" else torch.bfloat16).contiguous().view(torch.uint8).reshape(-1)
return QUANTIZE[type_name](w).reshape(-1)
The Q4_K encoder here is the straightforward one: each sub-block’s range from its min and max, then the 6-bit quantization of scales and mins. llama.cpp’s encoder also searches each sub-block’s range to minimize squared error (and, with an importance matrix, weighted error), which buys a few percent of quality. The test checks that llama.cpp’s decoder reads our blocks exactly as ours does.
Run it
python run.py gguf # export the test model and reload it
python run.py gguf --model-dir model.Q4_K_M.gguf # or run a real GGUF file
{"type": "Q8_0", "file_bits_per_param": 8.54, "resident_bits_per_quantized_weight": 10.0, "relative_logit_error": 0.0071, "same_greedy_next_token": true}
{"type": "Q4_K", "file_bits_per_param": 4.67, "resident_bits_per_quantized_weight": 6.0, "relative_logit_error": 0.0697, "same_greedy_next_token": true}
{"served_from_q4_k": [220, 296, 51, 500, 59, 328, 150, 402]}
The 4-layer test model (random weights, rows of 256) exported as Q8_0 is 8.54 bits per parameter on disk and changes its logits by 0.7%; as Q4_K, 4.67 bits and 7% (random weights have no structure for quantization to exploit, so real models fare better). The Q4_K model is then served by the engine core of Chapter 31, its linear layers running from codes through AffineQuantLinear.
Build it
Engine milestone 38: GGUF. Implement GGUFReader.value and load_gguf in engine/formats/gguf.py; scale_min_k4, dequant_q4_k and dequant_q6_k in engine/formats/ggml_quants.py; repack_ggml in engine/formats/runtime.py; and affine_matmul_kernel in engine/kernels/triton_formats.py (the other block decoders, the writer, the encoders, the tokenizer rebuild and the export are provided).
pytest tests/test_ch38_gguf.py
python run.py gguf --impl engine
The tests check every legacy type and every K-quant against llama.cpp’s gguf package, our encoders’ blocks in llama.cpp’s decoder, our files in llama.cpp’s reader and its files in ours, exact repacking with the kernel agreeing for six types, the inverse of llama.cpp’s Q/K permutation, a Qwen3 GGUF export and reload in Q8_0 and Q4_K (quantized and dense paths agreeing, logits close to the original, the tokenizer round-tripping), and a Qwen3-MoE GGUF round trip.
Stretch exercises
- ★ Load a real Qwen3 GGUF (for example Qwen3-0.6B Q8_0 and Q4_K_M) with
run.py gguf --model-dir, and compare its next-token probabilities on a few prompts with the BF16 safetensors checkpoint’s: KL divergence per token, by quantization type. Where:experiments/ch38.py(create it), usingengine.formats.gguf.load_ggufandengine.evaluation.logit_quality. - ★★★ Write a Triton kernel that reads Q4_K super-blocks directly: per (row, super-block), load
d,dminand the 12 scale bytes, unpack the 8 scales and mins in registers, and decode codes fromqs. Compare resident memory and decode bandwidth withAffineQuantLinear. Where: add a direct Q4_K kernel/launcher inengine/kernels/triton_formats.py; dispatch it fromengine/formats/runtime.py. - ★★ Add the SentencePiece-style GGUF tokenizer (
tokenizer.ggml.model = "llama"): merges come from token scores, highest first, with byte fallback. Where:tokenizer_from_ggufinengine/formats/gguf.py, with tokenization support inengine/serve/tokenizer.py. - ★★ Implement llama.cpp’s importance-weighted Q4_K encoder: given per-column activation statistics from calibration data, choose each sub-block’s scale and min to minimize the importance-weighted squared error. Measure KL against the plain encoder on a real model. Where:
quantize_q4_kinengine/formats/ggml_quants.py.
Check your understanding
- What does a GGUF file contain that a
model.safetensorsfile doesn’t? - Why does the reader memory-map the file rather than read it?
- In Q4_0, which weights share a byte? What would a decoder that paired weights $j$ and $j+1$ produce?
- Why do K-quants quantize their sub-block scales, and what does that buy over Q4_0’s one f16 scale per 32 weights?
- What is Q4_K_M, as opposed to Q4_K?
- Why can one kernel serve GGUF, GPTQ and AWQ weights, and what does the shared layout cost for Q4_K?
- Why must a Llama GGUF’s Q and K rows be un-permuted for a rotate-half model, but a Qwen GGUF’s not?
Going deeper
- The GGUF specification (
docs/gguf.mdin theggmlrepository) and llama.cpp’sgguf-pypackage, used as this chapter’s reference. - llama.cpp’s
ggml/src/ggml-quants.c(dequantize_row_q4_Kand friends) andggml-common.h(the block structs); its CUDA dequantization kernels inggml/src/ggml-cuda/. - Kawrakow’s pull requests introducing the K-quants (llama.cpp #1684, 2023) and the IQ types, with their perplexity tables.
convert_hf_to_gguf.py, for how each architecture’s tensors and metadata are named.
39. Quantized checkpoints: GPTQ, AWQ, FP8, FP4 and trellis codes
In this chapter
- GPTQ: quantizing a layer column by column while correcting the remaining columns for each error, using the inverse Hessian of the layer's output error.
- AWQ: scaling up the weights that large activations multiply before quantizing, and folding the scale into the previous operation.
- The checkpoint formats behind
"quant_method": "gptq" | "awq" | "fp8", read and written, and served through the same layer as GGUF. - Low-precision floating point: block-scaled FP8, and the FP4 microscaling formats MXFP4 and NVFP4.
- Trellis-coded quantization (QTIP, ExLlamaV3's EXL3): coding whole sequences of weights with a Viterbi search, and why it beats any scalar code at 2 bits.
You will build
gptq_quantize and unpack_gptq (engine/formats/gptq.py), search_scale, apply_awq and unpack_awq (awq.py), quantize_blocks (fp8.py), quantize_mxfp4 (mx.py), viterbi_encode (trellis.py), load_quantized (hf_quant.py) and the FP8 kernel.
Time: 8-10 hours. GPU: not needed.
Beyond round-to-nearest
Chapter 20 quantized each weight to the nearest point of a grid, with one scale per group, and noted that the methods people actually ship do better. This chapter builds them. They answer three different questions:
| question | methods |
|---|---|
| Given the grid, which grid point should each weight get? Not always the nearest. | GPTQ |
| Can we change the problem so that rounding hurts less? | AWQ (scale salient channels), QuIP#/QTIP (rotate), SmoothQuant |
| Is a uniform integer grid even the right code? | FP8, MXFP4, NVFP4 (floating-point grids); trellis codes (code sequences, not values) |
And for each, a checkpoint format that the engine must read, because users download quantized checkpoints far more often than they quantize models themselves. All of them, like GGUF, repack into Chapter 38’s AffineQuantLinear (code × scale + offset per group), except the floating-point formats, which get their own layer.
Throughout, the measure of damage is the output, not the weights: Chapter 20’s choose_clip already minimized output error on calibration activations, and that idea is the seed of everything here.
GPTQ
Quantize a layer’s weight $W$ (rows are outputs, columns are inputs) to minimize the squared error of its output on calibration inputs $X$:
$$ \min_{\hat W} ; \lVert X \hat W^\top - X W^\top \rVert^2 . $$
The error is a quadratic function of each row of $\hat W - W$, with the same Hessian for every row: $H = 2 X^\top X$. Its diagonal says how strongly each input is activated; its off-diagonal entries say which inputs move together. That second part is what round-to-nearest ignores. If inputs $i$ and $j$ are correlated, an error on weight $i$ can be partly cancelled by adjusting weight $j$.
GPTQ (Frantar et al., 2022), built on Optimal Brain Quantization, quantizes the columns one at a time and, after each, updates every column not yet quantized to absorb the error:
$$ e = \frac{w_i - \text{quant}(w_i)}{[H^{-1}]{ii}}, \qquad w{j} \leftarrow w_j - e , [H^{-1}]_{ij} \quad \text{for every later } j . $$
The inverse Hessian is computed once, through a Cholesky factorization (with a little damping on the diagonal for stability), and the updates are applied lazily: inside a block of 128 columns each error updates the block, and the rest of the matrix is updated with one matmul per block. That makes GPTQ fast enough to quantize a 70-billion-parameter model on one GPU in a few hours.
@torch.no_grad()
def gptq_quantize(weight, hessian, bits=4, group_size=128, act_order=False, damp=0.01, block=128):
"""weight [out, in], hessian [in, in] = 2 X^T X / n -> (codes [out, in], scales, zeros [out, groups], g_idx [in]).
(Your engine: Chapter 39)
Columns are processed in blocks of `block`; inside a block every column's error is pushed
onto the block's remaining columns at once, and onto the columns after the block in one
matmul when the block is done ("lazy batch updates"). With act_order, columns are visited
from the largest Hessian diagonal (the most activated input) down.
"""
W = weight.float().clone()
H = hessian.float().clone()
cols = W.shape[1]
dead = torch.diag(H) == 0 # inputs that are never active: their weights don't matter
H[dead, dead] = 1
W[:, dead] = 0
perm = torch.argsort(torch.diag(H), descending=True) if act_order else torch.arange(cols)
W, H = W[:, perm], H[perm][:, perm]
H += damp * torch.diag(H).mean() * torch.eye(cols)
Hinv = torch.linalg.cholesky(torch.cholesky_inverse(torch.linalg.cholesky(H)), upper=True)
codes = torch.zeros_like(W)
groups = -(-cols // group_size)
scales, zeros = torch.zeros(W.shape[0], groups), torch.zeros(W.shape[0], groups)
top = 2 ** bits - 1
for i1 in range(0, cols, block):
i2 = min(i1 + block, cols)
W1, Err = W[:, i1:i2].clone(), torch.zeros(W.shape[0], i2 - i1)
H1 = Hinv[i1:i2, i1:i2]
for i in range(i2 - i1):
col = i1 + i
if col % group_size == 0: # a new group: its range from the CURRENT (corrected) weights
g = col // group_size
ahead = torch.cat((W1[:, i:], W[:, i2:]), 1)[:, :group_size]
scales[:, g], zeros[:, g] = find_params(ahead, bits)
g = col // group_size
q = (W1[:, i] / scales[:, g] + zeros[:, g]).round().clamp(0, top)
codes[:, col] = q
error = (W1[:, i] - scales[:, g] * (q - zeros[:, g])) / H1[i, i]
W1[:, i:] -= error[:, None] * H1[i, i:][None, :] # correct the rest of the block
Err[:, i] = error
W[:, i2:] -= Err @ Hinv[i1:i2, i2:] # and everything after it
inverse = torch.argsort(perm)
g_idx = (torch.arange(cols) // group_size)[inverse] # group of each ORIGINAL column
return codes[:, inverse].to(torch.uint8), scales, zeros, g_idx.to(torch.int32)
Two details matter in practice:
- Group ranges come from the corrected weights. When the column loop enters a new group of 128, the group’s scale and zero point are computed from the weights as GPTQ has already adjusted them, not from the original ones.
- Act-order (
desc_actin configs) visits columns from the most activated (largest $H_{ii}$) to the least, so the important columns are quantized while the most freedom remains to compensate them. It helps quality, at a cost the format must carry: groups no longer cover contiguous columns, so the checkpoint storesg_idx, each input column’s group.
The format
qweight int32 [in / 8, out] eight 4-bit codes along the INPUT dimension per int32, low bits first
qzeros int32 [groups, out / 8] zero points along the OUTPUT dimension, stored minus one
scales f16 [groups, out]
g_idx int32 [in]
W[o, i] = scales[g_idx[i], o] * (q[i, o] - (qzeros[g_idx[i], o] + 1))
The “stored minus one” is a historical quirk of AutoGPTQ’s first packing code that every reader must replicate (newer tools call the corrected layout gptq_v2). To serve it, to_runtime sorts the input columns by g_idx once at load time, so that groups become contiguous, and permutes the activations to match at run time. ExLlama’s kernels do the same, permuting the previous layer’s output instead, which is free.
def pack_gptq(codes, scales, zeros, g_idx):
"""Our [out, in] codes -> the checkpoint tensors (4-bit)."""
q = codes.T.to(torch.int64) # [in, out]
qweight = torch.zeros(q.shape[0] // 8, q.shape[1], dtype=torch.int64)
for j in range(8):
qweight |= q[j::8] << (4 * j)
z = (zeros.T.to(torch.int64) - 1) & 0xF # [groups, out], stored minus 1
qzeros = torch.zeros(z.shape[0], z.shape[1] // 8, dtype=torch.int64)
for j in range(8):
qzeros |= z[:, j::8] << (4 * j)
return {"qweight": as_int32(qweight), "qzeros": as_int32(qzeros), "scales": scales.T.to(torch.float16).contiguous(),
"g_idx": g_idx.to(torch.int32)}
def as_int32(t):
"""32 packed bits held in an int64 -> int32 with the same bits (two's complement)."""
return ((t + 2 ** 31) % 2 ** 32 - 2 ** 31).to(torch.int32)
def unpack_gptq(qweight, qzeros, scales, g_idx):
"""Checkpoint tensors -> (codes [out, in], scales [out, groups], zeros [out, groups], g_idx). (Your engine: Chapter 39)"""
shifts = torch.arange(0, 32, 4, dtype=torch.int32)
q = (qweight[:, None, :] >> shifts[None, :, None]) & 0xF # [in/8, 8, out]
codes = q.reshape(-1, qweight.shape[1]).T # [out, in]
z = (qzeros[:, :, None] >> shifts[None, None, :]) & 0xF # [groups, out/8, 8]
zeros = (z.reshape(qzeros.shape[0], -1) + 1).T.float() # stored minus 1
return codes.to(torch.uint8), scales.T.float(), zeros, g_idx
AWQ
GPTQ corrects errors after the fact. AWQ (Lin et al., 2023) prevents some of them. Its observation: about 1% of a layer’s input channels carry activations an order of magnitude larger than the rest (Chapter 20’s outliers), and errors in the weights those channels multiply dominate the output error. Scale those weights up by $s > 1$ before quantizing: their quantization step stays the same in absolute terms (the group’s range barely changes), so their relative error shrinks. Then divide the activations by $s$, which leaves the product exactly unchanged before rounding:
$$ X W^\top = (X ,\text{diag}(s)^{-1}) , (W ,\text{diag}(s))^\top . $$
The division is free, because it folds into whatever produced $X$: an RMSNorm’s weight (for Q/K/V and gate/up, which read the norms’ outputs), or the previous linear layer’s output rows (V for O, up for down). The scale is $s = \overline{|x|}^{,\alpha}$ per input channel, with $\alpha$ chosen per group of layers by a 20-point grid search on calibration data. Since $\alpha = 0$ is plain round-to-nearest, the search can only help:
@torch.no_grad()
def search_scale(weights, inputs, bits=4, group_size=128, grid=20):
"""Best per-input-channel scale for the linears in `weights` (all reading `inputs`). (Your engine: Chapter 39)
For alpha in 0, 1/grid, ..., (grid-1)/grid: s = mean|x|^alpha (normalized), quantize W * s,
undo the scale, and measure the output error on the calibration inputs. alpha = 0 is plain
round-to-nearest, so the search can only help.
"""
x = inputs.reshape(-1, inputs.shape[-1]).float()
x_mean = x.abs().mean(0)
reference = [x @ w.float().T for w in weights]
best = (float("inf"), torch.ones_like(x_mean), 0.0)
for step in range(grid):
alpha = step / grid
s = x_mean.clamp_min(1e-4) ** alpha
s = s / (s.max() * s.min()).sqrt()
error = 0.0
for w, ref in zip(weights, reference):
approx = (x / s) @ fake_quantize(w.float() * s, bits, group_size).T
error += (approx - ref).pow(2).mean().item()
if error < best[0]:
best = (error, s, alpha)
return best[1], best[2]
@torch.no_grad()
def apply_awq(model, calibration_ids, bits=4, group_size=128):
"""Scale and quantize every Qwen3 layer in place, AutoAWQ's way. (Your engine: Chapter 39)
Scaling pairs (previous op -> linears that read its output):
input_layernorm -> q, k, v post_attention_layernorm -> gate, up
v_proj (rows) -> o_proj up_proj (rows) -> down_proj
Each layer's inputs are captured from a forward pass over the calibration tokens.
Returns {linear name: AffineQuantLinear} after replacing the modules.
"""
captured = {}
def hook(name):
def save(module, args, output):
captured.setdefault(name, []).append(args[0].detach().reshape(-1, args[0].shape[-1]))
return save
handles = []
for i, layer in enumerate(model.model.layers):
for name in ("self_attn.q_proj", "self_attn.o_proj", "mlp.gate_proj", "mlp.down_proj"):
handles.append(layer.get_submodule(name).register_forward_hook(hook(f"{i}.{name}")))
model(calibration_ids)
for h in handles:
h.remove()
replaced = {}
for i, layer in enumerate(model.model.layers):
attn, mlp = layer.self_attn, layer.mlp
x = lambda name: torch.cat(captured[f"{i}.{name}"]) # noqa: E731
pairs = [(layer.input_layernorm, None, [attn.q_proj, attn.k_proj, attn.v_proj], x("self_attn.q_proj")),
(layer.post_attention_layernorm, None, [mlp.gate_proj, mlp.up_proj], x("mlp.gate_proj")),
(None, attn.v_proj, [attn.o_proj], x("self_attn.o_proj")),
(None, mlp.up_proj, [mlp.down_proj], x("mlp.down_proj"))]
for norm, previous, linears, inputs in pairs:
if previous is attn.v_proj and attn.v_proj.out_features != attn.o_proj.in_features:
continue # GQA: v's rows feed several heads' columns
s, _ = search_scale([l.weight for l in linears], inputs, bits, group_size)
if norm is not None:
norm.weight.div_(s.to(norm.weight.dtype))
else:
previous.weight.div_(s[:, None].to(previous.weight.dtype))
for linear in linears:
linear.weight.mul_(s[None, :].to(linear.weight.dtype))
for i, layer in enumerate(model.model.layers):
for name in ("self_attn.q_proj", "self_attn.k_proj", "self_attn.v_proj", "self_attn.o_proj",
"mlp.gate_proj", "mlp.up_proj", "mlp.down_proj"):
parent_name, _, child = name.rpartition(".")
parent = layer.get_submodule(parent_name)
linear = getattr(parent, child)
codes, scale, offset = quantize_to_affine(linear.weight, bits, group_size)
quantized = AffineQuantLinear(pack_nibbles(codes), scale, offset, group_size, True,
dtype=linear.weight.dtype)
setattr(parent, child, quantized)
replaced[f"model.layers.{i}.{name}"] = quantized
return replaced
AutoAWQ’s checkpoint packs eight codes along the output dimension in an interleaved order, 0 2 4 6 1 3 5 7, chosen to match its CUDA kernel’s register layout. That’s the kind of detail a format reader can only get from the source, and only verify with a test that pins it down (output 1’s code sits in bits 16-19):
def pack_awq(codes, zeros):
"""[rows, cols] codes 0..15 -> int32 [rows, cols / 8] in AWQ's interleaved order."""
c = codes.to(torch.int64)
packed = torch.zeros(c.shape[0], c.shape[1] // 8, dtype=torch.int64)
for j, o in enumerate(AWQ_ORDER):
packed |= c[:, o::8] << (4 * j)
return ((packed + 2 ** 31) % 2 ** 32 - 2 ** 31).to(torch.int32)
def unpack_awq(packed):
"""(Your engine: Chapter 39)"""
out = torch.zeros(packed.shape[0], packed.shape[1] * 8, dtype=torch.uint8)
for j, o in enumerate(AWQ_ORDER):
out[:, o::8] = ((packed >> (4 * j)) & 0xF).to(torch.uint8)
return out
def to_checkpoint(layer):
"""AffineQuantLinear (from apply_awq) -> AutoAWQ GEMM tensors."""
codes = layer.weight_codes.T # [in, out]
scales = layer.scales.T # [groups, out]
zeros = (-layer.offsets / layer.scales).round().T # offset = -zero * scale
return {"qweight": pack_awq(codes, None), "qzeros": pack_awq(zeros.clamp(0, 15), None),
"scales": scales.to(torch.float16).contiguous()}
def from_checkpoint(qweight, qzeros, scales, group_size, bias=None, dtype=torch.bfloat16):
codes = unpack_awq(qweight).T # [out, in]
zeros = unpack_awq(qzeros).T.float() # [out, groups]
s = scales.T.float()
return AffineQuantLinear(pack_nibbles(codes), s, -zeros * s, group_size, True, bias=bias, dtype=dtype)
FP8 checkpoints
FP8 E4M3 has 3 mantissa bits and 4 exponent bits: 448 is its largest value, and its grid is relative: fine near zero, coarse near the top. With one scale per block that keeps each block’s largest value near 448, it reaches nearly BF16 quality at half the memory. DeepSeek-V3 was trained in FP8, and Qwen3 ships FP8 checkpoints, with one FP32 scale per 128 × 128 weight block (weight_scale_inv, which despite its name multiplies the stored values):
def quantize_blocks(weight, block=128):
"""[out, in] -> (fp8 weight, scale_inv [ceil(out/b), ceil(in/b)]) with one scale per block. (Your engine: Chapter 39)"""
out_f, in_f = weight.shape
rows, cols = -(-out_f // block), -(-in_f // block)
padded = F.pad(weight.float(), (0, cols * block - in_f, 0, rows * block - out_f))
tiles = padded.reshape(rows, block, cols, block)
amax = tiles.abs().amax(dim=(1, 3)).clamp_min(1e-12)
scale = amax / E4M3_MAX
q = (tiles / scale[:, None, :, None]).clamp(-E4M3_MAX, E4M3_MAX).to(torch.float8_e4m3fn)
return q.reshape(rows * block, cols * block)[:out_f, :in_f].contiguous(), scale
def dequantize_blocks(weight, scale_inv, block=128):
s = scale_inv.repeat_interleave(block, 0)[:weight.shape[0]].repeat_interleave(block, 1)[:, :weight.shape[1]]
return weight.float() * s
On Hopper and newer GPUs the matmul can run in FP8 on both sides: activations are quantized on the fly with one scale per token per 128 input channels, so each 128-wide slice of the product has one activation scale and one weight scale. w8a8_reference computes exactly what such a kernel computes, so you can measure its error on a CPU:
def quantize_activations(x, group=128):
"""Per token, per 128 input channels: the FP8 activations a W8A8 GEMM consumes."""
n, k = x.shape
g = x.float().reshape(n, k // group, group)
scale = g.abs().amax(-1).clamp_min(1e-12) / E4M3_MAX
return (g / scale[..., None]).to(torch.float8_e4m3fn).reshape(n, k), scale
def w8a8_reference(x, weight, scale_inv, block=128):
"""What an FP8 x FP8 tensor-core GEMM computes: both operands rounded to E4M3, products
accumulated in FP32, each 128-wide slice of K rescaled by its activation and weight scales."""
xq, xs = quantize_activations(x, block)
out = torch.zeros(x.shape[0], weight.shape[0])
for kb in range(x.shape[1] // block):
a = xq[:, kb * block:(kb + 1) * block].float() * xs[:, kb:kb + 1]
w = weight[:, kb * block:(kb + 1) * block].float() * scale_inv[:, kb].repeat_interleave(block)[:weight.shape[0], None]
out += a @ w.T
return out
The engine’s FP8 kernel is the portable W8A16 variant, which widens FP8 tiles to FP32 in registers and works on every GPU (and in the Triton interpreter):
@triton.jit
def fp8_block_kernel(x_ptr, w_ptr, s_ptr, y_ptr, M, N, K, s_cols,
BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr):
"""W8A16: FP8 weight tiles are widened in registers and scaled by their block's factor. (Your engine: Chapter 39)"""
pid_m, pid_n = tl.program_id(0), tl.program_id(1)
rm = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
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)
x = tl.load(x_ptr + rm[:, None] * K + rk[None, :], mask=(rm[:, None] < M) & (rk[None, :] < K), other=0.0)
w_ok = (rn[:, None] < N) & (rk[None, :] < K)
w = tl.load(w_ptr + rn[:, None] * K + rk[None, :], mask=w_ok, other=0.0).to(tl.float32)
s = tl.load(s_ptr + (rn[:, None] // 128) * s_cols + rk[None, :] // 128, mask=w_ok, other=0.0)
acc += tl.dot(x.to(tl.float32), tl.trans(w * s), input_precision="ieee")
tl.store(y_ptr + rm[:, None] * N + rn[None, :], acc.to(y_ptr.dtype.element_ty),
mask=(rm[:, None] < M) & (rn[None, :] < N))
def fp8_block_matmul(x, weight, scale_inv, block_m=16, block_n=32, block_k=32):
check_device(x, weight, scale_inv)
m, k = x.shape
n = weight.shape[0]
y = torch.empty((m, n), device=x.device, dtype=x.dtype)
grid = (triton.cdiv(m, block_m), triton.cdiv(n, block_n))
fp8_block_kernel[grid](x.contiguous(), weight.contiguous(), scale_inv.float().contiguous(), y, m, n, k,
scale_inv.shape[1], BLOCK_M=block_m, BLOCK_N=block_n, BLOCK_K=block_k)
return y
FP4: MXFP4 and NVFP4
Blackwell GPUs multiply 4-bit floats natively. E2M1, with 2 exponent bits and 1 mantissa bit, has only 16 values: $0, 0.5, 1, 1.5, 2, 3, 4, 6$ and their negatives. With so few, the scales do most of the work, and two formats differ exactly there:
| format | block | scale | bits / weight | used by |
|---|---|---|---|---|
| MXFP4 (OCP Microscaling) | 32 | a power of two (one E8M0 byte) | 4.25 | gpt-oss’s MoE weights, GGUF |
| NVFP4 | 16 | FP8 E4M3 per block × FP32 per tensor | 4.5 | NVIDIA’s Blackwell-optimized checkpoints |
A power-of-two scale can be off from the ideal by almost a factor of 2, wasting most of a binade; NVFP4’s FP8 scale is accurate to about 6% and its blocks are half as long, which is why it loses less.
def quantize_mxfp4(w, layout="interleaved"):
"""[..., K] (K % 32 == 0) -> (blocks uint8 [..., K/32, 16], scales uint8 [..., K/32]). (Your engine: Chapter 39)
The shared exponent puts the block's largest magnitude in E2M1's top binade:
e = floor(log2(amax)) - 2, since E2M1's largest value is 6 = 1.5 * 2^2.
"""
x = w.float().reshape(*w.shape[:-1], -1, 32)
amax = x.abs().amax(-1).clamp_min(2.0 ** -126)
exponent = (torch.floor(torch.log2(amax)) - 2).clamp(-127, 127)
codes = to_e2m1(x / torch.exp2(exponent)[..., None])
if layout == "interleaved":
packed = codes[..., 0::2] | (codes[..., 1::2] << 4)
else:
packed = codes[..., :16] | (codes[..., 16:] << 4)
return packed, (exponent + 127).to(torch.uint8)
def dequantize_mxfp4(blocks, scales, layout="interleaved"):
low, high = (blocks & 0xF).long(), (blocks >> 4).long()
if layout == "interleaved":
codes = torch.stack((low, high), -1).flatten(-2)
else:
codes = torch.cat((low, high), -1)
values = E2M1[codes] * torch.exp2(scales.float() - 127)[..., None]
return values.flatten(-2)
def quantize_nvfp4(w):
"""[..., K] (K % 16 == 0) -> (codes uint8 [..., K/16, 8], block scales fp8 [..., K/16], tensor scale)."""
x = w.float().reshape(*w.shape[:-1], -1, 16)
tensor_scale = x.abs().amax().clamp_min(1e-12) / (448.0 * 6.0) # E4M3's max x E2M1's max
block_scale = (x.abs().amax(-1) / 6.0 / tensor_scale).clamp(max=448.0).to(torch.float8_e4m3fn)
scale = block_scale.float() * tensor_scale
codes = to_e2m1(torch.where(scale[..., None] > 0, x / scale[..., None].clamp_min(1e-30), torch.zeros_like(x)))
return codes[..., 0::2] | (codes[..., 1::2] << 4), block_scale, tensor_scale
def dequantize_nvfp4(packed, block_scale, tensor_scale):
codes = torch.stack(((packed & 0xF).long(), (packed >> 4).long()), -1).flatten(-2)
return (E2M1[codes] * (block_scale.float() * tensor_scale)[..., None]).flatten(-2)
The two codes in each byte appear in two orders in the wild: gpt-oss interleaves them (elements $2j$ and $2j+1$), GGUF splits each block into halves (elements $j$ and $j + 16$). The test checks the second against Chapter 38’s GGUF decoder.
Trellis codes: QTIP and EXL3
At 2 bits per weight, every scalar code fails: four levels can’t represent a weight distribution well, however they’re placed. The best possible 2-bit scalar code for Gaussian data, the Lloyd-Max quantizer, still has a mean squared error of 0.118 per unit variance. Vector quantization does better by coding several weights jointly, but a codebook for 8 weights at 2 bits has $2^{16}$ entries to search and to keep in fast memory.
Trellis-coded quantization codes a long sequence with no stored codebook at all. QTIP (Tseng et al., 2024), which ExLlamaV3’s EXL3 format is built on, uses a bitshift trellis:
- a state is an $L$-bit integer, and each weight consumes $K$ new bits: $s_{t+1} = ((s_t \ll K) ,|, b_t) \bmod 2^L$;
- every state stands for a value, computed from its bits by a tiny hash. QTIP’s “1MAD” does one multiply-add mod $2^{32}$ and sums the four bytes of the result, and a sum of four roughly uniform bytes is roughly Gaussian;
- a group of 256 weights is coded by an $L$-bit start state and 255 $K$-bit symbols: $K$ bits per weight.
def code_1mad(states):
"""A pseudo-Gaussian value for each L-bit state: one multiply-add mod 2^32, then the sum of
the four bytes of the result (a sum of 4 roughly uniform bytes is roughly normal)."""
x = (states.to(torch.int64) * 34038481 + 76625530) % (2 ** 32)
total = sum((x >> (8 * i)) & 0xFF for i in range(4))
return (total.double() - 510.0) / 147.8
def codebook(L):
values = code_1mad(torch.arange(2 ** L))
return ((values - values.mean()) / values.std()).float() # exactly zero mean, unit variance
Which sequence of symbols reproduces the weights best? All $2^{L + 255K}$ sequences are candidates, but the cost decomposes along the path, so the best one is a shortest path through a graph of $2^L$ states per step: the Viterbi algorithm. Each state can be reached from exactly $2^K$ predecessors (those whose low $L - K$ bits equal its high bits), so each step is a vectorized minimum over a [2^L, 2^K] table:
@torch.no_grad()
def viterbi_encode(x, L=8, K=2):
"""x [T] (unit variance) -> (start state, K-bit symbols [T-1], reconstruction [T]). (Your engine: Chapter 39)
cost[s] = least squared error of any path ending in state s. A state s at step t can only
follow the 2^K states p whose low L - K bits are s's high bits: p = (s >> K) + j * 2^(L-K).
"""
table = codebook(L)
states = torch.arange(2 ** L)
T = x.shape[0]
cost = (x[0] - table) ** 2 # any start state
back = torch.zeros((T, 2 ** L), dtype=torch.long)
preds = (states >> K)[:, None] + (torch.arange(2 ** K) << (L - K))[None, :] # [S, 2^K]
for t in range(1, T):
options = cost[preds] # [S, 2^K]
best, which = options.min(-1)
back[t] = preds[states, which]
cost = best + (x[t] - table) ** 2
path = torch.zeros(T, dtype=torch.long)
path[-1] = cost.argmin()
for t in range(T - 1, 0, -1):
path[t - 1] = back[t, path[t]]
symbols = path[1:] & (2 ** K - 1) # the K new bits of each step
return int(path[0]), symbols, table[path]
Decoding needs no search and no table in memory: a shift, an OR and the hash per weight, which a GPU computes faster than it could load a codebook entry. The test checks Viterbi’s optimality by brute force on a small trellis (every one of the $2^4 \times 2^5$ codes), and that on Gaussian data its error is well below the Lloyd-Max optimum: about 0.08 at $L = 8$ and 0.07 at $L = 10$, against Lloyd-Max’s 0.118. QTIP reports about 0.07 at $L = 16$ with its tuned codes.
The codebook is shaped for Gaussian data, and real weight rows aren’t Gaussian: they have outliers and different scales. A random Hadamard transform (QuIP#, Tseng et al., 2024) fixes that. Multiplying a group by $H D / \sqrt{n}$ (a Hadamard matrix and random signs) is orthogonal, so it’s undone exactly at run time, and it spreads every outlier across all 256 coordinates, which then look like i.i.d. Gaussian samples:
def random_hadamard(n, seed=0):
"""An orthogonal n x n matrix H D / sqrt(n): Sylvester Hadamard times random signs."""
if n & (n - 1):
raise ValueError("n must be a power of two")
h = torch.ones(1, 1)
while h.shape[0] < n:
h = torch.cat((torch.cat((h, h), 1), torch.cat((h, -h), 1)), 0)
signs = torch.where(torch.rand(n, generator=torch.Generator().manual_seed(seed)) < 0.5, -1.0, 1.0)
return h * signs[None, :] / math.sqrt(n)
@torch.no_grad()
def quantize_trellis(weight, L=8, K=2, group=256, seed=0):
"""Rotate each row's groups to look Gaussian, code each group with the trellis, rotate back.
Returns the dequantized weight (what a decoder would produce) and the bits per weight."""
rows, cols = weight.shape
rotation = random_hadamard(group, seed)
x = weight.float().reshape(rows, -1, group) @ rotation.T # incoherence processing
scale = x.pow(2).mean(-1, keepdim=True).sqrt().clamp_min(1e-12)
out = torch.empty_like(x)
for r in range(rows):
for g in range(x.shape[1]):
_, _, recon = viterbi_encode(x[r, g] / scale[r, g], L, K)
out[r, g] = recon * scale[r, g]
bits = (L + K * (group - 1) + 16) / group # start state, symbols, an f16 scale
return (out @ rotation).reshape(rows, cols), bits
This is the algorithm. ExLlamaV3 adds its engineering: $L = 16$, a hash called “3INST” chosen for GPU instruction throughput, 16 × 16 tiles laid out for tensor cores, per-tensor bit widths from 1 to 8 chosen by sensitivity, and its own file format. This chapter’s code doesn’t read EXL3 files; it shows why EXL3 can make a 2.5-bit model usable where round-to-nearest can’t.
Loading and writing checkpoints
All three Hugging Face formats share a structure: config.json’s quantization_config names the method, and each quantized linear layer’s tensors sit under its name. One loader streams the tensors and builds each quantized layer as soon as all of its parts have arrived:
@torch.no_grad()
def load_quantized(directory, device="cpu", dtype=torch.bfloat16):
"""A Qwen3 (dense) model from a GPTQ, AWQ or FP8 checkpoint. (Your engine: Chapter 39)
Tensors are streamed; a quantized layer is built as soon as all of its parts have arrived,
then put in place of the model's nn.Linear.
"""
from ..loaders import assign, read_config
from ..qwen3 import Qwen3, Qwen3Config
from ..safetensors_io import snapshot_tensors
raw = read_config(directory)
q = raw.get("quantization_config") or {}
method = q.get("quant_method")
if method not in PARTS:
raise ValueError(f"Unsupported quant_method {method!r} (gptq, awq or fp8)")
cfg = Qwen3Config.from_hf(raw)
with torch.device("meta"):
model = Qwen3(cfg)
model = model.to_empty(device=device).to(dtype)
if cfg.tie_word_embeddings:
model.lm_head.weight = model.model.embed_tokens.weight
params = dict(model.named_parameters())
linears = {name for name, m in model.named_modules() if isinstance(m, torch.nn.Linear) and name != "lm_head"}
pending = {}
for name, value in snapshot_tensors(directory):
module, _, part = name.rpartition(".")
if module in linears and (part in PARTS[method] or part == "bias"):
pending.setdefault(module, {})[part] = value
if all(p in pending[module] for p in PARTS[method]):
parts = pending.pop(module)
layer = build(method, parts, q, dtype)
parent, _, child = module.rpartition(".")
setattr(model.get_submodule(parent), child, layer.to(device))
continue
if name == "lm_head.weight" and cfg.tie_word_embeddings:
continue
if name not in params:
raise ValueError(f"Unexpected tensor {name}")
assign(params[name], value, name)
if pending:
raise ValueError(f"Incomplete quantized layers: {sorted(pending)[:3]}")
return model.eval()
def build(method, parts, q, dtype):
bias = parts.get("bias")
if method == "gptq":
if q.get("bits", 4) != 4:
raise ValueError("Only 4-bit GPTQ is implemented")
codes, scales, zeros, g_idx = gptq.unpack_gptq(parts["qweight"], parts["qzeros"], parts["scales"], parts["g_idx"])
return gptq.to_runtime(codes, scales, zeros, g_idx, q["group_size"], bias, dtype)
if method == "awq":
return awq.from_checkpoint(parts["qweight"], parts["qzeros"], parts["scales"], q["group_size"], bias, dtype)
if q.get("weight_block_size", [128, 128]) != [128, 128]:
raise ValueError("Only 128 x 128 FP8 blocks are implemented")
return fp8.FP8BlockLinear(parts["weight"], parts["weight_scale_inv"], bias, dtype)
and one writer quantizes a dense model and produces the same layout, so a model you fine-tuned in Chapter 22 can ship as GPTQ, AWQ or FP8:
@torch.no_grad()
def save_quantized(model, directory, method, calibration_ids=None, group_size=128, act_order=False):
"""Quantize a dense Qwen3 and write it as a checkpoint that load_quantized (and vLLM's and
transformers' loaders for the same method) can read."""
from ..safetensors_io import save_file
directory = Path(directory)
directory.mkdir(parents=True, exist_ok=True)
raw = {k: v for k, v in vars(model.cfg).items()}
raw["model_type"] = "qwen3"
tensors = {}
inputs = {}
if method == "gptq":
handles = [m.register_forward_hook(lambda mod, a, o, n=n: inputs.setdefault(n, []).append(a[0].reshape(-1, a[0].shape[-1])))
for n, m in model.named_modules() if isinstance(m, torch.nn.Linear) and n != "lm_head"]
model(calibration_ids)
for h in handles:
h.remove()
if method == "awq":
replaced = awq.apply_awq(model, calibration_ids, 4, group_size)
for name, module in model.named_modules():
if name == "lm_head" or not isinstance(module, (torch.nn.Linear, awq.AffineQuantLinear)):
continue
if method == "gptq":
H = gptq.hessian_from(torch.cat(inputs[name]))
codes, scales, zeros, g_idx = gptq.gptq_quantize(module.weight, H, 4, group_size, act_order)
tensors.update({f"{name}.{k}": v for k, v in gptq.pack_gptq(codes, scales, zeros, g_idx).items()})
elif method == "awq":
tensors.update({f"{name}.{k}": v for k, v in awq.to_checkpoint(replaced[name]).items()})
else:
weight, scale_inv = fp8.quantize_blocks(module.weight)
tensors[f"{name}.weight"], tensors[f"{name}.weight_scale_inv"] = weight, scale_inv
for name, p in model.named_parameters():
module = name.rpartition(".")[0]
if module == "lm_head" and model.cfg.tie_word_embeddings:
continue
if not any(key.startswith(module + ".") for key in tensors) or module in ("lm_head", ""):
tensors.setdefault(name, p.detach().contiguous())
raw["quantization_config"] = ({"quant_method": "fp8", "weight_block_size": [128, 128], "activation_scheme": "dynamic"}
if method == "fp8" else
{"quant_method": method, "bits": 4, "group_size": group_size, "desc_act": act_order,
"sym": False, **({"version": "gemm", "zero_point": True} if method == "awq" else {})})
(directory / "config.json").write_text(json.dumps(raw))
save_file(tensors, directory / "model.safetensors")
The tests write and reload each format and check that the result stays close to the original model. They can’t check compatibility with the reference tools themselves, AutoGPTQ, AutoAWQ and vLLM’s FP8 loader, because those need downloads this book’s test environment doesn’t allow. The layouts follow those tools’ source code and are pinned by the format tests; Appendix F records what was and wasn’t verified.
Run it
python run.py quantformats --steps 300
The command trains a small Qwen3 (2 layers, 256 wide) on The Verdict as bytes, then quantizes every linear layer with each method and measures the KL divergence of the quantized model’s next-token distributions from the original’s on held-out text. KL is llama.cpp’s preferred metric because, unlike perplexity, it can’t be improved by noise: a small overfit model’s held-out perplexity sometimes drops when quantization perturbs it, which says nothing about quality.
{"format": "bf16 reference", "bits": 16, "held_out_perplexity": 21.63}
{"format": "RTN int4, g64", "bits": 4.5, "kl_nats": 0.02238}
{"format": "GPTQ int4, g64", "bits": 4.5, "kl_nats": 0.00588}
{"format": "GPTQ int4, g64, act-order", "bits": 4.5, "kl_nats": 0.00505}
{"format": "AWQ int4, g64", "bits": 4.5, "kl_nats": 0.02201}
{"format": "FP8 E4M3, 128x128 blocks", "bits": 8.0, "kl_nats": 0.00216}
{"format": "MXFP4", "bits": 4.25, "kl_nats": 0.04254}
{"format": "NVFP4", "bits": 4.5, "kl_nats": 0.02723}
{"format": "trellis (QTIP-style), L=8", "bits": 2.09, "kl_nats": 0.29767}
{"format": "trellis (QTIP-style), L=8", "bits": 3.08, "kl_nats": 0.08815}
{"format": "RTN int2, g64 (for comparison)", "bits": 2.5, "kl_nats": 0.52499}
GPTQ cuts round-to-nearest’s damage by a factor of four at the same size, and act-order helps a little more. AWQ barely beats round-to-nearest here, and that’s informative rather than disappointing: this small model, trained for a few minutes, has none of the large outlier channels that real LLMs develop and that AWQ was designed to protect (the test on synthetic outliers shows the effect). FP8 is nearly lossless. NVFP4 beats MXFP4, its finer scales worth more than their quarter-bit. At about 2 bits, trellis codes are much better than 2.5-bit round-to-nearest, while still well behind 4 bits: on this tiny model every bit matters; on large models, whose weights are more redundant, QTIP and EXL3 at 2-3 bits come much closer to the original.
Build it
Engine milestone 39: quantized checkpoints. Implement gptq_quantize and unpack_gptq (engine/formats/gptq.py); search_scale, apply_awq and unpack_awq (engine/formats/awq.py); quantize_blocks (engine/formats/fp8.py); quantize_mxfp4 (engine/formats/mx.py); viterbi_encode (engine/formats/trellis.py); load_quantized (engine/formats/hf_quant.py); and fp8_block_kernel (engine/kernels/triton_formats.py). The packers, NVFP4, Hadamard rotation, decoders and the writer are provided.
pytest tests/test_ch39_quantized_checkpoints.py
python run.py quantformats --impl engine
The tests check GPTQ against round-to-nearest with and without act-order, GPTQ’s packing round trip and its act-order run-time permutation, AWQ on synthetic outlier channels and its interleaved packing, saving and loading GPTQ, AWQ and FP8 checkpoints of a Qwen3, FP8 blocks with the kernel and the W8A8 reference, MXFP4 in both byte orders (one against the GGUF decoder), NVFP4, Viterbi’s optimality by brute force, and trellis codes against the Lloyd-Max optimum.
Stretch exercises
- ★★ Measure every method on a real model: quantize Qwen3-0.6B with RTN, GPTQ (with and without act-order) and AWQ, using 128 calibration sequences of 2,048 tokens, and compare KL and a downstream accuracy (Chapter 44’s harness). Which layers suffer most? Where:
experiments/ch39.py(create it), importingengine.formats.gptq,engine.formats.awqandengine.evaluation. - ★★ Write a Triton kernel that decodes trellis codes on the fly: each program replays the bitshift state for its tile’s weights and applies the inverse Hadamard transform to the activations instead of the weights ($x H^\top$ is cheap with a fast Walsh-Hadamard transform). Where: add a trellis kernel in
engine/kernels/triton_formats.py; connect it toengine/formats/trellis.py. - ★★★ Combine them: GPTQ’s column-by-column error feedback with trellis codes instead of scalar rounding (QTIP’s “BlockLDLQ”). Measure the improvement over plain trellis coding at 2 bits. Where: add a feedback-based trellis quantizer in
engine/formats/trellis.py, usingengine.formats.gptq. - ★ Load gpt-oss-20b’s MXFP4 expert tensors (
*_blocks,*_scales) withdequantize_mxfp4(layout="interleaved")and check a few values against the BF16 conversion that Hugging Face publishes. Where:experiments/ch39.py(create it), importingengine.formats.mx.dequantize_mxfp4.
Check your understanding
- What information does GPTQ use that round-to-nearest ignores, and how does a quantization error in one column change the others?
- Why does act-order need
g_idx, and how does the engine serve such a checkpoint without scattered groups? - Why does scaling a weight column up before quantization reduce its relative error, and why is dividing the activations by the same factor free?
- Why does FP8 need block scales at all, given its exponent bits?
- MXFP4 and NVFP4 both store E2M1 values. Why does NVFP4 lose less?
- Why can a trellis code beat the best scalar code at the same number of bits? What makes it cheap to decode?
- Why is KL divergence a better quantization metric than perplexity for a small model?
Going deeper
- Frantar, Ashkboos, Hoefler, Alistarh, GPTQ: Accurate Post-Training Quantization for Generative Pre-trained Transformers (ICLR 2023); Frantar and Alistarh, Optimal Brain Compression (2022).
- Lin et al., AWQ: Activation-aware Weight Quantization for LLM Compression and Acceleration (MLSys 2024); Xiao et al., SmoothQuant (ICML 2023).
- Tseng et al., QuIP#: Even Better LLM Quantization with Hadamard Incoherence and Lattice Codebooks (ICML 2024) and QTIP: Quantization with Trellises and Incoherence Processing (NeurIPS 2024); the ExLlamaV3 repository’s documentation of EXL3.
- DeepSeek-AI, DeepSeek-V3 Technical Report (2024), §3.3 on FP8 training with fine-grained block scaling; the Open Compute Project’s Microscaling Formats (MX) Specification v1.0 (2023); NVIDIA’s NVFP4 documentation.
- AutoGPTQ / GPTQModel, AutoAWQ and vLLM’s
vllm/model_executor/layers/quantization/for the reference formats and kernels (Marlin, Machete, the FP8 block GEMMs).
40. Running everywhere: CPUs, Apple, AMD and offloading
In this chapter
- One engine on many devices: what changes between NVIDIA, AMD, Apple, Intel and CPUs, and how the engine picks its defaults.
- Fast decode on a CPU: why it's bandwidth-bound like a GPU, and SIMD kernels for quantized weights with integer dot products (AVX2, AVX-512 VNNI, ARM dotprod), in C++, Rust and Python.
- Apple GPUs through Metal, AMD GPUs through ROCm and Triton, and what Vulkan would take.
- Models bigger than the GPU: splitting layers across devices, keeping MoE experts in host memory, and streaming layers, each with its cost model.
You will build
detect and engine_defaults (engine/platform.py), and SplitModel.forward, OffloadedExperts.fetch, StreamedModel.load and StreamedModel.forward (engine/offload.py). The CPU kernels are the C++ and Rust tracks' contribution: cpp/izh_cpu.hpp and rust/src/simd.rs.
Time: 6-8 hours. GPU: not needed (the CPU is the point; Apple and AMD sections need that hardware to run).
One engine, many machines
llama.cpp’s popularity comes as much from where it runs as from how fast: laptops without a discrete GPU, Macs, AMD cards, phones. vLLM and SGLang run on NVIDIA and AMD datacenter GPUs. The engine you’ve built so far is NVIDIA-first: Triton kernels, CUDA graphs, FP8 checks against compute capability. This chapter makes it portable and gives it a fast CPU path. Here’s what changes from platform to platform:
| platform | how PyTorch sees it | kernels | graphs | notes |
|---|---|---|---|---|
| NVIDIA | cuda | Triton, CUDA | yes | FP8 from compute capability 8.9, FP4 from 10.0 |
| AMD (ROCm) | cuda (with torch.version.hip set) | the same Triton kernels, compiled for AMD | yes (HIP graphs) | 64-thread wavefronts; MI300’s FP8 is e4m3fnuz, not OCP e4m3fn |
| Apple | mps | PyTorch’s Metal ops; Metal shaders | no | unified memory: CPU and GPU share all of RAM |
| Intel GPUs | xpu | Triton (Intel’s backend) | no | |
| CPU | cpu | SIMD C++ / Rust kernels | no | threads = physical cores; NUMA matters on servers |
def detect():
"""Inspect the machine. Override the choice with IZH_DEVICE=cpu|cuda|mps|xpu. (Your engine: Chapter 40)"""
forced = os.environ.get("IZH_DEVICE")
if (forced in (None, "cuda")) and torch.cuda.is_available():
props = torch.cuda.get_device_properties(0)
if torch.version.hip: # ROCm builds reuse the cuda device API
arch = getattr(props, "gcnArchName", "gfx").split(":")[0]
return Platform("rocm", props.name, arch=arch, memory_bytes=props.total_memory, features={
"bf16": True, "fp8": arch.startswith(("gfx94", "gfx95", "gfx12")), "triton": True, "graphs": True,
"warp_size": 64})
cap = (props.major, props.minor)
return Platform("cuda", props.name, cap, memory_bytes=props.total_memory, features={
"bf16": cap >= (8, 0), "fp8": cap >= (8, 9), "fp4": cap >= (10, 0), "triton": True, "graphs": True,
"warp_size": 32})
if (forced in (None, "mps")) and getattr(torch.backends, "mps", None) and torch.backends.mps.is_available():
return Platform("mps", "Apple GPU", memory_bytes=_system_memory(), features={
"bf16": True, "fp8": False, "triton": False, "graphs": False, "unified_memory": True})
if (forced in (None, "xpu")) and hasattr(torch, "xpu") and torch.xpu.is_available():
return Platform("xpu", torch.xpu.get_device_name(0), features={"bf16": True, "triton": True, "graphs": False})
return Platform("cpu", _cpu_name(), memory_bytes=_system_memory(), features={
"bf16": True, "triton": False, "graphs": False, "simd": _cpu_isa()})
def engine_defaults(platform):
"""EngineConfig fields that suit the hardware; explicit settings always win. (Your engine: Chapter 40)"""
if platform.kind in ("cuda", "rocm"):
return {"attention_backend": "triton", "cuda_graphs": "auto", "kv_cache_dtype": "auto"}
if platform.kind == "xpu":
return {"attention_backend": "triton", "cuda_graphs": "off"}
# MPS and CPU: no Triton (the interpreter is for testing, not serving), no graphs, a modest batch.
return {"attention_backend": "reference", "cuda_graphs": "off", "max_num_seqs": 16, "max_num_batched_tokens": 512}
def default_dtype(platform):
if platform.kind in ("cuda", "rocm"):
return torch.bfloat16 if platform.features.get("bf16") else torch.float16
if platform.kind == "mps":
return torch.float16 # the fastest MPS matmul path
return torch.bfloat16 if platform.kind == "cpu" and "avx512" in str(platform.features.get("simd")) else torch.float32
The defaults are a starting point that explicit settings override: Triton attention and CUDA graphs on NVIDIA and AMD GPUs; the reference attention backend, no graphs and a smaller batch on Apple and CPUs; BF16 where the hardware computes it natively. On AMD, nothing else in Parts VIII’s code needs to change: the Triton kernels of Chapters 32, 33 and 38 compile for AMD GPUs as they are. Their tile sizes were chosen for NVIDIA’s 32-thread warps, so autotuning them for 64-thread wavefronts (and MI300’s larger register file) is worth doing; vLLM and SGLang also ship AMD-specific kernels (the AITER library) for the hottest operations.
Decoding on a CPU
A CPU decodes like a GPU: every token reads every weight once, so the time per token is bytes read divided by memory bandwidth (Chapter 10). The difference is the bandwidth: a laptop’s dual-channel DDR5 delivers 60-100 GB/s, an Apple M-series chip 100-800 GB/s (shared with its GPU), a 12-channel server socket 400-600 GB/s, against an H100’s 3,350 GB/s. A CPU is a perfectly good decoder for a quantized model of a few billion parameters, and the same rule as everywhere applies: fewer bytes per weight, more tokens per second.
That makes quantized weights essential, and it creates a kernel problem. Dequantizing a whole weight matrix to float before each matmul (what AffineQuantLinear does without a kernel) reads the compact weights, writes 4 bytes per weight back to memory, and reads them again: slower than not quantizing at all. The weights must be decoded in registers, inside the dot product, as Chapter 20’s Triton kernel does on a GPU.
Integer dot products
CPUs have a faster trick than decoding to float: integer dot products. llama.cpp quantizes the activation vector too, to int8 in blocks of 32 with one float scale each, and then each 32-weight block’s contribution is
$$ y_n \mathrel{+}= \underbrace{s_{n,g}}{\text{weight scale}} \cdot \underbrace{d_b}{\text{activation scale}} \cdot \underbrace{\sum_{k \in b} q_{n,k}, a_k}{\text{integer dot product}} ;+; o{n,g} \sum_{k \in g} x_k , $$
where the integer dot product of 32 4-bit codes with 32 int8 activations is one or two instructions on modern CPUs:
| instruction set | instruction | does |
|---|---|---|
| x86 AVX2 | vpmaddubsw + vpmaddwd | 32 u8 × s8 products, summed in pairs to s16, then to s32 |
| x86 AVX-512 VNNI / AVX-VNNI | vpdpbusd | 32 u8 × s8 products summed four at a time straight into s32 |
| ARM v8.2 dotprod | sdot | 16 s8 × s8 products summed four at a time into s32 |
The offsets need only the activations’ per-group sums, computed once per call. Quantizing activations costs a little accuracy (about 1% relative error per matmul here), which is why llama.cpp offers it per type and why AffineQuantLinear uses these kernels only when asked (set_cpu_kernels(True)).
The kernels live in a header shared by the C++ track and a PyTorch extension, and in the Rust track. First the activations, quantized once per call, with each block’s codes stored as its 16 even positions then its 16 odd ones, to meet the engine’s packing (code $k$ in the low nibble when $k$ is even):
// x [K] -> int8 codes per block of 32, one float scale per block, and each block's float sum
// (the offsets multiply sum(x) over a group). For the q4 kernels each block's codes are stored
// as [the 16 even positions | the 16 odd positions], to meet the low and high nibbles.
struct Q8Activations {
std::vector<int8_t> q, q_split;
std::vector<float> scale, sum;
};
inline Q8Activations quantize_activations(const float* x, int K) {
Q8Activations a;
int blocks = K / 32;
a.q.resize(K), a.q_split.resize(K), a.scale.resize(blocks), a.sum.resize(blocks);
for (int b = 0; b < blocks; ++b) {
const float* xb = x + 32 * b;
float amax = 0, total = 0;
for (int i = 0; i < 32; ++i) amax = std::max(amax, std::fabs(xb[i])), total += xb[i];
float d = amax / 127.f, inv = d > 0 ? 1.f / d : 0.f;
a.scale[b] = d, a.sum[b] = total;
for (int i = 0; i < 32; ++i) {
int8_t v = int8_t(std::lround(xb[i] * inv));
a.q[32 * b + i] = v;
a.q_split[32 * b + (i % 2) * 16 + i / 2] = v; // evens first, then odds
}
}
return a;
}
Then one block’s integer dot product, for each instruction set. Each function is compiled for exactly its own instruction set with a target attribute, and the best one is chosen at run time, so one binary runs on any x86 CPU:
// One block of 32: integer dot product of 4-bit codes (16 bytes) with 32 int8 activations.
inline int32_t dot_q4_scalar(const uint8_t* c, const int8_t* x) {
int32_t s = 0;
for (int i = 0; i < 16; ++i) s += (c[i] & 15) * x[i] + (c[i] >> 4) * x[16 + i];
return s;
}
inline int32_t dot_q8_scalar(const int8_t* c, const int8_t* x) {
int32_t s = 0;
for (int i = 0; i < 32; ++i) s += c[i] * x[i];
return s;
}
#ifdef IZH_X86
__attribute__((target("avx2"))) inline __m256i nibbles_avx2(const uint8_t* c) {
__m128i b = _mm_loadu_si128(reinterpret_cast<const __m128i*>(c));
__m128i m = _mm_set1_epi8(15);
return _mm256_set_m128i(_mm_and_si128(_mm_srli_epi16(b, 4), m), _mm_and_si128(b, m)); // [low | high]
}
__attribute__((target("avx2"))) inline __m256i dot_q4_avx2(const uint8_t* c, const int8_t* x) {
__m256i p16 = _mm256_maddubs_epi16(nibbles_avx2(c), _mm256_loadu_si256(reinterpret_cast<const __m256i*>(x)));
return _mm256_madd_epi16(p16, _mm256_set1_epi16(1)); // 8 lanes of int32 partial sums
}
__attribute__((target("avx2"))) inline __m256i dot_q8_avx2(const int8_t* c, const int8_t* x) {
__m256i cv = _mm256_loadu_si256(reinterpret_cast<const __m256i*>(c));
__m256i xv = _mm256_loadu_si256(reinterpret_cast<const __m256i*>(x));
// maddubs wants an unsigned left operand: |c| x (x * sign(c)) has the same products.
__m256i p16 = _mm256_maddubs_epi16(_mm256_sign_epi8(cv, cv), _mm256_sign_epi8(xv, cv));
return _mm256_madd_epi16(p16, _mm256_set1_epi16(1));
}
__attribute__((target("avx512vnni,avx512vl,avx2"))) inline __m256i dot_q4_vnni(const uint8_t* c, const int8_t* x) {
return _mm256_dpbusd_epi32(_mm256_setzero_si256(), nibbles_avx2(c),
_mm256_loadu_si256(reinterpret_cast<const __m256i*>(x)));
}
__attribute__((target("avx2"))) inline float hsum(__m256 v) {
__m128 s = _mm_add_ps(_mm256_castps256_ps128(v), _mm256_extractf128_ps(v, 1));
s = _mm_hadd_ps(s, s);
return _mm_cvtss_f32(_mm_hadd_ps(s, s));
}
#endif
#ifdef IZH_NEON_DOT
inline int32_t dot_q4_neon(const uint8_t* c, const int8_t* x) {
uint8x16_t b = vld1q_u8(c);
int8x16_t lo = vreinterpretq_s8_u8(vandq_u8(b, vdupq_n_u8(15))), hi = vreinterpretq_s8_u8(vshrq_n_u8(b, 4));
int32x4_t s = vdotq_s32(vdupq_n_s32(0), lo, vld1q_s8(x));
return vaddvq_s32(vdotq_s32(s, hi, vld1q_s8(x + 16)));
}
#endif
#![allow(unused)]
fn main() {
/// x [K] -> int8 per block of 32 (evens then odds, to meet low and high nibbles), scales, sums.
pub fn quantize_activations(x: &[f32]) -> (Vec<i8>, Vec<f32>, Vec<f32>) {
let mut q = vec![0i8; x.len()];
let (mut scale, mut sum) = (vec![], vec![]);
for (b, block) in x.chunks(32).enumerate() {
let amax = block.iter().fold(0.0f32, |m, v| m.max(v.abs()));
let d = amax / 127.0;
let inv = if d > 0.0 { 1.0 / d } else { 0.0 };
for (i, v) in block.iter().enumerate() {
q[32 * b + (i % 2) * 16 + i / 2] = (v * inv).round() as i8;
}
scale.push(d);
sum.push(block.iter().sum());
}
(q, scale, sum)
}
/// One block: 16 bytes of codes against 32 activations (evens then odds).
pub fn dot_q4_scalar(codes: &[u8], x: &[i8]) -> i32 {
(0..16).map(|i| (codes[i] & 15) as i32 * x[i] as i32 + (codes[i] >> 4) as i32 * x[16 + i] as i32).sum()
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
unsafe fn dot_q4_avx2(codes: &[u8], x: &[i8]) -> i32 {
use std::arch::x86_64::*;
let b = _mm_loadu_si128(codes.as_ptr() as *const __m128i);
let m = _mm_set1_epi8(15);
let c = _mm256_set_m128i(_mm_and_si128(_mm_srli_epi16::<4>(b), m), _mm_and_si128(b, m));
let p16 = _mm256_maddubs_epi16(c, _mm256_loadu_si256(x.as_ptr() as *const __m256i));
let p32 = _mm256_madd_epi16(p16, _mm256_set1_epi16(1));
let s = _mm_add_epi32(_mm256_castsi256_si128(p32), _mm256_extracti128_si256::<1>(p32));
let s = _mm_hadd_epi32(s, s);
_mm_cvtsi128_si32(_mm_hadd_epi32(s, s))
}
/// The best dot product this CPU supports, decided once per call site by runtime detection.
pub fn dot_q4(codes: &[u8], x: &[i8]) -> i32 {
assert!(codes.len() >= 16 && x.len() >= 32);
#[cfg(target_arch = "x86_64")]
{
if is_x86_feature_detected!("avx2") {
return unsafe { dot_q4_avx2(codes, x) };
}
}
dot_q4_scalar(codes, x)
}
/// y [N] = W x for 4-bit codes [N, K/2], scales and offsets [N, K/G], rows split across threads.
pub fn gemv_q4(codes: &[u8], scales: &[f32], offsets: &[f32], n: usize, k: usize, g: usize, x: &[f32],
threads: usize) -> Vec<f32> {
let (q, xs, xsum) = quantize_activations(x);
let mut y = vec![0.0f32; n];
let per = n.div_ceil(threads.max(1));
std::thread::scope(|scope| {
for (t, chunk) in y.chunks_mut(per).enumerate() {
let (q, xs, xsum) = (&q, &xs, &xsum);
scope.spawn(move || {
for (j, out) in chunk.iter_mut().enumerate() {
let row = t * per + j;
let (c, s, o) = (&codes[row * k / 2..], &scales[row * (k / g)..], &offsets[row * (k / g)..]);
let mut acc = 0.0f32;
for b in 0..k / 32 {
acc += s[32 * b / g] * xs[b] * dot_q4(&c[16 * b..], &q[32 * b..]) as f32;
}
for grp in 0..k / g {
acc += o[grp] * xsum[grp * g / 32..(grp + 1) * g / 32].iter().sum::<f32>();
}
*out = acc;
}
});
}
});
y
}
}
Then the row loop, and here two details doubled the speed when this chapter’s first version measured only a third of the memory bandwidth:
- No division in the inner loop. The first version looked up each block’s scale as
scales[32 * b / G], an integer division by a runtime value, 20-40 cycles per block. Looping over groups outside and blocks inside removes it. - Independent accumulators, and four rows at a time. A single accumulator makes every fused multiply-add wait for the previous one (4 cycles of latency each). Four accumulators, and four rows sharing each block’s activation load, keep the multiply units busy.
// y[n] = sum over groups g of scale[n, g] * sum over blocks b in g of xscale[b] * dot_b + offset[n, g] * xsum[g]
// Two things decide the speed: no division in the inner loop (groups outside, blocks inside),
// and four independent accumulators, so consecutive fused multiply-adds don't wait for each other.
inline std::vector<float> group_sums(const Q8Activations& a, int K, int G) {
std::vector<float> s(K / G, 0.f);
for (int b = 0; b < K / 32; ++b) s[32 * b / G] += a.sum[b];
return s;
}
#ifdef IZH_X86
#define IZH_ROW_BODY(DOT, CODES_PER_BLOCK, XQ) \
__m256 acc[4] = {_mm256_setzero_ps(), _mm256_setzero_ps(), _mm256_setzero_ps(), _mm256_setzero_ps()}; \
int per = G / 32; \
for (int g = 0, b = 0; g < K / G; ++g) { \
float sg = scales[g]; \
for (int i = 0; i < per; ++i, ++b) { \
__m256i p = DOT(codes + CODES_PER_BLOCK * b, XQ + 32 * b); \
acc[b & 3] = _mm256_fmadd_ps(_mm256_cvtepi32_ps(p), _mm256_set1_ps(sg * a.scale[b]), acc[b & 3]); \
} \
} \
return hsum(_mm256_add_ps(_mm256_add_ps(acc[0], acc[1]), _mm256_add_ps(acc[2], acc[3])));
// Separate functions, not one template: each must be compiled for exactly its own instruction set,
// or the compiler may use AVX-512 instructions in the path meant for AVX2-only CPUs.
__attribute__((target("avx2,fma"))) inline float row_q4_avx2(const uint8_t* codes, const float* scales, int K, int G,
const Q8Activations& a) {
IZH_ROW_BODY(dot_q4_avx2, 16, a.q_split.data())
}
__attribute__((target("avx512vnni,avx512vl,avx2,fma"))) inline float row_q4_vnni(const uint8_t* codes, const float* scales,
int K, int G, const Q8Activations& a) {
IZH_ROW_BODY(dot_q4_vnni, 16, a.q_split.data())
}
__attribute__((target("avx2,fma"))) inline float row_q8_avx2(const int8_t* codes, const float* scales, int K, int G,
const Q8Activations& a) {
IZH_ROW_BODY(dot_q8_avx2, 32, a.q.data())
}
#undef IZH_ROW_BODY
// Four rows at once: each block's 32 activations are loaded once and used four times, and the four
// rows' independent dot products keep the multiply units busy. y[0..3] receive the four results.
#define IZH_ROWS4(NAME, TARGET, DOT) \
__attribute__((target(TARGET))) inline void NAME(const uint8_t* codes, size_t row_bytes, const float* scales, \
size_t row_groups, int K, int G, const Q8Activations& a, \
float* y) { \
__m256 acc[4] = {_mm256_setzero_ps(), _mm256_setzero_ps(), _mm256_setzero_ps(), _mm256_setzero_ps()}; \
int per = G / 32; \
for (int g = 0, b = 0; g < K / G; ++g) { \
__m256 s[4]; \
for (int r = 0; r < 4; ++r) s[r] = _mm256_set1_ps(scales[r * row_groups + g]); \
for (int i = 0; i < per; ++i, ++b) { \
__m256i xv = _mm256_loadu_si256(reinterpret_cast<const __m256i*>(a.q_split.data() + 32 * b)); \
__m256 xs = _mm256_set1_ps(a.scale[b]); \
for (int r = 0; r < 4; ++r) { \
__m256i p = DOT(nibbles_avx2(codes + r * row_bytes + 16 * b), xv); \
acc[r] = _mm256_fmadd_ps(_mm256_cvtepi32_ps(p), _mm256_mul_ps(s[r], xs), acc[r]); \
} \
} \
} \
for (int r = 0; r < 4; ++r) y[r] = hsum(acc[r]); \
}
__attribute__((target("avx2"))) inline __m256i madd_u8s8_avx2(__m256i c, __m256i x) {
return _mm256_madd_epi16(_mm256_maddubs_epi16(c, x), _mm256_set1_epi16(1));
}
__attribute__((target("avx512vnni,avx512vl,avx2"))) inline __m256i madd_u8s8_vnni(__m256i c, __m256i x) {
return _mm256_dpbusd_epi32(_mm256_setzero_si256(), c, x);
}
IZH_ROWS4(rows4_q4_avx2, "avx2,fma", madd_u8s8_avx2)
IZH_ROWS4(rows4_q4_vnni, "avx512vnni,avx512vl,avx2,fma", madd_u8s8_vnni)
#undef IZH_ROWS4
#endif
inline float add_offsets(float y, const float* offsets, const std::vector<float>& gsum) {
for (size_t g = 0; g < gsum.size(); ++g) y += offsets[g] * gsum[g];
return y;
}
inline float row_q4(const uint8_t* codes, const float* scales, const float* offsets, int K, int G,
const Q8Activations& a, const std::vector<float>& gsum, Isa isa) {
float y = 0;
switch (isa) {
#ifdef IZH_X86
case Isa::avx512_vnni: y = row_q4_vnni(codes, scales, K, G, a); break;
case Isa::avx2: y = row_q4_avx2(codes, scales, K, G, a); break;
#endif
#ifdef IZH_NEON_DOT
case Isa::neon_dot:
for (int b = 0; b < K / 32; ++b)
y += scales[32 * b / G] * a.scale[b] * float(dot_q4_neon(codes + 16 * b, a.q_split.data() + 32 * b));
break;
#endif
default:
for (int b = 0; b < K / 32; ++b)
y += scales[32 * b / G] * a.scale[b] * float(dot_q4_scalar(codes + 16 * b, a.q_split.data() + 32 * b));
}
return add_offsets(y, offsets, gsum);
}
// Rows n .. n+3 into y[0..3]; falls back to one row at a time where there's no 4-row kernel.
inline void rows4_q4(const uint8_t* codes, const float* scales, const float* offsets, int K, int G,
const Q8Activations& a, const std::vector<float>& gsum, Isa isa, float* y) {
size_t row_bytes = size_t(K) / 2, row_groups = size_t(K / G);
#ifdef IZH_X86
if (isa == Isa::avx512_vnni || isa == Isa::avx2) {
if (isa == Isa::avx512_vnni)
rows4_q4_vnni(codes, row_bytes, scales, row_groups, K, G, a, y);
else
rows4_q4_avx2(codes, row_bytes, scales, row_groups, K, G, a, y);
for (int r = 0; r < 4; ++r) y[r] = add_offsets(y[r], offsets + r * row_groups, gsum);
return;
}
#endif
for (int r = 0; r < 4; ++r)
y[r] = row_q4(codes + r * row_bytes, scales + r * row_groups, offsets + r * row_groups, K, G, a, gsum, isa);
}
inline float row_q8(const int8_t* codes, const float* scales, const float* offsets, int K, int G,
const Q8Activations& a, const std::vector<float>& gsum, Isa isa) {
float y = 0;
#ifdef IZH_X86
if (isa == Isa::avx2 || isa == Isa::avx512_vnni)
y = row_q8_avx2(codes, scales, K, G, a);
else
#endif
for (int b = 0; b < K / 32; ++b)
y += scales[32 * b / G] * a.scale[b] * float(dot_q8_scalar(codes + 32 * b, a.q.data() + 32 * b));
return add_offsets(y, offsets, gsum);
}
// y [N] = W x for one activation row, rows split across threads. Each thread streams its own
// slice of the weight matrix: decode is memory-bound, so more threads means more bandwidth,
// up to what the memory system delivers.
inline void gemv_q4(const uint8_t* codes, const float* scales, const float* offsets, int N, int K, int G,
const float* x, float* y, int threads = 0, Isa isa = best_isa()) {
Q8Activations a = quantize_activations(x, K);
std::vector<float> gsum = group_sums(a, K, G);
threads = threads > 0 ? threads : int(std::max(1u, std::thread::hardware_concurrency()));
std::vector<std::thread> pool;
int per = (N + threads - 1) / threads;
for (int t = 0; t < threads; ++t) {
pool.emplace_back([&, t] {
int n = t * per, end = std::min(N, (t + 1) * per);
for (; n + 4 <= end; n += 4)
rows4_q4(codes + size_t(n) * K / 2, scales + size_t(n) * (K / G), offsets + size_t(n) * (K / G),
K, G, a, gsum, isa, y + n);
for (; n < end; ++n)
y[n] = row_q4(codes + size_t(n) * K / 2, scales + size_t(n) * (K / G), offsets + size_t(n) * (K / G),
K, G, a, gsum, isa);
});
}
for (auto& th : pool) th.join();
}
The PyTorch extension (izh/kernels/cpu_ext.cpp) wraps the same functions, splits rows across threads with at::parallel_for (which needs -fopenmp, or it silently runs on one thread, the third lesson of this section’s measurements), and is compiled on first use by torch.utils.cpp_extension.load.
Measured
python run.py cpu
On the 4-core Xeon virtual machine used to write this book, with Qwen3-0.6B’s MLP weights for all 28 layers (528 million weights, far beyond the caches, as in real decoding):
{"kind": "cpu", "name": "Intel(R) Xeon(R) Processor @ 2.10GHz", "capability": null, "memory_gib": 15.7, "bf16": true, "triton": false, "graphs": false, "simd": "avx512_vnni"}
{"measured_copy_bandwidth_GBs": 40.6, "threads": 4}
{"kernel": "torch fp32", "weights_MB": 1057, "ms_per_token": 30.5, "effective_GBs": 34.6, "mlp_tokens_per_s": 32.7}
{"kernel": "torch bf16", "weights_MB": 528, "ms_per_token": 28.3, "effective_GBs": 18.7, "mlp_tokens_per_s": 35.3}
{"kernel": "SIMD int8 codes", "weights_MB": 297, "ms_per_token": 27.2, "effective_GBs": 10.9, "mlp_tokens_per_s": 36.7}
{"kernel": "SIMD int4 codes", "weights_MB": 165, "ms_per_token": 12.7, "effective_GBs": 13.0, "mlp_tokens_per_s": 78.8}
PyTorch’s FP32 matrix-vector product (Intel’s oneDNN underneath) reads its weights at 35 GB/s, close to the 41 GB/s this machine can copy, so it’s bandwidth-bound, as decode should be. BF16 halves the bytes but saves little time: this CPU has no native BF16 arithmetic path for this shape. The int4 kernel reads 6.4 times fewer bytes than FP32 and runs 2.4 times faster. It would be 6 times faster at full bandwidth; what stands in the way is overhead per call (56 Python-level calls per token, each quantizing its activations and launching threads) and the arithmetic per weight. That gap is exactly why llama.cpp runs its whole graph in C++ with a persistent thread pool: stretch exercise 1 moves the decode loop into the Rust track to close it.
Apple GPUs
PyTorch runs on Apple GPUs through the mps device, and every model in this book runs there unchanged with the reference attention backend: Triton doesn’t target Metal. Apple’s GPUs share unified memory with the CPU, so a Mac with 64 GB of RAM can hold a 64 GB model on its GPU without any offloading, which is why Macs became popular for local inference despite modest compute.
For quantized weights, the engine needs a Metal kernel. PyTorch (2.6 and later) compiles Metal shading language at run time with torch.mps.compile_shader; here’s the 4-bit GEMV, one SIMD group of 32 threads per output row, with simd_sum adding their partial sums:
"""A Metal kernel for Apple GPUs (Chapter 40): the 4-bit group-affine GEMV of AffineQuantLinear,
compiled at run time with torch.mps.compile_shader (PyTorch 2.6+).
Apple's GPUs share memory with the CPU (unified memory), so a model's weights are read at the
memory system's full bandwidth by either side; decode is bandwidth-bound exactly as on a discrete
GPU, and keeping weights in 4 bits matters just as much.
One SIMD group (32 threads) per output row: each thread takes every 32nd byte of the row, and
simd_sum adds the 32 partial sums. Not validated in the book's test environment (no Apple GPU);
the test is skipped there and runs on any Mac with MPS.
"""
from functools import lru_cache
SOURCE = r"""
#include <metal_stdlib>
using namespace metal;
kernel void affine_gemv_q4(device const uchar* codes [[buffer(0)]],
device const float* scales [[buffer(1)]],
device const float* offsets [[buffer(2)]],
device const float* x [[buffer(3)]],
device float* y [[buffer(4)]],
constant uint& K [[buffer(5)]],
constant uint& G [[buffer(6)]],
uint tid [[thread_position_in_grid]],
uint lane [[thread_index_in_simdgroup]]) {
uint row = tid / 32;
uint groups = K / G;
float acc = 0.0f;
for (uint j = lane; j < K / 2; j += 32) { // byte j holds codes 2j (low) and 2j + 1 (high)
uchar b = codes[row * (K / 2) + j];
uint k = 2 * j, g = k / G;
float s = scales[row * groups + g], o = offsets[row * groups + g];
float x0 = x[k], x1 = x[k + 1];
acc += (float(b & 15) * s + o) * x0 + (float(b >> 4) * s + o) * x1;
}
acc = simd_sum(acc);
if (lane == 0) y[row] = acc;
}
"""
@lru_cache(maxsize=1)
def library():
import torch
return torch.mps.compile_shader(SOURCE)
def affine_gemv_q4(x, codes, scales, offsets, group_size):
"""x [K] float32 on "mps" -> y [N]."""
import torch
n, k = codes.shape[0], x.shape[0]
y = torch.empty(n, device=x.device, dtype=torch.float32)
library().affine_gemv_q4(codes, scales.float(), offsets.float(), x.float(), y, k, group_size,
threads=n * 32, group_size=32)
return y
This kernel is not validated in this book’s test environment, which has no Apple hardware; its test runs on any Mac with MPS and is skipped elsewhere (Appendix F). Apple’s own MLX framework, and llama.cpp’s Metal backend, go further with fused kernels for each quantization type; studying the latter is stretch exercise 3.
What about Vulkan?
llama.cpp’s Vulkan backend runs the same GGUF kernels on any GPU with a Vulkan driver: AMD and Intel consumer cards, older NVIDIA cards, many phones. This engine doesn’t have one, and the reason is architectural rather than a matter of effort: everything here runs through PyTorch, which has no Vulkan device, and through Triton, which has no Vulkan target. A Vulkan backend would be a second kernel set (compute shaders in GLSL or Slang for every operation of Chapters 31-39) plus memory management that PyTorch currently does for you. If Vulkan devices matter to you, llama.cpp’s ggml-vulkan is the reference to read, and the GGUF support of Chapter 38 lets this engine and llama.cpp share model files.
Models bigger than the GPU
A 30-billion-parameter MoE at 4 bits is 17 GB; a consumer GPU has 8-24 GB. Three ways to use host memory, each with a different cost:
Splitting layers
Put the last $n$ layers on the GPU and the rest on the CPU (llama.cpp’s --n-gpu-layers). Each token’s hidden state crosses PCIe twice, a few kilobytes, which costs microseconds. But the CPU layers run at CPU bandwidth: with 40% of the layers on a CPU that’s 10 times slower, the model runs at about a fifth of its full-GPU speed. It’s the right tool when the model almost fits.
class SplitModel(FlatModel):
"""Layers on several devices; the hidden state follows them."""
def __init__(self, model, layer_devices, head_device=None, backends=None):
super().__init__(model)
if len(layer_devices) != len(self.layers):
raise ValueError("One device per layer")
self.devices = [torch.device(d) for d in layer_devices]
self.head_device = torch.device(head_device or layer_devices[-1])
self.backends = backends or {}
for layer, device in zip(self.layers, self.devices):
layer.to(device)
self.backbone.embed_tokens.to(self.devices[0])
self.backbone.norm.to(self.head_device)
self.model.lm_head.to(self.head_device)
def layer_device(self, i):
return self.devices[i]
def forward(self, input_ids, positions, kv_caches, meta, backend, embeds=None, **features):
"""FlatModel.forward, with the hidden state moved to each layer's device. (Your engine: Chapter 40)"""
x = self.backbone.embed_tokens(input_ids.to(self.devices[0])) if embeds is None else embeds
delta, per_device = None, {}
for i, layer in enumerate(self.layers):
device = self.devices[i]
if device not in per_device: # positions, RoPE tables, metadata: once per device
pos = positions.to(device)
per_device[device] = (pos, self.rope(pos, x.dtype), meta_on(meta, device),
self.backends.get(device.type, backend))
pos, rope, local_meta, local_backend = per_device[device]
x = x.to(device, non_blocking=True)
delta = None if delta is None else delta.to(device, non_blocking=True)
h, x = self.add_norm(layer.input_layernorm, x, delta)
delta = self.attention(i, layer.self_attn, h, pos, rope, kv_caches[i], local_meta, local_backend)
h, x = self.add_norm(layer.post_attention_layernorm, x, delta)
delta = self.mlp(layer.mlp, h)
x, delta = x.to(self.head_device), delta.to(self.head_device)
h, _ = self.add_norm(self.backbone.norm, x, delta)
return h
def split_by_count(num_layers, gpu_layers, gpu="cuda", cpu="cpu"):
"""llama.cpp's rule: the LAST gpu_layers layers go to the GPU (nearest the output head)."""
return [cpu] * (num_layers - gpu_layers) + [gpu] * gpu_layers
The model moves the hidden state to each layer’s device, computes the positions, RoPE tables and batch metadata once per device per step, and the runner (Chapter 31) allocates each layer’s KV pool on that layer’s device.
Experts in host memory
A MoE reads only its active experts per token: Qwen3-30B-A3B has 30 B parameters but uses 3 B per token. Keep attention, the router and the shared weights on the GPU, the experts in host memory, and either compute the routed tokens’ experts on the CPU (ktransformers; llama.cpp’s -ot exps=CPU), or copy the experts each step needs into a small GPU cache:
class OffloadedExperts(nn.Module):
"""Stacked expert weights kept in (pinned) host memory.
mode="cpu": send the step's routed rows to the CPU, run the experts there, send results back.
mode="stream": copy the experts this step needs into a GPU cache of `capacity` slots (least
recently used evicted), then compute on the GPU.
"""
def __init__(self, experts, device, mode="stream", capacity=8):
super().__init__()
if mode not in ("cpu", "stream"):
raise ValueError("mode is cpu or stream")
self.mode, self.device, self.capacity = mode, torch.device(device), capacity
pin = self.device.type == "cuda"
host = lambda t: t.detach().cpu().pin_memory() if pin else t.detach().cpu().clone() # noqa: E731
self.host_gate_up, self.host_down = host(experts.gate_up_proj), host(experts.down_proj)
self.slots_gate_up = torch.empty((capacity, *self.host_gate_up.shape[1:]), dtype=self.host_gate_up.dtype,
device=self.device)
self.slots_down = torch.empty((capacity, *self.host_down.shape[1:]), dtype=self.host_down.dtype, device=self.device)
self.slot_of, self.lru = {}, []
self.stats = {"hits": 0, "misses": 0, "bytes_copied": 0}
@property
def num_experts(self):
return self.host_gate_up.shape[0]
def fetch(self, e):
"""Make expert e resident in a GPU slot (LRU eviction); returns the slot. (Your engine: Chapter 40)"""
if e in self.slot_of:
self.stats["hits"] += 1
self.lru.remove(e)
self.lru.append(e)
return self.slot_of[e]
self.stats["misses"] += 1
if len(self.slot_of) < self.capacity:
slot = len(self.slot_of)
else:
victim = self.lru.pop(0)
slot = self.slot_of.pop(victim)
self.slots_gate_up[slot].copy_(self.host_gate_up[e], non_blocking=True)
self.slots_down[slot].copy_(self.host_down[e], non_blocking=True)
self.stats["bytes_copied"] += self.host_gate_up[e].nbytes + self.host_down[e].nbytes
self.slot_of[e] = slot
self.lru.append(e)
return slot
def expert_weights(self, e):
if self.mode == "cpu":
return self.host_gate_up[e], self.host_down[e]
slot = self.fetch(e)
return self.slots_gate_up[slot], self.slots_down[slot]
def forward_grouped(self, x, weights, experts):
device = x.device
if self.mode == "cpu":
x, weights, experts = x.cpu(), weights.cpu(), experts.cpu()
out = torch.zeros_like(x)
# Experts run one after another, each fetched just before it's used, so a step may route to
# more experts than the cache holds: a slot is reused once its last expert has run (on one
# CUDA stream, the copy into it is ordered after that computation).
for e in experts.unique().tolist():
token, slot = torch.where(experts == e)
gate_up, down = self.expert_weights(e)
gate, up = nn.functional.linear(x[token], gate_up.to(x.dtype)).chunk(2, dim=-1)
y = nn.functional.linear(nn.functional.silu(gate) * up, down.to(x.dtype))
out.index_add_(0, token, y * weights[token, slot, None])
return out.to(device, non_blocking=True)
forward_loop = forward_grouped
def offload_experts(model, device, mode="stream", capacity=8):
"""Replace every MoE layer's experts; returns the new modules (for their stats)."""
from .moe import SparseMoeBlock
replaced = []
for module in model.modules():
if isinstance(module, SparseMoeBlock):
module.experts = OffloadedExperts(module.experts, device, mode, capacity)
replaced.append(module.experts)
return replaced
Whether copying pays depends on reuse:
python run.py offload --new-tokens 24
{"requests": 1, "gpu_slots_per_layer": "4 of 32", "hit_rate": 0.076, "MB_copied_per_token": 3.167, "all_experts_MB": 50.3}
{"requests": 1, "gpu_slots_per_layer": "8 of 32", "hit_rate": 0.174, "MB_copied_per_token": 2.83, "all_experts_MB": 50.3}
{"requests": 1, "gpu_slots_per_layer": "16 of 32", "hit_rate": 0.377, "MB_copied_per_token": 2.135, "all_experts_MB": 50.3}
{"requests": 1, "gpu_slots_per_layer": "32 of 32", "hit_rate": 0.74, "MB_copied_per_token": 0.892, "all_experts_MB": 50.3}
{"requests": 4, "gpu_slots_per_layer": "4 of 32", "hit_rate": 0.0, "MB_copied_per_token": 2.38, "all_experts_MB": 50.3}
{"requests": 4, "gpu_slots_per_layer": "8 of 32", "hit_rate": 0.0, "MB_copied_per_token": 2.38, "all_experts_MB": 50.3}
{"requests": 4, "gpu_slots_per_layer": "16 of 32", "hit_rate": 0.363, "MB_copied_per_token": 1.517, "all_experts_MB": 50.3}
{"requests": 4, "gpu_slots_per_layer": "32 of 32", "hit_rate": 0.906, "MB_copied_per_token": 0.225, "all_experts_MB": 50.3}
The test model’s router is random, so a single request’s hit rate is about the fraction of experts cached: there’s no locality to exploit. With four requests in a batch, each step routes to most experts, and a small LRU cache thrashes: every expert is evicted just before it’s needed again, and the hit rate drops to zero. Real MoEs route with more locality than a random router, but the lesson holds: streaming experts suits single-user decoding with a cache that holds a good fraction of the experts; for batched serving, computing experts where they live (on the CPU, or on other GPUs with expert parallelism, Chapter 41) is better.
Streaming layers
When nothing else fits, keep every layer in pinned host memory and stream them through two GPU slots, copying layer $i + 1$ while layer $i$ computes:
class StreamedModel(FlatModel):
"""All layers in host memory; two GPU slots, the next layer copied while this one computes."""
def __init__(self, model, device):
super().__init__(model)
self.device = torch.device(device)
pin = self.device.type == "cuda"
self.host_layers = self.layers
for layer in self.host_layers:
layer.to("cpu")
if pin:
for p in layer.parameters():
p.data = p.data.pin_memory()
self.slots = [copy.deepcopy(self.host_layers[0]).to(self.device) for _ in range(2)]
self.backbone.embed_tokens.to(self.device)
self.backbone.norm.to(self.device)
self.model.lm_head.to(self.device)
self.copy_stream = torch.cuda.Stream(self.device) if self.device.type == "cuda" else None
self.ready = [None, None]
self.bytes_copied = 0
def layer_device(self, i):
return self.device
def load(self, i):
"""Copy layer i's weights into slot i % 2, on the copy stream when there is one."""
slot = self.slots[i % 2]
if self.copy_stream is not None:
self.copy_stream.wait_stream(torch.cuda.current_stream(self.device)) # the slot's last reader is done
with torch.cuda.stream(self.copy_stream):
for dst, src in zip(slot.parameters(), self.host_layers[i].parameters()):
dst.data.copy_(src.data, non_blocking=True)
self.ready[i % 2] = torch.cuda.Event()
self.ready[i % 2].record(self.copy_stream)
else:
for dst, src in zip(slot.parameters(), self.host_layers[i].parameters()):
dst.data.copy_(src.data)
self.bytes_copied += sum(p.nbytes for p in self.host_layers[i].parameters())
def forward(self, input_ids, positions, kv_caches, meta, backend, embeds=None, **features):
"""Run layer i from slot i % 2 while layer i + 1 loads into the other. (Your engine: Chapter 40)"""
x = self.backbone.embed_tokens(input_ids) if embeds is None else embeds
rope = self.rope(positions, x.dtype)
delta = None
self.load(0)
for i in range(len(self.host_layers)):
if self.ready[i % 2] is not None:
torch.cuda.current_stream(self.device).wait_event(self.ready[i % 2])
if i + 1 < len(self.host_layers):
self.load(i + 1) # overlaps with this layer's compute
layer = self.slots[i % 2]
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)
h, _ = self.add_norm(self.backbone.norm, x, delta)
return h
Each step moves the whole model over PCIe: a 70 B model at 4 bits is 35 GB, a bit over a second per step at 25-30 GB/s. That’s hopeless for interactive decode and fine for batch work: a 4,096-token prefill, or 256 sequences decoded together, amortize the same 35 GB over thousands of tokens. FlexGen (Sheng et al., 2023) built a whole throughput-oriented engine on this observation.
Build it
Engine milestone 40: everywhere. Implement detect and engine_defaults in engine/platform.py; SplitModel.forward, OffloadedExperts.fetch, StreamedModel.load and StreamedModel.forward in engine/offload.py. The CPU kernels (cpp/izh_cpu.hpp, rust/src/simd.rs) and the Metal shader are provided as the C++ and Rust tracks’ contribution; build and check them with build/cpp/izh 40 and cargo test --release simd.
pytest tests/test_ch40_everywhere.py
python run.py cpu --impl engine
python run.py offload --impl engine
The tests check platform detection and defaults; the SIMD kernels’ 4-bit and 8-bit paths against exact matmuls; a model split across devices, offloaded experts in both modes, and streamed layers, each producing exactly the plain engine’s tokens, with stream statistics and bytes copied; and, on a GPU, a real GPU/CPU split, and on a Mac, the Metal kernel.
Stretch exercises
- ★★★ Move decode into the Rust track: give
rust/src/qwen3.rsa path that keeps its linear layers as 4-bit codes and usessimd::gemv_q4with a persistent thread pool. Compare tokens per second withrun.py cpu’s Python-driven kernel, and with llama.cpp on the same GGUF file. Where:rust/src/qwen3.rsandrust/src/simd.rs. - ★★ Prefill on a CPU is compute-bound: write a tiled int8 GEMM (4 × 4 output tiles, VNNI) for batches of 64 tokens, or use Intel AMX tiles on Sapphire Rapids. Measure prompt tokens per second. Where: CPU SIMD kernels in
cpp/izh_cpu.hpporrust/src/simd.rs; expose the C++ path throughengine/kernels/cpu_ext.cppandengine/kernels/cpu.py. - ★★ Read llama.cpp’s Metal kernel for Q4_K (
ggml-metal.metal) and port its structure tokernels/metal.py. What does it do that the simple SIMD-group kernel doesn’t? Where: the Metal shader and launcher inengine/kernels/metal.py. - ★★ Prefetch experts: run the next layer’s router on the current layer’s output (an approximation some MoEs make exact with a pre-gating design) and start copying its likely experts early, on a second CUDA stream. Measure the stall time saved. Where:
OffloadedExperts.fetchand the cache lifecycle inengine/offload.py, with scheduling hooks in the model’s layer loop.
Check your understanding
- Why does a CPU decode at a rate set by memory bandwidth, and what does that imply about quantization on CPUs?
- Why is dequantizing the whole weight matrix before a matmul slower than not quantizing at all?
- What does quantizing the activations to int8 buy, and what does it cost?
- Why are the SIMD functions compiled with per-function
targetattributes instead of-march=native? - Why does an LRU expert cache with 8 of 32 slots thrash when four requests share a step?
- When does streaming every layer through the GPU make sense, and when doesn’t it?
Going deeper
- llama.cpp’s
ggml/src/ggml-cpu/(thevec_dot_q4_0_q8_0family for AVX2, AVX-512, NEON and others) andggml-metal/,ggml-vulkan/for its other backends. - Intel’s Intel® 64 and IA-32 Architectures Optimization Reference Manual (VNNI, AMX) and Arm’s Neon Programmer’s Guide (
sdot); Agner Fog’s Optimizing software in C++ for the latency and throughput reasoning of this chapter. - Sheng et al., FlexGen: High-Throughput Generative Inference of Large Language Models with a Single GPU (ICML 2023); Chen et al., KTransformers (SOSP 2025), for CPU/GPU hybrid MoE inference.
- AMD’s ROCm documentation for PyTorch and Triton, and vLLM’s ROCm notes (AITER, FP8 formats on MI300 and MI350); Apple’s Metal Shading Language Specification and the MLX documentation.
41. Many GPUs: tensor, pipeline, expert and data parallelism
In this chapter
- Why one GPU stops being enough, and the four ways to split serving across many: tensor, pipeline, expert and data parallelism, each with its communication bill.
- Tensor parallelism Megatron-style: column- and row-parallel layers, one all-reduce per sublayer, split KV heads and a split vocabulary.
- One scheduler driving many ranks in lockstep, as vLLM and SGLang do.
- Expert parallelism with all-to-all dispatch, a prefix-aware router for data-parallel replicas, prefill and decode on different machines with the KV cache shipped between them, and ring attention for contexts too long for one GPU.
You will build
engine/parallel.py: the parallel layers, shard_tensor_parallel, PipelineStage.forward, ExpertParallelMoE.forward, worker_loop, PrefixRouter.route, KVExporter.export, import_kv and ring_attention.
Time: 8-10 hours. GPU: not needed: every test runs several processes on the CPU with PyTorch's gloo backend. Two or more GPUs with NCCL run the same code.
When one GPU isn’t enough
There are three reasons to spread a model over several GPUs, and they lead to different designs:
- It doesn’t fit. Llama-3.1-70B in BF16 is 141 GB of weights; an H100 has 80 GB. Even when the weights fit, the KV cache needs room: every gigabyte left over is more concurrent requests (Chapter 31).
- It’s too slow. Decode reads every weight once per step (Chapter 10). Split the weights over 8 GPUs and each reads an eighth: the step takes (nearly) an eighth of the time, because the GPUs’ memory bandwidths add up.
- There’s too much traffic. One replica serves at most so many tokens per second. More replicas serve more.
The four kinds of parallelism split different things:
| splits | communication per layer | good for | |
|---|---|---|---|
| tensor (TP) | every weight matrix, across GPUs | 2 all-reduces of the activations | latency, fitting big models inside one NVLink domain |
| pipeline (PP) | the layers, into stages | one hidden vector per token per stage boundary | fitting models across nodes with slower links |
| expert (EP) | a MoE’s experts | 2 all-to-alls (or one all-reduce) | big MoEs: DeepSeek-V3, Qwen3-235B, Kimi K2 |
| data (DP) | the requests, over whole replicas | none (a router in front) | throughput |
Production deployments combine them: DeepSeek-V3 across 18 nodes runs attention data-parallel and experts expert-parallel; a Llama-405B deployment runs TP 8 inside each node and PP 2 across two. vLLM and SGLang expose all four with flags (--tensor-parallel-size, --pipeline-parallel-size, --enable-expert-parallel, --data-parallel-size); this chapter builds each one small enough to read.
The cost of talking
GPUs in one server talk over NVLink (900 GB/s per H100, each direction); servers talk over InfiniBand or RoCE (400 Gb/s = 50 GB/s per NIC, usually one NIC per GPU); consumer GPUs talk over PCIe (32-64 GB/s), often through the CPU. A ring all-reduce of $S$ bytes over $n$ GPUs sends $2\frac{n-1}{n}S$ bytes from each GPU and takes $2(n-1)$ latency-bound hops. For a decode step those messages are small (64 tokens × 8,192 features × 2 bytes = 1 MB), so latency, not bandwidth, sets their cost: a few microseconds per all-reduce on NVLink, which is why vLLM ships a custom one-shot all-reduce for small messages instead of NCCL’s ring, and why TP across PCIe or between nodes is slow.
The communicator
Four collectives cover the chapter: all-reduce (everyone gets the sum), all-gather (everyone gets everyone’s pieces), all-to-all (rank $i$ sends piece $j$ to rank $j$), and point-to-point send/receive.
class Comm:
"""The process group's collectives. With one process (or none initialized) each is a no-op."""
def __init__(self):
self.enabled = dist.is_available() and dist.is_initialized()
self.rank = dist.get_rank() if self.enabled else 0
self.world = dist.get_world_size() if self.enabled else 1
self.nccl = self.enabled and dist.get_backend() == "nccl"
self.traffic = {} # op -> [calls, bytes sent by this rank]
def count(self, op, t):
entry = self.traffic.setdefault(op, [0, 0])
entry[0] += 1
entry[1] += t.numel() * t.element_size()
def all_reduce(self, t, op=None):
"""Every rank ends with the elementwise sum (in place)."""
if self.world > 1:
self.count("all_reduce", t)
dist.all_reduce(t, op=op or dist.ReduceOp.SUM)
return t
def all_gather(self, t, dim=-1):
"""Every rank's tensor, concatenated along dim, on every rank."""
if self.world == 1:
return t
self.count("all_gather", t)
parts = [torch.empty_like(t) for _ in range(self.world)]
dist.all_gather(parts, t.contiguous())
return torch.cat(parts, dim=dim)
def all_to_all(self, chunks):
"""chunks[p] goes to rank p; returns the chunks every rank sent to this one. Row counts may
differ, so they are exchanged first. NCCL has a native all-to-all; gloo gets the same
result from pairwise sends and receives."""
if self.world == 1:
return [chunks[0]]
sizes = torch.tensor([c.shape[0] for c in chunks], dtype=torch.int64, device=chunks[0].device)
table = self.all_gather(sizes[None], dim=0) # table[src, dst] = rows src sends dst
tail, dtype = chunks[0].shape[1:], chunks[0].dtype
received = [torch.empty((int(table[src, self.rank]), *tail), dtype=dtype, device=chunks[0].device)
for src in range(self.world)]
if self.nccl:
dist.all_to_all(received, [c.contiguous() for c in chunks])
return received
received[self.rank] = chunks[self.rank]
ops = []
for peer in range(self.world):
if peer != self.rank:
ops.append(dist.isend(chunks[peer].contiguous(), peer))
ops.append(dist.irecv(received[peer], peer))
for op in ops:
op.wait()
return received
def broadcast_object(self, obj=None, src=0):
if self.world == 1:
return obj
box = [obj]
dist.broadcast_object_list(box, src)
return box[0]
def min(self, value):
t = torch.tensor([value], dtype=torch.int64, device="cuda" if self.nccl else "cpu")
return int(self.all_reduce(t, dist.ReduceOp.MIN).item())
def send(self, t, dst):
self.count("send", t)
dist.send(t.contiguous(), dst)
def recv(self, shape, dtype, src, device):
t = torch.empty(shape, dtype=dtype, device=device)
dist.recv(t, src)
return t
torch.distributed provides all of them on two backends: NCCL for GPUs and gloo for CPUs. Gloo lacks an all-to-all with uneven pieces, so all_to_all builds one from paired sends and receives. traffic counts calls and bytes, which the measurements below use.
spawn(fn, world) starts world processes on this machine, joined by one process group, the way torchrun does across machines; every test in this chapter runs through it.
Tensor parallelism
Megatron-LM’s observation (Shoeybi et al., 2019) is that a transformer’s matmuls come in pairs that split without communication in between. Take an MLP, $y = W_{down},\sigma(W_{up},x)$. Split $W_{up}$ by output rows across $n$ GPUs (column parallelism, after the column blocks of $W^T$): GPU $i$ computes features $[iI/n, (i+1)I/n)$ of the intermediate activation, needing all of $x$ but nothing from the other GPUs. The activation is elementwise, so each GPU applies it to its own slice. Then split $W_{down}$ by input columns (row parallelism): GPU $i$ multiplies its slice of the activation by its columns, producing a full-width partial sum. One all-reduce adds the partial sums, and every GPU holds $y$.
Attention splits the same way, by heads: Q, K and V are column-parallel (each GPU computes its heads), attention runs per head with no communication, and $W_o$ is row-parallel. So each layer costs two all-reduces, one after attention and one after the MLP, and every GPU holds $1/n$ of the weights and, because it only computes its own heads, $1/n$ of the KV cache.
@torch.no_grad()
def take_rows(linear, start, end):
"""Column parallelism: this rank computes output features [start, end) (no communication)."""
out = nn.Linear(linear.in_features, end - start, bias=linear.bias is not None,
device=linear.weight.device, dtype=linear.weight.dtype)
out.weight.copy_(linear.weight[start:end])
if linear.bias is not None:
out.bias.copy_(linear.bias[start:end])
return out
@torch.no_grad()
def take_cols(linear, start, end):
"""The matching row-parallel half: input features [start, end), a PARTIAL sum of the output."""
out = nn.Linear(end - start, linear.out_features, bias=False, device=linear.weight.device, dtype=linear.weight.dtype)
out.weight.copy_(linear.weight[:, start:end])
return out
class RowParallelLinear(nn.Module):
"""Partial products summed across ranks: the one all-reduce per sublayer."""
def __init__(self, linear, start, end, comm):
super().__init__()
self.local, self.comm = take_cols(linear, start, end), comm
self.bias = None if linear.bias is None else nn.Parameter(linear.bias.detach().clone())
def forward(self, x):
"""(Your engine: Chapter 41)"""
y = self.comm.all_reduce(self.local(x))
return y if self.bias is None else y + self.bias # added once, after the sum
def vocab_range(vocab, comm):
per = -(-vocab // comm.world) # padded: every rank holds `per` rows
return per, comm.rank * per, min(vocab, (comm.rank + 1) * per)
def vocab_slice(weight, comm):
per, lo, hi = vocab_range(weight.shape[0], comm)
local = weight.new_zeros((per, weight.shape[1]))
local[:hi - lo] = weight[lo:hi]
return nn.Parameter(local)
class VocabParallelEmbedding(nn.Module):
"""Each rank holds a slice of the table; tokens outside it look up zeros, and the all-reduce
assembles every row."""
def __init__(self, weight, vocab, comm):
super().__init__()
self.weight, self.comm = weight, comm # weight: this rank's slice (vocab_slice)
_, self.lo, self.hi = vocab_range(vocab, comm)
def forward(self, ids):
"""(Your engine: Chapter 41)"""
mine = (ids >= self.lo) & (ids < self.hi)
rows = F.embedding(torch.where(mine, ids - self.lo, 0), self.weight)
return self.comm.all_reduce(rows * mine[..., None].to(rows.dtype))
class ParallelLMHead(nn.Module):
"""Each rank scores its vocabulary slice; an all-gather assembles the full logits."""
def __init__(self, weight, vocab, comm):
super().__init__()
self.weight, self.vocab, self.comm = weight, vocab, comm
def forward(self, h):
"""(Your engine: Chapter 41)"""
return self.comm.all_gather(F.linear(h, self.weight), dim=-1)[..., :self.vocab]
Three details:
- The bias of a row-parallel layer is added once, after the all-reduce; added before, it would be counted $n$ times.
- The vocabulary is split too. Qwen3’s 151,936 × 4,096 embedding is 1.2 GB in BF16, worth splitting. Each rank holds a slice of rows; a token outside the slice looks up zeros, and the all-reduce assembles every token’s row. The LM head scores each rank’s slice and an all-gather builds the full logits. (vLLM all-gathers to one rank only, the one that samples; we gather everywhere, which is simpler and costs little.) The vocabulary is padded to a multiple of $n$.
- KV heads. Qwen3-8B has 32 query heads but 8 KV heads. Over 8 GPUs, each gets 4 query heads and 1 KV head. Over 16 GPUs, each gets 2 query heads, and they share a KV head with a neighbour: KV heads are replicated when there are fewer than ranks. Both ranks then cache the same KV head, so KV memory stops shrinking past $n = H_{kv}$.
@torch.no_grad()
def shard_tensor_parallel(model, comm):
"""Keep this rank's share of every layer. (Your engine: Chapter 41)
Attention: query heads split evenly; KV heads split too while there are at least as many as
ranks, and replicated beyond that (8 KV heads over 16 GPUs: each KV head on two ranks). Q, K
and V are column-parallel, o_proj row-parallel. MLP: gate/up column-parallel, down
row-parallel. Embedding and LM head: split by vocabulary. Norms and the router: replicated.
"""
from .moe import SparseMoeBlock
flat = model if isinstance(model, FlatModel) else FlatModel(model)
if hasattr(flat.model, "kv_spec"):
raise ValueError("Latent attention (Chapter 42) has one KV head: serve it with data-parallel attention")
c, tp, r = flat.cfg, comm.world, comm.rank
h, hkv, d = c.num_attention_heads, c.num_key_value_heads, c.head_dim
if h % tp or (hkv % tp and tp % hkv):
raise ValueError(f"{h} query / {hkv} KV heads can't be split over {tp} ranks")
hq, hk = h // tp, max(hkv // tp, 1)
kv0 = r * hk if hkv >= tp else (r * hq) // (h // hkv) # first KV head this rank needs
for layer in flat.layers:
a = layer.self_attn
if getattr(a, "qkv_proj", None) is not None:
raise ValueError("Shard before fuse(): fusing concatenates the projections this splits")
a.q_proj = take_rows(a.q_proj, r * hq * d, (r + 1) * hq * d)
a.k_proj = take_rows(a.k_proj, kv0 * d, (kv0 + hk) * d)
a.v_proj = take_rows(a.v_proj, kv0 * d, (kv0 + hk) * d)
a.o_proj = RowParallelLinear(a.o_proj, r * hq * d, (r + 1) * hq * d, comm)
m = layer.mlp
if isinstance(m, SparseMoeBlock):
layer.mlp = TensorParallelMoE(m, comm)
else:
lo, hi = shard_bounds(m.down_proj.in_features, comm)
m.gate_proj, m.up_proj = take_rows(m.gate_proj, lo, hi), take_rows(m.up_proj, lo, hi)
m.down_proj = RowParallelLinear(m.down_proj, lo, hi, comm)
embed = flat.backbone.embed_tokens.weight
tied = flat.model.lm_head.weight is embed
table = vocab_slice(embed, comm)
flat.backbone.embed_tokens = VocabParallelEmbedding(table, c.vocab_size, comm)
head = table if tied else vocab_slice(flat.model.lm_head.weight, comm)
flat.model.lm_head = ParallelLMHead(head, c.vocab_size, comm)
flat.cfg = replace(c, num_attention_heads=hq, num_key_value_heads=hk)
return flat
shard_tensor_parallel keeps only this rank’s share of each layer and records the local head counts in flat.cfg, so FlatModel.attention reshapes into the local heads, kv_spec reports the local KV heads, and the runner allocates a pool of the local size. Nothing else in the engine changes: the paged attention kernels from Chapter 32 run per head and never know other heads exist.
Sharding happens before fuse() (Chapter 33), because fusing concatenates exactly the projections that sharding splits. A real loader also never builds the full model on every rank: safetensors’ get_slice reads only the rows a rank keeps (stretch exercise 1).
MoE layers under tensor parallelism
A MoE layer could split each expert’s intermediate dimension like a dense MLP. With 128 experts of intermediate size 768 over 8 GPUs that leaves 96-wide slivers, too thin for efficient matmuls. vLLM’s --enable-expert-parallel instead gives each rank whole experts: every rank sees every token (attention is tensor parallel, so after its all-reduce every rank holds the same activations), runs the router, computes only the experts it owns, and the same single all-reduce adds everyone’s contributions. A shared expert splits like a dense MLP, and its partial sum joins the same all-reduce:
def local_experts(x, weights, experts, gate_up, down, first):
"""The routed experts this rank owns (first .. first + len(gate_up) - 1); others contribute
zero here and arrive in the all-reduce. A fused MoE kernel does the same with an expert map
that sends non-local experts to -1."""
out = torch.zeros_like(x)
owned = (experts >= first) & (experts < first + gate_up.shape[0])
for e in experts[owned].unique().tolist():
token, slot = torch.where(experts == e)
gate, up = F.linear(x[token], gate_up[e - first]).chunk(2, dim=-1)
y = F.linear(F.silu(gate) * up, down[e - first])
out.index_add_(0, token, y * weights[token, slot, None])
return out
class TensorParallelMoE(nn.Module):
"""A MoE layer under tensor parallelism: every rank sees every token, so experts are split
whole across ranks (expert parallelism inside the TP group) and the shared expert is split
by columns like a dense MLP. One all-reduce sums routed and shared partials together."""
def __init__(self, block, comm):
super().__init__()
e = block.experts.num_experts
if e % comm.world:
raise ValueError(f"{e} experts don't split over {comm.world} ranks")
per = e // comm.world
self.comm, self.first, self.gate = comm, comm.rank * per, block.gate
self.gate_up = nn.Parameter(block.experts.gate_up_proj[self.first:self.first + per].detach().clone())
self.down = nn.Parameter(block.experts.down_proj[self.first:self.first + per].detach().clone())
self.shared, self.shared_gate = None, block.shared_expert_gate
if block.shared_expert is not None:
s = block.shared_expert
lo, hi = shard_bounds(s.down_proj.in_features, comm)
self.shared = nn.ModuleDict({"gate_proj": take_rows(s.gate_proj, lo, hi), "up_proj": take_rows(s.up_proj, lo, hi),
"down_proj": take_cols(s.down_proj, lo, hi)})
def forward(self, x):
"""(Your engine: Chapter 41)"""
_, weights, experts = self.gate(x)
out = local_experts(x, weights, experts, self.gate_up, self.down, self.first)
if self.shared is not None:
s = self.shared
partial = s["down_proj"](F.silu(s["gate_proj"](x)) * s["up_proj"](x))
out = out + torch.sigmoid(self.shared_gate(x)) * partial # the gate is replicated: scaling commutes with the sum
return self.comm.all_reduce(out)
One scheduler, many workers
Who runs the scheduler? If every rank ran its own, they would have to agree on every decision (which requests, which blocks, which sampled token), which is fragile. vLLM and SGLang put the scheduler, block manager and sampler on one process and make the GPU workers dumb: each step, the driver broadcasts what to run, and every worker runs the same forward pass on its shard, meeting the others at each collective.
def portable(batch, device="cpu"):
"""The batch with its tensors on `device`, ready to pickle and broadcast."""
from .offload import meta_on
move = lambda t: t.to(device) if isinstance(t, torch.Tensor) else t # noqa: E731
extra = {k: move(v) for k, v in batch.extra.items()} if batch.extra else batch.extra
return replace(batch, input_ids=batch.input_ids.to(device), positions=batch.positions.to(device),
logits_indices=batch.logits_indices.to(device), meta=meta_on(batch.meta, device),
prompt_spans=None, extra=extra)
class _Swapped(dict):
def __init__(self, runner):
super().__init__()
self.runner = runner
def pop(self, key, default=None):
self.runner.comm.broadcast_object(("drop", key))
return self.runner.inner.swapped.pop(key, default)
class DistributedRunner:
"""Rank 0's runner. Every call that touches device state is broadcast first, so all ranks
run the same step in lockstep: the scheduler, block manager and sampler exist only here."""
def __init__(self, runner, comm):
self.inner, self.comm = runner, comm
self.swapped = _Swapped(self)
def __getattr__(self, name):
return getattr(self.inner, name)
def execute(self, batch):
self.comm.broadcast_object(("execute", portable(batch)))
return self.inner.execute(batch)
def swap_out(self, request, blocks):
self.comm.broadcast_object(("swap_out", request.request_id, blocks))
self.inner.swap_out(request, blocks)
def swap_in(self, request, blocks):
self.comm.broadcast_object(("swap_in", request.request_id, blocks))
self.inner.swap_in(request, blocks)
def shutdown(self):
self.comm.broadcast_object(("stop",))
def worker_loop(runner, comm):
"""Ranks 1..n-1: repeat whatever rank 0 does, until it says stop. (Your engine: Chapter 41)"""
while True:
command = comm.broadcast_object(None)
kind = command[0]
if kind == "execute":
runner.execute(portable(command[1], runner.device)) # collectives inside match rank 0's
elif kind in ("swap_out", "swap_in"):
getattr(runner, kind)(SimpleNamespace(request_id=command[1]), command[2])
elif kind == "drop":
runner.swapped.pop(command[1], None)
elif kind == "stop":
return
def parallel_engine(model, comm, config=None, mode="tensor", bounds=None, **engine_options):
"""Every rank calls this with the same model. Rank 0 gets an EngineCore to drive (call
engine.runner.shutdown() when done); the other ranks serve inside this call and get None."""
from .serve.engine import EngineCore, EngineConfig
from .serve.attention import get_backend
from .serve.runner import ModelRunner, num_blocks_for_memory
config = config or EngineConfig()
flat = shard_tensor_parallel(model, comm) if mode == "tensor" else PipelineStage(model, comm, bounds)
layers, kv_heads, head_dim = flat.kv_spec()
dtype = next(flat.parameters()).dtype
blocks = config.num_blocks or num_blocks_for_memory(config.kv_cache_bytes, layers, kv_heads, head_dim,
config.block_size, dtype)
config = replace(config, num_blocks=comm.min(blocks), cuda_graphs="off") # every rank: the same pool size
if comm.rank == 0:
engine = EngineCore(flat, config, **engine_options)
engine.runner = DistributedRunner(engine.runner, comm)
return engine
options = {"kv_cache_dtype": config.kv_cache_dtype} if config.kv_cache_dtype != "auto" else {}
worker_loop(ModelRunner(flat, config.num_blocks, config.block_size, get_backend(config.attention_backend, **options)),
comm)
return None
DistributedRunner wraps rank 0’s ModelRunner: execute, swap_out and swap_in are broadcast before they run locally, and everything else (prepare, the KV pool on rank 0) passes through. worker_loop is the whole of a worker. Every rank must have the same number of KV blocks, because block ids in the broadcast batch refer to all their pools at once, so parallel_engine takes the minimum over ranks.
Broadcasting a pickled batch costs a little CPU time per step: vLLM V1 writes the scheduler’s output into a shared-memory ring buffer that workers poll, and each worker builds its own input tensors. CUDA graphs (Chapter 33) capture NCCL all-reduces too, so the per-layer collectives replay with the rest of the step; parallel_engine turns graphs off because the gloo backend can’t be captured.
Pipeline parallelism
Tensor parallelism needs a fast link: two all-reduces per layer per step. Between servers, split the layers instead: rank $r$ owns a contiguous run of layers and forwards one hidden vector per token to rank $r + 1$.
class PipelineStage(FlatModel):
"""Rank r runs layers bounds[r] .. bounds[r+1] - 1. Rank 0 also embeds the tokens and, when
the residual stream comes back from the last stage, applies the final norm (the engine's
runner then applies the LM head there, where the sampler lives)."""
def __init__(self, model, comm, bounds=None):
super().__init__(model)
n, w = len(self.layers), comm.world
self.bounds = bounds or [round(i * n / w) for i in range(w + 1)]
self.comm = comm
mine = self.layers[self.bounds[comm.rank]:self.bounds[comm.rank + 1]]
self.layers = self.backbone.layers = nn.ModuleList(mine) # the other stages' layers are dropped
if comm.rank != 0:
self.backbone.embed_tokens = self.backbone.norm = self.model.lm_head = None
self.width, self.dtype = self.cfg.hidden_size, next(self.parameters()).dtype
def kv_spec(self):
return len(self.layers), self.cfg.num_key_value_heads, self.cfg.head_dim # this stage's layers only
def forward(self, input_ids, positions, kv_caches, meta, backend, embeds=None, **features):
"""This stage's layers; hidden states travel 0 -> 1 -> ... -> last -> 0. (Your engine: Chapter 41)"""
c, n = self.comm, positions.shape[0]
if c.rank == 0:
x = self.backbone.embed_tokens(input_ids) if embeds is None else embeds
else:
x = c.recv((n, self.width), self.dtype, c.rank - 1, positions.device)
rope = self.rope(positions, x.dtype)
delta = None
for i, layer in enumerate(self.layers):
h, x = self.add_norm(layer.input_layernorm, x, delta)
delta = self.attention(i, layer.self_attn, h, positions, rope, kv_caches[i], meta, backend)
h, x = self.add_norm(layer.post_attention_layernorm, x, delta)
delta = self.mlp(layer.mlp, h)
x = x if delta is None else x + delta # one tensor crosses each boundary: the residual stream
if c.world > 1:
c.send(x, (c.rank + 1) % c.world)
if c.rank != 0:
return None
x = c.recv((n, self.width), self.dtype, c.world - 1, positions.device)
h, _ = self.add_norm(self.backbone.norm, x, None)
return h
Only the residual stream crosses a boundary: x + delta, which the next stage’s first add_norm treats as a residual with nothing to add. The last stage sends it back to rank 0, which applies the final norm and the LM head where the sampler lives. (vLLM keeps the head on the last stage and returns sampled tokens instead: a few bytes rather than a hidden vector per sampled row. Sending the hidden state back keeps the engine unchanged here.)
As written, one batch is in flight at a time, so while stage 1 works, stage 0 waits: the pipeline bubble. With $p$ stages, each GPU is busy $1/p$ of the time. PP alone then buys memory, not speed. The fix is to keep $p$ batches in flight, each a step behind the one ahead: vLLM’s scheduler holds up to $p$ “virtual engines” whose batches never share a request. That is the async scheduling of Chapter 33, generalized from 1 to $p$ steps in flight (stretch exercise 2).
Measured
run.py dist runs a 4-layer model (hidden size 512, 8 query and 4 KV heads, 1,024-word vocabulary) with eight 64-token prompts and 16 new tokens each, on 1, 2 and 4 CPU processes over gloo, with the same 64 MB KV budget per rank:
{"mode": "tensor", "ranks": 1, "same_tokens": true, "steps": 16, "kv_blocks_per_rank": 512, "seconds": 0.27, "rank0_traffic": {}}
{"mode": "tensor", "ranks": 2, "same_tokens": true, "steps": 16, "kv_blocks_per_rank": 1024, "seconds": 0.6, "rank0_traffic": {"all_reduce": {"calls": 145, "MB": 11.65}, "all_gather": {"calls": 16, "MB": 0.26}}}
{"mode": "tensor", "ranks": 4, "same_tokens": true, "steps": 16, "kv_blocks_per_rank": 2048, "seconds": 0.79, "rank0_traffic": {"all_reduce": {"calls": 145, "MB": 11.65}, "all_gather": {"calls": 16, "MB": 0.13}}}
{"mode": "pipeline", "ranks": 2, "same_tokens": true, "steps": 16, "kv_blocks_per_rank": 1024, "seconds": 0.41, "rank0_traffic": {"all_reduce": {"calls": 1, "MB": 0.0}, "send": {"calls": 16, "MB": 1.29}}}
{"mode": "pipeline", "ranks": 4, "same_tokens": true, "steps": 16, "kv_blocks_per_rank": 2048, "seconds": 0.7, "rank0_traffic": {"all_reduce": {"calls": 1, "MB": 0.0}, "send": {"calls": 16, "MB": 1.29}}}
- Every configuration produces exactly the single-process tokens. (Exact equality holds here because FP32 sums in a different order still round to the same greedy choices on this model; in BF16 on GPUs, expect rare divergences after near-ties, as with any change of kernel.)
- KV capacity grows with the ranks. With a fixed budget per rank, TP 2 caches twice the blocks (each rank holds half the KV heads) and TP 4 four times. PP stages hold a fraction of the layers, so they gain the same way.
- The traffic matches the arithmetic. TP: 16 steps × (4 layers × 2 + 1 for the embedding) = 144 all-reduces, plus one to agree on the pool size. The first step (8 × 64 = 512 prompt tokens × 512 features × 4 bytes = 1 MB per all-reduce) dominates the 11.65 MB; a decode step moves 16 KB per all-reduce. PP sends one hidden vector per token per step: 1.29 MB in total.
- The times say nothing about GPUs. Four processes share a 4-core CPU, and each gloo all-reduce costs a millisecond or two. On GPUs with NVLink, TP 2 makes a bandwidth-bound decode step nearly twice as fast.
Expert parallelism with all-to-all
Tensor parallelism gives every rank every token. DeepSeek’s deployments do something else for big MoEs: attention runs data-parallel (each rank serves its own requests, with its own KV cache, so there’s no KV-head limit and no attention all-reduce), and experts run expert-parallel across all ranks. A token’s hidden vector travels to the ranks owning its $k$ experts and the results travel back: two all-to-alls per MoE layer.
class ExpertParallelMoE(nn.Module):
"""Experts split across ranks that hold DIFFERENT tokens (data-parallel attention, as in
DeepSeek's deployments). Each routed (token, expert) pair travels to the expert's rank and
back: two all-to-alls per MoE layer instead of an all-reduce of every token."""
def __init__(self, block, comm):
super().__init__()
e = block.experts.num_experts
self.per = e // comm.world
self.comm, self.first, self.gate = comm, comm.rank * self.per, block.gate
self.gate_up = nn.Parameter(block.experts.gate_up_proj[self.first:self.first + self.per].detach().clone())
self.down = nn.Parameter(block.experts.down_proj[self.first:self.first + self.per].detach().clone())
self.shared, self.shared_gate = block.shared_expert, block.shared_expert_gate # replicated
def forward(self, x):
"""Dispatch, compute, combine. (Your engine: Chapter 41)"""
_, weights, experts = self.gate(x)
k, flat = experts.shape[1], experts.reshape(-1)
owner = flat // self.per
order = owner.argsort(stable=True) # assignments grouped by destination rank
counts = torch.bincount(owner, minlength=self.comm.world).tolist()
rows = self.comm.all_to_all(list(x[order // k].split(counts)))
ids = self.comm.all_to_all(list(flat[order].split(counts)))
arrived, local = torch.cat(rows), torch.cat(ids) - self.first
y = torch.empty_like(arrived)
for e in local.unique().tolist(): # a grouped GEMM over the arrivals
idx = torch.where(local == e)[0]
gate, up = F.linear(arrived[idx], self.gate_up[e]).chunk(2, dim=-1)
y[idx] = F.linear(F.silu(gate) * up, self.down[e])
back = torch.cat(self.comm.all_to_all(list(y.split([r.shape[0] for r in rows]))))
out = torch.zeros_like(x).index_add_(0, order // k, back * weights.reshape(-1)[order, None])
if self.shared is not None:
out = out + torch.sigmoid(self.shared_gate(x)) * self.shared(x)
return out
Each rank sorts its (token, expert) assignments by destination rank, sends each destination its rows and expert ids, computes whatever arrives, and sends the results back in the same order, so the combine is an index_add_ with the router weights. The test gives two ranks different numbers of tokens and checks the result against the whole MoE on one process.
Moving $k$ hidden vectors per token twice per layer is the dominant cost at scale; DeepEP’s kernels overlap it with computation and use NVLink inside a node and RDMA between nodes. Routing is also never uniform: popular experts make their ranks the bottleneck, so DeepSeek’s EPLB replicates hot experts onto several ranks, rebalancing every few minutes from observed routing counts.
Data parallelism and prefix-aware routing
Data parallelism is the simplest: $N$ independent replicas behind a load balancer. Round-robin routing wastes the prefix cache, though: requests that share a long system prompt land on different replicas, and each replica prefills it. A router that remembers which prefixes it sent where can send the next such request to the replica that already holds its blocks, unless that replica is much busier.
class PrefixRouter:
"""Data parallelism: N independent engine replicas behind one router.
The router remembers which block hashes it has sent to each replica (an approximation of
that replica's prefix cache: it can't see evictions) and scores each replica by the prompt
tokens it would NOT have to prefill there, minus `balance` times the tokens it is already
working on. balance = 0 is pure cache affinity; a large balance is least-loaded routing.
"""
def __init__(self, num_replicas, block_size, balance=0.5, max_blocks=1 << 16):
self.block_size, self.balance, self.max_blocks = block_size, balance, max_blocks
self.seen = [OrderedDict() for _ in range(num_replicas)]
self.load = [0] * num_replicas
self.assigned = {}
def hashes(self, prompt, cache_key=()):
out, parent = [], None
for start in range(0, len(prompt) - self.block_size + 1, self.block_size):
parent = hash_block(parent, prompt[start:start + self.block_size], tuple(cache_key))
out.append(parent)
return out
def route(self, request_id, prompt, cost=None, cache_key=()):
"""Pick a replica for a request and charge it `cost` tokens. (Your engine: Chapter 41)"""
hashes = self.hashes(prompt, cache_key)
def score(r):
hit = 0
while hit < len(hashes) and hashes[hit] in self.seen[r]:
hit += 1
return hit * self.block_size - self.balance * self.load[r], -self.load[r], -r
best = max(range(len(self.seen)), key=score)
seen = self.seen[best]
for h in hashes:
seen[h] = None
seen.move_to_end(h)
while len(seen) > self.max_blocks:
seen.popitem(last=False)
cost = len(prompt) if cost is None else cost
self.load[best] += cost
self.assigned[request_id] = (best, cost)
return best
def finish(self, request_id):
replica, cost = self.assigned.pop(request_id, (None, 0))
if replica is not None:
self.load[replica] -= cost
class DataParallel:
"""Several EngineCores behind a PrefixRouter, stepped one after another (in a deployment,
each replica is its own process on its own GPUs, behind an HTTP router)."""
def __init__(self, engines, router):
self.engines, self.router = engines, router
self.owner = {}
def add_request(self, request_id, prompt, params):
replica = self.router.route(request_id, prompt, len(prompt) + params.max_tokens * params.n)
self.owner[request_id] = replica
return self.engines[replica].add_request(request_id, prompt, params)
def step(self):
outputs = []
for engine in self.engines:
for out in engine.step():
root = out.request_id.split(":")[0]
if out.finished and root in self.router.assigned and not any(
r.split(":")[0] == root for e in self.engines for r in e.scheduler.requests):
self.router.finish(root)
outputs.append(out)
return outputs
@property
def has_unfinished(self):
return any(e.has_unfinished for e in self.engines)
def generate(self, prompts, params):
results = {}
for rid, prompt in prompts.items():
self.add_request(rid, prompt, params)
while self.has_unfinished:
for out in self.step():
results.setdefault(out.request_id, []).extend(out.new_token_ids)
return results
The score of replica $r$ is the prompt tokens it would not have to prefill there minus balance × the tokens it is already working on. The router’s memory of cached blocks is approximate (it can’t see evictions); SGLang’s router keeps the same kind of approximate radix tree per worker, and llm-d subscribes to each vLLM replica’s KV-cache events to know exactly. DataParallel steps several in-process engines for the tests; in production each replica is its own server process (Chapter 36) and the router is an HTTP proxy.
Disaggregated prefill and decode
Prefill and decode want different things. A prefill chunk is compute-bound and makes every decode in the same step wait for it (Chapter 24’s interference, which chunked prefill only softens). Decode is bandwidth-bound and wants big batches. Disaggregation (DistServe, Splitwise, Mooncake) runs them on different GPUs: a prefill instance computes the prompt’s KV and the first token, ships the KV to a decode instance, and the decode instance carries on. Each pool is sized and tuned for its job, and the time to first token no longer depends on how busy decode is.
The engine already has every piece: preemption by swapping (Chapter 31) copies a request’s blocks to host memory and back. A prefill instance exports a finished one-token request exactly the way swap-out does, and a decode instance admits it exactly the way swap-in does:
@dataclass
class KVTransfer:
request_id: str
token_ids: list # prompt + the first output token
num_prompt_tokens: int
kv: list # per layer: (K blocks, V blocks[, scales]) in host memory
generator_state: object = None
class KVExporter:
"""On the PREFILL instance: when a request marked transfer_kv finishes its one-token run,
copy its blocks out before the scheduler frees them."""
def __init__(self, engine):
self.engine, self.ready = engine, {}
engine.scheduler.on_finish = self.export
def prefill(self, request_id, prompt, params):
request = self.engine.add_request(request_id, prompt, replace(params, max_tokens=1))
request.extra["transfer_kv"] = True
return request
def export(self, request):
"""(Your engine: Chapter 41)"""
if not request.extra.get("transfer_kv") or request.status is not Status.FINISHED_LENGTH:
return # stopped at its first token (EOS): nothing to decode
e, rid = self.engine, request.request_id
e.runner.swap_out(request, e.blocks.req_blocks[rid]) # the swap path already copies blocks out
generator = e.generators.get(rid)
self.ready[rid] = KVTransfer(rid, list(request.token_ids), request.num_prompt_tokens,
e.runner.swapped.pop(rid), generator.get_state() if generator else None)
def import_kv(engine, transfer, params):
"""On the DECODE instance: admit the request as if it had been swapped out after its prefill,
so the scheduler's swap-in path allocates blocks and copies the KV in. (Your engine: Chapter 41)"""
request = Request(transfer.request_id, transfer.token_ids[:transfer.num_prompt_tokens], params, engine.eos_token_id)
for token in transfer.token_ids[transfer.num_prompt_tokens:]:
request.append(token)
request.num_computed_tokens = transfer.num_prompt_tokens # the first output token's KV is computed next
request.swapped_out = True
engine.runner.swapped[request.request_id] = transfer.kv
if transfer.generator_state is not None: # seeded sampling continues the same stream
engine.generators[request.request_id] = torch.Generator(device=engine.device).set_state(transfer.generator_state)
elif params.seed is not None:
engine.generators[request.request_id] = torch.Generator(device=engine.device).manual_seed(params.seed)
engine.scheduler.add(request)
return request
def send_transfer(comm, transfer, dst):
"""Ship a transfer to another rank: metadata as a pickled object, KV as raw tensors."""
meta = replace(transfer, kv=[[(t.shape, t.dtype) for t in layer] for layer in transfer.kv])
dist.send_object_list([meta], dst)
for layer in transfer.kv:
for t in layer:
comm.send(t, dst)
def recv_transfer(comm, src):
box = [None]
dist.recv_object_list(box, src)
meta = box[0]
kv = [tuple(comm.recv(shape, dtype, src, "cpu") for shape, dtype in layer) for layer in meta.kv]
return replace(meta, kv=kv)
KVExporter hooks the scheduler’s on_finish, the moment before a finished request’s blocks are freed. The transfer carries the prompt, the first token, every layer’s blocks and, for seeded sampling, the random generator’s state, so the decode side continues the same random stream: the test checks that seeded sampling gives the same tokens as one engine. import_kv creates the request with num_computed_tokens set past the prompt and swapped_out = True, and the scheduler’s swap-in path allocates blocks and copies the KV in. The decode engine runs no prefill step at all.
Is shipping KV affordable? Qwen3-8B’s KV is 144 KB per token (Chapter 31), so a 2,000-token prompt is 295 MB: 6 ms over a 50 GB/s RDMA link, against roughly 100 ms to prefill it on an H100 (32 TFLOP at about 40% of peak). Real connectors (vLLM’s NIXL connector, Mooncake’s transfer engine, SGLang’s) move KV GPU-to-GPU over RDMA, layer by layer while the prefill is still running, and can transfer between instances with different TP sizes by re-slicing heads. send_transfer here goes over torch.distributed (gloo between CPU processes in the test).
Ring attention for very long contexts
A million-token prompt’s KV cache doesn’t fit on one GPU even for a small model, and its prefill is quadratic. Context parallelism splits the sequence: rank $r$ holds tokens $[rT, (r+1)T)$ and their Q, K and V. Attention needs every earlier key, so the K/V chunks travel around a ring of ranks: at each hop, each rank attends its queries to the chunk it holds, then passes the chunk on and receives the previous rank’s. After $n - 1$ hops every rank has seen every chunk.
def attention_partial(q, k, v, q_pos, k_pos, scale):
"""Causal attention of q [Tq, H, D] over k, v [Tk, H, D] -> (out, lse [Tq, H]); rows that see
no key get lse = -inf and contribute nothing when merged."""
scores = torch.einsum("qhd,khd->hqk", q.float(), k.float()) * scale
scores = scores.masked_fill(k_pos[None, None, :] > q_pos[None, :, None], float("-inf"))
lse = torch.logsumexp(scores, dim=-1) # [H, Tq]
probs = torch.exp(scores - torch.where(lse.isinf(), 0, lse)[..., None])
return torch.einsum("hqk,khd->qhd", probs, v.float()), lse.T
def merge_partials(out, lse, new_out, new_lse):
"""Chapter 32's split-KV merge, two at a time."""
top = torch.maximum(lse, new_lse)
top = torch.where(top.isinf(), 0, top)
a, b = torch.exp(lse - top), torch.exp(new_lse - top)
total = a + b
merged = (out * a[..., None] + new_out * b[..., None]) / torch.where(total > 0, total, 1)[..., None]
return merged, top + torch.log(total)
def ring_attention(q, k, v, comm, scale=None):
"""Each rank holds one contiguous chunk of a long sequence (q, k, v: [T, H, D]). K/V chunks
travel around the ring; after world - 1 hops every rank has attended its queries to every
earlier key without any rank holding the whole sequence. (Your engine: Chapter 41)"""
t, scale = q.shape[0], scale or q.shape[-1] ** -0.5
q_pos = torch.arange(comm.rank * t, (comm.rank + 1) * t, device=q.device)
out = torch.zeros(q.shape, dtype=torch.float32, device=q.device)
lse = torch.full((t, q.shape[1]), float("-inf"), device=q.device)
kv, owner = torch.stack((k, v)), comm.rank
for hop in range(comm.world):
if owner <= comm.rank: # chunks from later ranks are entirely masked
k_pos = torch.arange(owner * t, (owner + 1) * t, device=q.device)
o, l = attention_partial(q, kv[0], kv[1], q_pos, k_pos, scale)
out, lse = merge_partials(out, lse, o, l)
if hop + 1 < comm.world: # pass ours on, take the previous rank's
incoming = torch.empty_like(kv)
ops = [dist.isend(kv.contiguous(), (comm.rank + 1) % comm.world),
dist.irecv(incoming, (comm.rank - 1) % comm.world)]
for op in ops:
op.wait()
kv, owner = incoming, (owner - 1) % comm.world
return out.to(q.dtype)
Combining the partial results is exactly Chapter 32’s split-KV merge: each partial carries its log-sum-exp, and $o = \sum_s e^{\ell_s - \ell} o_s$ with $\ell = \log \sum_s e^{\ell_s}$. Causality makes the plain split unbalanced: rank 0’s queries see one chunk, the last rank’s see all of them. Production implementations split the sequence into $2n$ chunks and give rank $r$ chunks $r$ and $2n - 1 - r$ (“zigzag”), which evens out the work, and overlap each hop’s transfer with the previous chunk’s computation.
Build it
Engine milestone 41: many GPUs. Implement RowParallelLinear.forward, VocabParallelEmbedding.forward, ParallelLMHead.forward, TensorParallelMoE.forward, shard_tensor_parallel, PipelineStage.forward, ExpertParallelMoE.forward, worker_loop, PrefixRouter.route, KVExporter.export, import_kv and ring_attention in engine/parallel.py.
pytest tests/test_ch41_multi_gpu.py # 13 tests; several start 2-4 processes each
python run.py dist --impl engine
The tests check token-for-token equality with a single process under TP 2, TP 2 with tied embeddings, TP 4 (replicated KV heads), TP 2 for a MoE with seeded sampling, and PP 2; expert parallelism with uneven token counts; ring attention over three ranks; the router’s choices; and disaggregated serving in one process and across two.
Stretch exercises
- ★★ Load only the shard: give
shard_tensor_parallela path that reads each rank’s rows and columns straight from the safetensors files (safe_open(...).get_slice(name)[start:end]), so no rank ever holds the full model. Measure peak memory per rank. Where:shard_tensor_parallelinengine/parallel.py. - ★★★ Fill the pipeline bubble: let the engine keep
ppbatches in flight (each made of requests not in the others), using Chapter 33’s placeholder mechanism. Measure throughput on two GPUs against PP with one batch in flight. Where:DistributedRunner/worker_loopinengine/parallel.py, with scheduling inengine/serve/engine.py. - ★★ Two GPUs over PCIe: compare TP 2 and PP 2 on decode latency and throughput for a model that doesn’t fit on one GPU. Where does each win? Where:
experiments/ch41.py(create it), adaptingrun.py’scmd_distandengine.parallel.parallel_engine. - ★★ Tensor parallelism for Medusa and draft models (Chapter 37): which parts must be sharded, and which can stay replicated on rank 0? Where: paper first; implement sharding/wrappers in
engine/parallel.pyand drafter integration inengine/serve/spec.py. - ★★★ Layer-by-layer KV transfer: start sending layer $\ell$’s blocks as soon as the prefill finishes layer $\ell$ (a CUDA event per layer, a second stream for the copies). How much of the transfer hides behind the prefill? Where:
KVExporter,send_transferandrecv_transferinengine/parallel.py, with per-layer completion hooks inengine/serve/model.py.
Check your understanding
- Why can Q/K/V and the MLP’s gate and up projections be split by output without any communication, while $W_o$ and the down projection need an all-reduce?
- Why does TP’s KV memory per GPU stop shrinking once there are more ranks than KV heads?
- A decode step’s all-reduces are about 1 MB each. Is that bandwidth- or latency-bound on NVLink? On PCIe?
- Why does pipeline parallelism with one batch in flight not make decoding faster, and what fixes it?
- When is all-to-all expert parallelism cheaper than tensor parallelism’s all-reduce for a MoE layer?
- Why does a prefix-aware router need a load term at all?
- How does disaggregation reuse preemption by swapping, and what extra state does a seeded request need to carry?
Going deeper
- Shoeybi et al., Megatron-LM: Training Multi-Billion Parameter Language Models Using Model Parallelism (2019): the column/row-parallel construction.
- Huang et al., GPipe (NeurIPS 2019) and Narayanan et al., PipeDream and Efficient Large-Scale Language Model Training on GPU Clusters Using Megatron-LM (SC 2021), for pipeline schedules.
- DeepSeek-AI, DeepSeek-V3 Technical Report (2024) and the DeepEP, EPLB and “Day 6: inference system overview” releases (2025), for DP attention with EP.
- Zhong et al., DistServe (OSDI 2024); Patel et al., Splitwise (ISCA 2024); Qin et al., Mooncake (FAST 2025), for disaggregated serving.
- Liu, Zaharia and Abbeel, Ring Attention with Blockwise Transformers for Near-Infinite Context (2023).
- vLLM’s
vllm/distributed/(parallel state, custom all-reduce, KV connectors) andvllm/model_executor/layers/linear.py; SGLang’s router (sgl-router) and its DP-attention implementation.
42. More models: a registry, scaled RoPE, sliding windows and latent attention
In this chapter
- What actually differs between Llama, Mistral, Qwen2, Qwen3 and DeepSeek, and how one configurable decoder covers the dense ones.
- Long-context RoPE scaling (linear, Llama 3, YaRN), derived from what each pair of features does over the training context.
- Sliding-window attention in the engine's kernels.
- Multi-head latent attention: DeepSeek's KV cache, 71 times smaller than full multi-head attention would be, and the "absorbed" form that serves it as one wide query head.
- DeepSeek's grouped router with shared experts, a registry keyed on
config.json, and Llama-architecture GGUF files.
You will build
rope_inv_freq, DecoderAttention.forward, MLAAttention.paged_forward and GroupedRouter.forward in engine/models.py.
Time: 5-7 hours. GPU: not needed. The parity tests need transformers (they build random models locally; nothing is downloaded).
What differs between model families
Since Chapter 17 the engine has run Qwen3. vLLM and llama.cpp each support well over a hundred architectures. That sounds like a lot of code, but most of the decoder-only models people serve differ in a short list of switches:
| Llama 3 | Mistral 7B | Qwen2 / 2.5 | Qwen3 | DeepSeek-V3 | |
|---|---|---|---|---|---|
| Q/K/V biases | no (attention_bias adds them, and on o_proj) | no | Q, K, V | no | no |
| per-head Q/K norm | no | no | no | yes | no (norms on the latents) |
| RoPE | Llama 3 scaling | plain | plain or YaRN | plain or YaRN | YaRN, interleaved pairs |
| attention | GQA | GQA, 4,096-token window | GQA, optional windows | GQA | multi-head latent |
| MLP | SwiGLU | SwiGLU | SwiGLU | SwiGLU or MoE | first 3 layers dense, then 256 routed + 1 shared expert |
| tied embeddings | 1B, 3B | no | small models | small models | no |
Everything in the first four columns fits one class whose constructor reads these switches from config.json. DeepSeek needs a different attention module and router. Gemma (normalized embeddings, 1 + w norms, soft-capping), Phi and others need a handful more switches (stretch exercise 1).
@dataclass
class DecoderConfig(Qwen3Config):
model_type: str = "qwen3"
qkv_bias: bool = False # Qwen2
o_bias: bool = False
qk_norm: bool = True # Qwen3
rope_scaling: dict | None = None
sliding_window: int = 0
layer_types: list | None = None # "full_attention" / "sliding_attention" per layer
@classmethod
def from_hf(cls, raw):
kind = raw.get("model_type")
if kind not in ("llama", "mistral", "qwen2", "qwen3"):
raise ValueError(f"DecoderConfig doesn't cover model_type {kind!r}")
if raw.get("hidden_act", "silu") != "silu" or raw.get("mlp_bias"):
raise ValueError("Only SwiGLU MLPs without biases are implemented")
layers, heads = raw["num_hidden_layers"], raw["num_attention_heads"]
window = raw.get("sliding_window") or 0
types = raw.get("layer_types")
if types is None and window:
if kind == "mistral":
types = ["sliding_attention"] * layers
elif kind == "qwen2" and raw.get("use_sliding_window"):
start = raw.get("max_window_layers", layers)
types = ["full_attention" if i < start else "sliding_attention" for i in range(layers)]
bias = bool(raw.get("attention_bias", False))
return cls(vocab_size=raw["vocab_size"], hidden_size=raw["hidden_size"], intermediate_size=raw["intermediate_size"],
num_hidden_layers=layers, num_attention_heads=heads,
num_key_value_heads=raw.get("num_key_value_heads") or heads,
head_dim=raw.get("head_dim") or raw["hidden_size"] // heads, rms_norm_eps=raw.get("rms_norm_eps", 1e-6),
rope_theta=theta_of(raw), max_position_embeddings=raw.get("max_position_embeddings", 4096),
tie_word_embeddings=raw.get("tie_word_embeddings", False), model_type=kind,
qkv_bias=bias or kind == "qwen2", o_bias=bias, qk_norm=kind == "qwen3",
rope_scaling=scaling_of(raw), sliding_window=window if types and "sliding_attention" in types else 0,
layer_types=types)
def window(self, layer):
types = self.layer_types
return self.sliding_window if types and types[layer] == "sliding_attention" else 0
Two details hide in that table. Llama’s attention_bias adds biases to all four projections, Qwen2 always has them on Q, K and V and never on o_proj; and Qwen2’s sliding windows apply only to layers from max_window_layers on. Recent transformers configs spell the result out per layer in layer_types, which the config prefers when present.
class DecoderAttention(Qwen3Attention):
def __init__(self, cfg, window=0):
nn.Module.__init__(self)
self.cfg, self.sliding_window = cfg, window
d, hq, hkv = cfg.head_dim, cfg.num_attention_heads, cfg.num_key_value_heads
self.q_proj = nn.Linear(cfg.hidden_size, hq * d, bias=cfg.qkv_bias)
self.k_proj = nn.Linear(cfg.hidden_size, hkv * d, bias=cfg.qkv_bias)
self.v_proj = nn.Linear(cfg.hidden_size, hkv * d, bias=cfg.qkv_bias)
self.o_proj = nn.Linear(hq * d, cfg.hidden_size, bias=cfg.o_bias)
if cfg.qk_norm: # absent, not identity: FlatModel checks hasattr
self.q_norm = RMSNorm(d, cfg.rms_norm_eps)
self.k_norm = RMSNorm(d, cfg.rms_norm_eps)
def forward(self, x, positions, rope, cache=None, layer=0, rows=None):
"""Qwen3Attention.forward with optional norms and a sliding window. (Your engine: Chapter 42)"""
b, t, _ = x.shape
c = self.cfg
q = self.q_proj(x).view(b, t, c.num_attention_heads, c.head_dim)
k = self.k_proj(x).view(b, t, c.num_key_value_heads, c.head_dim)
v = self.v_proj(x).view(b, t, c.num_key_value_heads, c.head_dim).transpose(1, 2)
if hasattr(self, "q_norm"):
q, k = self.q_norm(q), self.k_norm(k)
cos, sin = rope
q, k = apply_rope(q.transpose(1, 2), cos, sin), apply_rope(k.transpose(1, 2), cos, sin)
key_positions = None
if cache is not None:
k, v, key_positions = cache.update(layer, k, v, positions, rows)
allowed = None
if self.sliding_window:
s = k.shape[2]
kp = key_positions if key_positions is not None else torch.arange(s, device=x.device)
qp = positions if positions is not None else torch.arange(s - t, s, device=x.device)
allowed = window_mask(qp, kp, self.sliding_window)
allowed = allowed[:, None] if allowed.ndim == 3 else allowed
y = causal_attention(q, k, v, positions, key_positions, allowed)
return self.o_proj(y.transpose(1, 2).reshape(b, t, -1))
class DecoderLayer(Qwen3Layer):
def __init__(self, cfg, layer, attention=None, mlp=None):
nn.Module.__init__(self)
self.input_layernorm = RMSNorm(cfg.hidden_size, cfg.rms_norm_eps)
self.self_attn = attention or DecoderAttention(cfg, cfg.window(layer))
self.post_attention_layernorm = RMSNorm(cfg.hidden_size, cfg.rms_norm_eps)
self.mlp = mlp or SwiGLU(cfg.hidden_size, cfg.intermediate_size)
class Backbone(nn.Module):
def __init__(self, cfg, make_layer):
super().__init__()
self.embed_tokens = nn.Embedding(cfg.vocab_size, cfg.hidden_size)
self.layers = nn.ModuleList(make_layer(i) for i in range(cfg.num_hidden_layers))
self.norm = RMSNorm(cfg.hidden_size, cfg.rms_norm_eps)
class Decoder(Qwen3):
"""Llama, Mistral, Qwen2 and Qwen3 in one class."""
def __init__(self, cfg, make_layer=None):
nn.Module.__init__(self)
self.cfg = cfg
self.model = Backbone(cfg, make_layer or (lambda i: DecoderLayer(cfg, i)))
self.lm_head = nn.Linear(cfg.hidden_size, cfg.vocab_size, bias=False)
if cfg.tie_word_embeddings:
self.lm_head.weight = self.model.embed_tokens.weight
self.inv_freq, self.attention_factor = rope_inv_freq(cfg.head_dim, cfg.rope_theta, cfg.rope_scaling,
cfg.max_position_embeddings)
def rope_tables(self, positions, dtype):
"""FlatModel's hook: [N, 1, dim] tables with this model's scaling."""
cos, sin = rope_tables(positions, self.inv_freq, self.attention_factor, dtype)
return cos[:, None], sin[:, None]
def forward(self, ids, cache=None, positions=None, rows=None, return_hidden=False):
if positions is None:
start = cache.length if cache is not None else 0
positions = torch.arange(start, start + ids.shape[1], device=ids.device)
x = self.model.embed_tokens(ids)
cos, sin = rope_tables(positions, self.inv_freq, self.attention_factor, x.dtype)
rope = (cos[None, None], sin[None, None]) if positions.ndim == 1 else (cos[:, None], sin[:, None])
for i, layer in enumerate(self.model.layers):
x = layer(x, positions, rope, cache, i, rows)
hidden = self.model.norm(x)
return hidden if return_hidden else self.lm_head(hidden)
DecoderAttention is Chapter 17’s attention with the norms optional, biases switchable and a window mask. It leaves q_norm undefined rather than an identity, because FlatModel (Chapter 31) checks hasattr(attn, "q_norm"). Decoder adds one hook, rope_tables, which FlatModel already looks for, so the serving path needs no model-specific code for any of these families.
Scaled RoPE
RoPE rotates feature pair $i$ by angle $p,\theta^{-2i/d}$ at position $p$ (Chapter 17). Pair 0 turns fast (once every $2\pi$ positions); the last pair turns slowly: for $\theta = 500{,}000$ and $d = 128$ its wavelength is about 2.6 million positions. A model trained on 8,192-token sequences has seen fast pairs turn thousands of times, but slow pairs only through a fraction of a turn. Positions past the training length push slow pairs into angles never seen in training, and quality collapses. Every long-context scheme is a way to keep angles in the trained range:
- Linear (position interpolation) divides every frequency by the extension factor $s$: position $p$ looks like $p/s$. Simple, but fast pairs, which distinguish neighbouring tokens, lose resolution.
- Llama 3 keeps pairs whose wavelength is shorter than (original context /
high_freq_factor), divides pairs longer than (original context /low_freq_factor) by $s$, and blends linearly in between. - YaRN chooses the same split by counting rotations: pairs that turn more than
beta_fast(32) times over the original context are kept, pairs that turn fewer thanbeta_slow(1) times are interpolated, with a linear ramp between. It also sharpens attention, multiplying cos and sin by $0.1 \ln s + 1$ (DeepSeek folds a version of this into the softmax scale instead).
def rope_inv_freq(dim, theta, scaling=None, max_positions=None):
"""Rotation frequency of each feature pair, with long-context scaling. (Your engine: Chapter 42)
Returns (inv_freq [dim/2], attention_factor), the factor multiplying cos and sin (YaRN only).
linear every frequency divided by `factor`: positions squeezed into the trained range
llama3 low frequencies (wavelength > original context / low_freq_factor) divided by
`factor`, high ones kept, a smooth blend between
yarn the same idea with a linear ramp over dimensions, chosen by how many full
rotations each pair makes over the original context, plus a temperature
"""
base = 1.0 / (theta ** (torch.arange(0, dim, 2, dtype=torch.int64, device="cpu").float() / dim)) # cpu: even under a meta-device build
s = scaling or {}
kind = s.get("rope_type", s.get("type", "default"))
if kind == "default":
return base, 1.0
factor = s.get("factor")
if kind == "linear":
return base / factor, 1.0
if kind == "llama3":
old, low, high = s["original_max_position_embeddings"], s["low_freq_factor"], s["high_freq_factor"]
wavelen = 2 * math.pi / base
scaled = torch.where(wavelen > old / low, base / factor, base)
smooth = (old / wavelen - low) / (high - low)
medium = (wavelen <= old / low) & (wavelen >= old / high)
return torch.where(medium, (1 - smooth) * base / factor + smooth * base, scaled), 1.0
if kind == "yarn":
old = s["original_max_position_embeddings"]
factor = factor or max_positions / old
def mscale(scale, m=1.0):
return 1.0 if scale <= 1 else 0.1 * m * math.log(scale) + 1.0
attention = s.get("attention_factor")
if attention is None:
attention = (mscale(factor, s["mscale"]) / mscale(factor, s["mscale_all_dim"])
if s.get("mscale") and s.get("mscale_all_dim") else mscale(factor))
def dim_for(rotations): # the pair index that turns `rotations` times over `old`
return dim * math.log(old / (rotations * 2 * math.pi)) / (2 * math.log(theta))
lo, hi = dim_for(s.get("beta_fast") or 32), dim_for(s.get("beta_slow") or 1)
if s.get("truncate", True):
lo, hi = math.floor(lo), math.ceil(hi)
lo, hi = max(lo, 0), min(hi, dim - 1)
ramp = ((torch.arange(dim // 2, dtype=torch.float32, device="cpu") - lo) / ((hi - lo) or 0.001)).clamp(0, 1)
keep = 1 - ramp # 1: fast pairs, extrapolated; 0: slow pairs, interpolated
return base / factor * (1 - keep) + base * keep, attention
raise ValueError(f"RoPE type {kind!r} is not implemented")
def rope_tables(positions, inv_freq, attention_factor, dtype):
"""cos and sin [..., dim] for any positions shape."""
angles = positions.float()[..., None] * inv_freq.to(positions.device)
angles = torch.cat((angles, angles), dim=-1)
return (angles.cos() * attention_factor).to(dtype), (angles.sin() * attention_factor).to(dtype)
All three change only the frequencies, so the cos/sin tables are computed once per step as before. (Qwen’s “dynamic NTK” and Phi-3’s LongRoPE are variants: the first changes $\theta$ with the sequence length, which breaks the prefix cache’s assumption that a token’s K never changes; the second learns per-pair factors.)
Sliding windows
Mistral 7B attends to the last 4,096 tokens only; Gemma 2 and 3, gpt-oss and Qwen2 alternate windowed and full layers. The engine’s kernels already take a window argument (Chapter 32): the reference backend masks keys more than $w - 1$ positions back, and the Triton kernel also starts its loop at the first block inside the oldest query’s window, so a windowed layer costs $O(w)$ per token however long the context. FlatModel.attention now passes each layer’s sliding_window through.
What the engine doesn’t do yet is free the blocks that fall out of every window. A model with only windowed layers needs just $\lceil w / \text{block size} \rceil + 1$ blocks per request; a model that mixes windowed and full layers needs both kinds. vLLM’s hybrid KV-cache manager gives each layer group its own block table so that windowed layers’ blocks can be recycled (stretch exercise 2).
Multi-head latent attention
DeepSeek-V3 has 128 attention heads with 192-wide keys and 128-wide values. Cached the ordinary way, that’s $61 \times 128 \times (192 + 128) \times 2$ bytes = 4.8 MB per token: one 32K-token conversation would need 150 GB. MLA (DeepSeek-V2, 2024) caches 68.6 KB per token instead:
- Each token’s hidden state is compressed to a latent $c = \mathrm{norm}(W_{dkv},x)$ of width 512.
- Every head’s key and value are linear functions of it: $k^{nope}_h = W^{uk}_h c$, $v_h = W^{uv}_h c$ (
kv_b_projholds all of them). - RoPE can’t pass through $W^{uk}$ (rotation depends on position, the matrix doesn’t), so each token also gets one small rotary key $k^{rope}$ of width 64, shared by all heads; queries have a matching rotary part.
Only $[c, k^{rope}]$, 576 numbers, is cached. run.py models prints the arithmetic for some real configurations:
{"model": "Llama-3.1-8B", "kv_KiB_per_token": 128.0, "GiB_for_32k_tokens": 4.0, "saving_vs_full_mha": "4.0x", "sliding_window": null}
{"model": "Mistral-7B-v0.1", "kv_KiB_per_token": 128.0, "GiB_for_32k_tokens": 4.0, "saving_vs_full_mha": "4.0x", "sliding_window": 4096}
{"model": "Qwen2.5-7B", "kv_KiB_per_token": 56.0, "GiB_for_32k_tokens": 1.75, "saving_vs_full_mha": "7.0x", "sliding_window": null}
{"model": "DeepSeek-V3", "kv_KiB_per_token": 68.6, "GiB_for_32k_tokens": 2.14, "saving_vs_full_mha": "71.1x", "sliding_window": null}
A 671-billion-parameter model whose KV cache per token is about half of Llama-3.1-8B’s.
Serving it: absorbing the up-projections
Decompressing every cached token into 128 heads’ keys and values at every decode step would cost far more than reading them. The trick is associativity. A head’s score against cached token $j$ is
$$q^{nope}_h \cdot (W^{uk}_h c_j) + q^{rope}_h \cdot k^{rope}_j = \big((W^{uk}_h)^T q^{nope}_h\big) \cdot c_j + q^{rope}_h \cdot k^{rope}_j .$$
Map each query into latent space once, $\tilde q_h = (W^{uk}_h)^T q^{nope}h$, concatenate its rotary part, and attention becomes multi-query attention with a single 576-wide KV head over the cached rows. The values are the latents themselves: the output in latent space, $\sum_j a{hj} c_j$, is mapped back per head by $W^{uv}_h$ after attention. So the cache needs one tensor, whose first 512 features serve as values. That’s exactly what DeepSeek’s FlashMLA kernel consumes.
class MLAAttention(nn.Module):
"""Multi-head latent attention (DeepSeek-V2). Keys and values of all heads are linear
functions of ONE compressed vector per token, c = norm(W_dkv x) of width kv_lora_rank, plus
a small rotary key shared by all heads. Only [c, k_rope] is cached."""
def __init__(self, cfg):
super().__init__()
self.cfg = cfg
h, qk = cfg.num_attention_heads, cfg.qk_nope_head_dim + cfg.qk_rope_head_dim
if cfg.q_lora_rank:
self.q_a_proj = nn.Linear(cfg.hidden_size, cfg.q_lora_rank, bias=cfg.attention_bias)
self.q_a_layernorm = RMSNorm(cfg.q_lora_rank)
self.q_b_proj = nn.Linear(cfg.q_lora_rank, h * qk, bias=False)
else:
self.q_proj = nn.Linear(cfg.hidden_size, h * qk, bias=False)
self.kv_a_proj_with_mqa = nn.Linear(cfg.hidden_size, cfg.kv_lora_rank + cfg.qk_rope_head_dim, bias=cfg.attention_bias)
self.kv_a_layernorm = RMSNorm(cfg.kv_lora_rank)
self.kv_b_proj = nn.Linear(cfg.kv_lora_rank, h * (cfg.qk_nope_head_dim + cfg.v_head_dim), bias=False)
self.o_proj = nn.Linear(h * cfg.v_head_dim, cfg.hidden_size, bias=cfg.attention_bias)
self.scale = qk ** -0.5
s = cfg.rope_scaling or {}
if s.get("mscale_all_dim"): # YaRN's temperature, applied to the softmax scale
m = 0.1 * s["mscale_all_dim"] * math.log(s["factor"]) + 1.0 if s["factor"] > 1 else 1.0
self.scale *= m * m
def queries(self, x):
c = self.cfg
q = self.q_b_proj(self.q_a_layernorm(self.q_a_proj(x))) if c.q_lora_rank else self.q_proj(x)
q = q.view(*x.shape[:-1], c.num_attention_heads, c.qk_nope_head_dim + c.qk_rope_head_dim)
return q.split([c.qk_nope_head_dim, c.qk_rope_head_dim], dim=-1)
def latent(self, x):
c = self.cfg
latent, k_rope = self.kv_a_proj_with_mqa(x).split([c.kv_lora_rank, c.qk_rope_head_dim], dim=-1)
return self.kv_a_layernorm(latent), k_rope
def forward(self, x, positions, rope, cache=None, layer=0, rows=None):
"""Training-style MLA: expand every head's K and V, then ordinary attention."""
if cache is not None:
raise NotImplementedError("MLA decodes through the engine's paged latent cache (Chapter 42)")
c = self.cfg
b, t, _ = x.shape
q_nope, q_rope = self.queries(x) # [b, t, H, *]
latent, k_rope = self.latent(x)
cos, sin = rope
q_rope = apply_rope(deinterleave(q_rope.transpose(1, 2)), cos, sin)
k_rope = apply_rope(deinterleave(k_rope[:, None]), cos, sin) # [b, 1, t, rope]
kv = self.kv_b_proj(latent).view(b, t, c.num_attention_heads, -1).transpose(1, 2)
k_nope, v = kv.split([c.qk_nope_head_dim, c.v_head_dim], dim=-1)
q = torch.cat((q_nope.transpose(1, 2), q_rope), dim=-1)
k = torch.cat((k_nope, k_rope.expand(-1, c.num_attention_heads, -1, -1)), dim=-1)
y = causal_attention(q, k, v, positions, scale=self.scale)
return self.o_proj(y.transpose(1, 2).reshape(b, t, -1))
def paged_forward(self, h, positions, rope, kv, meta, backend):
"""Serving-style MLA with the up-projections absorbed. (Your engine: Chapter 42)
score = q_nope . (W_uk c) + q_rope . k_rope = (W_uk^T q_nope) . c + q_rope . k_rope, so
each query is mapped into latent space once and attends straight over the cached
[c, k_rope] rows: multi-query attention with ONE head of width kv_lora_rank +
qk_rope_head_dim (576 in V3). The values are the latents themselves (the first
kv_lora_rank features of the same rows); W_uv maps the result back per head.
"""
c = self.cfg
n, r = h.shape[0], c.kv_lora_rank
q_nope, q_rope = self.queries(h) # [N, H, *]
latent, k_rope = self.latent(h)
cos, sin = rope
q_rope = apply_rope(deinterleave(q_rope), cos, sin)
k_rope = apply_rope(deinterleave(k_rope[:, None]), cos, sin) # [N, 1, rope]
w = self.kv_b_proj.weight.view(c.num_attention_heads, c.qk_nope_head_dim + c.v_head_dim, r)
q_latent = torch.einsum("nhd,hdr->nhr", q_nope, w[:, :c.qk_nope_head_dim])
row = torch.cat((latent[:, None], k_rope), dim=-1) # what the cache holds
backend.write(kv, row, row, meta)
out = backend.forward(torch.cat((q_latent, q_rope), dim=-1), kv, meta, scale=self.scale)[..., :r]
y = torch.einsum("nhr,hvr->nhv", out, w[:, c.qk_nope_head_dim:])
return self.o_proj(y.reshape(n, -1))
paged_forward is that computation over the paged pool: it writes the step’s $[c, k^{rope}]$ rows, runs the backend’s attention with 128 query heads against one KV head, keeps the first 512 output features and applies $W^{uv}$. forward is the training-style version (decompress, then ordinary attention), kept for parity tests. Two smaller details: DeepSeek’s checkpoints pair rotary features as adjacent elements $(0, 1), (2, 3), \ldots$, so deinterleave reorders them to the split-half pairing that apply_rope uses (applied to queries and keys alike, so dot products are unchanged). And YaRN’s temperature multiplies the softmax scale by $(0.1 \cdot \texttt{mscale_all_dim} \cdot \ln s + 1)^2$.
FlatModel.attention hands any attention module with a paged_forward method its own path. The model reports its pool shape through kv_spec (one “head” of width 576) and allocates it through allocate_kv, which returns the same tensor as both K and V, so the engine’s block manager, prefix cache, swapping and disaggregated transfer (Chapter 41) all work unchanged:
class MLADecoder(Decoder):
"""DeepSeek-V2/V3 (and Kimi K2): MLA in every layer; dense MLPs in the first
first_k_dense_replace layers, DeepseekMoE after."""
kv_pools = 1 # one latent pool per layer, not separate K and V
def __init__(self, cfg):
def layer(i):
moe = cfg.n_routed_experts and i >= cfg.first_k_dense_replace
return DecoderLayer(cfg, i, MLAAttention(cfg), DeepseekMoE(cfg) if moe else None)
super().__init__(cfg, layer)
def kv_spec(self):
c = self.cfg
return c.num_hidden_layers, 1, c.kv_lora_rank + c.qk_rope_head_dim
def allocate_kv(self, shape, dtype, device):
"""K and V are the same latent rows: one tensor serves as both."""
pool = torch.zeros(shape, dtype=dtype, device=device)
return pool, pool
Absorption isn’t always the cheaper form. For a long prefill, decompressing the chunk’s latents and running ordinary attention over 128 heads of width 192 does fewer FLOPs than 128 heads of width 576. vLLM and SGLang use the decompressed form for prefill and the absorbed form for decode (stretch exercise 3). MLA also doesn’t split over tensor parallelism the usual way: it has one KV head, so every rank would cache the whole latent. DeepSeek’s deployments run attention data-parallel instead (Chapter 41), and shard_tensor_parallel refuses MLA models.
DeepSeek’s router
DeepSeek-V3’s MoE layers have 256 routed experts (8 chosen per token) and one shared expert that every token uses. The router has two twists over Chapter 27’s:
- Group-limited routing. The 256 experts are split into 8 groups. Each group is scored by the sum of its two best experts, only the top 4 groups are eligible, and the 8 experts come from those. Under expert parallelism with groups mapped to nodes, a token’s experts then span at most 4 nodes, which bounds its all-to-all traffic.
- Auxiliary-loss-free balancing. Scores are sigmoids, not a softmax, and each expert has a bias added for selection only. During training the bias rises for underused experts and falls for overused ones; the combine weights use the raw scores, normalized over the chosen 8 and multiplied by
routed_scaling_factor(2.5).
DeepSeek-V2 uses a softmax and, for its large model, groups scored by their best expert. One class covers both:
class GroupedRouter(nn.Module):
"""DeepSeek's router. Experts are split into n_group groups; only the topk_group best groups
may be chosen (which bounds how many nodes a token's experts span under expert parallelism),
then the top k experts among them. V3 scores with a sigmoid and adds a per-expert bias,
adjusted during training to balance load, for SELECTION only; the weights use the raw scores."""
def __init__(self, cfg):
super().__init__()
self.cfg = cfg
self.weight = nn.Parameter(torch.zeros(cfg.n_routed_experts, cfg.hidden_size))
if cfg.scoring == "sigmoid":
self.e_score_correction_bias = nn.Parameter(torch.zeros(cfg.n_routed_experts), requires_grad=False)
def forward(self, x):
"""x [N, D] -> (logits, weights [N, k], experts [N, k]). (Your engine: Chapter 42)"""
c = self.cfg
logits = F.linear(x.float(), self.weight.float())
scores = logits.sigmoid() if c.scoring == "sigmoid" else logits.softmax(-1)
choice = scores + self.e_score_correction_bias if c.scoring == "sigmoid" else scores
if c.n_group > 1:
grouped = choice.view(-1, c.n_group, c.n_routed_experts // c.n_group)
group_score = grouped.topk(2, dim=-1).values.sum(-1) if c.group_score == "top2" else grouped.amax(-1)
keep = torch.zeros_like(group_score).scatter_(1, group_score.topk(c.topk_group, dim=-1).indices, 1)
choice = choice.masked_fill(~keep.bool().repeat_interleave(c.n_routed_experts // c.n_group, dim=1), float("-inf"))
experts = choice.topk(c.num_experts_per_tok, dim=-1).indices
weights = scores.gather(1, experts)
if c.norm_topk_prob:
weights = weights / (weights.sum(-1, keepdim=True) + 1e-20)
return logits, (weights * c.routed_scaling_factor).to(x.dtype), experts
class DeepseekMoE(nn.Module):
"""Routed experts (Chapter 27's stacked layout) plus always-on shared experts."""
def __init__(self, cfg):
super().__init__()
self.gate = GroupedRouter(cfg)
self.experts = Experts(cfg.n_routed_experts, cfg.hidden_size, cfg.moe_intermediate_size)
self.shared_experts = (SwiGLU(cfg.hidden_size, cfg.moe_intermediate_size * cfg.n_shared_experts)
if cfg.n_shared_experts else None)
def forward(self, x):
shape = x.shape
flat = x.reshape(-1, shape[-1])
_, weights, experts = self.gate(flat)
out = self.experts.forward_grouped(flat, weights, experts)
if self.shared_experts is not None:
out = out + self.shared_experts(flat)
return out.reshape(shape)
The registry
def build_model(raw):
"""An empty model for a config.json dict."""
kind = raw.get("model_type")
if kind in ("llama", "mistral", "qwen2", "qwen3"):
return Decoder(DecoderConfig.from_hf(raw))
if kind == "qwen3_moe":
from .moe import Qwen3Moe, Qwen3MoeConfig
return Qwen3Moe(Qwen3MoeConfig.from_hf(raw))
if kind in ("deepseek_v2", "deepseek_v3", "kimi_k2"):
return MLADecoder(MLAConfig.from_hf({**raw, "model_type": "deepseek_v3" if kind == "kimi_k2" else kind}))
raise ValueError(f"No model registered for model_type {kind!r}")
@torch.no_grad()
def load_weights(model, named_tensors):
"""Name-to-name copy. Per-expert checkpoint tensors (experts.E.gate_proj.weight) are stacked
into the [E, 2I, D] / [E, D, I] layout; layers past num_hidden_layers (DeepSeek-V3's
multi-token-prediction module) are skipped."""
from .loaders import assign
from .moe import stack_expert_tensors
cfg = model.cfg
params = dict(model.named_parameters())
loaded, pending = set(), {}
for name, value in named_tensors:
parts = name.split(".")
if name.startswith("model.layers.") and int(parts[2]) >= cfg.num_hidden_layers:
continue
if ".mlp.experts." in name and parts[5].isdigit():
pending.setdefault(int(parts[2]), {})[(int(parts[5]), parts[6])] = value
continue
if name == "lm_head.weight" and cfg.tie_word_embeddings:
continue
if name not in params:
raise ValueError(f"Unexpected tensor {name}")
assign(params[name], value, name)
loaded.add(name)
for layer, tensors in pending.items():
experts = model.model.layers[layer].mlp.experts
e, inter = experts.gate_up_proj.shape[0], experts.down_proj.shape[2]
gate_up, down = stack_expert_tensors(tensors, e, inter)
assign(experts.gate_up_proj, gate_up, f"layer {layer} experts")
assign(experts.down_proj, down, f"layer {layer} experts")
loaded |= {f"model.layers.{layer}.mlp.experts.gate_up_proj", f"model.layers.{layer}.mlp.experts.down_proj"}
missing = set(params) - loaded - ({"lm_head.weight"} if cfg.tie_word_embeddings else set())
if missing:
raise ValueError(f"Missing tensors: {sorted(missing)[:5]}")
return model.eval()
def load_registered(directory, device="cpu", dtype=torch.bfloat16):
"""Any registered architecture from a local snapshot (Chapter 18's loader, generalized)."""
from .loaders import read_config
from .safetensors_io import snapshot_tensors
raw = read_config(directory)
if raw.get("quantization_config"):
raise ValueError("Quantized checkpoints: see formats.hf_quant (Chapter 39)")
with torch.device("meta"):
model = build_model(raw)
model = model.to_empty(device=device).to(dtype)
if model.cfg.tie_word_embeddings:
model.lm_head.weight = model.model.embed_tokens.weight
return load_weights(model, snapshot_tensors(directory))
build_model maps config.json’s model_type to a class; load_weights is Chapter 18’s name-to-name copy, plus stacking for per-expert tensors (experts.17.gate_proj.weight) and skipping DeepSeek-V3’s multi-token-prediction layer, which follows the last regular layer (Chapter 37 could use it as a drafter). loaders.load_model now falls through to this registry for anything it doesn’t handle itself. Kimi K2 declares model_type: kimi_k2 with DeepSeek-V3’s architecture, so it maps to the same class.
GGUF files get the same coverage. Llama-architecture files (which include Mistral) store Q and K rows permuted for llama.cpp’s interleaved RoPE, so load_gguf undoes the permutation (Chapter 38), on quantized rows too: GGML quantizes each row independently, so permuting rows of codes and scales is exact. Llama 3’s frequency scaling arrives as a rope_freqs tensor of per-pair divisors; Qwen2 files carry attn_q.bias tensors. export_gguf writes both architectures, which the round-trip test uses.
Build it
Engine milestone 42: more models. Implement rope_inv_freq, DecoderAttention.forward, MLAAttention.paged_forward and GroupedRouter.forward in engine/models.py.
pytest tests/test_ch42_models.py
python run.py models
The tests compare logits with transformers for Llama 3 (scaled RoPE, tied embeddings), Llama with linear scaling and biases, Mistral (sliding window), Qwen2 (Q/K/V biases, windows on later layers, YaRN), DeepSeek-V3 (MLA with a low-rank query, YaRN, grouped sigmoid routing) and DeepSeek-V2-Lite (plain query, softmax routing, two shared experts); check that the engine generates the same tokens as a cache-free forward for each; load checkpoints written by save_pretrained; check the latent pool’s size; and round-trip Llama and Qwen2 GGUF files.
What hasn’t been checked here: real checkpoints. The test environment can’t download from the Hugging Face Hub, so every comparison uses randomly initialized models with the real architectures (Appendix F). Load a real Llama-3.2-1B or DeepSeek-V2-Lite and compare greedy outputs with transformers before trusting a new family.
Stretch exercises
- ★★ Add Gemma 3: RMSNorm with
(1 + weight), embeddings scaled by $\sqrt{d}$, GeGLU MLPs, Q/K norms, an extra norm after each sublayer, and five windowed layers for every full one. Compare withtransformers. Where: configuration, layers and registry inengine/models.py, with execution support inengine/serve/model.py. - ★★★ Free out-of-window blocks: give windowed and full layers separate block tables in
BlockManager, and recycle a request’s windowed-layer blocks once every query has moved past them. How many more Mistral requests fit in the same pool at 32K context? Where:BlockManagerinengine/serve/blocks.py, with per-layer tables inengine/serve/batch.pyand attention dispatch inengine/serve/model.py. - ★★ Implement MLA’s prefill path: decompress the context’s latents for a long prompt and run ordinary attention with 128 heads. At what prompt length does it beat the absorbed form on your hardware? Where:
MLAAttention.forward/paged_forwardinengine/models.py. - ★★ Load DeepSeek-V3’s FP8 checkpoint: extend Chapter 39’s
load_quantizedto the registry, buildingFP8BlockLinearmodules fromweightandweight_scale_inv. Where:load_quantized/buildinengine/formats/hf_quant.py, with registry loading inengine/models.py. - ★★★ Write a Triton MLA decode kernel: one program per (request, head group), with the 576-wide rows read once for 16 or more heads. Compare with the generic kernel run with one KV head. Where: new
engine/kernels/triton_mla.py, selected byMLAAttention.paged_forwardinengine/models.py.
Check your understanding
- Why must
DecoderAttentionleaveq_normundefined instead of setting it to an identity module? - Which RoPE pairs does Llama 3’s scaling leave alone, and why can it?
- Why does “dynamic NTK” scaling clash with prefix caching?
- How does MLA cache 576 numbers per token while giving 128 heads distinct keys and values?
- Why can’t the rotary part of MLA’s key be absorbed like the rest?
- When is the decompressed form of MLA cheaper than the absorbed form?
- What does group-limited routing bound, and why does that matter under expert parallelism?
Going deeper
- Chen et al., Extending Context Window of Large Language Models via Positional Interpolation (2023); Peng et al., YaRN: Efficient Context Window Extension of Large Language Models (ICLR 2024); Meta’s Llama 3 report, The Llama 3 Herd of Models (2024), for its RoPE scaling.
- Jiang et al., Mistral 7B (2023); Beltagy et al., Longformer (2020), for sliding-window attention.
- DeepSeek-AI, DeepSeek-V2 (2024) for MLA and DeepSeek-V3 Technical Report (2024) for auxiliary-loss-free balancing and group-limited routing; the FlashMLA repository.
- vLLM’s
vllm/model_executor/models/(one file per family, andregistry.py) andvllm/v1/attention/backends/mla/; llama.cpp’ssrc/llama-model.cppandconvert_hf_to_gguf.py.
43. Beyond text: images, embeddings, reranking and many LoRAs
In this chapter
- Turn an image into decoder input rows with a small vision transformer and projector.
- Give image patches three-dimensional rotary positions, while keeping the scheduler's physical token positions.
- Serve embeddings and query-document scores without an autoregressive loop.
- Run requests with different LoRA adapters in the same batch, using adapter slots and shrink/expand kernels.
- Make prefix reuse depend on everything that changes activations, including images and adapter weights.
You will build
VisionEncoder.forward, expand_images and mrope_cos_sin in engine/multimodal.py; pool_hidden in engine/pooling.py; grouped_lora and AdapterBank.load in engine/serve/adapters.py; ModelRunner.prepare_features; and the BGMV/SGMV kernels.
Time: 6-8 hours. GPU: optional; the kernels also run in Triton's interpreter.
What changes at the input and output
The decoder sees a matrix of hidden vectors, however those vectors were obtained. A text embedding lookup is one source. A vision encoder is another. The final hidden vectors can go to an LM head to generate tokens, to a pooling operation to produce an embedding, or to a trained scalar head to score a query-document pair.
LoRA changes something different: which function a linear layer computes for each row. Chapter 22 attached one adapter to a model. A service needs to keep a base model resident while requests choose different adapters, without splitting the batch into one base-model forward per adapter.
| request | input boundary | work in the model | output boundary |
|---|---|---|---|
| text generation | token IDs | causal decoder, repeated decode | incremental text |
| image + text | text IDs and projected image patches | the same causal decoder | incremental text |
| embeddings | token IDs and a padding mask | one encoder/decoder forward | pooled vector |
| reranking | tokenized query-document pairs | one encoder forward and score head | scalar scores |
| any supported LoRA generation | token IDs and adapter identity | base projections plus per-row low-rank updates | incremental text |
Images and adapters belong in Chapter 31’s request and batch metadata. Pooling belongs beside generation: it needs one bounded forward, not KV allocation and hundreds of decode steps.
An image becomes tokens
Take a 16 × 16 RGB image and divide it into 4 × 4 patches. There are 16 patches; each contains 48 pixel values. A convolution with kernel and stride 4 is a learned linear projection of those values into the vision width. Add a learned position vector to every patch, then run transformer layers with bidirectional attention: every patch may see every other patch. Finally, an MLP projects each patch into the decoder’s hidden width.
RGB pixels [B,3,16,16]
│ patch projection + learned positions
▼
16 patch rows [B,16,vision_width]
│ bidirectional transformer + normalization
│ trained projector
▼
decoder rows [B,16,hidden_size]
class VisionEncoder(nn.Module):
def __init__(self, config=None):
super().__init__()
self.cfg = c = config or VisionConfig()
if c.image_size % c.patch_size or c.width % c.heads:
raise ValueError("Image size must divide into patches, width into heads")
self.patch = nn.Conv2d(3, c.width, c.patch_size, stride=c.patch_size)
self.position = nn.Parameter(torch.randn((c.image_size // c.patch_size) ** 2, c.width) * 0.02)
self.blocks = nn.ModuleList(nn.TransformerEncoderLayer(
c.width, c.heads, 4 * c.width, dropout=0, activation="gelu", batch_first=True, norm_first=True)
for _ in range(c.layers))
self.norm = nn.LayerNorm(c.width)
self.projector = nn.Sequential(nn.Linear(c.width, c.output_width), nn.GELU(),
nn.Linear(c.output_width, c.output_width))
def forward(self, pixels):
"""Normalized RGB [B,3,H,W] -> projected patch rows [B,P,D]. (Your engine: Chapter 43)"""
if pixels.shape[1:] != (3, self.cfg.image_size, self.cfg.image_size):
raise ValueError("Pixels must match the configured RGB image size")
x = self.patch(pixels).flatten(2).transpose(1, 2) + self.position
for block in self.blocks:
x = block(x) # bidirectional attention within the image
return self.projector(self.norm(x))
There is no magic in the projector. It has to be trained so that its outputs mean something to the language model. Random weights are enough to test shapes and serving correctness; they do not make an image assistant. Real VLMs may resize into several crops, merge groups of patches, use a different vision tower or insert boundary tokens. Loading a Qwen-VL checkpoint requires all those exact conventions, alongside its weights.
At the text boundary, a prompt contains one reserved <image> placeholder for each image. The processor expands each placeholder into as many reserved IDs as there are projected patches, and records which absolute positions get which vectors:
@dataclass
class VisionPrompt:
token_ids: list
image_rows: dict # absolute token position -> projected embedding (CPU)
positions: torch.Tensor # [3,T]: time, height, width
delta: int # next text position minus physical prompt length
sections: tuple
def features(self):
return {"image_rows": self.image_rows, "mrope_positions": self.positions,
"mrope_delta": self.delta, "mrope_sections": self.sections}
def expand_images(ids, placeholder_id, embeddings, grids, sections):
"""One placeholder becomes T*H*W patch rows; text shares all three axes.
Resume text at max(image coordinates)+1, not after the number of patches. (Your engine: Chapter 43)
"""
if ids.count(placeholder_id) != len(embeddings) or len(grids) != len(embeddings):
raise ValueError("Each image needs exactly one placeholder and one grid")
tokens, axes, rows, cursor, index = [], [], {}, 0, 0
for token in ids:
if token != placeholder_id:
tokens.append(token), axes.append((cursor, cursor, cursor))
cursor += 1
continue
grid, features = grids[index], embeddings[index].detach().cpu().clone()
t, h, w = grid
if min(grid) < 1 or features.ndim != 2 or len(features) != t * h * w:
raise ValueError("Projected image length does not match its grid")
for tt in range(t):
for hh in range(h):
for ww in range(w):
rows[len(tokens)] = features[(tt * h + hh) * w + ww]
tokens.append(placeholder_id), axes.append((cursor + tt, cursor + hh, cursor + ww))
cursor += max(grid)
index += 1
return VisionPrompt(tokens, rows, torch.tensor(axes, dtype=torch.long).T.contiguous(),
cursor - len(tokens), tuple(sections))
The reserved IDs give the scheduler a physical sequence to count and page. Their embedding lookup is replaced by the projected rows before the decoder runs. A 16-patch image consumes 16 KV positions, even though it began as one placeholder. Expand before validating prompt length and reserving KV memory.
VisionProcessor supplies a fixed-resolution demo processor. It accepts base64 image data URIs, checks compressed bytes and pixel count, converts to RGB and normalizes to [-1, 1]. Its decoder inputs are detached CPU tensors, which can cross Chapter 36’s process queues. Install Pillow for this boundary (uv pip install pillow); tensor-only milestones don’t need it. The service does not fetch remote URLs. A deployed processor needs a controlled fetcher and the model’s trained preprocessing recipe if URLs are part of its API.
Two meanings of position
Ordinary RoPE has one coordinate, the text position. Image patches have spatial coordinates; videos also have time. M-RoPE assigns different frequency pairs to time, height and width. With a 12-wide head there are 6 rotary pairs. Sections (2, 2, 2) assign 2 pairs to each axis; the selected angles are then repeated for the head’s split-half pairing (Chapter 17).
For text, all three coordinates are equal. M-RoPE then reduces exactly to ordinary RoPE. For an image, flattening patch rows should not erase the fact that neighbouring rows may be vertically or horizontally adjacent:
physical row in A <image> B | content | time | height | width |
|---|---|---|---|---|
| 0 | A | 0 | 0 | 0 |
| 1 | image patch (0,0) | 1 | 1 | 1 |
| 2 | image patch (0,1) | 1 | 1 | 2 |
| 3 | image patch (1,0) | 1 | 2 | 1 |
| 4 | image patch (1,1) | 1 | 2 | 2 |
| 5 | B | 3 | 3 | 3 |
The next generated token has physical position 6 but rotary position 4. Store the continuation delta, here -2, and compute text rotary positions as physical position + delta after the prompt. Do not change slot_mapping, causal masking or sequence lengths: those still use the physical order.
def mrope_cos_sin(positions, head_dim, theta, sections, dtype):
"""positions [3,N], sections count frequency pairs on each axis. (Your engine: Chapter 43)
Split-half RoPE repeats the chosen angles for both halves of the head.
"""
if positions.ndim != 2 or positions.shape[0] != 3 or len(sections) != 3 \
or min(sections) < 0 or sum(sections) * 2 != head_dim:
raise ValueError("M-RoPE needs three sections totaling head_dim/2 pairs")
inv = theta ** (-torch.arange(0, head_dim, 2, device=positions.device).float() / head_dim)
selected, start = [], 0
for axis, count in enumerate(sections):
selected.append(positions[axis, :, None].float() * inv[start:start + count])
start += count
angles = torch.cat(selected, -1)
angles = torch.cat((angles, angles), -1)[:, None]
return angles.cos().to(dtype), angles.sin().to(dtype)
This implements the position mechanism described by Qwen2-VL, with a deliberately small image processor. The model must choose the sections; the engine fixes them on first use if they are not configured, then rejects mismatches. This path uses ordinary rotary frequencies and refuses scaled RoPE rather than silently mixing the two recipes.
Chunked prefill, cache hits and recomputation
An image can span several prefill chunks. Storing projected rows by absolute prompt position makes that case simple: when a request computes positions c:c+n, replace only image rows inside that interval. A cached prefix skips rows already computed; recompute preemption reads the same saved vectors again. M-RoPE coordinates are sliced the same way, with the continuation delta used during decode.
@torch.inference_mode()
def prepare_features(self, batch, scheduled):
"""Map absolute feature positions and adapter slots to this step's ragged rows. (Your engine: Chapter 43)
A prefill chunk can cut through an image; recomputation uses the same saved rows.
"""
extras = [r.extra for r, _ in scheduled]
if any("lora_slot" in e for e in extras):
slots = torch.full_like(batch.input_ids, -1)
for (request, n), start in zip(scheduled, batch.meta.query_start_loc_cpu):
slots[start:start + n] = request.extra.get("lora_slot", -1)
batch.extra["lora_ids"] = slots
mrope = [e for e in extras if "mrope_positions" in e]
if mrope:
sections = mrope[0]["mrope_sections"]
if any(e["mrope_sections"] != sections for e in mrope):
raise ValueError("All requests must use the model's M-RoPE sections")
axes = batch.positions[None].expand(3, -1).clone()
for (request, n), start in zip(scheduled, batch.meta.query_start_loc_cpu):
if "mrope_positions" not in request.extra:
continue
for offset in range(n):
pos = request.num_computed_tokens + offset
if pos < request.num_prompt_tokens:
axes[:, start + offset] = request.extra["mrope_positions"][:, pos].to(self.device)
else:
axes[:, start + offset] = pos + request.extra["mrope_delta"]
batch.extra.update(mrope_positions=axes, mrope_sections=sections)
replacements = []
for (request, n), start in zip(scheduled, batch.meta.query_start_loc_cpu):
for pos, row in request.extra.get("image_rows", {}).items():
offset = pos - request.num_computed_tokens
if 0 <= offset < n:
replacements.append((start + offset, row))
if replacements:
embeds = self.flat.backbone.embed_tokens(batch.input_ids)
for row, value in replacements:
embeds[row] = value.to(embeds)
batch.extra["embeds"] = embeds
Two requests can have identical placeholder IDs and entirely different images. Chapter 31’s token hash alone would alias their KV. validate_features snapshots projected rows and derives hashes from their actual values, positions, sections and continuation delta. Those hashes are added to Request.cache_key, so the block hash chain identifies the inputs the decoder actually consumed. The caller cannot assert that two images share a hash. This conservative implementation salts every block with all image inputs, including text blocks before an image; a more selective cache can begin salting at the first affected block.
Images are encoded in the frontend, then their CPU rows travel through AsyncLLM.generate(..., features=...). For a vision-capable Server, set multimodal_processor=VisionProcessor(...) and ensure the tokenizer recognizes its reserved marker. A text-only server rejects image parts. Dropping them and answering only the text would look successful while computing the wrong request.
Embeddings: one forward, one vector
An embedding model is trained to put related inputs near each other. Pooling determines which vector represents the sequence:
- Mean: sum the visible token vectors and divide by their count.
- Last: select the last visible token, useful for decoder embedding models.
- CLS: select the designated first token, useful for encoders trained with that convention.
For cosine similarity, normalize the pooled vector to unit length. Then the dot product equals cosine similarity. The pooling mode, query/document prefixes and normalization must match training; a random causal decoder with mean pooling proves the plumbing, not retrieval quality.
def pool_hidden(hidden, mask, mode="mean", normalize=True):
"""Exclude padding, select last/CLS or average, optionally unit-normalize. (Your engine: Chapter 43)"""
if hidden.ndim != 3 or mask.shape != hidden.shape[:2] or not mask.any(-1).all():
raise ValueError("Each sequence needs at least one visible token")
if mode == "mean":
result = (hidden.float() * mask[..., None]).sum(1) / mask.sum(1, keepdim=True)
elif mode == "last":
index = torch.arange(mask.shape[1], device=mask.device).expand_as(mask).masked_fill(~mask.bool(), -1).max(1).values
result = hidden[torch.arange(len(hidden), device=hidden.device), index].float()
elif mode == "cls":
if not mask[:, 0].all():
raise ValueError("CLS pooling requires an unpadded first token")
result = hidden[:, 0].float()
else:
raise ValueError("Pooling mode must be mean|last|cls")
return F.normalize(result, dim=-1) if normalize else result
Padding must contribute neither vectors nor a denominator. The mean of [h0, h1, padding] is (h0+h1)/2, not /3. Last pooling finds the last visible index, which also works with left padding. The supplied DecoderEncoder uses right padding and causal attention, so valid tokens cannot see subsequent padding. HFEncoder instead passes the attention mask to a bidirectional encoder.
class PoolingService:
def __init__(self, encoder, tokenizer, model_name, mode="mean", normalize=True, score_head=None,
max_model_len=512, max_batch=32, pad_id=0):
self.encoder, self.tokenizer, self.model_name = encoder.eval(), tokenizer, model_name
self.mode, self.normalize, self.score_head = mode, normalize, score_head
self.max_model_len, self.max_batch, self.pad_id = max_model_len, max_batch, pad_id
self.lock = asyncio.Lock() # bound device work; don't overlap calls on one encoder
def inputs(self, value):
if isinstance(value, str) or isinstance(value, list) and value and all(type(x) is int for x in value):
value = [value]
if not isinstance(value, list) or not 1 <= len(value) <= self.max_batch:
raise ValueError("Input must be a nonempty bounded batch")
ids = []
for item in value:
row = self.tokenizer.encode(item) if isinstance(item, str) else item
if not isinstance(row, list) or not row or not all(type(x) is int and x >= 0 for x in row):
raise ValueError("Each input must be text or nonempty token ids")
if len(row) > self.max_model_len:
raise ValueError("Embedding input exceeds the model context limit")
ids.append(row)
return ids
@torch.inference_mode()
def encode(self, ids):
device = next(self.encoder.parameters()).device
tokens = torch.full((len(ids), max(map(len, ids))), self.pad_id, dtype=torch.long, device=device)
mask = torch.zeros_like(tokens, dtype=torch.bool)
for i, row in enumerate(ids):
tokens[i, :len(row)] = torch.tensor(row, device=device)
mask[i, :len(row)] = True
hidden = self.encoder(tokens, mask)
pooled = pool_hidden(hidden, mask, self.mode, self.normalize)
if self.score_head is not None:
return self.score_head(pooled.to(next(self.score_head.parameters()))).float().flatten().cpu()
return pooled.cpu()
async def __call__(self, body):
if self.score_head is not None:
raise ValueError("This service is a reranker")
if body.get("dimensions") is not None:
raise ValueError("Dimension truncation needs a model trained for it")
fmt = body.get("encoding_format", "float")
if fmt not in ("float", "base64"):
raise ValueError("encoding_format must be float|base64")
ids = self.inputs(body.get("input"))
async with self.lock:
vectors = await asyncio.to_thread(self.encode, ids)
embeddings = vectors.tolist() if fmt == "float" else [
base64.b64encode(v.contiguous().numpy().astype("<f4").tobytes()).decode() for v in vectors]
count = sum(map(len, ids))
return {"object": "list", "model": self.model_name,
"data": [{"object": "embedding", "index": i, "embedding": v} for i, v in enumerate(embeddings)],
"usage": {"prompt_tokens": count, "total_tokens": count}}
async def rerank(self, body):
if self.score_head is None:
raise ValueError("Reranking needs a trained scalar score head")
query, docs = body.get("query"), body.get("documents")
if not isinstance(query, str) or not isinstance(docs, list) or not docs or not all(isinstance(d, str) for d in docs):
raise ValueError("Reranking needs a query and a nonempty list of text documents")
# The tokenizer owns the trained pair convention; do not invent a separator.
pair = getattr(self.tokenizer, "encode_pair", None)
if pair is None:
raise ValueError("Reranking requires a tokenizer with encode_pair(query, document)")
ids = self.inputs([pair(query, doc) for doc in docs])
top_n = body.get("top_n", len(docs))
if type(top_n) is not int or not 1 <= top_n <= len(docs):
raise ValueError("top_n must be between one and the number of documents")
async with self.lock:
scores = await asyncio.to_thread(self.encode, ids)
if scores.shape != (len(docs),):
raise ValueError("Score head must return one scalar per query-document pair")
order = sorted(range(len(docs)), key=lambda i: (-float(scores[i]), i))[:top_n]
return {"model": self.model_name, "results": [{"index": i, "relevance_score": float(scores[i]),
"document": {"text": docs[i]}} for i in order],
"usage": {"total_tokens": sum(map(len, ids))}}
PoolingService accepts text, IDs or a bounded batch of either. Its lock permits one device forward at a time; asyncio.to_thread keeps that forward off the HTTP event loop. The maximum batch and context limits bound each allocation. This is a separate padded encoder path, not continuous batching across HTTP embedding requests.
Chapter 36 already had an /v1/embeddings hook. Wire it to a trained local encoder:
from transformers import AutoModel
from izh.pooling import HFEncoder, PoolingService
# `server` is Chapter 36's Server; `tokenizer` must implement this model's text recipe.
encoder = AutoModel.from_pretrained("models/my-embedding-model", local_files_only=True).eval()
server.embedder = PoolingService(HFEncoder(encoder), tokenizer, server.model_name,
mode="mean", normalize=True, max_model_len=512)
The response preserves input order, counts non-padding prompt tokens and supports float lists or base64 little-endian FP32 vectors. Arbitrary dimension truncation is refused: it is a training property, not a free compression operation.
Reranking pairs a query with every candidate document, processes each pair jointly, then applies a trained scalar head to a pooled hidden vector. Unlike independent embeddings, query and document tokens interact inside attention. It costs one forward per pair and usually follows a cheap embedding search over many documents. Set server.reranker to a PoolingService with score_head, and provide tokenizer.encode_pair(query, document) with the trained pair convention. /v1/rerank returns scores sorted by descending relevance and original indices; it is an extension, not an OpenAI standard endpoint. A score need not be a calibrated probability.
Many adapters, one base-model batch
For token row $i$ selecting adapter $s_i$, Chapter 22’s equation becomes
$$y_i = W x_i + \frac{\alpha_{s_i}}{r_{s_i}} B_{s_i}(A_{s_i}x_i).$$
Compute the base projection over all rows once. Then shrink each row from width $D$ to its adapter’s rank, expand to the output width and add. Rows selecting no adapter get zero correction. Padding adapters to a common maximum rank lets the device buffers have fixed shapes; unused rows of A and columns of B are zero. The scale is folded into B at load time.
def grouped_lora(x, a, b, slots):
"""x [N,K], A [S,R,K], B [S,D,R], slots [N]; -1 means base only.
Group rows by adapter, shrink then expand, and scatter to original order. (Your engine: Chapter 43)
B already includes alpha/r. Padded ranks have zero weights.
"""
out = x.new_zeros((len(x), b.shape[1]))
for slot in range(a.shape[0]):
rows = torch.where(slots == slot)[0]
if rows.numel():
out[rows] = F.linear(F.linear(x[rows], a[slot]), b[slot])
return out
The PyTorch reference groups rows by adapter and scatters them back. Two kernel shapes cover serving:
- BGMV, batched gathered matrix-vector multiplication, handles decode rows, each with its own adapter slot.
- SGMV, segmented gathered matrix-vector multiplication, sorts prefill rows by adapter so several rows reuse the same weights. Segment offsets describe the runs.
@triton.jit
def bgmv_kernel(X, W, IDS, Y, K: tl.constexpr, D: tl.constexpr, BK: tl.constexpr):
"""One program per row and output coordinate, selected adapter's matrix. (Your engine: Chapter 43)"""
row, out = tl.program_id(0), tl.program_id(1)
slot = tl.load(IDS + row)
k = tl.arange(0, BK)
acc = tl.full((), 0, tl.float32)
for start in range(0, K, BK):
kk = start + k
x = tl.load(X + row * K + kk, kk < K, other=0).to(tl.float32)
w = tl.load(W + (tl.maximum(slot, 0) * D + out) * K + kk,
(kk < K) & (slot >= 0), other=0).to(tl.float32)
acc += tl.sum(x * w, 0)
tl.store(Y + row * D + out, acc)
@triton.jit
def sgmv_kernel(X, W, STARTS, Y, K: tl.constexpr, D: tl.constexpr, BK: tl.constexpr, BR: tl.constexpr):
"""Each adapter owns a contiguous segment; process BR rows together. (Your engine: Chapter 43)"""
tile, out, slot = tl.program_id(0), tl.program_id(1), tl.program_id(2)
start, end = tl.load(STARTS + slot), tl.load(STARTS + slot + 1)
rows = start + tile * BR + tl.arange(0, BR)
k = tl.arange(0, BK)
acc = tl.zeros((BR,), tl.float32)
for base in range(0, K, BK):
kk = base + k
x = tl.load(X + rows[:, None] * K + kk[None, :],
(rows[:, None] < end) & (kk[None, :] < K), other=0).to(tl.float32)
w = tl.load(W + (slot * D + out) * K + kk, kk < K, other=0).to(tl.float32)
acc += tl.sum(x * w[None, :], 1)
tl.store(Y + rows * D + out, acc, rows < end)
The kernels accumulate in FP32, mask odd widths and treat slot -1 as base-only. They perform both shrink and expand. These are readable vector-reduction kernels; the SGMV wrapper sorts at each projection and launches empty segments. Production implementations reuse grouping metadata and use tiled tensor-core GEMMs. Punica develops this batching approach; S-LoRA adds unified paging and adapter management at much larger scale.
Slots are storage; identities are content
class AdapterBank:
"""A fixed device budget for adapters; slots cannot be overwritten while leased."""
def __init__(self, flat, capacity=8, max_rank=16, targets=("q_proj", "v_proj"), backend="torch"):
if capacity < 1 or max_rank < 1 or backend not in ("torch", "triton"):
raise ValueError("Need positive capacity/rank and backend torch|triton")
if getattr(flat, "adapters", None) is not None:
raise ValueError("An adapter bank is already installed")
self.capacity, self.max_rank, self.backend = capacity, max_rank, backend
self.names, self.keys, self.refs, self.layers = {}, {}, [0] * capacity, {}
self.rows, self.decode = None, False
for name, module in list(flat.model.named_modules()):
if isinstance(module, nn.Linear) and name.rsplit(".", 1)[-1] in targets:
parent, _, attr = name.rpartition(".")
wrapped = AdapterLinear(module, self)
setattr(flat.model.get_submodule(parent), attr, wrapped)
self.layers[name] = wrapped
if not self.layers:
raise ValueError("No adapter targets found; install after projection fusion with matching target names")
flat.adapters = self
@torch.no_grad()
def load(self, name, weights, alpha=None):
"""weights: {module_name: (A [rank,in], B [out,rank])}. Validate all before writing.
A new generation's content digest prevents stale prefix-cache hits. (Your engine: Chapter 43)
"""
if not name or not weights or set(weights) - set(self.layers):
raise ValueError("An adapter needs a name and known target modules")
slot = self.names.get(name)
if slot is not None and self.refs[slot]:
raise ValueError("Cannot replace an adapter used by active requests")
if slot is None:
used = set(self.names.values())
slot = next((s for s in range(self.capacity) if s not in used), None)
if slot is None:
raise ValueError("Adapter slots full; unload an idle adapter first")
digest = hashlib.sha256()
prepared = {}
for target, (a, b) in sorted(weights.items()):
layer = self.layers[target]
rank = a.shape[0]
if rank < 1 or rank > self.max_rank or a.shape != (rank, layer.in_features) \
or b.shape != (layer.out_features, rank):
raise ValueError(f"{target}: incompatible LoRA shapes or rank")
scale = (alpha if alpha is not None else rank) / rank
a, b = a.detach().to(layer.a), b.detach().to(layer.b) * scale
if not torch.isfinite(a).all() or not torch.isfinite(b).all():
raise ValueError("Adapter contains non-finite weights")
for tensor in (a, b):
digest.update(str((target, tensor.shape, tensor.dtype)).encode())
digest.update(tensor.contiguous().cpu().view(torch.uint8).numpy().tobytes())
prepared[target] = a, b
for target, layer in self.layers.items():
layer.a[slot].zero_(), layer.b[slot].zero_()
if target in prepared:
a, b = prepared[target]
layer.a[slot, :len(a)].copy_(a)
layer.b[slot, :, :len(a)].copy_(b)
self.names[name], self.keys[name] = slot, digest.hexdigest()
return slot
def load_peft(self, name, directory):
"""The Chapter 22 PEFT safetensors layout, with explicit rejection of extra features."""
from ..safetensors_io import load_file
path = Path(directory)
config = json.loads((path / "adapter_config.json").read_text())
if config.get("peft_type") != "LORA" or config.get("use_dora") or config.get("use_rslora") \
or config.get("bias", "none") != "none" or config.get("modules_to_save") \
or config.get("rank_pattern") or config.get("alpha_pattern"):
raise ValueError("Only plain LoRA with uniform alpha, no trained bias or saved modules is supported")
tensors = load_file(path / "adapter_model.safetensors")
weights = {}
for key in tensors:
if not key.endswith(".lora_A.weight"):
continue
target = key.removeprefix("base_model.model.").removesuffix(".lora_A.weight")
weights[target] = tensors[key], tensors[key.replace(".lora_A.", ".lora_B.")]
expected = {f"base_model.model.{t}.lora_{part}.weight" for t in weights for part in ("A", "B")}
if set(tensors) != expected:
raise ValueError("Unsupported or incomplete adapter tensors")
return self.load(name, weights, config["lora_alpha"])
def unload(self, name):
slot = self.names[name]
if self.refs[slot]:
raise ValueError("Cannot unload an adapter used by active requests")
del self.names[name], self.keys[name]
def acquire(self, name):
if name not in self.names:
raise ValueError(f"Unknown LoRA adapter {name!r}")
slot = self.names[name]
self.refs[slot] += 1
return slot, ("lora", self.keys[name])
def release(self, slot):
if self.refs[slot] < 1:
raise RuntimeError("Adapter lease released twice")
self.refs[slot] -= 1
@contextmanager
def activate(self, rows, decode=False):
previous = self.rows, self.decode
self.rows, self.decode = rows, decode
try:
yield
finally:
self.rows, self.decode = previous
AdapterBank wraps selected linears with resident slot buffers. Loading validates every target and shape before writing any weights. Requests lease a slot at admission and release it on completion or cancellation; a preempted request keeps its lease. Replacing or unloading a leased adapter fails. Otherwise a request could prefill with adapter A and resume with B, keeping incompatible KV.
The prefix cache uses a digest of the adapter’s effective weights, not the integer slot. Reusing slot 0 for a different adapter must not reuse slot 0’s old KV. Conversely, two adapter names with identical effective weights can safely share prefixes. Changing base-model weights still requires draining the engine and clearing its caches.
load_peft reads the plain LoRA safetensors format used in Chapter 22. It refuses DoRA, rank-stabilized LoRA, trained biases, saved full modules and per-layer rank/alpha patterns. The launcher preloads named adapters:
python serve.py --model-dir models/Qwen3-0.6B --device cpu --dtype fp32 \
--lora support=runs/support-adapter --lora coding=runs/coding-adapter
Requests use a name: {"prompt":"...", "lora":"support", "max_tokens":32}. The name and feature inputs travel through the process queue; the core derives the slot and cache salt. Clients cannot supply slot numbers or load arbitrary paths. Each API choice from n>1 gets its own lease, while sharing compatible prefix blocks.
Install the bank after any projection fusion, targeting the resulting module names. The launcher uses unfused Q/V targets. Quantized adapter bases, expert adapters and distributed adapter sharding need further work. Feature batches take the eager runner path because graph buffers don’t yet hold these inputs; feature requests with a drafter are refused until the draft path also consumes them. The dense paged cache, batching, prefix reuse and preemption remain shared with text generation.
Build it
Engine milestone 43: beyond text. Implement:
- in
engine/multimodal.py:VisionEncoder.forward,expand_imagesandmrope_cos_sin; - in
engine/serve/runner.py:ModelRunner.prepare_features; - in
engine/serve/adapters.py:grouped_loraandAdapterBank.load; - in
engine/kernels/triton_lora.py:bgmv_kernelandsgmv_kernel; - in
engine/pooling.py:pool_hidden.
Then run:
pytest tests/test_ch43_multimodal_lora_embeddings.py
python run.py features
The tests check ViT batch equivalence; exact text RoPE equivalence and axis selection; image generation against a separate cache-free forward; image chunks mixed with text; warm-image prefix hits and changed-image misses; mixed adapters against separately merged models in synchronous and asynchronous scheduling; slot leases and cancellation; odd-shaped shrink/expand kernels; padding-independent embeddings, float/base64 HTTP responses and reranking order.
Stretch exercises
- ★★ Add a trained VLM’s patch merger and processor. Compare image rows, rotary positions and logits with its official implementation before testing text answers. Where:
VisionEncoder/VisionProcessorinengine/multimodal.py, with feature insertion inengine/serve/runner.py. - ★★ Cache vision-encoder outputs by image bytes, preprocessing configuration and encoder revision. Bound the host cache; measure repeated-image latency separately from decoder prefix hits. Where: feature caching around
VisionPrompt.featuresinengine/multimodal.py. - ★★★ Replace the SGMV reductions with tiled matmuls and reuse the row permutation across layers. Measure crossover by adapter rank, tokens per adapter and number of adapters. Where:
sgmv_kernelinengine/kernels/triton_lora.py, with permutation reuse inengine/serve/adapters.py. - ★★ Add dynamic batching for embeddings with a bounded queue and a maximum hold time. Compare throughput and p99 latency against the current per-call lock. Where:
PoolingServiceinengine/pooling.py, with endpoint dispatch inengine/serve/api.py. - ★★★ Add adapter inputs to graph buffers, bucket by maximum rank, and test slot replacement after every request in a captured bucket retires. Where: adapter buffers in
engine/serve/graphs.pyandengine/serve/runner.py, with slot leases inengine/serve/adapters.py.
Check your understanding
- Why does a placeholder expand before the engine checks the context limit?
- Why do KV positions and M-RoPE positions differ after an image?
- What must an image prefix-cache key identify besides placeholder IDs?
- Why doesn’t mean pooling turn an ordinary decoder into a useful embedding model?
- Why must padding be excluded from both the pooled sum and its divisor?
- Which work is shared across requests with different LoRAs, and which work depends on the adapter?
- Why is an adapter slot number an unsafe prefix-cache identity?
- Why does a preempted request keep its adapter lease?
Going deeper
- Dosovitskiy et al., An Image is Worth 16x16 Words, for ViT; Liu et al., Visual Instruction Tuning, for projecting image features into a language model.
- Wang et al., Qwen2-VL, for M-RoPE and image/video position conventions.
- Reimers and Gurevych, Sentence-BERT, for trained sentence embeddings and the distinction from cross-encoder scoring.
- Chen et al., Punica; Sheng et al., S-LoRA, for gathered low-rank kernels and serving many adapters.
44. Benchmark the whole engine, then harden it
In this chapter
- Measure user latency and throughput under an arrival process, rather than timing an isolated forward pass.
- Separate TTFT, TPOT, token ITL, streamed chunk gaps and end-to-end latency.
- Count goodput under explicit service-level objectives, with failures kept visible.
- Compare greedy trajectories, perplexity, KL and a small task evaluation before interpreting speed.
- Launch the same checkpoint in the book's engine and production engines, with a reproducible workload and an honest results table.
- Bound request bodies, deadlines, output buffers and shutdown; distinguish health from readiness.
You will build
arrival_offsets, request and summarize in engine/bench.py; logit_quality and evaluate_endpoint in engine/evaluation.py.
Time: 5-7 hours plus experiments. GPU: optional for the milestone; meaningful engine comparisons need representative hardware and real weights.
What are we trying to make faster?
Chapter 10 timed kernels. Chapter 19 timed one decoder. A user sends an HTTP request, waits in admission and scheduling queues, waits for prefill, then receives tokens. A kernel improvement matters only if it improves that experience, or lets the service handle more users at the same experience.
Consider two engines. One produces 1,000 tokens/s but stalls some requests for several seconds. Another produces 800 tokens/s and keeps every request below the latency target. Which is better depends on the service’s objective. Report throughput and latency together, at several offered loads, and identify the load at which the target starts failing.
Before timing, write an experiment record: hardware and power limits, engine revision, model/tokenizer revision, precision, quantization, cache dtype and capacity, context limit, parallelism, graph and speculation settings, prefix reuse policy, prompt/output length distributions, sampling settings and arrival process. --metadata attaches that JSON record to every result. Record the full launch command too. A model name and a tokens/s number are not a reproducible benchmark.
One request’s clocks
Use a monotonic clock on the client. Let $a$ be scheduled arrival, $d$ HTTP dispatch, $f$ first visible content, $l$ last visible content, $e$ completed stream, and $n$ the generated token count:
| metric | definition | what it captures |
|---|---|---|
| client queue | $d-a$ | waiting for the load generator’s connection/concurrency cap |
| TTFT | $f-a$ | client queue, transport, engine queue, prefill and first visible output |
| E2E | $e-a$ | everything through stream termination |
| TPOT | $(l-f)/(n-1)$, if $n>1$ | average time per output token after the first |
| ITL | gap between consecutive token arrivals | irregular decode progress |
| chunk gap | gap between consecutive content events | what the streaming client actually observed |
An SSE role header, keepalive or empty finish event is not the first token. The benchmark waits for nonempty content. A Unicode character may need several token bytes before becoming visible, and a speculative step may stream several tokens together. TTFT here is first visible text, which can differ from an engine’s first-token-ID timestamp. Keep that distinction when comparing server metrics with client metrics.
A chunk is not necessarily a token. The client records content-event timestamps and token counts from final usage. It reports token ITL only when logprobs prove that each event contains exactly one token and the total equals usage. Otherwise the ITL sample is empty, while chunk-gap percentiles remain available. It cannot recover true individual token arrival times by dividing a five-token chunk into five imaginary arrivals.
TPOT uses total token count and the first/last content times. If output arrives entirely in one chunk, it is unknown rather than zero. Usage is preferred; an optional tokenizer fallback retokenizes the visible text and labels the count as an estimate. Stop tokens, special tokens and invisible reasoning can make that estimate differ from the engine’s generated-token count.
An open-loop load generator
In a closed loop, a client submits another request when one finishes. A slower engine then sees less offered load, which hides its overload behavior. In an open loop, arrivals are chosen independently of completions:
- Constant arrivals are spaced by $1/\lambda$ seconds.
- Poisson arrivals use exponential gaps with mean $1/\lambda$; bursts arise naturally.
- Explicit bursts send $b$ requests together every $b/\lambda$ seconds.
def arrival_offsets(count, rate, kind="poisson", seed=0, burst_size=8):
"""Open loop: arrivals don't wait for completions. Poisson gaps have mean 1/rate.
Burst groups arrive together, with the same long-run offered rate. (Your engine: Chapter 44)
"""
if count < 1 or rate <= 0 or math.isnan(rate) or burst_size < 1 or kind not in ("poisson", "burst", "constant"):
raise ValueError("Need positive count/rate/burst_size and a known arrival process")
if math.isinf(rate):
return [0.0] * count
rng, at, out = random.Random(seed), 0.0, [0.0]
for i in range(1, count):
if kind == "poisson":
at += rng.expovariate(rate)
elif kind == "constant":
at = i / rate
elif i % burst_size == 0:
at += burst_size / rate
out.append(at)
return out
The first request arrives at zero; rate=inf sends the finite workload together. A seeded random generator makes arrival times reproducible. A semaphore caps simultaneous HTTP connections, but scheduled arrival stays unchanged, so local queueing is included in TTFT and E2E. Inspect client_queue: if it grows, the client cap is part of your result. Raise it, or declare that you are measuring a client-concurrency-limited service.
async def request(client, url, model, workload, trace, clock=time.perf_counter, tokenizer=None, api_key=None):
"""Time actual content events, require [DONE], retain final usage, and record errors.
The logprobs length can prove token/event counts; otherwise ITL is chunk-level. (Your engine: Chapter 44)
"""
payload = {"model": model, "prompt": workload["prompt"], "max_tokens": workload.get("max_tokens", 128),
"temperature": 0, "n": 1, "stream": True, "stream_options": {"include_usage": True}}
headers = {"Authorization": f"Bearer {api_key}"} if api_key else {}
done = False
try:
async with client.stream("POST", url, json=payload, headers=headers) as response:
response.raise_for_status()
async for data in sse_events(response.aiter_lines()):
if data == "[DONE]":
done = True
break
event = json.loads(data)
if event.get("error"):
raise ValueError(f"SSE error: {event['error']}")
usage = event.get("usage")
if usage:
trace.prompt_tokens, trace.output_tokens = usage.get("prompt_tokens"), usage.get("completion_tokens")
trace.token_count_source = "usage"
for choice in event.get("choices", []):
if choice.get("index", 0) != 0:
continue
text = choice.get("text", "")
delta = choice.get("delta") or {}
text += delta.get("content") or delta.get("reasoning_content") or ""
if text:
trace.content_times.append(clock())
trace.text += text
lp = choice.get("logprobs") or {}
tokens = lp.get("tokens", lp.get("content"))
trace.event_token_counts.append(len(tokens) if tokens is not None else None)
if not done:
raise ValueError("Stream ended without [DONE]")
if trace.output_tokens is None and tokenizer is not None:
trace.output_tokens = len(tokenizer.encode(trace.text))
trace.token_count_source = "retokenized_text"
if trace.output_tokens is not None and (type(trace.output_tokens) is not int or trace.output_tokens < 0):
raise ValueError("Invalid completion_tokens")
if not trace.content_times:
raise ValueError("No generated content received")
except asyncio.CancelledError:
raise
except Exception as error:
trace.error = f"{type(error).__name__}: {error}"
finally:
trace.finished = clock()
return trace
httpx.aiter_lines handles fragmented network bytes and UTF-8 decoding. sse_events assembles multiline data fields, ignores comments and recognizes [DONE]. A transport success without [DONE], a server error event, an HTTP 503 or a timeout becomes an explicit failed trace. Warmup failures stop the experiment. Requests have both transport timeouts and a total deadline; a stream that periodically sends bytes cannot run forever.
The JSON output preserves workload, configuration and every request’s timestamps, token counts, text and error. Those traces let you recompute percentiles or change SLOs after a run without sending the workload again. They may contain sensitive prompts and outputs; use a synthetic workload when sharing them.
Throughput and goodput
The measurement window begins at the first scheduled arrival and ends when the final request finishes or fails. Include the final drain, not just the time requests were submitted. Output throughput is successful requests’ generated tokens divided by that window. Request throughput is successful requests divided by it. Keep the number of requests with known token counts beside token throughput; a missing usage field cannot become an invented count.
Goodput counts successful requests that meet every configured SLO. A request meeting TTFT but missing TPOT fails the combined objective. Missing TPOT cannot pass a TPOT SLO; failed requests remain in the attainment denominator.
def summarize(traces, duration, ttft_slo=None, tpot_slo=None, e2e_slo=None):
"""Throughput counts successful work; goodput requires every configured SLO.
Failed requests remain in the denominator; missing token timing cannot pass a TPOT SLO. (Your engine: Chapter 44)
"""
if duration <= 0 or not traces:
raise ValueError("Need traces and a positive measurement window")
good = [t for t in traces if t.error is None]
met = [t for t in good if all(limit is None or value is not None and value <= limit
for value, limit in ((t.ttft, ttft_slo), (t.tpot, tpot_slo), (t.e2e, e2e_slo)))]
gaps, token_gaps = [], []
for t in good:
intervals = [b - a for a, b in zip(t.content_times, t.content_times[1:])]
gaps.extend(intervals)
# Only report token ITL when the stream proves every content event is one token,
# and there are no invisible tokens omitted from the events.
if t.event_token_counts and all(n == 1 for n in t.event_token_counts) \
and t.output_tokens == len(t.content_times):
token_gaps.extend(intervals)
measured = [t for t in good if t.output_tokens is not None]
return {"requests": len(traces), "successful": len(good), "failed": len(traces) - len(good),
"duration_seconds": duration, "request_throughput_per_second": len(good) / duration,
"output_token_throughput_per_second": sum(t.output_tokens for t in measured) / duration if measured else None,
"requests_with_token_counts": len(measured), "goodput_requests_per_second": len(met) / duration,
"slo_attainment": len(met) / len(traces),
"latency_seconds": {"ttft": percentiles([t.ttft for t in good]),
"tpot": percentiles([t.tpot for t in good]),
"itl": percentiles(token_gaps), "chunk_gap": percentiles(gaps),
"e2e": percentiles([t.e2e for t in good]),
"client_queue": percentiles([t.dispatched - t.scheduled for t in traces])}}
Worked arithmetic, not a performance measurement: over 2 seconds, successful requests produce 3 and 5 tokens and a third request fails. Output throughput is 4 tokens/s, request throughput is 1 request/s. If only the first request meets every SLO, goodput is 0.5 requests/s and attainment is 1/3. Reporting only successful requests would hide the failure.
Percentiles use linear interpolation between sorted samples. Run enough requests to support a tail claim: the p99 of ten samples is almost the maximum, with very little statistical evidence. Repeat trials, report trial variation, and sweep arrival rate until SLO attainment falls below the service target. Goodput usually rises, peaks and then falls as queueing overwhelms the engine.
Run a workload
Create prompts.jsonl with fixed text and output limits:
{"prompt":"Explain why a paged KV cache saves memory.","max_tokens":128}
{"prompt":"Write a short example of continuous batching.","max_tokens":128}
For a real run use hundreds or thousands of rows, with representative lengths. Keep the exact file and hash it. Test at least short-prefill/long-decode, long-prefill/short-decode, mixed lengths and deliberately shared prefixes. Distinguish a cold run (fresh process/cache and unrelated prompts) from a warm run (declared shared prefixes and cache state). Warmup also warms prefixes; do not call a workload cold just because compilation is finished.
python -m izh.bench --base-url http://127.0.0.1:8000 --model Qwen3-0.6B \
--workload prompts.jsonl --arrival poisson --rate 4 --seed 7 \
--max-concurrency 256 --warmup 4 --timeout 120 \
--ttft-slo 1.0 --tpot-slo 0.05 --e2e-slo 10 \
--metadata experiment.json --output runs/izh-poisson-4.json
python -m izh.bench --base-url http://127.0.0.1:8000 --model Qwen3-0.6B \
--workload prompts.jsonl --arrival burst --burst-size 16 --rate 4 \
--output runs/izh-burst-4.json
SLOs in those commands are experiment inputs, not claims about achieved latency. The client uses raw /v1/completions, temperature=0, n=1, no chat template, and each row’s max_tokens. EOS may finish early: record actual token counts. If testing fixed output lengths, configure every server’s supported EOS policy consistently and retain that configuration. Do not send an extension one engine ignores and assume it took effect.
Accuracy beside speed
First check greedy equality against an independent generator, starting with identical token IDs. Report the first divergent token. Once one token differs, later histories differ; a full-sequence logit comparison no longer diagnoses the first error. A candidate that returns immediately or uses the wrong tokenizer can look wonderfully fast.
Then compare logits on the same teacher-forced history:
def logit_quality(reference, candidate, ids, mask=None):
"""Logits [B,T,V], IDs [B,T]. Score next tokens, excluding optional padding.
KL is KL(reference || candidate), teacher-forced on the same history. (Your engine: Chapter 44)
"""
if reference.shape != candidate.shape or reference.shape[:2] != ids.shape or ids.shape[1] < 2:
raise ValueError("Quality comparison needs matching logits and at least two tokens")
targets = ids[:, 1:]
valid = torch.ones_like(targets, dtype=torch.bool) if mask is None else (mask[:, 1:] & mask[:, :-1]).bool()
if not valid.any():
raise ValueError("No next-token positions to score")
ref = reference[:, :-1].float().log_softmax(-1)
cand = candidate[:, :-1].float().log_softmax(-1)
rloss = -ref.gather(-1, targets[..., None]).squeeze(-1)[valid].mean()
closs = -cand.gather(-1, targets[..., None]).squeeze(-1)[valid].mean()
kl = (ref.exp() * (ref - cand)).sum(-1)[valid].mean()
return {"tokens": int(valid.sum()), "reference_nll": float(rloss), "candidate_nll": float(closs),
"reference_perplexity": math.exp(float(rloss)), "candidate_perplexity": math.exp(float(closs)),
"kl_reference_candidate": float(kl),
"teacher_forced_argmax_agreement": float((ref.argmax(-1) == cand.argmax(-1))[valid].float().mean()),
"max_abs_logit_difference": float((reference - candidate).abs().max())}
def greedy_equality(reference_generate, candidate_generate, prompts, max_tokens=16):
"""Call two independent {request_id: tokens} generators; report the first divergence."""
ref, cand = reference_generate(prompts, max_tokens), candidate_generate(prompts, max_tokens)
differences = []
for rid in prompts:
a, b = ref[rid], cand[rid]
if a != b:
at = next((i for i, (x, y) in enumerate(zip(a, b)) if x != y), min(len(a), len(b)))
differences.append({"request_id": rid, "position": at,
"reference": a[at] if at < len(a) else None, "candidate": b[at] if at < len(b) else None})
return {"requests": len(prompts), "equal": len(prompts) - len(differences), "differences": differences}
Next-token negative log likelihood excludes the first token and masked padding; perplexity is its exponential. KL(reference || candidate) measures how much the whole distribution changed, even if argmax did not. Use the same tokens, position conventions and mask in both models. Compare FP32 first, then deployment precision and quantization. Exact greedy equality is a strong regression test, but small floating-point differences near a tied argmax can change a trajectory without proving a major quality loss (Chapter 13).
Finally run a small task evaluation. The supplied runner accepts a local JSONL subset with question and answer fields, including GSM8K’s #### numeric answer convention:
async def evaluate_endpoint(client, base_url, model, examples, max_tokens=256):
"""A reproducible greedy, zero-shot numeric-answer smoke evaluation. (Your engine: Chapter 44)
Each row retains the question, expected answer, full response and HTTP failures.
"""
records = []
for i, row in enumerate(examples):
expected = numeric_answer(row["answer"])
if expected is None:
raise ValueError(f"Example {i} has no numeric reference answer")
prompt = f"Question: {row['question']}\nSolve step by step. End with #### followed by the numeric answer.\nAnswer:"
result = {"index": i, "question": row["question"], "answer": row["answer"], "prompt": prompt, "correct": False}
try:
response = await client.post(base_url.rstrip("/") + "/v1/completions",
json={"model": model, "prompt": prompt, "temperature": 0, "max_tokens": max_tokens})
response.raise_for_status()
result["output"] = response.json()["choices"][0]["text"]
result["correct"] = numeric_answer(result["output"]) == expected
except Exception as error:
result["error"] = f"{type(error).__name__}: {error}"
records.append(result)
if not records:
raise ValueError("Evaluation subset is empty")
return {"protocol": "zero-shot numeric extraction, greedy completions; not the official GSM8K protocol",
"model": model, "max_tokens": max_tokens, "examples": len(records),
"accuracy": sum(r["correct"] for r in records) / len(records), "records": records}
python -m izh.evaluation --base-url http://127.0.0.1:8000 --model Qwen3-0.6B \
--data data/gsm8k-subset.jsonl --limit 32 --max-tokens 256 --output runs/izh-eval.json
The reader supplies that file; it isn’t bundled or downloaded. The runner retains prompts, responses and failures, uses greedy zero-shot completions and extracts a numeric answer. This is a smoke protocol, not the official few-shot GSM8K evaluation. A small subset has large uncertainty and the extractor can accept a stray final number. For a published score, use the benchmark’s defined protocol and a reviewed harness. Preserve dataset revision, subset IDs and licenses with your experiment.
Launch the competitors with the same checkpoint
These recipes use one GPU, the same local Qwen3-0.6B checkpoint, FP16 weights and KV, a 4,096-token per-request limit and up to 16 active sequences. Start from one pinned HF revision and convert that directory for GGUF. Use each engine’s own supported environment; mixing all five dependency stacks into one environment is unnecessary. Record installed versions or git SHAs and actual resolved startup settings. CLI recipes were checked against the linked upstream documentation; competitor execution is recorded separately in Appendix F.
Run servers one at a time on port 8000 and query /v1/models for the model name the benchmark must send. Set CUDA_VISIBLE_DEVICES=0 in the server shell. Check prompt IDs/BOS behavior before comparing greedy outputs; identical text does not guarantee identical tokenization.
The book’s engine:
python serve.py --model-dir models/Qwen3-0.6B --served-model-name Qwen3-0.6B \
--device cuda --dtype fp16 --max-model-len 4096 --max-num-seqs 16 \
--max-num-batched-tokens 2048 --attention-backend triton --kv-cache-dtype auto \
--cuda-graphs auto --port 8000
vLLM (serve arguments):
vllm serve models/Qwen3-0.6B --served-model-name Qwen3-0.6B \
--dtype float16 --kv-cache-dtype auto --tensor-parallel-size 1 \
--max-model-len 4096 --max-num-seqs 16 --max-num-batched-tokens 2048 --port 8000
SGLang (server arguments):
python -m sglang.launch_server --model-path models/Qwen3-0.6B \
--served-model-name Qwen3-0.6B --dtype float16 --kv-cache-dtype auto \
--tp-size 1 --context-length 4096 --max-running-requests 16 \
--chunked-prefill-size 2048 --port 8000
llama.cpp (server arguments), from its checkout with a CUDA build:
python convert_hf_to_gguf.py /absolute/path/to/models/Qwen3-0.6B \
--outfile /absolute/path/to/models/Qwen3-0.6B-F16.gguf --outtype f16
build/bin/llama-server --model /absolute/path/to/models/Qwen3-0.6B-F16.gguf \
--alias Qwen3-0.6B --n-gpu-layers 999 --parallel 16 --kv-unified \
--kv-unified-per-slot 4096 --batch-size 2048 --ubatch-size 512 \
--cache-type-k f16 --cache-type-v f16 --port 8000
With these current flags, the unified pool is sized for 16 × 4,096 positions. Older revisions use different context/slot semantics; preserve the resolved per-slot context printed at startup rather than assuming --ctx-size 4096 gives every one of sixteen slots 4,096 tokens.
ExLlamaV3 through TabbyAPI (sample config), from its checkout, with config.yml:
network:
host: 127.0.0.1
port: 8000
disable_auth: true
allowed_origins: []
model:
model_dir: /absolute/path/to/models
model_name: Qwen3-0.6B
backend: exllamav3
max_seq_len: 4096
cache_size: 65536
cache_mode: FP16
max_batch_size: 16
tensor_parallel: false
python main.py
This local benchmark configuration loads the original unquantized checkpoint; ExLlamaV3’s linear implementation has an unquantized path. Its native quantized deployment deserves a separate EXL3 experiment, with bitrate and accuracy recorded, alongside separate GGUF, GPTQ/AWQ and FP8 experiments for the other engines. A 4-bit result is not the same precision experiment as the baseline.
TensorRT-LLM (CLI, LLM arguments), with trtllm-bench.yml:
dtype: float16
max_batch_size: 16
max_num_tokens: 2048
max_seq_len: 4096
tensor_parallel_size: 1
pipeline_parallel_size: 1
trtllm-serve models/Qwen3-0.6B --backend pytorch --config trtllm-bench.yml --port 8000
These match model, precision, context and sampling, but not every internal policy. Graph buckets, prefix caches, total KV allocation, prefill tiling and batch-token limits differ. Measure each engine’s normal optimized serving configuration first and disclose those differences. Then run controlled ablations (prefix cache, graphs, speculation), changing one setting at a time. Restart between cold trials; use deliberately reused prefixes for warm trials. For a memory-normalized comparison, match usable KV bytes explicitly rather than trusting unrelated default utilization fractions.
A results table that earns its numbers
There are no competitor measurements in this table. Fill a row only after saving its trace JSON, startup log, revisions, workload hash and accuracy result. A local random-model smoke run validates transport and metric plumbing; it cannot rank these engines.
| engine / revision | precision / KV | workload / rate | output tok/s | TTFT p50 / p99 | TPOT p50 / p99 | goodput req/s | failures | greedy / KL / task score |
|---|---|---|---|---|---|---|---|---|
| izh | FP16 / FP16 | not run | — | — | — | — | — | — |
| vLLM | FP16 / FP16 | not run | — | — | — | — | — | — |
| SGLang | FP16 / FP16 | not run | — | — | — | — | — | — |
| llama.cpp | F16 GGUF / F16 | not run | — | — | — | — | — | — |
| ExLlamaV3 / TabbyAPI | unquantized / FP16 | not run | — | — | — | — | — | — |
| TensorRT-LLM | FP16 / FP16 | not run | — | — | — | — | — | — |
Keep per-token ITL and chunk-gap percentiles beside this table when studying streaming stalls. Record CPU usage and GPU memory alongside throughput: moving a bottleneck into a tokenizer process or exhausting host memory is still a system bottleneck.
Harden what you have measured
The service now has a raw-body limit before JSON parsing, a body-read deadline, a total generation deadline, bounded pending requests and a bounded number of buffered stream updates. If a consumer falls behind, its request is aborted and its buffer becomes an error, releasing its KV and adapter lease. It does not let one stalled connection retain memory indefinitely.
AsyncLLM.stop marks the service draining, refuses new generation, gives existing consumers 30 seconds to finish, then aborts remaining requests, shuts down and joins the core worker. A process still alive after the join deadline is terminated; a Python thread cannot be forcibly stopped safely. /health reports fatal engine failure; /ready additionally requires loaded model metadata and refuses readiness while draining. The launcher emits structured request metadata (ID, path, status and elapsed time) without logging prompts. Stream failures after HTTP 200 still need stream/error counters; the HTTP status alone isn’t sufficient.
Deployment checklist:
| failure to exercise | behavior to verify | supplied mechanism / remaining integration |
|---|---|---|
| slow or oversized body | reject before expensive parsing/allocation | RequestLimits: byte cap and body deadline, including chunked bodies |
| overloaded admission | bounded queues; clear retryable error | max_pending, 503, scheduler context/KV limits; add per-tenant quotas |
| long or stuck generation | cancel at a total deadline | request_timeout, finally abort; test stream error delivery through your proxy |
| stalled/disconnected consumer | release KV and adapter leases | update cap, cancellation, core abort; test proxy disconnect propagation |
| deployment shutdown | stop admission, drain, then bounded abort/join | AsyncLLM.stop; align orchestrator/proxy grace periods and stop routing before teardown |
| failed model load / core crash | readiness fails; all affected clients fail | ready/fatal messages, /health, /ready; exercise process death and restart |
| malformed request or unknown adapter | explicit client error | schema/parameter validation; monitor 4xx and avoid retry storms |
| incident investigation | correlate request/engine events without text leakage | request IDs, structured metadata, Prometheus counters; connect to your log collector |
| checkpoint or adapter replacement | no stale KV or mid-request weight change | drain base-model updates; adapter leases and content keys |
Rate limiting, identity, rolling-deployment routing, persistent observability, host-level limits and recovery from device failure belong to the deployment around the engine. Test them with faults, not just a successful curl. Chapter 45 distinguishes the implemented mechanisms from the work still needed for an operational service.
Build it
Engine milestone 44: measure the service. Implement arrival_offsets, request and summarize in engine/bench.py, and logit_quality and evaluate_endpoint in engine/evaluation.py. The CLI entry points, trace properties, SSE parser and greedy-equality helper are provided.
pytest tests/test_ch44_benchmarking.py
python run.py benchmark
python -m izh.bench --help
python -m izh.evaluation --help
The tests use fixed clocks for denominators/SLOs, fragmented UTF-8 SSE streams for parsing, explicit HTTP and truncated-stream failures, seeded arrival sequences, independent cross-entropy calculations, first-divergence checks, a task fixture with a failed request, body limits and slow-consumer cancellation. Their timing fixtures are not benchmark measurements.
Stretch exercises
- ★★ Sweep offered rate, plot p99 TTFT and goodput, and repeat three trials at each point. Find the knee where waiting-queue growth begins. Where: terminal: the benchmark CLI recipes above; save trial results and plots in
experiments/. - ★★ Add separate cold/warm prefix workloads and compare equal memory budgets. How much of an observed gain comes from cache hits rather than faster kernels? Where: workload construction in
engine/bench.pyorexperiments/ch44.py(create it) using itsbenchmarkfunction. - ★★★ Add a chat or multimodal workload with the same rendered prompt/processor in every engine. Record image-encoder time separately from decoder prefill. Where: request construction in
engine/bench.py, with chat/image preparation shared byexperiments/ch44.py(create it). - ★★ Inject core-process failure and proxy disconnects during a burst. Verify bounded client failure time, memory release and readiness transitions. Where: new
experiments/service_failures.py, driving the server as a subprocess and an interruptible proxy. - ★★★ Compare INT4 and FP8 variants with teacher-forced KL and a full task protocol. Report a speed-quality-memory curve rather than one winning number. Where:
experiments/ch44.py(create it), usingengine.evaluation.logit_qualityandevaluate_endpointalongsideengine.bench.benchmark.
Check your understanding
- Why can a closed-loop client conceal overload?
- Which timestamp changes when the client waits for a connection, and which should remain fixed?
- Why can’t a five-token SSE chunk supply five real ITL measurements?
- Why is TPOT undefined for a one-token response?
- Why must failed requests remain in SLO attainment’s denominator?
- Why compare logits on a teacher-forced history after greedy outputs diverge?
- Why are a 4-bit model and an FP16 model different benchmark experiments?
- What should readiness report while a process is healthy but draining?
Going deeper
- The vLLM serving benchmark and SGLang benchmark guide, for workload generation and metric conventions.
- GSM8K and EleutherAI’s evaluation harness, for datasets, task definitions and reproducible evaluation protocols.
- Chapters 10, 13, 24 and 31: rooflines, numerical differences, continuous batching and token-budget scheduling explain the curves this client measures.
45. Where to go next
In this chapter
- Identify what the completed engine implements and where its boundaries remain.
- Use a repeatable bring-up method for new models, hardware and feature combinations.
- Choose a concrete next project using correctness evidence and measured bottlenecks.
Time: as long as you like.
What you’ve built
You began with tensors, an autograd engine and a tokenizer. You trained a GPT, loaded real model formats, wrote CUDA and Triton kernels, built cached decoding, changed models with fine-tuning and LoRA, and matched a frontier hybrid architecture to its reference. Part VIII then put serving mechanisms together: flattened mixed batches, paging and prefix caching, preemption, graph buckets, batched sampling, text boundaries, an asynchronous HTTP frontend, batched speculation, quantized checkpoint formats, offloading, distributed execution, model registration, latent attention, image inputs, pooling and many adapters. Chapter 44 added the client that measures the resulting service.
The most valuable result is the method: derive, implement the simple version, compare with an independent reference, optimize while keeping the comparison green, then measure the system. That method transfers to a model or device this book has never seen.
There is still a difference between implementing a mechanism and operating a dependable inference service. Tests on tiny random models establish useful invariants. They do not establish real-checkpoint coverage, tail latency under sustained load or recovery after a GPU fails. Appendix F records the evidence this book actually has.
What remains beyond the book
The earlier final chapter left distributed serving as a next project. Chapter 41 now implements its mechanics; Chapter 42 brings up more model families; Chapter 43 adds non-text boundaries. The next work is deeper coverage and integration:
| area | built here | next level |
|---|---|---|
| attention and matmul | paged unified/split-KV kernels, affine quantized matmul, FP8 weight reads | architecture-specific TMA/WGMMA pipelines, optimized MLA prefill/decode, tensor-core INT4/FP8/FP4 kernels and autotuning |
| formats | GGUF/selected GGML quants, GPTQ/AWQ layouts, block FP8, MX/NV FP4 references, a trellis lab | broad versioned interoperability, every GGML/i-quant variant and native EXL2/EXL3 loading; checkpoint-specific quality validation |
| hardware | CPU references and selected SIMD paths, CUDA/Triton, a Metal lab and platform detection | complete optimized Metal/Vulkan/ROCm backends, kernel tuning per device, reliable heterogeneous deployment |
| many GPUs | TP/PP/EP, ring attention, routing and KV transfer reference paths | overlapping communication, topology-aware collectives, RDMA KV transport, resharding and distributed failure recovery |
| models | dense registry, MoE and MLA, a separate hybrid capstone, a tiny vision path | trained multimodal checkpoint loaders, audio/video processors, native batched hybrid-state scheduling and broader model contracts |
| LoRA and pooling | named resident slots, gathered kernels, bounded embedding/reranking forwards | adapter paging, quantized/distributed adapters, graph feature buffers and cross-request embedding batching |
| service | cancellation, limits, deadlines, readiness, drain, metrics and load generator | tenant identity/quotas, rollout routing, durable telemetry, availability targets and fault-tested operations |
| combinations | shared paged scheduler with tested feature paths | a measured compatibility matrix for graphs, speculation, guides, quantization, LoRA, images, offload and parallelism together |
That last row matters. An adapter request currently takes an eager path, and a feature request with a drafter is refused. The capstone’s recurrent state is not automatically paged like dense K/V. A reference matmul reading quantized weights is not a fused low-bit tensor-core implementation. Read the explicit limitations in each chapter; don’t advertise a combination just because each feature exists in isolation.
Bring up a new architecture
Use the checklist that brought up the book’s model families:
- Read the config and reference code. Build Chapter 30’s ledger: equations, names, shapes, dtypes, masks, positions, state and trained input/output conventions.
- List behavior changes. A different vocabulary size is a parameter; a different norm, rotary layout, router or image processor is a new computation.
- Reject unsupported options. Never let a config field that changes behavior quietly fall through to a default. Add a test for the refusal too.
- Test components independently. Randomize all weights, including norms and gates that initialize to zero or one; use non-default dimensions and odd lengths.
- Save and reload a tiny reference checkpoint. Compare FP32 hidden states/logits at each layer, then at deployment precision. Check every tensor was consumed or explicitly excluded.
- Test state equivalence. Full forward, chunked prefill, decode, prefix reuse and preemption should describe the same history. For hybrid states, test snapshots and restoration.
- Test feature combinations. Mix request lengths and adapters; interrupt an image chunk; force memory pressure; cancel a request with shared prefixes; change batch membership. Pick invariants that catch errors rather than repeating implementation steps.
- Plan memory and load real weights. Account for activations, graph buffers, scale tensors, adapter slots and host copies as well as weights and KV. Compare real greedy outputs with a trusted implementation.
- Measure the service. Use Chapter 44’s workloads, a declared accuracy protocol and the hardware ceiling. Optimize the bottleneck you measured, then repeat the correctness check and experiment.
Read production engines as a request journey
Trace one request from HTTP parsing through scheduling, cache allocation, the model’s forward, attention dispatch, sampling and output. You now have names for each boundary:
| engine | useful entry point | book counterpart |
|---|---|---|
| vLLM | V1 scheduler, KV manager, model runner and attention backends | Chapters 31-34 and 42 |
| SGLang | tokenizer manager, scheduler and radix cache | Chapters 31, 35-37 |
| llama.cpp | model loader, GGML graph and backend kernels | Chapters 38-40 |
| ExLlamaV3 | loader, linear modules, cache and generator | Chapters 25, 37, 39 and 44 |
| TensorRT-LLM | LLM API, executor and kernel dispatch | Chapters 32-33, 39 and 41 |
Start from a pinned revision, since module paths move. Chapter 44 links the projects’ current launch documentation. Find where the real engine handles a corner case your reference refuses: adapter rank changes, graph fallback, a cache eviction during speculation, or a failed KV transfer. That is a good first contribution because you can state the invariant and write a focused regression test.
Open questions worth experimenting with
- Long context with less state. Compare sparse retrieval, recurrent/hybrid state and selective KV retention on recall tasks, not just cache size. Which information is lost, and at what context length?
- Low-bit everything. Study weight, activation and KV quantization separately before combining them. Add quality, scale-storage bytes and kernel traffic to the performance curve.
- Speculation for complex state. Tree verification, MTP and hybrid-state rollback need ownership and commit rules as much as faster draft models. Can an adapter-aware drafter stay useful across many adapters?
- Batch-invariant inference. Results should be reproducible under a declared contract when neighbouring requests change. Measure the throughput cost of deterministic reductions and routing.
- Scheduling heterogeneous work. An image encoder, an embedding batch, long prefill and a decode step have different resource footprints. Design admission and scheduling that protect decode SLOs without starving the other work.
- Kernel specialization with a fallback. Optimize one costly architecture/device pair while keeping a simple reference that verifies every supported shape and supplies correct results elsewhere.
Follow GPU Mode, model technical reports and the systems projects you use. Read release code with a specific question, then reproduce one result. A small measured experiment is more useful than collecting every new acronym.
A last exercise
Pick one concrete boundary in the table above. Bring up a trained VLM, capture adapter batches in graphs, move KV over an actual RDMA link, or optimize a quantized matmul on your device. Write the contract first. Match a trusted reference, exercise failure and state transitions, save a reproducible benchmark with accuracy beside speed, and publish the result with its limitations.
Where: create experiments/final_project/ in the code root for the contract, scripts and
results, with a README.md for the write-up. For the implementation, start with the targets in
Chapter 43 for a VLM or adapter graphs,
Chapter 41 for KV transfer, or
Chapter 39 for a quantized matmul.
That write-up is the proof that you can continue without a book telling you the next step.
A. The engine, tests and commands
Everything in this appendix is in the code download (the code/ folder of the book’s repository). Chapter 0 sets it up.
Layout
engine/ YOUR engine: every module, class and signature, with the chapter's functions left as TODO
izh/ the complete reference engine (the answer key): same modules, same names
tests/ one milestone test file per chapter, run against engine/ by default
run.py one demo command per chapter; --impl engine runs yours
serve.py OpenAI-compatible server launcher; see Chapter 36
izh/serve/ unified scheduler, block manager, runner, API and adapter slots
izh/bench.py, izh/evaluation.py
endpoint load generator and correctness/quality harnesses (Chapter 44)
data/ The Verdict, BALLM's 1,100 instructions (with fixed splits), steering prompts
rust/ the Rust track: CPU engine for Qwen3, tests against the Python reference
cpp/ the C++ track: header-only companions and standalone CUDA kernels
rust-cuda/ optional NVIDIA Rust GPU examples
finetune.py, quantize_checkpoint.py, edit_checkpoint.py, model_workflows.py
Hugging Face-based workflows for real checkpoints (Chapters 20-23)
engine/ is generated from izh/ by tools/make_engine.py: every function whose docstring says “(Your engine: Chapter N)” keeps its signature and docstring and has its body replaced with raise NotImplementedError("TODO(Chapter N): ..."). Triton kernels marked that way become pass, and the CUDA kernels in cuda_ops.cu lose their bodies. Everything not marked is provided, so you write the ideas and skip the boilerplate.
Warning
Regenerating overwrites
engine/. The book’s repository does this to keep the skeleton in sync with the reference; you never need to. If you do, commit your work first.
Where to make exercise changes
All exercise paths are relative to the extracted code download’s root: the directory containing
run.py, engine/ and tests/. In this repository that directory is
docs/inference-zero-to-hero/src/code/.
- Engine milestones: fill in the named TODOs in
engine/.izh/is the complete reference used for comparison. A code tab that includescode/izh/...shows that reference; the matchingengine/...file is your implementation target. - Stretch exercises: each Where note names the edit target. These extensions can add functions, classes or arguments beyond the milestone’s existing TODOs. A file marked new or create it is a suggested file to create, not a supplied starter.
- Experiments: create
experiments/chNN.pyfor a chapter’s measurements, plots or standalone comparisons. Import your completedenginemodules there. From the code root, run, for example,python -m experiments.ch02; this keeps the code root on Python’s import path. Use a notebook in the same code root if you prefer. Save handwritten predictions and observations alongside it. - Tests: run the chapter’s milestone suite after engine changes. Add extension checks in a
new
tests/test_chNN_stretch.py, importing the relevantenginemodules. Milestone tests cover the required implementation; they do not automatically validate an optional extension. - Native tracks: edit the explicitly named
cpp/orrust/src/file. For a new CUDA operation inengine/kernels/cuda_ops.cu, add its host launcher andPYBIND11_MODULEentry in that same file;engine/kernels/cuda.pyloads the extension. Existingcpp/kernels.cuis the standalone companion.
The Check your understanding questions are written answers; they require no engine edits. Appendix C contains answers and stretch-exercise hints.
The chapter loop
pytest tests/test_ch16_kv_cache.py # red: NotImplementedError("TODO(Chapter 16): ...")
$EDITOR engine/kv_cache.py # implement the TODOs
pytest tests/test_ch16_kv_cache.py # green
python run.py cache --impl engine # watch your code run
IZH_IMPL=izh pytest tests/test_ch16_kv_cache.py # the reference passes the same tests
Later chapters build on earlier ones: Chapter 17’s tests use your Chapter 5 attention and Chapter 16 cache. If a late test fails in an early function, fix the early one; its own tests may not have covered the case.
Environment variables:
| variable | effect |
|---|---|
IZH_IMPL=izh | tests use the reference instead of engine/ |
IZH_DEVICE=cpu | tests run on the CPU even if a GPU exists |
TRITON_INTERPRET=1 | Triton kernels run in the interpreter (set automatically without a GPU) |
Test markers: gpu (skipped without CUDA), reference (needs transformers; compares with the official implementation), slow. Run pytest -m "not slow" for a quick pass.
Milestones by chapter
| ch. | file | implement |
|---|---|---|
| 2 | tensors.py | contiguous_strides, element_offset, broadcast_shapes, matmul_loops, linear |
| 3 | autograd.py | Value.__add__, __mul__, __pow__, exp, log, relu, tanh, backward |
| 4 | tokenizer.py, data.py | pair_counts, merge_pair, BPETokenizer.train, _encode_chunk; windows |
| 5 | attention.py | split_heads, merge_heads, causal_attention |
| 6 | gpt.py | GPT.__init__, GPT.forward (and the attention and block classes) |
| 7 | train.py | lm_loss, evaluate, train |
| 8 | sampling.py | sample, generate_stream |
| 9 | safetensors_io.py, loaders.py | read_header, load_file; load_gpt2 |
| 10 | measure.py | measure_wall, matmul_intensity, attainable_flops, decode_ceiling |
| 11-13 | kernels/cuda_ops.cu | add_kernel, naive_matmul_kernel, tiled_matmul_kernel, warp_sum, block_sum, row_sum_kernel, rmsnorm_kernel |
| 12 | tiling.py | tiled_matmul |
| 13 | numerics.py | stable_softmax |
| 14 | kernels/triton_basics.py, triton_matmul.py | add_kernel, softmax_kernel, add_rmsnorm_kernel, matmul_kernel |
| 15 | attention.py, kernels/triton_flash.py | online_attention; flash_fwd_kernel |
| 16 | kv_cache.py | KVCache.update, KVCache.truncate |
| 17 | qwen3.py | rope_cos_sin, apply_rope, RMSNorm.forward and the model’s forward methods |
| 18 | loaders.py, engine.py | load_qwen3; LLM.stream |
| 19 | fast.py, kv_cache.py | sample_on_device, FastDecoder._step, generate; StaticKVCache.update |
| 20 | quant.py, kernels/triton_quant.py | quantize_int8_rows, quantize_groupwise, QuantLinear.dequantized_weight, forward; w4a16_kernel |
| 21 | lora.py, sft.py | instruction_batch; last_token_logits |
| 22 | lora.py | LoRALinear.forward, LoRALinear.merged |
| 23 | steering.py | direction_from_means, project_out |
| 24 | scheduler.py | Scheduler.plan, ContinuousBatchingEngine.step |
| 25 | paged.py, kernels/triton_paged.py | BlockAllocator.allocate, release; PagedKVCache.reserve, update, fork; paged_decode_kernel |
| 26 | speculative.py | accept_or_correct, speculative_generate |
| 27 | moe.py | TopKRouter.forward, Experts.forward_grouped |
| 28 | gdn.py | recurrent_gated_delta_rule, chunk_gated_delta_rule, causal_conv1d, GatedDeltaNet.forward |
| 29 | sparse.py | QSAIndexer.forward, GatedAttention.forward |
| 30 | flashnext.py | GatedResidual.read, NGramEmbedding.shifted, forward, PLELayer.forward, FlashNextLayer.forward, FlashNext.step, load_flashnext |
| 31 | serve/batch.py, blocks.py, scheduler.py, model.py, engine.py | ragged batches, hash chains, paging, token-budget scheduling and the unified step loop |
| 32 | kernels/triton_unified.py, serve/triton_backend.py | unified/split-KV attention, merging, quantized KV reads and writes |
| 33 | serve/model.py, graphs.py, engine.py, kernels/triton_fused.py | projection/MoE fusion, graph buckets and overlapping scheduling |
| 34 | serve/sampler.py, beam.py, structured.py | batched sampling, beam search, regex/schema constraints |
| 35 | serve/tokenizer.py, chat.py | tokenizer loading, incremental UTF-8, stops and chat/tool boundaries |
| 36 | serve/async_engine.py, api.py | core loop, async dispatch, generation and request preparation |
| 37 | serve/spec.py | speculation within the paged engine loop |
| 38 | formats/gguf.py, ggml_quants.py, runtime.py, kernels/triton_formats.py | GGUF I/O, GGML decoding/repacking and affine matmul |
| 39 | formats/gptq.py, awq.py, fp8.py, mx.py, trellis.py | quantized checkpoint layouts, quantization references and trellis decoding |
| 40 | platform.py, offload.py, kernels/cpu.py, metal.py | platform policies, CPU kernels and weight/layer offload |
| 41 | parallel.py | tensor/pipeline/expert parallelism, routing, KV transfer and ring attention |
| 42 | models.py | configurable decoders, scaled RoPE, sliding windows, MLA and grouped routing |
| 43 | multimodal.py, pooling.py, serve/adapters.py, runner.py, kernels/triton_lora.py | image rows/M-RoPE, pooling, adapter slots and gathered shrink/expand kernels |
| 44 | bench.py, evaluation.py | arrivals, stream clocks, goodput, teacher-forced quality and a task runner |
rg "TODO\(Chapter" engine/ lists what’s left. Chapters 0, 1 and 45 have no implementation milestone.
run.py commands
Every command runs on a CPU unless noted; --impl engine runs your code; python run.py <command> --help lists the options.
| command | chapter | what it shows |
|---|---|---|
tensors | 2 | shapes, strides, views and a gradient |
autograd | 3 | a tiny MLP trained with your scalar autograd |
bpe | 4 | a byte-level BPE trained on The Verdict |
attention | 5 | a causal attention weight matrix |
gpt | 6 | a GPT and its parameter count (and GPT-2 small’s) |
train | 7 | training on The Verdict, with train and validation loss |
generate | 8 | the trained model under several decoding policies |
gpt2 | 9 | real GPT-2 weights through your GPT (--model-dir) |
profile | 10 | wrong and right GPU timing, and a profiler trace |
kernels | 11-15 | custom CUDA and Triton kernels against PyTorch |
cache | 16 | cached against uncached generation time |
chat | 18, 30 | a real checkpoint through your engine (--model-dir; Flash-Next: --expert-bits 4 --ngram-mmap) |
fast | 19 | the plain loop against the static-buffer decoder (graph mode on CUDA) |
quant | 20 | quantizing a model’s linear layers: size and logit error |
classify | 21 | a GPT turned into a classifier (--data for SMS spam) |
sft | 21 | instruction tuning (GPT-2 with --model-dir, or a GPT from scratch) |
lora | 22 | LoRA on the sft model, then merging |
steer | 23 | the logit lens and a steering direction |
batch | 24 | continuous batching against one request at a time |
paged | 25 | paged memory, prefix sharing and copy-on-write forks |
speculate | 26 | speculative decoding with perfect and early-exit drafts |
moe | 27 | router load, dispatch strategies, parameters of real MoEs |
linear | 28 | associative recall, chunked against recurrent |
sparse | 29 | top-k error and the indexer’s savings |
flashnext | 30 | the real model’s memory plan and a tiny instance |
core | 31 | mixed ragged scheduling, prefix reuse and memory pressure |
backends | 32 | attention backends and split-KV comparisons |
overhead | 33 | projection fusion, graph buffers and async scheduling |
guided | 34 | batched sampling, constrained JSON and beam search |
text | 35 | tokenizer parity, UTF-8 streaming, stops and chat parsing |
serve | 36 | local random-model HTTP/SSE demo and client session |
drafters | 37 | n-gram, draft-model and Medusa serving |
gguf | 38 | GGUF export/load and quantized serving |
quantformats | 39 | checkpoint layouts and quality comparisons |
cpu, offload | 40 | quantized CPU mat-vec and expert-cache traffic |
dist | 41 | distributed greedy parity and communication traffic |
models | 42 | real configurations’ KV-memory arithmetic |
features | 43 | a tiny vision prompt, mixed adapter requests and pooling |
benchmark | 44 | worked SLO arithmetic and independent greedy parity; no performance claim |
Server and measurement commands
Install serve-requirements.txt for HTTP, tokenizer and benchmark labs; Pillow is optional for image data URIs.
python serve.py --model-dir models/Qwen3-0.6B --device cpu --dtype fp32
python serve.py --help
python -m izh.bench --help
python -m izh.evaluation --help
Chapter 36 documents generation/API settings, Chapter 43 the preloaded --lora NAME=DIR adapters and embedding/reranking hooks, and Chapter 44 the complete workload/launch recipes and service limits. engine has the same module CLIs after its TODOs are implemented.
Workflows for real checkpoints
| script | chapter | does |
|---|---|---|
quantize_checkpoint.py | 20 | bitsandbytes NF4 or LLM.int8() against BF16: response loss, KL, bytes, reload check |
finetune.py | 21-22 | full, LoRA or QLoRA fine-tuning of Qwen3 with response-only loss, manifest, reload |
edit_checkpoint.py | 23 | fit and ablate a residual direction; loss and KL against the original |
model_workflows.py | 20, 22, 23 | small self-contained labs: quantization trade-offs, LoRA, activation editing |
They use Hugging Face Transformers for the model so that the chapter can focus on data, measurement and evaluation; your own engine implements the same ideas from scratch elsewhere in the chapter.
B. The C++ and Rust tracks
Python is the book’s main language because it’s where models, kernels (Triton) and tools meet. But an inference engine is systems software, and seeing each mechanism without PyTorch underneath is the best test that you understand it. Throughout the book, the tabs next to the Python code show the same mechanism in C++ and Rust. This appendix explains what each track contains and how to build it.
What each track covers
| chapter | mechanism | C++ (cpp/izh.hpp) | Rust (rust/src/) |
|---|---|---|---|
| 2 | strided tensors, views, matmul | strided, matmul | tensor.rs |
| 3 | scalar autograd | value | autograd.rs |
| 4 | byte-level BPE | bpe | bpe.rs |
| 5, 15 | attention; online softmax | attention, online | attention.rs |
| 6, 17 | LayerNorm, RMSNorm, GELU, SiLU, RoPE | norms, activations, rope | ops.rs |
| 8 | sampling (temperature, top-k, top-p) | sample | sampling.rs, rng.rs |
| 9 | safetensors, BF16 | safetensors | safetensors.rs |
| 11-13 | CUDA kernels | kernels.cu | rust-cuda/ (optional) |
| 16 | KV cache | cache | kv_cache.rs |
| 17-18 | a complete Qwen3 engine on the CPU | qwen3.rs, main.rs | |
| 20 | groupwise INT4 and packing | quant | quant.rs |
| 25 | block tables | paged | paged.rs |
| 26 | speculative acceptance | accept | speculative.rs |
| 28 | delta-rule step | delta | delta.rs |
| 40 | quantized CPU mat-vec and SIMD dispatch | cpp/izh_cpu.hpp: activation quantization, Q4/Q8 rows, runtime ISA selection | simd.rs: Q4 GEMV and x86 AVX2/scalar dispatch |
The C++ track has small, readable companions in izh.hpp and Chapter 40’s quantized kernels in izh_cpu.hpp, each exercised by a self-checking demo. The Rust track grows into a real engine: from Chapter 18 on, it loads Qwen3-0.6B (or any dense Qwen3) from the downloaded safetensors and generates text on the CPU.
C++
Requirements: a C++17 compiler and CMake 3.18 or newer.
cmake -S cpp -B build/cpp -DCMAKE_BUILD_TYPE=Release
cmake --build build/cpp -j
build/cpp/izh all # every chapter's demo
build/cpp/izh 15 # just Chapter 15's
Each demo checks its results and prints one line, for example:
ch15 online softmax equals dense attention for every tile size ok
ch20 int4 group quantization, packing, error bound ok
ch26 rejection correction reproduces p ok
cpp/kernels.cu is a standalone CUDA program, with no PyTorch, for Chapter 11’s first kernels (vector add and naive matmul), checked against a CPU computation. Build it with the CUDA Toolkit:
cmake -S cpp -B build/cpp -DBUILD_CUDA=ON && cmake --build build/cpp -j
build/cpp/izh_cuda
The full set of Chapters 11-13 (vector add, naive and tiled matmul, warp and block reductions, row sums, RMSNorm) lives in engine/kernels/cuda_ops.cu, compiled as a PyTorch extension by engine/kernels/cuda.py with torch.utils.cpp_extension.load, so the tests can check every kernel against PyTorch.
Rust
Requirements: stable Rust (from rustup.rs). The crate’s only dependency is serde_json, for safetensors headers and config.json.
cd rust
cargo test --release
The tests are of two kinds. Unit tests in each module check the mechanisms (the same worked examples as the book). tests/fixtures.rs loads two tiny Qwen3 checkpoints, one in F32 and one in BF16, written by the Python reference (python rust/make_fixtures.py regenerates them), and checks that the Rust engine reproduces the reference’s logits at every position to $10^{-4}$ and its greedy tokens exactly.
Running Qwen3 on the CPU
The engine works on token IDs. A small Python helper handles the text boundary with the checkpoint’s own tokenizer:
cd rust
IDS=$(python tokenize_ids.py encode --model-dir ../models/Qwen3-0.6B "Why is the sky blue?") # chat template; --raw for plain text
OUT=$(cargo run --release -- generate --model-dir ../models/Qwen3-0.6B --ids "$IDS" --new-tokens 64)
python tokenize_ids.py decode --model-dir ../models/Qwen3-0.6B "$OUT"
cargo run --release -- bench --model-dir ../models/Qwen3-0.6B # tokens/s
How it’s built, in qwen3.rs:
- Weights stay in BF16 in memory, exactly as on disk. The mat-vec widens each BF16 value to F32 inside the dot product, so memory traffic is 2 bytes per weight: the decode ceiling of Chapter 10 applies directly. Compare your measured tokens/s with your machine’s memory bandwidth divided by 1.19 GB.
- The mat-vec splits output rows across threads with
std::thread::scope. Decode is memory-bound, so a few threads saturate the memory bandwidth; more don’t help. - The KV cache is a preallocated
Vec<f32>per layer (kv_cache.rs), as in Chapter 16.
Chapter 40 adds simd.rs, with scalar and x86 AVX2 Q4 dots, checked against the scalar result and an error bound against floating-point weights. cpp/izh_cpu.hpp additionally supplies Q4/Q8 paths with AVX2, AVX-512 VNNI and compile-enabled ARM dot-product support; the Python CPU extension uses that layout. ISA coverage depends on the machine and build flags (Appendix F). Further work: wire quantized checkpoint storage into the native Qwen3 loader, add broader SIMD coverage and use the tokenizers crate to remove the Python helper.
Part VIII’s scheduler, HTTP server, model registry and distributed orchestration live in Python. These native tracks teach the underlying CPU mechanisms; they are not alternate implementations of all 44 milestones.
Optional: NVIDIA’s Rust GPU toolchains
Two experimental NVIDIA projects compile Rust for GPUs. They change quickly, so the book pins exact revisions and keeps the examples small:
- cuda-oxide (
rust-cuda/oxide/): Rust kernels in the SIMT model, one thread’s view like CUDA C++. Chapter 11 shows its vector kernel next to the CUDA C++ and Triton versions.install_example.pycopies the example into a fresh checkout of the pinnedNVIDIA/cuda-rustrevision; thencargo oxide run izh_vecaddbuilds and runs it. - cuTile Rust (
rust-cuda/cutile/): a tile-based model closer to Triton, where each program owns a tile of the output.cargo run --releasein that folder builds against the pinnedNVlabs/cutile-rsrevision.
Both need a recent NVIDIA driver and CUDA Toolkit, and neither is required for any milestone.
C. Answers to check questions and stretch exercises
Try each question or exercise before reading its answer. Each chapter below keeps the original question and exercise numbering. Check answers explain the concepts; stretch answers give worked results or a solution approach and verification criteria. Measurements and plots depend on your checkpoint, data, seed and hardware: the guidance below describes what to measure and what to expect, rather than claiming unrun experimental results. If an answer surprises you, reread the corresponding chapter section.
1. The big picture
Check your understanding
- Each token is chosen from logits that depend on all previous tokens, including the ones just generated. Token 2 can’t be computed until token 1 exists, so 40 tokens need 40 dependent model calls (speculative decoding, Chapter 26, bends this rule by guessing).
- Both read the weights once, but prefill does 1,000 tokens’ worth of arithmetic per weight read, so it’s compute-bound and uses the tensor cores fully. One decode step does one token’s arithmetic per weight read and is memory-bound. The 1,000× extra arithmetic costs far less than 1,000× the time.
- 7B × 2 bytes = 14 GB per token; 1 TB/s ÷ 14 GB ≈ 71 tokens/s at most. At 4 bits (3.5 GB plus scales) the ceiling rises to about 270-285 tokens/s.
Stretch exercises
- Sampling at temperature 1.2 gives varied continuations and can choose less likely tokens; greedy decoding repeats the same choices for fixed logits. The prompt, weights, vocabulary and autoregressive dependency stay the same. Three sampled runs can still coincide, especially with a fixed seed.
- With warmed-up, synchronized runs and EOS disabled for this measurement, let the timings be $t_1$ and $t_{101}$. TTFT is approximately $t_1$, and mean decode time is $(t_{101}-t_1)/100$. Prefill usually dominates short answers; enough decode steps dominate long ones.
- FP32 doubles weight bytes relative to BF16, so an otherwise identical bandwidth-bound decode should take about twice as long. Actual slowdown depends on kernels, bandwidth and launch overhead; prefill also depends on the hardware’s FP32 compute path.
2. Tensors
Check your understanding
- Matmul contracts the last axis of the left operand (4) with the first axis of the right (4) and broadcasts the leading batch axis (2); the remaining axes 3 and 5 form the output.
- No: contiguous strides for
[3, 2]are(2, 1). Strides(1, 3)are what you get from transposing a contiguous[2, 3]tensor (.T), which swaps shape and strides without moving data. - A square matrix and its transpose have the same shape, so every shape check passes while the computation is wrong. Only a value check against a reference catches it.
- 28 × 2 × 8 × 4,096 × 128 × 2 bytes = 469,762,048 bytes = 448 MiB.
Stretch exercises
- The shapes are
[8, 7, 6, 5],[3, 3]and[4, 4, 5], respectively. Align broadcast dimensions from the right; the final expression forms every pairwise difference between the four rows. - Transpose swaps entries
aandbin both shape and strides, leaving the offset unchanged. For a valid nonnegativestartand positivestep, slicing sets the offset tooffset + start * strides[axis], the stride tostep * strides[axis], and the length tomax(0, ceil((shape[axis] - start)/step)). Normalize slice bounds first; PyTorch does not support negative-step tensor slices. - Measure the ratio of median loop time to median
@time with the same dtype and thread count. Python dispatch and scalar tensor operations dominate the loop, while@runs optimized compiled code; there is no machine-independent slowdown factor. - Broadcast the leading shapes, allocate
[..., M, N], and loop over each output batch index and(m, n, k). Map an output batch coordinate to 0 whenever an operand’s corresponding batch dimension is 1, including implicit leading dimensions. Compare values with@on unequal batch ranks and singleton dimensions.
3. How networks learn
Check your understanding
- The gradient points in the direction of steepest increase of the loss; stepping against it decreases the loss fastest for a small step.
- A value used in several places receives gradient contributions from each use (the multivariable chain rule sums them). Assigning with
=would keep only the last contribution. - Without nonlinearities, a stack of linear layers collapses into one linear map, so depth adds nothing: the network can only represent linear functions.
eval()changes module behavior (dropout off, normalization uses running statistics). It doesn’t stop autograd from recording the graph;no_grad()(orinference_mode()) does, saving memory and time.- About 16-18 bytes per parameter (weights, master copy, gradients, two AdamW moments) against 2 for BF16 inference: roughly 8-9 times as much, before activations.
Stretch exercises
- Compute
s = sigmoid(x.data), create a result with values, and have its backward closure adds * (1-s) * out.gradtox.grad. Use a stable sigmoid evaluation and compare with central differences at negative, zero and positive inputs. - With identical initialization, 0.005 usually progresses slowly, 0.05 often converges sooner, and 0.5 can oscillate or fail. These are expectations, not guaranteed outcomes: seed, loss and saturation affect the curves. Plot all runs over the same number of updates.
- Momentum carries a running gradient, for example
v = mu*v + grad; p -= lr*v. Adam tracks first and second moments, bias-corrects both, then updates withlr*m_hat/(sqrt(v_hat)+eps). Clear accumulated gradients each step and tune each optimizer fairly; steps to convergence depend on initialization and hyperparameters. - Use a nonlinear hidden layer between two
nn.Linearmodules and train on the four XOR pairs with AdamW. For 1B FP32 parameters, weights, gradients and two FP32 moments total about 16 GB, excluding activations and small step counters. A mixed-precision layout with BF16 weights/gradients and FP32 master weights/moments also totals 16 GB; FP32 gradients raise it to 18 GB. BF16 inference weights need 2 GB.
4. Text as numbers
Check your understanding
- Byte-level BPE starts from all 256 byte values, so any string, in any language or with any emoji, can be written as bytes and therefore as tokens. A word tokenizer has no ID for words it never saw.
- Later merges were learned on text where earlier merges had already been applied; they refer to tokens that only exist after those merges. Applying them in another order produces different (and unknown-to-the-model) token sequences.
- The model was trained to begin answers after
<|im_start|>assistant\n. Without it, the most likely continuation is often more of the user’s turn or a new role marker, so the model continues the prompt instead of answering. - Overlapping windows share text. If one window is in training and its neighbor in validation, the model has already seen most of the validation text, and validation loss underestimates the true error.
- An embedding maps each token ID to a vector regardless of where it appears. A permuted sequence gets the same set of vectors; order must come from positional embeddings or RoPE.
Stretch exercises
- Counts depend on the learned merge table: case, spaces and punctuation change which pairs were frequent enough to merge. Encode all three with the same trained tokenizer and print IDs and decoded pieces; no fixed counts follow from vocabulary size alone.
- Divide the story’s UTF-8 byte count by its token count for each trained vocabulary. More merges normally increase bytes per token on the training text, with diminishing gains once common substrings are represented. The knee depends on the corpus; evaluate held-out text too.
- Print both tokenizers’ decoded pieces for exactly the same paragraph. GPT-2’s larger vocabulary often preserves common words or leading-space word pieces that the 512-token tokenizer splits, but its pre-tokenizer also differs, so vocabulary size alone does not explain every difference.
- Keep live symbols in a linked list and heap entries
(merge_rank, position, version). Pop the best valid pair, merge it, invalidate stale entries and enqueue only its new neighbours. Preserve rank and left-to-right tie ordering, verify identical IDs, then benchmark the whole story.
5. Attention
Check your understanding
- Each query asks “how much should I read from each key?”, and its weights over the keys must form a distribution (sum to 1). That’s a normalization along each row, over keys.
- Masked positions must get exactly zero weight. Setting their scores to $-\infty$ before the softmax does that and renormalizes the rest; zeroing weights after the softmax would leave the row not summing to 1.
- Only the token at position 0 (its query, key and value). Causal masking hides every later position.
- Dot products of $d$-dimensional random vectors have variance proportional to $d$. Without the scaling, scores are large, the softmax saturates to nearly one-hot, and gradients through it vanish.
- The group size is 16 / 8 = 2, so query head 11 reads KV head 11 // 2 = 5.
Stretch exercises
-
Compute
scores = x @ x.T, mask entries above the diagonal to negative infinity, and apply row softmax with scale 1. The causal weight table, rounded to four decimals, is:1.0000 0 0 0 0 0 0.3680 0.6320 0 0 0 0 0.2284 0.3893 0.3822 0 0 0 0.2046 0.2956 0.2915 0.2084 0 0 0.1753 0.2250 0.2269 0.1570 0.2158 0 0.1385 0.2184 0.2128 0.1420 0.0988 0.1896Before rounding, every row sums to 1 and has zeros at future positions.
-
Fill
[B, T, H*d]with distinct values encoding position, head and feature. Splitting must reshape to[B, T, H, d]and transpose to[B, H, T, d]; assert individual head contents as well as the correct round trip. A wrong split and its matching wrong inverse can pass a round-trip check together. -
For absolute query/key positions, use
allowed = (key_pos <= query_pos) & (key_pos > query_pos - w). This includes the current key and at mostw-1earlier keys; handle queries with no valid keys explicitly. -
The score/probability tensors each have
B*H*T*Telements, so doubling T roughly quadruples their memory and attention arithmetic. Reset peak allocation statistics and synchronize each measurement; estimate the fit limit from measured peak memory, leaving room for weights and other tensors. Small-size timing may be dominated by overhead.
6. The transformer block
Check your understanding
- The MLP mixes features within each position; attention mixes information across positions.
- A pre-norm block computes
x + attn(norm(x))andx + mlp(norm(x)): if both branch outputs are zero, each addition returnsxunchanged. That’s why residual networks start close to the identity and train stably. - Storage: one vocabulary-by-width matrix instead of two (for GPT-2, 38.6M fewer parameters). Training: the matrix gets gradients from both uses, so input and output representations are coupled.
- LayerNorm normalizes and then applies a learned scale $\gamma$ and shift $\beta$, which can move the output to any mean and variance.
- Attention has four $D \times D$ projections (Q, K, V, output): $4D^2$. The MLP has $D \times 4D$ and $4D \times D$: $8D^2$. Total $12D^2$, ignoring biases and norms.
Stretch exercises
- With vocabulary 50,257, context 1,024 and a tied head, the count is $50{,}257(1{,}024)+1{,}024^2+24(12(1{,}024)^2+13(1{,}024))+2(1{,}024)=354{,}823{,}168$. Thus “355M” is rounded, not an exact count.
- Collect each block’s residual output before final normalization and plot mean vector norm over batch/positions. Independent zero-mean updates would make squared norms grow roughly linearly, giving square-root growth of norms; correlated updates and initialization can change this. Measure rather than assuming monotonic growth.
- Post-norm computes
a = LN1(x + attn(x)); y = LN2(a + mlp(a)). With zero branches it returnsLN2(LN1(x)), generally different fromx, because normalization remains on the residual path. - Exact GELU is $x(1+\operatorname{erf}(x/\sqrt{2}))/2$. A dense scalar sweep over [-6, 6] gives a maximum absolute difference of about $4.73\times10^{-4}$ from the tanh approximation. That is an activation error, not a logit bound: compare fixed-input logits after replacing every GELU, since later weights and layers can amplify or cancel it.
7. Training
Check your understanding
- It applies a numerically stable log-softmax internally (subtracting the maximum, using log-sum-exp). Passing probabilities would apply softmax twice, or take the log of rounded probabilities and lose precision.
- A uniform guess gives $\ln 50{,}257 \approx 10.8$. A much larger initial loss means the model is confidently wrong: initialization too large, a missing scale, or mismatched targets.
- Overfitting: the model memorizes training text, which lowers training loss, while its predictions on unseen text get worse.
- Teacher forcing always feeds the correct previous tokens. At generation time, the model feeds its own outputs, so one early mistake shifts it into contexts it never saw; loss also doesn’t measure qualities like repetition.
- The optimizer state (AdamW moments and step count), the learning-rate schedule position, the random number generator states, and the data order position.
Stretch exercises
- On each validation improvement, clone the state tensors or copy the state dict; save patience and best loss, then restore the best weights after stopping. Holding references to live tensors would overwrite the saved best state. Judge sample changes with fixed prompts/seeds and report best held-out loss; improvement in readability is not guaranteed.
- Change one setting per run with the same data split and training budget. More dropout can reduce overfitting but slow fitting; larger width increases capacity and cost; 1e-2 can converge faster or destabilize training compared with 1e-3. Record best validation loss rather than assuming any setting wins.
- With comparable data quality and enough optimization, more distinct text usually reduces memorization and the train/validation gap. Preserve document-level splits and report tokens processed as well as steps; a fixed step budget may undertrain the larger corpus.
- Zero gradients once, sum each microbatch’s token losses divided by the total valid-token count across all microbatches, then clip and step once. Compare with the concatenated batch using dropout off and identical initial optimizer state. Unequal token counts require token weighting, and different reduction order can introduce small floating-point differences.
8. Generation
Check your understanding
- Dividing logits by 0 is undefined, and the limit as temperature goes to 0 is the argmax. A separate branch computes the argmax directly, deterministically.
- The mass before token 0 is 0 (keep), before token 1 is 0.4 (keep, since 0.4 < 0.5), before token 2 is 0.7 (drop). Tokens 0 and 1 remain, renormalized to 4/7 and 3/7.
- So that a request’s samples depend only on its own seed, not on how many other requests used a shared generator before it. Reproducibility and isolation require it.
- The prompt’s last position already produces the logits for the next token during prefill.
Stretch exercises
- Decode incrementally and check the accumulated suffix, including stops spanning token boundaries and UTF-8 pieces. Hold back suffixes that could begin a stop string before streaming them. A stop may end inside one token’s text, so trim the decoded output at the match rather than relying only on token IDs.
- For output count
c_i, subtractpresence_penalty * (c_i > 0)andfrequency_penalty * c_ifrom logit i. Repetition penalty instead divides positive seen logits and multiplies negative ones by its factor. Compare repetition and held-out quality at matched settings; stronger penalties can suppress useful repeated words. - Keep four candidate histories, expand them using accumulated log-softmax scores, and retain the best four while tracking finished hypotheses separately. Compare summed logprob at fixed length or specify a length penalty. Beam search often finds higher-scoring sequences, but finite-width search need not beat greedy globally and probability need not predict readability.
- Define a grammar, for example
-?[0-9]+(\.[0-9]+)?, and track states for sign, integer, decimal point and fraction. For each state, allow a vocabulary token only if all its bytes follow valid transitions; allow EOS only in accepting states. A single context-independent “numeric token” flag cannot enforce the grammar.
9. Real weights
Check your understanding
- The header is a JSON block at the start of the file, preceded by its length, listing every tensor’s name, dtype, shape and byte range. Reading it doesn’t touch the data section.
- Pickle can execute arbitrary code during loading. A malicious checkpoint can run commands on your machine; safetensors contains only data.
- A square matrix and its transpose have the same shape, so the copy succeeds; the outputs are wrong. Only comparing values with a reference reveals it.
- With tied weights, the head is the embedding matrix, so storing it twice would be redundant. The loader must point the head at the embedding (or verify equality if both are present).
- The output of the first layer (or the embedding), then each layer in turn: the first point where values diverge contains the bug.
Stretch exercises
- Sum
numel * dtype_sizeover stored tensors and compare with the tied parameter formula times 4 for an FP32 checkpoint. Additional bytes may come from duplicated tied weights or stored buffers, such as attention masks in some checkpoint formats. Report names/dtypes rather than assuming every GPT-2 snapshot stores the same extras. - Load the same weights in both dtypes, evaluate fixed token IDs, and report max-absolute and relative-L2 logit error. BF16 rounding can flip near-tied argmax choices; after the first flip, free-running texts no longer isolate precision error. Neither identical text nor a particular error magnitude is guaranteed.
- The attention
c_projis square, so its shape check passes; reference layer/logit comparisons catch the wrong orientation. The MLPc_projis rectangular, so skipping its transpose also fails shape checks. Generated text can degrade, but text inspection alone is an unreliable correctness test. - Export a matching config and Hugging Face names, transposing linear weights into GPT-2’s Conv1D convention while leaving embeddings/norms unchanged. Preserve biases, tied-head semantics and tanh GELU. Reload with
transformersin eval mode and compare logits for several fixed lengths and batches.
10. GPU performance
Check your understanding
- Kernel launches are asynchronous. Without synchronizing first, earlier queued work is still running and gets counted in your timing.
- Profilers add overhead and change timing. Use the profiler to find where time goes, and separate clean runs to measure how much.
- Each weight byte supports about 16 FLOPs (2 FLOPs per BF16 weight × 16 sequences ÷ 2 bytes). That’s far below the H100’s ridge point of about 295 FLOPs per byte, so it’s memory-bound.
- Amdahl’s law: $1 / (0.7 + 0.3/3) = 1/0.8 = 1.25\times$.
- DGX Spark’s CPU and GPU share one LPDDR5x memory with 273 GB/s of bandwidth. CPU-heavy work consumes bandwidth the memory-bound decode needs.
Stretch exercises
- Unsynchronized wall time mostly measures enqueue/launch time and can miss the GPU computation. Synchronizing before and after gives end-to-end completed-operation time; CUDA events measure the interval on the GPU stream. Warm up first and use the same TF32 setting in every run.
- For FP32
x+y, count two reads and one write: achieved bandwidth is12*N/timebytes/s. Small vectors are launch-bound; large vectors approach a bandwidth plateau. Use repeated synchronized runs and distinguish cache-resident cases from DRAM traffic; the plateau’s fraction of peak is hardware-dependent. - Use $2N^3/t$ FLOP/s and approximately $N/3$ FLOPs/byte for square BF16 GEMM under a one-read-per-input, one-write-output traffic model. Small matrices underfill the GPU; large ones can approach compute peak. Whether any size reaches 70% depends on device, clocks and library kernels.
- Inspect a GPU timeline, count kernels in one warmed-up decode step, and compare their occupied-time union with the step interval. Gaps identify time without kernel execution, although dependencies and transfers also cause gaps. Measure final latency in a separate unprofiled run.
11. CUDA
Check your understanding
- Block sizes come in fixed sizes (multiples of 32, often 256), and $N$ is rarely a multiple. Rounding the grid up and having extra threads return keeps the code simple and the block size efficient.
- CUDA doesn’t guarantee that all blocks run at the same time or in any order. A block waiting for another that hasn’t been scheduled can deadlock.
- Addresses are 64 bytes apart, so the warp spans 32 × 64 = 2,048 bytes: 16 segments of 128 bytes. Only 32 × 4 = 128 bytes are used: 6.25% efficiency.
- Each inner iteration loads two 4-byte floats for 2 FLOPs: 0.25 FLOPs per byte. The roofline predicts speed limited to 0.25 × bandwidth, a tiny fraction of peak compute.
Stretch exercises
- Load/store four aligned FP32 elements per thread and handle the tail separately when N is not divisible by 4. Vectorization can reduce instruction overhead but does not reduce the three arrays’ total bytes; speedup can be small when scalar vector add already saturates bandwidth.
- Compare equal numbers of reads/writes using CUDA events. A stride of 32 spreads a warp’s reads across separate memory regions and wastes transactions.
32*t % Nrepeats indices for many N values, so also use a bijective strided permutation to distinguish coalescing effects from cache reuse. - With consecutive lanes mapped to rows, row-major A reads and output stores become strided instead of contiguous, increasing memory transactions. The precise slowdown depends on broadcasts, caches and matrix size; verify the output before timing.
- Load a 32×32 tile with contiguous lanes reading contiguous columns, synchronize, then swap tile/thread coordinates for contiguous output writes. Declare the shared tile as
[32][33]to avoid transpose bank conflicts. Mask edge loads/stores while ensuring every thread reaches the barrier.
12. Fast matrix multiplication
Check your understanding
- The first barrier ensures the whole tile is loaded before anyone computes with it. The second ensures everyone has finished computing before the next phase overwrites the tile. Without the second, fast threads overwrite data that slow threads are still reading.
- Larger tiles need more shared memory and registers per block, so fewer blocks fit on an SM (lower occupancy), and edge tiles waste more work. Past a point, lost parallelism outweighs the reuse.
- Column 5 is at word addresses $32i + 5$, all in bank 5: a 32-way conflict, serialized into 32 accesses. With row length 33 the addresses are $33i + 5$, in banks $(i + 5) \bmod 32$, all different: no conflict.
- Each weight is used exactly once in a matrix-vector product, so there’s no reuse for tiling to exploit. It stays memory-bound; only reading fewer bytes (quantization) or batching helps.
Stretch exercises
- For 256³ work, input loads are $2(256)^3/b$: 8,388,608, 4,194,304, 2,097,152 and 1,048,576 elements for tiles 4, 8, 16 and 32. Two 64×64 FP32 tiles need 32,768 bytes (32 KiB), before padding or double buffering. Compare with the per-block shared-memory limit; a 64×64 one-thread-per-output block would also exceed the usual thread limit.
- Give each thread four accumulators, reuse two A and two B register values per inner step, and write its 2×2 output region. This reduces repeated shared-memory reads at the cost of more registers. Check irregular boundaries and measure occupancy as well as time.
- A warp strides through one W row with consecutive lanes loading consecutive columns, accumulates in FP32, then reduces with warp shuffles and lets lane 0 store. Mask the K tail and count at least
2*N*Kweight bytes for bandwidth. Vector reuse and caching affect the full traffic accounting. - Stage correctly aligned BF16 tiles, load WMMA fragments, accumulate in FP32 and store the result with the required layout. Validate against cuBLAS with dtype-appropriate tolerances. cuBLAS generally has better pipelining and tuning; multiples of 16 remove edge handling but do not guarantee competitive speed.
13. Reductions and numerics
Check your understanding
- Softmax is invariant to adding a constant to all inputs: $e^{x_i - c} / \sum e^{x_j - c} = e^{x_i} / \sum e^{x_j}$.
- BF16 keeps FP32’s 8 exponent bits (same range) but only 7 mantissa bits (about 3 significant decimal digits) instead of 23.
- Adding a small number to a large BF16 sum loses the small number’s low bits, or all of it (swamping). FP32 accumulation keeps them; the result can be rounded to BF16 at the end.
- The maximum error can be dominated by one harmless outlier, or hide a systematic small error everywhere. An aggregate (relative L2, mean) describes the typical error.
- $-\infty$, the identity of max: it never changes the result.
Stretch exercises
- Use Kahan’s update
y = value - c; t = s + y; c = (t - s) - y; s = t, explicitly rounding every operation to BF16 for this experiment. With explicit round-to-nearest-even BF16 rounding after every operation, this recurrence returns 1,004 instead of the plain accumulator’s 32: an error of 4 relative to 1,000. The represented input is 0.10009765625, whose exact 10,000-term sum is 1,000.9765625: compensation does not undo input rounding or guarantee exact summation in such low precision. Compare errors against both reference sums. - Initialize active lanes from data and inactive lanes to negative infinity, then reduce with max. Padding with 0 makes every all-negative row incorrectly return 0. Test non-power-of-two widths as well as negative inputs.
- Evaluate fixed histories and report
max(abs(a-b))andnorm(a-b)/norm(a), with a defined convention for a zero reference norm. Also compare 50-step greedy continuations and record the first divergence. Precision, kernels and model determine the numbers; compare teacher-forced logits after divergence. - Load a row into registers, reduce its maximum, compute exponentials relative to that maximum, reduce their sum and normalize before storing. Use FP32 intermediates and mask inactive lanes. Ideal traffic is one read and one write per value; very long rows may spill registers or require a multi-block algorithm.
14. Triton
Check your understanding
- One program instance handles one block of data with vector operations; Triton maps it onto many threads (a few warps). A CUDA thread handles individual elements and you manage the cooperation yourself.
- Softmax is memory-bound, so time is proportional to memory passes. The fused kernel reads and writes each element once instead of about eight times.
- Masked lanes must not affect the maximum. Filling them with $-\infty$ makes them neutral; 0 would be wrong for rows of negative values.
- To verify the kernel on every input, to fall back when the kernel doesn’t support a shape or device, and to measure the kernel’s actual benefit.
Stretch exercises
- Apply GELU to the FP32 accumulator before its final cast/store. Match the reference’s exact or tanh convention; applying GELU before rounding the matmul output can differ slightly from
F.gelu(a @ b)with a low-precision intermediate. Compare with an appropriate reference and tolerance. - Autotune valid tile sizes, warp counts and stage counts keyed by matrix dimensions and dtype. Exclude compilation/tuning from steady-state timing, validate every selected configuration and compare with
torch.matmulunder identical settings. The winning configurations and speedups vary by GPU and shape. - Load both members of each rotary pair before writing either one, apply
a*cos-b*sinandb*cos+a*sin, and use absolute positions. Respect split-half pairing, rotary dimension, strides and differing Q/K head counts. Compare withapply_ropeon multiple positions and tails. - Reduce tile-local maxima/sums to obtain a row maximum and denominator, then reread and normalize the row. Alternatively update a running maximum/sum with exponential rescaling while scanning tiles, then make an output pass. Neither approach can finalize early tiles’ normalized outputs before the denominator is known.
15. FlashAttention
Check your understanding
- Its terms are measured relative to a running maximum that may still change. Normalizing early would need rescaling again; dividing once by $\ell$ at the end is cheaper and exact.
- Each tile’s softmax normalizes by that tile’s sum only. The true weights normalize by the sum over all tiles, so the per-tile results are wrongly weighted.
- Intermediate memory (from $O(T^2)$ to $O(T)$) and memory traffic. The arithmetic is unchanged.
- Decode has one query per sequence and head, too few programs to fill the GPU. Splitting the keys across programs and merging the partial results restores parallelism.
Stretch exercises
- Reset peak allocation statistics for each run and compare peaks above the same baseline. One FP32 4,096² score matrix costs 64 MiB per batch/head; materializing both scores and probabilities costs roughly twice that. Online attention replaces these with tile-sized scratch and row accumulators, but actual savings depend on implementation and live tensors.
- Let
m = max(m_s)over splits, then mergel = sum(exp(m_s-m)*l_s)anda = sum(exp(m_s-m)*a_s). Returna/l. Treat empty splits as zero mass and return zero for all-empty rows; compare different partitions with the unsplit reference. - For a window containing w keys including the current position, valid keys satisfy
q_pos-w < k_pos <= q_pos. Skip tiles whose keys all lie outside this range and mask the boundary tiles. Verify queries at the start of the sequence and after an existing prefix. - Reconstruct $P=\exp(S-\operatorname{LSE})$ tilewise using the original mask and score scale. Compute $dV=P^\top dO$, $dP=dO V^\top$, $dS=P\odot(dP-\sum_j P\odot dP)$, then $dQ=dS K/\sqrt d$ and $dK=dS^\top Q/\sqrt d$. Accumulate contributions across tiles, sum shared-head gradients for GQA and compare with autograd.
16. The KV cache
Check your understanding
- A query is used only once, by its own position, at the step it’s computed. Keys and values are read by every later position.
- Prefill: the prompt’s last position predicts the first new token.
- With a 3-token prefix, the chunk’s queries are at positions 3 and 4 and keys at 0-4. A square triangle aligned at the top-left would let query 3 see only key 0 and query 4 only keys 0-1. The rule
key_pos <= query_posgives the correct rectangle. - The cache stores KV heads, which GQA reduces. Every query head still computes scores against the shared keys, so the number of scores depends on query heads.
- Each decode step reads the whole cache once, and the cache grows with context while the weights don’t. At long contexts the KV bytes exceed the weight bytes.
Stretch exercises
- Repeated concatenation copies the growing history: total copied elements scale as
P*N + N*(N-1)/2for prompt length P and N new tokens, per cached stream. Preallocated writes copy only the new entries. Measure synchronized generation at each length; the point where quadratic copies dominate depends on model and device. - Qwen3-0.6B’s BF16 cache at 32,768 tokens is $28\times2\times8\times32{,}768\times128\times2=3{,}758{,}096{,}384$ bytes (3.5 GiB). An ideal decode reads roughly this amount plus about 1.2 GB of BF16 weights, so KV is about three quarters of their combined traffic. Report actual stored weights, context length and measured reuse separately.
- Store position p in physical slot
p % w, but retain absolute positions for RoPE and the causal/window mask. Gather in logical order or pass each slot’s absolute key position; physical slot indices cannot stand in for positions. Protect keys still needed by earlier queries in a multi-token chunk. - Store FP8 values and scales, dequantize inside attention, and ensure a changing scale never reinterprets older entries encoded with a different scale. Fixed calibrated per-head scales or explicit rescaling are possible choices. Compare against BF16 on the same 1,000 teacher-forced positions and report errors by context length; free-running divergence is a separate result.
17. Qwen3
Check your understanding
- A reshape assumes a layout. If the projection’s output features are ordered differently from how you split them into heads (for example, head-major versus interleaved), the shapes match but heads get the wrong features.
- Project with
k_proj, reshape into KV heads, apply QK-norm (RMSNorm per head), then RoPE at its absolute position. The rotated, normalized key is what’s cached. - Rotations compose: rotating $q$ by $m\theta$ and $k$ by $n\theta$ gives a dot product that depends only on $(m - n)\theta$. Shifting both positions by the same amount leaves the difference unchanged.
- SwiGLU computes a gate and a value from the input with two matrices, multiplies them elementwise ($\operatorname{SiLU}(xW_g) \odot xW_u$), then projects down with a third. GPT-2’s MLP has one up projection and one down.
- The architecture is the computation, independent of the weight values. Identical random weights in both implementations exercise every operation, and any mismatch is a bug in the computation. Loading real weights then only adds name mapping, which has its own checks.
Stretch exercises
- Each rotary pair is multiplied by an orthogonal 2×2 rotation, so its squared norm is unchanged. At equal positions, $R_p^\top R_p=I$, making the rotated Q/K dot product equal to the unrotated one. Verify with floating-point tolerances across several positions.
- Rotate features
(0,1), (2,3), ...instead of split-half pairs. Using unchanged Qwen3 weights applies different positional geometry and changes nonzero-position logits; compare fixed-input max-absolute and relative-L2 errors. An appropriate Q/K row permutation can reconcile conventions, so pairing and weight layout must change together. - Linear extension by factor f uses
position/f, equivalently frequencies divided by f, multiplying wavelengths by f. Multiplying positions by f instead shortens wavelengths. YaRN scales slow frequencies while preserving fast ones and blending between them; match the model’s attention scaling too. - For each Q/K head, reduce mean square in FP32, apply its learned RMSNorm weights, then rotate paired components and store once. Handle different Q and KV head counts and the checkpoint’s epsilon/pairing. Keep a reference backend and compare fixed-input logits before benchmarking the fused path.
18. Engine v1
Check your understanding
- Generated text can look fine with subtly wrong computation, and once one token differs everything after differs. Comparing layer outputs on fixed inputs finds where the computation first diverges.
- The prompt adds a KV cache proportional to its length, plus activations and logits proportional to it during prefill (all 4,096 × 151,936 logits in FP32 is 2.5 GB). Weights fitting says nothing about those.
- Prefill processes all prompt tokens into the cache but emits one token; later steps process one token each. The last emitted token is never processed into the cache unless generation continues.
to_emptyallocates fresh storage for every parameter, including the head, which breaks the shared storage. The head must be pointed back at the embedding.<|im_end|>(end of the assistant turn) and<|endoftext|>(end of document). With only<|endoftext|>, the model ends its turn and continues by writing an imaginary next turn.
Stretch exercises
- Slice final normalized hidden states to the last position before the head. With vocabulary 151,936 and FP32 logits, avoiding 4,095 rows saves $4{,}095\times151{,}936\times4=2{,}488{,}711{,}680$ bytes, about 2.32 GiB, plus unnecessary head computation. This does not remove prefill activations or KV.
- Register corresponding layer hooks, detach outputs, align their layouts/dtypes, and report the earliest layer outside absolute/relative tolerances. Compare embeddings first and use identical token IDs, positions and eval settings. Remove hooks in
finallyand avoid retaining all GPU outputs when memory is tight. - Pass the thinking option through the checkpoint’s chat template and record correctness and generated-token counts on the same five questions. Thinking mode often uses more tokens for intermediate reasoning, but quality improvements are task-dependent. State whether the budget includes reasoning tokens and inspect the rendered prompts.
- Parse vocabulary, ranked merges, added tokens, byte mapping and the exact Unicode-aware pre-tokenizer. Match normalization and added-token rules before BPE. Compare all 1,000 lines with
AutoTokenizer, including whitespace, Unicode and reserved tokens; Chapter 35 develops this complete boundary.
19. Fast decode
Check your understanding
int(token)makes the CPU wait until the GPU finishes, emptying the GPU’s queue. The next step’s kernels then start only as fast as the CPU can launch them, and the GPU idles between them. Without the sync, the CPU queues work ahead.- A graph replays recorded kernels; it has no CPU logic. The
ifwould need the value on the CPU, a sync that isn’t allowed during capture, and its branch would be frozen at capture time. - Unwritten slots sit at positions greater than every live query’s position, so the causal rule
key_pos <= query_posmasks them. - CPU operations run synchronously, so there’s no asynchronous queue to keep full and nothing for sync removal to gain.
- The graph keeps reading the old tensor’s address, so every replay feeds the same stale token. In-place updates keep the address the graph recorded.
Stretch exercises
- Typical v1 synchronizations include GPU scalar extraction with
int,.item()or.tolist(), CPU copies for output, and host-dependent stop checks. The fast loop defers these until a batched result transfer. Record the actual warning sites; sync-debug warnings do not exhaustively identify every synchronization. - Concatenate Q/K/V weights along their output dimension and gate/up likewise, then split the outputs using their original widths. This replaces five projection matmuls with two per layer, saving three matmul launches before other fusions. Actual profiler kernel counts also depend on the library implementations.
- Use the fused residual-add/RMSNorm operation while preserving both the updated residual and normalized input required by later branches. Verify layer outputs and logits, including epsilon and dtype, then time warmed graph replay. Gains depend on how much of decode is launch-bound.
- Right-pad to the bucket so real causal queries cannot attend to padding; keep real positions unchanged and give padding isolated/scratch cache slots. Use a validity mask and select logits at the last real position, not the bucket’s last row. Padding positions alone do not protect cache writes or output selection.
- Use fixed-shaped sort/top-k tensors and mask logits in place; keep the token crossing the top-p threshold by testing cumulative mass before each token. Sample with fixed-shaped device operations and per-request RNG. Test empirical frequencies against the reference rather than expecting identical RNG draws across different algorithms.
20. Quantization
Check your understanding
- Decode is memory-bound, so time is proportional to bytes read. Large-batch prefill is compute-bound, and its arithmetic doesn’t shrink with fewer bits (unless the math itself runs at lower precision).
- Each scale covers fewer values, so a large value raises the step size for fewer neighbors. The cost is more scales: metadata bytes and work in the kernel.
- Clipping saturates rare large weights (more weight error) but gives all others a finer grid. If the clipped weights multiply inputs that are usually small, the output improves.
- It materializes the full-precision weight matrix and reads it again, so as many bytes cross the memory bus as without quantization (more, counting the dequantization pass).
- Their errors are amplified: the head’s errors land directly on the logits, and a router’s small error can flip a discrete expert choice. They’re also small, so quantizing them saves little.
Stretch exercises
- For b-bit codes and one s-bit scale per group of g weights, effective bits/weight are approximately
b + s/g, plus any zeros, padding or other metadata. Smaller groups generally reduce held-out output error but increase metadata. Plot error on held-out activations, not only weight reconstruction error. - Dispatch CUDA calls to a kernel that unpacks and dequantizes tiles in registers while multiplying BF16 activations. Avoid constructing a full BF16 weight matrix, keep graph buffers stable and retain a reference path. Validate logits and compare memory and warmed decode speed; a small launch-bound model need not realize a fourfold gain.
- Clamp activation statistics away from zero, search alpha on calibration data, quantize
W*diag(s)and evaluate withx/s. Select by output error, then evaluate on held-out inputs. Fold inverse scaling into a compatible preceding operation only when every affected consumer is adjusted consistently. - Build a damped activation Hessian, factor its inverse, and for column i quantize
w_i, forme_i = (w_i-q_i)/U_ii, then subtracte_i*U_i,jfrom unprocessed columns. Handle poorly observed columns and group boundaries. Compare held-out output error with RTN at matched bits/group size; Chapter 39 gives the full implementation.
21. Fine-tuning
Check your understanding
- Under the causal mask, only the last token has attended to the whole input. The first token has seen only itself.
- The last prompt token’s position: its target is the first response token.
- Masking by ID would also mask the real EOS that ends the response, so the model would never learn to stop.
- Training needs gradients and optimizer state for every trained parameter (and an FP32 master copy in mixed precision): about 18 bytes per parameter against 2, plus activations saved for the backward pass.
- Fine-tuning on 935 examples teaches a pattern (the answer format and when to stop). Facts need far more text than that; a pretrained model already has them.
Stretch exercises
- Prompt-inclusive loss spends updates learning prompt tokens as well as responses; response-only loss targets the desired completion behavior directly. Under the same budget it often improves response learning efficiency, but there is no guaranteed winner. Evaluate both with the same response-only held-out mask and generation prompts.
- Use the last non-padding hidden state for classification, keep one fixed train/validation/test split, and count only parameters with gradients enabled. Head-only training is cheapest; unfreezing the last block or all layers adds adaptability, memory cost and overfitting risk. Select hyperparameters on validation data before reporting test accuracy.
- Concatenate examples, restart their position IDs, and allow attention only within each example to earlier/equal positions. Mask cross-example next-token targets and retain response-only loss boundaries. Compare token-weighted loss with unpacked examples; packing helps most when padding waste was large.
- Let $\Delta=\log\pi_\theta(y_c|x)-\log\pi_\theta(y_r|x)-\log\pi_{ref}(y_c|x)+\log\pi_{ref}(y_r|x)$. Minimize $-\log\sigma(\beta\Delta)$ using response-token logprob sums and a frozen reference. Mask prompt/padding tokens and compare held-out preference accuracy and ordinary response quality.
22. LoRA
Check your understanding
- The adapter’s contribution is $\frac{\alpha}{r} BAx$, which is zero when $B = 0$.
- $\partial L/\partial A$ contains $B^\top$, which is zero. $B$’s gradient isn’t zero, so after one step $B \ne 0$ and $A$ starts receiving gradients.
- Merging costs nothing at inference (an ordinary weight matrix). Separate adapters let one base model serve many fine-tunes, swapped per request, each a few MB.
- The base is stored in 4 bits (about a quarter of the memory), but each forward and backward pass must dequantize it, which costs time.
- The adapter learned a correction relative to specific weights. On different weights, the same correction no longer corresponds to the intended change.
Stretch exercises
- Each adapted linear layer adds
r*(d_in+d_out)parameters, so parameter count grows linearly with rank. Run identical data/step budgets and multiple seeds; plot held-out loss against the actual count. The rank where gains flatten is task- and layer-selection-dependent. - LoRA reduces trainable gradients and optimizer state substantially, while still needing base weights and activations for adapter training. Compare held-out loss, peak allocation and steady-state step time with matched data and steps. Lower memory is expected; equal quality or faster steps is not guaranteed.
- Compute the shared base projection once, gather each row’s A/B matrices and add its scaled adapter correction. Compare with separate row runs and cover a no-adapter row. Preserve each adapter’s rank/scale and prevent one row’s weights from affecting another.
- Initialize m to the base weight’s norm along the exercise’s column convention and compute the normalized direction from
W + BA, with a safe denominator. Count m’s extra parameters when choosing a rank for the comparison. Check that zero adapter correction reproduces W; translate the normalization axis if using[out, in]storage with a different convention.
23. Inside the model
Check your understanding
- The residual stream at every layer lives in the same space (each layer adds to it), and the head was trained to read that space. Early layers aren’t trained to be decoded, so the lens is approximate there.
- The direction is a difference of means. Anything else that differs between the groups (topic, length) also ends up in the difference and gets steered or removed along with the property.
- Adding moves every activation along the direction by a fixed amount. Projecting out sets the component along the direction to exactly zero, whatever it was.
- The residual stream is a sum of what the embedding and each layer write. If none of them can write along $d$, the stream never has a component along $d$.
- Measure the target behavior on held-out prompts, and the KL divergence from the original and held-out loss on unrelated data. Small changes elsewhere show a targeted edit.
Stretch exercises
- Learn the direction using training examples at each layer, apply the same 1.0× mean-gap convention and evaluate all-caps probability on held-out prompts. Plot the baseline as well as the steered probability. The best layer is empirical and can change with model or prompt distribution.
- At every layer, apply the final norm and head to the last prompt position, then rank the exact token ID for
Paris. Report the first top-five entry and its probability; it may leave and reenter. The layer depends on the checkpoint and tokenizer, so it cannot be inferred from architecture alone. - Select a layer on held-out calibration prompts, then generate from identical test prompts with
lambda0, 4 and 8. Larger additions usually strengthen the selected style but can distort content or fluency. Record the direction’s normalization and compare unrelated-task loss/KL as well as the five answers. - For unit direction d, project each output writer with
W' = (I-dd.T)Wandb' = (I-dd.T)b; project embedding rows similarly. The runtime comparison must project those same writes, not just one selected layer. Preserve head tying consistently: if editing a tied embedding also edits the head, apply the equivalent head change in the hooked reference. Save/reload and compare fixed-input logits.
24. Batching
Check your understanding
- All requests share the weights, so one read serves the whole batch. Each request has its own KV cache, so cache reads grow with the batch.
- Finished requests’ slots idle until the whole batch finishes, and new requests wait for the slowest member. Iteration-level scheduling refills slots every step.
- Decode latency is what streaming users feel; a large prefill in the same step would stall every running request. Serving decode first and filling the remaining budget with prefill chunks keeps time per token steady.
- Each row has its own
positionsentry and its own cache row (slot). Attention masks bykey_pos <= query_posper row, so rows at different lengths need no padding. - So that a request’s samples don’t depend on which other requests share its batch, keeping results reproducible and isolated.
Stretch exercises
- Throughput generally rises as more requests reuse each weight read, then plateaus as compute, KV traffic or capacity becomes limiting. Mean TPOT can grow even while aggregate throughput improves. Use the same lengths and sampling settings and report both curves; the saturation slot count is hardware-dependent.
- Draw interarrival times from an exponential distribution with mean
1/rate, and measure TTFT from the scheduled arrival. Smaller token budgets limit prefill interference but can increase queueing/prefill steps; larger budgets trade better utilization for longer decode stalls. Report p50/p99 at the same offered load and include failures. - Give each request an output queue, submit work to one owning engine loop, and stream incremental text as SSE events. Validate/admit before sending headers and abort on disconnect. Keep tokenization and detokenization from blocking decode; Chapter 36 develops the process-separated server.
- Save prompt/output IDs, progress and RNG state, retire the victim’s slot, then rebuild its KV on readmission without sampling its already committed tokens again. Compare with solo execution using per-request generators and fixed settings. Bitwise equality may also require shape-independent numerics or a deterministic reference backend.
25. Paged attention
Check your understanding
- The output length isn’t known until generation stops, so exact allocation is impossible; reserving the maximum wastes memory, and growing contiguous buffers means copies and fragmentation.
- In physical block
table[p // b], at offsetp % b. - A block’s keys and values depend on every earlier token. Two blocks with the same tokens but different histories have different K/V.
- When a sequence is about to write into a block whose reference count is above 1. Usually the partially filled last block of a shared prompt.
- The prompt’s last position must run through the model to produce the logits for the first output token.
Stretch exercises
- Utilization is
sum(lengths)/(block_size*sum(ceil(length/block_size)))for nonempty requests without sharing. Smaller blocks reduce tail waste; larger blocks reduce block-table and allocator overhead. Plot the same length set for every size and report metadata separately; a utilization peak alone does not choose the fastest size. - Reserve every chunk before model execution, release evictable zero-reference prefixes first, and preempt a live request if capacity is still insufficient. Keep block reference counts and copy-on-write correct during adoption and writes. Verify request outputs and that every block is reclaimed after completion/abort.
- Prefill once, fork references to shared blocks, and give each child independent sampling state. A child writing into a shared partial block must copy it first. Compare each branch with a solo request using its own seed and verify freeing one child does not invalidate the others.
- Map each query’s absolute position through its request’s block table and stream paged K/V using online softmax. Include the newly written chunk only up to each query’s causal position and keep ragged requests isolated. Compare against gathered dense attention across partial blocks and nonzero prefixes.
26. Speculative decoding
Check your understanding
- At batch 1, a decode step is memory-bound: reading the weights dominates. Five tokens read the same weights once, and the extra arithmetic is small.
- Through acceptance, $y$ is output with probability $q(y)\min(1, p(y)/q(y)) = \min(p(y), q(y))$. A rejection happens with probability $\sum_z \max(0, p(z) - q(z))$ and then outputs $y$ with probability $\max(0, p(y)-q(y))$ divided by that same sum. The total is $\min(p, q) + \max(0, p - q) = p(y)$.
- If the first proposal is rejected, the target’s own choice (or the residual sample) replaces it; if all are accepted, the bonus token is added.
- The target and draft wrote K/V for proposals that may have been rejected. Both are truncated to the committed tokens except the newest one, which hasn’t been fed yet: base length + accepted + 1.
- At large batches decode is compute-bound, so verifying $k+1$ tokens costs about $k+1$ times as much, and rejected tokens are wasted compute.
Stretch exercises
- Measure acceptance, draft time and verification time for each k. With approximately constant conditional acceptance alpha, expected emitted tokens per round are $E=(1-\alpha^{k+1})/(1-\alpha)$, or k+1 at alpha=1; predicted speedup is
E*t_decode/(t_draft(k)+t_verify(k)). Observed acceptance varies by depth and prompt, so the formula is a model, not a guaranteed speedup. - Search for an earlier occurrence of the trailing three-token sequence that has following tokens available, and propose up to k of those tokens. Treat each proposal as deterministic and verify with the target’s greedy or exact sampling rule. Rewriting often reuses source spans; empty matches or changed wording reduce acceptance.
- Estimate recent acceptance and draft/verify cost for candidate k values, then choose the largest expected tokens-per-second ratio, including k=0. Smooth estimates and periodically probe alternatives. Adaptation should respond to both acceptance and load, since longer verification can cost more than it saves.
- Flatten the 14 draft nodes with depth-based absolute positions and mask each node to the shared prefix and its ancestors, including itself. Match target argmax choices along a single root-to-leaf path, commit only that path and obtain a correction/bonus from its endpoint. Keep sibling KV separate and compare with ordinary greedy decoding.
27. Mixture of experts
Check your understanding
- Each token only runs, and so only reads, its active experts. But any token may choose any expert, so all must be resident.
- Routing is a discrete choice. A rounding difference between two nearly tied experts can swap them and change the output completely.
- Each expert’s tokens become contiguous rows, so one grouped GEMM can process all experts with dense matrix multiplications and no per-expert gather.
- With $B$ tokens each choosing $k$ of $E$ experts, the number of distinct experts touched, $E(1 - (1 - k/E)^B)$, approaches $E$ quickly; by batch 32, Qwen3-30B-A3B touches 112 of 128 experts per layer.
- An unbalanced router overloads a few experts and wastes the rest during training. At inference, the busiest expert sets the step time, and in expert-parallel serving the busiest GPU does.
Stretch exercises
- Count top-k routed assignments per expert and layer on representative text, then compare dispersion, max/mean load or entropy. Each token contributes k assignments, so totals must match
tokens*k. There is no universal early-versus-late balance pattern; use trained routers and report the observed distribution. - Group routed rows by expert, build prefix offsets and launch output tiles for each nonempty expert range. Use the matching expert weights, then scatter and combine outputs with routing weights. Handle repeated token destinations and empty experts; validate against per-expert reference matmuls before timing.
- Quantize each expert’s matrices with local groups/scales while leaving the router and attention in BF16. Expert storage falls toward one quarter of BF16 plus metadata, but total model reduction depends on the expert fraction. Compare fixed-input logits and routing choices; later layers’ routes can still change through altered hidden states.
- Keep host weights pinned, map expert IDs to resident GPU slots and copy misses before their kernels consume the slots. Do not evict an expert with in-flight users. Record transferred bytes, cache hits and stall time over the conversation; latency can be PCIe-bound despite a smaller GPU footprint.
28. Linear attention
Check your understanding
- Without the exponential, $\sum_i (q_t \cdot k_i) v_i = q_t^\top \sum_i k_i v_i^\top$: the sum over the past can be accumulated once into a matrix and updated per token.
- A query reads $\sum_i (k_i \cdot k_j) v_i$: every stored value leaks in proportionally to its key’s similarity. At most $d_k$ keys can be mutually orthogonal, which limits clean storage to $d_k$ associations.
- The error $v_t - S^\top k_t$, scaled by $\beta$. With $\beta = 1$ and a unit key, after the update $S^\top k_t = v_t$ exactly.
- Within a chunk, each position’s write depends on what earlier positions in the same chunk wrote (the prediction uses the updated state). That dependency is a unit lower-triangular linear system.
- Every token’s update mixes into the same matrix; there’s no per-position entry to forget. Engines save snapshots of the state and restore one instead.
Stretch exercises
- An untouched contribution decays as $0.95^n$, with a half-life of about 13.5 steps and about 0.006 of its initial amplitude after 100 steps. In an isolated recall example, error relative to the original value grows roughly as
1-0.95**n. Later corrective writes and key interference change this simple relationship. - Small chunks incur many sequential chunk transitions; large chunks increase local triangular work and temporary storage. A dense C×C triangular solve with C right-hand sides has cubic work per chunk, giving roughly quadratic per-position cost for that component. Time the proposed sizes at T=4,096; the best C depends on dimensions, implementation and device.
- Fuse decay, old-state prediction, gated correction, state update and output read for each sequence/head, keeping the state tile on chip when it fits. Use FP32 state arithmetic as in the reference and avoid reading partially overwritten state for the prediction. Compare multi-step state and outputs; large states can cause register spills.
- Remove the prediction-subtraction correction and use $S_t=\alpha_t S_{t-1}+k_t v_t^\top$ with readout $q_t^\top S_t$. Within a chunk, write contributions are weighted by products of intervening alpha values and the entering state by the full prefix product. This is the simplified selective recurrence in the exercise; removing correction does not by itself implement every component of a full Mamba-2 block.
29. Sparse attention
Check your understanding
- To find the top $k$ scores it must compute all of them, which reads every key: the expensive part.
- Blocks reduce the number of selection decisions and index keys by the block size, and the selected keys are contiguous in memory, which makes reads efficient.
- The partially filled block containing the query’s own position. It’s always readable so recent tokens are never lost to selection.
- All keys must stay stored, because a future query may select any of them. Only the reads per step shrink.
- Selection is discrete: a tiny difference can pick a different block and change the output a lot. If the two paths broke ties differently, cached decoding wouldn’t reproduce the full forward.
Stretch exercises
- Union the first four valid keys with the selected keys and remove duplicates before softmax. For flat scores, most probability mass may still lie outside the sparse set; for peaked scores, retaining the important peaks gives much smaller error. Ordinary retained “sink” keys here are different from Chapter 32’s learned denominator-only sink.
- Maintain absolute positions and RoPE while mapping storage into
position % w. Mask by the logical window, not by physical slot order, and prevent chunk writes from overwriting keys earlier queries still need. Compare multiple wraps and chunk sizes against the full-cache windowed reference. - Return selected block IDs, gather each block’s K/V and apply causal/token-valid masks within boundary blocks. Remove duplicates and define a sentinel for fewer-than-budget selections. Values should match the dense mask path within tolerance; actual sparse reads, rather than just a sparse mask, provide the traffic saving.
- Iterate only selected blocks, load each through its physical mapping, and update the running maximum, denominator and weighted value sum. Apply causal and partial-block masks and handle empty selections. Compare with the explicit gathered reference and measure bytes/time as selection budget varies.
30. Capstone
Check your understanding
- Gated DeltaNet (28), gated sparse attention with QSA and partial RoPE (29, 17), MoE with a shared expert (27), loading and quantization (9, 18, 20). New: hyper-connections and the n-gram memory.
- Read weights are per stream and per feature (four 2,560-vectors), in (0, 1) through a sigmoid. Write weights are one scalar per stream, in (0, 2) through $2\sigma(\cdot)$.
- Each head is an independent hash; distinct prime sizes make collisions in different heads independent. Two n-grams sharing a row in one head almost never share rows in all eight.
- Each token reads 16 rows of 160 BF16 values, about 5 KB. Disk and page-cache reads of that size are small next to the gigabytes of weights read per token.
- Keeping the old state is free, so a rejected speculation is undone by restoring it. A recurrent state can’t be truncated.
- The dense parameters in BF16 (attention, linear attention, shared experts and LM head): about 8.5 GB of the 9.7 GB per token.
Stretch exercises
- Snapshot paged KV progress and recurrent state before verification; after rejection restore both and rerun only the committed tokens to reconstruct the correct state. Synchronize the draft’s history/cache too. Compare complete greedy outputs with the nonspeculative path; a recurrent state cannot be fixed by KV truncation alone.
- INT8 roughly halves the dense BF16 weight bytes, before scales, while leaving already-quantized routed experts unchanged. In the chapter’s estimate, halving 8.5 GB of a 9.7 GB/token weight read leaves about 5.45 GB/token, suggesting at most roughly 1.8× improvement if that traffic alone dominates. Validate fixed-input logits and measure actual latency; quantization kernels and non-weight work reduce the gain.
- Inspect
mtp.tensors/configuration, match the checkpoint’s embedding/hidden fusion, normalization and transformer module, and validate its forward computation before drafting. Use the target’s embedding/head where specified and maintain separate draft state. Report acceptance by draft depth and total latency; matching prefixes alone does not establish a correct MTP implementation. - Keep quantized expert codes/scales in pinned memory, stage only the routed experts and synchronize copy completion before use. Exact next-layer routes require its input; overlap using a verified pre-gating mechanism or speculative prefetch with a fallback for misses. Compare transferred bytes divided by achieved PCIe bandwidth with measured token latency.
- Assign each request its own recurrent state and KV/block metadata, gather recurrent states into batches and scatter updated states back to their owners. Attention remains ragged and must isolate histories. Test joining, finishing, preemption and reordering against solo runs so slot reuse cannot mix recurrent states.
31. Engine v2
Check your understanding
- Each computes the next
nuncomputed positions of one request. Prompt/output boundaries change whether logits are sampled, not the scheduler’s work representation. - Attention needs sequence boundaries, lengths, block tables and slots; linear layers and tokenwise norms just see one matrix. Combining all prompt chunks gives larger matmuls and amortizes launches/weight reads.
- The last prompt position must be computed to obtain the first output token’s logits. KV alone does not contain those logits.
- Head-first insertion would make early prefix blocks older in the free queue and evict them sooner. Tail-first preserves the portion most likely to be reused.
- Images and adapters change activations even when IDs match. Cryptographic hashes prevent deliberately crafted collisions from aliasing another request’s KV; all activation-changing inputs must enter the identity.
- Preemption proves the pool is already under pressure. Admitting more work immediately can undo the reclaimed capacity and cause repeated eviction without progress.
- Swap is likely cheaper for a modest context: moving a few KiB per cached token can cost far less than running billions of parameters over its history again. Compare context size, transfer bandwidth and recompute FLOPs on the target hardware; it is not universal.
Stretch exercises
- Count adopted cached tokens, newly computed tokens and physically evicted cached tokens with clearly defined units. Multi-turn reuse improves when the pool preserves earlier full blocks, but the last position still needs logits and new text needs computation. Plot aggregate hit rate against pool size and avoid double-counting lookup hits that were not adopted.
- Swap wins when host round-trip transfer plus bookkeeping costs less than recomputing the evicted history, especially when the history has little surviving prefix reuse. Prefix caching can make recomputation mostly a hit, reducing swap’s advantage. Force the same pressure/workload and compare total runtime and p99 TPOT with caching on and off.
- Rank waiting requests by reusable prefix length, then compare hit rate, TTFT and throughput with FCFS. Popular-prefix traffic can indefinitely delay requests with short/no hits. Add age or bounded bypasses and measure worst wait rather than judging the scheduler only by hit rate.
- Store token-block paths in a radix tree, split compressed edges at divergence, track live references and evict unreferenced leaves in LRU order. Copy a matched partial block before extending it. Full-block hash chains already share exact prefixes; any extra hit-rate benefit should come from partial matches or changed eviction behavior, not from the tree representation alone.
32. Attention kernels
Check your understanding
- One decode query reads a large K/V history for relatively little arithmetic. Its memory lower bound is total required KV bytes across layers/requests divided by achievable bandwidth, before launch and other work.
- Query heads sharing a KV head reuse the same loaded keys/values. Several heads become matrix rows, increasing reuse and offering a useful matmul tile even with one token per request.
- Four requests give too few independent programs to occupy the GPU; partitioning each history supplies parallel work. Sixty-four requests may already supply enough programs, making split/merge overhead unnecessary.
- Split s contributes partition sum
exp(L_s)to the total softmax denominator, so its normalized output receives weightexp(L_s - logsumexp(L)). Partition sums and weighted numerators add associatively; numerical roundoff can still depend on order. - An empty split has zero probability mass, whose log is negative infinity. If every split is empty, subtracting negative infinity from negative infinity produces NaN; padding rows must instead return zero.
- Different tokens and heads have different ranges. Local scales spend more of the quantizer’s levels on each range, reducing error at the cost of scale storage and reads.
Stretch exercises
- Autotune only valid configurations using representative ragged lengths, head groups and hardware. Count useful KV bytes and compare achieved bandwidth with both attainable bandwidth and the predicted time floor. Exclude tuning/compilation, validate each configuration and include enough splits to occupy the GPU at small batches.
- Compute shared-prefix attention and private-suffix attention separately and merge their outputs using log-sum-exp weights. Ideally the 4,000-token shared prefix is loaded once per query tile instead of independently for 64 requests, though real reuse is limited by tiling/cache capacity. Distinct queries still require distinct scores; measure traffic as well as latency.
- Use the same held-out token sequence and teacher-forced loss with all cache formats, specifying scales, group axes and any unquantized residual region. INT8/FP8 often stay closer to BF16 than 4-bit caches, but actual perplexity changes are model-dependent. Per-channel key scaling addresses outliers; compare downstream behavior too.
- For sink logit s, initialize the running maximum to s, denominator to 1 and value accumulator to zero, then process the real keys normally. This gives $O=\sum_j e^{z_j}v_j/(e^s+\sum_j e^{z_j})$. For split attention, include the sink exactly once in the final denominator, not once per split; test against dense attention with an appended zero-valued sink.
33. Overhead
Check your understanding
- The smaller model’s memory work is short enough that many kernel launches and host decisions dominate. Reading the larger model’s weights takes much longer, so the same launch cost is a smaller fraction.
- It saves launches, repeated input reads and separate matmul setup, while making one larger operation. The concatenated weight storage replaces the originals rather than duplicating them.
- Data-dependent Python grouping and dynamically sized tensors/synchronizations cannot be replayed as a fixed graph. Device-side alignment uses fixed-capacity buffers and masked padded expert tiles instead.
- Dummy rows read zero-length contexts and write only scratch slots. Sequence metadata isolates requests, so they cannot write real KV or join another request’s attention.
- Captured kernels replay the bucket’s fixed program shape. Changing the host count underneath it would mismatch the captured metadata and actual padded rows.
- Step s+1 may already have launched; its speculative extra token is discarded. Finishing frees the request’s references after the work is ordered, without exposing its placeholder KV as a reusable prefix.
- The placeholder is not the actual token, so its hash does not identify the activation history. Only known committed IDs can be published under prefix keys.
Stretch exercises
- Fusion lowers kernel count and repeated reads; graphs lower host launch gaps while replaying the captured kernels. Profile the same warmed-up step before/after and report occupied GPU time and idle gaps, then measure latency without profiling. Graph capture alone does not imply fewer numerical kernels.
- Keep fixed-capacity metadata buffers keyed by stable request slots and update membership, lengths, token IDs and affected block-table rows in place. Lengths/positions of active decode requests still advance every step even when membership is unchanged. Compare host time at 256 requests and check that retired rows cannot retain live references.
- Capture fixed-shape norm/projection/MLP pieces into token-count buckets, with eager attention between pieces. Keep interface buffers/addresses stable and isolate padding rows’ metadata and cache writes. Benchmark TTFT and decode latency under mixed load; extra boundaries and padding can offset launch savings.
- Rotate normalized K with its absolute position, then write K and V directly into their assigned paged slots. Load each rotary pair before storing, respect head/stride layouts and mask padding/scratch rows. Compare pool contents and subsequent attention with
apply_ropefollowed bywrite_kv.
34. Sampling and structured output
Check your understanding
- Raw model logprobs describe the model distribution; temperature and truncation describe a sampling policy. Changing the policy should not change the reported underlying model score.
- Repetition penalty rescales logits of seen IDs with sign-aware multiplication/division; frequency penalty subtracts a term proportional to output counts. Avoid indiscriminately penalizing source vocabulary in summarization; frequent source words may be exactly what the answer needs.
- Independent exponential waiting times with rates
p_ihave token i as their minimum with probabilityp_i / sum(p). The race uses fixed-shaped device tensor operations and no host-selected index. - A shared generator advances when any neighbouring request samples, making a request’s output depend on batch membership. A per-request generator isolates its random stream, though numerical batch effects remain possible.
- A prefix hit supplies KV, not logits for each prompt position. Prompt logprobs require recomputing those positions (unless separately cached).
- Four normal requests arrive together; children adopt already published full prompt blocks of the first child. The final partial block, or last full block held back for logits, still runs for each child.
- Token pieces can split Unicode characters and include arbitrary byte sequences. A character-level state cannot consistently advance across such token boundaries.
- A fixed supported schema can bound structure/recursion and compile to a regular language. Arbitrarily nested JSON needs unbounded stack state and cannot be recognized by an ordinary finite automaton.
Stretch exercises
- Zero a slot’s count row on admission, update counts for each committed output token and clear it on retirement/reuse. Prompt counts, if needed for repetition penalty, should be tracked according to that penalty’s semantics. INT32 storage costs
4*max_num_seqs*vocabbytes, so measure device memory alongside reduced host rebuilding time. - Store bit i in word
i // 32, extract it with a shift/mask and replace forbidden logits with negative infinity. A packed mask uses about one eighth of a one-byte boolean mask’s memory, before padding. Validate vocab tails and signed bit 31; measure expansion/application time as well as mask transfer cost. - Combine byte-level JSON lexical states with a stack recording object/array context and expected separators/values. Precompute stack-independent transitions and simulate full token byte strings for stack-sensitive candidates. Handle strings, escapes, numbers and EOS only in accepting states, while bounding output length/nesting for resource control.
- Find a forced byte continuation, then use a tokenizer-compatible continuation that preserves already committed token boundaries. Append forced tokens and update guide/text state, but process those tokens through the model before the next unconstrained choice so KV stays correct. This saves sampling/guide steps and can batch forced-position computation; reporting skipped model positions would misstate the work.
35. The text boundary
Check your understanding
- The printable mapping makes every byte, including whitespace/control bytes, representable in JSON vocabulary strings without losing reversibility.
- Added tokens are literal reserved boundaries. Normalization or regex splitting first could alter or split them so that their trained IDs are never emitted.
- The regex determines which substrings BPE can merge. Different splits produce different IDs from the same merge table and therefore different model inputs.
- One token can end halfway through a UTF-8 character. The incremental decoder buffers unfinished bytes until later pieces complete the character, instead of replacing each incomplete piece.
- It holds back the trailing ST. It releases it when subsequent text disproves the stop prefix, or at normal finish; if OP follows, the completed stop is removed unless inclusion was requested.
- Templates come from outside the server’s code and run on request data. A sandbox restricts object access and execution; it does not remove the need to bound rendering work.
- A supported constraint masks every continuation outside the allowed tool-call syntax/name/schema. It guarantees membership in that language, not that a tool is appropriate or its arguments are factually correct.
Stretch exercises
- Compare exact IDs on mixed whitespace, Unicode, special tokens and long strings, then inspect the first divergent substring. Differences can come from normalization, added-token rules, regex boundaries, merge ranking or byte mapping. Match the checkpoint’s supported tokenizer semantics rather than changing outputs to fit only the sampled corpus.
- Maintain an incremental parser across chunks, emit the completed name once with a stable tool-call index/ID, then emit only new argument-text fragments. Retain escape/quote state so boundaries inside JSON strings are handled correctly. Test malformed and truncated calls and ensure fragments reconstruct the original valid arguments.
- Use a linked symbol list and heap ordered by merge rank and position; attach versions to discard stale candidate pairs. Only update the neighbours of a merge. Confirm byte-for-byte token ID equality on long repetitive and Unicode inputs before measuring the 50,000-character case.
- Build a DAG over input positions with vocabulary-piece edges weighted by log probability, maximize cumulative score with Viterbi and backtrack the pieces. Match normalization, unknown-token behavior and byte fallback as well as segmentation. Compare with deterministic
tokenizersencoding, with any segmentation sampling disabled.
36. The server
Check your understanding
- Frontend tokenization, JSON and detokenization no longer contend for the engine’s Python interpreter. The core can keep scheduling/launching while the API processes unpredictable text work.
- Invalid request parameters fail one request before mutation. A model-step/core crash may leave shared state inconsistent, so the core fails all affected requests and stops accepting work.
- Stops are text patterns crossing token boundaries, requiring incremental detokenization. The core owns IDs and numerical state; the frontend owns decoded text and visibility.
- Disconnect cancels response consumption; merging cancels producers; generation’s
finallysends abort; the core retires the request and frees KV/adapter references. The frontend removes its request queue. - Once HTTP 200 and response headers are sent, it cannot replace them with 503. Admission must precede streaming; later errors are stream events.
- The checkpoint’s authors chose defaults for that model’s behavior. Applying unrelated API defaults can change quality and reasoning behavior; explicit client values still override them.
Stretch exercises
- Parse the timeout at admission and set a monotonic deadline covering queueing and generation. On expiry, abort core work and finalize collected text with the exercise’s
finish_reason: "length"; if streaming has begun, send the final event rather than attempting a new HTTP response. Test expiry while queued, during decode and on disconnect. - Refill each key’s token bucket over time and reserve
prompt_tokens + max_tokensbefore admission. Reject insufficient reservations with 429 and a retry delay based on the deficit/refill rate; optionally refund unused output allowance at completion. Make reservations/refunds atomic so concurrent requests cannot overspend. - On SIGTERM mark readiness false, reject new admissions with 503, and keep existing streams/core work alive until they finish. Bound shutdown time and abort/release remaining requests if the bound expires. Test an in-flight response, a new rejected request and final process/resource cleanup.
- Preserve request IDs, bounded backpressure, abort/error propagation and readiness handling when replacing queues with sockets and msgpack. Define the transport’s message ordering and payload contract explicitly. Compare frontend/core latency and CPU cost at the same 256-stream load; serialization changes alone do not guarantee lower overhead.
37. Speculation in serving
Check your understanding
- Uncomputed-token counts, ragged verification rows, paged block allocation/trimming and per-request sampling settings already express speculative work. Verification changes commit rules rather than building a separate engine.
- A deterministic proposal has
q(d)=1, so acceptance isp(d). Acceptance contributes massp(d)to d; rejection samples the normalized remaining target mass, reconstructing p exactly. - If the request had L known tokens before verification, accepting two drafts leaves L+2 computed positions; the new correction/bonus token remains uncomputed. Blocks entirely beyond the committed computed prefix can be freed; its last partial block remains private.
- Drafts were chosen for a particular next step/history. Preemption changes when/how that step is computed, so stale proposals must be regenerated.
- A same-position draft runner has separate KV storage with compatible logical block addresses, while the target’s block manager owns allocation decisions. Draft state must be synchronized/restored for newly adopted prefixes.
- Larger batches already reuse weights well and can be compute-bound. Extra verification rows, rejected work, draft cost and buffers may cost more than the steps they save.
Stretch exercises
- For each k, estimate accepted-plus-correction/bonus tokens and divide by measured/predicted draft-and-verification time at the current batch size. Include k=0, smooth per-request acceptance and probe alternatives. High load can make k=0 best even when acceptance stays high.
- Bucket by both request count and k, giving
(1+k)*bucket_sizeverification rows with fixed buffers. Mask padded requests and keep allocation, rejection/commit decisions and ragged exceptions outside capture. Test output/cache equivalence and measure whether padding and verification compute outweigh replay savings. - Train the fused embedding/target-hidden input projection and draft layer on aligned target histories, with compatible vocabulary and positional semantics. Use a separate draft KV pool addressed by the target’s logical block tables and restore it after rejection/preemption. Compare held-out acceptance and total latency with Medusa; teacher-forced training quality alone does not measure rollout acceptance.
- Flatten the draft tree, assign depth positions, and let each node attend only to the committed prefix and its ancestor chain. For greedy decoding, traverse target-agreeing edges and commit the longest path plus its target correction/bonus, retaining only committed KV. Exact stochastic tree verification needs an explicit probability-correct acceptance rule beyond choosing the longest path.
38. GGUF
Check your understanding
- GGUF includes model/tokenizer metadata and per-tensor GGML types/layouts. A standalone safetensors weight file normally relies on separate config/tokenizer files.
- Memory mapping avoids eagerly copying a large file and permits tensor views/lazy page access. Some runtime repacking still needs allocated device/host storage.
- A byte pairs weights j and j+16 within a 32-weight block, in low/high nibbles. Pairing adjacent weights scrambles each block while retaining plausible shapes.
- Compressing sub-block scales reduces metadata overhead while adapting to local ranges. One scale for the whole legacy block cannot fit those ranges as flexibly.
- Q4_K is a tensor encoding. Q4_K_M is a model-level recipe mixing encodings for selected tensors, not a new block layout.
- They can be repacked into codes plus affine group scales/offsets. That shared layout simplifies dispatch, but expands some packed scale metadata and can use more memory than native Q4_K.
- Llama conversion permutes Q/K rows for its interleaved rotary convention. A split-half decoder must invert that permutation; the supported Qwen conversion convention already matches its row pairing.
Stretch exercises
- Compare each GGUF variant with BF16 on identical prompts and subsequent teacher-forced tokens, using a stated KL direction such as
KL(p_BF16 || p_quant). Confirm tokenizer IDs and architecture conventions first. Q8_0 often has lower divergence than Q4_K_M, but report measured per-token means/tails rather than assuming a fixed quality gap. - Decode each super-block’s scales/mins and packed codes in registers and multiply without repacking into expanded affine metadata. Native Q4_K is 144 bytes per 256 weights, or 4.5 bits/weight; the chapter’s affine representation costs about 6. Validate exact layout reconstruction and compare achieved bandwidth, residency and output error.
- Implement the tokenizer’s normalization and space marker, then repeatedly merge the highest-scored eligible adjacent piece with the specified tie rule; apply byte fallback when no vocabulary piece covers input. Scores are not BPE merge-rank entries. Compare with the checkpoint’s tokenizer on spaces, Unicode and added tokens.
- Weight reconstruction errors by calibration-derived per-column importance, fit each sub-block’s affine scale/minimum, and account for quantization of the stored scales/minima themselves. Evaluate on held-out activations and teacher-forced KL with the same model recipe. Lower weighted calibration error does not guarantee every prompt improves.
39. Quantized checkpoints
Check your understanding
- GPTQ uses activation-derived curvature to estimate how column errors affect outputs. It compensates remaining columns for each committed quantization error, rather than rounding independently.
- Act-order changes column processing/group assignments;
g_idxidentifies the resulting groups. Reordering codes and the corresponding activation columns makes runtime groups contiguous while preserving the linear operation. - Scaling salient columns gives them more effective quantization resolution. Inverse input scaling can be absorbed into the preceding normalization/projection; it isn’t free if left as a separate runtime operation.
- Exponent bits have finite range and spacing. Block scales place local values in a useful portion of that range, avoiding overflow or excessive rounding of small values.
- NVFP4 uses smaller blocks and more flexible scale representation/global scaling than MXFP4’s power-of-two block scales. E2M1 values alone do not determine error.
- Trellis states share constraints/codewords across a sequence, optimizing joint error beyond independent scalar choices. A small fixed-state recurrence and lookup tables keep decoding manageable.
- KL compares the full candidate distribution with a reference and is useful when a small/random model’s task loss is uninformative. Perplexity on representative real data remains necessary; neither metric replaces a task evaluation.
Stretch exercises
- Use the same 128 calibration sequences and separate evaluation text/task data, keeping group sizes, eligible layers and packing comparable. Record per-layer output error, teacher-forced KL and downstream accuracy. Sensitive layers and the relative RTN/GPTQ/AWQ ranking are empirical; ensure act-order metadata is interpreted correctly before comparing.
- Start each tile with the correct trellis state or reconstruct the required prefix state, then replay code transitions and lookups without fully materializing weights. Move the orthogonal transform to activations using the exact normalization, signs and transpose convention. Validate dense equivalence; tile boundaries cannot reset the recurrence arbitrarily.
- Factor the activation curvature into blocks and apply error feedback between blocks while jointly selecting trellis codes within each block. Include the trellis state constraints and transform conventions in reconstruction. Compare held-out quality with plain trellis at the same 2-bit budget, including metadata; the extra encoding complexity must produce a measured benefit.
- Read the expert blocks and exponent scales with the checkpoint’s interleaved layout and compare sampled dequantized entries to the BF16 conversion. Cover multiple experts, blocks and scale values; check shape/layout before blaming arithmetic. Match BF16 rounding in the reference comparison.
40. More hardware
Check your understanding
- A small-batch mat-vec reuses each weight little, so moving weights dominates. Quantization reduces those bytes if the kernel consumes codes directly.
- It reads codes, writes an expanded matrix and reads that matrix again, adding traffic and allocations. Dequantize tiles in registers or consume integer codes inside the dot product.
- Int8 activations enable fast integer dot-product instructions; local scales restore magnitude. Quantization/scale work and additional rounding are the costs.
- Runtime dispatch can choose supported instructions while leaving the rest of the binary portable. Globally compiling for one machine can emit illegal instructions even before dispatch on another CPU.
- Several requests touch many distinct experts, often more than eight per layer. The next step evicts experts the following request needs, causing repeated transfers.
- Streaming helps when weights don’t fit, compute can hide transfers, or prefill has enough rows to amortize them. Batch-1 decode over a slow host link is usually dominated by moving every layer’s weights again.
Stretch exercises
- Keep quantized codes/scales resident, use
simd::gemv_q4directly and reuse a persistent worker pool across the entire Rust decode loop. Verify logits against the Python reference before comparing warmed tokens/s. Match model file, threads, context and sampling with llama.cpp so conversion or workload differences do not masquerade as loop-overhead gains. - Quantize activation tiles to int8 with scales, reuse weight tiles across 64 prompt rows and accumulate integer products safely before rescaling. Use runtime ISA dispatch for VNNI/AMX and handle matrix tails. Compare prompt throughput and output error; tile packing and activation quantization must be included in end-to-end timing.
- Inspect the exact llama.cpp revision’s Q4_K layout handling, scale/minimum unpacking, cooperative loads and SIMD-group reductions, then port the relevant structure. Its native packed-format path avoids the simple kernel’s generic unpacking/repacking costs. Validate on actual Apple hardware; performance cannot be established by non-Metal tests.
- Predict next-layer expert needs on another stream, copy candidates early and fall back for actual routing misses. Current-layer output is only an approximation to the next router’s true input unless the architecture explicitly supports pre-gating. Measure accuracy of predictions, wasted bytes and exposed copy stalls while preserving exact final routes.
41. Multi-GPU
Check your understanding
- Independent output rows need only the common input. Output/down projections sum partial products over split input features, so their contributions must be reduced to reconstruct the full result.
- Once KV heads are fewer than ranks, heads are replicated across ranks. Further TP splits query work, not the size of each replicated KV head.
- Compare payload/link bandwidth with collective startup/round-trip latency. At roughly 1 MB, fast NVLink often makes startup a large fraction; PCIe transfer can be comparable or dominant. Topology and collective algorithm decide the crossover.
- A single batch traverses stages serially and pays transfers. Multiple microbatches/requests in flight fill the pipeline; this increases throughput, not automatically one request’s latency.
- When routing is sufficiently sparse/balanced and dispatch payloads are smaller than repeated projection reductions. Expert placement, duplicate top-k destinations and link topology matter.
- Cache affinity alone can overload the replica holding a popular prefix. A load term trades some reuse for bounded queues and available capacity.
- Export/import moves logical KV via the existing swap-copy contract, alongside tokens and progress. Seeded sampling must carry generator state, not just the original seed, to continue its random stream exactly.
Stretch exercises
- Determine each rank’s projection slices from the architecture before allocation, then read only those safetensors rows/columns and shared small tensors. Respect GQA head replication and shard before fusion. Measure peak host and device memory per rank, including transient slices, rather than only final shard size.
- Maintain pp disjoint request batches with stage-local buffers and placeholder/progress state, and schedule their stages concurrently. Preserve per-request order and stop/abort cleanup for in-flight work. Throughput improves by filling bubbles, while single-request latency still includes every stage and transfer.
- TP splits each layer’s work but pays frequent collectives; PP uses fewer boundary transfers but needs concurrent batches to fill stages. Over PCIe, PP can benefit throughput at sufficient concurrency while TP may reduce compute time enough to help latency. Measure both with the same model, memory constraints and offered load; topology determines the crossover.
- Shard large draft-model projections and KV according to their own head geometry if the draft runs across TP ranks. Medusa’s small heads can be replicated or evaluated on rank 0 if the required target hidden states are available there. Broadcast draft IDs/commit decisions and synchronize draft progress; account for hidden-state gathers and vocabulary-head sharding.
- Record an event after each layer produces stable KV, transfer on a separate stream and let the decode receiver wait on its layer’s completion. Preserve block ownership and metadata until copies finish. Overlap is bounded by remaining prefill compute and transfer bandwidth; report exposed tail-transfer time and total TTFT against a non-overlapped baseline.
42. More models
Check your understanding
FlatModelchecks for Q/K normalization by attribute presence. Defining only a Q identity would select the Q/K-normalization path and require a K module too; leaving both absent expresses the model’s actual contract.- Short-wavelength, fast-rotating pairs already experienced many full rotations during training. They keep local positional resolution; slow pairs are interpolated and the middle region blended.
- Its frequencies can change with sequence length. Cached keys rotated under earlier frequencies then disagree with recomputed keys at the longer length, violating prefix/state equivalence.
- Each head’s nonrotary key/value is an up-projection of the shared latent. Only the latent and shared rotary key are cached; distinct projections give distinct head scores/outputs.
- Rotation depends on token position and does not commute with a fixed up-projection. The rotary query/key term must remain explicit.
- Long prefill can benefit from decompressing once and attending with narrower ordinary heads. Decode avoids repeatedly decompressing the whole history by absorbing projections around latent attention.
- It bounds the number of expert groups a token can visit. When groups map to nodes, that bounds communication destinations and cross-node routing traffic.
Stretch exercises
- Implement every listed Gemma-specific norm, embedding scale, GeGLU and layer-type rule from the matching config, preserving operation order and window masks. Compare embeddings, intermediate residuals and fixed-input logits with the official model. Generic RMSNorm or SwiGLU substitutions can have plausible shapes but wrong values.
- Keep separate layer-group pools/tables and reclaim a window block only when no live query or in-flight chunk can read it. For purely windowed layers, storage scales with the window plus boundary/chunk allowance instead of full context; the ideal capacity gain is roughly
32768/windowfor those layers alone. Full-attention layers and shared overhead reduce the whole-engine gain. - Decompress latents into per-head K/V once for the prompt, apply the rotary term correctly and run ordinary causal attention. Compare with the absorbed latent path on identical prompts and include decompression/scratch costs. The crossover depends on head widths, prompt length and kernels; it is not a universal token count.
- Recognize FP8 tensors and their block scale metadata, construct
FP8BlockLinearwith the exact scale direction and preserve the architecture registry’s names/layouts. Validate selected dequantized blocks and layer outputs before full-model logits. Unsupported shapes or devices need an explicit reference fallback. - Tile many query heads against each shared 576-wide latent row, reuse loaded cache data and accumulate online softmax/output in FP32. Include the distinct rotary and absorbed value projections correctly. Compare against the generic one-KV-head reference across context lengths, head-group sizes and split counts; larger groups can increase register pressure.
43. Beyond text
Check your understanding
- Expanded patch rows, not the original marker count, consume decoder/KV positions. Length checks before expansion can underreserve memory and exceed the model limit.
- KV addresses and causal order follow flattened rows; M-RoPE follows spatial/time coordinates and resumes text after the maximum coordinate. Its continuation delta connects the two.
- The actual projected inputs and position recipe, including encoder/processor effects, axes, sections and continuation. The implementation hashes those values and snapshots them.
- A pooling rule selects/averages vectors; retrieval geometry must be learned with the model’s training objective and text conventions.
- Padding is not input content. Including it changes the direction and magnitude with batch padding length, making the same input depend on its neighbours.
- The base weight projection runs once over the whole token matrix. Shrink/expand corrections select A/B weights per adapter and scatter back to the original rows.
- Slots are reused and overwritten. A stable digest of effective weights prevents old KV from being reused after a different adapter occupies the same slot.
- It will resume using cached/recomputed activations for the same adapter. Replacing those weights during preemption would change the request’s model mid-history.
Stretch exercises
- Match the trained processor’s resize/crop/normalization, patch ordering and merger weights. Compare projected image rows, all rotary axes/continuation positions and fixed-input decoder logits with the official model. Textual answer plausibility alone cannot detect a wrong spatial layout.
- Key encoder outputs by image bytes, preprocessing settings and encoder revision; include merger identity if caching projected rather than raw encoder features. Use a byte-bounded LRU and preserve the position metadata needed by the decoder. Measure encoder-cache and decoder-prefix-cache hits separately, including host/device transfer costs.
- Group token rows by adapter, reuse the row permutation across layers and compute shrink/expand with tiled matmuls. Restore row order before operations that depend on sequence layout. Compare with SGMV over rank, rows per adapter and adapter count, since grouping overhead can dominate small groups.
- Queue embedding inputs up to a token/item limit or maximum wait, then run a padded/ragged batch with correct masks and per-input pooling. Bound queue size and propagate cancellation/errors. Measure throughput and p99 from admission, including the batching wait, at matched offered load.
- Capture fixed-capacity adapter buffers and rank buckets, zero unused rank lanes and update weights in place only after every old lease and replay user retires. Keep effective adapter identity in prefix keys. Test slot replacement, mixed/no-adapter rows and rank changes against eager execution to catch stale graph pointers or weights.
44. Benchmarking
Check your understanding
- Slower responses delay the client’s next submissions, lowering offered load exactly when the engine falls behind. Open-loop arrivals reveal queue growth.
- Dispatch time moves later. Scheduled arrival stays fixed so waiting for the client cap remains visible in TTFT/E2E.
- There is one observed arrival timestamp; all five token times inside the chunk are unknown. Report chunk gaps and average TPOT when its prerequisites hold.
- There are no post-first-token intervals and its denominator is zero. Missing data is clearer than claiming zero decode latency.
- Failure is part of service behavior. Excluding failed requests makes overload appear to meet its SLOs by simply refusing slow work.
- Once greedy outputs differ, histories differ too. Teacher forcing holds history fixed so logit/loss/KL differences measure the computation rather than cascading prompt changes.
- Quantization changes memory traffic, kernel dispatch and numerical quality. Report it as a separate speed-quality-memory experiment with the same source checkpoint.
- Not ready (503), even if health remains good: it is intentionally refusing new work while admitted requests drain.
Stretch exercises
- Run three open-loop trials at each rate and plot p99 TTFT, SLO goodput and queue length over time. The knee is where sustained offered load exceeds service capacity and queues begin growing; short bursts alone do not establish it. Keep failures in the population and record trial duration/warmup.
- Separate genuinely cold unique prefixes from warmed repeated prefixes, using equal total memory budgets and the same prompt/output lengths. Report cache hit tokens and compute time alongside TTFT/throughput. Cache-disabled or matched-hit controls help attribute gains to reuse versus kernels.
- Send equivalent rendered prompts, tokenizer settings, image preprocessing and generation budgets to every engine. Report vision-encoder time, decoder prefill and end-to-end TTFT separately. Different templates or image resolutions change the workload and invalidate a direct speed comparison.
- Kill the core and disconnect the proxy during controlled bursts, then measure time until every affected client fails or completes. Verify readiness drops, admission stops appropriately and KV/adapter/queue resources are released. Bound client timeouts so hung streams cannot silently disappear from the results.
- Use one source checkpoint and identical teacher-forced histories/task protocol, then report quality, peak/resident memory, latency and throughput for INT4/FP8. Plot a trade-off curve with failed requests and SLOs included. A faster variant with worse task performance is a different operating point, not an unconditional winner.
D. Notation and glossary
Notation
| symbol | meaning |
|---|---|
| $B$ | batch size (sequences processed together) |
| $T$ | number of new tokens in a forward call (query length) |
| $S$ | number of keys visible (cached plus new) |
| $D$ or $d_\text{model}$ | hidden size (width of the residual stream) |
| $H$, $H_{kv}$ | query heads, key/value heads |
| $d$ or $D_h$ | head dimension |
| $V$ | vocabulary size |
| $L$ | number of layers |
| $E$, $k$ | number of experts, experts per token (Chapter 27); draft length (Chapter 26) |
| $r$, $\alpha$ | LoRA rank and scale (Chapter 22) |
| $W$ | a weight matrix, stored [out_features, in_features] as in nn.Linear |
| $\sigma$ | the logistic sigmoid $1/(1+e^{-x})$ |
| $\odot$ | elementwise product |
Tensor shapes are written in brackets, [B, T, D], with the fastest-varying axis last.
Glossary
Activation. Any intermediate value computed by a model during a forward pass; also the nonlinear functions (GELU, SiLU) applied in MLPs.
AdamW. The standard optimizer for transformers: per-parameter adaptive step sizes from running averages of the gradient and its square, with decoupled weight decay (Chapter 3).
Arithmetic intensity. FLOPs performed per byte moved from memory. Below the hardware’s ridge point an operation is memory-bound; above it, compute-bound (Chapter 10).
Attention. The operation that lets each position read from other positions: softmax-normalized query-key scores weight a sum of values (Chapter 5).
Autograd. Automatic differentiation by recording operations in a graph and applying the chain rule backwards (Chapter 3).
Bandwidth. Bytes per second a memory system delivers. Decode speed at small batch is bandwidth ÷ bytes per token (Chapters 1, 10).
Bank conflict. Several threads of a warp accessing different addresses in the same shared-memory bank, which serializes the accesses (Chapter 12).
BF16 (bfloat16). 16-bit floating point with FP32’s 8-bit exponent and a 7-bit mantissa: FP32’s range, about 3 significant digits (Chapter 13).
Block (CUDA). A group of threads that run on one SM, can share memory and synchronize with __syncthreads() (Chapter 11). In paged attention, a fixed-size chunk of the KV cache (Chapter 25).
BPE (byte-pair encoding). A tokenizer that starts from bytes and repeatedly merges the most frequent adjacent pair into a new token (Chapter 4).
Causal mask. The rule that position $t$ may attend only to positions $\le t$. This book implements it as key_pos <= query_pos, which also covers caches and chunks (Chapters 5, 16).
Chat template. The model-specific format that wraps conversation turns in special tokens (Chapter 4).
Chunked prefill. Splitting a long prompt into pieces processed over several steps, to bound per-step latency (Chapter 24).
Coalescing. Combining a warp’s memory accesses into few wide transactions when consecutive threads read consecutive addresses (Chapter 11).
Continuous batching. Scheduling at every decode step: finished requests leave and new ones join the running batch (Chapter 24).
Copy-on-write. Sharing a cache block between sequences until one writes to it, then copying (Chapter 25).
CUDA graph. A recorded sequence of GPU work replayed with one launch (Chapter 19).
Decode. Generating one token per step after prefill; memory-bound at small batch (Chapters 1, 16).
Delta rule. A state update that writes the error between a target value and what the memory currently returns for the key (Chapter 28).
Dequantization. Converting integer codes back to floating point with their scales (Chapter 20).
Embedding. A learned table mapping token IDs to vectors (Chapter 4).
Expert (MoE). One of several MLPs in a mixture-of-experts layer; a router chooses which experts process each token (Chapter 27).
FlashAttention. Exact attention computed tile by tile with an online softmax, never storing the full score matrix (Chapter 15).
FLOP. A floating-point operation; a multiply-add counts as 2.
Fusion. Combining several operations into one kernel so intermediates stay in registers or shared memory (Chapter 14).
Gated DeltaNet. A linear-attention layer with a decaying, delta-rule-updated matrix state (Chapter 28).
GQA (grouped-query attention). Several query heads share each key/value head, shrinking the KV cache (Chapters 5, 16).
Greedy decoding. Always choosing the highest-probability token (Chapter 8).
Hook. A function PyTorch runs after (or before) a module’s forward pass, able to read or replace its output (Chapter 23).
Hyper-connections. Several parallel residual streams with learned read and write weights per sublayer (Chapter 30).
KV cache. Stored keys and values of past positions, so each decode step computes only the new token (Chapter 16).
Kernel. A function that runs on the GPU, launched over a grid of threads or programs (Chapter 11).
Launch overhead. CPU time to issue a kernel; dominates when kernels are tiny (Chapter 19).
Layer normalization / RMSNorm. Rescaling each position’s vector to a fixed size; RMSNorm omits mean subtraction (Chapters 6, 17).
Logits. Unnormalized scores over the vocabulary; softmax turns them into probabilities (Chapter 1).
Logit lens. Applying the final norm and head to intermediate residual streams to see what each layer predicts (Chapter 23).
LoRA. Fine-tuning through a trainable low-rank update $\frac{\alpha}{r}BA$ added to frozen weights (Chapter 22).
Memory-bound / compute-bound. Limited by bytes moved, or by arithmetic throughput (Chapter 10).
MTP (multi-token prediction). Extra heads trained to predict tokens beyond the next; usable as speculative drafts (Chapters 26, 30).
N-gram memory. A hashed lookup table of embeddings for recent token pairs and triples, injected into the residual streams (Chapter 30).
NF4. A 4-bit data type with levels at quantiles of a normal distribution, used for QLoRA bases (Chapters 20, 22).
Occupancy. The fraction of an SM’s thread slots in use; higher occupancy hides memory latency (Chapter 12).
Online softmax. Computing softmax-weighted sums in one pass with a running maximum and rescaling (Chapter 15).
Paged attention. A KV cache in fixed-size blocks from a shared pool, addressed through per-sequence block tables (Chapter 25).
Perplexity. $e^{\text{loss}}$: the effective number of equally likely choices the model is uncertain between (Chapter 7).
Prefill. Processing the prompt in one forward pass to fill the cache and produce the first token’s logits; compute-bound (Chapters 1, 16).
Prefix caching. Reusing the KV cache of a shared prompt prefix across requests (Chapter 25).
QLoRA. LoRA on a frozen 4-bit (NF4) base model (Chapter 22).
QSA (Qwen Sparse Attention). Flash-Next’s block indexer that selects which keys full attention reads (Chapter 29).
Quantization. Storing values with fewer bits as integer codes plus scales (Chapter 20).
Residual stream. The running hidden state each layer reads from and adds to (Chapters 6, 23).
Roofline. The plot of attainable FLOP/s against arithmetic intensity: a bandwidth slope meeting a compute ceiling (Chapter 10).
RoPE (rotary position embedding). Encoding position by rotating query and key feature pairs by angles proportional to position (Chapter 17).
Safetensors. A checkpoint format with a JSON header and raw tensor bytes, safe to load (Chapter 9).
Sampling. Choosing the next token at random from a (possibly filtered) distribution: temperature, top-k, top-p, min-p (Chapter 8).
Shared memory. Fast on-chip memory shared by the threads of a block (Chapter 11).
SM (streaming multiprocessor). A GPU core: runs warps, holds registers and shared memory (Chapter 10).
Speculative decoding. A cheap draft proposes tokens and the target verifies them in one pass, with an exact acceptance rule (Chapter 26).
SwiGLU. The gated MLP used by Qwen and Llama: $W_\text{down}(\operatorname{SiLU}(xW_\text{gate}) \odot xW_\text{up})$ (Chapter 17).
Teacher forcing. Training on the true previous tokens rather than the model’s own outputs (Chapter 7).
Tensor core. Hardware that multiplies small matrix tiles per instruction, at much higher throughput than ordinary arithmetic (Chapter 12).
Tiling. Splitting a computation into blocks that fit in fast memory, to reuse each loaded value many times (Chapter 12).
Token. One unit of the tokenizer’s vocabulary; the model sees only token IDs (Chapter 4).
TTFT / TPOT. Time to first token; time per output token (Chapters 18, 24).
Warp. 32 threads that execute together on NVIDIA GPUs (Chapter 11).
Weight tying. Using the token embedding matrix as the output head (Chapters 6, 9).
Serving, formats and model coverage
Adapter slot / lease. Resident LoRA storage and a request’s right to keep its contents unchanged until completion; cache identity comes from effective weights, not the slot (Chapter 43).
AWQ / GPTQ. Activation-aware column scaling and curvature-aware weight-error compensation, respectively, with checkpoint-specific packed layouts (Chapter 39).
BGMV / SGMV. Batched or segmented gathered matrix-vector multiplication: low-rank shrink/expand operations selected by adapter per row or row segment (Chapter 43).
Chunked prefill. Computing a prompt in bounded token chunks mixed with other requests’ decode work (Chapters 24, 31).
Disaggregated prefill/decode. Separate engines for prompt processing and generation, with KV and request state transferred between them (Chapter 41).
E2E / ITL. End-to-end request latency and inter-token latency. Streamed chunk gaps are reported separately when token arrival times are unknown (Chapter 44).
Expert / tensor / pipeline / context parallelism. Splitting experts, projection matrices, layer stages or context positions across devices, with different communication patterns (Chapter 41).
FP8 block / MXFP4 / NVFP4. Floating-point quantization with block scales; the formats differ in value bits, scale encoding and block size (Chapter 39).
GGUF / GGML quants. A model file containing metadata and tensors, and the tensor block formats commonly carried inside it (Chapter 38).
Goodput / SLO. Successful requests per second meeting all configured service-level objectives; failures remain in attainment’s denominator (Chapter 44).
Guided decoding. Masking tokens according to a byte-level constraint automaton, enforcing a supported regex or schema language (Chapter 34).
Hash-chained prefix key. A fixed-size cryptographic digest of preceding blocks, current tokens and activation-changing inputs such as images and adapters (Chapters 31, 43).
MLA. Multi-head latent attention: cached compressed latent vectors plus shared rotary keys, with head-specific projections absorbed into queries and outputs during decode (Chapter 42).
M-RoPE. Rotary frequency pairs assigned to time, height and width coordinates; text uses equal coordinates on every axis (Chapter 43).
Open-loop arrivals. Request submission times chosen independently of earlier completions; used to expose queueing and overload (Chapter 44).
Pooling / cross-encoder reranking. Producing one vector from token hidden states, or scoring jointly encoded query-document pairs with a trained head (Chapter 43).
Preemption. Releasing a request’s device KV by dropping it for recomputation or saving it in host memory for later restoration (Chapter 31).
Ragged batch. One flattened token matrix whose sequence boundaries, lengths and cache slots are described by metadata (Chapter 31).
Readiness / drain. Whether a process can accept new work, and the shutdown phase that stops admission while existing work completes under a deadline (Chapter 44).
Split-KV. Partitioning long-context attention across programs, then merging outputs using each partition’s log-sum-exp normalizer (Chapter 32).
Trellis quantization. A finite-state sequence of codes optimized jointly over a weight vector, rather than independent scalar rounding (Chapter 39).
E. Sources, credits and further reading
This book stands on three main sources. Each chapter’s Going deeper section points to the exact sections; this appendix gives the overall map, so you can read the sources alongside the book.
- BALLM: Sebastian Raschka, Build a Large Language Model (From Scratch), Manning, 2024. The modeling foundations: tokenization, attention, GPT, pretraining, fine-tuning. Companion code: github.com/rasbt/LLMs-from-scratch (Apache-2.0).
- PMPP: Wen-mei W. Hwu, David B. Kirk and Izzat El Hajj, Programming Massively Parallel Processors: A Hands-on Approach, 5th edition, Morgan Kaufmann. The GPU foundations: CUDA, memory, tiling, reductions, scan, and (Chapter 20) attention and KV caching.
- GPU Mode: the lecture series and community at github.com/gpu-mode/lectures and youtube.com/@GPUMODE. Kernels, profiling and inference systems, taught by the people who build them.
Chapter map
| book chapter | BALLM | PMPP | GPU Mode |
|---|---|---|---|
| 1. The big picture | Ch. 1 | §20.4, §20.6 | L1 |
| 2. Tensors | App. A §§A.1-A.3 | ||
| 3. How networks learn | App. A §§A.3-A.7 | L6 | |
| 4. Text as numbers | Ch. 2 | ||
| 5. Attention | Ch. 3 | §20.2 | |
| 6. The transformer block and GPT | Ch. 4 | §20.1 | |
| 7. Training | Ch. 5 §§5.1-5.2 | ||
| 8. Generation | Ch. 5 §5.3 | ||
| 9. Real weights | Ch. 5 §§5.4-5.5 | ||
| 10. GPU performance | Ch. 1, Ch. 4, Ch. 5 (roofline) | L1, L8, L16 | |
| 11. CUDA | Ch. 2-4 | L2, L3, L4, L5 | |
| 12. Fast matrix multiplication | Ch. 5-6, Ch. 15 | L5, L8, L23 | |
| 13. Reductions and numerics | Ch. 10, App. A | L9 | |
| 14. Triton | §15.8 | L14, L18, L28, L29 | |
| 15. FlashAttention | §20.5 | L12, L13, L36 | |
| 16. KV cache | §20.4, §§20.6-20.7 | ||
| 17. Qwen3 | ch05/11_qwen3 (repository) | ||
| 18. Engine v1 | Ch. 5 §5.5 | §20.6 | |
| 19. Fast decode | L1, L6, L16, L35 | ||
| 20. Quantization | App. A | L7, L30, L33 | |
| 21. Fine-tuning | Ch. 6, Ch. 7 | ||
| 22. LoRA and QLoRA | App. E | L32 | |
| 23. Inside the model | |||
| 24. Batching | §20.6 | L35 | |
| 25. Paged attention | §20.7 | L35, L40 | |
| 26. Speculative decoding | §20.6 | L22 | |
| 27. Mixture of experts | L11 | ||
| 28. Linear attention | Ch. 11 (scan) | L20, L21, L24 | |
| 29. Sparse attention | |||
| 30. Capstone | |||
| 31. Engine v2 | §§20.6-20.7 | L35, L40 | |
| 32. Attention backends | §20.5 | L12, L13, L36 | |
| 33. Fusion, graphs, async scheduling | Ch. 4, Ch. 15 | L1, L6, L35 | |
| 34. Sampling and structured output | Ch. 5 §5.3 | ||
| 35. Text boundary | Ch. 2 | ||
| 36. HTTP server | §20.6 | L35 | |
| 37. Batched speculation | §20.6 | L22 | |
| 38. GGUF and GGML | App. A | L7, L30 | |
| 39. Quantized checkpoints | App. A | L7, L30, L33 | |
| 40. Hardware and offload | Ch. 4, Ch. 5 | L8, L16 | |
| 41. Multi-GPU | L17 | ||
| 42. Model coverage and MLA | |||
| 43. Images, pooling and multi-LoRA | App. E (LoRA) | L32 | |
| 44. Benchmarking and hardening | Ch. 5 (measurement) | L1, L16 | |
| 45. What’s next |
Papers and documentation by topic
Transformers and language models. Vaswani et al., Attention Is All You Need (2017); Radford et al., Language Models are Unsupervised Multitask Learners (GPT-2, 2019); the Qwen3 technical report (2025) and the Qwen3-Next, Qwen3.5 and Qwen3.8-Flash-Next model cards; Su et al., RoFormer (RoPE, 2021); Shazeer, GLU Variants Improve Transformer (2020).
Kernels and performance. Williams, Waterman and Patterson, Roofline (2009); Tillet et al., Triton (2019); Dao et al., FlashAttention 1-3 (2022-2024); Milakov and Gimelshein, Online normalizer calculation for softmax (2018); the CUDA C++ Programming Guide and Best Practices Guide; the Triton tutorials.
Inference systems. Pope et al., Efficiently Scaling Transformer Inference (2022); Yu et al., Orca (2022); Kwon et al., PagedAttention (2023); Zheng et al., SGLang (2024); Agrawal et al., Sarathi-Serve (2024); Zhong et al., DistServe (2024); Leviathan et al. and Chen et al. on speculative decoding (2023); the PyTorch gpt-fast blog posts.
Quantization and adaptation. Dettmers et al., LLM.int8() (2022) and QLoRA (2023); Frantar et al., GPTQ (2022); Lin et al., AWQ (2023); Xiao et al., SmoothQuant (2022); Hu et al., LoRA (2021).
Architectures. Shazeer et al. (2017) and Fedus et al. (2021) on mixtures of experts; DeepSeek-V3 technical report; Katharopoulos et al. (2020), Schlag et al. (2021), Yang et al. (2023-2024) on linear attention and Gated DeltaNet; Gu and Dao on Mamba (2023-2024); Zhu et al., Hyper-Connections (2024); DeepSeek’s Native Sparse Attention and Engram (2025-2026).
Interpretability. Elhage et al., A Mathematical Framework for Transformer Circuits (2021); nostalgebraist, the logit lens (2020); Turner et al., Activation Addition (2023); Arditi et al., Refusal Is Mediated by a Single Direction (2024).
Credits
- The Verdict by Edith Wharton (1908), public domain, used as BALLM uses it for the training chapters.
- The 1,100-example instruction dataset (
data/instruction-data.json) comes from BALLM Chapter 7’s companion repository (Apache-2.0). - The SMS Spam Collection (Almeida and Hidalgo, UCI Machine Learning Repository, CC BY 4.0) is referenced for Chapter 21 and downloaded by the reader.
- Kernel structure follows the official Triton tutorials (fused softmax, matmul with grouped ordering, fused attention) and PMPP’s CUDA examples, adapted and simplified.
- The Flash-Next implementation was written from the Transformers
qwen4_expreference implementation (Apache-2.0) and is tested against it. - Diagrams and interactive explorers are original to this book.
Part VIII primary sources
The foundations above remain useful, but Part VIII also follows the code and papers of the systems it compares. Chapter-level links give the particular mechanism; the following groups are the starting points.
- Scheduling and serving: vLLM, SGLang, PagedAttention, Sarathi-Serve and DistServe. Chapters 31-37 implement the shared mechanisms with smaller interfaces.
- Attention: FlashAttention, FlashInfer and FlashMLA. Their architecture-specific kernels are comparison/reference targets, not claims that the teaching kernels reach the same performance.
- Formats and hardware: llama.cpp/GGML, GGUF specification, GPTQ, AWQ, QTIP and ExLlamaV3. Chapters 38-40 distinguish file interchange, repacked runtime layouts and native kernels.
- More models: Llama 3, YaRN, DeepSeek-V2 and DeepSeek-V3, together with their official model/config implementations.
- Images and retrieval: ViT, LLaVA, Qwen2-VL and Sentence-BERT. Chapter 43 implements a small vision path and pooling contracts; it does not validate trained VLM/retrieval quality.
- Many adapters: Punica and S-LoRA, for gathered low-rank kernels and adapter-aware memory management.
- Evaluation and reproducible launches: GSM8K, lm-evaluation-harness, vLLM serve, SGLang arguments, llama-server, TabbyAPI configuration and TensorRT-LLM serve. Chapter 44 provides launch recipes and a blank competitor-results template; a documentation check is not an executed benchmark.
F. Validation log
A book that teaches “test against a trusted reference” should say what it tested itself. This appendix records how the code and the numbers in this book were checked, in which environment, and what could not be checked there.
Policy
- Every program output shown in the book was produced by running the command shown, unless the text explicitly calls it illustrative or quotes another source.
- Every engine milestone is checked by tests that the reference implementation (
izh/) passes and that the generated skeleton (engine/) fails with aTODO(Chapter N)message. - Architectures are checked against independent reference implementations (Hugging Face Transformers) with random weights in FP32 before any real checkpoint is involved.
Completion environment (2026-10-05)
The original book validation ran on an x86-64 CPU-only environment. This completion pass ran on a DGX Spark with an available CUDA GPU; the results below separate these records.
| component | version / coverage |
|---|---|
| OS / CPU | Ubuntu 24.04.5 LTS, Linux aarch64 |
| GPU / driver | NVIDIA GB10, driver 580.178.04 |
| Python | 3.14.7 |
| PyTorch | 2.14.1+cu132, CUDA runtime 13.2 |
| CUDA compiler | nvcc 13.0.88 (/usr/local/cuda) |
| Triton | 3.8.0; native CUDA and a separate CPU-interpreter run |
| Transformers / tokenizers | 5.18.0 / 0.23.2 |
| FastAPI / httpx | 0.142.2 / 0.28.1 |
| GGUF reference package | 0.19.0 |
| Rust | rustc / cargo 1.99.0 |
| C++ / mdBook | g++ 13.3 / mdBook 0.5.4 |
No pretrained checkpoint was downloaded during this pass. The original environment reported blocked Hugging Face Hub access; current correctness checks use locally generated checkpoints, fixtures and random weights, rather than claiming real-checkpoint validation.
Completion results
| check | result |
|---|---|
IZH_IMPL=izh OMP_NUM_THREADS=2 pytest -rs (all chapters, CUDA available) | 374 passed, 2 skipped; skipped: flash_attn package unavailable and Metal requires an Apple GPU |
Chapters 38, 39, 43 and 44 with CUDA_VISIBLE_DEVICES='' TRITON_INTERPRET=1 IZH_DEVICE=cpu | 63 passed, 2 skipped; only the two explicit CUDA LoRA-kernel cases skip |
| generated skeleton spot checks | Chapter 43 M-RoPE and Chapter 44 arrival tests fail at their intended TODO(Chapter 43/44) functions |
tools/make_engine.py --check | all 87 generated package files match the reference/stub generator |
cargo test --release --offline | 18 unit + 2 checkpoint-fixture tests passed |
C++ Release build and izh all | 14 demos passed; CPU extension and native demo selected scalar kernels on this ARM build |
run.py features, run.py benchmark | ran on CPU; image/adapter/pooling plumbing, worked SLO arithmetic and independent greedy parity succeed |
| real loopback HTTP/SSE benchmark smoke | Poisson and burst runs, 8 requests each, 6 tokens/request, all 16 succeed; 48 output tokens per run; no competitor or production-performance claim |
| new HTTP contracts | mixed image + named LoRA matches offline generation; embedding float/base64, reranking, readiness, byte limits, deadline abort and slow-consumer cleanup pass |
| GGUF interoperability | legacy/K quant decoding, packed-to-affine repacking, reader/writer exchange and local model round trips checked with the upstream gguf package |
tools/check_book.py, mdbook build | 53 pages, 30 widget kinds; includes/links/anchors and HTML build pass |
prepare_book.py and lab sync | code download regenerated; source code artifacts mirrored to lab/inference-zero-to-hero-code, preserving local environment files |
The benchmark smoke uses a tiny random CPU model and a character-level test tokenizer. Its timings prove that requests, SSE parsing, usage, timestamps and shutdown work end to end. They are not representative throughput measurements and are not entered into Chapter 44’s competitor table. The quality smoke fixtures likewise do not establish trained retrieval, vision or GSM8K accuracy.
Earlier CPU-only records
The previous log recorded 144 passing and 13 GPU-skipped tests before Part VIII, 17 Rust unit tests plus 2 fixtures, and 13 C++ demos. It also recorded a full pre-Part-VIII skeleton run (125 expected failures, 5 fixture errors, 14 provided-code passes). Those historical counts are not the current suite size. The completion run checks generated-file equality and the two new chapter TODO boundaries rather than claiming another full skeleton run.
Parity with reference implementations (FP32, random weights)
| model / component | reference | max abs logit difference |
|---|---|---|
| GPT-2 | GPT2LMHeadModel | 1.8e-7 |
| Qwen3 (tied and untied heads, randomized norm weights) | Qwen3ForCausalLM | 3.0e-7 |
| Qwen3-MoE (with and without renormalization, 6 experts) | Qwen3MoeForCausalLM | 1.5e-7 |
| Gated DeltaNet layer | Qwen4ExpTextGatedDeltaNet | within the test’s 1e-5 tolerance |
| Qwen3.8-Flash-Next, 8 layers, all components, EOS mid-sequence | Qwen4ExpForCausalLM | 1.5e-8 (logits of magnitude 0.12) |
| Flash-Next with memory-mapped n-gram tables | own FP32 load | 0 (bit-identical) |
Bugs found by these checks
Recorded because each one is a lesson, and each is now covered by a test:
- Chunked Gated DeltaNet kept the wrong triangle of its system matrix. Chunk size 1 passed; sizes 5, 8 and 64 against the recurrent form caught it (Chapter 28).
- QSA top-k tie-breaking: masking invisible blocks with $-\infty$ in a longer vector broke ReLU-zero ties differently from the reference. Fixed by running
topkover exactly the visible blocks (Chapter 29). - Qwen3-MoE config: Transformers 5 writes
num_local_experts; the loader silently used its default of 8 experts and passed tests that happened to use 8. Fixed to accept both names and refuse missing fields; tests now use non-default sizes (Chapter 27). - safetensors:
torch.frombufferrejects zero-length tensors; empty tensors are now created directly (Chapter 9). - Full fine-tuning in BF16 with AdamW at learning rate 1e-5 loses most updates to rounding;
finetune.pynow keeps FP32 master weights with BF16 autocast for full fine-tuning (Chapters 13, 21). - Test collection: a test module that built models at import time turned one missing function into a collection error for the whole suite; models are now built inside tests (Chapter 16).
What was not validated here
Passing a teaching test establishes the tested contract and shapes, not every deployment or feature combination. The following boundaries remain explicit:
| path | completion status |
|---|---|
| CUDA kernels and Triton kernels | tiny correctness cases ran on GB10; broad architecture tuning and representative performance unmeasured |
| CUDA graphs and compile (Chapters 19, 33) | available GPU correctness tests passed; feature batches still take eager fallback, and production graph/memory tuning is unmeasured |
flash_attn adapter (Chapter 32) | skipped: package unavailable |
profile_num_blocks automatic GPU memory sizing | not independently exercised; tests specify pool sizes |
| GPU/CPU split and offload (Chapter 40) | tiny split-model correctness passed; representative transfer overlap and memory pressure unmeasured |
| NCCL, multiple physical GPUs, RDMA KV transfer | not run; Chapter 41 uses gloo CPU processes and local transfer contracts |
| Metal, ROCm, Vulkan, XPU | no target hardware validation; Metal test skipped; complete Vulkan/ROCm backends remain beyond this implementation |
| AVX2, AVX-512 VNNI, NEON dot-product kernels | not exercised by this ARM build; scalar CPU path passed. The original x86 record does not establish current cross-ISA coverage |
| real HF/GGUF checkpoints | no pretrained weights downloaded; architecture loaders tested with locally generated checkpoints and tiny fixtures |
| external GPTQ/AWQ/FP8/FP4 checkpoint tools, native EXL2/EXL3 | local synthetic layouts/reference equations tested; broad vendor/checkpoint interoperability not established; trellis lab is not native EXL3 loading |
| trained VLMs, embedding/retrieval quality, production rerankers | architecture and HTTP plumbing only; exact trained processors, towers, pair conventions and weights need independent validation |
| LoRA combinations | dense resident slots, mixed batches, cache keys and plain PEFT-layout loading tested; quantized/distributed/expert adapters, graph feature buffers and adapter-aware speculation are not validated |
| benchmark competitors | no vLLM, SGLang, llama.cpp, ExLlamaV3/TabbyAPI or TensorRT-LLM run; launch recipes checked against upstream docs, table intentionally blank |
| GSM8K/task accuracy | extractor and runner tested with fixtures; no pretrained-model benchmark score claimed |
| sustained load, device/process failures, orchestrated drain | unit/local HTTP checks only; multi-host recovery, rollout routing and long-duration SLO evidence require deployment tests |
| Chapter 1 probabilities, Chapter 18 chat transcript, expected GPU tables | retain their illustrative/source-attributed labels |
| fine-tuning/editing and bitsandbytes workflows | original tiny-checkpoint results retained; not rerun in this pass on real models |
| Flash-Next memory/deployment estimates | arithmetic, not measured serving of its pretrained weights |
The original CUDA compile-only and CPU-interpreter records remain historical evidence; current GPU correctness tests add coverage without turning estimates into measurements.
Reproducing
From the code/ directory:
IZH_IMPL=izh OMP_NUM_THREADS=2 pytest -rs # the reference, including available GPU paths
pytest # your engine: red until you implement it
(cd rust && cargo test --release)
cmake -S cpp -B build/cpp && cmake --build build/cpp -j && build/cpp/izh all
CUDA_VISIBLE_DEVICES='' TRITON_INTERPRET=1 IZH_DEVICE=cpu IZH_IMPL=izh pytest \
tests/test_ch38_gguf.py tests/test_ch39_quantized_checkpoints.py \
tests/test_ch43_multimodal_lora_embeddings.py tests/test_ch44_benchmarking.py
python run.py features --device cpu
python run.py benchmark --device cpu