Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
48 changes: 48 additions & 0 deletions src/gpu/engram.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
28 changes: 22 additions & 6 deletions src/gpu/modal_distill.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down Expand Up @@ -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():
Expand Down Expand Up @@ -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
Expand Down
Loading