MiniMind — Train a Small Language Model Yourself, End to End
A KV cache remembers what need not be recomputed
In one line
Generation means running the model once for every token it picks. Without a cache, each time it feeds in everything so far and recomputes, and with a KV cache it feeds in only the one new token and pulls the K and V of the earlier tokens from the cache. The result does not differ by a single token, and only the computation is reduced. How you pick (greedy, temperature, top-p) separately decides "which token to choose". In this module, you measure both in numbers with MiniMind's generate.
Why this was needed
In the Transformer course, you computed the principle of the KV cache and the sampling formulas by hand. Here you look at how they appear in a model you actually trained and in an actual generation loop. Most serving cost comes from the generation stage, and the cache decides the shape of that cost. And the same model gives the same answer every time with greedy decoding and a different answer each time when you raise the temperature, and whether that difference is "creativity" or "nonsense" shows up especially starkly in a small model.
How it works
MiniMind's generate runs like this.
for _ in range(max_new_tokens):
past_len = past_key_values[0][0].shape[1] if past_key_values else 0
outputs = self.forward(input_ids[:, past_len:], past_key_values=past_key_values, use_cache=use_cache)
logits = outputs.logits[:, -1, :] / temperature
... top_k · top_p 로 자르기 ...
next_token = torch.multinomial(softmax(logits), 1) if do_sample else argmax(logits)
input_ids = torch.cat([input_ids, next_token], -1)
past_key_values = outputs.past_key_values if use_cache else None
KV cache. In the first step, the P prompt tokens are fed in all at once (prefill). With a cache, from then on only one token is fed in per step, so the tokens that go into the forward while producing N tokens are P + (N − 1). Without a cache, P, P+1, … are fed in again at every step, so it is P·N + N(N−1)/2. Producing 64 tokens from a 183-token prompt is 246 versus 13,728 — 56 times. Thanks to the causal mask, the K and V of earlier tokens do not change even when later tokens appear, so the result is the same when you pull them from the cache.
Cache size. The cache holds K and V for each layer, with shape (batch, length, KV heads, head_dim). MiniMind puts the K and V from before copying with repeat_kv into the cache, so it shrinks by exactly as much as GQA reduced. For 100 tokens, it is 2 × 4 layers × 100 × 2 heads × 32 × 4 bytes = 204,800 bytes. In serving, this number decides the number of concurrent users.
Temperature. It divides the logits by the temperature and then applies softmax. Below 1, the distribution gets sharper and concentrates on the most plausible token, and above 1 it flattens and rare tokens get picked too. Near temperature 0, it is the same as greedy.
top-p (nucleus). Sort by probability from largest, keep only up to the point where the cumulative probability exceeds p, and erase the rest to −∞. MiniMind shifts the mask by one slot to keep up to the token that first exceeds p, and always keeps the top token. When the distribution is sharp, only a few remain, and when it is flat, many remain — this differs from top-k, which keeps a fixed number. MiniMind's generate defaults are temperature 0.85, top_p 0.85, and top_k 50.
What it looks like in the field
A bug report saying "I turned on the cache and the answer changed" is usually a problem on the sampling side, not the cache — the random numbers were not fixed, or only one side has do_sample. If you first check with greedy decoding that the two methods do not differ by even one token, you cut the possible causes in half. Conversely, a long time to first token (TTFT) when summarizing a long document is not reduced by the cache — prefill has to compute the whole prompt anyway. What the cache reduces is the time between the tokens after that.
How this course differs from the original MiniMind
MiniMind's eval_llm.py and web demo sample with temperature 0.85 and top_p 0.95, and you can apply a repetition penalty (repetition_penalty). To keep experiments clean, this course turns off top_k (top_k=0) and changes only one thing at a time. Also, MiniMind's generate remembers which rows have finished in each batch (finished), stops when all are finished, and keeps filling eos into finished rows. That is why eos follows a short answer when you generate 30 at once. Real serving engines (such as vLLM) add to this a mechanism to bundle caches of different lengths per request into one batch — the principle is the same, and management just gets more complicated.
What you will do in the next lab
With the reference pretrained model, you check whether the same tokens come out with and without the cache, count the tokens fed into the forward and match them to the formula, and measure the time. You match the cache's actual bytes to the formula, change the temperature with the SFT model and generate 30 times each to count the number of distinct answers, and count the candidates top-p leaves.