Knowledge distillation
Knowledge distillation trains a small model to copy a big model's full opinion rather than its final answer, which lets a 249x smaller network land within a point of its teacher.
- 13 min read
- 3 reading levels
- Updated
Read these first
On this page 9
One lesson, three depths. Pick the one that fits you today — you can switch any time.
Beginner — No maths. Plain English.
The short answer
Knowledge distillation trains a small model to copy a big model, including everything the big model was unsure about.
The analogy you have already lived
Two teachers hand back the same maths paper.
The first writes one red cross. You know you were wrong, and nothing else.
The second writes: "You nearly had it. You treated this minus as a plus. That mistake is easy to make here."
You learn far more from the second teacher. You learned where the traps are.
The correct answer alone is the first teacher. A big model's full opinion is the second teacher. It says how sure it was, and what it nearly said instead.
Why it exists
The usual way to train a small model is to show it examples with the right label attached. The label says "this is a 7", and nothing else.
A big trained model looking at the same image says something richer. "I am 90 out of 100 sure this is a 7. There is some chance it is a 1. There is almost no chance it is an 8."
That extra information is real knowledge. It says sevens and ones look alike, and sevens and eights do not. A plain label can never say that.
Knowledge distillation feeds the small model the big model's whole opinion. The small model learns the shape of the problem, not only the answers.
How it works
1. train a big model normally (the TEACHER)
2. show the teacher lots of inputs
3. record its full opinion for each one ("90% a 7, 8% a 1, 2% other")
4. train a small model to produce the same opinions (the STUDENT)
5. ship the student to the phoneThe teacher's full opinion has a name: soft labels. They are probabilities across all the answers, instead of one hard answer.
There is one dial: temperature. Raising it flattens the teacher's opinion, so the small differences between the wrong answers become visible to the student. That is where most of the extra teaching lives.
The part that makes it worth the trouble
Step 2 says "lots of inputs". It does not say "lots of labelled inputs".
The teacher can label them for you. So distillation works even when you have very little hand-labelled data, as long as you have plenty of raw examples. For most real projects, raw examples are cheap and labels are expensive.
Where you have already seen it
- DistilBERT is a smaller copy of BERT, made this way, and it powers a lot of on-device text work.
- Small versions of large chat models are frequently trained on the bigger model's outputs.
- Phone camera models are routinely distilled from far larger models that could never run on a phone.
The honest part
Distillation needs a good teacher first. If you do not have one, this technique gives you nothing.
It also needs a full training run for the student. Quantisation takes seconds; distillation takes as long as training the small model from scratch, which it is.
And the student rarely matches the teacher exactly. You are trading a point or two of quality for a model that is a hundred times smaller. Sometimes that trade is plainly right, and sometimes it is not.
Remember this
- The student copies the teacher's whole opinion, not the final answer.
- It works without hand labels, because the teacher labels the data.
- It costs a full training run, and needs a good teacher to exist first.
What to learn next
- ONNX — packaging the student so a device can run it.
- Model compression — how distillation stacks with the other three levers.
- Fine-tuning — the other way to get a model to do your task.
Developer — Code and libraries.
Setup
pip install torch scikit-learnThe script below trains four models and finishes in about ten seconds on a CPU. Nothing is downloaded.
Distillation with no labels at all
The experiment compares four things: the teacher, a student trained on 100 hand labels, the same student trained on all 1257 hand labels, and a student trained on nothing but the teacher's opinions.
import torch, torch.nn as nn, torch.nn.functional as F
from sklearn.datasets import load_digits
from sklearn.model_selection import train_test_split
torch.set_num_threads(1)
X, y = load_digits(return_X_y=True) # 1797 tiny 8x8 images, bundled with scikit-learn
Xtr, Xte, ytr, yte = train_test_split(X / 16.0, y, test_size=0.3, random_state=0, stratify=y)
Xtr = torch.tensor(Xtr, dtype=torch.float32); ytr = torch.tensor(ytr, dtype=torch.long)
Xte = torch.tensor(Xte, dtype=torch.float32); yte = torch.tensor(yte, dtype=torch.long)
teacher_net = lambda: nn.Sequential(nn.Linear(64, 512), nn.ReLU(),
nn.Linear(512, 512), nn.ReLU(), nn.Linear(512, 10))
student_net = lambda: nn.Sequential(nn.Linear(64, 16), nn.ReLU(), nn.Linear(16, 10))
def fit_on_labels(model, Xs, ys, epochs):
opt = torch.optim.Adam(model.parameters(), lr=1e-3)
for _ in range(epochs):
for i in torch.randperm(len(Xs)).split(64):
opt.zero_grad(); F.cross_entropy(model(Xs[i]), ys[i]).backward(); opt.step()
return model.eval()
def fit_on_teacher(model, Xs, teacher, T=3.0, epochs=200):
opt = torch.optim.Adam(model.parameters(), lr=1e-3)
with torch.no_grad():
targets = F.softmax(teacher(Xs) / T, dim=1) # soft labels: the teacher's whole opinion
for _ in range(epochs):
for i in torch.randperm(len(Xs)).split(64):
# KL divergence measures how far the student's opinion is from the teacher's.
# The T*T puts the gradients back on the scale they had before the temperature.
loss = F.kl_div(F.log_softmax(model(Xs[i]) / T, 1), targets[i], reduction="batchmean") * T * T
opt.zero_grad(); loss.backward(); opt.step()
return model.eval()
def accuracy(m):
with torch.no_grad(): return (m(Xte).argmax(1) == yte).float().mean().item()
def size(m): return sum(p.numel() for p in m.parameters())
torch.manual_seed(0)
teacher = fit_on_labels(teacher_net(), Xtr, ytr, epochs=30)
few = torch.randperm(len(Xtr), generator=torch.Generator().manual_seed(0))[:100]
torch.manual_seed(100); s_few = fit_on_labels(student_net(), Xtr[few], ytr[few], epochs=200)
torch.manual_seed(100); s_all = fit_on_labels(student_net(), Xtr, ytr, epochs=200)
torch.manual_seed(100); s_dist = fit_on_teacher(student_net(), Xtr, teacher)
print("model params test accuracy")
print("teacher (512-512) %7d %.4f" % (size(teacher), accuracy(teacher)))
print("student, 100 hand labels %7d %.4f" % (size(s_few), accuracy(s_few)))
print("student, all 1257 hand labels %7d %.4f" % (size(s_all), accuracy(s_all)))
print("student, distilled from teacher %7d %.4f" % (size(s_dist), accuracy(s_dist)))
print("teacher is %.0fx bigger than the student" % (size(teacher) / size(s_dist)))model params test accuracy teacher (512-512) 301066 0.9759 student, 100 hand labels 1210 0.8185 student, all 1257 hand labels 1210 0.9722 student, distilled from teacher 1210 0.9704 teacher is 249x bigger than the student
What those four rows say
The headline. A 1,210-parameter student reached 0.9704 while its 301,066-parameter teacher reached 0.9759. That is 249 times smaller for half a percentage point.
The distilled student used zero hand labels. Read fit_on_teacher again — ytr never appears in it. Only the input images and the teacher's opinions. Compare that with the 100-label student at 0.8185, and you can see what the teacher was worth: about 15 points of accuracy, for free, from data you already had.
The honest comparison is row three, not row two. With all 1257 real labels, the same tiny student reaches 0.9722 — marginally better than the distilled 0.9704. On this easy dataset, distillation did not beat full supervision. It matched it without needing any labels.
That is the honest claim for distillation on small, clean problems. The win is labels, not accuracy. The accuracy wins reported in the literature come from harder tasks, bigger gaps between teacher and student, and much longer training.
The temperature dial
At temperature 1 the teacher's opinion is often almost one-hot: 0.9995 for the right class, and everything interesting buried in the last decimal places.
Dividing the scores by T before the softmax flattens the distribution, so the relative sizes of the small probabilities become visible. This is where "which wrong answers were close" actually lives.
The * T * T matters. Softening shrinks the gradients by roughly 1/T², so multiplying by T² keeps the soft-label term on the same scale as an ordinary loss term. Drop it and your learning rate silently becomes nine times too small at T=3.
Set T to 1, 3 and 8 and rerun. On this dataset you get 0.9722, 0.9704 and 0.9685 — a spread of two test images, which is noise. Do not conclude that temperature never matters; conclude that this task is too easy to show it. A confident teacher on an easy problem has little dark knowledge to hand over. There is no universal best value, and 2 to 5 is the usual starting range.
Mixing in the real labels
When you do have labels, the standard recipe uses both:
loss = alpha * soft_loss + (1 - alpha) * F.cross_entropy(student_scores, hard_labels)alpha between 0.5 and 0.9 is typical. On this dataset the mixed version does not beat the pure versions — the task is too easy for the extra signal to matter. On a harder one it usually helps.
Common mistakes
Forgetting the T² scale factor. Silent, and it makes distillation look useless.
Distilling from a teacher you have not evaluated. The student inherits the teacher's errors and its biases, faithfully. A teacher that fails on one accent or one skin tone produces a student that fails the same way.
Using a teacher that is too far ahead. A gigantic teacher and a minuscule student can be worse than a moderate teacher and the same student. The student has to be able to represent what it is shown.
Distilling on the wrong inputs. The teacher's opinions are only useful on data resembling what the student will meet. Distilling on clean studio photos and deploying to a scratched phone camera transfers the wrong knowledge.
Assuming the teacher's licence lets you. Many model licences restrict training other models on their outputs. Check before you build a product on it.
Try it yourself
Change the student's hidden width from 16 to 8, then to 4, and rerun each time. The distilled accuracy goes 0.9704, then 0.9500, then 0.7630. Somewhere between 8 and 4 units the student stops being able to hold what the teacher is teaching. That boundary is the real capacity floor for this task, and finding it takes four minutes.
What to learn next
- ONNX — packaging the student so a device can run it.
- Model compression — how distillation stacks with the other three levers.
- Fine-tuning — the other way to get a model to do your task.
Researcher — Mathematics and papers.
The objective
Hinton, Vinyals and Dean (2015) train the student on a convex combination of a soft target term and a hard-label term:
$$ \mathcal{L} = \alpha T^2 \, \mathrm{KL}!\left(\sigma!\left(\frac{z_t}{T}\right) \,\Big|\, \sigma!\left(\frac{z_s}{T}\right)\right) \;+\; (1-\alpha)\, \mathcal{H}!\left(y, \sigma(z_s)\right) $$
- $z_t, z_s \in \mathbb{R}^C$ — teacher and student logits over $C$ classes.
- $\sigma$ — the softmax function.
- $T > 0$ — the temperature. $T = 1$ recovers the ordinary softmax; larger $T$ flattens both distributions.
- $\alpha \in [0,1]$ — the weight on the soft term.
- $\mathcal{H}$ — cross-entropy against the one-hot label $y$.
- $T^2$ — the gradient-rescaling factor derived below.
Why $T^2$, derived
For a single logit $z_{s,i}$, the gradient of the soft cross-entropy term is
$$ \frac{\partial \mathcal{L}{\text{soft}}}{\partial z{s,i}} = \frac{1}{T}\left(\sigma_i!\left(\frac{z_s}{T}\right) - \sigma_i!\left(\frac{z_t}{T}\right)\right) $$
In the high-temperature limit, expanding the softmax to first order with zero-mean logits gives
$$ \frac{\partial \mathcal{L}{\text{soft}}}{\partial z{s,i}} \approx \frac{1}{C T^2}\left(z_{s,i} - z_{t,i}\right) $$
The gradient therefore scales as $T^{-2}$, and multiplying the loss by $T^2$ restores it. Without that factor, $\alpha$ and $T$ interact and the effective learning rate on the soft term changes whenever you tune $T$.
That limit also shows what distillation reduces to at high $T$: matching logits, in the least-squares sense. This is precisely the objective of Ba and Caruana (2014), which predates the temperature formulation.
Why soft targets help
Three explanations, all with evidence, none complete.
Dark knowledge. The relative magnitudes of the non-target probabilities encode a similarity structure over classes that a one-hot label discards. Hinton's original framing.
Label smoothing plus example weighting. Yuan et al. (2020) show that distillation is close to an adaptive label smoothing, and that a worse teacher — even a poorly trained one — still helps, which the dark-knowledge story alone does not predict.
Variance reduction. Menon et al. (2021) argue the soft targets approximate the Bayes class-probabilities $p(y \mid x)$, and that regressing on a lower-variance estimate of the target reduces the student's own estimation variance. This predicts that a better calibrated teacher distils better than one that is only more accurate, which is observed.
Variants
| Variant | What is matched | Note |
|---|---|---|
| Response-based (Hinton, 2015) | output logits | Architecture-agnostic; the default |
| Feature-based (FitNets, Romero et al., 2015) | intermediate activations | Needs a projection when widths differ |
| Attention transfer (Zagoruyko and Komodakis, 2017) | spatial attention maps | Strong for convolutional vision models |
| Relational (Park et al., 2019) | pairwise distances between samples | Transfers structure, not point values |
| Self-distillation | a model to itself | Improves accuracy with no compression |
| Online / deep mutual (Zhang et al., 2018) | peers, trained together | Removes the need for a pretrained teacher |
Self-distillation is the strangest result in this area: a student with identical architecture to its teacher frequently outperforms it (Furlanello et al., 2018, Born-Again Neural Networks). Mobahi et al. (2020) analyse repeated self-distillation as progressively increasing regularisation in a Hilbert-space setting, which also predicts that it eventually degrades — as observed.
At language-model scale
Two distinct practices share the name and should not be confused.
Logit distillation requires access to the teacher's full output distribution. DistilBERT (Sanh et al., 2019) retained roughly 97% of BERT-base's GLUE performance with 40% fewer parameters and 60% faster inference, using a triple loss of soft targets, masked language modelling and cosine embedding alignment.
Sequence-level distillation on generated text (Kim and Rush, 2016) trains on the teacher's decoded outputs rather than its distributions. This is what "training on GPT-4 outputs" means in practice, and it is a weaker signal — one sample from the distribution rather than the distribution itself. Gu et al. (2024), MiniLLM, replace the forward KL with a reverse KL to stop the student wasting capacity on regions the teacher considers unlikely, which matters when the student cannot cover the teacher's full support.
Reading
- Hinton, Vinyals and Dean, Distilling the Knowledge in a Neural Network, 2015 — arxiv.org/abs/1503.02531
- Ba and Caruana, Do Deep Nets Really Need to be Deep?, NeurIPS 2014 — arxiv.org/abs/1312.6184
- Romero et al., FitNets, ICLR 2015 — arxiv.org/abs/1412.6550
- Sanh et al., DistilBERT, 2019 — arxiv.org/abs/1910.01108
- Furlanello et al., Born-Again Neural Networks, ICML 2018 — arxiv.org/abs/1805.04770
- Menon et al., A Statistical Perspective on Distillation, ICML 2021.
- Gou et al., Knowledge Distillation: A Survey, IJCV 2021 — arxiv.org/abs/2006.05525
What to learn next
- ONNX — packaging the student so a device can run it.
- Model compression — how distillation stacks with the other three levers.
- Fine-tuning — the other way to get a model to do your task.