The most accurate models are often too large, slow or expensive to deploy on a phone, in a browser or at scale. Knowledge distillation, popularised by Hinton, Vinyals and Dean (2015), trains a compact student model to mimic a large teacher. The student often performs far better than the same small model trained on labels alone. Distillation is behind many efficient deployed models, including compact language models such as DistilBERT.
Dark knowledge in soft targets#
A one-hot label for an image of a "2" says only "this is a 2". A trained teacher's output might say: 2 with probability 0.90, 3 with 0.06, 7 with 0.03, and nearly zero for others. These small probabilities encode similarity structure โ this "2" looks a bit like a 3 and a 7, and nothing like a 4. Hinton called this dark knowledge. Learning from full distributions gives the student far more information per example than hard labels.
Temperature#
Softmax outputs of a confident teacher are nearly one-hot, hiding the small probabilities. We soften them with a temperature $T > 1$:
Higher $T$ produces a softer distribution that reveals the relative ranking of wrong classes.
The distillation loss#
The student is trained on a weighted combination of two terms:
- The first term matches the softened teacher distribution.
- The second term is ordinary cross-entropy with the true labels.
- The factor $T^2$ keeps the gradient magnitude of the soft term roughly independent of $T$, since softened gradients scale as $1/T^2$.
Typical settings: $T \in [2, 10]$, $\alpha \in [0.5, 0.9]$.
import torch
import torch.nn.functional as F
def distillation_loss(student_logits, teacher_logits, labels, 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 * T)
hard = F.cross_entropy(student_logits, labels)
return alpha * soft + (1 - alpha) * hard
# Training loop sketch
teacher.eval()
for xb, yb in loader:
with torch.no_grad():
t_logits = teacher(xb) # teacher is frozen
s_logits = student(xb)
loss = distillation_loss(s_logits, t_logits, yb)
opt.zero_grad(); loss.backward(); opt.step()Why does it work so well?#
- Richer supervision โ each example carries a full probability vector rather than one bit of class information.
- Regularisation โ soft targets act like label smoothing informed by real class similarity.
- Easier function โ the teacher has already smoothed away label noise and found a simpler decision function that the student can approximate.
- Unlabelled data โ the teacher can label unlimited unlabelled (or augmented) data for the student.
Variants#
- Feature (hint) distillation โ FitNets match intermediate representations of teacher and student via a small projection; attention-transfer matches attention maps.
- Relational distillation โ match the similarity structure between examples rather than individual outputs.
- Self-distillation โ a model distils into a copy of the same architecture, which surprisingly often improves it; born-again networks iterate this.
- Online / mutual distillation โ several students teach each other during training, without a pre-trained teacher.
- Data-free distillation โ synthesise inputs when the original training data cannot be shared.
Distillation for language models#
For large language models, distillation takes several forms:
- Logit distillation โ match the teacher's next-token distributions (as in DistilBERT, which retained most of BERT's performance with about 40% fewer parameters and higher speed).
- Sequence-level distillation โ train the student on text generated by the teacher (e.g. teacher-written answers or reasoning traces). Many small instruction-following and reasoning models are trained this way.
- On-policy distillation โ the student generates, and the teacher provides token-level feedback on the student's own outputs, reducing the mismatch between training and use.
Where distillation fits among compression methods#
| Method | Reduces | Needs retraining |
|---|---|---|
| Distillation | Architecture size (new smaller model) | Yes (train the student) |
| Pruning | Parameters/structures in the same model | Usually fine-tuning |
| Quantisation | Bits per parameter | Optional (calibration or QAT) |
| Low-rank factorisation | Matrix sizes | Usually fine-tuning |
They combine well: distil into a small student, then prune and quantise it for deployment.