Knowledge distillation
Be able to train a small model on a large model's output and measure the loss of quality.
Prerequisites
- EFine-tuning language modelsrequired
- EInformation theory: entropy and KL divergencerequired
Intuition
Distillation trains a small model (the student) to imitate a large one (the teacher).
The surprising part is that the student often becomes better than if it had been trained directly on the same labelled data. The explanation lies in what the teacher's output contains.
A hard label says: «this is a 7».
The teacher's soft distribution says: «this is a 7 (0.82), but it resembles a 1 (0.11) and a little a 9 (0.04)».
The second carries information about how the classes relate to one another — Hinton called it «dark knowledge». An image that is hard to tell from a 1 teaches the student something the hard label cannot.
| A hard label | The teacher's distribution | |
|---|---|---|
| Information per example | log₂(10) ≈ 3.3 bits | much more |
| Says something about uncertainty | no | yes |
| Requires labelled data | yes | no — unlabelled is enough |
The last row is often the most important in practice: distillation works on unlabelled data, since the teacher produces the targets.
Formal
The loss combines two terms:
| The part | Means |
|---|---|
| the temperature, typically 2–10 — it softens both distributions | |
| it compensates for the gradients being scaled by | |
| the weight between the teacher and the hard labels, often 0.5–0.9 |
The temperature is the whole point. At the teacher's distribution is often nearly one-hot and carries little more information than the label. At the small probabilities emerge, and that is where the information is.
The factor is often missed. The gradient of the KL term is scaled by , so without the compensation the teacher term becomes negligible at a high temperature.
Three forms of distillation:
| The form | What the student matches | A comment |
|---|---|---|
| Response-based | the output distribution | the simplest, works well |
| Feature-based | the intermediate representations | it requires a projection between different dimensions |
| Sequence level | the teacher's generated sequences | the most common for language models |
For language models the third dominates: let the teacher generate answers to many prompts and fine-tune the student on them. That is in practice what most small «instruct» models have been trained with.
What you can realistically expect:
| The compression | A typical loss of quality |
|---|---|
| 2× | nearly none |
| 4× | small |
| 10× | noticeable, often acceptable |
| 50×+ | large |
Two legal and practical caveats:
- The terms. Many commercial APIs expressly forbid using the output to train competing models. Read them.
- Faults are inherited. The student learns the teacher's mistakes, biases and hallucinations — including those nobody has discovered. Evaluate the student independently, not just against the teacher.
Combine with other things. Distillation, quantisation and pruning are orthogonal: distil into a smaller architecture, quantise it, and prune if needed. Together they give considerably more than each on its own.
Code
import torch, torch.nn as nn, torch.nn.functional as F
def distillation_loss(student_logits, teacher_logits, y=None, T=4.0, alpha=0.7):
soft = F.kl_div(F.log_softmax(student_logits / T, dim=-1),
F.softmax(teacher_logits / T, dim=-1),
reduction="batchmean") * (T ** 2) # T² compensates the gradient scaling
if y is None:
return soft # pure distillation, unlabelled data
hard = F.cross_entropy(student_logits, y)
return alpha * soft + (1 - alpha) * hard
def train_student(student, teacher, dataloader, opt, T=4.0, alpha=0.7, epochs=3):
teacher.eval()
for _ in range(epochs):
for x, y in dataloader:
with torch.no_grad():
tl = teacher(x)
loss = distillation_loss(student(x), tl, y, T=T, alpha=alpha)
opt.zero_grad(); loss.backward(); opt.step()
# Why the temperature matters — the same logits, a different information content
logits = torch.tensor([[8.0, 2.0, 1.5, 0.2, -1.0]])
for T in (1.0, 2.0, 4.0, 8.0):
p = F.softmax(logits / T, dim=-1)[0]
entropy = float(-(p * p.clamp_min(1e-12).log()).sum())
print(f"T={T}: {[round(float(v), 4) for v in p]} entropy {entropy:.3f}")
# T=1.0: [0.9955, 0.0025, 0.0015, 0.0004, 0.0001] entropy 0.033
# T=4.0: [0.5996, 0.1338, 0.1181, 0.0853, 0.0632] entropy 1.213
# ↑ at T=1 the distribution is nearly one-hot; at T=4 the class relations show
# Sequence level for language models: let the teacher generate the training data
def generate_distillation_data(teacher, tokenizer, prompts, max_tokens=512):
pairs = []
for p in prompts:
ids = tokenizer(p, return_tensors="pt")
with torch.no_grad():
out = teacher.generate(**ids, max_new_tokens=max_tokens, do_sample=False)
answer = tokenizer.decode(out[0, ids["input_ids"].shape[1]:], skip_special_tokens=True)
pairs.append({"instruction": p, "answer": answer})
return pairs
# Measure the loss of quality AGAINST the gain in cost — both sides are needed
def evaluate_distillation(teacher, student, testset, run):
results = {}
for name, m in (("teacher", teacher), ("student", student)):
right = sum(run(m, f["question"]) == f["key"] for f in testset) / len(testset)
params = sum(p.numel() for p in m.parameters())
results[name] = {"accuracy": round(right, 4), "parameters_m": round(params / 1e6, 1)}
t, s = results["teacher"], results["student"]
results["compression"] = round(t["parameters_m"] / s["parameters_m"], 1)
results["quality_loss_pp"] = round((t["accuracy"] - s["accuracy"]) * 100, 2)
return results
# {'teacher': {...}, 'student': {...}, 'compression': 8.0, 'quality_loss_pp': 1.8}
Evaluate the student independently, not just against the teacher. A student that perfectly imitates a teacher with a systematic fault has learnt the fault just as well as everything else.
Mastery means
- Trains a student model on a teacher's output
- Explains why soft targets carry more information
- Measures the loss of quality against the gain in cost
Sign in to do the exercises and build your mastery up.
Sources
- arXiv — Distilling the Knowledge in a Neural Network — arXiv (open access; licence per article)
- arXiv — DistilBERT, a distilled version of BERT — arXiv (open access; licence per article)
- Hugging Face — dokumentation (Apache-2.0) — Apache-2.0