primer.ml.inference

Inference: what happens, and what it costs, when a model generates

Run: python -m primer.ml.inference

New to the notation? primer.notation explains every symbol used here from zero. This lesson builds on attention from primer.ml.attention and the transformer from primer.ml.transformer.

Level 1: The practitioner's guide

In one sentence. Inference is everything that happens when a trained model answers a request, and its cost and speed are governed by one fact: each generated token reads every weight from memory, so the levers that matter (caching, batching, quantization, speculation, sampling settings) are all ways to read less or to reuse each read.

When you need it. You need this lesson the day a model leaves the notebook: when you set temperature and top_p on a request, when someone asks why the first token takes two seconds, when a GPU bill arrives, or when you decide whether to serve a model yourself. The tell is a latency or cost question you can only answer by guessing. The numbers here replace the guess. This lesson's model of an 8-billion-parameter network at 16 bits on an H100-class GPU (3.35 TB/s of memory bandwidth, about 10¹⁵ operations per second, the constants in this module) reads a 1,000-token prompt in about 16 ms but then produces at most about 209 tokens per second for a single user, because every token needs all 16 GB of weights read again (4.78 ms each). A lone request uses well under 1% of the GPU's arithmetic. You do not need this lesson for a prototype at ten requests a day; you need it before the first load test.

Your options. The levers a practitioner can pull, roughly from the cheapest to the most involved:

Lever What it does What it buys you What it costs Where it lives
Sampling settings Temperature reshapes the token probabilities; top-k and top-p cut the unlikely tail Control over variety: T = 0 for extraction and tool calls, higher for ideas Nothing in compute; wrong settings cost quality The request
Prompt caching Reuses the prefill of a shared prefix (system prompt, tools, documents) across calls This lesson's 100 calls with a 10,000-token prefix: \$0.48 instead of \$3.15, an 85% saving, and a faster first token A cache write at about 1.25× the input price on the first call (Anthropic's prompt caching docs) The provider, or your server's prefix cache
Quantization Stores weights at 8 or 4 bits instead of 16 A 70B model in 35 GB instead of 140 GB, and faster memory-bound decode 8-bit is nearly lossless; naive 4-bit loses small weights (the lesson's 0.02 rounds to 0) unless a smarter method such as GPTQ or AWQ is used The model files and the server
A model with grouped-query attention Fewer key/value heads means a smaller KV cache per token 8 KV heads instead of 32 fit four times as many long conversations per GPU A model choice made at training time; you can only pick a model that has it The model architecture
Continuous batching Seats a new request the moment any slot frees, instead of waiting for the whole batch In this lesson's 32-request simulation, 82% slot utilisation instead of 60%, in fewer steps Nothing beyond a server that does it (they all do now) The inference server
Speculative decoding A small model drafts several tokens; the big one verifies them in one pass With an 80% acceptance rate and 4 drafts, 3.36 tokens per big-model pass instead of 1, with the same output distribution A draft model to run, and gains that shrink when the draft guesses badly or the server is already batch-saturated The inference server
A hosted API Someone else runs all of the above No GPUs to size, caching and batching done for you A per-token price, and less control over settings and residency The provider

How to choose. Start from the symptom.

  • Slow first token: the prompt is long. Trim it, or put its stable part first and let prompt caching skip its prefill.
  • Slow streaming: decode is memory-bound. Quantize, batch more requests together, or add speculative decoding.
  • Running out of GPU memory as traffic grows: it is the KV cache, not the weights. Do the arithmetic (weights plus cache per token times context times concurrent requests) before renting a bigger card; prefer a model with grouped-query attention and cap the context you allow.
  • Answers that vary when you want them stable: temperature 0 and a validator, not a hope. Answers that all sound the same when you want range: raise the temperature and let top-p keep the tail sane.
  • Deciding whether to self-host: only when volume, privacy or a model the APIs do not offer justifies owning the batching and memory problems above.
  • Whatever you pick, measure time to first token and tokens per second separately. They are set by different phases and fixed by different levers.

What it costs. Money follows tokens, and tokens follow decode. Prefill of 1,000 tokens costs 16 ms of a GPU's full compute; each output token costs a full read of the weights, so output tokens are the expensive ones, and providers price them that way. Memory sets capacity: on this lesson's Llama-3-8B-shaped model each token of context holds about 128 KB of keys and values, a 32,000-token conversation holds 4.2 GB, and an 80 GB GPU with 16 GB of weights fits 15 such conversations at once. Batching is what makes serving economical: one read of the weights serves every request in the batch, which is why the roofline figure in Level 2 shows a batch of 64 reaching 21% of the GPU's arithmetic where a single user reaches 0.3%. Prompt caching costs a little on the first call and saves most of the input bill after it; Anthropic prices a five-minute cache write at 1.25× and a read at 0.1× the input price, with a one-hour write at 2× (its prompt caching docs). Quantization costs a little quality for a large memory saving. Speculative decoding costs a second model and a more complex server.

What breaks.

  • Temperature 0 still varies. Floating-point addition depends on order, and GPU kernels change their order with the batch they land in. One published measurement found 80 distinct completions in 1,000 runs at temperature 0, identical for the first 102 tokens and then diverging (Thinking Machines, Defeating Nondeterminism in LLM Inference). Treat determinism as reduced, not guaranteed, and validate outputs.
  • A timestamp at the top of the prompt silently disables prompt caching, because a cache matches only up to the first differing token. Stable content first, volatile content last; on Anthropic's API the order is tools, then system, then messages.
  • Short prompts are not cached. Providers set a minimum cacheable length (Anthropic's is between 512 and 4,096 tokens depending on the model) and return no error below it; the bill just does not fall.
  • A long-context feature exhausts memory. Doubling the context you allow doubles the cache per request and halves the requests that fit.
  • Naive 4-bit quantization erases small weights. Use a method that compensates (GPTQ, AWQ) and check quality on your own evaluation, not on the model card.
  • Top-k with a fixed k cuts too much when the model is unsure and too little when it is confident; top-p adapts, which is why it is the usual default.
  • Speculation that guesses badly costs more than it saves: the big model's pass still runs, and every rejected draft is wasted work.

In the wild. vLLM's documentation lists continuous batching, chunked prefill, prefix caching, PagedAttention for KV memory, speculative decoding (n-gram and EAGLE drafts among others) and quantization from FP8 to INT4, GPTQ and AWQ, behind an OpenAI-compatible API. SGLang offers the same set with RadixAttention for prefix caching; NVIDIA's TensorRT-LLM does it with custom kernels and FP8 and FP4 formats on NVIDIA GPUs; llama.cpp runs quantized models from 1.5-bit to 8-bit on CPUs and Apple silicon. Hugging Face's Transformers exposes the sampling knobs (greedy by default, sampling with do_sample, beam search with num_beams). Hosted APIs expose prompt caching explicitly, with Anthropic's linked in Further reading. The ideas come from the papers at the end of this lesson: speculative decoding (Leviathan, Kalman and Matias), PagedAttention (Kwon et al.), FlashAttention (Dao et al.), LLM.int8() and GPTQ for quantization, and nucleus sampling (Holtzman et al.).

Go deeper. Level 2 builds each lever from nothing: the roofline that explains why decode is memory-bound, a tiny decoder with and without a KV cache whose outputs match to ten decimal places, the temperature and top-p arithmetic, the accept-or-reject rule that makes speculative decoding exact, a quantizer in five lines, and simulations of both batching policies. If you only needed to size a deployment or set a request, you are done.

Level 2: How it works, from scratch.

Level 2: How it works, from scratch

Training happens once; inference happens every time anyone uses the model, so this is where the money goes. Generating text has two very different phases, one memory trick that makes it affordable (the KV cache), a few knobs that decide which token comes out (sampling), and a toolbox of speedups: quantization, speculative decoding, continuous batching and prompt caching. Every one of them follows from one fact:

Generating one token requires reading every weight of the model from memory, and memory is much slower than arithmetic.

1. Two phases: prefill and decode

Everyday picture. You are handed a letter and asked to reply. Reading the letter is fast: your eyes take in whole lines at once. Writing the reply is slow: one word at a time, and before each word you must walk to a filing cabinet and flip through an entire reference binder. The walk, not the thinking, is what takes the time.

The model is the same. Prefill reads the whole prompt in one parallel pass. Decode then produces the answer one token at a time, and every single token requires streaming all the weights from GPU memory.

Tiny worked example. An 8-billion-parameter model stored at 16 bits is 16 GB of weights. On a GPU that moves 3.35 TB/s from memory and does about 10¹⁵ 16-bit operations per second:

  • Prefill of a 1,000-token prompt: 2 × 8×10⁹ × 1,000 = 1.6×10¹³ operations, about 16 ms.
  • Decode: read 16 GB for every token, 16×10⁹ / 3.35×10¹² s = 4.78 ms per token, at most about 209 tokens per second for a single user.
sequenceDiagram participant U as User participant M as Model participant C as KV cache U->>M: Prompt of 1,000 tokens M->>C: Prefill: store K and V for all 1,000 M-->>U: First token loop Each new token M->>C: Read cached K and V, add one entry M-->>U: Next token end

Reading it: time runs downward. The first arrow is prefill: one big parallel pass over the whole prompt, which also fills the KV cache (section 2). It decides the time to first token. Everything inside the loop is decode: one small pass per token, each reading the cache and adding one entry to it. It decides tokens per second. Long prompts slow the first token; long answers slow the total.

The reason the phases behave so differently is arithmetic intensity: how many operations you do for each byte you fetch from memory.

Level 3: the formula and its symbols

$$ I = \frac{2\,n}{b} \qquad t_{\text{decode}} \approx \frac{P\,b}{\text{BW}} \qquad t_{\text{prefill}} \approx \frac{2\,P\,n}{\text{FLOPS}} $$

Symbols

Symbol Meaning here Shape / range
$I$ arithmetic intensity: operations per byte of weights read FLOPs/byte
$n$ tokens processed in one pass (1 when decoding, the prompt length when prefilling) ≥ 1
$2$ one multiply plus one add per weight per token
$b$ bytes per weight (2 at 16-bit, 0.5 at 4-bit)
$P$ number of parameters (weights) e.g. 8×10⁹
BW memory bandwidth: bytes the GPU can read per second e.g. 3.35×10¹²
FLOPS arithmetic throughput: operations per second e.g. 10¹⁵
$\approx$ "roughly": these are lower bounds that ignore overheads

In words: decode does two operations per two-byte weight it reads, so its speed is set by memory bandwidth; prefill reuses each weight for every prompt token, so its speed is set by arithmetic.

On the worked example: decode I = 2 × 1 / 2 = 1 FLOP per byte; prefill of 1,000 tokens I = 1,000. The GPU breaks even at 10¹⁵ / 3.35×10¹² ≈ 299 FLOPs per byte, so decode sits far below it (memory-bound) and prefill far above (compute-bound).

Level 3: in Python

In Python:

# 8 billion weights at 2 bytes each (16-bit)
P, b = 8e9, 2
# bytes read per second, operations per second
BW, FLOPS = 3.35e12, 1e15
def I(n):
    # operations per byte of weights read
    return 2 * n / b
# decode, then a 1,000-token prefill
I(1), I(1000)  # → (1.0, 1000.0)
# the break-even intensity
round(FLOPS / BW)  # → 299
# t_decode, in milliseconds per token
round(P * b / BW * 1000, 2)  # → 4.78
# t_prefill for 1,000 tokens, in milliseconds
round(2 * P * 1000 / FLOPS * 1000, 1)  # → 16.0

Decode for 1 user uses 0.3% of the GPU's compute and a batch of 64 reaches 21%, while a 1,000-token prefill passes the 299 break-even to run at full speed

Reading it: the x-axis is arithmetic intensity (log scale); the y-axis is the speed the GPU can actually reach. The sloped part of the roof is the memory limit (bandwidth × intensity); the flat part is the arithmetic limit. The corner is the break-even point, ~299. Decode at batch size 1 sits at intensity 1, deep in the memory-bound region, using well under 1% of the GPU's arithmetic. Batching many users together moves decode to the right, because one read of the weights then serves every user in the batch. That is the economic reason inference servers batch aggressively.

In code: arithmetic_intensity computes I, ridge_point finds the break-even, and bottleneck says which side a pass falls on; decode_seconds_per_token and prefill_seconds give the two time bounds.

Why it matters in practice. When a system feels slow, ask which phase dominates. Slow first token: the prompt is long (trim it, or cache it). Slow streaming: decode is memory-bound (quantize, batch, speculate).

2. The KV cache: take notes instead of rereading

Everyday picture. Reading a mystery novel, you don't reread the whole book before each new sentence; you keep notes on every character and clue. For each new sentence you glance at your notes and add one line.

Attention needs every earlier token's key and value (see primer.ml.attention). Because of causal masking, an earlier token's keys and values never change when later tokens arrive, so they can be computed once and kept.

Tiny worked example. An 8-token prompt, then 16 generated tokens.

  • Without a cache, step k re-processes the whole sequence so far: 8 + 9 + … + 23 = 248 token positions.
  • With a cache: prefill the 8 prompt tokens once, then 1 new position for each of the next 15 tokens = 23 token positions.

On TinyDecoder that is about 11× fewer operations for this short run, and the gap grows with length: without a cache the total work grows with the square of the length.

flowchart LR subgraph NOCACHE["No cache: every step starts over"] A1[step 1: tokens 1..8] --> A2[step 2: tokens 1..9] --> A3[step 3: tokens 1..10] end subgraph CACHE["KV cache: every step adds one"] B1[prefill: tokens 1..8<br/>store K,V] --> B2[token 9 only<br/>read K,V 1..8] --> B3[token 10 only<br/>read K,V 1..9] end

Reading it: the top row recomputes a sequence that grows by one each step. The bottom row computes each token exactly once: prefill stores keys and values for the prompt, and each decode step computes only the newest token's query, key and value, reading everything older from the cache. The outputs are identical (TinyDecoder checks this to 10 decimal places); only the cost differs.

Without a cache, total work curves upward to about 20 times the cached total after 32 tokens; with the cache it grows in a straight line

Reading it: the x-axis is how many tokens have been generated after an 8-token prompt; the y-axis is total operations spent so far, counted inside TinyDecoder. Without a cache the curve bends upward (each step costs more than the last); with a cache it is a straight line (each step costs about the same). The gap between them is pure waste the cache removes.

In code: TinyDecoder.forward_full processes a whole sequence (and fills a cache during prefill), TinyDecoder.forward_step processes one new token against the cache from TinyDecoder.new_cache, and TinyDecoder.generate runs either way, returning a Generation that holds the tokens and the work counted.

Why it matters in practice. The cache trades memory for speed, and that memory is what limits how many users one GPU can serve. Section 3 does the math.

3. Memory math: will it fit?

Everyday picture. Packing for a trip. The suitcase is GPU memory. The model's weights are the big fixed items that always go in. Every active conversation adds a bag of notes (its KV cache) whose size grows with the conversation's length. Once the suitcase is full, the next customer waits.

Tiny worked example.

  • Weights: 70×10⁹ parameters × 2 bytes (16-bit) = 140 GB: more than one 80 GB GPU. At 4 bits (half a byte) it is 35 GB and fits on one.
  • KV cache per token for a Llama-3-8B-shaped model (32 layers, 8 KV heads, 128 dimensions per head, 16-bit): 2 × 32 × 8 × 128 × 2 = 131,072 bytes, about 128 KB.
  • A 32,000-token conversation: 32,000 × 131,072 ≈ 4.2 GB of cache.
  • An 80 GB GPU holding 16 GB of weights has 64 GB left: room for 15 such conversations at once.
Level 3: the formula and its symbols

$$ \text{weight bytes} = P \times \frac{\text{bits}}{8} \qquad \text{KV bytes per token} = 2 \times L \times H_{kv} \times d_h \times b $$

Symbols

Symbol Meaning here Shape / range
$P$ number of parameters e.g. 7×10¹⁰
bits / 8 bytes per parameter (16 bits = 2 bytes)
$2$ one key and one value per token
$L$ number of transformer layers, each with its own cache e.g. 32
$H_{kv}$ number of key/value heads (fewer than query heads with grouped-query attention) e.g. 8
$d_h$ dimensions per head e.g. 128
$b$ bytes per stored number 2 at 16-bit

In words: weights cost parameters times bytes each; the cache costs, for every token, one key and one value per layer per KV head.

On the worked example: 7×10¹⁰ × 16/8 = 1.4×10¹¹ bytes = 140 GB; and 2 × 32 × 8 × 128 × 2 = 131,072 bytes per token.

Level 3: in Python

In Python:

P = 7e10
# weight GB at 16 bits, then at 4 bits
P * 16 / 8 / 1e9, P * 4 / 8 / 1e9  # → (140.0, 35.0)
L, H_kv, d_h, b = 32, 8, 128, 2
# KV bytes per token: a key and a value, per layer, per KV head
2 * L * H_kv * d_h * b  # → 131072
# GB of cache for a 32,000-token conversation
round(32_000 * 131_072 / 1e9, 1)  # → 4.2

At 128k tokens, 32 KV heads need 67 GB, more than the 64 GB free, while 8 KV heads need 17 GB, so three such requests fit

Reading it: the x-axis is context length per request; the y-axis is KV cache memory for one request. The steep line is a model with 32 KV heads (classic multi-head attention); the shallow one has 8 (grouped-query attention, like Llama 3). The dashed line is the 64 GB left on an 80 GB GPU after 16 GB of weights. With 32 KV heads a single 128k-token request would not fit; with 8 it fits three times over. Doing this arithmetic out loud is the fastest way to size a deployment.

Try it: pick a model shape, then drag the context length and the number of requests. The bar is one GPU's memory: the grey part is the weights, the coloured part the KV cache. Watch how quickly a long context pushes the cache past the weights, and how much further 8 KV heads go than 32.

In code: weight_bytes and kv_cache_bytes_per_token are the two formulas, kv_cache_bytes scales the cache to a context and batch, and max_concurrent_requests counts how many conversations fit beside the weights.

4. Sampling: from scores to one token

Everyday picture. Choosing where to eat. Greedy always picks the top-rated place. Sampling holds a lottery weighted by rating. Temperature is how adventurous you feel: low means you nearly always pick the favourite, high means long shots get a real chance. Top-k says "only consider the top 3". Top-p says "consider just enough places to cover 90% of my enthusiasm": one place if you have a clear favourite, several if you're torn.

Tiny worked example. Scores (logits) 2, 1, 0 for three tokens:

temperature probabilities
0 (greedy) 1, 0, 0
0.5 0.867, 0.117, 0.016
1 0.665, 0.245, 0.090
2 0.506, 0.307, 0.186

With probabilities 0.5, 0.3, 0.15, 0.05: top-k = 2 keeps 0.5 and 0.3, renormalised to 0.625 and 0.375. Top-p = 0.9 keeps 0.5, 0.3 and 0.15 (the first set whose total reaches 0.9) and cuts the 0.05 tail.

flowchart LR L[Logits, one per<br/>vocabulary token] --> T[Divide by<br/>temperature T] T --> S[Softmax<br/>probabilities] S --> K[Top-k: keep the<br/>k likeliest] K --> P[Top-p: keep the smallest set<br/>reaching probability p] P --> R[Renormalise<br/>and draw one token]

Reading it: the model only ever produces the scores on the left; everything after that is a choice you make at request time. Temperature reshapes the whole distribution; top-k and top-p then cut off the unreliable tail before the draw, so a rare nonsense token can't be picked by bad luck.

Level 3: the formula and its symbols

$$ p_i = \frac{e^{z_i / T}}{\sum_j e^{z_j / T}} $$

Symbols

Symbol Meaning here Shape / range
$z_i$ the model's score (logit) for vocabulary token i any real
$T$ temperature > 0; T → 0 approaches greedy
$e^{x}$ the exponential function; makes every score positive and stretches gaps
$\sum_j$ sum over every token j in the vocabulary, so the $p_i$ add to 1
$p_i$ probability of drawing token i 0 … 1

In words: divide every score by the temperature, exponentiate, and divide by the total so the results sum to one.

On the worked example: T = 0.5 turns (2, 1, 0) into (4, 2, 0); e⁴ = 54.6, e² = 7.39, e⁰ = 1, total 63.0; probabilities 0.867, 0.117, 0.016.

Level 3: in Python

In Python:

import math
z, T = [2.0, 1.0, 0.0], 0.5
# e^(z_i / T)
exps = [math.exp(z_i / T) for z_i in z]
# Σ_j e^(z_j / T)
round(sum(exps), 1)  # → 63.0
# p_i
[round(e / sum(exps), 3) for e in exps]  # → [0.867, 0.117, 0.016]

Five tokens at three temperatures: at T = 0.5 the favourite takes 79% and top-p 0.9 cuts three tokens; at T = 2 it falls to 38% and only one is cut

Reading it: five candidate tokens with fixed scores, shown at three temperatures. At T = 0.5 (left) the favourite takes 79%; at T = 2 (right) it falls to 38% and the rest spread out, down to 7% for the least likely. The hatched bars are the tokens top-p = 0.9 would cut: three at T = 0.5, two at T = 1, one at T = 2. A confident model reaches 90% with fewer tokens, so top-p cuts more of them; an unsure one needs more tokens to reach 90%, so it cuts fewer. That's why top-p adapts where a fixed top-k can't.

Temperature 0 is not a determinism guarantee. Floating-point addition isn't associative: in 32-bit floats, (10⁸ + 1) − 10⁸ = 0 but (10⁸ − 10⁸) + 1 = 1. On a GPU the order of additions can depend on the kernel chosen and on what else is in the batch, so two nearly tied tokens can swap places between otherwise identical requests.

In code: temperature_probs applies the formula, top_k_filter and top_p_filter cut the tail, and sample_next chains them into one draw. float32_sum and greedy_pick_with_summation_order show a near tie flipping with the order of additions.

Why it matters in practice. Use low temperature for extraction, classification and tool calls; higher for brainstorming and creative text.

5. Speculative decoding: a junior drafts, a senior checks

Everyday picture. A junior writer drafts the next few sentences quickly. A senior editor reads the whole draft at once, keeps every sentence they would have written themselves, rewrites the first one they wouldn't, and throws away the rest. The senior's reading is fast; their writing is slow. So the team moves at the junior's speed but produces the senior's text.

This works because decode is memory-bound: checking 5 draft tokens in one pass of the big model costs about the same as generating 1.

Tiny worked example. The big model's next-token probabilities are (0.5, 0.3, 0.2); the small model's are (0.3, 0.3, 0.4). A draft token survives with probability Σ min = 0.3 + 0.3 + 0.2 = 0.8. Drafting 4 tokens per round yields on average (1 − 0.8⁵) / (1 − 0.8) = 3.36 tokens per pass of the big model instead of 1.

flowchart TD D[Small model drafts γ tokens<br/>one by one, cheap] --> V[Big model scores all γ positions<br/>in ONE parallel pass] V --> C{For each draft x in order:<br/>keep with probability min 1, p/q} C -->|kept| N[Next draft] N --> C C -->|rejected| F[Replace x with a draw from<br/>max 0, p − q, renormalised. Stop.] C -->|all kept| B[Bonus: draw one more<br/>token from the big model]

Reading it: the loop in the middle walks the drafts left to right. A draft the big model likes at least as much as the small one did (p ≥ q) is always kept; one it likes less is kept only with probability p/q. The first rejection is replaced by a token drawn from exactly the probability the small model under-proposed, which is what makes the final output statistically identical to sampling from the big model alone. The test suite checks this empirically over 20,000 rounds.

Level 3: the formula and its symbols

$$ P(\text{keep } x) = \min!\left(1, \frac{p(x)}{q(x)}\right) \qquad \alpha = \sum_x \min\big(p(x), q(x)\big) \qquad \mathbb{E}[\text{tokens per pass}] = \frac{1 - \alpha^{\gamma+1}}{1 - \alpha} $$

Symbols

Symbol Meaning here Shape / range
$p(x)$ the big (target) model's probability for token x 0 … 1
$q(x)$ the small (draft) model's probability for token x 0 … 1
$\min(a, b)$ the smaller of a and b
$\alpha$ alpha, the acceptance rate: how often a draft survives 0 … 1
$\gamma$ gamma, how many tokens the small model drafts per round e.g. 4
$\mathbb{E}[\cdot]$ expected value: the long-run average

In words: keep a draft with probability "how much the big model likes it compared with the small one, capped at 1"; the average number of tokens per big-model pass is a geometric series in the acceptance rate.

On the worked example: α = 0.8, γ = 4: (1 − 0.8⁵)/(1 − 0.8) = (1 − 0.328)/0.2 = 3.36.

Level 3: in Python

In Python:

# big model
p = [0.5, 0.3, 0.2]
# small model
q = [0.3, 0.3, 0.4]
# P(keep x) for each token
[round(min(1.0, p_x / q_x), 2) for p_x, q_x in zip(p, q)]  # → [1.0, 1.0, 0.5]
alpha = sum(min(p_x, q_x) for p_x, q_x in zip(p, q))
round(alpha, 2)  # → 0.8
gamma = 4
# expected tokens per big-model pass
round((1 - alpha ** (gamma + 1)) / (1 - alpha), 2)  # → 3.36

In code: acceptance_rate computes α and expected_tokens_per_round the geometric series; speculative_round runs the draft, verify and replace loop once, and speculative_generate repeats it until enough tokens exist.

6. Quantization: fewer bits per weight

Everyday picture. Writing prices to the nearest dollar instead of the nearest cent: shorter to store, slightly less exact. Using one ruler per row of a spreadsheet, instead of one for the whole sheet, keeps a single huge number in one row from making every other row coarse.

Tiny worked example. A row of weights (0.5, −1.27, 0.02).

  • int8: the largest magnitude, 1.27, maps to 127, so the step size is 0.01. Codes: 50, −127, 2. Decoding gives back exactly (0.5, −1.27, 0.02).
  • int4: only 15 levels (−7 … 7), step size 1.27/7 = 0.181. Codes: 3, −7, 0. Decoding gives (0.544, −1.27, 0): the small weight vanished.
Level 3: the formula and its symbols

$$ s = \frac{\max_j |w_j|}{2^{\,\text{bits}-1} - 1} \qquad c_j = \operatorname{round}!\left(\frac{w_j}{s}\right) \qquad \hat{w}_j = s \, c_j $$

Symbols

Symbol Meaning here Shape / range
$w_j$ the j-th original weight in the row real
$\max_j w_j $ the largest absolute value in the row ≥ 0
$2^{\text{bits}-1} - 1$ the largest code: 127 for 8 bits, 7 for 4 bits
$s$ the scale (step size), one per row > 0
$c_j$ the stored integer code −127…127 or −7…7
$\hat{w}_j$ the weight as reconstructed at inference time ("w-hat") real

In words: pick a step size so the largest weight lands on the largest code, store each weight as the nearest whole number of steps, and multiply back at run time.

On the worked example: int4: s = 1.27/7 = 0.181; 0.5/0.181 = 2.76 → 3; 3 × 0.181 = 0.544.

Level 3: in Python

In Python:

w = [0.5, -1.27, 0.02]
def quantize(w, bits):
    # the step size
    s = max(abs(w_j) for w_j in w) / (2 ** (bits - 1) - 1)
    # c_j: whole steps
    return s, [round(w_j / s) for w_j in w]
s, c = quantize(w, bits=8)
round(s, 3), c  # → (0.01, [50, -127, 2])
s, c = quantize(w, bits=4)
round(s, 3), c  # → (0.181, [3, -7, 0])
# ŵ_j = s · c_j: the 0.02 is gone
[round(s * c_j, 3) for c_j in c]  # → [0.544, -1.27, 0.0]

In code: quantize returns the integer codes and one scale per row, dequantize multiplies them back, and quantization_error measures how far the round trip lands from the original weights.

Why it matters in practice. 8-bit weights are nearly lossless; 4-bit methods with smarter rounding (GPTQ, AWQ) keep most quality at a quarter of the memory. Fewer bytes per weight also means faster memory-bound decode.

7. Continuous batching: seat the next party as soon as a table frees

Everyday picture. A restaurant with two tables. Static batching seats two parties and won't seat anyone new until both have left, so a table sits empty while one slow diner lingers. Continuous batching seats the next party the moment any table frees.

Tiny worked example. Four requests needing 4, 1, 1 and 1 decode steps, two slots. Static: {4, 1} runs 4 steps with one slot idle for 3, then {1, 1} runs 1: 5 steps, 70% of slot-steps busy. Continuous: the short requests slide into the freed slot while the long one runs: 4 steps, 87.5% busy.

With 32 requests on 8 slots, static batching leaves idle gaps and needs 433 steps at 60% busy; continuous batching needs 318 steps at 82%

Reading it: each row is a GPU batch slot and each column is one decode step; colour identifies the request occupying the slot, and white is an idle slot. On a realistic mix of 32 requests of varied length over 8 slots, static batching (top) leaves white holes wherever a short request finished early; continuous batching (bottom) keeps nearly every cell busy and finishes the same work in fewer steps (about 82% vs. 60% utilisation).

Level 3: the formula and its symbols

$$ U = \frac{\text{useful slot-steps}}{\text{steps} \times \text{slots}} $$

Symbols

Symbol Meaning here Shape / range
$U$ utilisation: the share of slot-steps doing real work 0 … 1
useful slot-steps total decode steps all requests need integer
steps × slots the capacity the GPU offered while serving them integer

In words: utilisation is the work the requests needed divided by the work capacity the server spent serving them.

On the worked example: 7 useful slot-steps; static 7/(5×2) = 0.70, continuous 7/(4×2) = 0.875.

Level 3: in Python

In Python:

# decode steps the four requests need
useful = 4 + 1 + 1 + 1
slots = 2
# static takes 5 steps, continuous 4
useful / (5 * slots), useful / (4 * slots)  # → (0.7, 0.875)

In code: simulate_static_batching and simulate_continuous_batching play out the two policies step by step, each returning a ServingRun that holds the slot timeline and its utilisation U.

Why it matters in practice. Continuous batching (together with paged KV-cache memory, as in vLLM) is a large part of why modern inference servers reach high throughput.

8. Prompt caching: reuse the prefill of a shared prefix

Everyday picture. A kitchen that pre-chops the ingredients every order uses, instead of chopping them again for each plate.

The KV cache normally lives for one request. Prompt caching keeps it across requests: if many calls start with the same long system prompt, tool definitions or reference document, the provider stores that prefix's keys and values and skips its prefill next time. It is only valid up to the first differing token, because every token's keys depend on everything before it, so put stable content first and volatile content last.

Tiny worked example. 100 requests, each a 10,000-token shared prefix plus a 500-token question, at \$3 per million input tokens, with cache writes at 1.25× and cache reads at 0.1× (typical of providers; check current pricing):

  • Without caching: 100 × 10,500 × \$3/10⁶ = \$3.15.
  • With caching: the first call writes the cache (\$0.039); each of the other 99 costs \$0.0045. Total \$0.48, an 85% saving, and each cached call also skips 10,000 tokens of prefill, so its first token arrives sooner.
flowchart LR subgraph Prompt["One request's prompt, in order"] S[System prompt<br/>stable] --> T[Tool definitions<br/>stable] --> D[Reference docs<br/>stable] --> Q[User question<br/>changes every call] end S & T & D -.->|cached after the first call| C[(Prefix KV cache)]

Reading it: the prompt is read left to right, and the cache can cover only an unbroken run from the very start. Stable parts go first so that every request shares the longest possible prefix; the part that changes on every call goes last. A timestamp or request ID placed at the top would break the match on the first token and silently disable caching.

In code: reusable_prefix_tokens counts how many leading tokens two prompts share, and prompt_cache_cost prices a run of requests with and without the cache.

In 20 seconds

  • Prefill reads the prompt in parallel (compute-bound, sets time to first token); decode writes one token at a time (memory-bound, sets tokens per second).
  • The KV cache stores every earlier token's keys and values so each step computes one token; it trades GPU memory for speed.
  • Memory math: weights = parameters × bytes; KV per token = 2 × layers × KV heads × head dim × bytes. 70B at 16-bit is 140 GB; 32k tokens of Llama-3-8B cache is about 4 GB.
  • Temperature reshapes the distribution; top-k and top-p cut the tail. Temperature 0 reduces but does not guarantee determinism.
  • Speculative decoding: a small model drafts, the big one verifies in one pass; output is identical in distribution.
  • Quantization, continuous batching and prompt caching are the other big serving levers.

Self-test questions

Q: Why is decoding slow even on a huge GPU? A: Each token needs every weight read from memory but does only about one operation per byte read, far below the GPU's break-even (~300 FLOPs/byte). The arithmetic units mostly wait on memory.

Q: What is the KV cache, and why does it matter for serving cost? A: Stored keys and values of all earlier tokens, so each new token is computed once instead of recomputing the whole sequence. It grows with context length and batch size, and that memory caps how many requests a GPU serves at once.

Q: How much memory does a 70B model need at 16-bit and at 4-bit? A: 140 GB and 35 GB for weights alone, plus KV cache and activations.

Q: Estimate the KV cache for one 32k-token request on a Llama-3-8B-shaped model. A: 2 × 32 × 8 × 128 × 2 = 131,072 bytes per token; × 32,000 ≈ 4.2 GB.

Q: Why does grouped-query attention make serving cheaper? A: It shares each key/value head among several query heads, shrinking the KV cache (4× for 32 → 8 KV heads), so more requests fit per GPU.

Q: Why isn't temperature 0 perfectly deterministic? A: Floating-point addition isn't associative, and GPU reduction order can change with kernels and batch composition, so nearly tied logits can flip.

Q: How can speculative decoding be faster yet produce the same distribution? A: Verifying several draft tokens costs one memory-bound pass of the big model, about the same as generating one. The min(1, p/q) accept rule plus resampling from max(0, p − q) on rejection makes the output exactly the big model's distribution.

Q: What does continuous batching fix? A: Idle slots: static batches wait for their longest request, while continuous batching refills any freed slot immediately, raising utilisation.

Q: How do you structure a prompt to benefit from prompt caching? A: Put stable content (system prompt, tool definitions, reference documents) first and anything that varies (the question, timestamps, IDs) last, because a cache is valid only up to the first differing token.

The papers behind this lesson

  • Leviathan, Kalman & Matias, Fast Inference from Transformers via Speculative Decoding (2022): https://arxiv.org/abs/2211.17192. Introduced the draft-then-verify scheme with an accept/reject rule that provably preserves the target model's output distribution. annotated companion
  • Kwon et al., Efficient Memory Management for Large Language Model Serving with PagedAttention (2023): https://arxiv.org/abs/2309.06180. Stored the KV cache in fixed-size pages like an operating system's virtual memory, eliminating fragmentation and enabling vLLM's high-throughput continuous batching. annotated companion
  • Dao et al., FlashAttention (2022): https://arxiv.org/abs/2205.14135. Computed exact attention in GPU-memory-sized tiles, cutting the slow memory traffic that dominates long-context inference. annotated companion
  • Dettmers et al., LLM.int8(): 8-bit Matrix Multiplication for Transformers at Scale (2022): https://arxiv.org/abs/2208.07339. Showed that a few outlier features break naive 8-bit quantization and handled them separately.
  • Frantar et al., GPTQ: Accurate Post-Training Quantization for Generative Pre-trained Transformers (2022): https://arxiv.org/abs/2210.17323. Made 3-4-bit weight quantization practical with error-compensating rounding.
  • Holtzman et al., The Curious Case of Neural Text Degeneration (2019): https://arxiv.org/abs/1904.09751. Introduced nucleus (top-p) sampling.

Further reading

on GitHub
   1r"""
   2# Inference: what happens, and what it costs, when a model generates
   3
   4Run: `python -m primer.ml.inference`
   5
   6New to the notation? `primer.notation` explains every symbol used here from
   7zero. This lesson builds on attention from `primer.ml.attention` and the
   8transformer from `primer.ml.transformer`.
   9
  10## Level 1: The practitioner's guide
  11
  12**In one sentence.** Inference is everything that happens when a trained
  13model answers a request, and its cost and speed are governed by one fact:
  14each generated token reads every weight from memory, so the levers that
  15matter (caching, batching, quantization, speculation, sampling settings)
  16are all ways to read less or to reuse each read.
  17
  18**When you need it.** You need this lesson the day a model leaves the
  19notebook: when you set `temperature` and `top_p` on a request, when someone
  20asks why the first token takes two seconds, when a GPU bill arrives, or when
  21you decide whether to serve a model yourself. The tell is a latency or cost
  22question you can only answer by guessing. The numbers here replace the
  23guess. This lesson's model of an 8-billion-parameter network at 16 bits on
  24an H100-class GPU (3.35 TB/s of memory bandwidth, about 10¹⁵ operations per
  25second, the constants in this module) reads a 1,000-token prompt in about
  2616 ms but then produces at most about 209 tokens per second for a single
  27user, because every token needs all 16 GB of weights read again (4.78 ms
  28each). A lone request uses well under 1% of the GPU's arithmetic. You do not
  29need this lesson for a prototype at ten requests a day; you need it before
  30the first load test.
  31
  32**Your options.** The levers a practitioner can pull, roughly from the
  33cheapest to the most involved:
  34
  35| Lever | What it does | What it buys you | What it costs | Where it lives |
  36|---|---|---|---|---|
  37| Sampling settings | Temperature reshapes the token probabilities; top-k and top-p cut the unlikely tail | Control over variety: T = 0 for extraction and tool calls, higher for ideas | Nothing in compute; wrong settings cost quality | The request |
  38| Prompt caching | Reuses the prefill of a shared prefix (system prompt, tools, documents) across calls | This lesson's 100 calls with a 10,000-token prefix: \$0.48 instead of \$3.15, an 85% saving, and a faster first token | A cache write at about 1.25× the input price on the first call (Anthropic's prompt caching docs) | The provider, or your server's prefix cache |
  39| Quantization | Stores weights at 8 or 4 bits instead of 16 | A 70B model in 35 GB instead of 140 GB, and faster memory-bound decode | 8-bit is nearly lossless; naive 4-bit loses small weights (the lesson's 0.02 rounds to 0) unless a smarter method such as GPTQ or AWQ is used | The model files and the server |
  40| A model with grouped-query attention | Fewer key/value heads means a smaller KV cache per token | 8 KV heads instead of 32 fit four times as many long conversations per GPU | A model choice made at training time; you can only pick a model that has it | The model architecture |
  41| Continuous batching | Seats a new request the moment any slot frees, instead of waiting for the whole batch | In this lesson's 32-request simulation, 82% slot utilisation instead of 60%, in fewer steps | Nothing beyond a server that does it (they all do now) | The inference server |
  42| Speculative decoding | A small model drafts several tokens; the big one verifies them in one pass | With an 80% acceptance rate and 4 drafts, 3.36 tokens per big-model pass instead of 1, with the same output distribution | A draft model to run, and gains that shrink when the draft guesses badly or the server is already batch-saturated | The inference server |
  43| A hosted API | Someone else runs all of the above | No GPUs to size, caching and batching done for you | A per-token price, and less control over settings and residency | The provider |
  44
  45**How to choose.** Start from the symptom.
  46
  47- Slow first token: the prompt is long. Trim it, or put its stable part
  48  first and let prompt caching skip its prefill.
  49- Slow streaming: decode is memory-bound. Quantize, batch more requests
  50  together, or add speculative decoding.
  51- Running out of GPU memory as traffic grows: it is the KV cache, not the
  52  weights. Do the arithmetic (weights plus cache per token times context
  53  times concurrent requests) before renting a bigger card; prefer a model
  54  with grouped-query attention and cap the context you allow.
  55- Answers that vary when you want them stable: temperature 0 and a
  56  validator, not a hope. Answers that all sound the same when you want
  57  range: raise the temperature and let top-p keep the tail sane.
  58- Deciding whether to self-host: only when volume, privacy or a model the
  59  APIs do not offer justifies owning the batching and memory problems above.
  60- Whatever you pick, measure time to first token and tokens per second
  61  separately. They are set by different phases and fixed by different
  62  levers.
  63
  64**What it costs.** Money follows tokens, and tokens follow decode. Prefill
  65of 1,000 tokens costs 16 ms of a GPU's full compute; each output token costs
  66a full read of the weights, so output tokens are the expensive ones, and
  67providers price them that way. Memory sets capacity: on this lesson's
  68Llama-3-8B-shaped model each token of context holds about 128 KB of keys and
  69values, a 32,000-token conversation holds 4.2 GB, and an 80 GB GPU with 16
  70GB of weights fits 15 such conversations at once. Batching is what makes
  71serving economical: one read of the weights serves every request in the
  72batch, which is why the roofline figure in Level 2 shows a batch of 64
  73reaching 21% of the GPU's arithmetic where a single user reaches 0.3%. Prompt
  74caching costs a little on the first call and saves most of the input bill
  75after it; Anthropic prices a five-minute cache write at 1.25× and a read at
  760.1× the input price, with a one-hour write at 2× (its prompt caching
  77docs). Quantization costs a little quality for a large memory saving.
  78Speculative decoding costs a second model and a more complex server.
  79
  80**What breaks.**
  81
  82- **Temperature 0 still varies.** Floating-point addition depends on order,
  83  and GPU kernels change their order with the batch they land in. One
  84  published measurement found 80 distinct completions in 1,000 runs at
  85  temperature 0, identical for the first 102 tokens and then diverging
  86  (Thinking Machines, *Defeating Nondeterminism in LLM Inference*). Treat
  87  determinism as reduced, not guaranteed, and validate outputs.
  88- **A timestamp at the top of the prompt** silently disables prompt
  89  caching, because a cache matches only up to the first differing token.
  90  Stable content first, volatile content last; on Anthropic's API the order
  91  is tools, then system, then messages.
  92- **Short prompts are not cached.** Providers set a minimum cacheable
  93  length (Anthropic's is between 512 and 4,096 tokens depending on the
  94  model) and return no error below it; the bill just does not fall.
  95- **A long-context feature exhausts memory.** Doubling the context you
  96  allow doubles the cache per request and halves the requests that fit.
  97- **Naive 4-bit quantization erases small weights.** Use a method that
  98  compensates (GPTQ, AWQ) and check quality on your own evaluation, not on
  99  the model card.
 100- **Top-k with a fixed k** cuts too much when the model is unsure and too
 101  little when it is confident; top-p adapts, which is why it is the usual
 102  default.
 103- **Speculation that guesses badly** costs more than it saves: the big
 104  model's pass still runs, and every rejected draft is wasted work.
 105
 106**In the wild.** vLLM's documentation lists continuous batching, chunked
 107prefill, prefix caching, PagedAttention for KV memory, speculative decoding
 108(n-gram and EAGLE drafts among others) and quantization from FP8 to INT4,
 109GPTQ and AWQ, behind an OpenAI-compatible API. SGLang offers the same set
 110with RadixAttention for prefix caching; NVIDIA's TensorRT-LLM does it with
 111custom kernels and FP8 and FP4 formats on NVIDIA GPUs; llama.cpp runs
 112quantized models from 1.5-bit to 8-bit on CPUs and Apple silicon. Hugging
 113Face's Transformers exposes the sampling knobs (greedy by default, sampling
 114with `do_sample`, beam search with `num_beams`). Hosted APIs expose prompt
 115caching explicitly, with Anthropic's linked in Further reading. The ideas
 116come from the papers at the end of this lesson: speculative decoding
 117(Leviathan, Kalman and Matias), PagedAttention (Kwon et al.), FlashAttention
 118(Dao et al.), LLM.int8() and GPTQ for quantization, and nucleus sampling
 119(Holtzman et al.).
 120
 121**Go deeper.** Level 2 builds each lever from nothing: the roofline that
 122explains why decode is memory-bound, a tiny decoder with and without a KV
 123cache whose outputs match to ten decimal places, the temperature and top-p
 124arithmetic, the accept-or-reject rule that makes speculative decoding exact,
 125a quantizer in five lines, and simulations of both batching policies. If you
 126only needed to size a deployment or set a request, you are done.
 127
 128## Level 2: How it works, from scratch
 129
 130Training happens once; inference happens every time anyone uses the model,
 131so this is where the money goes. Generating text has two very different
 132phases, one memory trick that makes it affordable (the KV cache), a few
 133knobs that decide *which* token comes out (sampling), and a toolbox of
 134speedups: quantization, speculative decoding, continuous batching and
 135prompt caching. Every one of them follows from one fact:
 136
 137> **Generating one token requires reading every weight of the model from
 138> memory, and memory is much slower than arithmetic.**
 139
 140## 1. Two phases: prefill and decode
 141
 142**Everyday picture.** You are handed a letter and asked to reply. Reading the
 143letter is fast: your eyes take in whole lines at once. Writing the reply is
 144slow: one word at a time, and before each word you must walk to a filing
 145cabinet and flip through an entire reference binder. The walk, not the
 146thinking, is what takes the time.
 147
 148The model is the same. **Prefill** reads the whole prompt in one parallel
 149pass. **Decode** then produces the answer one token at a time, and every
 150single token requires streaming all the weights from GPU memory.
 151
 152**Tiny worked example.** An 8-billion-parameter model stored at 16 bits is
 15316 GB of weights. On a GPU that moves 3.35 TB/s from memory and does about
 15410¹⁵ 16-bit operations per second:
 155
 156* Prefill of a 1,000-token prompt: 2 × 8×10⁹ × 1,000 = 1.6×10¹³ operations,
 157  about **16 ms**.
 158* Decode: read 16 GB for every token, 16×10⁹ / 3.35×10¹² s = **4.78 ms per
 159  token**, at most about 209 tokens per second for a single user.
 160
 161```mermaid
 162sequenceDiagram
 163  participant U as User
 164  participant M as Model
 165  participant C as KV cache
 166  U->>M: Prompt of 1,000 tokens
 167  M->>C: Prefill: store K and V for all 1,000
 168  M-->>U: First token
 169  loop Each new token
 170    M->>C: Read cached K and V, add one entry
 171    M-->>U: Next token
 172  end
 173```
 174
 175**Reading it:** time runs downward. The first arrow is prefill: one big
 176parallel pass over the whole prompt, which also fills the KV cache (section
 1772). It decides the **time to first token**. Everything inside the loop is
 178decode: one small pass per token, each reading the cache and adding one
 179entry to it. It decides **tokens per second**. Long prompts slow the first
 180token; long answers slow the total.
 181
 182The reason the phases behave so differently is **arithmetic intensity**:
 183how many operations you do for each byte you fetch from memory.
 184
 185$$
 186I = \frac{2\,n}{b} \qquad
 187t_{\text{decode}} \approx \frac{P\,b}{\text{BW}} \qquad
 188t_{\text{prefill}} \approx \frac{2\,P\,n}{\text{FLOPS}}
 189$$
 190
 191**Symbols**
 192
 193| Symbol | Meaning here | Shape / range |
 194|---|---|---|
 195| $I$ | arithmetic intensity: operations per byte of weights read | FLOPs/byte |
 196| $n$ | tokens processed in one pass (1 when decoding, the prompt length when prefilling) | ≥ 1 |
 197| $2$ | one multiply plus one add per weight per token | |
 198| $b$ | bytes per weight (2 at 16-bit, 0.5 at 4-bit) | |
 199| $P$ | number of parameters (weights) | e.g. 8×10⁹ |
 200| BW | memory bandwidth: bytes the GPU can read per second | e.g. 3.35×10¹² |
 201| FLOPS | arithmetic throughput: operations per second | e.g. 10¹⁵ |
 202| $\approx$ | "roughly": these are lower bounds that ignore overheads | |
 203
 204**In words:** decode does two operations per two-byte weight it reads, so
 205its speed is set by memory bandwidth; prefill reuses each weight for every
 206prompt token, so its speed is set by arithmetic.
 207
 208**On the worked example:** decode I = 2 × 1 / 2 = 1 FLOP per byte; prefill
 209of 1,000 tokens I = 1,000. The GPU breaks even at 10¹⁵ / 3.35×10¹² ≈ **299**
 210FLOPs per byte, so decode sits far below it (**memory-bound**) and prefill
 211far above (**compute-bound**).
 212
 213**In Python:**
 214
 215```python
 216# 8 billion weights at 2 bytes each (16-bit)
 217P, b = 8e9, 2
 218# bytes read per second, operations per second
 219BW, FLOPS = 3.35e12, 1e15
 220def I(n):
 221    # operations per byte of weights read
 222    return 2 * n / b
 223# decode, then a 1,000-token prefill
 224I(1), I(1000)  # → (1.0, 1000.0)
 225# the break-even intensity
 226round(FLOPS / BW)  # → 299
 227# t_decode, in milliseconds per token
 228round(P * b / BW * 1000, 2)  # → 4.78
 229# t_prefill for 1,000 tokens, in milliseconds
 230round(2 * P * 1000 / FLOPS * 1000, 1)  # → 16.0
 231```
 232
 233![Decode for 1 user uses 0.3% of the GPU's compute and a batch of 64 reaches 21%, while a 1,000-token prefill passes the 299 break-even to run at full speed](figures/primer.ml.inference.roofline.svg)
 234
 235**Reading it:** the x-axis is arithmetic intensity (log scale); the y-axis
 236is the speed the GPU can actually reach. The sloped part of the roof is the
 237memory limit (bandwidth × intensity); the flat part is the arithmetic limit.
 238The corner is the break-even point, ~299. Decode at batch size 1 sits at
 239intensity 1, deep in the memory-bound region, using well under 1% of the
 240GPU's arithmetic. Batching many users together moves decode to the right,
 241because one read of the weights then serves every user in the batch. That is
 242the economic reason inference servers batch aggressively.
 243
 244**In code:** `arithmetic_intensity` computes I, `ridge_point` finds the
 245break-even, and `bottleneck` says which side a pass falls on;
 246`decode_seconds_per_token` and `prefill_seconds` give the two time bounds.
 247
 248**Why it matters in practice.** When a system feels slow, ask which phase
 249dominates. Slow first token: the prompt is long (trim it, or cache it).
 250Slow streaming: decode is memory-bound (quantize, batch, speculate).
 251
 252## 2. The KV cache: take notes instead of rereading
 253
 254**Everyday picture.** Reading a mystery novel, you don't reread the whole book
 255before each new sentence; you keep notes on every character and clue. For
 256each new sentence you glance at your notes and add one line.
 257
 258Attention needs every earlier token's key and value (see `primer.ml.attention`).
 259Because of causal masking, an earlier token's keys and values never change
 260when later tokens arrive, so they can be computed once and kept.
 261
 262**Tiny worked example.** An 8-token prompt, then 16 generated tokens.
 263
 264* Without a cache, step k re-processes the whole sequence so far: 8 + 9 + … +
 265  23 = **248** token positions.
 266* With a cache: prefill the 8 prompt tokens once, then 1 new position for
 267  each of the next 15 tokens = **23** token positions.
 268
 269On `TinyDecoder` that is about 11× fewer operations for this short run, and
 270the gap grows with length: without a cache the total work grows with the
 271*square* of the length.
 272
 273```mermaid
 274flowchart LR
 275  subgraph NOCACHE["No cache: every step starts over"]
 276    A1[step 1: tokens 1..8] --> A2[step 2: tokens 1..9] --> A3[step 3: tokens 1..10]
 277  end
 278  subgraph CACHE["KV cache: every step adds one"]
 279    B1[prefill: tokens 1..8<br/>store K,V] --> B2[token 9 only<br/>read K,V 1..8] --> B3[token 10 only<br/>read K,V 1..9]
 280  end
 281```
 282
 283**Reading it:** the top row recomputes a sequence that grows by one each
 284step. The bottom row computes each token exactly once: prefill stores keys
 285and values for the prompt, and each decode step computes only the newest
 286token's query, key and value, reading everything older from the cache. The
 287outputs are identical (`TinyDecoder` checks this to 10 decimal places); only
 288the cost differs.
 289
 290![Without a cache, total work curves upward to about 20 times the cached total after 32 tokens; with the cache it grows in a straight line](figures/primer.ml.inference.kv_cache_work.svg)
 291
 292**Reading it:** the x-axis is how many tokens have been generated after an
 2938-token prompt; the y-axis is total operations spent so far, counted inside
 294`TinyDecoder`. Without a cache the curve bends upward (each step costs more
 295than the last); with a cache it is a straight line (each step costs about the
 296same). The gap between them is pure waste the cache removes.
 297
 298**In code:** `TinyDecoder.forward_full` processes a whole sequence (and
 299fills a cache during prefill), `TinyDecoder.forward_step` processes one new
 300token against the cache from `TinyDecoder.new_cache`, and
 301`TinyDecoder.generate` runs either way, returning a `Generation` that holds
 302the tokens and the work counted.
 303
 304**Why it matters in practice.** The cache trades memory for speed, and that
 305memory is what limits how many users one GPU can serve. Section 3 does the
 306math.
 307
 308## 3. Memory math: will it fit?
 309
 310**Everyday picture.** Packing for a trip. The suitcase is GPU memory. The
 311model's weights are the big fixed items that always go in. Every active
 312conversation adds a bag of notes (its KV cache) whose size grows with the
 313conversation's length. Once the suitcase is full, the next customer waits.
 314
 315**Tiny worked example.**
 316
 317* Weights: 70×10⁹ parameters × 2 bytes (16-bit) = **140 GB**: more than one
 318  80 GB GPU. At 4 bits (half a byte) it is **35 GB** and fits on one.
 319* KV cache per token for a Llama-3-8B-shaped model (32 layers, 8 KV heads,
 320  128 dimensions per head, 16-bit): 2 × 32 × 8 × 128 × 2 = **131,072 bytes**,
 321  about 128 KB.
 322* A 32,000-token conversation: 32,000 × 131,072 ≈ **4.2 GB** of cache.
 323* An 80 GB GPU holding 16 GB of weights has 64 GB left: room for **15** such
 324  conversations at once.
 325
 326$$
 327\text{weight bytes} = P \times \frac{\text{bits}}{8} \qquad
 328\text{KV bytes per token} = 2 \times L \times H_{kv} \times d_h \times b
 329$$
 330
 331**Symbols**
 332
 333| Symbol | Meaning here | Shape / range |
 334|---|---|---|
 335| $P$ | number of parameters | e.g. 7×10¹⁰ |
 336| bits / 8 | bytes per parameter (16 bits = 2 bytes) | |
 337| $2$ | one key and one value per token | |
 338| $L$ | number of transformer layers, each with its own cache | e.g. 32 |
 339| $H_{kv}$ | number of key/value heads (fewer than query heads with grouped-query attention) | e.g. 8 |
 340| $d_h$ | dimensions per head | e.g. 128 |
 341| $b$ | bytes per stored number | 2 at 16-bit |
 342
 343**In words:** weights cost parameters times bytes each; the cache costs, for
 344every token, one key and one value per layer per KV head.
 345
 346**On the worked example:** 7×10¹⁰ × 16/8 = 1.4×10¹¹ bytes = 140 GB; and
 3472 × 32 × 8 × 128 × 2 = 131,072 bytes per token.
 348
 349**In Python:**
 350
 351```python
 352P = 7e10
 353# weight GB at 16 bits, then at 4 bits
 354P * 16 / 8 / 1e9, P * 4 / 8 / 1e9  # → (140.0, 35.0)
 355L, H_kv, d_h, b = 32, 8, 128, 2
 356# KV bytes per token: a key and a value, per layer, per KV head
 3572 * L * H_kv * d_h * b  # → 131072
 358# GB of cache for a 32,000-token conversation
 359round(32_000 * 131_072 / 1e9, 1)  # → 4.2
 360```
 361
 362![At 128k tokens, 32 KV heads need 67 GB, more than the 64 GB free, while 8 KV heads need 17 GB, so three such requests fit](figures/primer.ml.inference.kv_memory.svg)
 363
 364**Reading it:** the x-axis is context length per request; the y-axis is KV
 365cache memory for one request. The steep line is a model with 32 KV heads
 366(classic multi-head attention); the shallow one has 8 (grouped-query
 367attention, like Llama 3). The dashed line is the 64 GB left on an 80 GB GPU
 368after 16 GB of weights. With 32 KV heads a single 128k-token request would
 369not fit; with 8 it fits three times over. Doing this arithmetic out
 370loud is the fastest way to size a deployment.
 371
 372**Try it:** pick a model shape, then drag the context length and the number
 373of requests. The bar is one GPU's memory: the grey part is the weights, the
 374coloured part the KV cache. Watch how quickly a long context pushes the cache
 375past the weights, and how much further 8 KV heads go than 32.
 376
 377<div class="viz" data-viz="kv-cache" aria-label="KV cache memory calculator"></div>
 378
 379**In code:** `weight_bytes` and `kv_cache_bytes_per_token` are the two
 380formulas, `kv_cache_bytes` scales the cache to a context and batch, and
 381`max_concurrent_requests` counts how many conversations fit beside the
 382weights.
 383
 384## 4. Sampling: from scores to one token
 385
 386**Everyday picture.** Choosing where to eat. *Greedy* always picks the
 387top-rated place. *Sampling* holds a lottery weighted by rating. *Temperature*
 388is how adventurous you feel: low means you nearly always pick the favourite,
 389high means long shots get a real chance. *Top-k* says "only consider the top
 3903". *Top-p* says "consider just enough places to cover 90% of my
 391enthusiasm": one place if you have a clear favourite, several if you're torn.
 392
 393**Tiny worked example.** Scores (logits) 2, 1, 0 for three tokens:
 394
 395| temperature | probabilities |
 396|---|---|
 397| 0 (greedy) | 1, 0, 0 |
 398| 0.5 | 0.867, 0.117, 0.016 |
 399| 1 | 0.665, 0.245, 0.090 |
 400| 2 | 0.506, 0.307, 0.186 |
 401
 402With probabilities 0.5, 0.3, 0.15, 0.05: top-k = 2 keeps 0.5 and 0.3,
 403renormalised to 0.625 and 0.375. Top-p = 0.9 keeps 0.5, 0.3 and 0.15 (the
 404first set whose total reaches 0.9) and cuts the 0.05 tail.
 405
 406```mermaid
 407flowchart LR
 408  L[Logits, one per<br/>vocabulary token] --> T[Divide by<br/>temperature T]
 409  T --> S[Softmax<br/>probabilities]
 410  S --> K[Top-k: keep the<br/>k likeliest]
 411  K --> P[Top-p: keep the smallest set<br/>reaching probability p]
 412  P --> R[Renormalise<br/>and draw one token]
 413```
 414
 415**Reading it:** the model only ever produces the scores on the left;
 416everything after that is a choice *you* make at request time. Temperature
 417reshapes the whole distribution; top-k and top-p then cut off the unreliable
 418tail before the draw, so a rare nonsense token can't be picked by bad luck.
 419
 420$$
 421p_i = \frac{e^{z_i / T}}{\sum_j e^{z_j / T}}
 422$$
 423
 424**Symbols**
 425
 426| Symbol | Meaning here | Shape / range |
 427|---|---|---|
 428| $z_i$ | the model's score (logit) for vocabulary token i | any real |
 429| $T$ | temperature | > 0; T → 0 approaches greedy |
 430| $e^{x}$ | the exponential function; makes every score positive and stretches gaps | |
 431| $\sum_j$ | sum over every token j in the vocabulary, so the $p_i$ add to 1 | |
 432| $p_i$ | probability of drawing token i | 0 … 1 |
 433
 434**In words:** divide every score by the temperature, exponentiate, and
 435divide by the total so the results sum to one.
 436
 437**On the worked example:** T = 0.5 turns (2, 1, 0) into (4, 2, 0);
 438e⁴ = 54.6, e² = 7.39, e⁰ = 1, total 63.0; probabilities 0.867, 0.117, 0.016.
 439
 440**In Python:**
 441
 442```python
 443import math
 444z, T = [2.0, 1.0, 0.0], 0.5
 445# e^(z_i / T)
 446exps = [math.exp(z_i / T) for z_i in z]
 447# Σ_j e^(z_j / T)
 448round(sum(exps), 1)  # → 63.0
 449# p_i
 450[round(e / sum(exps), 3) for e in exps]  # → [0.867, 0.117, 0.016]
 451```
 452
 453![Five tokens at three temperatures: at T = 0.5 the favourite takes 79% and top-p 0.9 cuts three tokens; at T = 2 it falls to 38% and only one is cut](figures/primer.ml.inference.sampling.svg)
 454
 455**Reading it:** five candidate tokens with fixed scores, shown at three
 456temperatures. At T = 0.5 (left) the favourite takes 79%; at T = 2 (right) it
 457falls to 38% and the rest spread out, down to 7% for the least likely. The
 458hatched bars are the tokens top-p = 0.9 would cut: three at T = 0.5, two at
 459T = 1, one at T = 2. A confident model reaches 90% with fewer tokens, so
 460top-p cuts more of them; an unsure one needs more tokens to reach 90%, so it
 461cuts fewer. That's why top-p adapts where a fixed top-k can't.
 462
 463**Temperature 0 is not a determinism guarantee.** Floating-point addition
 464isn't associative: in 32-bit floats, (10⁸ + 1) − 10⁸ = 0 but (10⁸ − 10⁸) + 1 = 1.
 465On a GPU the order of additions can depend on the kernel chosen and on what
 466else is in the batch, so two nearly tied tokens can swap places between
 467otherwise identical requests.
 468
 469**In code:** `temperature_probs` applies the formula, `top_k_filter` and
 470`top_p_filter` cut the tail, and `sample_next` chains them into one draw.
 471`float32_sum` and `greedy_pick_with_summation_order` show a near tie flipping
 472with the order of additions.
 473
 474**Why it matters in practice.** Use low temperature for extraction,
 475classification and tool calls; higher for brainstorming and creative text.
 476
 477## 5. Speculative decoding: a junior drafts, a senior checks
 478
 479**Everyday picture.** A junior writer drafts the next few sentences quickly.
 480A senior editor reads the whole draft *at once*, keeps every sentence they
 481would have written themselves, rewrites the first one they wouldn't, and
 482throws away the rest. The senior's reading is fast; their writing is slow.
 483So the team moves at the junior's speed but produces the senior's text.
 484
 485This works because decode is memory-bound: checking 5 draft tokens in one
 486pass of the big model costs about the same as generating 1.
 487
 488**Tiny worked example.** The big model's next-token probabilities are
 489(0.5, 0.3, 0.2); the small model's are (0.3, 0.3, 0.4). A draft token
 490survives with probability Σ min = 0.3 + 0.3 + 0.2 = **0.8**. Drafting 4
 491tokens per round yields on average (1 − 0.8⁵) / (1 − 0.8) = **3.36** tokens
 492per pass of the big model instead of 1.
 493
 494```mermaid
 495flowchart TD
 496  D[Small model drafts γ tokens<br/>one by one, cheap] --> V[Big model scores all γ positions<br/>in ONE parallel pass]
 497  V --> C{For each draft x in order:<br/>keep with probability min 1, p/q}
 498  C -->|kept| N[Next draft]
 499  N --> C
 500  C -->|rejected| F[Replace x with a draw from<br/>max 0, p − q, renormalised. Stop.]
 501  C -->|all kept| B[Bonus: draw one more<br/>token from the big model]
 502```
 503
 504**Reading it:** the loop in the middle walks the drafts left to right. A
 505draft the big model likes at least as much as the small one did
 506(p ≥ q) is always kept; one it likes less is kept only with probability
 507p/q. The first rejection is replaced by a token drawn from exactly the
 508probability the small model *under*-proposed, which is what makes the final
 509output statistically identical to sampling from the big model alone. The
 510test suite checks this empirically over 20,000 rounds.
 511
 512$$
 513P(\text{keep } x) = \min\!\left(1, \frac{p(x)}{q(x)}\right) \qquad
 514\alpha = \sum_x \min\big(p(x), q(x)\big) \qquad
 515\mathbb{E}[\text{tokens per pass}] = \frac{1 - \alpha^{\gamma+1}}{1 - \alpha}
 516$$
 517
 518**Symbols**
 519
 520| Symbol | Meaning here | Shape / range |
 521|---|---|---|
 522| $p(x)$ | the big (target) model's probability for token x | 0 … 1 |
 523| $q(x)$ | the small (draft) model's probability for token x | 0 … 1 |
 524| $\min(a, b)$ | the smaller of a and b | |
 525| $\alpha$ | alpha, the acceptance rate: how often a draft survives | 0 … 1 |
 526| $\gamma$ | gamma, how many tokens the small model drafts per round | e.g. 4 |
 527| $\mathbb{E}[\cdot]$ | expected value: the long-run average | |
 528
 529**In words:** keep a draft with probability "how much the big model likes it
 530compared with the small one, capped at 1"; the average number of tokens per
 531big-model pass is a geometric series in the acceptance rate.
 532
 533**On the worked example:** α = 0.8, γ = 4: (1 − 0.8⁵)/(1 − 0.8) =
 534(1 − 0.328)/0.2 = 3.36.
 535
 536**In Python:**
 537
 538```python
 539# big model
 540p = [0.5, 0.3, 0.2]
 541# small model
 542q = [0.3, 0.3, 0.4]
 543# P(keep x) for each token
 544[round(min(1.0, p_x / q_x), 2) for p_x, q_x in zip(p, q)]  # → [1.0, 1.0, 0.5]
 545alpha = sum(min(p_x, q_x) for p_x, q_x in zip(p, q))
 546round(alpha, 2)  # → 0.8
 547gamma = 4
 548# expected tokens per big-model pass
 549round((1 - alpha ** (gamma + 1)) / (1 - alpha), 2)  # → 3.36
 550```
 551
 552**In code:** `acceptance_rate` computes α and `expected_tokens_per_round` the
 553geometric series; `speculative_round` runs the draft, verify and replace
 554loop once, and `speculative_generate` repeats it until enough tokens exist.
 555
 556## 6. Quantization: fewer bits per weight
 557
 558**Everyday picture.** Writing prices to the nearest dollar instead of the
 559nearest cent: shorter to store, slightly less exact. Using one ruler per
 560row of a spreadsheet, instead of one for the whole sheet, keeps a single
 561huge number in one row from making every other row coarse.
 562
 563**Tiny worked example.** A row of weights (0.5, −1.27, 0.02).
 564
 565* **int8**: the largest magnitude, 1.27, maps to 127, so the step size is
 566  0.01. Codes: **50, −127, 2**. Decoding gives back exactly
 567  (0.5, −1.27, 0.02).
 568* **int4**: only 15 levels (−7 … 7), step size 1.27/7 = 0.181. Codes:
 569  **3, −7, 0**. Decoding gives (0.544, −1.27, 0): the small weight vanished.
 570
 571$$
 572s = \frac{\max_j |w_j|}{2^{\,\text{bits}-1} - 1} \qquad
 573c_j = \operatorname{round}\!\left(\frac{w_j}{s}\right) \qquad
 574\hat{w}_j = s \, c_j
 575$$
 576
 577**Symbols**
 578
 579| Symbol | Meaning here | Shape / range |
 580|---|---|---|
 581| $w_j$ | the j-th original weight in the row | real |
 582| $\max_j |w_j|$ | the largest absolute value in the row | ≥ 0 |
 583| $2^{\text{bits}-1} - 1$ | the largest code: 127 for 8 bits, 7 for 4 bits | |
 584| $s$ | the scale (step size), one per row | > 0 |
 585| $c_j$ | the stored integer code | −127…127 or −7…7 |
 586| $\hat{w}_j$ | the weight as reconstructed at inference time ("w-hat") | real |
 587
 588**In words:** pick a step size so the largest weight lands on the largest
 589code, store each weight as the nearest whole number of steps, and multiply
 590back at run time.
 591
 592**On the worked example:** int4: s = 1.27/7 = 0.181; 0.5/0.181 = 2.76 → 3;
 5933 × 0.181 = 0.544.
 594
 595**In Python:**
 596
 597```python
 598w = [0.5, -1.27, 0.02]
 599def quantize(w, bits):
 600    # the step size
 601    s = max(abs(w_j) for w_j in w) / (2 ** (bits - 1) - 1)
 602    # c_j: whole steps
 603    return s, [round(w_j / s) for w_j in w]
 604s, c = quantize(w, bits=8)
 605round(s, 3), c  # → (0.01, [50, -127, 2])
 606s, c = quantize(w, bits=4)
 607round(s, 3), c  # → (0.181, [3, -7, 0])
 608# ŵ_j = s · c_j: the 0.02 is gone
 609[round(s * c_j, 3) for c_j in c]  # → [0.544, -1.27, 0.0]
 610```
 611
 612**In code:** `quantize` returns the integer codes and one scale per row,
 613`dequantize` multiplies them back, and `quantization_error` measures how far
 614the round trip lands from the original weights.
 615
 616**Why it matters in practice.** 8-bit weights are nearly lossless; 4-bit
 617methods with smarter rounding (GPTQ, AWQ) keep most quality at a quarter
 618of the memory. Fewer bytes per weight also means faster memory-bound decode.
 619
 620## 7. Continuous batching: seat the next party as soon as a table frees
 621
 622**Everyday picture.** A restaurant with two tables. *Static batching* seats
 623two parties and won't seat anyone new until *both* have left, so a table
 624sits empty while one slow diner lingers. *Continuous batching* seats the next
 625party the moment any table frees.
 626
 627**Tiny worked example.** Four requests needing 4, 1, 1 and 1 decode steps,
 628two slots. Static: {4, 1} runs 4 steps with one slot idle for 3, then
 629{1, 1} runs 1: **5 steps, 70%** of slot-steps busy. Continuous: the short
 630requests slide into the freed slot while the long one runs: **4 steps,
 63187.5%** busy.
 632
 633![With 32 requests on 8 slots, static batching leaves idle gaps and needs 433 steps at 60% busy; continuous batching needs 318 steps at 82%](figures/primer.ml.inference.batching.svg)
 634
 635**Reading it:** each row is a GPU batch slot and each column is one decode
 636step; colour identifies the request occupying the slot, and white is an idle
 637slot. On a realistic mix of 32 requests of varied length over 8 slots,
 638static batching (top) leaves white holes wherever a short request finished
 639early; continuous batching (bottom) keeps nearly every cell busy and
 640finishes the same work in fewer steps (about 82% vs. 60% utilisation).
 641
 642$$
 643U = \frac{\text{useful slot-steps}}{\text{steps} \times \text{slots}}
 644$$
 645
 646**Symbols**
 647
 648| Symbol | Meaning here | Shape / range |
 649|---|---|---|
 650| $U$ | utilisation: the share of slot-steps doing real work | 0 … 1 |
 651| useful slot-steps | total decode steps all requests need | integer |
 652| steps × slots | the capacity the GPU offered while serving them | integer |
 653
 654**In words:** utilisation is the work the requests needed divided by the
 655work capacity the server spent serving them.
 656
 657**On the worked example:** 7 useful slot-steps; static 7/(5×2) = 0.70,
 658continuous 7/(4×2) = 0.875.
 659
 660**In Python:**
 661
 662```python
 663# decode steps the four requests need
 664useful = 4 + 1 + 1 + 1
 665slots = 2
 666# static takes 5 steps, continuous 4
 667useful / (5 * slots), useful / (4 * slots)  # → (0.7, 0.875)
 668```
 669
 670**In code:** `simulate_static_batching` and `simulate_continuous_batching`
 671play out the two policies step by step, each returning a `ServingRun` that
 672holds the slot timeline and its utilisation U.
 673
 674**Why it matters in practice.** Continuous batching (together with paged
 675KV-cache memory, as in vLLM) is a large part of why modern inference servers
 676reach high throughput.
 677
 678## 8. Prompt caching: reuse the prefill of a shared prefix
 679
 680**Everyday picture.** A kitchen that pre-chops the ingredients every order
 681uses, instead of chopping them again for each plate.
 682
 683The KV cache normally lives for one request. Prompt caching keeps it
 684*across* requests: if many calls start with the same long system prompt,
 685tool definitions or reference document, the provider stores that prefix's
 686keys and values and skips its prefill next time. It is only valid up to the
 687first differing token, because every token's keys depend on everything
 688before it, so **put stable content first and volatile content last**.
 689
 690**Tiny worked example.** 100 requests, each a 10,000-token shared prefix
 691plus a 500-token question, at \$3 per million input tokens, with cache
 692writes at 1.25× and cache reads at 0.1× (typical of providers; check
 693current pricing):
 694
 695* Without caching: 100 × 10,500 × \$3/10⁶ = **\$3.15**.
 696* With caching: the first call writes the cache (\$0.039); each of the other
 697  99 costs \$0.0045. Total **\$0.48**, an 85% saving, and each cached call
 698  also skips 10,000 tokens of prefill, so its first token arrives sooner.
 699
 700```mermaid
 701flowchart LR
 702  subgraph Prompt["One request's prompt, in order"]
 703    S[System prompt<br/>stable] --> T[Tool definitions<br/>stable] --> D[Reference docs<br/>stable] --> Q[User question<br/>changes every call]
 704  end
 705  S & T & D -.->|cached after the first call| C[(Prefix KV cache)]
 706```
 707
 708**Reading it:** the prompt is read left to right, and the cache can cover
 709only an unbroken run from the very start. Stable parts go first so that
 710every request shares the longest possible prefix; the part that changes on
 711every call goes last. A timestamp or request ID placed at the top would
 712break the match on the first token and silently disable caching.
 713
 714**In code:** `reusable_prefix_tokens` counts how many leading tokens two
 715prompts share, and `prompt_cache_cost` prices a run of requests with and
 716without the cache.
 717
 718## In 20 seconds
 719- Prefill reads the prompt in parallel (compute-bound, sets time to first
 720  token); decode writes one token at a time (memory-bound, sets tokens per
 721  second).
 722- The KV cache stores every earlier token's keys and values so each step
 723  computes one token; it trades GPU memory for speed.
 724- Memory math: weights = parameters × bytes; KV per token = 2 × layers ×
 725  KV heads × head dim × bytes. 70B at 16-bit is 140 GB; 32k tokens of
 726  Llama-3-8B cache is about 4 GB.
 727- Temperature reshapes the distribution; top-k and top-p cut the tail.
 728  Temperature 0 reduces but does not guarantee determinism.
 729- Speculative decoding: a small model drafts, the big one verifies in one
 730  pass; output is identical in distribution.
 731- Quantization, continuous batching and prompt caching are the other big
 732  serving levers.
 733
 734## Self-test questions
 735
 736**Q: Why is decoding slow even on a huge GPU?**
 737A: Each token needs every weight read from memory but does only about one
 738operation per byte read, far below the GPU's break-even (~300 FLOPs/byte).
 739The arithmetic units mostly wait on memory.
 740
 741**Q: What is the KV cache, and why does it matter for serving cost?**
 742A: Stored keys and values of all earlier tokens, so each new token is
 743computed once instead of recomputing the whole sequence. It grows with
 744context length and batch size, and that memory caps how many requests a GPU
 745serves at once.
 746
 747**Q: How much memory does a 70B model need at 16-bit and at 4-bit?**
 748A: 140 GB and 35 GB for weights alone, plus KV cache and activations.
 749
 750**Q: Estimate the KV cache for one 32k-token request on a Llama-3-8B-shaped
 751model.**
 752A: 2 × 32 × 8 × 128 × 2 = 131,072 bytes per token; × 32,000 ≈ 4.2 GB.
 753
 754**Q: Why does grouped-query attention make serving cheaper?**
 755A: It shares each key/value head among several query heads, shrinking the
 756KV cache (4× for 32 → 8 KV heads), so more requests fit per GPU.
 757
 758**Q: Why isn't temperature 0 perfectly deterministic?**
 759A: Floating-point addition isn't associative, and GPU reduction order can
 760change with kernels and batch composition, so nearly tied logits can flip.
 761
 762**Q: How can speculative decoding be faster yet produce the same
 763distribution?**
 764A: Verifying several draft tokens costs one memory-bound pass of the big
 765model, about the same as generating one. The min(1, p/q) accept rule plus
 766resampling from max(0, p − q) on rejection makes the output exactly the big
 767model's distribution.
 768
 769**Q: What does continuous batching fix?**
 770A: Idle slots: static batches wait for their longest request, while
 771continuous batching refills any freed slot immediately, raising utilisation.
 772
 773**Q: How do you structure a prompt to benefit from prompt caching?**
 774A: Put stable content (system prompt, tool definitions, reference documents)
 775first and anything that varies (the question, timestamps, IDs) last, because
 776a cache is valid only up to the first differing token.
 777
 778## The papers behind this lesson
 779
 780- Leviathan, Kalman & Matias, *Fast Inference from Transformers via
 781  Speculative Decoding* (2022): https://arxiv.org/abs/2211.17192. Introduced
 782  the draft-then-verify scheme with an accept/reject rule that provably
 783  preserves the target model's output distribution.
 784  [annotated companion](../../papers/speculative-decoding.html)
 785- Kwon et al., *Efficient Memory Management for Large Language Model Serving
 786  with PagedAttention* (2023): https://arxiv.org/abs/2309.06180. Stored the KV
 787  cache in fixed-size pages like an operating system's virtual memory,
 788  eliminating fragmentation and enabling vLLM's high-throughput continuous
 789  batching. [annotated companion](../../papers/paged-attention.html)
 790- Dao et al., *FlashAttention* (2022): https://arxiv.org/abs/2205.14135.
 791  Computed exact attention in GPU-memory-sized tiles, cutting the slow
 792  memory traffic that dominates long-context inference.
 793  [annotated companion](../../papers/flashattention.html)
 794- Dettmers et al., *LLM.int8(): 8-bit Matrix Multiplication for Transformers
 795  at Scale* (2022): https://arxiv.org/abs/2208.07339. Showed that a few
 796  outlier features break naive 8-bit quantization and handled them
 797  separately.
 798- Frantar et al., *GPTQ: Accurate Post-Training Quantization for Generative
 799  Pre-trained Transformers* (2022): https://arxiv.org/abs/2210.17323. Made
 800  3-4-bit weight quantization practical with error-compensating rounding.
 801- Holtzman et al., *The Curious Case of Neural Text Degeneration* (2019):
 802  https://arxiv.org/abs/1904.09751. Introduced nucleus (top-p) sampling.
 803
 804## Further reading
 805- vLLM documentation: https://docs.vllm.ai/
 806- Hugging Face, *Text generation strategies*: https://huggingface.co/docs/transformers/generation_strategies
 807- Anthropic, *Prompt caching*: https://docs.claude.com/en/docs/build-with-claude/prompt-caching
 808- Thinking Machines Lab, *Defeating Nondeterminism in LLM Inference*: https://thinkingmachines.ai/blog/defeating-nondeterminism-in-llm-inference/
 809- Grattafiori et al., *The Llama 3 Herd of Models* (2024): https://arxiv.org/abs/2407.21783
 810"""
 811
 812from __future__ import annotations
 813
 814import numpy as np
 815
 816from primer._show import banner, matrix, say, table, takeaway
 817
 818# Roughly a current datacenter GPU (H100-class): ~1 PFLOP/s of dense 16-bit
 819# math and ~3.35 TB/s of memory bandwidth. Change these to model other hardware.
 820PEAK_FLOPS = 1e15
 821HBM_BANDWIDTH = 3.35e12
 822
 823# ---------------------------------------------------------------------------
 824# 1. Memory math: weights and the KV cache
 825# ---------------------------------------------------------------------------
 826
 827
 828def weight_bytes(params: float, bits: int = 16) -> float:
 829    """Memory for the weights alone: parameters × bytes per parameter."""
 830    return params * bits / 8
 831
 832
 833def kv_cache_bytes_per_token(layers: int, kv_heads: int, head_dim: int, bits: int = 16) -> int:
 834    """2 (a key and a value) × layers × KV heads × head dimension × bytes.
 835
 836    Every layer stores one key vector and one value vector per KV head for
 837    every token in the context, so the cache grows by this much per token.
 838    """
 839    return int(2 * layers * kv_heads * head_dim * bits // 8)
 840
 841
 842def kv_cache_bytes(context: int, layers: int, kv_heads: int, head_dim: int, bits: int = 16, batch: int = 1) -> int:
 843    """KV cache for `batch` requests of `context` tokens each."""
 844    return batch * context * kv_cache_bytes_per_token(layers, kv_heads, head_dim, bits)
 845
 846
 847def max_concurrent_requests(gpu_gb: float, weights_gb: float, context: int, layers: int, kv_heads: int, head_dim: int, bits: int = 16) -> int:
 848    """How many full-context requests fit in the memory left after the weights."""
 849    free = (gpu_gb - weights_gb) * 1e9
 850    return int(free // kv_cache_bytes(context, layers, kv_heads, head_dim, bits))
 851
 852
 853# ---------------------------------------------------------------------------
 854# 2. Prefill vs. decode: compute-bound vs. memory-bound
 855# ---------------------------------------------------------------------------
 856
 857
 858def arithmetic_intensity(tokens_in_parallel: int, bits: int = 16) -> float:
 859    """FLOPs done per byte of weights read from memory.
 860
 861    Each weight is read once per forward pass and used for one multiply-add
 862    (2 FLOPs) *per token processed in that pass*. Prefill processes the whole
 863    prompt in one pass; decode processes one new token per pass.
 864    """
 865    return 2 * tokens_in_parallel / (bits / 8)
 866
 867
 868def ridge_point(flops_per_second: float = PEAK_FLOPS, bytes_per_second: float = HBM_BANDWIDTH) -> float:
 869    """The intensity at which math and memory take equally long (the roofline's ridge)."""
 870    return flops_per_second / bytes_per_second
 871
 872
 873def bottleneck(tokens_in_parallel: int, bits: int = 16) -> str:
 874    """Below the ridge point the GPU waits on memory; above it, on arithmetic."""
 875    return "compute-bound" if arithmetic_intensity(tokens_in_parallel, bits) >= ridge_point() else "memory-bound"
 876
 877
 878def decode_seconds_per_token(params: float, bits: int = 16, bytes_per_second: float = HBM_BANDWIDTH) -> float:
 879    """Lower bound on time per generated token at batch size 1: read every weight once."""
 880    return weight_bytes(params, bits) / bytes_per_second
 881
 882
 883def prefill_seconds(params: float, n_tokens: int, flops_per_second: float = PEAK_FLOPS) -> float:
 884    """Lower bound on prefill time: 2 FLOPs per parameter per prompt token."""
 885    return 2 * params * n_tokens / flops_per_second
 886
 887
 888# ---------------------------------------------------------------------------
 889# 3. The KV cache, on a tiny but complete decoder
 890# ---------------------------------------------------------------------------
 891
 892
 893def _rms_norm(x: np.ndarray) -> np.ndarray:
 894    # Rescale each token vector to unit root-mean-square (no learned gain, for brevity).
 895    return x / np.sqrt(np.mean(x**2, axis=-1, keepdims=True) + 1e-6)
 896
 897
 898def _softmax(x: np.ndarray) -> np.ndarray:
 899    e = np.exp(x - x.max(axis=-1, keepdims=True))
 900    return e / e.sum(axis=-1, keepdims=True)
 901
 902
 903class Generation:
 904    """Result of `TinyDecoder.generate`: the new tokens and how much work they took."""
 905
 906    def __init__(self, tokens: list[int], positions_processed: int, flops: int):
 907        self.tokens = tokens
 908        self.positions_processed = positions_processed
 909        self.flops = flops
 910
 911
 912class TinyDecoder:
 913    """A 2-layer decoder-only transformer with random weights, small enough to read.
 914
 915    Each layer: RMS-norm -> multi-head causal self-attention -> residual add,
 916    then RMS-norm -> ReLU feed-forward -> residual add. Token embeddings plus
 917    learned position embeddings go in; the embedding matrix is reused to turn
 918    the final vectors into one score (logit) per vocabulary token.
 919
 920    Two ways to run it:
 921      * `forward_full(tokens)`: process a whole sequence at once (prefill, or
 922        naive generation that recomputes everything every step).
 923      * `forward_step(token, pos, cache)`: process ONE new token, reading the
 924        keys and values of all earlier tokens from `cache` instead of
 925        recomputing them. That is the KV cache.
 926    """
 927
 928    def __init__(self, vocab: int = 16, d_model: int = 16, n_heads: int = 2, n_layers: int = 2, max_len: int = 64, seed: int = 0):
 929        rng = np.random.default_rng(seed)
 930        s = 1 / np.sqrt(d_model)
 931        self.vocab, self.d, self.h, self.dh = vocab, d_model, n_heads, d_model // n_heads
 932        self.embed = rng.normal(0, 1, (vocab, d_model))
 933        self.pos = rng.normal(0, 0.5, (max_len, d_model))
 934        self.layers = [
 935            dict(
 936                Wq=rng.normal(0, s, (d_model, d_model)), Wk=rng.normal(0, s, (d_model, d_model)),
 937                Wv=rng.normal(0, s, (d_model, d_model)), Wo=rng.normal(0, s, (d_model, d_model)),
 938                W1=rng.normal(0, s, (d_model, 4 * d_model)), W2=rng.normal(0, s / 2, (4 * d_model, d_model)),
 939            )
 940            for _ in range(n_layers)
 941        ]
 942        self.flops = 0  # multiply-adds × 2, counted by _mm
 943
 944    def _mm(self, a: np.ndarray, b: np.ndarray) -> np.ndarray:
 945        # Matrix multiply that also counts its cost: 2 FLOPs per multiply-add.
 946        out = a @ b
 947        self.flops += 2 * a.shape[-1] * out.size
 948        return out
 949
 950    def _heads(self, x: np.ndarray) -> np.ndarray:
 951        # (seq, d_model) -> (heads, seq, d_head)
 952        return x.reshape(x.shape[0], self.h, self.dh).transpose(1, 0, 2)
 953
 954    def _attend(self, q: np.ndarray, k: np.ndarray, v: np.ndarray, causal: bool) -> np.ndarray:
 955        # q: (heads, n_q, d_head); k, v: (heads, n_k, d_head). Returns (n_q, d_model).
 956        scores = self._mm(q, k.transpose(0, 2, 1)) / np.sqrt(self.dh)
 957        if causal:
 958            n_q, n_k = scores.shape[-2:]
 959            # Query i sits at absolute position n_k - n_q + i and may see keys up to there.
 960            allowed = np.arange(n_k)[None, :] <= (np.arange(n_q)[:, None] + n_k - n_q)
 961            scores = np.where(allowed, scores, -np.inf)
 962        out = self._mm(_softmax(scores), v)  # (heads, n_q, d_head)
 963        return out.transpose(1, 0, 2).reshape(q.shape[1], self.d)
 964
 965    def _block(self, x: np.ndarray, layer: dict, k_all: np.ndarray, v_all: np.ndarray) -> np.ndarray:
 966        q = self._heads(self._mm(_rms_norm(x), layer["Wq"]))
 967        x = x + self._mm(self._attend(q, k_all, v_all, causal=True), layer["Wo"])
 968        return x + self._mm(np.maximum(0, self._mm(_rms_norm(x), layer["W1"])), layer["W2"])
 969
 970    def _kv(self, x: np.ndarray, layer: dict) -> tuple[np.ndarray, np.ndarray]:
 971        n = _rms_norm(x)
 972        return self._heads(self._mm(n, layer["Wk"])), self._heads(self._mm(n, layer["Wv"]))
 973
 974    def forward_full(self, tokens: np.ndarray, cache: list[dict] | None = None) -> np.ndarray:
 975        """Logits (seq, vocab) for every position, computing everything from scratch.
 976
 977        Pass an empty `cache` to fill it in the same pass: that is *prefill*,
 978        processing the whole prompt in parallel and keeping every key and value.
 979        """
 980        x = self.embed[tokens] + self.pos[: len(tokens)]
 981        for i, layer in enumerate(self.layers):
 982            k, v = self._kv(x, layer)
 983            if cache is not None:
 984                cache[i]["k"], cache[i]["v"] = k, v
 985            x = self._block(x, layer, k, v)
 986        return self._mm(_rms_norm(x), self.embed.T)
 987
 988    def new_cache(self) -> list[dict]:
 989        """One empty key store and value store per layer, each (heads, 0, d_head)."""
 990        empty = np.zeros((self.h, 0, self.dh))
 991        return [dict(k=empty, v=empty) for _ in self.layers]
 992
 993    def forward_step(self, token: int, pos: int, cache: list[dict]) -> np.ndarray:
 994        """Logits (vocab,) for ONE new token, appending its key and value to `cache`."""
 995        x = self.embed[[token]] + self.pos[[pos]]  # (1, d_model)
 996        for layer, store in zip(self.layers, cache):
 997            k_new, v_new = self._kv(x, layer)  # only the new token's K and V are computed
 998            store["k"] = np.concatenate([store["k"], k_new], axis=1)
 999            store["v"] = np.concatenate([store["v"], v_new], axis=1)
1000            x = self._block(x, layer, store["k"], store["v"])
1001        return self._mm(_rms_norm(x), self.embed.T)[0]
1002
1003    def generate(self, prompt: list[int], n_new: int, use_cache: bool = True) -> Generation:
1004        """Greedy generation of `n_new` tokens, with or without the KV cache."""
1005        self.flops = 0
1006        seq, out, processed = list(prompt), [], 0
1007        if use_cache:
1008            cache = self.new_cache()
1009            logits = self.forward_full(np.array(seq), cache)[-1]  # prefill: whole prompt, one pass
1010            processed += len(seq)
1011            for i in range(n_new):
1012                tok = int(np.argmax(logits))
1013                out.append(tok)
1014                seq.append(tok)
1015                if i < n_new - 1:  # decode: one new position per step, earlier K/V read from the cache
1016                    logits = self.forward_step(tok, len(seq) - 1, cache)
1017                    processed += 1
1018        else:
1019            for _ in range(n_new):
1020                logits = self.forward_full(np.array(seq))[-1]  # recompute the whole sequence each time
1021                processed += len(seq)
1022                tok = int(np.argmax(logits))
1023                out.append(tok)
1024                seq.append(tok)
1025        return Generation(out, processed, self.flops)
1026
1027
1028# ---------------------------------------------------------------------------
1029# 4. Sampling: turning scores into one next token
1030# ---------------------------------------------------------------------------
1031
1032
1033def temperature_probs(logits: np.ndarray, temperature: float) -> np.ndarray:
1034    """softmax(logits / T). T < 1 sharpens, T > 1 flattens, T = 0 means greedy (argmax)."""
1035    if temperature == 0:
1036        out = np.zeros_like(logits, dtype=float)
1037        out[int(np.argmax(logits))] = 1.0
1038        return out
1039    return _softmax(logits / temperature)
1040
1041
1042def top_k_filter(probs: np.ndarray, k: int) -> np.ndarray:
1043    """Keep the k most likely tokens, zero the rest, renormalise."""
1044    keep = np.argsort(probs)[::-1][:k]
1045    out = np.zeros_like(probs)
1046    out[keep] = probs[keep]
1047    return out / out.sum()
1048
1049
1050def top_p_filter(probs: np.ndarray, p: float) -> np.ndarray:
1051    """Nucleus sampling: keep the smallest set of top tokens whose total reaches p."""
1052    order = np.argsort(probs)[::-1]
1053    cumulative = np.cumsum(probs[order])
1054    # Include the token that crosses p: count how many cumulative sums are still below p, plus one.
1055    n_keep = int(np.searchsorted(cumulative, p - 1e-12) + 1)
1056    out = np.zeros_like(probs)
1057    out[order[:n_keep]] = probs[order[:n_keep]]
1058    return out / out.sum()
1059
1060
1061def sample_next(logits: np.ndarray, rng: np.random.Generator, temperature: float = 1.0, top_k: int | None = None, top_p: float | None = None) -> int:
1062    """The usual pipeline: temperature, then top-k, then top-p, then draw one token."""
1063    probs = temperature_probs(logits, temperature)
1064    if top_k is not None:
1065        probs = top_k_filter(probs, top_k)
1066    if top_p is not None:
1067        probs = top_p_filter(probs, top_p)
1068    return int(rng.choice(len(probs), p=probs))
1069
1070
1071def float32_sum(values: list[float]) -> float:
1072    """Add values left to right in 32-bit floats, the way one GPU reduction order might."""
1073    total = np.float32(0.0)
1074    for v in values:
1075        total = np.float32(total + np.float32(v))
1076    return float(total)
1077
1078
1079def greedy_pick_with_summation_order(contributions: list[float], other_logit: float) -> str:
1080    """Token A's logit is a float32 sum of `contributions`; token B's is `other_logit`.
1081
1082    On a GPU, the order of a reduction depends on kernel choice and on what
1083    else is in the batch. When two logits are nearly tied, that order alone
1084    can change which token wins, even at temperature 0.
1085    """
1086    return "A" if float32_sum(contributions) > other_logit else "B"
1087
1088
1089# ---------------------------------------------------------------------------
1090# 5. Speculative decoding: a small model drafts, the big model verifies
1091# ---------------------------------------------------------------------------
1092
1093
1094def _draw(p: np.ndarray, rng: np.random.Generator) -> int:
1095    # Inverse-CDF sampling: faster than rng.choice for many tiny draws.
1096    return int(min(np.searchsorted(np.cumsum(p), rng.random(), side="right"), len(p) - 1))
1097
1098
1099def acceptance_rate(p: np.ndarray, q: np.ndarray) -> float:
1100    """Probability a token drawn from the draft q survives verification against target p: Σ min(p, q)."""
1101    return float(np.minimum(p, q).sum())
1102
1103
1104def expected_tokens_per_round(alpha: float, n_draft: int) -> float:
1105    """Expected tokens per target pass when each draft is accepted with probability alpha.
1106
1107    Accepting i drafts in a row has probability alpha^i; every round also
1108    yields one token from the target (a correction or a bonus). Summing gives
1109    1 + alpha + ... + alpha^n_draft = (1 - alpha^(n+1)) / (1 - alpha).
1110    """
1111    if alpha == 1:
1112        return n_draft + 1.0
1113    return (1 - alpha ** (n_draft + 1)) / (1 - alpha)
1114
1115
1116def speculative_round(last_token: int, target: np.ndarray, draft: np.ndarray, n_draft: int, rng: np.random.Generator) -> list[int]:
1117    """One round of speculative decoding for a toy bigram language.
1118
1119    `target[i]` and `draft[i]` are next-token distributions after token i.
1120    1. The draft model proposes n_draft tokens, one after another (cheap).
1121    2. The target model scores every proposed position. In a real system
1122       this is ONE parallel forward pass, costing about the same as
1123       generating a single token, because decode is memory-bound.
1124    3. Walk the drafts left to right. Accept draft x with probability
1125       min(1, p(x) / q(x)). On the first rejection, draw a replacement from
1126       the leftover distribution max(0, p - q), renormalised, and stop.
1127    4. If every draft survives, draw one bonus token from the target.
1128    The accept/reject rule guarantees the output has exactly the target's
1129    distribution; the draft only changes *speed*.
1130    """
1131    drafts, qs, prev = [], [], last_token
1132    for _ in range(n_draft):
1133        q = draft[prev]
1134        x = _draw(q, rng)
1135        drafts.append(x)
1136        qs.append(q)
1137        prev = x
1138    ps = [target[last_token]] + [target[x] for x in drafts]  # the target's view of every position
1139    out = []
1140    for i, x in enumerate(drafts):
1141        if rng.random() < min(1.0, ps[i][x] / qs[i][x]):
1142            out.append(x)
1143            continue
1144        leftover = np.maximum(ps[i] - qs[i], 0)
1145        out.append(_draw(leftover / leftover.sum(), rng))
1146        return out
1147    out.append(_draw(ps[n_draft], rng))
1148    return out
1149
1150
1151def speculative_generate(start: int, target: np.ndarray, draft: np.ndarray, n_tokens: int, n_draft: int, rng: np.random.Generator) -> list[int]:
1152    """Run speculative rounds until at least `n_tokens` tokens exist; return the first n_tokens."""
1153    out: list[int] = []
1154    last = start
1155    while len(out) < n_tokens:
1156        new = speculative_round(last, target, draft, n_draft, rng)
1157        out += new
1158        last = new[-1]
1159    return out[:n_tokens]
1160
1161
1162# ---------------------------------------------------------------------------
1163# 6. Quantization: fewer bits per weight
1164# ---------------------------------------------------------------------------
1165
1166
1167def quantize(W: np.ndarray, bits: int = 8, per_channel: bool = True) -> tuple[np.ndarray, np.ndarray]:
1168    """Symmetric integer quantization: w ≈ scale × code, code in [-(2^(b-1)-1), 2^(b-1)-1].
1169
1170    Per-channel means one scale per output row, so an outlier in one row
1171    doesn't coarsen every other row. Returns (integer codes, scales).
1172    """
1173    qmax = 2 ** (bits - 1) - 1  # 127 for int8, 7 for int4
1174    absmax = np.abs(W).max(axis=1, keepdims=True) if per_channel else np.abs(W).max(keepdims=True)
1175    scales = absmax / qmax
1176    codes = np.clip(np.round(W / scales), -qmax, qmax).astype(np.int32)
1177    return codes, scales
1178
1179
1180def dequantize(codes: np.ndarray, scales: np.ndarray) -> np.ndarray:
1181    return codes * scales
1182
1183
1184def quantization_error(W: np.ndarray, bits: int = 8, per_channel: bool = True) -> float:
1185    """Relative error ||W - dequant(quant(W))|| / ||W||."""
1186    codes, scales = quantize(W, bits, per_channel)
1187    return float(np.linalg.norm(W - dequantize(codes, scales)) / np.linalg.norm(W))
1188
1189
1190# ---------------------------------------------------------------------------
1191# 7. Serving many users: continuous batching and prompt caching
1192# ---------------------------------------------------------------------------
1193
1194
1195class ServingRun:
1196    """Result of a batching simulation."""
1197
1198    def __init__(self, steps: int, useful_slot_steps: int, slots: int, timeline: list[list[int | None]]):
1199        self.steps = steps
1200        self.utilization = useful_slot_steps / (steps * slots)
1201        self.timeline = timeline  # timeline[step][slot] = request id, or None if idle
1202
1203
1204def simulate_static_batching(lengths: list[int], slots: int) -> ServingRun:
1205    """Fill every slot, run until the LONGEST request in the batch finishes, repeat.
1206
1207    Finished slots sit idle until the whole batch is done.
1208    """
1209    timeline: list[list[int | None]] = []
1210    for start in range(0, len(lengths), slots):
1211        batch = list(range(start, min(start + slots, len(lengths))))
1212        for t in range(max(lengths[i] for i in batch)):
1213            row = [i if t < lengths[i] else None for i in batch]
1214            timeline.append(row + [None] * (slots - len(row)))
1215    return ServingRun(len(timeline), sum(lengths), slots, timeline)
1216
1217
1218def simulate_continuous_batching(lengths: list[int], slots: int) -> ServingRun:
1219    """After every decode step, hand any freed slot to the next waiting request."""
1220    queue = list(range(len(lengths)))
1221    remaining = list(lengths)
1222    active: list[int | None] = [None] * slots
1223    timeline: list[list[int | None]] = []
1224    while queue or any(a is not None for a in active):
1225        for s in range(slots):  # refill free slots before the step
1226            if active[s] is None and queue:
1227                active[s] = queue.pop(0)
1228        timeline.append(list(active))
1229        for s, r in enumerate(active):  # one decode step for every active request
1230            if r is not None:
1231                remaining[r] -= 1
1232                if remaining[r] == 0:
1233                    active[s] = None
1234    return ServingRun(len(timeline), sum(lengths), slots, timeline)
1235
1236
1237def reusable_prefix_tokens(previous: list[int], new: list[int]) -> int:
1238    """How many leading tokens two prompts share: the part whose KV cache can be reused.
1239
1240    Attention makes every token's keys and values depend on ALL earlier
1241    tokens, so the cache is valid only up to the first difference.
1242    """
1243    n = 0
1244    for a, b in zip(previous, new):
1245        if a != b:
1246            break
1247        n += 1
1248    return n
1249
1250
1251def prompt_cache_cost(prefix: int, suffix: int, requests: int, price_per_million: float,
1252                      write_multiplier: float = 1.25, read_multiplier: float = 0.1) -> tuple[float, float]:
1253    """Input cost of `requests` calls sharing a `prefix`, without and with prompt caching.
1254
1255    The multipliers are typical of providers that price cache writes at a
1256    small premium and cache reads at a steep discount; check your provider's
1257    current pricing. Returns (uncached dollars, cached dollars).
1258    """
1259    per_token = price_per_million / 1e6
1260    uncached = requests * (prefix + suffix) * per_token
1261    first = (prefix * write_multiplier + suffix) * per_token
1262    later = (prefix * read_multiplier + suffix) * per_token
1263    return uncached, first + (requests - 1) * later
1264
1265
1266# ---------------------------------------------------------------------------
1267# 8. Figures (rendered into docs/figures by `make figures`)
1268# ---------------------------------------------------------------------------
1269
1270EXAMPLE_PROMPT = [3, 1, 4, 1, 5, 9, 2, 6]
1271
1272
1273def _cumulative_flops(use_cache: bool, n_new: int = 32) -> list[int]:
1274    return [TinyDecoder(seed=0).generate(EXAMPLE_PROMPT, n, use_cache).flops for n in range(1, n_new + 1)]
1275
1276
1277def _realistic_lengths(n: int = 32, seed: int = 0) -> list[int]:
1278    return [int(x) for x in np.random.default_rng(seed).integers(5, 120, n)]
1279
1280
1281def viz_data() -> dict:
1282    """The numbers the site's interactive KV-cache calculator starts from."""
1283    # Shapes of openly published models; the widget recomputes everything
1284    # else with the same formulas as kv_cache_bytes and weight_bytes.
1285    return {
1286        "kv-cache": {
1287            "gpu_bytes": 80e9,
1288            "presets": [
1289                {"name": "8B, 8 KV heads (Llama 3 8B shape)", "params": 8e9, "layers": 32, "kv_heads": 8, "head_dim": 128},
1290                {"name": "7B, 32 KV heads (no grouped-query attention)", "params": 7e9, "layers": 32, "kv_heads": 32, "head_dim": 128},
1291                {"name": "70B, 8 KV heads (Llama 3 70B shape)", "params": 70e9, "layers": 80, "kv_heads": 8, "head_dim": 128},
1292            ],
1293        }
1294    }
1295
1296
1297def figures() -> dict:
1298    """Data figures for this lesson, keyed by the name used in the docstring."""
1299    import matplotlib
1300
1301    matplotlib.use("Agg")
1302    import matplotlib.pyplot as plt
1303
1304    figs = {}
1305
1306    # Roofline.
1307    x = np.logspace(-0.5, 4, 300)
1308    roof = np.minimum(PEAK_FLOPS, HBM_BANDWIDTH * x)
1309    fig, ax = plt.subplots(figsize=(6.5, 4.2))
1310    ax.loglog(x, roof, lw=2, color="black")
1311    for n, label in [(1, "decode, 1 user"), (64, "decode, batch of 64"), (1000, "prefill, 1,000 tokens")]:
1312        i = arithmetic_intensity(n)
1313        ax.plot(i, min(PEAK_FLOPS, HBM_BANDWIDTH * i), "o", ms=8, label=f"{label} (I = {i:g})")
1314    ax.axvline(ridge_point(), ls=":", color="grey")
1315    ax.text(ridge_point() * 1.1, 2e12, f"break-even ≈ {ridge_point():.0f} FLOPs/byte", fontsize=8)
1316    ax.set(xlabel="arithmetic intensity (FLOPs per byte of weights read)", ylabel="attainable FLOP/s",
1317           title="Decode is memory-bound; prefill is compute-bound")
1318    ax.legend(loc="lower right", fontsize=8)
1319    figs["roofline"] = fig
1320
1321    # KV cache work.
1322    steps = np.arange(1, 33)
1323    fig, ax = plt.subplots(figsize=(6.5, 4))
1324    ax.plot(steps, _cumulative_flops(False), lw=2, label="no cache: recompute the whole sequence")
1325    ax.plot(steps, _cumulative_flops(True), lw=2, label="KV cache: compute only the new token")
1326    ax.set(xlabel="tokens generated after an 8-token prompt", ylabel="total FLOPs so far (TinyDecoder)",
1327           title="The KV cache turns quadratic work into linear work")
1328    ax.legend()
1329    figs["kv_cache_work"] = fig
1330
1331    # KV memory vs. context.
1332    ctx = np.linspace(0, 128_000, 200)
1333    fig, ax = plt.subplots(figsize=(6.5, 4))
1334    for heads, label in [(32, "32 KV heads (multi-head attention)"), (8, "8 KV heads (grouped-query, Llama-3-8B)")]:
1335        ax.plot(ctx / 1000, [kv_cache_bytes(int(c), 32, heads, 128) / 1e9 for c in ctx], lw=2, label=label)
1336    ax.axhline(64, ls="--", color="grey", label="free memory: 80 GB GPU − 16 GB weights")
1337    ax.set(xlabel="context length per request (thousands of tokens)", ylabel="KV cache for one request (GB)",
1338           title="Long contexts eat GPU memory; fewer KV heads help")
1339    ax.legend(fontsize=8)
1340    figs["kv_memory"] = fig
1341
1342    # Sampling at three temperatures, with top-p cut-offs.
1343    logits = np.array([3.0, 2.2, 1.5, 0.5, -0.5])
1344    temps = [0.5, 1.0, 2.0]
1345    fig, axes = plt.subplots(1, 3, figsize=(9, 3.4), sharey=True)
1346    for ax, t in zip(axes, temps):
1347        probs = temperature_probs(logits, t)
1348        kept = top_p_filter(probs, 0.9) > 0
1349        ax.bar(range(5), probs, color=["C0" if k else "white" for k in kept], edgecolor="C0",
1350               hatch=None, label="kept by top-p 0.9")
1351        for j in np.where(~kept)[0]:
1352            ax.bar(j, probs[j], color="white", edgecolor="C3", hatch="//")
1353        ax.set(title=f"T = {t}", xlabel="candidate token", xticks=range(5))
1354    axes[0].set_ylabel("probability")
1355    fig.suptitle("Temperature reshapes the distribution; top-p (hatched = cut) trims the tail")
1356    figs["sampling"] = fig
1357
1358    # Batching timelines.
1359    lengths = _realistic_lengths()
1360    fig, axes = plt.subplots(2, 1, figsize=(9, 4.5), sharex=True)
1361    for ax, (name, run) in zip(axes, [("static", simulate_static_batching(lengths, 8)),
1362                                      ("continuous", simulate_continuous_batching(lengths, 8))]):
1363        grid = np.array([[np.nan if r is None else r for r in row] for row in run.timeline]).T
1364        ax.imshow(grid, aspect="auto", cmap="tab20", interpolation="nearest")
1365        ax.set(ylabel="slot", title=f"{name} batching: {run.steps} steps, {run.utilization:.0%} of slot-steps busy")
1366    axes[1].set_xlabel("decode step")
1367    figs["batching"] = fig
1368
1369    for f in figs.values():
1370        f.tight_layout()
1371    return figs
1372
1373
1374# ---------------------------------------------------------------------------
1375# 9. Narrated walkthrough
1376# ---------------------------------------------------------------------------
1377
1378
1379def demo() -> None:
1380    banner("1. Prefill vs. decode on an 8B model (16-bit, H100-class GPU)")
1381    table(
1382        ["phase", "tokens per pass", "FLOPs per byte", "bound by", "time"],
1383        [
1384            ("prefill 1,000 tokens", 1000, arithmetic_intensity(1000), bottleneck(1000), f"{prefill_seconds(8e9, 1000) * 1e3:.1f} ms total"),
1385            ("decode", 1, arithmetic_intensity(1), bottleneck(1), f"{decode_seconds_per_token(8e9) * 1e3:.2f} ms per token"),
1386        ],
1387        floatfmt=".0f",
1388    )
1389    say(f"The GPU breaks even at {ridge_point():.1f} FLOPs per byte. Decode does 1, so the arithmetic units mostly wait on memory.")
1390    takeaway("Prefill sets time to first token; decode sets tokens per second, and decode is memory-bound.")
1391
1392    banner("2. The KV cache: identical output, far less work")
1393    model = TinyDecoder(seed=0)
1394    a = model.generate(EXAMPLE_PROMPT, 16, use_cache=False)
1395    b = model.generate(EXAMPLE_PROMPT, 16, use_cache=True)
1396    table(
1397        ["", "tokens generated", "positions processed", "FLOPs"],
1398        [("no cache", a.tokens[:6], a.positions_processed, f"{a.flops:,}"), ("KV cache", b.tokens[:6], b.positions_processed, f"{b.flops:,}")],
1399    )
1400    say(f"Same tokens (the weights are random, so the tokens themselves mean nothing), {a.flops / b.flops:.1f}× fewer operations. 248 = 8 + 9 + ... + 23; 23 = 8 prompt + 15 new.")
1401    takeaway("The KV cache trades GPU memory for speed: each new token is computed exactly once.")
1402
1403    banner("3. Memory math")
1404    table(
1405        ["quantity", "calculation", "result"],
1406        [
1407            ("70B weights, 16-bit", "70e9 × 2 bytes", f"{weight_bytes(70e9, 16) / 1e9:.0f} GB"),
1408            ("70B weights, 4-bit", "70e9 × 0.5 bytes", f"{weight_bytes(70e9, 4) / 1e9:.0f} GB"),
1409            ("KV per token (Llama-3-8B shape)", "2 × 32 × 8 × 128 × 2", f"{kv_cache_bytes_per_token(32, 8, 128):,} bytes"),
1410            ("KV for 32,000 tokens", "32,000 × 131,072", f"{kv_cache_bytes(32_000, 32, 8, 128) / 1e9:.2f} GB"),
1411            ("32k requests per 80 GB GPU", "(80 − 16) GB / 4.19 GB", max_concurrent_requests(80, 16, 32_000, 32, 8, 128)),
1412        ],
1413    )
1414
1415    banner("4. Sampling")
1416    z = np.array([2.0, 1.0, 0.0])
1417    table(["temperature", "p(token 1)", "p(token 2)", "p(token 3)"], [(t, *temperature_probs(z, t)) for t in (0.0, 0.5, 1.0, 2.0)], floatfmt=".3f")
1418    probs = np.array([0.5, 0.3, 0.15, 0.05])
1419    say(f"From (0.5, 0.3, 0.15, 0.05): top-k=2 gives {np.round(top_k_filter(probs, 2), 3)}; top-p=0.9 gives {np.round(top_p_filter(probs, 0.9), 3)}.")
1420    say(
1421        f"""
1422        Float32: (1e8 + 1) − 1e8 = {float32_sum([1e8, 1.0, -1e8])}, but (1e8 − 1e8) + 1 = {float32_sum([1e8, -1e8, 1.0])}.
1423        Reduction order alone can flip a near-tie, so temperature 0 is not a determinism guarantee.
1424        """
1425    )
1426
1427    banner("5. Speculative decoding")
1428    target = np.array([[0.5, 0.3, 0.2], [0.1, 0.6, 0.3], [0.3, 0.3, 0.4]])
1429    draft = np.array([[0.3, 0.3, 0.4], [0.2, 0.5, 0.3], [0.6, 0.2, 0.2]])
1430    rng = np.random.default_rng(0)
1431    firsts = np.bincount([speculative_round(0, target, draft, 3, rng)[0] for _ in range(20_000)], minlength=3) / 20_000
1432    table(["", "token 0", "token 1", "token 2"], [("big model alone", *target[0]), ("small model alone", *draft[0]), ("speculative output", *firsts)], floatfmt=".3f")
1433    say(f"Acceptance rate Σ min(p, q) = {acceptance_rate(target[0], draft[0]):.1f}; with 4 drafts that is {expected_tokens_per_round(0.8, 4):.2f} tokens per big-model pass.")
1434    takeaway("The draft model changes speed, never the output distribution.")
1435
1436    banner("6. Quantization")
1437    row = np.array([[0.5, -1.27, 0.02]])
1438    for bits in (8, 4):
1439        codes, scales = quantize(row, bits)
1440        print(f"int{bits}: codes {codes[0].tolist()}, scale {scales[0, 0]:.4f}, back to {np.round(dequantize(codes, scales)[0], 3).tolist()}")
1441    print()
1442
1443    banner("7. Continuous batching")
1444    for name, run in [("static", simulate_static_batching([4, 1, 1, 1], 2)), ("continuous", simulate_continuous_batching([4, 1, 1, 1], 2))]:
1445        print(f"{name:10s} {run.steps} steps, utilisation {run.utilization:.1%}, timeline {run.timeline}")
1446    print()
1447
1448    banner("8. Prompt caching")
1449    uncached, cached = prompt_cache_cost(10_000, 500, 100, 3.0)
1450    say(f"100 calls with a 10,000-token shared prefix: ${uncached:.2f} uncached vs. ${cached:.2f} cached ({1 - cached / uncached:.0%} less).")
1451    takeaway("Stable content first, volatile content last: a cache is valid only up to the first differing token.")
1452
1453
1454if __name__ == "__main__":
1455    demo()
Level 3: the code, function by function.
PEAK_FLOPS = 1000000000000000.0
HBM_BANDWIDTH = 3350000000000.0
def weight_bytes(params: float, bits: int = 16) -> float: on GitHub
829def weight_bytes(params: float, bits: int = 16) -> float:
830    """Memory for the weights alone: parameters × bytes per parameter."""
831    return params * bits / 8

Memory for the weights alone: parameters × bytes per parameter.

def kv_cache_bytes_per_token(layers: int, kv_heads: int, head_dim: int, bits: int = 16) -> int: on GitHub
834def kv_cache_bytes_per_token(layers: int, kv_heads: int, head_dim: int, bits: int = 16) -> int:
835    """2 (a key and a value) × layers × KV heads × head dimension × bytes.
836
837    Every layer stores one key vector and one value vector per KV head for
838    every token in the context, so the cache grows by this much per token.
839    """
840    return int(2 * layers * kv_heads * head_dim * bits // 8)

2 (a key and a value) × layers × KV heads × head dimension × bytes.

Every layer stores one key vector and one value vector per KV head for every token in the context, so the cache grows by this much per token.

def kv_cache_bytes( context: int, layers: int, kv_heads: int, head_dim: int, bits: int = 16, batch: int = 1) -> int: on GitHub
843def kv_cache_bytes(context: int, layers: int, kv_heads: int, head_dim: int, bits: int = 16, batch: int = 1) -> int:
844    """KV cache for `batch` requests of `context` tokens each."""
845    return batch * context * kv_cache_bytes_per_token(layers, kv_heads, head_dim, bits)

KV cache for batch requests of context tokens each.

def max_concurrent_requests( gpu_gb: float, weights_gb: float, context: int, layers: int, kv_heads: int, head_dim: int, bits: int = 16) -> int: on GitHub
848def max_concurrent_requests(gpu_gb: float, weights_gb: float, context: int, layers: int, kv_heads: int, head_dim: int, bits: int = 16) -> int:
849    """How many full-context requests fit in the memory left after the weights."""
850    free = (gpu_gb - weights_gb) * 1e9
851    return int(free // kv_cache_bytes(context, layers, kv_heads, head_dim, bits))

How many full-context requests fit in the memory left after the weights.

def arithmetic_intensity(tokens_in_parallel: int, bits: int = 16) -> float: on GitHub
859def arithmetic_intensity(tokens_in_parallel: int, bits: int = 16) -> float:
860    """FLOPs done per byte of weights read from memory.
861
862    Each weight is read once per forward pass and used for one multiply-add
863    (2 FLOPs) *per token processed in that pass*. Prefill processes the whole
864    prompt in one pass; decode processes one new token per pass.
865    """
866    return 2 * tokens_in_parallel / (bits / 8)

FLOPs done per byte of weights read from memory.

Each weight is read once per forward pass and used for one multiply-add (2 FLOPs) per token processed in that pass. Prefill processes the whole prompt in one pass; decode processes one new token per pass.

def ridge_point( flops_per_second: float = 1000000000000000.0, bytes_per_second: float = 3350000000000.0) -> float: on GitHub
869def ridge_point(flops_per_second: float = PEAK_FLOPS, bytes_per_second: float = HBM_BANDWIDTH) -> float:
870    """The intensity at which math and memory take equally long (the roofline's ridge)."""
871    return flops_per_second / bytes_per_second

The intensity at which math and memory take equally long (the roofline's ridge).

def bottleneck(tokens_in_parallel: int, bits: int = 16) -> str: on GitHub
874def bottleneck(tokens_in_parallel: int, bits: int = 16) -> str:
875    """Below the ridge point the GPU waits on memory; above it, on arithmetic."""
876    return "compute-bound" if arithmetic_intensity(tokens_in_parallel, bits) >= ridge_point() else "memory-bound"

Below the ridge point the GPU waits on memory; above it, on arithmetic.

def decode_seconds_per_token( params: float, bits: int = 16, bytes_per_second: float = 3350000000000.0) -> float: on GitHub
879def decode_seconds_per_token(params: float, bits: int = 16, bytes_per_second: float = HBM_BANDWIDTH) -> float:
880    """Lower bound on time per generated token at batch size 1: read every weight once."""
881    return weight_bytes(params, bits) / bytes_per_second

Lower bound on time per generated token at batch size 1: read every weight once.

def prefill_seconds( params: float, n_tokens: int, flops_per_second: float = 1000000000000000.0) -> float: on GitHub
884def prefill_seconds(params: float, n_tokens: int, flops_per_second: float = PEAK_FLOPS) -> float:
885    """Lower bound on prefill time: 2 FLOPs per parameter per prompt token."""
886    return 2 * params * n_tokens / flops_per_second

Lower bound on prefill time: 2 FLOPs per parameter per prompt token.

class Generation: on GitHub
904class Generation:
905    """Result of `TinyDecoder.generate`: the new tokens and how much work they took."""
906
907    def __init__(self, tokens: list[int], positions_processed: int, flops: int):
908        self.tokens = tokens
909        self.positions_processed = positions_processed
910        self.flops = flops

Result of TinyDecoder.generate: the new tokens and how much work they took.

Generation(tokens: list[int], positions_processed: int, flops: int) on GitHub
907    def __init__(self, tokens: list[int], positions_processed: int, flops: int):
908        self.tokens = tokens
909        self.positions_processed = positions_processed
910        self.flops = flops
tokens
flops
class TinyDecoder: on GitHub
 913class TinyDecoder:
 914    """A 2-layer decoder-only transformer with random weights, small enough to read.
 915
 916    Each layer: RMS-norm -> multi-head causal self-attention -> residual add,
 917    then RMS-norm -> ReLU feed-forward -> residual add. Token embeddings plus
 918    learned position embeddings go in; the embedding matrix is reused to turn
 919    the final vectors into one score (logit) per vocabulary token.
 920
 921    Two ways to run it:
 922      * `forward_full(tokens)`: process a whole sequence at once (prefill, or
 923        naive generation that recomputes everything every step).
 924      * `forward_step(token, pos, cache)`: process ONE new token, reading the
 925        keys and values of all earlier tokens from `cache` instead of
 926        recomputing them. That is the KV cache.
 927    """
 928
 929    def __init__(self, vocab: int = 16, d_model: int = 16, n_heads: int = 2, n_layers: int = 2, max_len: int = 64, seed: int = 0):
 930        rng = np.random.default_rng(seed)
 931        s = 1 / np.sqrt(d_model)
 932        self.vocab, self.d, self.h, self.dh = vocab, d_model, n_heads, d_model // n_heads
 933        self.embed = rng.normal(0, 1, (vocab, d_model))
 934        self.pos = rng.normal(0, 0.5, (max_len, d_model))
 935        self.layers = [
 936            dict(
 937                Wq=rng.normal(0, s, (d_model, d_model)), Wk=rng.normal(0, s, (d_model, d_model)),
 938                Wv=rng.normal(0, s, (d_model, d_model)), Wo=rng.normal(0, s, (d_model, d_model)),
 939                W1=rng.normal(0, s, (d_model, 4 * d_model)), W2=rng.normal(0, s / 2, (4 * d_model, d_model)),
 940            )
 941            for _ in range(n_layers)
 942        ]
 943        self.flops = 0  # multiply-adds × 2, counted by _mm
 944
 945    def _mm(self, a: np.ndarray, b: np.ndarray) -> np.ndarray:
 946        # Matrix multiply that also counts its cost: 2 FLOPs per multiply-add.
 947        out = a @ b
 948        self.flops += 2 * a.shape[-1] * out.size
 949        return out
 950
 951    def _heads(self, x: np.ndarray) -> np.ndarray:
 952        # (seq, d_model) -> (heads, seq, d_head)
 953        return x.reshape(x.shape[0], self.h, self.dh).transpose(1, 0, 2)
 954
 955    def _attend(self, q: np.ndarray, k: np.ndarray, v: np.ndarray, causal: bool) -> np.ndarray:
 956        # q: (heads, n_q, d_head); k, v: (heads, n_k, d_head). Returns (n_q, d_model).
 957        scores = self._mm(q, k.transpose(0, 2, 1)) / np.sqrt(self.dh)
 958        if causal:
 959            n_q, n_k = scores.shape[-2:]
 960            # Query i sits at absolute position n_k - n_q + i and may see keys up to there.
 961            allowed = np.arange(n_k)[None, :] <= (np.arange(n_q)[:, None] + n_k - n_q)
 962            scores = np.where(allowed, scores, -np.inf)
 963        out = self._mm(_softmax(scores), v)  # (heads, n_q, d_head)
 964        return out.transpose(1, 0, 2).reshape(q.shape[1], self.d)
 965
 966    def _block(self, x: np.ndarray, layer: dict, k_all: np.ndarray, v_all: np.ndarray) -> np.ndarray:
 967        q = self._heads(self._mm(_rms_norm(x), layer["Wq"]))
 968        x = x + self._mm(self._attend(q, k_all, v_all, causal=True), layer["Wo"])
 969        return x + self._mm(np.maximum(0, self._mm(_rms_norm(x), layer["W1"])), layer["W2"])
 970
 971    def _kv(self, x: np.ndarray, layer: dict) -> tuple[np.ndarray, np.ndarray]:
 972        n = _rms_norm(x)
 973        return self._heads(self._mm(n, layer["Wk"])), self._heads(self._mm(n, layer["Wv"]))
 974
 975    def forward_full(self, tokens: np.ndarray, cache: list[dict] | None = None) -> np.ndarray:
 976        """Logits (seq, vocab) for every position, computing everything from scratch.
 977
 978        Pass an empty `cache` to fill it in the same pass: that is *prefill*,
 979        processing the whole prompt in parallel and keeping every key and value.
 980        """
 981        x = self.embed[tokens] + self.pos[: len(tokens)]
 982        for i, layer in enumerate(self.layers):
 983            k, v = self._kv(x, layer)
 984            if cache is not None:
 985                cache[i]["k"], cache[i]["v"] = k, v
 986            x = self._block(x, layer, k, v)
 987        return self._mm(_rms_norm(x), self.embed.T)
 988
 989    def new_cache(self) -> list[dict]:
 990        """One empty key store and value store per layer, each (heads, 0, d_head)."""
 991        empty = np.zeros((self.h, 0, self.dh))
 992        return [dict(k=empty, v=empty) for _ in self.layers]
 993
 994    def forward_step(self, token: int, pos: int, cache: list[dict]) -> np.ndarray:
 995        """Logits (vocab,) for ONE new token, appending its key and value to `cache`."""
 996        x = self.embed[[token]] + self.pos[[pos]]  # (1, d_model)
 997        for layer, store in zip(self.layers, cache):
 998            k_new, v_new = self._kv(x, layer)  # only the new token's K and V are computed
 999            store["k"] = np.concatenate([store["k"], k_new], axis=1)
1000            store["v"] = np.concatenate([store["v"], v_new], axis=1)
1001            x = self._block(x, layer, store["k"], store["v"])
1002        return self._mm(_rms_norm(x), self.embed.T)[0]
1003
1004    def generate(self, prompt: list[int], n_new: int, use_cache: bool = True) -> Generation:
1005        """Greedy generation of `n_new` tokens, with or without the KV cache."""
1006        self.flops = 0
1007        seq, out, processed = list(prompt), [], 0
1008        if use_cache:
1009            cache = self.new_cache()
1010            logits = self.forward_full(np.array(seq), cache)[-1]  # prefill: whole prompt, one pass
1011            processed += len(seq)
1012            for i in range(n_new):
1013                tok = int(np.argmax(logits))
1014                out.append(tok)
1015                seq.append(tok)
1016                if i < n_new - 1:  # decode: one new position per step, earlier K/V read from the cache
1017                    logits = self.forward_step(tok, len(seq) - 1, cache)
1018                    processed += 1
1019        else:
1020            for _ in range(n_new):
1021                logits = self.forward_full(np.array(seq))[-1]  # recompute the whole sequence each time
1022                processed += len(seq)
1023                tok = int(np.argmax(logits))
1024                out.append(tok)
1025                seq.append(tok)
1026        return Generation(out, processed, self.flops)

A 2-layer decoder-only transformer with random weights, small enough to read.

Each layer: RMS-norm -> multi-head causal self-attention -> residual add, then RMS-norm -> ReLU feed-forward -> residual add. Token embeddings plus learned position embeddings go in; the embedding matrix is reused to turn the final vectors into one score (logit) per vocabulary token.

Two ways to run it:

  • forward_full(tokens): process a whole sequence at once (prefill, or naive generation that recomputes everything every step).
  • forward_step(token, pos, cache): process ONE new token, reading the keys and values of all earlier tokens from cache instead of recomputing them. That is the KV cache.
TinyDecoder( vocab: int = 16, d_model: int = 16, n_heads: int = 2, n_layers: int = 2, max_len: int = 64, seed: int = 0) on GitHub
929    def __init__(self, vocab: int = 16, d_model: int = 16, n_heads: int = 2, n_layers: int = 2, max_len: int = 64, seed: int = 0):
930        rng = np.random.default_rng(seed)
931        s = 1 / np.sqrt(d_model)
932        self.vocab, self.d, self.h, self.dh = vocab, d_model, n_heads, d_model // n_heads
933        self.embed = rng.normal(0, 1, (vocab, d_model))
934        self.pos = rng.normal(0, 0.5, (max_len, d_model))
935        self.layers = [
936            dict(
937                Wq=rng.normal(0, s, (d_model, d_model)), Wk=rng.normal(0, s, (d_model, d_model)),
938                Wv=rng.normal(0, s, (d_model, d_model)), Wo=rng.normal(0, s, (d_model, d_model)),
939                W1=rng.normal(0, s, (d_model, 4 * d_model)), W2=rng.normal(0, s / 2, (4 * d_model, d_model)),
940            )
941            for _ in range(n_layers)
942        ]
943        self.flops = 0  # multiply-adds × 2, counted by _mm
embed
pos
layers
flops
def forward_full( self, tokens: numpy.ndarray, cache: list[dict] | None = None) -> numpy.ndarray: on GitHub
975    def forward_full(self, tokens: np.ndarray, cache: list[dict] | None = None) -> np.ndarray:
976        """Logits (seq, vocab) for every position, computing everything from scratch.
977
978        Pass an empty `cache` to fill it in the same pass: that is *prefill*,
979        processing the whole prompt in parallel and keeping every key and value.
980        """
981        x = self.embed[tokens] + self.pos[: len(tokens)]
982        for i, layer in enumerate(self.layers):
983            k, v = self._kv(x, layer)
984            if cache is not None:
985                cache[i]["k"], cache[i]["v"] = k, v
986            x = self._block(x, layer, k, v)
987        return self._mm(_rms_norm(x), self.embed.T)

Logits (seq, vocab) for every position, computing everything from scratch.

Pass an empty cache to fill it in the same pass: that is prefill, processing the whole prompt in parallel and keeping every key and value.

def new_cache(self) -> list[dict]: on GitHub
989    def new_cache(self) -> list[dict]:
990        """One empty key store and value store per layer, each (heads, 0, d_head)."""
991        empty = np.zeros((self.h, 0, self.dh))
992        return [dict(k=empty, v=empty) for _ in self.layers]

One empty key store and value store per layer, each (heads, 0, d_head).

def forward_step(self, token: int, pos: int, cache: list[dict]) -> numpy.ndarray: on GitHub
 994    def forward_step(self, token: int, pos: int, cache: list[dict]) -> np.ndarray:
 995        """Logits (vocab,) for ONE new token, appending its key and value to `cache`."""
 996        x = self.embed[[token]] + self.pos[[pos]]  # (1, d_model)
 997        for layer, store in zip(self.layers, cache):
 998            k_new, v_new = self._kv(x, layer)  # only the new token's K and V are computed
 999            store["k"] = np.concatenate([store["k"], k_new], axis=1)
1000            store["v"] = np.concatenate([store["v"], v_new], axis=1)
1001            x = self._block(x, layer, store["k"], store["v"])
1002        return self._mm(_rms_norm(x), self.embed.T)[0]

Logits (vocab,) for ONE new token, appending its key and value to cache.

def generate( self, prompt: list[int], n_new: int, use_cache: bool = True) -> Generation: on GitHub
1004    def generate(self, prompt: list[int], n_new: int, use_cache: bool = True) -> Generation:
1005        """Greedy generation of `n_new` tokens, with or without the KV cache."""
1006        self.flops = 0
1007        seq, out, processed = list(prompt), [], 0
1008        if use_cache:
1009            cache = self.new_cache()
1010            logits = self.forward_full(np.array(seq), cache)[-1]  # prefill: whole prompt, one pass
1011            processed += len(seq)
1012            for i in range(n_new):
1013                tok = int(np.argmax(logits))
1014                out.append(tok)
1015                seq.append(tok)
1016                if i < n_new - 1:  # decode: one new position per step, earlier K/V read from the cache
1017                    logits = self.forward_step(tok, len(seq) - 1, cache)
1018                    processed += 1
1019        else:
1020            for _ in range(n_new):
1021                logits = self.forward_full(np.array(seq))[-1]  # recompute the whole sequence each time
1022                processed += len(seq)
1023                tok = int(np.argmax(logits))
1024                out.append(tok)
1025                seq.append(tok)
1026        return Generation(out, processed, self.flops)

Greedy generation of n_new tokens, with or without the KV cache.

def temperature_probs(logits: numpy.ndarray, temperature: float) -> numpy.ndarray: on GitHub
1034def temperature_probs(logits: np.ndarray, temperature: float) -> np.ndarray:
1035    """softmax(logits / T). T < 1 sharpens, T > 1 flattens, T = 0 means greedy (argmax)."""
1036    if temperature == 0:
1037        out = np.zeros_like(logits, dtype=float)
1038        out[int(np.argmax(logits))] = 1.0
1039        return out
1040    return _softmax(logits / temperature)

softmax(logits / T). T < 1 sharpens, T > 1 flattens, T = 0 means greedy (argmax).

def top_k_filter(probs: numpy.ndarray, k: int) -> numpy.ndarray: on GitHub
1043def top_k_filter(probs: np.ndarray, k: int) -> np.ndarray:
1044    """Keep the k most likely tokens, zero the rest, renormalise."""
1045    keep = np.argsort(probs)[::-1][:k]
1046    out = np.zeros_like(probs)
1047    out[keep] = probs[keep]
1048    return out / out.sum()

Keep the k most likely tokens, zero the rest, renormalise.

def top_p_filter(probs: numpy.ndarray, p: float) -> numpy.ndarray: on GitHub
1051def top_p_filter(probs: np.ndarray, p: float) -> np.ndarray:
1052    """Nucleus sampling: keep the smallest set of top tokens whose total reaches p."""
1053    order = np.argsort(probs)[::-1]
1054    cumulative = np.cumsum(probs[order])
1055    # Include the token that crosses p: count how many cumulative sums are still below p, plus one.
1056    n_keep = int(np.searchsorted(cumulative, p - 1e-12) + 1)
1057    out = np.zeros_like(probs)
1058    out[order[:n_keep]] = probs[order[:n_keep]]
1059    return out / out.sum()

Nucleus sampling: keep the smallest set of top tokens whose total reaches p.

def sample_next( logits: numpy.ndarray, rng: numpy.random._generator.Generator, temperature: float = 1.0, top_k: int | None = None, top_p: float | None = None) -> int: on GitHub
1062def sample_next(logits: np.ndarray, rng: np.random.Generator, temperature: float = 1.0, top_k: int | None = None, top_p: float | None = None) -> int:
1063    """The usual pipeline: temperature, then top-k, then top-p, then draw one token."""
1064    probs = temperature_probs(logits, temperature)
1065    if top_k is not None:
1066        probs = top_k_filter(probs, top_k)
1067    if top_p is not None:
1068        probs = top_p_filter(probs, top_p)
1069    return int(rng.choice(len(probs), p=probs))

The usual pipeline: temperature, then top-k, then top-p, then draw one token.

def float32_sum(values: list[float]) -> float: on GitHub
1072def float32_sum(values: list[float]) -> float:
1073    """Add values left to right in 32-bit floats, the way one GPU reduction order might."""
1074    total = np.float32(0.0)
1075    for v in values:
1076        total = np.float32(total + np.float32(v))
1077    return float(total)

Add values left to right in 32-bit floats, the way one GPU reduction order might.

def greedy_pick_with_summation_order(contributions: list[float], other_logit: float) -> str: on GitHub
1080def greedy_pick_with_summation_order(contributions: list[float], other_logit: float) -> str:
1081    """Token A's logit is a float32 sum of `contributions`; token B's is `other_logit`.
1082
1083    On a GPU, the order of a reduction depends on kernel choice and on what
1084    else is in the batch. When two logits are nearly tied, that order alone
1085    can change which token wins, even at temperature 0.
1086    """
1087    return "A" if float32_sum(contributions) > other_logit else "B"

Token A's logit is a float32 sum of contributions; token B's is other_logit.

On a GPU, the order of a reduction depends on kernel choice and on what else is in the batch. When two logits are nearly tied, that order alone can change which token wins, even at temperature 0.

def acceptance_rate(p: numpy.ndarray, q: numpy.ndarray) -> float: on GitHub
1100def acceptance_rate(p: np.ndarray, q: np.ndarray) -> float:
1101    """Probability a token drawn from the draft q survives verification against target p: Σ min(p, q)."""
1102    return float(np.minimum(p, q).sum())

Probability a token drawn from the draft q survives verification against target p: Σ min(p, q).

def expected_tokens_per_round(alpha: float, n_draft: int) -> float: on GitHub
1105def expected_tokens_per_round(alpha: float, n_draft: int) -> float:
1106    """Expected tokens per target pass when each draft is accepted with probability alpha.
1107
1108    Accepting i drafts in a row has probability alpha^i; every round also
1109    yields one token from the target (a correction or a bonus). Summing gives
1110    1 + alpha + ... + alpha^n_draft = (1 - alpha^(n+1)) / (1 - alpha).
1111    """
1112    if alpha == 1:
1113        return n_draft + 1.0
1114    return (1 - alpha ** (n_draft + 1)) / (1 - alpha)

Expected tokens per target pass when each draft is accepted with probability alpha.

Accepting i drafts in a row has probability alpha^i; every round also yields one token from the target (a correction or a bonus). Summing gives 1 + alpha + ... + alpha^n_draft = (1 - alpha^(n+1)) / (1 - alpha).

def speculative_round( last_token: int, target: numpy.ndarray, draft: numpy.ndarray, n_draft: int, rng: numpy.random._generator.Generator) -> list[int]: on GitHub
1117def speculative_round(last_token: int, target: np.ndarray, draft: np.ndarray, n_draft: int, rng: np.random.Generator) -> list[int]:
1118    """One round of speculative decoding for a toy bigram language.
1119
1120    `target[i]` and `draft[i]` are next-token distributions after token i.
1121    1. The draft model proposes n_draft tokens, one after another (cheap).
1122    2. The target model scores every proposed position. In a real system
1123       this is ONE parallel forward pass, costing about the same as
1124       generating a single token, because decode is memory-bound.
1125    3. Walk the drafts left to right. Accept draft x with probability
1126       min(1, p(x) / q(x)). On the first rejection, draw a replacement from
1127       the leftover distribution max(0, p - q), renormalised, and stop.
1128    4. If every draft survives, draw one bonus token from the target.
1129    The accept/reject rule guarantees the output has exactly the target's
1130    distribution; the draft only changes *speed*.
1131    """
1132    drafts, qs, prev = [], [], last_token
1133    for _ in range(n_draft):
1134        q = draft[prev]
1135        x = _draw(q, rng)
1136        drafts.append(x)
1137        qs.append(q)
1138        prev = x
1139    ps = [target[last_token]] + [target[x] for x in drafts]  # the target's view of every position
1140    out = []
1141    for i, x in enumerate(drafts):
1142        if rng.random() < min(1.0, ps[i][x] / qs[i][x]):
1143            out.append(x)
1144            continue
1145        leftover = np.maximum(ps[i] - qs[i], 0)
1146        out.append(_draw(leftover / leftover.sum(), rng))
1147        return out
1148    out.append(_draw(ps[n_draft], rng))
1149    return out

One round of speculative decoding for a toy bigram language.

target[i] and draft[i] are next-token distributions after token i.

  1. The draft model proposes n_draft tokens, one after another (cheap).
  2. The target model scores every proposed position. In a real system this is ONE parallel forward pass, costing about the same as generating a single token, because decode is memory-bound.
  3. Walk the drafts left to right. Accept draft x with probability min(1, p(x) / q(x)). On the first rejection, draw a replacement from the leftover distribution max(0, p - q), renormalised, and stop.
  4. If every draft survives, draw one bonus token from the target. The accept/reject rule guarantees the output has exactly the target's distribution; the draft only changes speed.
def speculative_generate( start: int, target: numpy.ndarray, draft: numpy.ndarray, n_tokens: int, n_draft: int, rng: numpy.random._generator.Generator) -> list[int]: on GitHub
1152def speculative_generate(start: int, target: np.ndarray, draft: np.ndarray, n_tokens: int, n_draft: int, rng: np.random.Generator) -> list[int]:
1153    """Run speculative rounds until at least `n_tokens` tokens exist; return the first n_tokens."""
1154    out: list[int] = []
1155    last = start
1156    while len(out) < n_tokens:
1157        new = speculative_round(last, target, draft, n_draft, rng)
1158        out += new
1159        last = new[-1]
1160    return out[:n_tokens]

Run speculative rounds until at least n_tokens tokens exist; return the first n_tokens.

def quantize( W: numpy.ndarray, bits: int = 8, per_channel: bool = True) -> tuple[numpy.ndarray, numpy.ndarray]: on GitHub
1168def quantize(W: np.ndarray, bits: int = 8, per_channel: bool = True) -> tuple[np.ndarray, np.ndarray]:
1169    """Symmetric integer quantization: w ≈ scale × code, code in [-(2^(b-1)-1), 2^(b-1)-1].
1170
1171    Per-channel means one scale per output row, so an outlier in one row
1172    doesn't coarsen every other row. Returns (integer codes, scales).
1173    """
1174    qmax = 2 ** (bits - 1) - 1  # 127 for int8, 7 for int4
1175    absmax = np.abs(W).max(axis=1, keepdims=True) if per_channel else np.abs(W).max(keepdims=True)
1176    scales = absmax / qmax
1177    codes = np.clip(np.round(W / scales), -qmax, qmax).astype(np.int32)
1178    return codes, scales

Symmetric integer quantization: w ≈ scale × code, code in [-(2^(b-1)-1), 2^(b-1)-1].

Per-channel means one scale per output row, so an outlier in one row doesn't coarsen every other row. Returns (integer codes, scales).

def dequantize(codes: numpy.ndarray, scales: numpy.ndarray) -> numpy.ndarray: on GitHub
1181def dequantize(codes: np.ndarray, scales: np.ndarray) -> np.ndarray:
1182    return codes * scales
def quantization_error(W: numpy.ndarray, bits: int = 8, per_channel: bool = True) -> float: on GitHub
1185def quantization_error(W: np.ndarray, bits: int = 8, per_channel: bool = True) -> float:
1186    """Relative error ||W - dequant(quant(W))|| / ||W||."""
1187    codes, scales = quantize(W, bits, per_channel)
1188    return float(np.linalg.norm(W - dequantize(codes, scales)) / np.linalg.norm(W))

Relative error ||W - dequant(quant(W))|| / ||W||.

class ServingRun: on GitHub
1196class ServingRun:
1197    """Result of a batching simulation."""
1198
1199    def __init__(self, steps: int, useful_slot_steps: int, slots: int, timeline: list[list[int | None]]):
1200        self.steps = steps
1201        self.utilization = useful_slot_steps / (steps * slots)
1202        self.timeline = timeline  # timeline[step][slot] = request id, or None if idle

Result of a batching simulation.

ServingRun( steps: int, useful_slot_steps: int, slots: int, timeline: list[list[int | None]]) on GitHub
1199    def __init__(self, steps: int, useful_slot_steps: int, slots: int, timeline: list[list[int | None]]):
1200        self.steps = steps
1201        self.utilization = useful_slot_steps / (steps * slots)
1202        self.timeline = timeline  # timeline[step][slot] = request id, or None if idle
steps
utilization
timeline
def simulate_static_batching(lengths: list[int], slots: int) -> ServingRun: on GitHub
1205def simulate_static_batching(lengths: list[int], slots: int) -> ServingRun:
1206    """Fill every slot, run until the LONGEST request in the batch finishes, repeat.
1207
1208    Finished slots sit idle until the whole batch is done.
1209    """
1210    timeline: list[list[int | None]] = []
1211    for start in range(0, len(lengths), slots):
1212        batch = list(range(start, min(start + slots, len(lengths))))
1213        for t in range(max(lengths[i] for i in batch)):
1214            row = [i if t < lengths[i] else None for i in batch]
1215            timeline.append(row + [None] * (slots - len(row)))
1216    return ServingRun(len(timeline), sum(lengths), slots, timeline)

Fill every slot, run until the LONGEST request in the batch finishes, repeat.

Finished slots sit idle until the whole batch is done.

def simulate_continuous_batching(lengths: list[int], slots: int) -> ServingRun: on GitHub
1219def simulate_continuous_batching(lengths: list[int], slots: int) -> ServingRun:
1220    """After every decode step, hand any freed slot to the next waiting request."""
1221    queue = list(range(len(lengths)))
1222    remaining = list(lengths)
1223    active: list[int | None] = [None] * slots
1224    timeline: list[list[int | None]] = []
1225    while queue or any(a is not None for a in active):
1226        for s in range(slots):  # refill free slots before the step
1227            if active[s] is None and queue:
1228                active[s] = queue.pop(0)
1229        timeline.append(list(active))
1230        for s, r in enumerate(active):  # one decode step for every active request
1231            if r is not None:
1232                remaining[r] -= 1
1233                if remaining[r] == 0:
1234                    active[s] = None
1235    return ServingRun(len(timeline), sum(lengths), slots, timeline)

After every decode step, hand any freed slot to the next waiting request.

def reusable_prefix_tokens(previous: list[int], new: list[int]) -> int: on GitHub
1238def reusable_prefix_tokens(previous: list[int], new: list[int]) -> int:
1239    """How many leading tokens two prompts share: the part whose KV cache can be reused.
1240
1241    Attention makes every token's keys and values depend on ALL earlier
1242    tokens, so the cache is valid only up to the first difference.
1243    """
1244    n = 0
1245    for a, b in zip(previous, new):
1246        if a != b:
1247            break
1248        n += 1
1249    return n

How many leading tokens two prompts share: the part whose KV cache can be reused.

Attention makes every token's keys and values depend on ALL earlier tokens, so the cache is valid only up to the first difference.

def prompt_cache_cost( prefix: int, suffix: int, requests: int, price_per_million: float, write_multiplier: float = 1.25, read_multiplier: float = 0.1) -> tuple[float, float]: on GitHub
1252def prompt_cache_cost(prefix: int, suffix: int, requests: int, price_per_million: float,
1253                      write_multiplier: float = 1.25, read_multiplier: float = 0.1) -> tuple[float, float]:
1254    """Input cost of `requests` calls sharing a `prefix`, without and with prompt caching.
1255
1256    The multipliers are typical of providers that price cache writes at a
1257    small premium and cache reads at a steep discount; check your provider's
1258    current pricing. Returns (uncached dollars, cached dollars).
1259    """
1260    per_token = price_per_million / 1e6
1261    uncached = requests * (prefix + suffix) * per_token
1262    first = (prefix * write_multiplier + suffix) * per_token
1263    later = (prefix * read_multiplier + suffix) * per_token
1264    return uncached, first + (requests - 1) * later

Input cost of requests calls sharing a prefix, without and with prompt caching.

The multipliers are typical of providers that price cache writes at a small premium and cache reads at a steep discount; check your provider's current pricing. Returns (uncached dollars, cached dollars).

EXAMPLE_PROMPT = [3, 1, 4, 1, 5, 9, 2, 6]
def viz_data() -> dict: on GitHub
1282def viz_data() -> dict:
1283    """The numbers the site's interactive KV-cache calculator starts from."""
1284    # Shapes of openly published models; the widget recomputes everything
1285    # else with the same formulas as kv_cache_bytes and weight_bytes.
1286    return {
1287        "kv-cache": {
1288            "gpu_bytes": 80e9,
1289            "presets": [
1290                {"name": "8B, 8 KV heads (Llama 3 8B shape)", "params": 8e9, "layers": 32, "kv_heads": 8, "head_dim": 128},
1291                {"name": "7B, 32 KV heads (no grouped-query attention)", "params": 7e9, "layers": 32, "kv_heads": 32, "head_dim": 128},
1292                {"name": "70B, 8 KV heads (Llama 3 70B shape)", "params": 70e9, "layers": 80, "kv_heads": 8, "head_dim": 128},
1293            ],
1294        }
1295    }

The numbers the site's interactive KV-cache calculator starts from.

def figures() -> dict: on GitHub
1298def figures() -> dict:
1299    """Data figures for this lesson, keyed by the name used in the docstring."""
1300    import matplotlib
1301
1302    matplotlib.use("Agg")
1303    import matplotlib.pyplot as plt
1304
1305    figs = {}
1306
1307    # Roofline.
1308    x = np.logspace(-0.5, 4, 300)
1309    roof = np.minimum(PEAK_FLOPS, HBM_BANDWIDTH * x)
1310    fig, ax = plt.subplots(figsize=(6.5, 4.2))
1311    ax.loglog(x, roof, lw=2, color="black")
1312    for n, label in [(1, "decode, 1 user"), (64, "decode, batch of 64"), (1000, "prefill, 1,000 tokens")]:
1313        i = arithmetic_intensity(n)
1314        ax.plot(i, min(PEAK_FLOPS, HBM_BANDWIDTH * i), "o", ms=8, label=f"{label} (I = {i:g})")
1315    ax.axvline(ridge_point(), ls=":", color="grey")
1316    ax.text(ridge_point() * 1.1, 2e12, f"break-even ≈ {ridge_point():.0f} FLOPs/byte", fontsize=8)
1317    ax.set(xlabel="arithmetic intensity (FLOPs per byte of weights read)", ylabel="attainable FLOP/s",
1318           title="Decode is memory-bound; prefill is compute-bound")
1319    ax.legend(loc="lower right", fontsize=8)
1320    figs["roofline"] = fig
1321
1322    # KV cache work.
1323    steps = np.arange(1, 33)
1324    fig, ax = plt.subplots(figsize=(6.5, 4))
1325    ax.plot(steps, _cumulative_flops(False), lw=2, label="no cache: recompute the whole sequence")
1326    ax.plot(steps, _cumulative_flops(True), lw=2, label="KV cache: compute only the new token")
1327    ax.set(xlabel="tokens generated after an 8-token prompt", ylabel="total FLOPs so far (TinyDecoder)",
1328           title="The KV cache turns quadratic work into linear work")
1329    ax.legend()
1330    figs["kv_cache_work"] = fig
1331
1332    # KV memory vs. context.
1333    ctx = np.linspace(0, 128_000, 200)
1334    fig, ax = plt.subplots(figsize=(6.5, 4))
1335    for heads, label in [(32, "32 KV heads (multi-head attention)"), (8, "8 KV heads (grouped-query, Llama-3-8B)")]:
1336        ax.plot(ctx / 1000, [kv_cache_bytes(int(c), 32, heads, 128) / 1e9 for c in ctx], lw=2, label=label)
1337    ax.axhline(64, ls="--", color="grey", label="free memory: 80 GB GPU − 16 GB weights")
1338    ax.set(xlabel="context length per request (thousands of tokens)", ylabel="KV cache for one request (GB)",
1339           title="Long contexts eat GPU memory; fewer KV heads help")
1340    ax.legend(fontsize=8)
1341    figs["kv_memory"] = fig
1342
1343    # Sampling at three temperatures, with top-p cut-offs.
1344    logits = np.array([3.0, 2.2, 1.5, 0.5, -0.5])
1345    temps = [0.5, 1.0, 2.0]
1346    fig, axes = plt.subplots(1, 3, figsize=(9, 3.4), sharey=True)
1347    for ax, t in zip(axes, temps):
1348        probs = temperature_probs(logits, t)
1349        kept = top_p_filter(probs, 0.9) > 0
1350        ax.bar(range(5), probs, color=["C0" if k else "white" for k in kept], edgecolor="C0",
1351               hatch=None, label="kept by top-p 0.9")
1352        for j in np.where(~kept)[0]:
1353            ax.bar(j, probs[j], color="white", edgecolor="C3", hatch="//")
1354        ax.set(title=f"T = {t}", xlabel="candidate token", xticks=range(5))
1355    axes[0].set_ylabel("probability")
1356    fig.suptitle("Temperature reshapes the distribution; top-p (hatched = cut) trims the tail")
1357    figs["sampling"] = fig
1358
1359    # Batching timelines.
1360    lengths = _realistic_lengths()
1361    fig, axes = plt.subplots(2, 1, figsize=(9, 4.5), sharex=True)
1362    for ax, (name, run) in zip(axes, [("static", simulate_static_batching(lengths, 8)),
1363                                      ("continuous", simulate_continuous_batching(lengths, 8))]):
1364        grid = np.array([[np.nan if r is None else r for r in row] for row in run.timeline]).T
1365        ax.imshow(grid, aspect="auto", cmap="tab20", interpolation="nearest")
1366        ax.set(ylabel="slot", title=f"{name} batching: {run.steps} steps, {run.utilization:.0%} of slot-steps busy")
1367    axes[1].set_xlabel("decode step")
1368    figs["batching"] = fig
1369
1370    for f in figs.values():
1371        f.tight_layout()
1372    return figs

Data figures for this lesson, keyed by the name used in the docstring.

def demo() -> None: on GitHub
1380def demo() -> None:
1381    banner("1. Prefill vs. decode on an 8B model (16-bit, H100-class GPU)")
1382    table(
1383        ["phase", "tokens per pass", "FLOPs per byte", "bound by", "time"],
1384        [
1385            ("prefill 1,000 tokens", 1000, arithmetic_intensity(1000), bottleneck(1000), f"{prefill_seconds(8e9, 1000) * 1e3:.1f} ms total"),
1386            ("decode", 1, arithmetic_intensity(1), bottleneck(1), f"{decode_seconds_per_token(8e9) * 1e3:.2f} ms per token"),
1387        ],
1388        floatfmt=".0f",
1389    )
1390    say(f"The GPU breaks even at {ridge_point():.1f} FLOPs per byte. Decode does 1, so the arithmetic units mostly wait on memory.")
1391    takeaway("Prefill sets time to first token; decode sets tokens per second, and decode is memory-bound.")
1392
1393    banner("2. The KV cache: identical output, far less work")
1394    model = TinyDecoder(seed=0)
1395    a = model.generate(EXAMPLE_PROMPT, 16, use_cache=False)
1396    b = model.generate(EXAMPLE_PROMPT, 16, use_cache=True)
1397    table(
1398        ["", "tokens generated", "positions processed", "FLOPs"],
1399        [("no cache", a.tokens[:6], a.positions_processed, f"{a.flops:,}"), ("KV cache", b.tokens[:6], b.positions_processed, f"{b.flops:,}")],
1400    )
1401    say(f"Same tokens (the weights are random, so the tokens themselves mean nothing), {a.flops / b.flops:.1f}× fewer operations. 248 = 8 + 9 + ... + 23; 23 = 8 prompt + 15 new.")
1402    takeaway("The KV cache trades GPU memory for speed: each new token is computed exactly once.")
1403
1404    banner("3. Memory math")
1405    table(
1406        ["quantity", "calculation", "result"],
1407        [
1408            ("70B weights, 16-bit", "70e9 × 2 bytes", f"{weight_bytes(70e9, 16) / 1e9:.0f} GB"),
1409            ("70B weights, 4-bit", "70e9 × 0.5 bytes", f"{weight_bytes(70e9, 4) / 1e9:.0f} GB"),
1410            ("KV per token (Llama-3-8B shape)", "2 × 32 × 8 × 128 × 2", f"{kv_cache_bytes_per_token(32, 8, 128):,} bytes"),
1411            ("KV for 32,000 tokens", "32,000 × 131,072", f"{kv_cache_bytes(32_000, 32, 8, 128) / 1e9:.2f} GB"),
1412            ("32k requests per 80 GB GPU", "(80 − 16) GB / 4.19 GB", max_concurrent_requests(80, 16, 32_000, 32, 8, 128)),
1413        ],
1414    )
1415
1416    banner("4. Sampling")
1417    z = np.array([2.0, 1.0, 0.0])
1418    table(["temperature", "p(token 1)", "p(token 2)", "p(token 3)"], [(t, *temperature_probs(z, t)) for t in (0.0, 0.5, 1.0, 2.0)], floatfmt=".3f")
1419    probs = np.array([0.5, 0.3, 0.15, 0.05])
1420    say(f"From (0.5, 0.3, 0.15, 0.05): top-k=2 gives {np.round(top_k_filter(probs, 2), 3)}; top-p=0.9 gives {np.round(top_p_filter(probs, 0.9), 3)}.")
1421    say(
1422        f"""
1423        Float32: (1e8 + 1) − 1e8 = {float32_sum([1e8, 1.0, -1e8])}, but (1e8 − 1e8) + 1 = {float32_sum([1e8, -1e8, 1.0])}.
1424        Reduction order alone can flip a near-tie, so temperature 0 is not a determinism guarantee.
1425        """
1426    )
1427
1428    banner("5. Speculative decoding")
1429    target = np.array([[0.5, 0.3, 0.2], [0.1, 0.6, 0.3], [0.3, 0.3, 0.4]])
1430    draft = np.array([[0.3, 0.3, 0.4], [0.2, 0.5, 0.3], [0.6, 0.2, 0.2]])
1431    rng = np.random.default_rng(0)
1432    firsts = np.bincount([speculative_round(0, target, draft, 3, rng)[0] for _ in range(20_000)], minlength=3) / 20_000
1433    table(["", "token 0", "token 1", "token 2"], [("big model alone", *target[0]), ("small model alone", *draft[0]), ("speculative output", *firsts)], floatfmt=".3f")
1434    say(f"Acceptance rate Σ min(p, q) = {acceptance_rate(target[0], draft[0]):.1f}; with 4 drafts that is {expected_tokens_per_round(0.8, 4):.2f} tokens per big-model pass.")
1435    takeaway("The draft model changes speed, never the output distribution.")
1436
1437    banner("6. Quantization")
1438    row = np.array([[0.5, -1.27, 0.02]])
1439    for bits in (8, 4):
1440        codes, scales = quantize(row, bits)
1441        print(f"int{bits}: codes {codes[0].tolist()}, scale {scales[0, 0]:.4f}, back to {np.round(dequantize(codes, scales)[0], 3).tolist()}")
1442    print()
1443
1444    banner("7. Continuous batching")
1445    for name, run in [("static", simulate_static_batching([4, 1, 1, 1], 2)), ("continuous", simulate_continuous_batching([4, 1, 1, 1], 2))]:
1446        print(f"{name:10s} {run.steps} steps, utilisation {run.utilization:.1%}, timeline {run.timeline}")
1447    print()
1448
1449    banner("8. Prompt caching")
1450    uncached, cached = prompt_cache_cost(10_000, 500, 100, 3.0)
1451    say(f"100 calls with a 10,000-token shared prefix: ${uncached:.2f} uncached vs. ${cached:.2f} cached ({1 - cached / uncached:.0%} less).")
1452    takeaway("Stable content first, volatile content last: a cache is valid only up to the first differing token.")