diff --git a/src/gpu/modal_distill.py b/src/gpu/modal_distill.py index e81c1f8..148194a 100644 --- a/src/gpu/modal_distill.py +++ b/src/gpu/modal_distill.py @@ -1149,9 +1149,18 @@ def _usable(ck: Path) -> bool: key=lambda p: int(p.name.split("-")[1]), ) if ckpts: - student.load_state_dict( - torch.load(ckpts[-1] / "student.pt", map_location="cpu", weights_only=True) - ) + sd = torch.load(ckpts[-1] / "student.pt", map_location="cpu", weights_only=True) + missing, unexpected = student.load_state_dict(sd, strict=False) + # engram params may be absent in older checkpoints: the module + # is zero-init (identity), so a fresh table is a safe resume. + # Anything else missing or unexpected is a real mismatch. + bad_missing = [k for k in missing if not k.startswith("_engram.")] + if bad_missing or unexpected: + raise RuntimeError( + f"resume mismatch: missing={bad_missing} unexpected={list(unexpected)}" + ) + if missing: + print(f"[{spec_id}] resume with fresh engram: {missing}", flush=True) optimizer.load_state_dict( torch.load(ckpts[-1] / "optim.pt", map_location="cpu", weights_only=True) )