An annotated companion · AI Primer

Distilling the Knowledge in a Neural Network, annotated

About this page. This is a companion, not a copy. It follows the paper section by section, quotes only a sentence or two per section (clearly marked), and explains everything in its own words. Equations are reproduced with every symbol decoded. Diagrams and demos are drawn from scratch, and only selected numbers from the paper's tables are shown, each with attribution. Read the original alongside: every section links to it.

How to read this page

  • Any dotted word explains itself on hover, focus or tap; so does every symbol in every equation.
  • The temperature slider in §2 is the heart of the paper: drag it and watch hidden knowledge appear.

Each idea climbs the ladder: everyday picture, tiny example, diagram, math, why it matters today. The training stages lesson implements the distillation loss from scratch, and the cost lesson shows why distilled models are often the biggest saving in production.

Abstract · original

“Caruana and his collaborators have shown that it is possible to compress the knowledge in an ensemble into a single model which is much easier to deploy and we develop this approach further using a different compression technique.”Hinton, Vinyals and Dean (2015), Abstract

Everyday picture

A master chef can't open a restaurant in every town, so they train apprentices. Handing an apprentice only the finished dishes (“this is a soufflé”) teaches a little. Letting them hear the master think aloud (“this is mostly soufflé, a bit like a mousse, nothing like a steak”) teaches far more. Distillation trains a small model (the student) on a big model's (the teacher's) full spread of opinions instead of on bare right answers.

What the paper claims

  • Averaging many models (an ensemble) is accurate but too expensive to deploy. Its knowledge can be squeezed into one small model.
  • The trick is to raise the temperature of the teacher's softmax, so its small probabilities become visible, and train the student to match them.
  • It works on MNIST digits and on a speech recognizer used by Android voice search.
  • A new kind of ensemble, one generalist plus many cheap “specialists”, helps on a dataset of 100 million images.

Why it matters today

Many of the small, fast models you use every day were distilled from larger ones. Whenever someone says “we trained the small model on the big model's outputs”, this paper is the root.

1 Introduction · original

“An image of a BMW, for example, may only have a very small chance of being mistaken for a garbage truck, but that mistake is still many times more probable than mistaking it for a carrot.”Hinton, Vinyals and Dean (2015), §1

Everyday picture

The paper opens with insects: a larva is built for eating and growing, an adult for flying and reproducing, and the two forms look nothing alike. Machine learning usually uses the same model for learning (where you can afford to be huge and slow) and for serving millions of users (where you must be small and fast). Distillation separates the two: learn with a cumbersome model, then transfer what it learned into a compact one.

Tiny example: dark knowledge

A trained model shown a BMW might say 0.99 BMW, 0.009 garbage truck, 0.001 carrot, …. The “right answer” label only says BMW. But the ratios of the wrong answers, garbage truck nine times likelier than carrot, reveal how the model sees the world: cars resemble trucks more than vegetables. That similarity structure is the knowledge worth transferring, and it lives in numbers so small that normal training barely notices them. It has been nicknamed dark knowledge.

Why it matters today

The key conceptual step is to stop thinking of a model's knowledge as its weights and think of it as its input-to-output mapping. Then you can move the knowledge into a model with completely different weights, and a completely different shape.

2 Distillation · original

Everyday picture

A confident teacher's opinions are like a photo so over-exposed that everything except the brightest object is white: 0.997 for the right answer and nothing visible anywhere else. Raising the temperature is like turning down the exposure: the bright object dims and the faint details appear. Train the student on the dimmed photo, at the same temperature, and it sees everything the teacher saw.

Tiny example

Scores (logits) of (4, 1, 0). At temperature 1, e4 = 54.6, e1 = 2.72, e0 = 1, total 58.3, giving probabilities (0.936, 0.047, 0.017). At temperature 4, divide the scores first: (1, 0.25, 0), so e1 = 2.718, e0.25 = 1.284, e0 = 1, total 5.002, giving (0.543, 0.257, 0.200). Same ranking, but the second and third options now carry real information.

In words: “divide every score by the temperature, then apply the usual softmax; a temperature above 1 flattens the probabilities, below 1 sharpens them.”

With the numbers: logits (4, 1, 0) at T = 4: q1 = e1 / (e1 + e0.25 + e0) = 2.718 / 5.002 = 0.543.

In Python:

import math
z, T = [4, 1, 0], 4
# Σ_j exp(z_j / T)
denominator = sum(math.exp(z_j / T) for z_j in z)
round(denominator, 3)  # → 5.002
# q_i for every class i
q = [math.exp(z_i / T) / denominator for z_i in z]
[round(q_i, 3) for q_i in q]  # → [0.543, 0.257, 0.2]

Try it: turn down the exposure

Teacher is shown:
1.0

Soft targets from the teacher at temperature T

0123456789

The hard target (the label), identical for both images

Reading it: the upper bars are the teacher's probabilities for the ten digits; the lower bars are the plain label, which only says “2”. The teacher's scores are illustrative, chosen to mimic the paper's example, not taken from a trained model. At T = 1 the teacher is almost certain: 2 gets over 99% and every other bar is invisible, so the soft targets say barely more than the label. Drag T up. Around 4 to 5, the runner-up appears, and switching between the two images reveals the point the paper makes about MNIST: the labels are identical, but one “2” looks like a 3 and the other like a 7. That is exactly the information a student can learn from soft targets and never from labels. Push T to 20 and everything flattens towards 10% each: too hot, and the useful ranking starts to wash out.

The training objective

The student is trained on a weighted mix of two losses: cross-entropy with the teacher's soft targets, both softened at the same high temperature T, and cross-entropy with the true labels at T = 1. The paper found the best results with a considerably lower weight on the true-label term. After training, the student runs at T = 1.

In words: “match the teacher's softened opinions at temperature T, multiplied by T² to keep the gradients at a sensible size, plus a small amount of matching the true labels at normal temperature.” (The paper describes this weighted average in prose; this is our compact notation for it.)

With the numbers: in the speech experiment of §4 the best temperature was T = 2 and the hard-label term had a relative weight of 0.5, so L = 4 · H(p(2), q(2)) + 0.5 · H(y, q(1)). On three classes, with teacher logits (4, 1, 0), student logits (2, 1, 0) and the true class first: p(2) = (0.736, 0.164, 0.100) and q(2) = (0.506, 0.307, 0.186) give H(p(2), q(2)) = 0.862; at T = 1 the student gives the true class 0.665, so H(y, q(1)) = −ln 0.665 = 0.408; and L = 4 × 0.862 + 0.5 × 0.408 = 3.652.

In Python:

import math
def softmax(z, T=1):
    e = [math.exp(z_i / T) for z_i in z]
    return [e_i / sum(e) for e_i in e]
T, lam = 2, 0.5
# teacher, softened: p^(T)
p_T = softmax([4, 1, 0], T)
# student, softened: q^(T)
q_T = softmax([2, 1, 0], T)
[round(x, 3) for x in p_T], [round(x, 3) for x in q_T]  # → ([0.736, 0.164, 0.1], [0.506, 0.307, 0.186])
# H(p^(T), q^(T))
H_soft = -sum(p * math.log(q) for p, q in zip(p_T, q_T))
# student at T = 1: q^(1)
q_1 = softmax([2, 1, 0])
# H(y, q^(1)) with y = the first class
H_hard = -math.log(q_1[0])
round(H_soft, 3), round(q_1[0], 3), round(H_hard, 3)  # → (0.862, 0.665, 0.408)
# L = T² · H(p, q) + λ · H(y, q)
round(T ** 2 * H_soft + lam * H_hard, 3)  # → 3.652

Why multiply by T²?

Hover or tap to read the gradient size at each temperature.

Reading it: computed live for the “2” above, with an untrained student whose logits are all zero. The x-axis is the temperature; the y-axis is the size of the soft-target gradient, the push training gives the student's logits, on a log scale. Without correction (blue, solid) the push shrinks roughly as 1/T², so at T = 20 it is hundreds of times weaker than at T = 1 and the true-label term would drown it out. Multiplied by T² (green, dashed), it stays in the same range at every temperature, so you can change T while experimenting without re-balancing the two losses. That is the paper's stated reason for the T² factor.

Why it matters today

“Soft targets at a raised temperature plus a small hard-label term, scaled by T²” is still the textbook recipe, exactly as written here.

2.1 Matching logits is a special case of distillation · original

Everyday picture

An earlier method (Caruana and colleagues) had the student copy the teacher's raw scores directly, by minimizing the squared difference between logits. This section shows that copying raw scores is what distillation turns into when the temperature is very high. At moderate temperatures, distillation mostly ignores the teacher's very negative scores, which can be noise, since nothing in the teacher's training pinned them down.

Tiny example

Three classes, T = 10, student logits z = (1, 0, −1) and teacher logits v = (2, 0, −2), both averaging zero. The exact gradient for the first logit is (q1 − p1)/T = −0.00346. The high-temperature approximation gives (z1 − v1) / (N T²) = (1 − 2)/(3 × 100) = −0.00333, within 4%.

In words: “the push on each student logit is the gap between the student's and the teacher's softened probabilities, divided by T; when T is large compared with the logits, and the logits average zero, that is just the gap between the logits themselves, divided by N T², which is what minimizing squared logit differences does.”

With the numbers: q = softmax((0.1, 0, −0.1)), p = softmax((0.2, 0, −0.2)); (q1 − p1)/10 = −0.00346 ≈ (1 − 2)/300 = −0.00333.

In Python:

import math
def softmax(z):
    e = [math.exp(z_i) for z_i in z]
    return [e_i / sum(e) for e_i in e]
# student logits z, teacher logits v
T, z, v = 10, [1, 0, -1], [2, 0, -2]
N = len(z)
# softmax((0.1, 0, -0.1))
q = softmax([z_i / T for z_i in z])
# softmax((0.2, 0, -0.2))
p = softmax([v_i / T for v_i in v])
# exact: (1/T)(q_i − p_i)
round((q[0] - p[0]) / T, 5)  # → -0.00346
# approximation: (z_i − v_i) / (N T²)
round((z[0] - v[0]) / (N * T ** 2), 5)  # → -0.00333

The paper's MNIST results suggest that when the student is much too small to absorb everything, intermediate temperatures work best, a hint that ignoring the teacher's large negative logits helps.

Why it matters today

“Match the logits” and “match the softened probabilities” are both still used for distilling language models; this section explains how they relate.

3 Preliminary experiments on MNIST · original

Everyday picture

A test of whether a student can inherit what a teacher learned from experiences the student never had. The teacher trained on images jiggled by up to two pixels; the student never sees a jiggled image, yet inherits the robustness through the soft targets.

Selected results from §3 of Hinton, Vinyals and Dean (2015): test errors out of 10,000 MNIST digits
ModelTest errors
Big teacher: 2 × 1,200 units, dropout, weight constraints, jittered images67
Small net: 2 × 800 units, no regularization, trained on labels146
Same small net trained to match the teacher's soft targets at T = 2074

Reading it: fewer errors is better. The same small network goes from 146 errors (labels alone) to 74 (soft targets), nearly matching the big teacher's 67, with no dropout and no jittered images of its own. The numbers are from §3 of the paper.

The mythical digit

Next, every 3 was removed from the student's training data. The student made 206 test errors, 133 of them on the 1,010 threes. Most of those mistakes came from its learned bias for the class 3 being too low. After raising that bias by 3.5, it made 109 errors, only 14 on threes: it got 98.6% of the 3s right without ever having seen one. Everything it knew about 3s came from the soft targets on other digits (“this 5 looks a bit like a 3”). Trained on only the 7s and 8s, and with those two biases lowered by 7.6, its error across all ten digits was 13.2%.

Temperature mattered with a tiny student: with 300 or more units per layer, any T above 8 worked about equally; with only 30 units per layer, T between 2.5 and 4 was clearly best.

Why it matters today

The mythical-digit result is the most vivid demonstration that soft targets carry information about the whole data distribution, not just the labelled answer.

4 Experiments on speech recognition · original

Everyday picture

A real product test. The acoustic model of a speech recognizer listens to about a quarter-second of audio and guesses which of 14,000 sound states is being spoken. Ten copies trained from different random starts, averaged, beat any one of them; the question is whether one model the size of a single copy can capture most of the ten's advantage.

Tiny example

The baseline gets 58.9% of frames right, the 10-model ensemble 61.1%, and the distilled single model 60.8%. The ensemble's gain is 61.1 − 58.9 = 2.2 points; the distilled model recovers 60.8 − 58.9 = 1.9 of them, 1.9 / 2.2 = 86%, which is the paper's “more than 80%”. On the metric users care about, word error rate, both the ensemble and the distilled model reach 10.7% against the baseline's 10.9%.

Table 1 of Hinton, Vinyals and Dean (2015), reproduced in summary with attribution
SystemTest frame accuracyWord error rate
Baseline (8 hidden layers × 2,560 units, about 85 million parameters)58.9%10.9%
Ensemble of 10 such models61.1%10.7%
Distilled single model, same size as the baseline60.8%10.7%

Training data: about 2,000 hours of spoken English, roughly 700 million training examples. Temperatures of 1, 2, 5 and 10 were tried; 2 was best.

Why it matters today

Ensembles are the most reliable way to squeeze out extra accuracy and the least deployable. This result turned “train an ensemble, ship a single model” into a standard recipe.

5 Training ensembles of specialists on very big datasets · original

Everyday picture

A hospital does not train every doctor in everything. A general practitioner sees everyone and refers you to a specialist when your case falls in a confusing corner: the specialist knows 300 kinds of mushroom apart but lumps everything else into “not my area”. The paper builds exactly that: one generalist model, plus many small specialists, each focused on a cluster of classes the generalist tends to confuse.

Tiny example

JFT, an internal Google dataset: 100 million labelled images in 15,000 classes. The generalist had taken about six months to train. Each of 61 specialists handled 300 classes plus one “dustbin” class for everything else, started from the generalist's weights, and trained in a few days, all in parallel. Specialist training sets were half examples from the specialist's own classes and half random others, so afterwards the dustbin's score is raised by the log of how much the specialist's classes were oversampled, correcting the bias.

image x generalist (15,000 classes) its top class k specialist A spec. B others: idle combine: min Σ KL final prediction q

Hover or tap a part. Start with the image at the bottom.

Inference with one generalist and many specialists (§5.4 of Hinton, Vinyals and Dean, 2015), drawn from scratch.

Reading it: read upwards. The generalist looks at the image first and names its most likely class k. Only the specialists whose class cluster contains k wake up (orange); the rest stay idle, which is what keeps this cheap. The final prediction is the probability distribution q that disagrees least, in total, with the generalist and every active specialist, measured by KL divergence. The paper finds q by gradient descent separately for each image.

In words: “find the single distribution over all classes that is, in total, closest to the generalist's opinion and to every active specialist's opinion.”

With the numbers: if the generalist says “bridge” and specialist A covers bridges, viaducts and chimneys, then A_k = {A}, and q balances the generalist's 15,000-class view with A's fine distinctions within its 300 classes. On a toy with three classes (bridge, viaduct, everything else), say pg = (0.5, 0.3, 0.2) and pA = (0.8, 0.1, 0.1). Choosing q = pg scores 0 + 0.197 = 0.197; choosing q = pA scores 0.233 + 0 = 0.233; their average, q = (0.65, 0.2, 0.15), scores 0.048 + 0.056 = 0.104, the lowest. For this sum of KL divergences the best q is always the average of the opinions.

In Python:

import math
# KL divergence: Σ p log(p / q)
def KL(p, q):
    return sum(p_c * math.log(p_c / q_c) for p_c, q_c in zip(p, q))
# the generalist's opinion
p_g = [0.5, 0.3, 0.2]
# specialist A, the only member of A_k
p_A = [0.8, 0.1, 0.1]
# KL(p^g, q) + Σ over m in A_k of KL(p^m, q)
def objective(q):
    return KL(p_g, q) + sum(KL(p_m, q) for p_m in [p_A])
round(objective(p_g), 3), round(objective(p_A), 3)  # → (0.197, 0.233)
# the average opinion
q = [(a + b) / 2 for a, b in zip(p_g, p_A)]
[round(q_c, 2) for q_c in q], round(KL(p_g, q), 3), round(KL(p_A, q), 3), round(objective(q), 3)  # → ([0.65, 0.2, 0.15], 0.048, 0.056, 0.104)
Selected results from Table 3 of Hinton, Vinyals and Dean (2015): top-1 accuracy on the JFT development set
SystemConditional accuracyTest accuracy
Baseline generalist43.1%25.0%
+ 61 specialists45.9%26.1%

Test accuracy rose from 25.0% to 26.1%, a 4.4% relative improvement (1.1 / 25.0). Accuracy gains grew with the number of specialists covering the correct class, which is encouraging because specialists are trivially parallel to train. The paper notes that it had not yet shown the specialists could be distilled back into one model.

Why it matters today

Routing inputs to specialists is the idea behind modern mixture-of-experts models, although those learn the routing jointly (§7 explains the difference).

6 Soft targets as regularizers · original

“This shows that soft targets are a very effective way of communicating the regularities discovered by a model trained on all of the data to another model.”Hinton, Vinyals and Dean (2015), §6

Everyday picture

A student with only 3% of the textbook but a tutor who has read all of it can still learn most of the subject: the tutor's graded, nuanced feedback carries far more than 3% of the information.

Tiny example

The 85-million-parameter speech model trained on only 3% of the data (about 20 million examples) with ordinary labels badly overfits: training frame accuracy 67.3%, test accuracy only 44.5%, and that was with early stopping. The same model on the same 3%, trained with soft targets from a model that saw all the data, reaches 57.0% test accuracy, within 2 points of the 58.9% achieved with the full data, and it simply converged, with no early stopping needed.

Reading it: each pair of bars is one training setup: training frame accuracy (grey) and test frame accuracy (blue, striped). With 3% of the data and hard labels, the grey bar is high but the striped blue one collapses: memorising, not learning. With soft targets on the same 3%, the gap nearly closes and test accuracy almost matches the full-data baseline. Numbers from Table 5 of the paper.

The paper suggests the same fix for specialists: train them with soft targets from the generalist for the classes outside their specialty, so they keep their general knowledge while specializing.

Why it matters today

This is the other half of distillation's value: not only a smaller model, but far less labelled data needed, because the teacher's soft outputs are a much richer signal than labels.

7 Relationship to mixtures of experts · original

Everyday picture

A mixture of experts learns, at the same time, what each expert is good at and which expert to send each case to. That is powerful but hard to parallelize: every expert's workload keeps shifting as the router learns. The specialist scheme decides the assignment once, up front, from what the generalist confuses, and then trains every specialist completely independently.

Tiny example

With 61 specialists and 61 machines, the specialists train in the time of one, because no specialist needs anything from another. A jointly trained mixture must keep all experts and the router in sync throughout.

Why it matters today

Modern language models took the other branch: mixture-of-experts layers with a learned router, trained jointly at enormous scale, made practical by specialized infrastructure. The trade-off this section names, joint learning versus easy parallelism, is still the right way to think about it. See the transformer lesson.

8 Discussion · original

Distillation transfers an ensemble's or a big regularized model's knowledge into a small model; on MNIST it works even when whole classes are missing from the transfer set; on speech, nearly all of an ensemble's gain fits into one model of the original size; and specialists can improve a very large model that could not feasibly be ensembled. The open question it leaves, distilling specialists back into a single model, is the ancestor of much later work.

What changed since 2015

In the paperCommon todayWhyWhere to learn more
Distill an ensemble into a same-size modelDistill a large model into a much smaller one (for example DistilBERT: about 40% smaller and 60% faster, keeping about 97% of BERT's language-understanding performance)Serving cost scales with model size; distillation is often the largest single cost savingcost lesson
Match softened class probabilitiesFor language models: match the teacher's next-token distributions, or fine-tune on text the teacher generatedGenerated text is a convenient transfer set when only the teacher's outputs are availabletraining stages lesson
Specialists with a fixed assignmentMixture-of-experts layers with learned routingSpecialized infrastructure made joint training practical at scaletransformer lesson
Temperature on the softmaxTemperature is the same knob used for sampling from language modelsOne formula, two jobs: soften targets for training, or control randomness in generationinference lesson

Glossary

Every term with hover guidance on this page, in one place.