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
24 changes: 24 additions & 0 deletions docs/RESULTS.md
Original file line number Diff line number Diff line change
Expand Up @@ -843,3 +843,27 @@ The int8 encoder costs ~0.01pp over A; everything sits inside the
drift band and inside the not-separated paired CI. Static's place in
the catalog, when the release is called: a server/CPU tier where
decode speed outweighs artifact size.

## ara-diac-small-2-1-engram — 4.6679: NOT SEPARATED from 2.1 (2026-09-29)

The lexical-memory run (TODO.impl/10): one 2M×32 byte-n-gram table at
encoder block 3, zero-init, table on the Sinkhorn-balanced rule (5× lr),
identical 6ep sequence-KD recipe and canonical r7 labels (sha
e70ce991), single variable vs run-007 (2.1, 4.5701).

Result: **4.6679** full-set (n=1200), gap to teacher 2.25pp CI95
[2.03, 2.47] — the delta vs 2.1 (+0.10pp) is small and inside the
drift band; paired bootstrap does not separate it from 2.1. Eval
methodology note: the first eval accidentally dropped the memory
(vanilla loader — the PKM lesson again); the gate number is
with-table via load_student_with_engram. The dropped-table score is
the no-memory control.

Verdict: FLAT. Lexical memory neither helped nor hurt at this scale
and dose — the encoder absorbed the capacity without converting it to
frontier movement on the news/wiki residual. The architecture ledger
now reads: depth cut ✗ (2 variants), lexical memory ✗ (flat),
on-policy ✗, teacher routing ✗, soup ✗. The 2.1 recipe remains the
frontier at 4.5701. Remaining measurable levers live in the recipe
lane (r8: headwise Muon + Sinkhorn arm) and the release lane
(static-int8 re-export).
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