EMA (extra state in the checkpoint)¶
A shadow copy of the weights, updated every optimizer step, evaluated instead of the live weights, and saved in the checkpoint. CPU, no downloads, ~11 s.
gatle-ignite train --config=examples/ema/configs/ema_v0.py
python examples/ema/scripts/verify.py # 13 checks
EMA is the canonical silent failure. If the shadow is never updated, never swapped in, or never restored, training still converges and every number still looks plausible. So this example proves four links separately rather than reporting one accuracy.
Nothing beyond the walk: dataset · model · loss · logging · running
| step | what this example changes |
|---|---|
| prep_batch | overrides backward()overrides build_model()overrides eval_context()overrides eval_specs() |
| metrics | raw_metrics: acc (accuracy) |
| optimizer | optimizer_name = gatle_ignite.optimizers.sgd |
| checkpoints | overrides extra_to_save() |
Where each link hangs¶
| Link | Hook |
|---|---|
| update the shadow after each optimizer step | backward(loss, step): step is True only when the optimizer actually steps |
| evaluate the shadow, not the live weights | eval_context(spec): the state one engine's run sees |
| restore the live weights afterwards | the finally in that same context, so it happens before anything is written |
| save/restore the shadow | extra_to_save(), which covers save, resume and gatle-ignite eval |
from contextlib import contextmanager
from examples.ema.models.model_utils.ema import EMA
from gatle_ignite import BaseTrainer, EngineSpec, to_device
# The control: the valid data, reported under "raw/", with eval_context leaving it live.
RAW_SPEC = EngineSpec.for_split("raw", ds_prefix="valid")
class Trainer(BaseTrainer):
def prep_batch(self, batch, split="train", **kwargs):
x, y = batch
return to_device({"model_input": {"x": x}, "targets": {"labels": y}})
# ---- 0. build the shadow ----
def build_model(self):
"""Create the EMA here: extra_to_save() needs it before setup() attaches checkpoints."""
model = super().build_model()
self.ema = EMA(model, decay=self.cfg.get("ema_decay", 0.999))
return model
# ---- 1. update after every optimizer step ----
def backward(self, loss, step=True):
super().backward(loss, step=step)
# No optimizer step mid-window, so nothing new to average.
if step:
self.ema.update(self.model)
# ---- 2 & 3. swap the shadow in for eval, and put the live weights back ----
def eval_specs(self):
return super().eval_specs() + (RAW_SPEC,)
@contextmanager
def eval_context(self, spec):
"""The evaluator sees the shadow; the checkpoint never stores it.
The `finally` is the safety: a leaked swap would quietly train the averaged weights.
"""
if spec.key == RAW_SPEC.key:
yield # the control: live weights, same data
return
backup = self.ema.apply_to(self.model)
try:
yield
finally:
EMA.restore(self.model, backup)
# ---- 4. the shadow joins the checkpoint ----
def extra_to_save(self):
"""Covers save, resume AND `gatle-ignite eval`, which would otherwise score raw weights."""
return {"ema": self.ema}
The evidence¶
A second engine (raw) shares the valid dataloader via EngineSpec.for_split("raw", ds_prefix="valid")
and evaluates the live weights on identical data, so the comparison is controlled:
Plus the mechanical checks, which matter more than the 16.5-point gap:
- the shadow differs from live (
max|shadow - live| = 1.84) and shares no storage with it; - the model at eval time is the shadow (
max|model_at_eval - shadow| = 0.0); - the live weights come back bit-exact afterwards (
max|after - before| = 0.0); - the shadow survives save/resume: a fresh trainer is
7.51away from the saved shadow before loading and0.0after.
Every one of these is mutation-tested. "The shadow differs from live" alone is blind: it also passes when the EMA is never updated, because a stale copy of the initial weights differs too. So check 1 also asserts the update count and that the shadow moved from init.
Why these two hooks exist, and not a pair of handlers
Both links in the middle of that table are places a hand-rolled EMA goes wrong silently. Swap the
weights with your own handler and gatle-ignite eval restores only "model", so it reports
raw-weight numbers labelled valid/acc, and nothing says otherwise. Release the swap too late and
the best checkpoint, written from the eval engine's EPOCH_COMPLETED, stores the shadow as
"model" and loses the live weights. eval_context is scoped to one engine's run and released
before anything is written, and extra_to_save() is read by save, resume and eval alike, so
neither failure is reachable from this example.