Trust is earned, not given

A different perspective

2024-09-25 · Projects

AI fundamentals, part 5: Inference — how an LLM writes one token at a time

Part 5from the AI fundamentals series · 8 parts in all

Parts 1–4 built and scaled a model. Part 5 runs it: the generation loop, temperature and top-p sampling, the KV cache that makes it 100× faster, and why your chatbot streams. Everything here is implementable in ~80 lines.

The loop that produces "intelligence"

A model answers by repeatedly predicting one next token, appending it, and predicting again. Autoregression is the entire trick:

def generate(token_ids, params, max_new=50):
    # Naive autoregressive generation. Returns the extended id sequence.
    out = list(token_ids)
    for _ in range(max_new):
        logits = gpt_forward(out, params)      # forward over ALL tokens so far
        next_id = sample(logits[-1])           # only the LAST position matters
        out.append(next_id)                    # append and repeat
        if next_id == EOS_TOKEN: break
    return out

Note the inefficiency — every step recomputes attention over the whole prefix. That's what the KV cache fixes below.

Temperature: the creativity dial

Raw logits become probabilities via softmax. Temperature reshapes that distribution before sampling: low T sharpens it (safe, repetitive), high T flattens it (surprising, incoherent past ~1.5). Divide logits by T, then softmax.

def sample(logits, temperature=1.0, top_p=0.95, rng=None):
    # Turn final-position logits into one sampled token id.
    rng = rng or np.random.default_rng()

    logits = logits / max(temperature, 1e-6)        # sharpen or flatten

    # top-p (nucleus): keep only the smallest set of tokens whose cumulative
    # probability reaches p — cuts the tail of unlikely tokens adaptively.
    probs = softmax(logits)
    order = np.argsort(-probs)                      # best token first
    cum = np.cumsum(probs[order])
    keep = order[cum <= top_p]
    keep = np.append(keep, order[len(keep)])        # always keep at least one

    mask = np.full_like(probs, -1e9)
    mask[keep] = logits[keep]                       # zero out the tail
    probs = softmax(mask)

    return int(rng.choice(len(probs), p=probs))     # one weighted dice roll

Practical anchors: T≈0.0–0.3 for extraction/classification (near-greedy), T≈0.7–1.0 for chat, top_p≈0.9–0.95 with T≈1.0 for creative work. Greedy (argmax) is just T→0 — and is why "temperature 0" answers can still loop.

The KV cache: the 100× trick

Generation is slow because each new token re-attends over the entire prefix. But the keys and values of past tokens never change — only the new token's query is unknown. So cache every past K and V; each step processes one token against the cache:

def generate_cached(token_ids, params, max_new=50):
    # Generation with a KV cache: O(T) total attention work instead of O(T^2).
    out = list(token_ids)
    cache = []                                     # list of (K, V) per layer
    x = None
    for step in range(max_new):
        new_tok = out[-1:] if step > 0 else out    # feed 1 token (or the prompt)
        x = params["tok_emb"][new_tok] + params["pos_emb"][len(out)-len(new_tok):len(out)]
        new_cache = []
        for i, blk in enumerate(params["blocks"]):
            prev_k, prev_v = cache[i] if cache else (None, None)
            x, K, V = block_with_cache(x, blk, prev_k, prev_v)  # attend to cache+self
            new_cache.append((K, V))               # remember this position too
        cache = new_cache
        x = layer_norm(x, params["lnfg"], params["lnfb"])
        logits = x[-1] @ params["tok_emb"].T       # last position's logits only
        nxt = sample(logits)
        out.append(nxt)
        if nxt == EOS_TOKEN: break
    return out

This is why prompt processing and generation are billed separately by APIs: the prompt is one big parallel forward pass (compute-bound); generation is a long sequential loop (memory-bandwidth-bound, reading the whole cache per token). It's also why longer contexts cost more and why context length is an engineering brag: the cache grows linearly with context and must fit in GPU memory.

Why streaming

Time-to-first-token is a full forward pass; each subsequent token is one small step. That's why chat UIs stream: first word lands in ~300 ms, the rest flows at 30–200 tokens/second — perceived latency is dominated by the first token, not the total.

One more part after this: how do we measure whether any of this works? Benchmarks, their gaming, and honest evaluation — the part most practitioners skip and most regret.