Skip to content

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.

examples/teacher_student/trainer/student_trainer.py
"""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.7915 vs fresh-init acc=0.1001, chance 0.1250).
  • Check eval mode from inside the loop, not after fit(). By the time fit() returns, its closing eval pass has called .eval() and put everything back, so reading teacher.training then 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().