Part 3 trained a 5.8M-parameter GPT and sampled from it with a loop that is slow on purpose: for every new token, it reruns every previous token through the whole model. A KV cache removes that waste, and it is the standard trick inside every inference engine. This post adds one, checks that it changes nothing, and measures it.
The measurement is the surprise. On this model the cache sped the laptop CPU up 4.4x at batch 1. It did almost nothing for the RTX 3070 (1.0x to 1.1x), and with the cache on, the CPU (250 tokens per second) out-generated the GPU (223). At this size the GPU spends most of each step waiting on Python and kernel launches, not doing math. Batch 16 brings the GPU’s gain back (2.2x). After that, sampling settings decide whether the output reads like Holmes or like noise.
Train a Tiny GPT at Home series: Part 1: Tokenizers · Part 2: Attention · Part 3: Training · Part 4: Inference (you are here)
Full example: Clone the working files at github.com/KingPin/sumguy-examples/…/part-4-inference
Run everything from llm/tiny-gpt/ after the setup in the series README. You need Part 1’s tokenizer.json and a Part 3 checkpoint (gpu1500.pt).
cd part-4-inferencepython sampling.py # self-check, under 1 spython cached.py ../part-3-training/gpu1500.pt # cache vs no cache, about 10 spython bench.py cudapython bench.py cpu 8python sample_grid.py ../part-3-training/gpu1500.pt "Holmes"All numbers come from the same laptop as Part 3: an NVIDIA RTX 3070 Laptop GPU (8 GB) and an Intel Core i7-11800H, using 8 threads, with PyTorch 2.14.1 in the pytorch/pytorch:2.14.1-cuda13.2-cudnn9-runtime Docker image. The container needs pip install --break-system-packages regex==2026.9.29 first, exactly as in Part 3, and the Part 4 README has the full docker run line. The benchmark uses untrained weights, because speed does not depend on what the model learned. It uses greedy picks, fp32, and an 8-token random prompt. Repeat runs vary by about 5%.
What the No-Cache Loop Wastes
Generating token 201 needs one thing: the model’s scores for the next token after position 200. Part 3’s loop gets them by feeding all 200 tokens through all six blocks and throwing away 199 of the 200 outputs.
Most of that work repeats. In each attention layer, every token produces a key and a value. Those depend only on the token and the tokens before it. Token 100’s key at layer 3 is identical on step 200 and on step 201. The only new thing on each step is the newest token’s query, key, and value.
So the cache keeps each layer’s keys and values from earlier steps. A decode step then runs just one token through the model. Its query looks at the cached keys of everything before it, plus its own. The cost of a step drops from “256 tokens through six blocks” to “1 token through six blocks.”
The Cache, in Code
cached.py subclasses Part 3’s RoPEAttention and adds an optional cache dict per layer:
class CachedAttention(RoPEAttention): def __init__(self, cfg: Config): super().__init__(cfg) self.context = cfg.context
def forward(self, x, cache: dict | None = None, start: int = 0): """With no cache this is Part 3's attention, so model(idx) still works.""" B, T, C = x.shape q, k, v = self.qkv(x).split(C, dim=-1) q, k, v = (t.view(B, T, self.n_heads, -1).transpose(1, 2) for t in (q, k, v)) with torch.device(q.device): q, k = rope_at(q, start).to(v.dtype), rope_at(k, start).to(v.dtype) if cache is not None: if "k" in cache: assert T == 1, "after the prompt, feed one token at a time" k = torch.cat([cache["k"], k], dim=2) v = torch.cat([cache["v"], v], dim=2) cache["k"], cache["v"] = k[:, :, -self.context:], v[:, :, -self.context:] y = F.scaled_dot_product_attention(q, k, v, is_causal=T > 1) return self.out(y.transpose(1, 2).reshape(B, T, C))The code does three things. It computes q, k, and v for the new tokens only. It appends the new k and v to the stored ones. It trims the cache to the last 256 positions, the model’s context window. (I trimmed some of the file’s comments here; the full text explains each line.)
CachedGPT swaps this attention into every block and adds a step() method that runs only the new tokens:
class CachedGPT(GPT): def __init__(self, cfg: Config): super().__init__(cfg) for block in self.blocks: # same parameter names, so Part 3 checkpoints load block.attn = CachedAttention(cfg)
def step(self, idx, caches, start: int): """Run only the new tokens through the model. Returns last-position logits.""" x = self.emb(idx) for block, cache in zip(self.blocks, caches): x = x + block.attn(block.norm1(x), cache, start) x = x + block.mlp(block.norm2(x)) return self.head(self.norm(x[:, -1]))Because the new attention has the same parameter names as the old one, the Part 3 checkpoint loads without retraining. Generation then has two phases:
caches = [{} for _ in self.blocks]prompt = idx[:, -self.cfg.context:]logits = self.step(prompt, caches, 0) # prefill: the whole prompt in one passpos = prompt.shape[1]for _ in range(n_tokens): nxt = pick(logits, idx)[:, None] idx = torch.cat([idx, nxt], dim=1) logits = self.step(nxt, caches, pos) pos += 1Prefill runs the whole prompt in one pass and fills the cache. After that, each step feeds one token and passes its position in pos.
Two Gotchas That Break It Quietly
Both of these produce a model that runs and prints text. Neither raises an error. I broke each one on purpose and compared 100 greedy tokens against the correct cached output: the RoPE bug drifted off after 7 tokens, the mask bug after 3.
Gotcha 1: RoPE starts at position 0. Part 2’s rope() builds its angle table from torch.arange(T), so the first row is always position 0. In a cached step, T is 1. Without an offset, every new token gets rotated as if it sat at the very start of the text, so its distance to every cached key comes out wrong. The fix is rope_at(x, start), which begins the arange at start:
angles = torch.arange(start, start + T, dtype=torch.float32)[:, None] * freqsPart 2’s file stays unchanged. The cached model carries its own copy of the function.
Gotcha 2: is_causal=True with one query. F.scaled_dot_product_attention(..., is_causal=True) aligns its mask to the top left. With a single query row, that mask hides every key except the first. One new query is allowed to see the whole cache, so the code passes is_causal=T > 1: the mask applies during prefill (many queries) and switches off during decode (one query).
Does the Cache Change the Output?
It should not, and python cached.py checks. It generates 600 tokens greedily, with and without the cache, and compares the token IDs:
first 253 tokens identical with and without the cachefirst difference at position 322 of 603Inside the 256-token window, the outputs match exactly. Same token IDs, so the cache is a pure speedup there. On this run they actually stay identical until position 322, but the script only asserts the first 253, the part that must match. It fails loudly if that ever breaks.
Past the window, the two loops drift apart, and the script reports where instead of asserting. The reason is subtle. The uncached loop reruns the last 256 tokens from scratch each time. The rolling cache keeps entries that were computed while now-evicted tokens were still visible, so an old entry carries a bit of context that a fresh rerun would not have. RoPE itself stays correct as the window slides, because attention scores depend only on the distance between query and key. The content of old entries is what differs. Both are valid ways to slide a window, and they compute different things.
The Benchmark
bench.py generates tokens with and without the cache and reports tokens per second. The no-cache cost per token grows with context until the 256-token window fills, then stays flat. The 1024-token rows average over both phases. GPU first:
batch tokens no cache tok/s cache tok/s speedup 1 64 228 218 1.0x 1 256 227 230 1.0x 1 1024 194 223 1.1x 16 64 3,500 3,465 1.0x 16 256 2,073 3,592 1.7x 16 1024 1,369 3,059 2.2xThen the CPU, 8 threads:
batch tokens no cache tok/s cache tok/s speedup 1 64 164 279 1.7x 1 256 86 260 3.0x 1 1024 57 250 4.4x 16 64 506 2,055 4.1x 16 256 150 1,758 11.7x 16 1024 93 1,705 18.4xThe CPU behaves the way textbooks say. The no-cache speed falls as the context grows (164 to 57 tok/s at batch 1), the cached speed stays near 250 to 280, and the speedup climbs to 4.4x at batch 1 and 18.4x at batch 16.
The GPU does not. At batch 1 the cache gains nothing: 228 vs 218, 227 vs 230, 194 vs 223. And the cached GPU (223 tok/s at 1024 tokens) loses to the cached CPU (250).
Why the GPU Shrugs
A look at the profiler explains it. I ran torch.profiler on one cached decode step on the GPU at batch 1. That one step launched 317 CUDA kernels. The GPU was busy for about 1.0 ms of a step that took roughly 4 to 6 ms of wall time. The GPU sat idle about 80% of the step.
Each kernel is tiny at this size: a 256-wide matrix multiply on one token. The GPU finishes it faster than Python can launch the next one. Python and the CUDA launch queue set the step time. The cache removes arithmetic, so it removes nothing that mattered.
The no-cache numbers say the same thing from the other side. The uncached GPU step processes up to 256 tokens in the same wall time the cached step needs for 1 token. If the GPU were math-bound, 256 times the work would not cost the same time.
The CPU is different because it does the math itself. Remove 255 of 256 tokens of work and the step gets shorter. That shows up directly in the clock.
Fixes for launch overhead exist: CUDA graphs and torch.compile both cut the number of separate launches. They are out of scope here, and I am not claiming numbers for them.
Where the GPU Wins Back
Batch 16 changes the picture. Sixteen sequences share every kernel launch, so the launch cost spreads over 16 tokens instead of 1. The GPU then does real work between launches, and the cache matters again: 1,369 tok/s without it, 3,059 with it at 1024 tokens (2.2x). Total throughput at batch 16 reaches about 3,000 to 3,600 tok/s on the GPU against about 1,700 to 2,100 on the CPU. A server answering many requests at once is exactly this case, which is why batching is the standard serving trick.
So the rule from this table is plain. For one user and a model this small, the CPU is fine, and the cache is what makes it fine. For many users, the GPU wins, and the cache helps there too.
What the Cache Costs in Memory
The cache stores two tensors (K and V) per layer, one row per token, one d_model-wide vector per row:
cache bytes = 2 (K and V) x layers x tokens x d_model x bytes per valueFor this model at 256 tokens and batch 1, that is 2 x 6 x 256 x 256 x 4 bytes = 3.00 MiB. It is nothing next to 23 MB of fp32 weights.
The same formula is why large models are a different story. Layers, width, and context length all multiply, and every concurrent user gets their own cache. At long contexts the cache can reach gigabytes. KV Cache Quantization covers shrinking it, and GPU Memory Math walks through where it fits in a full VRAM budget.
Sampling: Where the Output Gets Its Character
The model’s job ends at a score for each of the 4,096 tokens. Everything about whether the text reads as careful or wild is decided afterward, by pick() in sampling.py:
def pick(logits, temperature=1.0, top_k=0, top_p=1.0, penalty=1.0, history=None): logits = logits.float().clone() if penalty != 1.0 and history is not None: seen = logits.gather(-1, history) # Divide positive scores, multiply negative ones: both push the token down. logits.scatter_(-1, history, torch.where(seen > 0, seen / penalty, seen * penalty)) if temperature == 0: return logits.argmax(-1) logits /= temperature if top_k: kth = logits.topk(top_k, dim=-1).values[:, -1:] logits[logits < kth] = float("-inf") if top_p < 1.0: sorted_logits, order = logits.sort(dim=-1, descending=True) cum = sorted_logits.softmax(-1).cumsum(-1) # Drop a token if the tokens ranked above it already reach p. # The top token always survives. drop = cum - sorted_logits.softmax(-1) >= top_p logits.scatter_(-1, order, sorted_logits.masked_fill(drop, float("-inf"))) return torch.multinomial(logits.softmax(-1), 1)[:, 0](Docstring trimmed.) Four knobs:
- Temperature divides the scores before the softmax. Zero means always take the top token (greedy). Below 1 sharpens the distribution, above 1 flattens it.
- Top-k keeps only the k highest-scoring tokens and drops the rest.
- Top-p keeps the smallest set of tokens whose probabilities add up to p. A token is dropped when the tokens ranked above it already reach p, so the top token always survives.
- Repetition penalty divides the positive scores, and multiplies the negative scores, of every token already in the history. Both moves push the token down. The history includes the prompt.
A self-check at the bottom of the file asserts the properties that matter. Temperature 0, top_k=1, and a tiny top_p all equal argmax. top_k=5 never leaves the top 5. A top_p of 0.75 over probabilities [0.5, 0.3, 0.15, 0.05] only ever picks the first two. A penalty of 2.0 pushes a repeated top token below the runner-up.
Same Model, Same Seed, Different Settings
sample_grid.py runs one checkpoint (gpu1500.pt from Part 3), the prompt “Holmes”, seed 1, 60 tokens, on the CPU. The script prints more settings than I show here. These five tell the story. The text is verbatim.
--- greedyHolmes, and I am sure that I am a very good friend of the kindness which I have already explained.
“I have been very much surprised, Mr. Holmes, that I have not been able to see you in the morning.”
“I have no doubt that you have been
--- temperature 1.0Holmes rose to the door and stood by his feetfropping himself to his sun and short an icely, with a wooden line oflamps drawn in nose, with his eyesug eyes which fled behind Holmes, andit seemed to me that your heart had lined Mrs
--- temperature 1.5Holmes rose side LimeROBLY TUBTER—”
sun held short an mistakenodge, absorbsful happens, long succincts came like thirtyug eyes within help. At Holmes’ here time he curtained your object that four French Mrs
--- top-p 0.9Holmes rose to the door and stood by his feetfropping himself to his watch.
“What was it, then?” asked Baldwin, chuckling. “We have reason to cometo the help. We have, here we’ll keep him. I think that I was here
--- greedy + penalty 1.3Holmes, and I am sure that you are not areprivate man. You will find me in the same way of your own mind.”
“I’m sorry,” said Holmes, “that it is all right to my wife; but I havenot seen her since thenWhat the output shows:
- Greedy is fluent but repetitive. It leans on “I have” phrasing again and again (“I have been”, “I have not”, “I have no doubt”).
- Temperature 1.0 and top-p 0.9 share an opening line. Same seed, and top-p only trimmed the unlikely tail at first. Then they diverge.
- Temperature 1.5 falls apart. It invents fragments like “LimeROBLY TUBTER” and “absorbsful”.
- A penalty of 1.3 cuts the repetition in greedy. The text still starts the same way, then goes somewhere new, and the “I have been / I have not / I have no doubt” chain does not recur.
None of this proves one setting is best. A 5.8M-parameter model trained on 1M tokens will invent words under any of them. The settings move you along a line from “safe and dull” to “creative and broken,” and where you stop depends on the job. Temperature 0.5 to 0.8 with top-p 0.9 is a reasonable place to start for prose like this. Raise the penalty only if you see loops.
What’s Next
We now have a trained model and a fast sampler. Part 5 is optional and still in planning: compare this from-scratch model against a LoRA fine-tune of Qwen 3.6 on the same Sherlock Holmes text, and see what 5.8M parameters buy you against a model that already knows English.
Common Questions
Does a KV cache change the model’s output?
No, not inside the context window. With greedy decoding, this 5.8M-parameter model produced identical token IDs with and without the cache for the first 253 generated tokens. Past the window, a rolling cache keeps older entries that a fresh rerun would not, so the two outputs eventually differ.
Why is my GPU not faster than my CPU for LLM inference?
Small models at batch 1 leave the GPU idle. In this series, a cached decode step launched 317 CUDA kernels and kept the GPU busy about 1.0 ms of a 4 to 6 ms step. Python and kernel launches set the speed. Larger batches or larger models give the GPU enough math to win.
What is a good temperature for text generation?
Start between 0.5 and 0.8. Temperature 0 is greedy and repeats itself. Temperature 1.0 gave varied but already garbled text from this tiny model, and 1.5 fell apart into invented fragments. Pair a temperature near 0.8 with top-p 0.9 to trim the unlikely tail, and tune by reading the output.
How much memory does a KV cache use?
The formula is 2 (K and V) x layers x tokens x d_model x bytes per value. For this model at 256 tokens and batch 1, that is 2 x 6 x 256 x 256 x 4 bytes, or 3.00 MiB. Every concurrent sequence needs its own cache, so memory grows with context length and batch size.