Post-training and Alignment

Distilling a large model into a small one

A big model teaches a small one by answering questions for it, and the small model ends up far better than its own size and data would allow.

On this page 9
  1. The short answer
  2. The analogy you have already lived
  3. Why it exists
  4. How it works
  5. The two ways to teach
  6. What is honestly hard here
  7. Where you have already seen this
  8. Remember this
  9. What to learn next

One lesson, three depths. Pick the one that fits you today — you can switch any time.

Beginner — No maths. Plain English.

The short answer

A big model answers a huge pile of questions, and a small model is trained on those answers.

The analogy you have already lived

Think about a topper in your class making notes for a junior. The junior does not attend the topper's coaching classes and does not have the topper's years of practice.

What the junior gets is the worked answers. Hundreds of them, on the exact kinds of question that will come up.

The junior will not become the topper. But with those notes, the junior does far better than a student left alone with the textbook.

Why it exists

Big models are expensive. Every question you ask one costs money and time. Running one on a phone is out of the question.

Small models are cheap and fast. They are also worse, because they were trained on the same public text with less capacity to absorb it.

Distillation closes part of that gap without making the small model bigger. The teacher supplies answers the student could never have produced alone. The student learns from those.

How it works

   a big pile of questions (no answers needed)
                |
                v
      [ big teacher model ]
                |
        answers, one per question
                |
                v
   [ small student model ] trained on those pairs

Notice the important detail. Those questions never needed human answers. A person can write ten answers an hour. The teacher writes ten thousand.

The scarce thing was never the questions. It was the answers. Distillation makes answers cheap.

The two ways to teach

Give the final answer. The teacher writes out its answer, and the student trains on it as if a human had written it. This is what almost all modern distillation does.

Give the whole opinion. The teacher hands over its full ranking. Mostly cat, a bit dog, definitely not aeroplane. That extra shading is sometimes called dark knowledge.

The second sounds better. In the developer section it measurably is not, on a small student. Which of the two wins depends on the setup, and it is worth measuring rather than assuming.

What is honestly hard here

The student inherits the teacher's mistakes, and inherits them confidently.

If the teacher is wrong about something, every single one of those ten thousand answers repeats the error. There is no second opinion anywhere in the pipeline.

There is also a limit nobody has removed. A distilled model is good at the kinds of question the teacher covered. Elsewhere it is no better than its size.

Where you have already seen this

  • A senior's notes passed down before an exam.
  • A trainee shadowing an experienced worker for a month.
  • Recorded coaching classes, which are cheaper than the teacher's time.

Remember this

  • The teacher answers a big pile of questions, and the student trains on the answers.
  • Human answers are the expensive part, and this removes that cost.
  • The student copies the teacher's errors, and only improves where the teacher taught it.

What to learn next

Developer — Code and libraries.

Setup

bash
pip install torch

Runs on a CPU in about ninety seconds.

Measuring where the gain actually comes from

Four training regimes for the same tiny student, on a task with ordered classes so neighbouring labels are genuinely similar.

import torch
import torch.nn as nn
import torch.nn.functional as F

# Task: read 6 digits, say which of 10 ordered buckets their sum falls into.
V, LEN, C = 10, 6, 10


def data(n, gen):
    x = torch.randint(0, V, (n, LEN), generator=gen)
    y = (x.sum(1) * C) // (V * LEN - LEN + 1)
    return F.one_hot(x, V).float().reshape(n, LEN * V), y.long()


def mlp(hidden, seed):
    torch.manual_seed(seed)
    if hidden == 0:
        return nn.Linear(LEN * V, C)                 # the student: a linear model
    return nn.Sequential(nn.Linear(LEN * V, hidden), nn.ReLU(),
                         nn.Linear(hidden, hidden), nn.ReLU(),
                         nn.Linear(hidden, C))


gen = torch.Generator().manual_seed(0)
labelled_x, labelled_y = data(200, gen)          # all the HUMAN labels we have
transfer_x, _ = data(20000, gen)                 # unlabelled inputs: cheap and plentiful
teacher_x, teacher_y = data(20000, gen)          # the teacher's own training set
test_x, test_y = data(5000, torch.Generator().manual_seed(9))


def train(model, x, y=None, steps=800, soft=None, T=4.0, alpha=1.0, lr=3e-3):
    opt = torch.optim.AdamW(model.parameters(), lr=lr)
    for _ in range(steps):
        logits = model(x)
        loss = 0.0
        if y is not None:
            loss = loss + (1 - alpha if soft is not None else 1.0) * F.cross_entropy(logits, y)
        if soft is not None:
            loss = loss + alpha * T * T * F.kl_div(
                (logits / T).log_softmax(-1), soft, reduction="batchmean")
        opt.zero_grad()
        loss.backward()
        opt.step()
    with torch.no_grad():
        return (model(test_x).argmax(-1) == test_y).float().mean().item()


teacher = mlp(256, seed=1)
t_acc = train(teacher, teacher_x, teacher_y, steps=1200)
print(f"teacher: 256-wide MLP, {sum(p.numel() for p in teacher.parameters()):,} params -> "
      f"test accuracy {t_acc:.3f}")

with torch.no_grad():
    logits_T = teacher(transfer_x)
    hard_from_teacher = logits_T.argmax(-1)

print(f"\nstudent: linear, {sum(p.numel() for p in mlp(0, 2).parameters()):,} params")
print(f"  200 human labels only                {train(mlp(0, 2), labelled_x, labelled_y):.3f}")
print(f"  20000 teacher HARD labels            "
      f"{train(mlp(0, 2), transfer_x, hard_from_teacher):.3f}")
for T in (1.0, 2.0, 4.0):
    with torch.no_grad():
        soft = (logits_T / T).softmax(-1)
    print(f"  20000 teacher SOFT targets (T={T:.0f})     "
          f"{train(mlp(0, 2), transfer_x, hard_from_teacher, soft=soft, T=T, alpha=0.9):.3f}")
print(f"  20000 TRUE labels (an upper bound)   "
      f"{train(mlp(0, 2), teacher_x, teacher_y):.3f}")

with torch.no_grad():
    probs = (teacher(labelled_x[:1]) / 4.0).softmax(-1)[0]
print(f"\nteacher's softened view of one example (true class {labelled_y[0].item()}):")
print("  " + " ".join(f"{i}:{p:.3f}" for i, p in enumerate(probs.tolist())))
Output
teacher: 256-wide MLP, 83,978 params -> test accuracy 0.998

student: linear, 610 params
  200 human labels only                0.445
  20000 teacher HARD labels            0.726
  20000 teacher SOFT targets (T=1)     0.726
  20000 teacher SOFT targets (T=2)     0.639
  20000 teacher SOFT targets (T=4)     0.547
  20000 TRUE labels (an upper bound)   0.736

teacher's softened view of one example (true class 5):
  0:0.000 1:0.000 2:0.000 3:0.000 4:0.106 5:0.894 6:0.000 7:0.000 8:0.000 9:0.000

Written against PyTorch 2.5.1, CPU, all seeds fixed — reproducible on this build.

The headline result, and the one that contradicts the folklore

Distillation recovered almost the whole gap, with zero extra human labels. 0.445 with 200 human labels, 0.726 with 20,000 teacher labels, against an upper bound of 0.736 obtained with 20,000 true labels. The student got 96% of the way to the ceiling using labels a machine produced.

The transfer set is doing the work, not the softness. That is the single most useful takeaway here. Distillation's benefit in this run is "more labelled examples", and the mechanism is the teacher turning cheap unlabelled inputs into training data.

Soft targets did not help, and hotter temperatures hurt. 0.726 at T=1, 0.639 at T=2, 0.547 at T=4. A 610-parameter linear student cannot match a 10-way softened distribution and get the argmax right; forcing it to try spends capacity on the wrong objective.

That is the honest picture, and it is not the picture folklore gives. Soft-target gains are real in some settings — small capacity gaps, well-calibrated teachers, image classification — and absent in others. Measure it on your own setup instead of assuming.

The teacher's soft view is nearly one-hot anyway. 5: 0.894, 4: 0.106, everything else zero, even at T=4. A confident teacher has little dark knowledge left to transfer. Teacher calibration is a precondition for soft-target distillation, not a detail.

How LLM distillation is actually done

Almost all of it is the "hard label" column above, applied to sequences.

python
# 1. collect prompts (no answers needed)
# 2. generate answers with the teacher
outputs = teacher.generate(**prompts, max_new_tokens=512, do_sample=True, temperature=0.7)
# 3. filter: drop anything a verifier says is wrong
# 4. run ordinary SFT on the surviving (prompt, answer) pairs

No output block — this needs a large model and a GPU, and its outputs vary by sampling seed.

Step 3 is what separates a good distillation dataset from a bad one. For maths and code you can verify the teacher's answer and keep only the correct ones — rejection sampling, and it is the highest-value step in the recipe.

DeepSeek-R1's distilled models were built exactly this way: around 800k samples generated by R1, then plain SFT onto Qwen and Llama base models, with no RL stage on the students. The reported result is striking — a 7B distilled model beating much larger non-reasoning models on maths benchmarks.

The logit-matching variant does appear at pretraining scale, where the teacher is available offline and the capacity gap is modest. Gemma 2's smaller models and Llama 3.2's 1B and 3B were both trained with teacher logits as a supervision signal alongside next-token prediction.

Common mistakes

No filtering. Teacher answers are wrong some of the time. Unfiltered, those errors become permanent student behaviour. Verify what you can, sample several answers and keep the majority where you cannot.

Distilling from a model whose licence forbids it. Many API terms prohibit using outputs to train a competing model. This is a legal question, not a technical one, and it is worth answering before you spend the compute.

Ignoring tokenizer mismatch. Logit distillation requires the teacher and student to share a vocabulary. Sequence-level distillation does not, which is another reason it dominates in practice.

Distilling only the final answers of a reasoning model. The reasoning trace is the valuable part. Keep it, and train the student to produce it.

Assuming the student generalises beyond the transfer set. It is good at what the teacher covered. Build your transfer set from the distribution you will actually serve.

Evaluating the student on the teacher's outputs. That measures imitation. Evaluate on held-out data with real labels.

Try it yourself

Reduce transfer_x from 20000 to 2000 and re-run. Watch the hard-label row fall. Then reduce it to 200 — the same size as the human-labelled set — and watch it approach the first row. That sweep isolates exactly what distillation is buying you.

What to learn next

Researcher — Mathematics and papers.

Two objectives

Logit-level (Hinton et al., 2015). Match the teacher's softened output distribution:

$$ \mathcal{L} = (1-\alpha)\,\mathcal{L}_{\mathrm{CE}}(y, \sigma(z_s)) + \alpha\,T^2\, D_{\mathrm{KL}}!\left(\sigma(z_t/T)\ |\ \sigma(z_s/T)\right) $$

$z_s$ and $z_t$ are student and teacher logits, $T$ the temperature, $\alpha$ the mixing weight. The $T^2$ factor exists because the gradient of the softened cross-entropy scales as $1/T^2$; without it, changing $T$ silently rescales the learning rate.

Hinton et al.'s argument for $T>1$: the relative probabilities of the incorrect classes encode the teacher's learned similarity structure, and a sharp softmax hides them. The code above shows the counterweight — a low-capacity student forced to match a full distribution can do worse than one matching the argmax.

Sequence-level (Kim and Rush, 2016). For sequence models, matching per-token distributions is not the same as matching the sequence distribution. Their sequence-level KD approximates the intractable sum over all sequences with the teacher's own beam-search output, reducing to ordinary supervised training on teacher-generated text. They report it outperforming word-level KD and allowing a 1000× smaller decoder with a small BLEU drop.

This is why modern LLM distillation is sequence-level. It needs no shared vocabulary, no teacher logits at training time, and it composes with rejection sampling.

On-policy distillation and the exposure-bias fix

Offline sequence-level KD trains the student on the teacher's distribution, so at inference the student conditions on its own prefixes and drifts.

  • MiniLLM (Gu et al., 2024) minimises reverse KL $D_{\mathrm{KL}}(\pi_s | \pi_t)$ rather than forward KL, arguing the mode-seeking direction is right when the student cannot cover the teacher's full distribution — it prevents the student assigning mass to regions the teacher considers improbable.
  • GKD (Agarwal et al., 2024, Generalized Knowledge Distillation) samples from the student during training and scores those samples with the teacher, directly removing the train/inference mismatch. It also generalises the divergence to the Jensen–Shannon family.

Both consistently beat offline KD at equal budget, at the cost of running teacher inference inside the training loop.

Distillation scaling laws

Busbridge et al., 2025 (Distillation Scaling Laws, Apple) fit the student's cross-entropy as a function of teacher size, student size, and distillation token budget. Their findings are the most practically useful in this area:

  • Distillation beats supervised pretraining only when the compute or token budget for the student is below a threshold that grows with student size. Above it, plain pretraining wins.
  • There is a capacity-gap effect: a stronger teacher does not always give a better student, and past a point a larger teacher makes the student worse.
  • If the teacher must be trained specifically for this purpose, its cost has to be counted, and distillation is often no longer the better option.

That last point is why distillation is most attractive when a strong teacher already exists for other reasons.

Rejection sampling and self-improvement

Zelikman et al., 2022 (STaR) bootstrap reasoning by sampling rationales, keeping those that reach the correct answer, and fine-tuning on them — self-distillation with a verifier. Yuan et al., 2023 (RFT) and Llama 2's iterative rejection-sampling stage apply the same idea at scale.

The failure mode is model collapse when the loop is closed without fresh data or verification. Shumailov et al., 2024 (The Curse of Recursion / AI models collapse when trained on recursively generated data, Nature 2024) show that indiscriminate training on self-generated data degrades tail behaviour and eventually the whole distribution. Verification and human data in the mixture are what break the cycle.

Weak-to-strong generalisation

The reverse direction is an open problem with practical consequences for alignment. Burns et al., 2023 (OpenAI) fine-tune GPT-4 on labels from a GPT-2-level supervisor and find the strong student exceeds its weak supervisor, recovering a substantial fraction of the gap to full supervision. The framing matters: if humans are the weak supervisors of superhuman models, this is the mechanism that would have to work.

Papers

What to learn next