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
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
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.
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
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]
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.
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
- vLLM documentation: https://docs.vllm.ai/
- Hugging Face, Text generation strategies: https://huggingface.co/docs/transformers/generation_strategies
- Anthropic, Prompt caching: https://docs.claude.com/en/docs/build-with-claude/prompt-caching
- Thinking Machines Lab, Defeating Nondeterminism in LLM Inference: https://thinkingmachines.ai/blog/defeating-nondeterminism-in-llm-inference/
- Grattafiori et al., The Llama 3 Herd of Models (2024): https://arxiv.org/abs/2407.21783
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 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 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 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 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 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()
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.
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.
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.
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.
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.
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).
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.
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.
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.
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.
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 fromcacheinstead of recomputing them. That is the KV cache.
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
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.
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).
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.
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.
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).
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.
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.
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.
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.
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.
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).
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).
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.
- The draft model proposes n_draft tokens, one after another (cheap).
- 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.
- 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.
- 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.
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.
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).
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||.
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.
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.
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.
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.
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).
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.
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.
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.")