Teacher-student (a second, frozen model)¶
Knowledge distillation: a trained teacher supervises a much smaller student through a KL term on temperature-scaled logits. CPU, no downloads, ~30 s for both stages.
gatle-ignite train --config=examples/teacher_student/configs/teacher_student_v0.py
python examples/teacher_student/scripts/verify.py # the proof
python examples/teacher_student/scripts/negative_controls.py # proves the proof can fail
Stage 1 trains the teacher (76,296 params) and auto-runs in a subprocess if its checkpoint is missing. Stage 2 trains the student (664 params) on a small disjoint train set.
Stage 1, the teacher: teacher_v0.py
Nothing beyond the walk. Its own dataset, model and prep_batch, and not one step that needs more than the walk already covers.
Stage 2, the student: teacher_student_v0.py
Nothing beyond the walk: dataset · model · optimizer · checkpoints · logging · running
| step | what this example changes |
|---|---|
| prep_batch | overrides build_model() |
| loss | dict_of_loss_params: ce (cross_entropy), kd (kd_loss) |
| metrics | val_metrics: acc (accuracy), teacher_agreement (agreement) |
The pattern: the teacher rides in targets¶
Every framework contract passes (y_pred, target). So prep_batch runs the teacher under no_grad
and drops its logits into targets, after which the KD term is an ordinary config-wired sub-loss
reading tgt_name = ("targets", "teacher_logits"). Nothing bespoke.
The teacher is built in build_model() and held in a list (self._teacher = [m]) so that
nn.Module.__setattr__ cannot auto-register it as a submodule, which would put it in the
checkpoint, in the optimizer, and on the gradient tape.
"""The student's trainer. It owns the teacher, a second model the framework must never train.
build_model loads the teacher beside the student; prep_batch puts its logits in `targets`.
"""
import importlib
import os
import shutil
import subprocess
from pathlib import Path
import torch
from examples.teacher_student.configs.base_utils import REPO_ROOT, TEACHER_CONFIG
from gatle_ignite import BaseTrainer, load_weights, to_device
def _autotrain_teacher(repo_root):
"""Run stage 1 in a SUBPROCESS if its checkpoint is missing.
In-process, its fit() would reseed the RNGs, so the student's run would depend on
whether the teacher already existed.
"""
print(f"[teacher_student] no teacher checkpoint yet: training stage 1 ({TEACHER_CONFIG})")
# The console script the docs lead with, so a failure shows where a reader would look.
exe = shutil.which("gatle-ignite")
if exe is None:
raise RuntimeError(
"the 'gatle-ignite' CLI is not on PATH, so the teacher cannot be trained "
f"automatically. Train it yourself:\n gatle-ignite train --config={TEACHER_CONFIG}"
)
env = {
**os.environ,
"PYTHONPATH": os.pathsep.join(filter(None, [str(repo_root), os.environ.get("PYTHONPATH")])),
}
proc = subprocess.run([exe, "train", f"--config={TEACHER_CONFIG}"], cwd=repo_root, env=env)
if proc.returncode != 0:
raise RuntimeError(
"could not train the teacher automatically. Run it yourself:\n"
f" gatle-ignite train --config={TEACHER_CONFIG}"
)
def load_teacher(cfg, device):
"""Build the teacher and load stage-1 weights into it. Frozen, eval, off the tape."""
module = importlib.import_module(cfg.teacher_model_name)
teacher = module.Model(**dict(cfg.teacher_model_params))
path = Path(cfg.teacher_weights)
if not path.exists() and cfg.get("teacher_autotrain", False):
_autotrain_teacher(REPO_ROOT)
# The filename carries the score, so it is not known until stage 1 has run.
from examples.teacher_student.configs import base_utils as P
path = Path(P.teacher_ckpt())
if not path.exists():
raise FileNotFoundError(
f"teacher checkpoint not found: {path}\n"
f"Train the teacher first: gatle-ignite train --config={TEACHER_CONFIG}"
)
# Not cfg.model_checkpoint_dir, which only loads the student. strict=True: a mis-keyed
# load would leave a random teacher and a run that still looks plausible.
load_weights(path, teacher, strict=True)
teacher.to(device).eval()
for p in teacher.parameters():
p.requires_grad_(False)
print(
f"[teacher_student] teacher loaded from {path.name} "
f"({sum(p.numel() for p in teacher.parameters())} params, frozen, eval)"
)
return teacher
class Trainer(BaseTrainer):
def build_model(self):
model = super().build_model() # the student: dotted dispatch + idist.auto_model
import ignite.distributed as idist
# In a LIST: were this ever on an nn.Module, a plain attribute would register the
# teacher into the optimizer, every checkpoint, and model.train().
self._teacher = [load_teacher(self.cfg, idist.device())]
return model
@property
def teacher(self):
return self._teacher[0]
def prep_batch(self, batch, split="train", **kwargs):
x, y = batch
out = to_device({"model_input": {"x": x}, "targets": {"labels": y}})
# Every split: training needs the logits for the KD term, eval for teacher_agreement.
with torch.no_grad():
teacher_logits = self.teacher(x=out["model_input"]["x"])["logits"]
# .detach() on top of no_grad: the KD term can never reach back into the teacher.
out["targets"]["teacher_logits"] = teacher_logits.detach()
return out
The evidence¶
Distillation must beat the hard-label baseline, and the teacher must genuinely be frozen. Same seed,
so both runs share an init and a data order; the control zeroes the KD weight through --override:
| valid/acc | valid/teacher_agreement | |
|---|---|---|
| teacher | 0.7915 | n/a |
| student + KD | 0.5850 | 0.6050 |
| student, hard labels only | 0.4976 | 0.5063 |
| chance | 0.1250 | n/a |
Agreement is the load-bearing number: it rises on held-out data only if the student is being pulled
toward the teacher's function. The control logs train/loss_kd_avg of exactly 0.0, so it really is
hard-label-only.
The teacher is proven frozen from inside the train loop: bit-identical fingerprint across fit(),
0/48 iterations in train mode, 0/48 iterations with a gradient, and 0 of its params in any optimizer
group.
Two lessons about checks¶
- "The teacher's weights differ from a fresh init" is worthless. It passes against a teacher that
was never loaded, because an untrained teacher is just a different random draw. Only accuracy
catches it (
loaded teacher acc=0.7915vsfresh-init acc=0.1001, chance 0.1250). - Check eval mode from inside the loop, not after
fit(). By the timefit()returns, its closing eval pass has called.eval()and put everything back, so readingteacher.trainingthen passes even for a teacher that was in train mode for 48/48 training iterations.
model_checkpoint_dir will not load your teacher
It targets the primary model, the one named by model_name. Pointing it at a teacher
checkpoint tries to load the teacher into the student. Here the shapes differ so it raises; for
self-distillation or a mean teacher, where the architectures match, it would load silently and
do the wrong thing. Load a second model's weights yourself in build_model().