Skip to content
AI-grafen
FAI engineeringModel training and fine-tuning· about 90 min· fast-moving, sources checked often· verified 2026-09-21· EN

Knowledge distillation

Be able to train a small model on a large model's output and measure the loss of quality.

Prerequisites

Intuition

Distillation trains a small model (the student) to imitate a large one (the teacher).

The surprising part is that the student often becomes better than if it had been trained directly on the same labelled data. The explanation lies in what the teacher's output contains.

A hard label says: «this is a 7».

The teacher's soft distribution says: «this is a 7 (0.82), but it resembles a 1 (0.11) and a little a 9 (0.04)».

The second carries information about how the classes relate to one another — Hinton called it «dark knowledge». An image that is hard to tell from a 1 teaches the student something the hard label cannot.

A hard labelThe teacher's distribution
Information per examplelog₂(10) ≈ 3.3 bitsmuch more
Says something about uncertaintynoyes
Requires labelled datayesno — unlabelled is enough

The last row is often the most important in practice: distillation works on unlabelled data, since the teacher produces the targets.

Formal

The loss combines two terms:

L=α T2⋅KL ⁣(σ(zt/T) ∥ σ(zs/T))+(1−α)⋅CE(y,zs)\mathcal{L} = \alpha\, T^2 \cdot \mathrm{KL}\!\left(\sigma(z_t/T) \,\|\, \sigma(z_s/T)\right) + (1-\alpha)\cdot \mathrm{CE}(y, z_s)

The partMeans
TTthe temperature, typically 2–10 — it softens both distributions
T2T^2it compensates for the gradients being scaled by 1/T21/T^2
α\alphathe weight between the teacher and the hard labels, often 0.5–0.9

The temperature is the whole point. At T=1T = 1 the teacher's distribution is often nearly one-hot and carries little more information than the label. At T=4T = 4 the small probabilities emerge, and that is where the information is.

The T2T^2 factor is often missed. The gradient of the KL term is scaled by 1/T21/T^2, so without the compensation the teacher term becomes negligible at a high temperature.

Three forms of distillation:

The formWhat the student matchesA comment
Response-basedthe output distributionthe simplest, works well
Feature-basedthe intermediate representationsit requires a projection between different dimensions
Sequence levelthe teacher's generated sequencesthe most common for language models

For language models the third dominates: let the teacher generate answers to many prompts and fine-tune the student on them. That is in practice what most small «instruct» models have been trained with.

What you can realistically expect:

The compressionA typical loss of quality
2×nearly none
4×small
10×noticeable, often acceptable
50×+large

Two legal and practical caveats:

  1. The terms. Many commercial APIs expressly forbid using the output to train competing models. Read them.
  2. Faults are inherited. The student learns the teacher's mistakes, biases and hallucinations — including those nobody has discovered. Evaluate the student independently, not just against the teacher.

Combine with other things. Distillation, quantisation and pruning are orthogonal: distil into a smaller architecture, quantise it, and prune if needed. Together they give considerably more than each on its own.

Code

import torch, torch.nn as nn, torch.nn.functional as F

def distillation_loss(student_logits, teacher_logits, y=None, 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 ** 2)      # T² compensates the gradient scaling
    if y is None:
        return soft                                        # pure distillation, unlabelled data
    hard = F.cross_entropy(student_logits, y)
    return alpha * soft + (1 - alpha) * hard

def train_student(student, teacher, dataloader, opt, T=4.0, alpha=0.7, epochs=3):
    teacher.eval()
    for _ in range(epochs):
        for x, y in dataloader:
            with torch.no_grad():
                tl = teacher(x)
            loss = distillation_loss(student(x), tl, y, T=T, alpha=alpha)
            opt.zero_grad(); loss.backward(); opt.step()

# Why the temperature matters — the same logits, a different information content
logits = torch.tensor([[8.0, 2.0, 1.5, 0.2, -1.0]])
for T in (1.0, 2.0, 4.0, 8.0):
    p = F.softmax(logits / T, dim=-1)[0]
    entropy = float(-(p * p.clamp_min(1e-12).log()).sum())
    print(f"T={T}: {[round(float(v), 4) for v in p]}  entropy {entropy:.3f}")
# T=1.0: [0.9955, 0.0025, 0.0015, 0.0004, 0.0001]  entropy 0.033
# T=4.0: [0.5996, 0.1338, 0.1181, 0.0853, 0.0632]  entropy 1.213
#  ↑ at T=1 the distribution is nearly one-hot; at T=4 the class relations show

# Sequence level for language models: let the teacher generate the training data
def generate_distillation_data(teacher, tokenizer, prompts, max_tokens=512):
    pairs = []
    for p in prompts:
        ids = tokenizer(p, return_tensors="pt")
        with torch.no_grad():
            out = teacher.generate(**ids, max_new_tokens=max_tokens, do_sample=False)
        answer = tokenizer.decode(out[0, ids["input_ids"].shape[1]:], skip_special_tokens=True)
        pairs.append({"instruction": p, "answer": answer})
    return pairs

# Measure the loss of quality AGAINST the gain in cost — both sides are needed
def evaluate_distillation(teacher, student, testset, run):
    results = {}
    for name, m in (("teacher", teacher), ("student", student)):
        right = sum(run(m, f["question"]) == f["key"] for f in testset) / len(testset)
        params = sum(p.numel() for p in m.parameters())
        results[name] = {"accuracy": round(right, 4), "parameters_m": round(params / 1e6, 1)}
    t, s = results["teacher"], results["student"]
    results["compression"] = round(t["parameters_m"] / s["parameters_m"], 1)
    results["quality_loss_pp"] = round((t["accuracy"] - s["accuracy"]) * 100, 2)
    return results
# {'teacher': {...}, 'student': {...}, 'compression': 8.0, 'quality_loss_pp': 1.8}

Evaluate the student independently, not just against the teacher. A student that perfectly imitates a teacher with a systematic fault has learnt the fault just as well as everything else.

Mastery means

  • Trains a student model on a teacher's output
  • Explains why soft targets carry more information
  • Measures the loss of quality against the gain in cost

Sign in to do the exercises and build your mastery up.

Sources

All the sources and licences