Edge and On-device AI

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.

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 part that makes it worth the trouble
  6. Where you have already seen it
  7. The honest part
  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

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 phone

The 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

bash
pip install torch scikit-learn

The 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.

distil.py
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)))
Output
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:

python
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

VariantWhat is matchedNote
Response-based (Hinton, 2015)output logitsArchitecture-agnostic; the default
Feature-based (FitNets, Romero et al., 2015)intermediate activationsNeeds a projection when widths differ
Attention transfer (Zagoruyko and Komodakis, 2017)spatial attention mapsStrong for convolutional vision models
Relational (Park et al., 2019)pairwise distances between samplesTransfers structure, not point values
Self-distillationa model to itselfImproves accuracy with no compression
Online / deep mutual (Zhang et al., 2018)peers, trained togetherRemoves 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

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.

What to learn next

These follow on from what you just read.

  • Edge and On-device AI

    ONNX

    ONNX is a single open file format for trained models, so a model built in PyTorch can run in C++, Java, JavaScript or on a phone without shipping PyTorch with it.

  • Edge and On-device AI

    TensorFlow Lite

    TensorFlow Lite, now called LiteRT, converts a trained model into a small flat file that a few-megabyte runtime can execute on Android, iOS and microcontrollers.

  • Edge and On-device AI

    Core ML

    Core ML is Apple's on-device model format, which automatically routes each part of a model to the CPU, GPU or Neural Engine — and only runs on Apple hardware.