Lab · runnable experiments
Knowledge Distillation: What a Teacher Actually Teaches
Created Sep 12, 2026 Updated Sep 18, 2026
Read the parent noteThe image experiments use CIFAR-10: ResNet teachers trained from scratch, small convolutional students trained against them. The language experiments use a character-level transformer trained on the text of this site.
Part 1 asks why a teacher’s probabilities can be better training targets than the true labels, and takes the answer apart one ingredient at a time. Part 2 corrupts a fraction of the labels the teacher learns from, and asks how much of that reaches the student. Part 3 varies the teacher and the loss: the teacher’s size, a teacher assistant between a large teacher and a tiny student, matching features and relations instead of outputs, the direction of the KL divergence, the temperature, and how long the student trains. Part 4 varies the data the teacher is queried on: which images, which view of them, and which fraction is worth paying for. Part 5 asks whether a student can beat its teacher, plants mistakes in a teacher to see which ones the student inherits, and then separates copying the teacher from being right across every image student in the lab. Part 6 is language models: token-level, top-k and sequence-level distillation on the same teacher. Part 7 is the cost — of distilling a small student, and of the alternative of serving a bigger one.
Setup
pip install torch torchvision numpy pandas matplotlibimport os, time, math, json, datetime, pathlib, glob, re, hashlib
import numpy as np, pandas as pd, torch, torch.nn as nn, torch.nn.functional as F
import matplotlib.pyplot as plt
from torchvision import datasets
DEV = "mps" if torch.backends.mps.is_available() else ("cuda" if torch.cuda.is_available() else "cpu")
CACHE = pathlib.Path(".lab_cache"); CACHE.mkdir(exist_ok=True)
# LAB_SMOKE=1 shrinks every run to a few steps to check the notebook end to end; the published run has it unset
SMOKE = os.environ.get("LAB_SMOKE") == "1"
SCALE = 0.01 if SMOKE else 1.0
steps_ = lambda n: max(5, int(n * SCALE))
INK, LINE = "#2a2a2a", "#d9d3c7"
BLUE, EMBER, FOREST, GRAY, PLUM, GOLD = "#3b6ea5", "#c4521e", "#2e7d5b", "#9a9384", "#7a4f8c", "#b8860b"
plt.rcParams.update({"font.size": 9, "axes.edgecolor": LINE, "axes.spines.top": False,
"axes.spines.right": False, "figure.dpi": 120})
RESULTS = {"meta": {"executed": datetime.datetime.now().isoformat(timespec="seconds"), "device": DEV,
"torch": torch.__version__, "smoke": SMOKE}}
T0 = time.time()
print(RESULTS["meta"]){'executed': '2026-09-18T01:02:00', 'device': 'mps', 'torch': '2.7.1', 'smoke': False}
# CIFAR-10 (and CIFAR-100 images, used only as unlabeled "other" data), normalized, on the device
DATA = str(CACHE / "data")
def load(cls, train):
ds = cls(DATA, train=train, download=True)
return torch.tensor(ds.data).permute(0, 3, 1, 2).float().div(255), torch.tensor(ds.targets), ds.classes
Xtr, Ytr, CLASSES = load(datasets.CIFAR10, True)
Xte, Yte, _ = load(datasets.CIFAR10, False)
X100, _, _ = load(datasets.CIFAR100, True)
MEAN, STD = Xtr.mean((0, 2, 3), keepdim=True), Xtr.std((0, 2, 3), keepdim=True)
RAW_TE = (Xte * 255).round().byte()
norm = lambda x: ((x - MEAN) / STD).to(DEV)
Xtr, Xte, X100 = norm(Xtr), norm(Xte), norm(X100)
Ytr, Yte = Ytr.to(DEV), Yte.to(DEV)
K = 10
print(CLASSES, tuple(Xtr.shape), tuple(Xte.shape), tuple(X100.shape))['airplane', 'automobile', 'bird', 'cat', 'deer', 'dog', 'frog', 'horse', 'ship', 'truck'] (50000, 3, 32, 32) (10000, 3, 32, 32) (50000, 3, 32, 32)
The networks. Teachers are ResNets with three stages; students are plain convolutional stacks. All of them end in a linear layer over the ten classes, so any two can be compared logit for logit.
def cnn(w, depth=3):
layers, c = [], 3
for i in range(depth):
layers += [nn.Conv2d(c, w * 2 ** i, 3, padding=1), nn.BatchNorm2d(w * 2 ** i), nn.ReLU(), nn.MaxPool2d(2)]
c = w * 2 ** i
return nn.Sequential(*layers, nn.AdaptiveAvgPool2d(1), nn.Flatten(), nn.Linear(c, K))
class Block(nn.Module):
def __init__(self, cin, cout, stride):
super().__init__()
self.c1, self.b1 = nn.Conv2d(cin, cout, 3, stride, 1, bias=False), nn.BatchNorm2d(cout)
self.c2, self.b2 = nn.Conv2d(cout, cout, 3, 1, 1, bias=False), nn.BatchNorm2d(cout)
self.sc = (nn.Sequential() if stride == 1 and cin == cout else
nn.Sequential(nn.Conv2d(cin, cout, 1, stride, bias=False), nn.BatchNorm2d(cout)))
def forward(self, x):
return F.relu(self.b2(self.c2(F.relu(self.b1(self.c1(x))))) + self.sc(x))
def resnet(w, n=3):
layers, c = [nn.Conv2d(3, w, 3, 1, 1, bias=False), nn.BatchNorm2d(w), nn.ReLU()], w
for i, mult in enumerate([1, 2, 4]):
for j in range(n):
layers.append(Block(c, w * mult, 2 if (j == 0 and i > 0) else 1)); c = w * mult
return nn.Sequential(*layers, nn.AdaptiveAvgPool2d(1), nn.Flatten(), nn.Linear(c, K))
ARCH = {"cnn8": lambda: cnn(8), "cnn16": lambda: cnn(16), "cnn32": lambda: cnn(32),
"resnet16": lambda: resnet(16), "resnet32": lambda: resnet(32), "resnet64": lambda: resnet(64)}
n_params = lambda m: sum(p.numel() for p in m.parameters())
print({k: f"{n_params(f()):,}" for k, f in ARCH.items()}){'cnn8': '6,474', 'cnn16': '24,458', 'cnn32': '94,986', 'resnet16': '272,474', 'resnet32': '1,084,586', 'resnet64': '4,327,754'}
One training loop serves every run. A run is a model, a set of input images, and a loss that sees the model’s logits for a batch together with whatever that batch’s targets are — true labels, teacher logits, or both. Teachers train with random crops and flips. Students train with flips only, because their teacher’s logits are computed once, for every image and its mirror image, and looked up rather than recomputed.
def crop_flip(x, g):
B = len(x)
p = F.pad(x, (4, 4, 4, 4), mode="reflect")
i = torch.randint(0, 9, (B,), generator=g).to(DEV); j = torch.randint(0, 9, (B,), generator=g).to(DEV)
ar = torch.arange(32, device=DEV)
out = p[torch.arange(B, device=DEV)[:, None, None], :, (i[:, None] + ar)[:, :, None], (j[:, None] + ar)[:, None, :]]
out = out.permute(0, 3, 1, 2)
flip = (torch.rand(B, generator=g) < 0.5).to(DEV)
return torch.where(flip[:, None, None, None], out.flip(3), out)
def fit(model, X, loss_fn, steps, lr=0.05, wd=5e-4, bs=256, seed=0, teacher_aug=False):
"""loss_fn(logits, idx, flip) -> loss. Returns the model and the wall time.
`model` is either a module or a factory returning one. A factory is called *after* the seed is
set, so a run's initial weights depend on its own seed and on nothing that ran before it — which
is what makes a number here reproducible no matter where in the notebook its cell sits."""
torch.manual_seed(seed)
g = torch.Generator().manual_seed(seed)
model = (model if isinstance(model, nn.Module) else model()).to(DEV)
opt = torch.optim.SGD(model.parameters(), lr=lr, momentum=0.9, nesterov=True, weight_decay=wd)
sched = torch.optim.lr_scheduler.OneCycleLR(opt, max_lr=lr, total_steps=steps, pct_start=0.15)
n, t0 = len(X), time.time()
perm, pos = torch.randperm(n, generator=g), 0
model.train()
for _ in range(steps):
if pos + bs > n:
perm, pos = torch.randperm(n, generator=g), 0
idx = perm[pos:pos + bs].to(DEV); pos += bs
if teacher_aug:
x, flip = crop_flip(X[idx], g), None
else:
flip = (torch.rand(bs, generator=g) < 0.5).to(DEV)
x = torch.where(flip[:, None, None, None], X[idx].flip(3), X[idx])
loss = loss_fn(model(x), idx, flip)
opt.zero_grad(set_to_none=True); loss.backward(); opt.step(); sched.step()
model.eval()
if DEV == "mps": torch.mps.synchronize()
return model, time.time() - t0
@torch.no_grad()
def logits_of(model, X, bs=1000):
model.eval()
return torch.cat([model(X[i:i + bs]) for i in range(0, len(X), bs)])
def bank(model, X):
"""Teacher logits for every image and for its mirror image."""
return logits_of(model, X), logits_of(model, X.flip(3))
def look(b, idx, flip):
f = flip.view(-1, *[1] * (b[0].dim() - 1))
return torch.where(f, b[1][idx], b[0][idx])
def ece(logits, y, bins=15):
p = logits.softmax(1); conf, pred = p.max(1)
edges = torch.linspace(0, 1, bins + 1, device=logits.device); e = 0.0
for lo, hi in zip(edges[:-1], edges[1:]):
m = (conf > lo) & (conf <= hi)
if m.any():
e += m.float().mean().item() * abs((pred[m] == y[m]).float().mean().item() - conf[m].mean().item())
return e
def evaluate(model, teacher_te=None):
lg = logits_of(model, Xte); pred = lg.argmax(1)
out = {"acc": (pred == Yte).float().mean().item(), "ece": ece(lg, Yte),
"entropy": (-(lg.log_softmax(1).exp() * lg.log_softmax(1)).sum(1)).mean().item(),
"per_class": [(pred[Yte == c] == c).float().mean().item() for c in range(K)]}
if teacher_te is not None:
tp = teacher_te.argmax(1); twrong = tp != Yte
out["agreement"] = (pred == tp).float().mean().item()
out["kl_to_teacher"] = F.kl_div(lg.log_softmax(1), teacher_te.log_softmax(1), log_target=True,
reduction="batchmean").item()
out["copied_mistakes"] = (pred[twrong] == tp[twrong]).float().mean().item()
return out, lgThe losses. Classic distillation minimizes the KL divergence from the teacher’s softened distribution to the student’s, scaled by T2 so the gradient size does not shrink as the temperature grows. soft_ce is cross-entropy against any target distribution; the ablations below feed it different targets.
def kd(s, t, T=4.0):
"""T² · KL(teacher_T || student_T)."""
return F.kl_div((s / T).log_softmax(1), (t / T).log_softmax(1), log_target=True, reduction="batchmean") * T * T
def reverse_kd(s, t, T=4.0):
"""T² · KL(student_T || teacher_T)."""
ls, lt = (s / T).log_softmax(1), (t / T).log_softmax(1)
return (ls.exp() * (ls - lt)).sum(1).mean() * T * T
def soft_ce(s, p, T=4.0):
"""T² · CE(p, student_T) for an arbitrary target distribution p."""
return -(p * (s / T).log_softmax(1)).sum(1).mean() * T * T
def run_student(arch, X, loss_fn, teacher_te=None, steps=6000, seed=0, lr=0.1, **kw):
model, secs = fit(ARCH[arch], X, loss_fn, steps_(steps), lr=lr, seed=seed, **kw)
res, lg = evaluate(model, teacher_te)
res.update(seconds=secs, seed=seed)
return res, model, lgTeachers
Four teachers of increasing size, trained on the 50,000 labelled CIFAR-10 training images with crops and flips, plus two deliberately flawed teachers used in Part 5. A trained teacher’s weights are cached, so re-running the notebook does not retrain it; the cache is a speed-up, never the record — delete .lab_cache and every teacher is trained again from the same seed.
def train_teacher(name, arch, labels=None, X=None, epochs=30, lr=0.1, seed=0):
X = Xtr if X is None else X
labels = Ytr if labels is None else labels
steps = steps_(epochs * len(X) // 256)
path = CACHE / f"teacher-{name}-{arch}-{steps}-{seed}.pt"
if path.exists():
model = ARCH[arch]().to(DEV)
model.load_state_dict(torch.load(path, map_location=DEV)); secs = json.loads(path.with_suffix(".json").read_text())["seconds"]
else:
model, secs = fit(ARCH[arch], X, lambda s, idx, flip: F.cross_entropy(s, labels[idx]), steps, lr=lr, seed=seed, teacher_aug=True)
torch.save(model.state_dict(), path); path.with_suffix(".json").write_text(json.dumps({"seconds": secs}))
model.eval()
return model, secs
TEACHERS = {"cnn32": "cnn32", "resnet16": "resnet16", "resnet32": "resnet32", "resnet64": "resnet64"}
teachers, rows = {}, []
for name, arch in TEACHERS.items():
m, secs = train_teacher(name, arch)
t_bank, t_te = bank(m, Xtr), logits_of(m, Xte)
res, _ = evaluate(m)
teachers[name] = {"model": m, "bank": t_bank, "te": t_te, "params": n_params(m), "acc": res["acc"], "ece": res["ece"],
"train_seconds": secs}
rows.append({"teacher": name, "parameters": n_params(m), "test accuracy": res["acc"], "ECE": res["ece"],
"train accuracy": (t_bank[0].argmax(1) == Ytr).float().mean().item(), "train minutes": secs / 60})
print(pd.DataFrame(rows).to_string(index=False, float_format=lambda v: f"{v: .4f}"))
TL = teachers["resnet64"] # "the teacher" everywhere a single teacher is needed
RESULTS["teachers"] = rows teacher parameters test accuracy ECE train accuracy train minutes
cnn32 94986 0.8312 0.0208 0.8654 0.7201
resnet16 272474 0.9052 0.0247 0.9634 2.3528
resnet32 1084586 0.9233 0.0324 0.9920 5.0113
resnet64 4327754 0.9367 0.0297 0.9987 15.5475
The four teachers reach 83.1%, 90.5%, 92.3% and 93.7% on the test set; the largest fits its training set almost perfectly (99.9%) and is reasonably calibrated (expected calibration error 0.030).
What a teacher’s probabilities look like
The same test image, the teacher’s distribution over the ten classes at several temperatures. At T = 1 almost all the mass sits on one class; the ranking of the rest — which wrong answers the teacher considers nearly right — only becomes visible as the temperature rises.
picks = []
for c in range(K): # one confidently right and one uncertain test image per class
idx = (Yte == c).nonzero().squeeze(1)
p = TL["te"][idx].softmax(1)
conf = p.max(1).values
right = (p.argmax(1) == c)
order = conf.argsort(descending=True)
sure = idx[order[right[order]][0]].item()
unsure_pool = idx[(conf < 0.7) & right]
picks.append(sure)
if len(unsure_pool):
picks.append(unsure_pool[0].item())
picks = picks[:16]
temps = [1, 2, 4, 8]
fig, axes = plt.subplots(2, 4, figsize=(9, 3.6))
for ax, i in zip(axes.flat, picks[:8]):
for T, col in zip(temps, [INK, BLUE, EMBER, GRAY]):
ax.plot((TL["te"][i] / T).softmax(0).cpu(), "o-", ms=2.5, lw=1, color=col, label=f"T={T}")
ax.set_title(f"{CLASSES[Yte[i]]}", fontsize=8); ax.set_xticks(range(K)); ax.set_xticklabels([c[:3] for c in CLASSES], fontsize=5.5, rotation=90)
axes.flat[0].legend(fontsize=6, frameon=False); plt.tight_layout(); plt.show()
# the dark knowledge in aggregate: average teacher probability of every wrong class, by true class, at T = 4
P4 = (TL["te"] / 4).softmax(1)
sim = torch.stack([P4[Yte == c].mean(0) for c in range(K)]).cpu()
for c in range(K):
s = sim[c].clone(); s[c] = -1
print(f"{CLASSES[c]:>10}: nearest wrong classes {CLASSES[s.argsort(descending=True)[0]]}, {CLASSES[s.argsort(descending=True)[1]]}")
RESULTS["kd-temperature"] = {
"classes": CLASSES, "temperatures": temps,
"examples": [{"label": int(Yte[i]), "logits": [round(v, 4) for v in TL["te"][i].tolist()],
"pixels": RAW_TE[i].permute(1, 2, 0).reshape(-1).tolist()} for i in picks],
"class_similarity_T4": [[round(v, 5) for v in row] for row in sim.tolist()],
"teacher": {"arch": "resnet64", "params": TL["params"], "acc": TL["acc"]},
} airplane: nearest wrong classes bird, ship
automobile: nearest wrong classes truck, ship
bird: nearest wrong classes cat, airplane
cat: nearest wrong classes dog, bird
deer: nearest wrong classes cat, dog
dog: nearest wrong classes cat, horse
frog: nearest wrong classes cat, bird
horse: nearest wrong classes dog, deer
ship: nearest wrong classes airplane, truck
truck: nearest wrong classes automobile, airplane
Averaged over the test set, the teacher’s closest wrong classes are the ones a person would name: automobile and truck, cat and dog, horse and deer, airplane and ship. None of that is in a label.
Part 1 — why a teacher can teach better than the labels
The student is cnn16 (24 thousand parameters), the teacher resnet64 (4.3 million). A label says only which class is right. A teacher’s distribution says three more things at once: how sure it is about this image, which wrong classes are close, and — because it is a smooth function of the image — something about the images near this one. On top of that, any soft target is a regularizer: it asks for less than full confidence.
To take those apart, the student is trained against a ladder of targets, each rung adding one ingredient. Every soft-target rung uses the same loss (soft_ce at T = 4), so only the target changes.
| rung | target | what it adds |
|---|---|---|
labels |
the true class, cross-entropy at T = 1 | — |
smoothing |
true class with mass 1 − ε, the rest spread evenly; ε equal to the teacher’s average non-top mass | softness as such |
confidence |
the teacher’s top probability on its top class, the rest spread evenly | how sure the teacher is about this image |
shuffled |
the teacher’s full distribution with the wrong-class probabilities randomly permuted | the same per-image entropy, but wrong-class structure destroyed |
teacher |
the teacher’s full distribution | which wrong classes are close |
teacher labels |
the teacher’s top class as a hard label, T = 1 | the teacher’s decisions without its probabilities |
Two amounts of data: all 50,000 training images, and the first 5,000. Same number of training steps in both, three seeds each.
SUB = 5000
# one fixed random key per training image decides where its wrong-class probabilities go under 'shuffled',
# so an image gets the same shuffled target every time it is seen (a fresh shuffle each step would average out)
KEY = torch.rand(len(Xtr), K, generator=torch.Generator().manual_seed(0)).to(DEV)
def targets(kind, tl, y, idx):
"""Soft target distributions at T = 4 for one batch; tl = teacher logits, y = true labels, idx = image ids."""
p = (tl / 4).softmax(1)
top = p.argmax(1)
if kind == "teacher":
return p
if kind == "smoothing":
return torch.full_like(p, EPS / (K - 1)).scatter(1, y[:, None], 1 - EPS)
if kind == "confidence":
pm = p.max(1, keepdim=True).values
return ((1 - pm) / (K - 1)).expand_as(p).clone().scatter(1, top[:, None], pm)
if kind == "shuffled":
classes = torch.arange(K, device=p.device, dtype=p.dtype).expand_as(p)
others = classes.scatter(1, top[:, None], float(K)).argsort(1)[:, :K - 1] # the nine wrong classes, in class order
dest = KEY[idx].scatter(1, top[:, None], 2.0).argsort(1)[:, :K - 1] # the same nine, in this image's fixed random order
out = torch.zeros_like(p).scatter(1, dest, p.gather(1, others))
return out.scatter(1, top[:, None], p.gather(1, top[:, None]))
raise ValueError(kind)
p_all = (TL["bank"][0] / 4).softmax(1)
EPS = (1 - p_all.max(1).values).mean().item()
print(f"teacher's average non-top probability at T = 4: ε = {EPS:.3f}")
def ladder_loss(kind, X_bank, Y):
if kind == "labels":
return lambda s, idx, flip: F.cross_entropy(s, Y[idx])
if kind == "teacher labels":
return lambda s, idx, flip: F.cross_entropy(s, look(X_bank, idx, flip).argmax(1))
return lambda s, idx, flip: soft_ce(s, targets(kind, look(X_bank, idx, flip), Y[idx], idx))
# a check that 'shuffled' keeps each image's entropy and top probability and changes nothing else
tl = TL["bank"][0][:512]; y = Ytr[:512]
ar512 = torch.arange(512, device=DEV)
pt, ps = targets("teacher", tl, y, ar512), targets("shuffled", tl, y, ar512)
H = lambda p: -(p * p.clamp_min(1e-12).log()).sum(1)
print(f"shuffled vs teacher: max |Δ entropy| {(H(pt) - H(ps)).abs().max():.2e}, "
f"max |Δ top prob| {(pt.max(1).values - ps.max(1).values).abs().max():.2e}, "
f"same wrong-class ranking in {(pt.argsort(1) == ps.argsort(1)).all(1).float().mean():.1%} of images")
LADDER = ["labels", "smoothing", "confidence", "shuffled", "teacher", "teacher labels"]
ladder = []
for n in (SUB, 50000):
Xn, Yn, bn = Xtr[:n], Ytr[:n], (TL["bank"][0][:n], TL["bank"][1][:n])
for kind in LADDER:
for seed in range(3):
res, model, lg = run_student("cnn16", Xn, ladder_loss(kind, bn, Yn), TL["te"], seed=seed)
fit_targets = logits_of(model, Xn).argmax(1)
res.update(rung=kind, n=n, fits_own_targets=(fit_targets == (bn[0].argmax(1) if kind == "teacher labels" else Yn)).float().mean().item())
ladder.append(res)
print(pd.DataFrame([r for r in ladder if r["n"] == n]).groupby("rung", sort=False)[["acc", "agreement", "ece"]].agg(["mean", "std"]).to_string(float_format=lambda v: f"{v:.4f}"))
RESULTS["kd-ladder"] = {"eps": EPS, "rows": ladder, "rungs": LADDER, "student": "cnn16", "teacher": "resnet64"}teacher's average non-top probability at T = 4: ε = 0.278
shuffled vs teacher: max |Δ entropy| 2.38e-07, max |Δ top prob| 0.00e+00, same wrong-class ranking in 0.0% of images
acc agreement ece
mean std mean std mean std
rung
labels 0.6376 0.0034 0.6445 0.0029 0.1852 0.0037
smoothing 0.5901 0.0040 0.5960 0.0046 0.2615 0.0051
confidence 0.6219 0.0098 0.6278 0.0093 0.2414 0.0110
shuffled 0.6094 0.0036 0.6132 0.0037 0.2569 0.0030
teacher 0.6468 0.0042 0.6537 0.0036 0.2239 0.0038
teacher labels 0.6374 0.0051 0.6445 0.0044 0.1838 0.0037
acc agreement ece
mean std mean std mean std
rung
labels 0.7772 0.0047 0.7864 0.0029 0.0160 0.0027
smoothing 0.7897 0.0015 0.7990 0.0034 0.1033 0.0015
confidence 0.7927 0.0038 0.8035 0.0050 0.1084 0.0015
shuffled 0.7957 0.0045 0.8066 0.0052 0.1068 0.0057
teacher 0.7992 0.0030 0.8103 0.0023 0.1045 0.0012
teacher labels 0.7780 0.0037 0.7886 0.0039 0.0141 0.0014
lad = pd.DataFrame(ladder)
fig, ax = plt.subplots(figsize=(7, 3))
for n, col in [(SUB, EMBER), (50000, FOREST)]:
d = lad[lad.n == n].groupby("rung", sort=False)["acc"]
ax.errorbar(range(len(LADDER)), d.mean(), d.std(), fmt="o-", color=col, capsize=3, label=f"{n:,} images")
ax.axhline(TL["acc"], color=GRAY, ls="--", lw=1, label="teacher")
ax.set_xticks(range(len(LADDER))); ax.set_xticklabels(LADDER); ax.set_ylabel("student test accuracy")
ax.legend(frameon=False, fontsize=8); plt.tight_layout(); plt.show()With all 50,000 images the teacher’s distribution is worth 2.1 points over the labels (77.9% → 80.0%). About half of that is softness as such: uniform smoothing alone gives 78.9%. The teacher’s per-image confidence adds 0.4, and the real ranking of the wrong classes another 0.5 over a scrambled one. The teacher’s decisions as hard labels are worth 0.2 — the teacher fits these training images almost perfectly, so its decisions and the labels are nearly the same thing.
With 5,000 images the picture changes. Uniform smoothing hurts, by 4.3 points: a strong prior that every wrong class is equally likely is a bad one when the data cannot overrule it. The teacher’s per-image confidence recovers most of that, a scrambled ranking of the wrong classes costs 1.5 points again, and the real ranking is worth 3.8 points over the scrambled one. On little data, which wrong classes are close is the part of the teacher’s answer that matters.
Two side effects are visible in the calibration error: every student trained on soft targets at T = 4 is less well calibrated at T = 1 (about 0.10) than a student trained on labels (0.015). Its outputs are nearly as confident as the teacher’s, while its accuracy is thirteen points lower.
Part 2 — when the labels are wrong
Everything so far ran on CIFAR-10’s own labels, which are clean. Real training sets are not, and a teacher is a model fitted to whatever labels it was given. This part corrupts a fraction of the training labels — each corrupted image keeps its picture and is given a different, uniformly random class — trains a teacher on them, and asks what the student gets.
NOISE = [0.0, 0.1, 0.2, 0.4]
def corrupt(p, seed=1234):
"""Symmetric label noise: a fraction p of the training labels become a different random class."""
g = torch.Generator().manual_seed(seed)
n = len(Ytr)
hit = torch.rand(n, generator=g).to(DEV) < p
shift = torch.randint(1, K, (n,), generator=g).to(DEV) # 1..K-1, so never the class it already had
return torch.where(hit, (Ytr + shift) % K, Ytr), hit
noise_rows, noise_teachers = [], {}
for p in NOISE:
yn, hit = corrupt(p)
name = "resnet16" if p == 0 else f"noise{int(p * 100)}" # p = 0 is the Part 3 teacher, already trained
tm, _ = train_teacher(name, "resnet16", labels=yn)
tb, tte = bank(tm, Xtr), logits_of(tm, Xte)
tres, _ = evaluate(tm)
tp = tb[0].argmax(1)
noise_teachers[p] = {"labels": yn, "hit": hit, "bank": tb, "te": tte, "acc": tres["acc"], "ece": tres["ece"],
# of the corrupted images: how many the teacher answers with the label it was given,
# and how many it answers with the class the picture actually shows
"memorized": (tp[hit] == yn[hit]).float().mean().item() if p else 0.0,
"recovered": (tp[hit] == Ytr[hit]).float().mean().item() if p else 0.0}
for seed in range(3):
for mode in ["labels", "distilled", "mix"]:
lf = {"labels": lambda s, idx, flip: F.cross_entropy(s, yn[idx]),
"distilled": lambda s, idx, flip: kd(s, look(tb, idx, flip)),
"mix": lambda s, idx, flip: 0.5 * F.cross_entropy(s, yn[idx]) + 0.5 * kd(s, look(tb, idx, flip))}[mode]
res, _, _ = run_student("cnn16", Xtr, lf, teacher_te=tte, seed=seed)
noise_rows.append({"noise": p, "mode": mode, "seed": seed, "teacher_acc": tres["acc"], **res})
nd = pd.DataFrame(noise_rows)
print(pd.DataFrame([{"noise": p, "teacher acc": t["acc"], "teacher ECE": t["ece"],
"memorized the wrong label": t["memorized"], "answered the true class": t["recovered"]}
for p, t in noise_teachers.items()]).to_string(index=False, float_format=lambda v: f"{v: .4f}"))
print()
print(nd.groupby(["noise", "mode"], sort=False)[["acc", "ece"]].agg(["mean", "std"]).to_string(float_format=lambda v: f"{v:.4f}")) noise teacher acc teacher ECE memorized the wrong label answered the true class
0.0000 0.9052 0.0247 0.0000 0.0000
0.1000 0.8865 0.0601 0.0342 0.8664
0.2000 0.8729 0.1407 0.0335 0.8581
0.4000 0.8442 0.3006 0.0370 0.8281
acc ece
mean std mean std
noise mode
0.0 labels 0.7772 0.0047 0.0160 0.0027
distilled 0.7994 0.0043 0.0733 0.0019
mix 0.7964 0.0004 0.0682 0.0008
0.1 labels 0.7602 0.0016 0.1233 0.0012
distilled 0.7691 0.0040 0.1434 0.0043
mix 0.7737 0.0010 0.1519 0.0005
0.2 labels 0.7454 0.0055 0.2020 0.0054
distilled 0.7638 0.0047 0.2487 0.0047
mix 0.7601 0.0023 0.2398 0.0010
0.4 labels 0.7080 0.0036 0.3206 0.0034
distilled 0.7497 0.0068 0.3808 0.0053
mix 0.7316 0.0054 0.3578 0.0054
The teacher barely learns the corrupted labels at all. At every noise level it answers with the label it was given for only 3–4% of the corrupted images, and answers with the class the picture actually shows for 83–87% of them: its probabilities are cleaner than the data it was trained on. That is what the student receives, and it is why the distilled student pulls ahead by more as the labels get worse — and why mixing the true (that is, noisy) labels back in, which cost nothing on clean data, now costs points.
Whether a teacher launders noise this way is not a given. It is a consequence of how the teacher itself was regularized, which the next cell varies while holding the corrupted labels fixed at 20%.
def teacher_variant(name, labels, epochs, aug, arch="resnet16", seed=0):
"""train_teacher with the augmentation and the length of training as knobs."""
steps = steps_(epochs * len(Xtr) // 256)
path = CACHE / f"teacher-{name}-{arch}-{steps}-aug{int(aug)}-{seed}.pt"
if path.exists():
m = ARCH[arch]().to(DEV)
m.load_state_dict(torch.load(path, map_location=DEV))
else:
m, _ = fit(ARCH[arch], Xtr, lambda s, idx, flip: F.cross_entropy(s, labels[idx]), steps, lr=0.1, seed=seed, teacher_aug=aug)
torch.save(m.state_dict(), path)
m.eval()
return m
yn20, hit20 = corrupt(0.2)
mem_rows, mem_teachers = [], []
for label, epochs, aug in [("crops + flips, 30 epochs", 30, True), ("flips only, 30 epochs", 30, False),
("flips only, 100 epochs", 100, False)]:
tm = teacher_variant(f"n20-{epochs}-{int(aug)}", yn20, epochs, aug)
tb, tte = bank(tm, Xtr), logits_of(tm, Xte)
tres, _ = evaluate(tm)
tp, tprob = tb[0].argmax(1), tb[0].softmax(1)
mem_teachers.append({"teacher": label, "acc": tres["acc"], "ece": tres["ece"],
"memorized": (tp[hit20] == yn20[hit20]).float().mean().item(),
"recovered": (tp[hit20] == Ytr[hit20]).float().mean().item(),
"p(wrong label)": tprob[hit20, yn20[hit20]].mean().item(),
"p(true class)": tprob[hit20, Ytr[hit20]].mean().item()})
for seed in range(3):
res, _, _ = run_student("cnn16", Xtr, lambda s, idx, flip: kd(s, look(tb, idx, flip)), teacher_te=tte, seed=seed)
mem_rows.append({"teacher": label, "seed": seed, "teacher_acc": tres["acc"], **res})
print(pd.DataFrame(mem_teachers).to_string(index=False, float_format=lambda v: f"{v: .4f}"))
print()
print(pd.DataFrame(mem_rows).groupby("teacher", sort=False)[["acc", "ece", "agreement"]].agg(["mean", "std"]).to_string(float_format=lambda v: f"{v:.4f}")) teacher acc ece memorized recovered p(wrong label) p(true class)
crops + flips, 30 epochs 0.8740 0.1390 0.0341 0.8522 0.0551 0.6451
flips only, 30 epochs 0.8018 0.0656 0.2241 0.6022 0.1964 0.4101
flips only, 100 epochs 0.7288 0.1108 0.9030 0.0471 0.7606 0.0672
acc ece agreement
mean std mean std mean std
teacher
crops + flips, 30 epochs 0.7619 0.0009 0.2478 0.0019 0.7791 0.0016
flips only, 30 epochs 0.7555 0.0008 0.1779 0.0021 0.7255 0.0011
flips only, 100 epochs 0.7548 0.0036 0.0252 0.0010 0.6610 0.0006
Take the crops away and the teacher memorizes six times as much of the noise; take them away and train it for a hundred epochs and it memorizes nine tenths of it, putting 0.76 of its probability on labels it was shown and 0.07 on what the picture shows. Its own test accuracy collapses. Its student is still better than it is — a 24-thousand-parameter network cannot reproduce a memorized list of arbitrary answers, so it averages them away, which is the born-again mechanism of Part 5 in a much cruder form.
The last cell in this part asks whether the size of the student changes the picture, since a student that can memorize the noise itself has more to gain from never seeing it.
big_rows = []
for p in [0.2, 0.4]:
t = noise_teachers[p]
for seed in range(2):
for mode in ["labels", "distilled"]:
lf = (lambda s, idx, flip: F.cross_entropy(s, t["labels"][idx])) if mode == "labels" else \
(lambda s, idx, flip: kd(s, look(t["bank"], idx, flip)))
res, m, _ = run_student("cnn32", Xtr, lf, teacher_te=t["te"], seed=seed)
sp = logits_of(m, Xtr).argmax(1)
big_rows.append({"noise": p, "mode": mode, "seed": seed,
"memorized_by_student": (sp[t["hit"]] == t["labels"][t["hit"]]).float().mean().item(), **res})
print(pd.DataFrame(big_rows).groupby(["noise", "mode"], sort=False)[["acc", "ece", "memorized_by_student"]].agg(["mean", "std"]).to_string(float_format=lambda v: f"{v:.4f}"))
RESULTS["kd-noise"] = {"sweep": noise_rows, "teachers": {str(p): {k: t[k] for k in ("acc", "ece", "memorized", "recovered")}
for p, t in noise_teachers.items()},
"memorize": {"teachers": mem_teachers, "rows": mem_rows}, "big_student": big_rows} acc ece memorized_by_student
mean std mean std mean std
noise mode
0.2 labels 0.7851 0.0011 0.1945 0.0000 0.0570 0.0016
distilled 0.8201 0.0000 0.2553 0.0004 0.0189 0.0005
0.4 labels 0.7446 0.0012 0.3224 0.0014 0.0572 0.0004
distilled 0.8038 0.0002 0.3989 0.0002 0.0215 0.0004
A 95-thousand-parameter student gains far more from distillation on noisy labels than the 24-thousand-parameter one does, and ends up repeating the corrupted labels about a third as often as when it was trained on them directly. The size of the gain is set by how much of the noise the student would otherwise have been able to memorize.
Part 3 — which teacher
Bigger is not automatically better: the capacity gap
Two students — cnn8 (6.5 thousand parameters) and cnn16 (24 thousand) — each distilled from the four teachers, pure distillation at T = 4 on all 50,000 images, three seeds. The same students trained on the labels are the baseline. Then the teacher assistant route: resnet16 is first distilled from resnet64, and the small students are distilled from that assistant instead.
cap = []
base = {}
for arch in ("cnn8", "cnn16"):
base[arch] = [run_student(arch, Xtr, lambda s, idx, flip: F.cross_entropy(s, Ytr[idx]), TL["te"], seed=s)[0] for s in range(3)]
for tname, t in teachers.items():
for seed in range(3):
res, _, _ = run_student(arch, Xtr, lambda s, idx, flip, b=t["bank"]: kd(s, look(b, idx, flip)), t["te"], seed=seed)
res.update(student=arch, teacher=tname, teacher_params=t["params"], teacher_acc=t["acc"], route="direct")
cap.append(res)
# the assistant: resnet16 distilled from resnet64, then used as the teacher
ta_model, ta_secs = fit(ARCH["resnet16"], Xtr, lambda s, idx, flip: kd(s, look(TL["bank"], idx, flip)), steps_(6000), lr=0.1, seed=0)
ta_res, _ = evaluate(ta_model, TL["te"])
ta_bank, ta_te = bank(ta_model, Xtr), logits_of(ta_model, Xte)
print(f"assistant (resnet16 distilled from resnet64): test accuracy {ta_res['acc']:.4f}, "
f"agreement with resnet64 {ta_res['agreement']:.4f} (resnet16 trained on labels: {teachers['resnet16']['acc']:.4f})")
for arch in ("cnn8", "cnn16"):
for seed in range(3):
res, _, _ = run_student(arch, Xtr, lambda s, idx, flip: kd(s, look(ta_bank, idx, flip)), TL["te"], seed=seed)
res.update(student=arch, teacher="resnet64 → resnet16", teacher_params=TL["params"], teacher_acc=ta_res["acc"], route="assistant")
cap.append(res)
capd = pd.DataFrame(cap)
print(capd.groupby(["student", "teacher"], sort=False)[["acc", "agreement", "kl_to_teacher"]].agg(["mean", "std"]).to_string(float_format=lambda v: f"{v:.4f}"))
for arch in ("cnn8", "cnn16"):
print(f"{arch} on labels: {np.mean([r['acc'] for r in base[arch]]):.4f} ± {np.std([r['acc'] for r in base[arch]]):.4f}")
RESULTS["kd-capacity"] = {"rows": cap, "baseline": {a: base[a] for a in base}, "assistant": {**ta_res, "seconds": ta_secs},
"teachers": [{k: v for k, v in t.items() if k in ("params", "acc", "ece")} | {"name": n} for n, t in teachers.items()]}assistant (resnet16 distilled from resnet64): test accuracy 0.9065, agreement with resnet64 0.9164 (resnet16 trained on labels: 0.9052)
acc agreement kl_to_teacher
mean std mean std mean std
student teacher
cnn8 cnn32 0.6959 0.0023 0.7536 0.0027 0.3393 0.0025
resnet16 0.7130 0.0023 0.7286 0.0025 0.7457 0.0101
resnet32 0.7132 0.0019 0.7260 0.0021 0.8869 0.0125
resnet64 0.7131 0.0012 0.7247 0.0002 0.9710 0.0033
cnn16 cnn32 0.7781 0.0009 0.8547 0.0015 0.1469 0.0061
resnet16 0.7994 0.0043 0.8161 0.0042 0.4403 0.0087
resnet32 0.7992 0.0039 0.8138 0.0026 0.5522 0.0095
resnet64 0.8012 0.0041 0.8108 0.0037 0.6157 0.0191
cnn8 resnet64 → resnet16 0.7141 0.0013 0.7242 0.0035 0.9640 0.0083
cnn16 resnet64 → resnet16 0.7993 0.0030 0.8099 0.0026 0.6210 0.0124
cnn8 on labels: 0.7021 ± 0.0019
cnn16 on labels: 0.7772 ± 0.0038
No systematic capacity-gap penalty shows up here. For cnn16 the gain levels off at resnet16: a teacher at 90.5% gives 80.0%, one at 93.7% gives 79.9%, and resnet32 in between gives 79.4% ± 0.6 — differences the size of the spread between seeds. cnn8 gains 1.3 points from cnn32 to resnet16 and 0.5 more from everything after it. The assistant route adds nothing, which is consistent with there being no gap for it to bridge. What does grow with the teacher’s size is the distance between student and teacher: the KL divergence from teacher to cnn8 rises from 0.34 with cnn32 to 0.96 with resnet64, and the student is no less accurate for it.
What to match: outputs, features, relations
Output distillation matches ten numbers per image. Feature distillation also asks the student’s intermediate activations to resemble the teacher’s, and that raises two practical questions: the teacher’s layers are wider than the student’s, and nothing says which of its layers corresponds to which of the student’s. The standard answer to the first is a projector — a learned 1×1 convolution that maps the student’s channels to the teacher’s, trained with the student and thrown away after. Relational distillation sidesteps both: it matches how similar the images in a batch are to each other, which is a batch×batch matrix whatever the widths.
Here both layer pairings are tried with a projector, spatially pooled to 4×4 and compared as unit vectors, at a light and a heavy weight β on top of ordinary distillation; the relational loss matches the cosine-similarity matrix of the final pooled features. Teacher resnet64, students cnn8 and cnn16, three seeds; cnn16 again on 5,000 images.
TAP_T = {"mid": 8, "late": 11, "pen": 13} # resnet64: end of stage 2 (128×16×16), stage 3 (256×8×8), pooled 256
TAP_S = {"mid": 6, "late": 10, "pen": 13} # cnn: second conv (2w×16×16), third conv (4w×8×8), pooled 4w
@torch.no_grad()
def teacher_features(model, X):
out = {k: [[], []] for k in TAP_T}
for f, flip in enumerate((False, True)):
for i in range(0, len(X), 500):
x = X[i:i + 500].flip(3) if flip else X[i:i + 500]
for j, layer in enumerate(model):
x = layer(x)
for k, tap in TAP_T.items():
if j == tap:
out[k][f].append((F.adaptive_avg_pool2d(x, 4) if x.dim() == 4 else x).half())
return {k: (torch.cat(v[0]), torch.cat(v[1])) for k, v in out.items()}
class Tapped(nn.Module):
"""A student that keeps its intermediate activations, plus the projectors into the teacher's widths."""
def __init__(self, net, proj):
super().__init__(); self.net, self.proj, self.feats = net, nn.ModuleDict(proj), {}
def forward(self, x):
self.feats = {}
for j, layer in enumerate(self.net):
x = layer(x)
for k, tap in TAP_S.items():
if j == tap:
self.feats[k] = x
return x
unit = lambda v: F.normalize(v.flatten(1).float(), dim=1)
def hint_loss(model, where, TF, idx, flip):
s = F.adaptive_avg_pool2d(model.proj[where](model.feats[where]), 4)
return (unit(s) - unit(look(TF[where], idx, flip))).pow(2).sum(1).mean()
def relation_loss(model, TF, idx, flip):
s, t = unit(model.feats["pen"]), unit(look(TF["pen"], idx, flip))
return (s @ s.T - t @ t.T).pow(2).mean()
TF = teacher_features(TL["model"], Xtr)
widths = {k: TF[k][0].shape[1] for k in ("mid", "late")}
print({k: tuple(v[0].shape) for k, v in TF.items()})
def feature_student(arch, match, beta, X, seed):
built, track = {}, {}
def build():
net = ARCH[arch](); w = net[0].out_channels
built["m"] = Tapped(net, {"mid": nn.Conv2d(2 * w, widths["mid"], 1), "late": nn.Conv2d(4 * w, widths["late"], 1)})
return built["m"]
def loss(s, idx, flip):
m = built["m"]
k = kd(s, look(TL["bank"], idx, flip))
h = (torch.zeros((), device=DEV) if match == "outputs only" else relation_loss(m, TF, idx, flip) if match == "relations"
else hint_loss(m, match.split()[0], TF, idx, flip))
track["match"] = h.detach()
return k + beta * h
model, secs = fit(build, X, loss, steps_(6000), lr=0.1, seed=seed)
res, _ = evaluate(model, TL["te"])
res.update(student=arch, match=match, beta=beta, n=len(X), final_match_loss=float(track["match"]), seconds=secs)
return res
FEAT = [("outputs only", 0.0), ("late features", 1.0), ("late features", 30.0), ("mid features", 1.0), ("relations", 10.0)]
feat_rows = []
for arch in ("cnn8", "cnn16"):
for match, beta in FEAT:
for seed in range(3):
feat_rows.append(feature_student(arch, match, beta, Xtr, seed))
for match, beta in [("outputs only", 0.0), ("late features", 1.0), ("late features", 30.0)]:
for seed in range(3):
feat_rows.append(feature_student("cnn16", match, beta, Xtr[:SUB], seed))
fd_ = pd.DataFrame(feat_rows)
print(fd_.groupby(["n", "student", "match", "beta"], sort=False)[["acc", "agreement", "final_match_loss"]].agg(["mean", "std"]).to_string(float_format=lambda v: f"{v:.4f}"))
RESULTS["kd-features"] = {"rows": feat_rows, "teacher_widths": widths}
del TF{'mid': (50000, 128, 4, 4), 'late': (50000, 256, 4, 4), 'pen': (50000, 256)}
acc agreement final_match_loss
mean std mean std mean std
n student match beta
50000 cnn8 outputs only 0.0 0.7131 0.0012 0.7247 0.0002 0.0000 0.0000
late features 1.0 0.7158 0.0048 0.7253 0.0046 0.6529 0.0083
30.0 0.6881 0.0016 0.6995 0.0027 0.5424 0.0082
mid features 1.0 0.7202 0.0051 0.7311 0.0056 0.3210 0.0021
relations 10.0 0.7115 0.0069 0.7234 0.0059 0.0406 0.0014
cnn16 outputs only 0.0 0.8012 0.0041 0.8108 0.0037 0.0000 0.0000
late features 1.0 0.7997 0.0015 0.8105 0.0006 0.5567 0.0066
30.0 0.7927 0.0010 0.8063 0.0019 0.4350 0.0100
mid features 1.0 0.7973 0.0036 0.8067 0.0034 0.2691 0.0013
relations 10.0 0.8000 0.0021 0.8110 0.0016 0.0319 0.0011
5000 cnn16 outputs only 0.0 0.6499 0.0012 0.6561 0.0041 0.0000 0.0000
late features 1.0 0.6557 0.0032 0.6626 0.0037 0.6128 0.0006
30.0 0.7129 0.0037 0.7203 0.0041 0.4142 0.0031
With all 50,000 images, matching features adds nothing measurable for cnn16, at either layer or either weight, and neither do relations. For the 6-thousand-parameter cnn8, heavy matching of the late features costs 2.7 points (71.4% → 68.7%): its final matching loss stays well above cnn16’s, consistent with a network that small being unable to reproduce 256 channels of the teacher’s features and giving up some accuracy in the attempt. On 5,000 images the same heavy matching is worth 6.7 points (64.6% → 71.3%): each image now supplies thousands of target numbers instead of ten.
The direction of the KL divergence
KL(pT ∥ q) = ∑kpT(k)log pT(k) − ∑kpT(k)log q(k): the first term does not depend on the student at all, so minimizing the forward KL and minimizing cross-entropy against the teacher’s probabilities are the same optimization — same gradient, losses a constant apart. The check below computes both gradients. The reverse direction, KL(q ∥ pT), is a different objective, and the difference is easiest to see where the student cannot represent the teacher: a teacher with two separated modes over forty ordered outcomes, and a student that can only produce one bump.
torch.manual_seed(0)
s = torch.randn(64, K, requires_grad=True); t = torch.randn(64, K)
g_kl = torch.autograd.grad(kd(s, t), s)[0]
g_ce = torch.autograd.grad(soft_ce(s, (t / 4).softmax(1)), s)[0]
const = (soft_ce(s, (t / 4).softmax(1)) - kd(s, t)).item()
H_t = (-((t / 4).softmax(1) * (t / 4).log_softmax(1)).sum(1)).mean().item() * 16
print(f"max |∇KL − ∇CE| = {(g_kl - g_ce).abs().max():.2e}; CE − KL = {const:.6f}, T²·H(teacher) = {H_t:.6f}")
# the toy: 40 ordered outcomes, a two-mode teacher, a one-bump student (a discretized Gaussian: two parameters)
xs = torch.arange(40.0)
def two_modes(sep, w=0.5):
c1, c2 = 19.5 - sep / 2, 19.5 + sep / 2
logits = torch.logsumexp(torch.stack([math.log(w) - (xs - c1) ** 2 / 8, math.log(1 - w) - (xs - c2) ** 2 / 8]), 0)
return logits.log_softmax(0)
def fit_bump(lp, direction, starts=(10.0, 19.5, 29.0)):
best = None
for mu0 in starts:
mu = torch.tensor(mu0, requires_grad=True); ls = torch.tensor(1.0, requires_grad=True)
opt = torch.optim.Adam([mu, ls], lr=0.05)
for _ in range(3000):
lq = (-(xs - mu) ** 2 / (2 * ls.exp() ** 2)).log_softmax(0)
loss = (lp.exp() * (lp - lq)).sum() if direction == "forward" else (lq.exp() * (lq - lp)).sum()
opt.zero_grad(); loss.backward(); opt.step()
if best is None or loss.item() < best[0]:
best = (loss.item(), lq.detach().exp().tolist(), mu.item(), ls.exp().item())
return best
toy = []
for sep in [0, 2, 4, 6, 8, 10, 12, 14, 16, 18, 20]:
lp = two_modes(sep)
f, r = fit_bump(lp, "forward"), fit_bump(lp, "reverse")
toy.append({"sep": sep, "teacher": lp.exp().tolist(), "forward": f[1], "reverse": r[1],
"forward_kl": f[0], "reverse_kl": r[0], "forward_sigma": f[3], "reverse_sigma": r[3]})
print(f"separation {sep:>2}: forward-KL student σ = {f[3]:.2f} (KL {f[0]:.3f}), reverse-KL student σ = {r[3]:.2f} (KL {r[0]:.3f})")
kl_dir = []
for name, fn in [("forward KL", kd), ("reverse KL", reverse_kd)]:
for seed in range(3):
res, _, lg = run_student("cnn16", Xtr, lambda s, idx, flip, fn=fn: fn(s, look(TL["bank"], idx, flip)), TL["te"], seed=seed)
p, pt = lg.softmax(1), TL["te"].softmax(1)
top2 = pt.topk(2, 1).indices
res.update(direction=name, second_choice_mass=p.gather(1, top2[:, 1:]).mean().item(),
teacher_second_choice_mass=pt.gather(1, top2[:, 1:]).mean().item())
kl_dir.append(res)
print(pd.DataFrame(kl_dir).groupby("direction")[["acc", "agreement", "entropy", "ece", "second_choice_mass"]].mean().to_string(float_format=lambda v: f"{v:.4f}"))
print(f"teacher: entropy {(-(TL['te'].softmax(1) * TL['te'].log_softmax(1)).sum(1)).mean():.4f}, "
f"mass on its own second choice {kl_dir[0]['teacher_second_choice_mass']:.4f}")
RESULTS["kd-kl-direction"] = {"toy": toy, "cifar": kl_dir, "grad_check": {"max_abs_diff": (g_kl - g_ce).abs().max().item(),
"ce_minus_kl": const, "teacher_entropy_T2": H_t}}max |∇KL − ∇CE| = 2.33e-09; CE − KL = 36.359882, T²·H(teacher) = 36.359882
separation 0: forward-KL student σ = 2.00 (KL -0.000), reverse-KL student σ = 2.00 (KL 0.000)
separation 2: forward-KL student σ = 2.24 (KL 0.000), reverse-KL student σ = 2.24 (KL 0.000)
separation 4: forward-KL student σ = 2.83 (KL 0.010), reverse-KL student σ = 2.79 (KL 0.011)
separation 6: forward-KL student σ = 3.61 (KL 0.063), reverse-KL student σ = 3.43 (KL 0.074)
separation 8: forward-KL student σ = 4.47 (KL 0.172), reverse-KL student σ = 4.09 (KL 0.226)
separation 10: forward-KL student σ = 5.39 (KL 0.314), reverse-KL student σ = 4.79 (KL 0.480)
separation 12: forward-KL student σ = 6.38 (KL 0.460), reverse-KL student σ = 2.05 (KL 0.689)
separation 14: forward-KL student σ = 7.52 (KL 0.593), reverse-KL student σ = 2.01 (KL 0.692)
separation 16: forward-KL student σ = 8.95 (KL 0.704), reverse-KL student σ = 2.00 (KL 0.693)
separation 18: forward-KL student σ = 11.01 (KL 0.791), reverse-KL student σ = 2.00 (KL 0.693)
separation 20: forward-KL student σ = 14.82 (KL 0.852), reverse-KL student σ = 2.00 (KL 0.693)
acc agreement entropy ece second_choice_mass
direction
forward KL 0.8012 0.8108 0.2573 0.1036 0.1053
reverse KL 0.8018 0.8125 0.2349 0.1107 0.1013
teacher: entropy 0.0945, mass on its own second choice 0.0265
The gradient check is exact to float precision, and the difference between the two losses is the teacher’s entropy (times T2) to six decimals. In the toy, the two directions agree while the modes overlap; from a separation of 12 the forward-KL student keeps widening to cover both modes, while the reverse-KL student collapses onto one of them. On CIFAR-10 the directions give the same accuracy (79.9% and 80.0%): with ten classes a 24-thousand-parameter student can come close enough to the teacher that there is little mass to leave out, and the reverse-KL student is only slightly more confident.
Temperature and the label weight
The loss most implementations use mixes the two signals: α CE(y, q) + (1 − α) T2 KL(pT∥qT). A small grid over both, cnn16 from resnet64 on all images, two seeds.
tgrid = []
for T in (1, 2, 4, 8):
for alpha in (0.0, 0.5):
for seed in range(2):
loss = lambda s, idx, flip, T=T, a=alpha: a * F.cross_entropy(s, Ytr[idx]) + (1 - a) * kd(s, look(TL["bank"], idx, flip), T)
res, _, _ = run_student("cnn16", Xtr, loss, TL["te"], seed=seed)
res.update(T=T, alpha=alpha); tgrid.append(res)
print(pd.DataFrame(tgrid).pivot_table(index="T", columns="alpha", values="acc", aggfunc="mean").to_string(float_format=lambda v: f"{v:.4f}"))
RESULTS["kd-temperature-grid"] = tgridalpha 0.0 0.5
T
1 0.7754 0.7778
2 0.7889 0.7829
4 0.8025 0.7972
8 0.7959 0.7889
T = 1 is barely better than the labels (78.1% against 77.9%); T = 4 is the best of the grid (80.1%). Adding the labels at α = 0.5 is slightly worse at every temperature.
How long the student trains
Every student in this notebook trains for 6,000 steps, about 30 passes over the data, and the same budget is given to the label-trained baseline it is compared against. That is a choice, and it can be wrong in either direction: if one of the two objectives is still improving when the budget runs out, the comparison measures the budget as much as the objective.
BUDGETS = [6000, 12000, 24000, 48000]
pat_rows = []
for mode, lf in [("labels", lambda s, idx, flip: F.cross_entropy(s, Ytr[idx])),
("distilled", lambda s, idx, flip: kd(s, look(TL["bank"], idx, flip)))]:
for steps in BUDGETS:
for seed in range(2):
m, secs = fit(ARCH["cnn16"], Xtr, lf, steps_(steps), lr=0.1, seed=seed)
res, _ = evaluate(m, TL["te"])
pat_rows.append({"mode": mode, "steps": steps, "epochs": round(steps * 256 / len(Xtr)), "seed": seed,
"seconds": secs, "train_acc": (logits_of(m, Xtr).argmax(1) == Ytr).float().mean().item(), **res})
pt = pd.DataFrame(pat_rows)
print(pt.groupby(["mode", "steps"], sort=False)[["acc", "train_acc", "ece", "seconds"]].agg(["mean", "std"]).to_string(float_format=lambda v: f"{v:.4f}"))
RESULTS["kd-patience"] = {"budgets": pat_rows} acc train_acc ece seconds
mean std mean std mean std mean std
mode steps
labels 6000 0.7774 0.0066 0.8340 0.0002 0.0168 0.0032 32.6314 1.0141
12000 0.7894 0.0011 0.8598 0.0003 0.0092 0.0021 62.4976 0.5841
24000 0.7926 0.0035 0.8812 0.0005 0.0101 0.0021 124.1348 1.5239
48000 0.7957 0.0045 0.8930 0.0024 0.0151 0.0033 253.8078 3.8221
distilled 6000 0.8025 0.0049 0.8487 0.0004 0.1026 0.0047 34.3197 0.1098
12000 0.8068 0.0020 0.8642 0.0033 0.1016 0.0003 68.7435 0.1117
24000 0.8156 0.0058 0.8782 0.0007 0.0953 0.0067 140.2894 0.3101
48000 0.8183 0.0022 0.8831 0.0004 0.0928 0.0016 268.7522 11.1506
Both objectives keep paying, and the distilled one pays faster: doubling the budget is worth about two points to it. It levels off after roughly 120 epochs while the label-trained student is still climbing at 240, and the gap between them is widest in between. The training accuracies say why — the label-trained student is fitting its own training set harder and harder, the distilled one much less, so the soft target keeps being useful for longer. Nothing in the extra budget lets the labels catch up: the distilled student at the standard 30 epochs is still ahead of a label-trained student given eight times as long.
The last question is where the student starts. Distillation is often bolted onto a model that has already been trained on labels, and it is worth knowing whether that start is a shortcut or a trap.
init_rows = []
for seed in range(2):
warm, _ = fit(ARCH["cnn16"], Xtr, lambda s, idx, flip: F.cross_entropy(s, Ytr[idx]), steps_(6000), lr=0.1, seed=seed)
r0, _ = evaluate(warm, TL["te"])
warm, _ = fit(warm, Xtr, lambda s, idx, flip: kd(s, look(TL["bank"], idx, flip)), steps_(6000), lr=0.1, seed=seed)
r1, _ = evaluate(warm, TL["te"])
cold, _ = fit(ARCH["cnn16"], Xtr, lambda s, idx, flip: kd(s, look(TL["bank"], idx, flip)), steps_(12000), lr=0.1, seed=seed)
r2, _ = evaluate(cold, TL["te"])
for how, r in [("labels only, 6k steps", r0), ("labels 6k then distilled 6k", r1), ("distilled from scratch, 12k steps", r2)]:
init_rows.append({"how": how, "seed": seed, **r})
print(pd.DataFrame(init_rows).groupby("how", sort=False)[["acc", "agreement"]].agg(["mean", "std"]).to_string(float_format=lambda v: f"{v:.4f}"))
RESULTS["kd-patience"]["init"] = init_rows acc agreement
mean std mean std
how
labels only, 6k steps 0.7774 0.0066 0.7870 0.0038
labels 6k then distilled 6k 0.8007 0.0017 0.8119 0.0010
distilled from scratch, 12k steps 0.8068 0.0020 0.8174 0.0033
A warm start is neither. Half the budget on labels and half on the teacher lands where the whole budget on the teacher lands, within the spread of two runs. What the student was initialized from does not matter here; how many steps it took against the teacher does.
Part 4 — the data the teacher is asked about
A teacher can only say something about the images it is shown. The teacher and the student stay fixed — resnet64 and cnn16, pure distillation, no labels anywhere — and only the transfer set changes:
- all: the 50,000 CIFAR-10 training images;
- 5k: the first 5,000 of them;
- no cats: every training image except the cats;
- animals only: the six animal classes, no vehicle ever shown;
- CIFAR-100: 50,000 natural images from a hundred other classes, none of them a CIFAR-10 class;
- noise: 50,000 images of Gaussian noise;
- mixup of 5k: the same 5,000 images, recombined into 50,000 blends of two images each.
The last one asks a question the others cannot: whether more queries of the teacher help when there are no more images.
def blends(X, n, seed=0):
g = torch.Generator().manual_seed(seed)
a, b = torch.randint(0, len(X), (n,), generator=g), torch.randint(0, len(X), (n,), generator=g)
lam = torch.rand(n, generator=g)[:, None, None, None].to(DEV)
return lam * X[a.to(DEV)] + (1 - lam) * X[b.to(DEV)]
ANIMALS = [CLASSES.index(c) for c in ("bird", "cat", "deer", "dog", "frog", "horse")]
CAT = CLASSES.index("cat")
def cat_shift(model, lg):
"""Read the omitted class once more after shifting its logit. The shift is chosen without any labels and without
the test set: on the full, unlabeled CIFAR-10 training images (cats included), it is the one number that makes the
student answer "cat" as often as the teacher does on those images. Then it is applied to the test logits lg."""
share = (TL["bank"][0].argmax(1) == CAT).float().mean()
s_tr = logits_of(model, Xtr)
lo, hi = -20.0, 20.0
for _ in range(40):
mid = (lo + hi) / 2; adj = s_tr.clone(); adj[:, CAT] += mid
(lo, hi) = (mid, hi) if (adj.argmax(1) == CAT).float().mean() < share else (lo, mid)
adj = lg.clone(); adj[:, CAT] += (lo + hi) / 2
return {"shift": (lo + hi) / 2, "cat_acc": (adj.argmax(1)[Yte == CAT] == CAT).float().mean().item(),
"acc": (adj.argmax(1) == Yte).float().mean().item()}
g = torch.Generator().manual_seed(0)
transfer = {
"all": Xtr,
"5k": Xtr[:SUB],
"no cats": Xtr[Ytr != CAT],
"animals only": Xtr[torch.isin(Ytr.cpu(), torch.tensor(ANIMALS)).to(DEV)],
"CIFAR-100": X100,
"noise": torch.randn(50000, 3, 32, 32, generator=g).to(DEV),
"mixup of 5k": blends(Xtr[:SUB], 50000),
}
data_rows = []
for name, Xs in transfer.items():
b = bank(TL["model"], Xs)
for seed in range(3):
res, model, lg = run_student("cnn16", Xs, lambda s, idx, flip, b=b: kd(s, look(b, idx, flip)), TL["te"], seed=seed)
res.update(transfer=name, n_images=len(Xs), cat_shift=cat_shift(model, lg) if name == "no cats" else None,
vehicle_acc=float(np.mean([res["per_class"][c] for c in range(K) if c not in ANIMALS])),
animal_acc=float(np.mean([res["per_class"][c] for c in ANIMALS])))
data_rows.append(res)
del b
dd = pd.DataFrame(data_rows)
dd["cat_acc"] = dd.per_class.map(lambda v: v[CAT])
print(dd.groupby("transfer", sort=False)[["acc", "agreement", "cat_acc", "vehicle_acc", "animal_acc"]].agg(["mean", "std"]).to_string(float_format=lambda v: f"{v:.4f}"))
nc = [r["cat_shift"] for r in data_rows if r["cat_shift"]]
print(f"'no cats', with the cat logit shifted (shift chosen on unlabeled training images): cat accuracy {np.mean([c['cat_acc'] for c in nc]):.4f}, overall {np.mean([c['acc'] for c in nc]):.4f}")
# the control: a label-trained student on the same no-cat images, put through the same shift. Its cat output was
# only ever pushed down, so whatever it recognizes after the shift comes from the procedure, not from a teacher
Y_nocat = Ytr[Ytr != CAT]
hard_nocat = []
for seed in range(3):
res, model, lg = run_student("cnn16", transfer["no cats"], lambda s, idx, flip: F.cross_entropy(s, Y_nocat[idx]), TL["te"], seed=seed)
res.update(cat_shift=cat_shift(model, lg)); hard_nocat.append(res)
hc = [r["cat_shift"] for r in hard_nocat]
print(f"labels, no cats: accuracy {np.mean([r['acc'] for r in hard_nocat]):.4f}, cat accuracy {np.mean([r['per_class'][CAT] for r in hard_nocat]):.4f}; "
f"with the same shift: cat accuracy {np.mean([c['cat_acc'] for c in hc]):.4f}, overall {np.mean([c['acc'] for c in hc]):.4f}")
RESULTS["kd-data"] = {"rows": data_rows, "hard_no_cats": hard_nocat, "classes": CLASSES, "animals": ANIMALS, "cat": CAT} acc agreement cat_acc vehicle_acc animal_acc
mean std mean std mean std mean std mean std
transfer
all 0.8012 0.0041 0.8108 0.0037 0.6377 0.0025 0.8712 0.0056 0.7544 0.0031
5k 0.6499 0.0012 0.6561 0.0041 0.4673 0.0180 0.7450 0.0053 0.5864 0.0017
no cats 0.7557 0.0030 0.7651 0.0035 0.0000 0.0000 0.8769 0.0034 0.6748 0.0055
animals only 0.4710 0.0021 0.4771 0.0010 0.6647 0.0080 0.0008 0.0010 0.7844 0.0037
CIFAR-100 0.6792 0.0047 0.6938 0.0062 0.6573 0.0042 0.7762 0.0013 0.6145 0.0074
noise 0.1313 0.0159 0.1318 0.0155 0.6370 0.3620 0.0071 0.0023 0.2142 0.0274
mixup of 5k 0.7064 0.0015 0.7163 0.0022 0.4853 0.0146 0.8068 0.0043 0.6396 0.0044
'no cats', with the cat logit shifted (shift chosen on unlabeled training images): cat accuracy 0.4297, overall 0.7684
labels, no cats: accuracy 0.7380, cat accuracy 0.0000; with the same shift: cat accuracy 0.2583, overall 0.7211
- Omitting a class. The student distilled without a single cat never outputs cat (0.0% on test cats), yet overall it is better than a label-trained student on the same images (75.4% against 73.9%). Shifting its cat logit, so that on the unlabeled training images it predicts cat as often as the teacher does, makes it recognize 42.6% of test cats, and overall accuracy rises to 76.5%. The shift is never chosen on the test set. The control — the label-trained student put through the same shift — recognizes 25.6% of test cats while its overall accuracy falls to 72.2%. The procedure alone recovers some cats; the teacher adds 17 points of them and turns the overall loss into a gain.
- Animals only. No vehicle transfers (0.1% on vehicles); the animal classes, relieved of any vehicle confusion, reach 79.1%.
- Other images. CIFAR-100, which contains no CIFAR-10 class, gives 67.1% — better than 5,000 images of the real classes (64.6%).
- Noise gives 12.7%.
- More queries, same images. 50,000 blends of the same 5,000 images give 71.3%, 6.7 points more than the 5,000 images themselves.
Which view of the image the teacher is asked about
Every student so far trained with flips only, because the teacher’s logits were computed once, for each image and its mirror, and looked up rather than recomputed. That is the cheap implementation, and it has a cost that is usually left unmeasured: as soon as the student’s augmentation does anything the lookup does not cover — a random crop, say — the target is the teacher’s answer about a different picture from the one the student is looking at.
This part compares that against the expensive alternative, the teacher running inside the training loop and answering about exactly the view the student sees. One training loop serves both: flips are always explicit, so a cached bank can still be indexed, the crop is a knob, and the loss is handed the augmented batch as well.
def fit_view(model, X, loss_fn, steps, lr=0.1, wd=5e-4, bs=256, seed=0, crop=False):
"""The lab's fit, with the augmented batch passed on: loss_fn(logits, idx, flip, x)."""
torch.manual_seed(seed)
g = torch.Generator().manual_seed(seed)
model = (model if isinstance(model, nn.Module) else model()).to(DEV)
opt = torch.optim.SGD(model.parameters(), lr=lr, momentum=0.9, nesterov=True, weight_decay=wd)
sched = torch.optim.lr_scheduler.OneCycleLR(opt, max_lr=lr, total_steps=steps, pct_start=0.15)
n, t0, ar = len(X), time.time(), torch.arange(32, device=DEV)
perm, pos = torch.randperm(n, generator=g), 0
model.train()
for _ in range(steps):
if pos + bs > n:
perm, pos = torch.randperm(n, generator=g), 0
idx = perm[pos:pos + bs].to(DEV); pos += bs
flip = (torch.rand(bs, generator=g) < 0.5).to(DEV)
x = torch.where(flip[:, None, None, None], X[idx].flip(3), X[idx])
if crop:
pad = F.pad(x, (4, 4, 4, 4), mode="reflect")
i = torch.randint(0, 9, (bs,), generator=g).to(DEV); j = torch.randint(0, 9, (bs,), generator=g).to(DEV)
x = pad[torch.arange(bs, device=DEV)[:, None, None], :, (i[:, None] + ar)[:, :, None],
(j[:, None] + ar)[:, None, :]].permute(0, 3, 1, 2).contiguous()
loss = loss_fn(model(x), idx, flip, x)
opt.zero_grad(set_to_none=True); loss.backward(); opt.step(); sched.step()
model.eval()
if DEV == "mps": torch.mps.synchronize()
return model, time.time() - t0
@torch.no_grad()
def teacher_now(x):
return TL["model"](x)
@torch.no_grad()
def averaged_bank(X, n_crops=8, seed=7):
"""The teacher's logits averaged over n random crops — a better cached target, at no cost to the student."""
g = torch.Generator().manual_seed(seed)
out = []
for src in (X, X.flip(3)):
acc = torch.zeros(len(src), K, device=DEV)
for _ in range(n_crops):
for i in range(0, len(src), 1000):
acc[i:i + 1000] += logits_of(TL["model"], crop_flip(src[i:i + 1000], g).contiguous(), bs=1000)
out.append(acc / n_crops)
return tuple(out)
AVG = averaged_bank(Xtr, 2 if SMOKE else 8)
def view_losses(Y, B, A=None):
"""The six regimes, bound to one transfer set: its labels, its cached teacher bank, its averaged bank."""
d = {"labels": (False, lambda s, idx, flip, x: F.cross_entropy(s, Y[idx])),
"labels + crops": (True, lambda s, idx, flip, x: F.cross_entropy(s, Y[idx])),
"cached, same view": (False, lambda s, idx, flip, x: kd(s, look(B, idx, flip))),
"cached, crops": (True, lambda s, idx, flip, x: kd(s, look(B, idx, flip))),
"teacher in the loop": (True, lambda s, idx, flip, x: kd(s, teacher_now(x)))}
if A is not None:
d["cached, 8-crop average"] = (False, lambda s, idx, flip, x: kd(s, look(A, idx, flip)))
return d
view_rows = []
for n, X, Y, B, A in [(len(Xtr), Xtr, Ytr, TL["bank"], AVG),
(SUB, Xtr[:SUB], Ytr[:SUB], (TL["bank"][0][:SUB], TL["bank"][1][:SUB]), None)]:
for name, (crop, lf) in view_losses(Y, B, A).items():
for seed in range(3):
m, secs = fit_view(ARCH["cnn16"], X, lf, steps_(6000), lr=0.1, seed=seed, crop=crop)
res, _ = evaluate(m, TL["te"])
view_rows.append({"n": n, "regime": name, "seed": seed, "seconds": secs,
"train_acc": (logits_of(m, X).argmax(1) == Y).float().mean().item(), **res})
vd = pd.DataFrame(view_rows)
print(vd.groupby(["n", "regime"], sort=False)[["acc", "agreement", "seconds"]].agg(["mean", "std"]).to_string(float_format=lambda v: f"{v:.4f}")) acc agreement seconds
mean std mean std mean std
n regime
50000 labels 0.7772 0.0047 0.7864 0.0029 31.9639 0.7489
labels + crops 0.7793 0.0040 0.7909 0.0028 37.5914 1.1268
cached, same view 0.8012 0.0041 0.8108 0.0037 32.8688 0.5635
cached, crops 0.7987 0.0029 0.8098 0.0015 40.5665 2.0658
teacher in the loop 0.8007 0.0027 0.8135 0.0009 294.5059 14.4856
cached, 8-crop average 0.8014 0.0033 0.8112 0.0015 33.0179 0.3041
5000 labels 0.6376 0.0034 0.6445 0.0029 31.3926 0.3969
labels + crops 0.7068 0.0027 0.7129 0.0036 38.5774 2.6497
cached, same view 0.6499 0.0012 0.6561 0.0041 35.1934 2.2932
cached, crops 0.7402 0.0028 0.7502 0.0039 40.6568 0.6322
teacher in the loop 0.7493 0.0063 0.7595 0.0071 294.0674 28.3662
On all 50,000 images every distillation route lands within a fifth of a point of every other: the teacher in the loop, at roughly eight times the training cost, buys nothing over a lookup table, and so does averaging the teacher over eight crops. On 5,000 images the crop itself is worth ten points — and the mismatch it introduces, the teacher answering about the uncropped picture, costs only a fraction of one. The cheap implementation keeps almost all of the benefit.
A crop, then, does not really ask the teacher anything new. A blend of two images might. The part above found that recombining 5,000 images into 50,000 blends was worth 6.7 points, with the teacher asked about each blend. The cell below asks whether that is the asking or only the blending: the same blends, with the teacher’s answer replaced by an interpolation of the two answers it had already given about the originals — which costs no teacher passes at all.
def blend_parts(X, n, seed=0):
g = torch.Generator().manual_seed(seed)
a = torch.randint(0, len(X), (n,), generator=g).to(DEV)
b = torch.randint(0, len(X), (n,), generator=g).to(DEV)
lam = torch.rand(n, generator=g)[:, None, None, None].to(DEV)
return a, b, lam
# the transfer sets above are the largest thing on the device and are finished with; the blends
# below are just as large, so they are built in chunks and the old ones are freed first
del transfer, X100
if DEV == "mps": torch.mps.empty_cache()
Xs, Bs = Xtr[:SUB], (TL["bank"][0][:SUB], TL["bank"][1][:SUB])
a, b, lam = blend_parts(Xs, 50000)
l2 = lam[:, :, 0, 0]
Xb = torch.empty(len(a), *Xs.shape[1:], device=DEV)
for i in range(0, len(a), 5000):
sl = slice(i, i + 5000)
torch.lerp(Xs[b[sl]], Xs[a[sl]], lam[sl], out=Xb[sl])
Bb = bank(TL["model"], Xb) # the teacher asked about each blend
MIXL = tuple(l2 * Bs[v][a] + (1 - l2) * Bs[v][b] for v in (0, 1)) # its two old answers, mixed as logits
MIXP = tuple(l2 * (Bs[v][a] / 4).softmax(1) + (1 - l2) * (Bs[v][b] / 4).softmax(1) for v in (0, 1)) # mixed as probabilities
blend_check = {"kl": F.kl_div((MIXL[0] / 4).log_softmax(1), (Bb[0] / 4).log_softmax(1), log_target=True, reduction="batchmean").item(),
"top1_agreement": (MIXL[0].argmax(1) == Bb[0].argmax(1)).float().mean().item()}
print(f"interpolated answer against the teacher's real answer about the blend: KL {blend_check['kl']:.3f}, "
f"same top class {blend_check['top1_agreement']:.3f}")
blend_rows = []
for name, lf in [("teacher asked about the blend", lambda s, idx, flip: kd(s, look(Bb, idx, flip))),
("answers mixed as logits", lambda s, idx, flip: kd(s, look(MIXL, idx, flip))),
("answers mixed as probabilities", lambda s, idx, flip: soft_ce(s, look(MIXP, idx, flip)))]:
for seed in range(3):
res, _, _ = run_student("cnn16", Xb, lf, teacher_te=TL["te"], seed=seed)
blend_rows.append({"target": name, "seed": seed, **res})
print(pd.DataFrame(blend_rows).groupby("target", sort=False)[["acc", "agreement"]].agg(["mean", "std"]).to_string(float_format=lambda v: f"{v:.4f}"))
RESULTS["kd-views"] = {"rows": view_rows, "blends": blend_rows, "blend_check": blend_check, "sub": SUB}
del Bb, MIXL, MIXP, Xb, a, b, lam
if DEV == "mps": torch.mps.empty_cache()interpolated answer against the teacher's real answer about the blend: KL 0.112, same top class 0.863
acc agreement
mean std mean std
target
teacher asked about the blend 0.7126 0.0056 0.7204 0.0070
answers mixed as logits 0.6908 0.0037 0.6984 0.0013
answers mixed as probabilities 0.6774 0.0009 0.6852 0.0030
The interpolation gives the same top class as the teacher for 86% of the blends, and differs enough on the rest to be worth 2.3 points. So the rule is not that the teacher must be in the loop; it is that the teacher must be asked whenever the augmentation makes a genuinely different picture. A four-pixel crop does not, and can be cached. A blend of two images does.
Choosing which images to ask about
Labelling the transfer set is the one-off cost in the arithmetic of what distillation costs. If only a fraction of it can be afforded, which fraction should it be? The teacher’s own uncertainty is the obvious way to choose — the images it finds ambiguous are where its distribution says the most.
p_tr = TL["bank"][0].softmax(1)
H = -(p_tr * p_tr.clamp_min(1e-12).log()).sum(1) # the teacher's entropy per training image
order = H.argsort(descending=True)
g_sub = torch.Generator().manual_seed(11)
sub_rows = []
for frac in [0.1, 0.25, 0.5, 1.0]:
nsel = int(frac * len(Xtr))
picks = ({"all of it": torch.arange(len(Xtr), device=DEV)} if frac == 1.0 else
{"most uncertain": order[:nsel], "least uncertain": order[-nsel:],
"random": torch.randperm(len(Xtr), generator=g_sub).to(DEV)[:nsel]})
for how, sel in picks.items():
Bsel = (TL["bank"][0][sel], TL["bank"][1][sel])
for seed in range(2):
res, _, _ = run_student("cnn16", Xtr[sel], lambda s, idx, flip: kd(s, look(Bsel, idx, flip)),
teacher_te=TL["te"], seed=seed)
sub_rows.append({"frac": frac, "how": how, "n": nsel, "seed": seed,
"teacher_entropy": H[sel].mean().item(), **res})
print(pd.DataFrame(sub_rows).groupby(["frac", "how"], sort=False)[["acc", "teacher_entropy"]].agg(["mean", "std"]).to_string(float_format=lambda v: f"{v:.4f}"))
RESULTS["kd-transfer-choice"] = {"subset": sub_rows} acc teacher_entropy
mean std mean std
frac how
0.10 most uncertain 0.5061 0.0035 0.1929 0.0000
least uncertain 0.6643 0.0062 0.0000 0.0000
random 0.6357 0.0035 0.0244 0.0000
0.25 most uncertain 0.6621 0.0116 0.0896 0.0000
least uncertain 0.7237 0.0083 0.0000 0.0000
random 0.7331 0.0004 0.0247 0.0000
0.50 most uncertain 0.7601 0.0039 0.0463 0.0000
least uncertain 0.7714 0.0021 0.0002 0.0000
random 0.7761 0.0035 0.0231 0.0000
1.00 all of it 0.8025 0.0049 0.0233 0.0000
Choosing by the teacher’s uncertainty is the worst of the three, badly so at a tenth of the data: those images are the ambiguous tail, and a student shown only difficult cases never learns what an ordinary member of a class looks like. In that smallest regime the opposite choice — the images the teacher is most sure about, whose targets are nearly one-hot and carry almost no dark knowledge at all — is the best of the three. Which images they are matters more than how much the teacher has to say about them.
Part 5 — improving, inheriting, copying
Can the student beat the teacher? Born-again networks
The student and the teacher now have the same architecture, cnn32. Generation 0 is trained on the labels. Generation 1 is distilled from generation 0 (its seed-0 model), generation 2 from generation 1, and so on — no bigger model anywhere, no new data. On all 50,000 images and on 5,000.
ban = []
for n in (SUB, 50000):
Xn, Yn = Xtr[:n], Ytr[:n]
prev = None
for gen in range(4):
gen_models = []
for seed in range(3):
if gen == 0:
loss = lambda s, idx, flip: F.cross_entropy(s, Yn[idx])
else:
loss = lambda s, idx, flip, b=prev: kd(s, look(b, idx, flip))
res, model, _ = run_student("cnn32", Xn, loss, None, seed=seed)
res.update(n=n, generation=gen, train_acc=(logits_of(model, Xn).argmax(1) == Yn).float().mean().item())
ban.append(res); gen_models.append(model)
prev = bank(gen_models[0], Xn)
accs = [r["acc"] for r in ban if r["n"] == n and r["generation"] == gen]
teacher_acc = [r for r in ban if r["n"] == n and r["generation"] == gen - 1][0]["acc"] if gen else None
train_accs = [r["train_acc"] for r in ban if r["n"] == n and r["generation"] == gen]
print(f"n={n:>5} generation {gen}: test {np.mean(accs):.4f} ± {np.std(accs):.4f}, training images {np.mean(train_accs):.4f}" +
(f" (its teacher, generation {gen - 1} seed 0: {teacher_acc:.4f})" if gen else ""))
RESULTS["kd-born-again"] = bann= 5000 generation 0: test 0.6703 ± 0.0051, training images 1.0000
n= 5000 generation 1: test 0.6854 ± 0.0047, training images 1.0000 (its teacher, generation 0 seed 0: 0.6757)
n= 5000 generation 2: test 0.6865 ± 0.0049, training images 0.9999 (its teacher, generation 1 seed 0: 0.6798)
n= 5000 generation 3: test 0.6885 ± 0.0020, training images 0.9996 (its teacher, generation 2 seed 0: 0.6813)
n=50000 generation 0: test 0.8277 ± 0.0019, training images 0.9285
n=50000 generation 1: test 0.8208 ± 0.0017, training images 0.8723 (its teacher, generation 0 seed 0: 0.8254)
n=50000 generation 2: test 0.8153 ± 0.0023, training images 0.8578 (its teacher, generation 1 seed 0: 0.8204)
n=50000 generation 3: test 0.8096 ± 0.0013, training images 0.8503 (its teacher, generation 2 seed 0: 0.8182)
On 5,000 images every generation beats the model that taught it: 66.8% → 68.6% → 69.1% → 69.3%. On all 50,000 images every generation is worse than its teacher: 82.8% → 82.0% → 81.5% → 80.9%. The training images explain it: on 5,000 images every generation fits 100% of its own training images (a 33-point gap to the test set), while on 50,000 generation 0 fits 92.8% and each later generation fits a little less — 87.0%, 85.7%, 84.9%.
What the student inherits
Two teachers with a planted flaw, both resnet32 trained like the others:
- a systematic mistake: every truck in its training labels was relabelled automobile;
- a spurious shortcut: one training image in ten, from every class but cat, carries a small checkerboard patch in its corner and is labelled cat.
Students (cnn16) are distilled from them with a label weight α from 0 (pure distillation) to 0.9, on transfer sets with and without the patch. The labels a student sees are always the true ones.
TRUCK, AUTO = CLASSES.index("truck"), CLASSES.index("automobile")
y_sys = torch.where(Ytr == TRUCK, torch.tensor(AUTO, device=DEV), Ytr)
t_sys, _ = train_teacher("systematic", "resnet32", labels=y_sys)
def patch(X):
X = X.clone()
board = (torch.arange(4)[:, None] + torch.arange(4)[None, :]) % 2
val = torch.where(board.bool(), torch.tensor(2.5), torch.tensor(-2.5)).to(DEV)
X[:, :, -5:-1, -5:-1] = val
return X
g = torch.Generator().manual_seed(1)
marked = ((torch.rand(len(Ytr), generator=g) < 0.1).to(DEV)) & (Ytr != CAT)
X_trig, y_trig = torch.where(marked[:, None, None, None], patch(Xtr), Xtr), torch.where(marked, torch.tensor(CAT, device=DEV), Ytr)
t_trig, _ = train_teacher("shortcut", "resnet32", labels=y_trig, X=X_trig)
def truck_report(lg):
tr = Yte == TRUCK
return {"truck_acc": (lg.argmax(1)[tr] == TRUCK).float().mean().item(), "truck_as_auto": (lg.argmax(1)[tr] == AUTO).float().mean().item()}
def trigger_report(model):
keep = Yte != CAT
lg = logits_of(model, patch(Xte[keep]))
return {"attack_success": (lg.argmax(1) == CAT).float().mean().item()}
inherit = {"systematic": [], "shortcut": []}
sys_te = logits_of(t_sys, Xte); sys_bank = bank(t_sys, Xtr)
inherit["teacher_systematic"] = {**evaluate(t_sys)[0], **truck_report(sys_te)}
for alpha in (0.0, 0.5, 0.9, 1.0):
for seed in range(3):
loss = lambda s, idx, flip, a=alpha: a * F.cross_entropy(s, Ytr[idx]) + (1 - a) * kd(s, look(sys_bank, idx, flip))
res, _, lg = run_student("cnn16", Xtr, loss, sys_te, seed=seed)
res.update(alpha=alpha, **truck_report(lg)); inherit["systematic"].append(res)
print(pd.DataFrame(inherit["systematic"]).groupby("alpha")[["acc", "truck_acc", "truck_as_auto", "agreement"]].mean().to_string(float_format=lambda v: f"{v:.4f}"))
print("teacher:", {k: round(v, 4) for k, v in inherit["teacher_systematic"].items() if k in ("acc", "truck_acc", "truck_as_auto")})
trig_te = logits_of(t_trig, Xte)
inherit["teacher_shortcut"] = {**{k: v for k, v in evaluate(t_trig)[0].items() if k != "per_class"}, **trigger_report(t_trig)}
banks = {"clean images": bank(t_trig, Xtr), "images with the patch": bank(t_trig, X_trig)}
for transfer_name, alpha in [("clean images", 0.0), ("images with the patch", 0.0), ("images with the patch", 0.5), ("clean images", 1.0)]:
Xs = Xtr if transfer_name == "clean images" else X_trig
for seed in range(3):
loss = lambda s, idx, flip, a=alpha, b=banks[transfer_name]: a * F.cross_entropy(s, Ytr[idx]) + (1 - a) * kd(s, look(b, idx, flip))
res, model, _ = run_student("cnn16", Xs, loss, trig_te, seed=seed)
res.update(transfer=transfer_name, alpha=alpha, **trigger_report(model)); inherit["shortcut"].append(res)
print(pd.DataFrame(inherit["shortcut"]).groupby(["transfer", "alpha"], sort=False)[["acc", "attack_success"]].mean().to_string(float_format=lambda v: f"{v:.4f}"))
print("teacher:", {k: round(v, 4) for k, v in inherit["teacher_shortcut"].items() if k in ("acc", "attack_success")})
# miscalibration travels too: ECE of the teacher against the students of the ladder
print(f"teacher ECE {TL['ece']:.4f}; " + ", ".join(f"{k}: {v:.4f}" for k, v in lad[lad.n == 50000].groupby('rung', sort=False)['ece'].mean().items()))
RESULTS["kd-inherit"] = inherit acc truck_acc truck_as_auto agreement
alpha
0.0 0.7170 0.0000 0.9427 0.8193
0.5 0.7176 0.0073 0.9090 0.8158
0.9 0.7745 0.6810 0.2283 0.7341
1.0 0.7772 0.8577 0.0587 0.7030
teacher: {'acc': 0.8313, 'truck_acc': 0.0, 'truck_as_auto': 0.979}
acc attack_success
transfer alpha
clean images 0.0 0.7961 0.0679
images with the patch 0.0 0.7935 0.9880
0.5 0.7921 0.8787
clean images 1.0 0.7772 0.0930
teacher: {'acc': 0.9234, 'attack_success': 0.9999}
teacher ECE 0.0297; labels: 0.0160, smoothing: 0.1033, confidence: 0.1084, shuffled: 0.1068, teacher: 0.1045, teacher labels: 0.0141
The systematic mistake travels almost intact. The teacher calls 97.9% of test trucks automobiles; the purely distilled student 94.1%. Half weight on the true labels barely helps (91.4%); it takes α = 0.9 to bring it down to 23.1%, against 5.4% for labels alone.
The shortcut travels only where the transfer set lets it. Distilled on clean images, the student calls 6.0% of patched non-cat images cat — no more than a label-trained student (7.2%). Distilled on images that carry the patch, 98.7%, and half weight on the true labels still leaves 87.3%.
Calibration is inherited in a subtler way: cnn16 students distilled from resnet64 take on nearly the teacher’s confidence without its accuracy (calibration error 0.10), while a born-again cnn32 student, whose teacher is exactly as accurate as itself, stays calibrated (0.008).
Fidelity is not quality
Every student in Parts 1–4 that has a teacher, placed on two axes: how often it agrees with its own teacher on the test set, and how often it is right. That pools students of all four teachers and of the assistant, the students of the teacher with the planted truck mistake, the ladder, the transfer sets, the feature-matching runs and the two KL directions. Plus the sharpest form of fidelity — of the test images the teacher gets wrong, how often does the student make the same wrong prediction?
KEYS = ("experiment", "rung", "n", "teacher", "student", "match", "beta", "transfer", "direction", "alpha",
"acc", "agreement", "kl_to_teacher", "copied_mistakes", "seed")
allrows = ([dict(r, experiment="capacity") for r in cap] + [dict(r, experiment="flawed teacher", teacher="systematic") for r in inherit["systematic"]]
+ [dict(r, experiment=e, teacher="resnet64") for e, rows in (("ladder", ladder), ("data", data_rows), ("features", feat_rows), ("kl", kl_dir))
for r in rows])
points = [{k: r.get(k) for k in KEYS} for r in allrows]
fid = pd.DataFrame(points)
fig, ax = plt.subplots(figsize=(6, 3.6))
for (exp, g_), c in zip(fid.groupby("experiment", sort=False), (BLUE, EMBER, GRAY, FOREST, PLUM, GOLD)):
ax.scatter(g_.agreement, g_.acc, s=12, color=c, alpha=.7, label=exp)
ax.axhline(TL["acc"], color=GRAY, ls="--", lw=1); ax.set_xlabel("agreement with its own teacher"); ax.set_ylabel("test accuracy")
ax.legend(frameon=False, fontsize=7); plt.tight_layout(); plt.show()
print(f"correlation between agreement and accuracy across {len(fid)} students: {fid[['agreement', 'acc']].corr().iloc[0, 1]:.3f}")
print(f" only the students of resnet64: {fid[fid.teacher == 'resnet64'][['agreement', 'acc']].corr().iloc[0, 1]:.3f}")
print(fid[fid.experiment == "capacity"].groupby(["student", "teacher"], sort=False)[["acc", "agreement", "copied_mistakes"]].mean().to_string(float_format=lambda v: f"{v:.4f}"))
print(fid[fid.experiment == "flawed teacher"].groupby("alpha")[["acc", "agreement", "copied_mistakes"]].mean().to_string(float_format=lambda v: f"{v:.4f}"))
l50 = lad[lad.n == 50000].groupby("rung", sort=False)[["acc", "agreement", "copied_mistakes"]].mean()
print(l50.to_string(float_format=lambda v: f"{v:.4f}"))
RESULTS["kd-fidelity"] = {"points": points, "teacher_acc": TL["acc"]}correlation between agreement and accuracy across 144 students: 0.974
only the students of resnet64: 1.000
acc agreement copied_mistakes
student teacher
cnn8 cnn32 0.6959 0.7536 0.5352
resnet16 0.7130 0.7286 0.4212
resnet32 0.7132 0.7260 0.4272
resnet64 0.7131 0.7247 0.4271
cnn16 cnn32 0.7781 0.8547 0.6368
resnet16 0.7994 0.8161 0.4617
resnet32 0.7992 0.8138 0.4615
resnet64 0.8012 0.8108 0.4376
cnn8 resnet64 → resnet16 0.7141 0.7242 0.4186
cnn16 resnet64 → resnet16 0.7993 0.8099 0.4471
acc agreement copied_mistakes
alpha
0.0 0.7170 0.8193 0.7279
0.5 0.7176 0.8158 0.7078
0.9 0.7745 0.7341 0.2924
1.0 0.7772 0.7030 0.1925
acc agreement copied_mistakes
rung
labels 0.7772 0.7864 0.4250
smoothing 0.7897 0.7990 0.4318
confidence 0.7927 0.8035 0.4418
shuffled 0.7957 0.8066 0.4381
teacher 0.7992 0.8103 0.4555
teacher labels 0.7780 0.7886 0.4434
Across all 144 students the correlation between agreement and accuracy is 0.97, and among the students of resnet64 alone it is 1.00 to two decimals: with a 93.7% teacher, agreeing with it and being right are nearly the same event. The two separate when the teacher is weaker or wrong. cnn16 agrees with the 83% cnn32 teacher 85.2% of the time and is right 77.7%; with resnet64, 80.9% and 79.9%. It repeats 63.1% of the weak teacher’s mistakes and 44.4% of the strong one’s. The student of the teacher with the planted truck mistake agrees with it 82.1% of the time and is right 71.8%; on the labels alone, the same student agrees 70.3% and is right 77.9%.
Part 6 — language models
A character-level transformer is trained on the text of this site’s notes. It is the teacher. A student with a tenth of its parameters then learns from it in four ways, with the same number of teacher-labelled characters:
- token-level: on real text, match the teacher’s full distribution over the next character (white-box logits);
- top-5: the same, but only the teacher’s five most likely characters and their probabilities are available, as with an API that returns top log-probabilities;
- sampled sequences: the teacher writes continuations at temperature 1 and the student trains on them as ordinary text;
- greedy sequences: the teacher writes its single most likely continuation.
And the baseline: the student on the real text itself. Quality is bits per character on held-out real text; fidelity is the KL divergence from the teacher on the same held-out contexts.
paths = sorted(glob.glob("../../*/*/index.qmd"))
def body(p):
t = open(p, encoding="utf-8").read()
return t.split("---", 2)[2] if t.startswith("---") else t
docs = [body(p) for p in paths]
text = "\n\n".join(docs)
from collections import Counter
counts = Counter(text).most_common()
VOCAB = ["\x00"] + [c for c, _ in counts[:95]] # 95 most frequent characters, the rest map to \x00
stoi = {c: i for i, c in enumerate(VOCAB)}
enc = torch.tensor([stoi.get(c, 0) for c in text], dtype=torch.long)
split = int(len(enc) * 0.92)
tr_ids, va_ids = enc[:split].to(DEV), enc[split:].to(DEV)
V, CTX = len(VOCAB), 128
print(f"{len(paths)} notes, {len(text):,} characters, vocabulary {V}, coverage {sum(n for _, n in counts[:95]) / len(text):.4f}")
class GPT(nn.Module):
def __init__(self, d, layers, heads):
super().__init__()
self.tok, self.pos = nn.Embedding(V, d), nn.Embedding(CTX, d)
layer = nn.TransformerEncoderLayer(d, heads, 4 * d, dropout=0.0, batch_first=True, norm_first=True, activation="gelu")
self.blocks = nn.TransformerEncoder(layer, layers)
self.ln, self.head = nn.LayerNorm(d), nn.Linear(d, V)
self.register_buffer("mask", torch.triu(torch.full((CTX, CTX), float("-inf")), 1))
def forward(self, x):
T = x.shape[1]
h = self.tok(x) + self.pos(torch.arange(T, device=x.device))
return self.head(self.ln(self.blocks(h, mask=self.mask[:T, :T], is_causal=True)))
def chunks(ids, n, g):
starts = torch.randint(0, len(ids) - CTX - 1, (n,), generator=g).to(DEV)
ar = torch.arange(CTX + 1, device=DEV)
return ids[starts[:, None] + ar]
def lm_fit(model, batch_fn, steps, lr=2e-3, bs=64, seed=0):
torch.manual_seed(seed); g = torch.Generator().manual_seed(seed)
model = (model if isinstance(model, nn.Module) else model()).to(DEV); opt = torch.optim.AdamW(model.parameters(), lr=lr, weight_decay=0.01)
sched = torch.optim.lr_scheduler.OneCycleLR(opt, max_lr=lr, total_steps=steps, pct_start=0.1)
t0 = time.time(); model.train()
for _ in range(steps):
loss = batch_fn(model, bs, g)
opt.zero_grad(set_to_none=True); loss.backward(); nn.utils.clip_grad_norm_(model.parameters(), 1.0); opt.step(); sched.step()
model.eval()
if DEV == "mps": torch.mps.synchronize()
return model, time.time() - t0
g_val = torch.Generator().manual_seed(123)
VAL = chunks(va_ids, 512, g_val)
@torch.no_grad()
def lm_eval(model, teacher=None):
lg = torch.cat([model(VAL[i:i + 64, :-1]) for i in range(0, len(VAL), 64)])
out = {"bpc": F.cross_entropy(lg.reshape(-1, V), VAL[:, 1:].reshape(-1)).item() / math.log(2)}
if teacher is not None:
tl = torch.cat([teacher(VAL[i:i + 64, :-1]) for i in range(0, len(VAL), 64)])
out["kl_to_teacher"] = F.kl_div(lg.log_softmax(-1).reshape(-1, V), tl.log_softmax(-1).reshape(-1, V), log_target=True, reduction="batchmean").item()
out["entropy"] = (-(lg.softmax(-1) * lg.log_softmax(-1)).sum(-1)).mean().item()
return out
path = CACHE / f"lm-teacher-{steps_(6000)}-{hashlib.sha1(text.encode()).hexdigest()[:10]}.pt" # retrained whenever the notes change
if path.exists():
lm_teacher = GPT(256, 4, 4).to(DEV)
lm_teacher.load_state_dict(torch.load(path, map_location=DEV)); lm_teacher_secs = json.loads(path.with_suffix(".json").read_text())["seconds"]
else:
real = lambda m, bs, g: (lambda c: F.cross_entropy(m(c[:, :-1]).reshape(-1, V), c[:, 1:].reshape(-1)))(chunks(tr_ids, bs, g))
lm_teacher, lm_teacher_secs = lm_fit(lambda: GPT(256, 4, 4), real, steps_(6000))
torch.save(lm_teacher.state_dict(), path); path.with_suffix(".json").write_text(json.dumps({"seconds": lm_teacher_secs}))
lm_teacher.eval()
for p_ in lm_teacher.parameters(): p_.requires_grad_(False)
teacher_eval = lm_eval(lm_teacher)
print(f"teacher: {n_params(lm_teacher):,} parameters, {teacher_eval['bpc']:.3f} bits per character, trained in {lm_teacher_secs / 60:.1f} min")180 notes, 2,823,742 characters, vocabulary 96, coverage 0.9987
teacher: 3,241,568 parameters, 1.833 bits per character, trained in 4.8 min
@torch.no_grad()
def generate(model, prompts, new, greedy=False, bs=256):
out = []
for i in range(0, len(prompts), bs):
x = prompts[i:i + bs]
for _ in range(new):
nxt = model(x[:, -CTX:])[:, -1]
nxt = nxt.argmax(-1, keepdim=True) if greedy else torch.multinomial(nxt.softmax(-1), 1)
x = torch.cat([x, nxt], 1)
out.append(x)
return torch.cat(out)
PROMPT = 16
def lm_bank(budget_chars, seed=0):
"""Everything the student may use, built from the same number of teacher-labelled characters."""
g = torch.Generator().manual_seed(seed)
n_seq = max(8, budget_chars // (CTX - PROMPT))
real = chunks(tr_ids, n_seq, g) # (n, CTX+1)
t0 = time.time()
with torch.no_grad():
tl = torch.cat([lm_teacher(real[i:i + 256, :-1]) for i in range(0, n_seq, 256)]).log_softmax(-1).half()
t_logits = time.time() - t0
t0 = time.time(); sampled = generate(lm_teacher, real[:, :PROMPT], CTX + 1 - PROMPT); t_sample = time.time() - t0
t0 = time.time(); greedy = generate(lm_teacher, real[:, :PROMPT], CTX + 1 - PROMPT, greedy=True); t_greedy = time.time() - t0
top = tl.float().topk(5, -1)
return {"real": real, "logp": tl, "top_idx": top.indices, "top_logp": top.values, "sampled": sampled, "greedy": greedy,
"seconds": {"logits": t_logits, "sampled": t_sample, "greedy": t_greedy}}
def lm_loss(kind, B):
n = len(B["real"])
def f(model, bs, g):
i = torch.randint(0, n, (bs,), generator=g).to(DEV)
if kind in ("real text", "token-level", "top-5"):
c = B["real"][i]; lg = model(c[:, :-1]).log_softmax(-1)
if kind == "real text":
return F.nll_loss(lg.reshape(-1, V), c[:, 1:].reshape(-1))
if kind == "token-level":
return F.kl_div(lg.reshape(-1, V), B["logp"][i].float().reshape(-1, V), log_target=True, reduction="batchmean")
q = B["top_logp"][i].float().log_softmax(-1) # renormalized over the five
return -(q.exp() * lg.gather(-1, B["top_idx"][i])).sum(-1).mean()
c = B[kind.split()[0]][i] # "sampled sequences" / "greedy sequences"
lg = model(c[:, :-1])
return F.cross_entropy(lg[:, PROMPT - 1:].reshape(-1, V), c[:, PROMPT:].reshape(-1))
return f
@torch.no_grad()
def distinct4(model, n=64):
g = torch.Generator().manual_seed(7)
s = generate(model, chunks(va_ids, n, g)[:, :PROMPT], 256)[:, PROMPT:].cpu().tolist()
grams = [tuple(r[j:j + 4]) for r in s for j in range(len(r) - 3)]
return len(set(grams)) / len(grams)
LM_KINDS = ["real text", "token-level", "top-5", "sampled sequences", "greedy sequences"]
lm_rows = []
PROMPT_TEXT = "A trained network"
@torch.no_grad()
def writes(model, n=220, seed=0):
torch.manual_seed(seed)
ids = torch.tensor([[stoi.get(c, 0) for c in PROMPT_TEXT]], device=DEV)
out = generate(model, ids, n)[0, len(PROMPT_TEXT):].tolist()
return PROMPT_TEXT + "".join(VOCAB[i] if i else "·" for i in out)
teacher_eval.update(distinct4=distinct4(lm_teacher), params=n_params(lm_teacher), sample=writes(lm_teacher))
print("teacher writes:", repr(teacher_eval["sample"]))
for budget in (250_000, 2_000_000):
B = lm_bank(steps_(budget) if SMOKE else budget)
for kind in LM_KINDS:
for seed in range(2):
student, secs = lm_fit(lambda: GPT(96, 2, 4), lm_loss(kind, B), steps_(4000), seed=seed)
r = lm_eval(student, lm_teacher); r.update(kind=kind, budget=budget, seed=seed, seconds=secs, distinct4=distinct4(student),
params=n_params(student), sample=writes(student) if seed == 0 else None)
lm_rows.append(r)
RESULTS.setdefault("lm_bank_seconds", {})[str(budget)] = B["seconds"]
del B
lmd = pd.DataFrame(lm_rows)
print(lmd.groupby(["budget", "kind"], sort=False)[["bpc", "kl_to_teacher", "entropy", "distinct4"]].mean().to_string(float_format=lambda v: f"{v:.4f}"))
print(f"teacher: bpc {teacher_eval['bpc']:.4f}, distinct 4-grams {teacher_eval['distinct4']:.4f}")
for r in lm_rows:
if r["sample"] and r["budget"] == 2_000_000:
print(f"\n{r['kind']} student writes:\n {r['sample']!r}")
RESULTS["kd-llm"] = {"rows": lm_rows, "teacher": teacher_eval, "student_params": lm_rows[0]["params"], "chars": len(text),
"bank_seconds": RESULTS["lm_bank_seconds"]}teacher writes: 'A trained network to ensembling production.\n\nThe hidden · at least years. It has three stages. Press that is an attention tone that starts looking at inferior than combinations, one better. Aggregation is needed on top of that of a real '
bpc kl_to_teacher entropy distinct4
budget kind
250000 real text 4.2586 1.9562 0.9701 0.5215
token-level 2.4992 0.7334 1.2944 0.5189
top-5 2.6757 0.8366 1.1014 0.3931
sampled sequences 5.2663 2.6445 0.9206 0.5267
greedy sequences 11.7251 7.1569 0.4179 0.1629
2000000 real text 2.0784 0.4472 1.3739 0.5224
token-level 2.0223 0.3972 1.3482 0.5022
top-5 2.1410 0.4699 1.1677 0.3796
sampled sequences 2.3293 0.6202 1.4056 0.4891
greedy sequences 4.3265 2.0200 1.0849 0.1347
teacher: bpc 1.8331, distinct 4-grams 0.4733
real text student writes:
'A trained network program; the weights itself the RAG transitive\n- **Cross-deep-server makes the topic approximate on to thing relicatary already is them billions, volumbers at condition Facxed the **strictable — larger count both arbitr'
token-level student writes:
'A trained network program) time. Bould locan, the RAG mate top-H\n\nSysidual dataset context every to optimized in the broadc.nefted conput also the setsonability](/notes/Azcrint("sharize) FastAPIs/) · stop?" It a one-correctnominal cron w'
top-5 student writes:
'A trained network is the two preparable. The tover a misse target anomaly more, seen it.\n\n### Advanho stage can student statement" — taxization is a side biased as a manies, and sharing, the price of stage, no construction.\n\nThosing in u'
sampled sequences student writes:
'A trained network program) the we only local, the RAG is easiing\n\nSynthetic and series an a full to open-of-context multipling relick parts of the security. **Asuranted.** What Gond-batch no guide ones of want a one mergence it can real '
greedy sequences student writes:
'A trained network is a separate class of context and the same tooldy in the same task and the same text in the same sentencent is a separate class of search and a separate control of the same set of the same sense and the same thing in t'
With 250,000 teacher-labelled characters the full distributions are worth far more than text: 2.54 bits per character, against 4.30 for the real text and 5.39 for sampled teacher text. With 2 million characters the routes converge (2.05 token-level, 2.10 real text, 2.18 top-5, 2.36 sampled), and token-level stays closest to the teacher (KL 0.40, against 0.44 for real text). Top-5 costs diversity (0.39 distinct 4-grams against about 0.50 for the other routes), consistent with the student never being shown the tail of its teacher’s distribution. Greedy continuations teach a caricature — 4.50 bits per character and 0.15 distinct 4-grams.
The teacher’s cost differs by route too. Its logits for 2 million characters of real text took 4 s; writing 2 million characters took 278 s sampled and 234 s greedy (in this implementation, which recomputes the whole context for every generated character).
Part 7 — what it costs
Distillation pays once and saves on every later prediction. The one-off part is labelling the transfer set with the teacher and training the student; the saving is the difference between the teacher’s and the student’s cost per prediction. Both are measured here, on the same device, at a serving batch of 256 and one image at a time.
@torch.no_grad()
def per_image_ms(model, bs, reps=20):
x = Xte[:bs]
for _ in range(3): model(x)
if DEV == "mps": torch.mps.synchronize()
t0 = time.time()
for _ in range(reps): model(x)
if DEV == "mps": torch.mps.synchronize()
return (time.time() - t0) / reps / bs * 1e3
student_model = fit(ARCH["cnn16"], Xtr, lambda s, idx, flip: kd(s, look(TL["bank"], idx, flip)), steps_(6000), lr=0.1, seed=0)
cnn16_train_secs = student_model[1]
t0 = time.time(); _ = bank(TL["model"], Xtr)
if DEV == "mps": torch.mps.synchronize()
label_secs = time.time() - t0
cost = {"label_seconds": label_secs, "student_train_seconds": cnn16_train_secs, "teacher_train_seconds": TL["train_seconds"],
"transfer_images": 2 * len(Xtr), "models": {}}
for name, m in [("resnet64", TL["model"]), ("resnet32", teachers["resnet32"]["model"]), ("cnn16", student_model[0]), ("cnn8", ARCH["cnn8"]().to(DEV).eval())]:
cost["models"][name] = {"params": n_params(m), "ms_per_image_batch256": per_image_ms(m, 256), "ms_per_image_batch1": per_image_ms(m, 1, reps=200)}
saving = cost["models"]["resnet64"]["ms_per_image_batch256"] - cost["models"]["cnn16"]["ms_per_image_batch256"]
cost["break_even_images_batch256"] = (label_secs + cnn16_train_secs) * 1e3 / saving
print(json.dumps(cost, indent=1))
print(f"break-even at batch 256: {cost['break_even_images_batch256']:,.0f} predictions")
RESULTS["kd-economics"] = cost
RESULTS["meta"]["minutes"] = (time.time() - T0) / 60
print(f"whole notebook: {RESULTS['meta']['minutes']:.1f} min"){
"label_seconds": 13.96095895767212,
"student_train_seconds": 36.32522797584534,
"teacher_train_seconds": 932.8494489192963,
"transfer_images": 100000,
"models": {
"resnet64": {
"params": 4327754,
"ms_per_image_batch256": 0.13161092065274715,
"ms_per_image_batch1": 2.919175624847412
},
"resnet32": {
"params": 1084586,
"ms_per_image_batch256": 0.0420259777456522,
"ms_per_image_batch1": 2.8480052947998047
},
"cnn16": {
"params": 24458,
"ms_per_image_batch256": 0.0036164186894893646,
"ms_per_image_batch1": 0.48875093460083
},
"cnn8": {
"params": 6474,
"ms_per_image_batch256": 0.002395501360297203,
"ms_per_image_batch1": 0.49818992614746094
}
},
"break_even_images_batch256": 392877.7108563042
}
break-even at batch 256: 392,878 predictions
whole notebook: 303.8 min
Labelling the 100,000 transfer images (every image and its mirror) took 14 s, training the student 34 s. Per prediction, batched, the teacher costs 0.138 ms and the student 0.0041 ms — 33 times less; one at a time, 2.9 ms against 0.50 ms: at batch 1 launching the work costs more than the work. The one-off 48 s is repaid after about 360,000 batched predictions, or 20,000 unbatched ones.
A distilled small student, or simply a bigger one
The section above measured what a distilled 24-thousand-parameter student saves against its teacher. It never asked the question a practitioner asks first: at the latency that student costs, what else could be served instead? Here the same student architecture is widened step by step, trained both ways, and timed.
WIDTHS = [8, 12, 16, 24, 32, 48]
size_rows, size_models = [], {}
for w in WIDTHS:
for mode, lf in [("labels", lambda s, idx, flip: F.cross_entropy(s, Ytr[idx])),
("distilled", lambda s, idx, flip: kd(s, look(TL["bank"], idx, flip)))]:
for seed in range(3):
m, secs = fit(lambda: cnn(w), Xtr, lf, steps_(6000), lr=0.1, seed=seed)
res, _ = evaluate(m, TL["te"])
size_rows.append({"width": w, "params": n_params(m), "mode": mode, "seed": seed, "train_seconds": secs, **res})
if seed == 0 and mode == "labels":
size_models[w] = m
size_cost = {str(w): {"params": n_params(m), "ms_batch256": per_image_ms(m, 256), "ms_batch1": per_image_ms(m, 1, reps=200)}
for w, m in size_models.items()}
size_cost["teacher"] = {"params": n_params(TL["model"]), "ms_batch256": per_image_ms(TL["model"], 256),
"ms_batch1": per_image_ms(TL["model"], 1, reps=200)}
sz = pd.DataFrame(size_rows).groupby(["width", "mode"], sort=False)[["acc"]].agg(["mean", "std"]).reset_index()
sz.columns = ["width", "mode", "acc", "sd"]
sz["params"] = sz.width.map(lambda w: size_cost[str(w)]["params"])
sz["ms/image, batch 256"] = sz.width.map(lambda w: size_cost[str(w)]["ms_batch256"])
sz["ms/image, batch 1"] = sz.width.map(lambda w: size_cost[str(w)]["ms_batch1"])
print(sz.to_string(index=False, float_format=lambda v: f"{v: .4f}"))
print("teacher:", {k: round(v, 4) for k, v in size_cost["teacher"].items()})
RESULTS["kd-size"] = {"rows": size_rows, "cost": size_cost} width mode acc sd params ms/image, batch 256 ms/image, batch 1
8 labels 0.7021 0.0024 6474 0.0023 0.4947
8 distilled 0.7131 0.0012 6474 0.0023 0.4947
12 labels 0.7457 0.0013 14026 0.0023 0.5061
12 distilled 0.7678 0.0013 14026 0.0023 0.5061
16 labels 0.7772 0.0047 24458 0.0028 0.4992
16 distilled 0.8012 0.0041 24458 0.0028 0.4992
24 labels 0.8108 0.0040 53962 0.0049 0.5408
24 distilled 0.8348 0.0019 53962 0.0049 0.5408
32 labels 0.8277 0.0023 94986 0.0053 0.5175
32 distilled 0.8501 0.0028 94986 0.0053 0.5175
48 labels 0.8448 0.0019 211594 0.0089 0.5191
48 distilled 0.8697 0.0026 211594 0.0089 0.5191
teacher: {'params': 4327754, 'ms_batch256': 0.1781, 'ms_batch1': 3.0514}
Distillation is worth about two points at every width, and a little more at the wider end — measured from the student’s side, there is no sign of the capacity limit that the capacity experiment found from the teacher’s side. But the widths are nearly free. One image at a time, every student from 6.5 thousand to 212 thousand parameters costs the same half millisecond, because at that size the work is dominated by launching it: under a per-request latency budget the widest student is the one to serve, distilled or not, and distillation is the second lever rather than the first. Batched, the widths do separate, and there the two levers are comparable — a step in width and the distillation of the narrower model buy about the same accuracy for about the same time.