๐Ÿ”— Deep Learning ยท Lecture 34 of 38

Knowledge Distillation: Teaching Small Models with Large Ones

A large "teacher" model's soft predictions contain rich information that can train a much smaller "student". We derive the distillation loss with temperature, discuss dark knowledge, and survey feature and LLM distillation.

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$:

$$ p_i^{(T)} = \frac{\exp(z_i/T)}{\sum_j\exp(z_j/T)} $$

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:

$$ \mathcal{L} = \alpha\,T^2\,D_{\text{KL}}\left(\mathbf{p}_{\text{teacher}}^{(T)}\,\Big\|\,\mathbf{p}_{\text{student}}^{(T)}\right) + (1 - \alpha)\,\text{CE}\left(\mathbf{y},\, \mathbf{p}_{\text{student}}^{(1)}\right) $$
  • 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]$.

python
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:

  1. 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).
  2. 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.
  3. 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#

MethodReducesNeeds retraining
DistillationArchitecture size (new smaller model)Yes (train the student)
PruningParameters/structures in the same modelUsually fine-tuning
QuantisationBits per parameterOptional (calibration or QAT)
Low-rank factorisationMatrix sizesUsually fine-tuning

They combine well: distil into a small student, then prune and quantise it for deployment.

JA
Written by

Janin A Apurba

B.Sc. in CSE, AUST ยท Advanced ICT Officer, CNRS-UNHCR. Teaching AI, ML and Deep Learning to the next generation of engineers and researchers.

Keep learning

Related lectures

๐Ÿ”— Deep Learning

Model Compression: Pruning and Quantisation

Neural networks are highly redundant. We remove unnecessary weights with pruning, represent the rest with fewer bits via quantisation, and discuss the lottery ticket hypothesis and deployment on edge devices.

Advancedโฑ 6 min#131
๐Ÿ”— Deep Learning

Graph Neural Networks: Learning on Relational Data

Molecules, social networks, road maps and knowledge graphs are graphs. GNNs learn from them by passing messages between neighbours. We derive message passing, GCN and GAT layers, and survey node, edge and graph-level tasks.

Advancedโฑ 6 min#129
๐Ÿ”— Deep Learning

Distributed Training: Data, Model, Pipeline and Sharded Parallelism

Large models and datasets need many accelerators. We explain data parallelism with all-reduce, sharded data parallelism (ZeRO/FSDP), tensor and pipeline model parallelism, and how they combine at scale.

Advancedโฑ 6 min#128