At a glance
Key takeaways
- 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.
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.
Chapter 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.
Figure 1 · Diagram
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
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
Symbols
| Symbol | Meaning here | Shape / range |
|---|---|---|
| arithmetic intensity: operations per byte of weights read | FLOPs/byte | |
| tokens processed in one pass (1 when decoding, the prompt length when prefilling) | ≥ 1 | |
| one multiply plus one add per weight per token | ||
| bytes per weight (2 at 16-bit, 0.5 at 4-bit) | ||
| 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¹⁵ |
| "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
# 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
Figure 2 · Drawn from the lesson's code
Decode for 1 user uses 0.3% of the GPU's compute and a batch of 64 reaches 21%, while a 1,000-token prefill passes the 299 break-even to run at full speed
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).
Chapter 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.
Figure 3 · Diagram
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
TinyDecoder checks this to 10 decimal places); only
the cost differs.Figure 4 · Drawn from the lesson's code
Without a cache, total work curves upward to about 20 times the cached total after 32 tokens; with the cache it grows in a straight line
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.
Chapter 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
Symbols
| Symbol | Meaning here | Shape / range |
|---|---|---|
| number of parameters | e.g. 7×10¹⁰ | |
| bits / 8 | bytes per parameter (16 bits = 2 bytes) | |
| one key and one value per token | ||
| number of transformer layers, each with its own cache | e.g. 32 | |
| number of key/value heads (fewer than query heads with grouped-query attention) | e.g. 8 | |
| dimensions per head | e.g. 128 | |
| 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
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
Figure 5 · Drawn from the lesson's code
At 128k tokens, 32 KV heads need 67 GB, more than the 64 GB free, while 8 KV heads need 17 GB, so three such requests fit
Figure 6 · Interactive · computed from the lesson's code
KV cache memory calculator
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.
Chapter 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.
Figure 7 · Diagram
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]
Level 3: the formula and its symbols
Symbols
| Symbol | Meaning here | Shape / range |
|---|---|---|
| the model's score (logit) for vocabulary token i | any real | |
| temperature | > 0; T → 0 approaches greedy | |
| the exponential function; makes every score positive and stretches gaps | ||
| sum over every token j in the vocabulary, so the add to 1 | ||
| 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
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]
Figure 8 · Drawn from the lesson's code
Five tokens at three temperatures: at T = 0.5 the favourite takes 79% and top-p 0.9 cuts three tokens; at T = 2 it falls to 38% and only one is cut
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.
Chapter 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.
Figure 9 · Diagram
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]
Level 3: the formula and its symbols
Symbols
| Symbol | Meaning here | Shape / range |
|---|---|---|
| the big (target) model's probability for token x | 0 … 1 | |
| the small (draft) model's probability for token x | 0 … 1 | |
| the smaller of a and b | ||
| alpha, the acceptance rate: how often a draft survives | 0 … 1 | |
| gamma, how many tokens the small model drafts per round | e.g. 4 | |
| 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
# 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.
Chapter 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
Symbols
| Symbol | Meaning here | Shape / range |
|---|---|---|
| the j-th original weight in the row | real | |
| the largest absolute value in the row | ≥ 0 | |
| the largest code: 127 for 8 bits, 7 for 4 bits | ||
| the scale (step size), one per row | > 0 | |
| the stored integer code | −127…127 or −7…7 | |
| 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
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.
Chapter 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.
Figure 10 · Drawn from the lesson's code
With 32 requests on 8 slots, static batching leaves idle gaps and needs 433 steps at 60% busy; continuous batching needs 318 steps at 82%
Level 3: the formula and its symbols
Symbols
| Symbol | Meaning here | Shape / range |
|---|---|---|
| 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
# 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.
Test yourself
9 questions
Answer each one out loud or on paper before you open it. If you can explain it, you know it.
Question 1Q: Why is decoding slow even on a huge GPU?Think it through, then reveal
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.
Question 2Q: What is the KV cache, and why does it matter for serving cost?Think it through, then reveal
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.
Question 3Q: How much memory does a 70B model need at 16-bit and at 4-bit?Think it through, then reveal
A: 140 GB and 35 GB for weights alone, plus KV cache and activations.
Question 4Q: Estimate the KV cache for one 32k-token request on a Llama-3-8B-shaped model.Think it through, then reveal
A: 2 × 32 × 8 × 128 × 2 = 131,072 bytes per token; × 32,000 ≈ 4.2 GB.
Question 5Q: Why does grouped-query attention make serving cheaper?Think it through, then reveal
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.
Question 6Q: Why isn't temperature 0 perfectly deterministic?Think it through, then reveal
A: Floating-point addition isn't associative, and GPU reduction order can change with kernels and batch composition, so nearly tied logits can flip.
Question 7Q: How can speculative decoding be faster yet produce the same distribution?Think it through, then reveal
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.
Question 8Q: What does continuous batching fix?Think it through, then reveal
A: Idle slots: static batches wait for their longest request, while continuous batching refills any freed slot immediately, raising utilisation.
Question 9Q: How do you structure a prompt to benefit from prompt caching?Think it through, then reveal
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.
Primary sources
The papers behind this lesson
Leviathan, Kalman & Matias, Fast Inference from Transformers via Speculative Decoding (2022): Introduced the draft-then-verify scheme with an accept/reject rule that provably preserves the target model's output distribution.
Read the annotated companion →The paper ↗Kwon et al., Efficient Memory Management for Large Language Model Serving with PagedAttention (2023): 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.
Read the annotated companion →The paper ↗Dao et al., FlashAttention (2022): Computed exact attention in GPU-memory-sized tiles, cutting the slow memory traffic that dominates long-context inference.
Read the annotated companion →The paper ↗Dettmers et al., LLM.int8(): 8-bit Matrix Multiplication for Transformers at Scale (2022): Showed that a few outlier features break naive 8-bit quantization and handled them separately.
The paper ↗Frantar et al., GPTQ: Accurate Post-Training Quantization for Generative Pre-trained Transformers (2022): Made 3-4-bit weight quantization practical with error-compensating rounding.
The paper ↗Holtzman et al., The Curious Case of Neural Text Degeneration (2019): Introduced nucleus (top-p) sampling.
The paper ↗Researcher's shelf
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
About this lesson. This is the illustrated edition of a lesson from the open-source AI Primer. Its text, figures and numbers are generated from the Primer's source at commit c8d5c21, so the two always agree: the explanation, the code that builds it and the tests that prove it.