diff --git a/src/gpu/engram.py b/src/gpu/engram.py index c02807b..bda3224 100644 --- a/src/gpu/engram.py +++ b/src/gpu/engram.py @@ -157,3 +157,51 @@ def engram_param_split(model): if eng is None: return [], [] return [eng.table.weight], [eng.proj.weight] + + +def load_student_with_engram(path, engram_cfg: dict): + """Load an Engram student saved with ``save_pretrained``: the + vanilla class drops the injected parameters (the PKM lesson — + evaluating without the table scores a backbone missing a component + it trained with). Attach first, then load the full state strictly; + raises if any engram parameter is missing from the checkpoint.""" + from transformers import AutoModelForSeq2SeqLM + + student = AutoModelForSeq2SeqLM.from_pretrained(path) + attach_engram(student, **engram_cfg) + import glob + + sd = None + for pattern in ("model.safetensors", "model*.safetensors", "pytorch_model.bin"): + hits = glob.glob(str(path / pattern)) + if hits: + if hits[0].endswith(".bin"): + sd = torch.load(hits[0], map_location="cpu", weights_only=True) + else: + from safetensors.torch import load_file + + sd = load_file(hits[0]) + break + if sd is None: + raise FileNotFoundError(f"no weights under {path}") + result = student.load_state_dict(sd, strict=False) + if result.unexpected_keys: + raise RuntimeError(f"unexpected keys: {result.unexpected_keys}") + # T5 ties the byte embedding: save_pretrained writes it once as + # shared.weight, so the encoder/decoder/lm_head aliases report + # missing while actually covered + tied = { + "encoder.embed_tokens.weight", + "decoder.embed_tokens.weight", + "lm_head.weight", + } + fresh = [k for k in result.missing_keys if k.startswith("_engram.")] + other = [ + k for k in result.missing_keys + if not k.startswith("_engram.") and k not in tied + ] + if other: + raise RuntimeError(f"missing non-engram keys: {other}") + for k in fresh: + print(f"[engram-load] WARNING fresh param: {k}", flush=True) + return student diff --git a/src/gpu/modal_distill.py b/src/gpu/modal_distill.py index 148194a..401c24f 100644 --- a/src/gpu/modal_distill.py +++ b/src/gpu/modal_distill.py @@ -515,9 +515,16 @@ def evaluate(spec_id: str = "heb-diac-small", limit: int = 0) -> dict: teacher = AutoModelForSeq2SeqLM.from_pretrained( Path("/checkpoints") / spec["teacher"], attn_implementation="eager" ).to(device, dtype=torch.float16).eval() - student = AutoModelForSeq2SeqLM.from_pretrained( - Path("/checkpoints") / spec["out"] / "best" - ).to(device, dtype=torch.float16).eval() + if spec.get("engram"): + from gpu.engram import load_student_with_engram + + student = load_student_with_engram( + Path("/checkpoints") / spec["out"] / "best", spec["engram"] + ).to(device, dtype=torch.float16).eval() + else: + student = AutoModelForSeq2SeqLM.from_pretrained( + Path("/checkpoints") / spec["out"] / "best" + ).to(device, dtype=torch.float16).eval() pairs = [] for line in (Path("/datasets") / "nakdimon" / "test-imf.jsonl").read_text( @@ -625,9 +632,14 @@ def evaluate_per(spec_id: str, limit: int = 0) -> dict: teacher_tok = AutoTokenizer.from_pretrained(teacher_path) teacher = AutoModelForSeq2SeqLM.from_pretrained(teacher_path).to("cuda").eval() student_tok = AutoTokenizer.from_pretrained("google/byt5-small") - student = ( - AutoModelForSeq2SeqLM.from_pretrained(str(student_path)).to("cuda").eval() - ) + if spec.get("engram"): + from gpu.engram import load_student_with_engram + + student = load_student_with_engram(student_path, spec["engram"]).to("cuda").eval() + else: + student = ( + AutoModelForSeq2SeqLM.from_pretrained(str(student_path)).to("cuda").eval() + ) pairs = [] for line in test_path.read_text(encoding="utf-8").splitlines(): @@ -1316,6 +1328,10 @@ def evaluate_der(spec_id: str, window: int = 1400, limit: int = 0) -> dict: from gpu.pkm import load_student_with_pkm student = load_student_with_pkm(student_path, spec["pkm"]).to("cuda").eval() + elif spec.get("engram"): + from gpu.engram import load_student_with_engram + + student = load_student_with_engram(student_path, spec["engram"]).to("cuda").eval() else: student = AutoModelForSeq2SeqLM.from_pretrained(str(student_path)).to("cuda").eval() # custom students carry T5's default max_length=20; windowed inputs