llm-inference-explained
← /learn · 01

From forward pass to generation loop

Prefill, decode, sampling and stopping: why generation is sequential.

Concept

A decoder-only language model has exactly one job: given a sequence of tokens, predict a distribution over the next token. Generating text is that job in a loop.

  1. Tokenise the prompt. Real models split text into sub-word pieces with a learned vocabulary of tens of thousands of entries (Llama 3 uses 128,256). The tiny model on this page uses one token per character, with 64 of them, so you can see every token.
  2. Prefill. Run the whole prompt through the model in one forward pass. Every position is processed at once, as matrices, which is what GPUs are good at. Only the last position's logits matter for what comes next.
  3. Sample a token from those logits: greedy (the most likely), temperature, top-k or top-p, exactly as in the explainer's Sampling chapter.
  4. Decode. Append the token and run the model again to get the next distribution. Repeat.
  5. Stop when the model emits an end-of-sequence token, when the output matches a stop string the caller asked for, or when it hits a length limit.

Why can't step 4 be done in parallel, like prefill? Because token t + 1 depends on token t, and token t doesn't exist until the previous step has sampled it. Generation is sequential by construction. A 500-token answer needs 500 forward passes one after another, however many GPUs you own. Everything else in these chapters is about making each of those passes cheaper, or about filling the GPU with other users' passes while you wait.

Try it below. Prefill is one step that takes in the whole prompt; each decode step takes in one token. The step log counts the arithmetic.

Interactive

Prefill once, then one token per step

A real 2-block decoder (d_model 16, untrained random weights, so the text is gibberish) running in your browser. Prefill processes the whole prompt in one pass; every decode step processes one token.

Sampling
the␣cat␣sat

Green: prompt (prefill). Blue: generated (decode). Stops on ‘.’ or at 32 tokens.

Concept

Two things in that log matter for the rest of the site.

  • Prefill and decode are different workloads. Prefill pushes many tokens through each weight matrix at once. Decode pushes one token per sequence through the same matrices. Chapter 3 shows that this makes prefill limited by arithmetic and decode limited by memory bandwidth.
  • Decode steps don't redo the prompt. Each one processes a single token, yet its output depends on every earlier token. That works because the model keeps each earlier token's keys and values: the KV cache, the subject of chapter 2.

Users see the two phases as two latencies: the wait for the first token (mostly prefill) and the time per output token after that (decode). Chapter 10 makes those precise.

Maths

A decoder defines the probability of a sequence autoregressively:

p(x1,…,xT)=∏t=1Tp(xt∣x<t).p(x_1, \dots, x_T) = \prod_{t=1}^{T} p(x_t \mid x_{<t}).

Generation draws xt∼p(⋅∣x<t)x_{t} \sim p(\cdot \mid x_{<t}) for t=P+1,P+2,…t = P+1, P+2, \dots after a prompt of length PP. With logits ℓ∈RV\ell \in \mathbb{R}^{V} at the last position and temperature τ\tau,

p(v∣x<t)=exp⁡(ℓv/τ)∑uexp⁡(ℓu/τ).p(v \mid x_{<t}) = \frac{\exp(\ell_v / \tau)}{\sum_{u} \exp(\ell_u / \tau)}.

Counting only matrix multiplies (2 FLOPs per multiply-add), a full forward pass over SS tokens of a model with LL blocks, width DD, MLP width FF and vocabulary VV costs

L (8SD2+4S2D+4SDF)⏟blocks+2SDV⏟output head\underbrace{L\,(8SD^2 + 4S^2D + 4SDF)}_{\text{blocks}} + \underbrace{2SDV}_{\text{output head}}

FLOPs: the four attention projections, the two attention products and the two MLP matrices. Decoding without a cache would pay this again at every step with SS one larger each time. The widget's FLOP column is this formula (for prefill) and its one-token version (for decode); a unit test checks both against FLOPs counted inside the matrix multiplies.

Code

// src/components/interactive/GenerationLoopWidget.tsx (simplified)
const { logits, cache } = prefill(promptIds, config, weights); // one pass, S tokens
let next = sample(logits[promptIds.length - 1], mode, rng);
while (next !== STOP_ID && length < MAX_LEN) {
  const l = decodeStep(next, cache, config, weights); // one token in
  next = sample(l, mode, rng);
}

prefill and decodeStep live in src/lib/transformer/kvcache.ts, next to the explainer's vendored transformer. Chapter 2 explains what cache holds.