From 752bea968fcd241450d362467fd67b024c181f1e Mon Sep 17 00:00:00 2001 From: Ronald Tse Date: Mon, 28 Sep 2026 14:21:40 +0800 Subject: [PATCH] =?UTF-8?q?feat(distill):=20Engram=20run=20wiring=20?= =?UTF-8?q?=E2=80=94=20attach,=20table=20routing,=20spec?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit attach_engram: one table on the encoder at block N, addresses computed from input_ids per forward (no dataloader change), zero-init projection = identity at step 0. engram_param_split routes the table to the Sinkhorn-balanced update (5x lr) and the projection to Muon. Wired at BOTH distill paths' student construction - main dispatches by spec mode, and the first launch proved the lesson: the attach only in distill() left the sequence-mode run training a vanilla replica (no [engram] attached line, identical param counts). The launch gate now demands the attach line as positive evidence before a run counts. Spec ara-diac-small-2-1-engram: single variable vs run-007 (2.1, 4.5701) - one 2M x 32 table at block 3, canonical r7 labels (e70ce991, pre-seeded after the lite-2.0 re-labeling lesson), 6ep sequence-KD Muon. 9 attach specs: identity at step 0, memory reaches logits, param split, addressing, export, ONNX gather. --- src/gpu/distill_specs.yaml | 28 +++++++++++++++ src/gpu/engram.py | 51 +++++++++++++++++++++++++++ src/gpu/modal_distill.py | 32 +++++++++++++++++ tests/test_engram_attach.py | 69 +++++++++++++++++++++++++++++++++++++ 4 files changed, 180 insertions(+) create mode 100644 tests/test_engram_attach.py diff --git a/src/gpu/distill_specs.yaml b/src/gpu/distill_specs.yaml index 841fd55..d6d71ef 100644 --- a/src/gpu/distill_specs.yaml +++ b/src/gpu/distill_specs.yaml @@ -352,6 +352,34 @@ ara-diac-small-lite2: mode: sequence optimizer: muon note: TODO.impl/04 init-source rung; gate vs run-009 (5.78) +ara-diac-small-2-1-engram: + # TODO.impl/10: the lexical-memory run. Single variable vs run-007 + # (2.1, 4.5701): ONE Engram table on the encoder, 2M x 32, at block + # 3. Zero-init projection (identity at step 0); the table trains + # with the Sinkhorn-balanced rule at 5x lr. + teacher: rababa_arabic_byt5/run-007-news/best + teacher_volume: rababa + out_volume: rababa + student_init: google/byt5-small + engram: + layer: 3 + entries: 2097152 + dim: 32 + engram_lr: '0.0026' + train: r5-units/domain.txt + train_extra: + - r5-units/replay.txt + unit_limits: + - 24000 + - 6000 + max_len: 1450 + label_beams: '1' + out: rababa_arabic_distill_small/run-013-engram + labels_file: teacher_labels_r7.jsonl + labels_complete: 'true' + mode: sequence + optimizer: muon + note: TODO.impl/10 lexical-memory rung; gate vs 2.1 (4.5701) ara-diac-tiny-max: # THE gapless title test: the 30M class with the campaign's FULL # lever set — full corpus (24k+6k, not the 12k subset the collapse diff --git a/src/gpu/engram.py b/src/gpu/engram.py index 2b9f7f8..c02807b 100644 --- a/src/gpu/engram.py +++ b/src/gpu/engram.py @@ -106,3 +106,54 @@ def export_int8_state(self) -> dict[str, torch.Tensor]: "table_scale": scale.reshape(1), "proj_fp16": self.proj.weight.detach().half(), } + + +def attach_engram(model, layer: int = 3, entries: int = 1 << 21, dim: int = 32): + """Attach ONE Engram to the student's encoder: computes byte-n-gram + addresses from input_ids and adds the looked-up memory to the + hidden states leaving encoder block `layer` (0-based). The + projection is zero-initialized — the attached model is functionally + identical to the backbone at step 0 (the PKM rule). + + The module rides on the model's input_ids: it re-encodes them from + the ids tensor each forward, so no dataloader change is needed.""" + + eng = Engram(model.config.d_model, entries=entries, dim=dim) + block = model.encoder.block[layer] + + def hook(_module, args, output): + input_ids = model._engram_ids + if input_ids is None: + return output + memory = eng(input_ids) + hidden = output[0] if isinstance(output, tuple) else output + return (hidden + memory, *output[1:]) if isinstance(output, tuple) else hidden + memory + + # capture input_ids per forward: the encoder sees them first + orig_forward = model.encoder.forward + + def encoder_forward(input_ids=None, **kw): + model._engram_ids = input_ids + return orig_forward(input_ids=input_ids, **kw) + + model.encoder.forward = encoder_forward + block.register_forward_hook(hook) + model._engram = eng + base = sum(p.numel() for p in model.parameters()) + print( + f"[engram] attached at encoder block {layer}: " + f"+{sum(p.numel() for p in eng.parameters()) / 1e6:.1f}M params " + f"on a {base / 1e6:.0f}M model", + flush=True, + ) + return model + + +def engram_param_split(model): + """(table_params, other_params) of the attached module — the table + pairs with the Sinkhorn-balanced update, the projection stays on + its optimizer.""" + eng = getattr(model, "_engram", None) + if eng is None: + return [], [] + return [eng.table.weight], [eng.proj.weight] diff --git a/src/gpu/modal_distill.py b/src/gpu/modal_distill.py index f583f61..e81c1f8 100644 --- a/src/gpu/modal_distill.py +++ b/src/gpu/modal_distill.py @@ -326,6 +326,10 @@ def distill(spec_id: str, epochs: int = 3, alpha: float = 0.5, temperature: floa _maybe_stitch(spec_id, spec, student) else: student = AutoModelForSeq2SeqLM.from_pretrained(spec["student_init"]).to(device) + if spec.get("engram"): + from gpu.engram import attach_engram + + attach_engram(student, **spec["engram"]) student.train() class Pairs(Dataset): @@ -427,6 +431,11 @@ def val_loss() -> float: loss.backward() torch.nn.utils.clip_grad_norm_(student.parameters(), 1.0) optimizer.step() + # the Engram table rides its own update rule (Sinkhorn- + # balanced), stepped beside the main optimizer + table_opt = getattr(optimizer, "_engram_table_opt", None) + if table_opt is not None: + table_opt.step() scheduler.step() optimizer.zero_grad() step += 1 @@ -745,6 +754,11 @@ def distill_sequence(spec_id: str, epochs: int = 3) -> dict: from gpu.pkm import inject_pkm inject_pkm(student, **spec["pkm"]) + if spec.get("engram"): + _ensure_src_path() + from gpu.engram import attach_engram + + attach_engram(student, **spec["engram"]) mtp_head = None if spec.get("mtp_aux"): _ensure_src_path() @@ -1078,6 +1092,14 @@ def __getitem__(self, i): named += list(mtp_named(mtp_head)) muon_params, adamw_params = split_parameters(named) + sinkhorn_tables = [] + if spec.get("engram"): + from gpu.engram import engram_param_split + + tables, projs = engram_param_split(student) + sinkhorn_tables = tables + proj_ids = {id(p) for p in projs} + muon_params = [p for p in muon_params if id(p) not in proj_ids] headwise = [] if spec.get("headwise_muon"): from gpu.muon import qk_named @@ -1092,6 +1114,11 @@ def __getitem__(self, i): if headwise: heads = int(spec.get("student_config", {}).get("num_heads", 6)) optimizer.add_headwise_group(headwise, heads=heads) + if sinkhorn_tables: + from gpu.sinkhorn_update import SinkhornUpdate + + table_opt = SinkhornUpdate(sinkhorn_tables, lr=float(spec.get("engram_lr", 5e-4))) + optimizer._engram_table_opt = table_opt # stepped alongside optimizer.add_adamw_group(adamw_params, lr=1e-4, weight_decay=0.0) print( f"[{spec_id}] muon: {len(muon_params)} matrix / " @@ -1189,6 +1216,11 @@ def _usable(ck: Path) -> bool: loss.backward() torch.nn.utils.clip_grad_norm_(student.parameters(), 1.0) optimizer.step() + # the Engram table rides its own update rule (Sinkhorn- + # balanced), stepped beside the main optimizer + table_opt = getattr(optimizer, "_engram_table_opt", None) + if table_opt is not None: + table_opt.step() scheduler.step() optimizer.zero_grad() step += 1 diff --git a/tests/test_engram_attach.py b/tests/test_engram_attach.py new file mode 100644 index 0000000..2a171fe --- /dev/null +++ b/tests/test_engram_attach.py @@ -0,0 +1,69 @@ +"""Engram attachment specs (TODO.impl/10's run): zero-init identity, +hook plumbing, param split for optimizer routing.""" + +from __future__ import annotations + +import sys +from pathlib import Path + +sys.path.insert(0, str(Path(__file__).resolve().parent.parent / "src")) + +import pytest + +torch = pytest.importorskip("torch") +transformers = pytest.importorskip("transformers") + +from gpu.engram import attach_engram, engram_param_split # noqa: E402 + + +def _tiny_t5(): + from transformers import T5Config, T5ForConditionalGeneration + + return T5ForConditionalGeneration( + T5Config( + vocab_size=259, d_model=32, d_ff=64, d_kv=16, + num_layers=4, num_decoder_layers=4, num_heads=2, + feed_forward_proj="relu", decoder_start_token_id=0, + ) + ) + + +def test_attached_model_is_identity_at_step0() -> None: + """The PKM rule: zero-init projection means the attached model's + forward equals the backbone's, bit-for-bit.""" + torch.manual_seed(0) + base = _tiny_t5().eval() + with torch.no_grad(): + for p in base.parameters(): + p.copy_(torch.randn_like(p).mul(0.1)) + attached = _tiny_t5().eval() + attached.load_state_dict(base.state_dict()) + attach_engram(attached, layer=2, entries=1024, dim=8) + + ids = torch.tensor([[5, 6, 7, 8, 1]]) + with torch.no_grad(): + out_base = base(input_ids=ids, decoder_input_ids=torch.tensor([[0]])) + out_att = attached(input_ids=ids, decoder_input_ids=torch.tensor([[0]])) + torch.testing.assert_close(out_att.logits, out_base.logits) + + +def test_memory_flows_once_projection_trains() -> None: + attached = _tiny_t5() + attach_engram(attached, layer=1, entries=4096, dim=8) + with torch.no_grad(): + attached._engram.proj.weight.normal_(std=0.1) + ids = torch.tensor([[5, 6, 7, 8, 1]]) + out = attached(input_ids=ids, decoder_input_ids=torch.tensor([[0]])).logits + ids2 = torch.tensor([[5, 6, 9, 8, 1]]) + out2 = attached(input_ids=ids2, decoder_input_ids=torch.tensor([[0]])).logits + assert not torch.allclose(out, out2), "memory must reach the logits" + + +def test_param_split_routes_table_and_projection() -> None: + model = _tiny_t5() + attach_engram(model, layer=0, entries=512, dim=8) + tables, others = engram_param_split(model) + assert len(tables) == 1 and len(others) == 1 + assert tables[0].shape == (512, 8) + assert tuple(others[0].shape) == (model.config.d_model, 8) + assert not any(p.requires_grad is False for p in tables + others)