Skip to content

Commit 752bea9

Browse files
author
Ronald Tse
committed
feat(distill): Engram run wiring — attach, table routing, spec
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.
1 parent ee7abb1 commit 752bea9

4 files changed

Lines changed: 180 additions & 0 deletions

File tree

‎src/gpu/distill_specs.yaml‎

Lines changed: 28 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -352,6 +352,34 @@ ara-diac-small-lite2:
352352
mode: sequence
353353
optimizer: muon
354354
note: TODO.impl/04 init-source rung; gate vs run-009 (5.78)
355+
ara-diac-small-2-1-engram:
356+
# TODO.impl/10: the lexical-memory run. Single variable vs run-007
357+
# (2.1, 4.5701): ONE Engram table on the encoder, 2M x 32, at block
358+
# 3. Zero-init projection (identity at step 0); the table trains
359+
# with the Sinkhorn-balanced rule at 5x lr.
360+
teacher: rababa_arabic_byt5/run-007-news/best
361+
teacher_volume: rababa
362+
out_volume: rababa
363+
student_init: google/byt5-small
364+
engram:
365+
layer: 3
366+
entries: 2097152
367+
dim: 32
368+
engram_lr: '0.0026'
369+
train: r5-units/domain.txt
370+
train_extra:
371+
- r5-units/replay.txt
372+
unit_limits:
373+
- 24000
374+
- 6000
375+
max_len: 1450
376+
label_beams: '1'
377+
out: rababa_arabic_distill_small/run-013-engram
378+
labels_file: teacher_labels_r7.jsonl
379+
labels_complete: 'true'
380+
mode: sequence
381+
optimizer: muon
382+
note: TODO.impl/10 lexical-memory rung; gate vs 2.1 (4.5701)
355383
ara-diac-tiny-max:
356384
# THE gapless title test: the 30M class with the campaign's FULL
357385
# lever set — full corpus (24k+6k, not the 12k subset the collapse

‎src/gpu/engram.py‎

Lines changed: 51 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -106,3 +106,54 @@ def export_int8_state(self) -> dict[str, torch.Tensor]:
106106
"table_scale": scale.reshape(1),
107107
"proj_fp16": self.proj.weight.detach().half(),
108108
}
109+
110+
111+
def attach_engram(model, layer: int = 3, entries: int = 1 << 21, dim: int = 32):
112+
"""Attach ONE Engram to the student's encoder: computes byte-n-gram
113+
addresses from input_ids and adds the looked-up memory to the
114+
hidden states leaving encoder block `layer` (0-based). The
115+
projection is zero-initialized — the attached model is functionally
116+
identical to the backbone at step 0 (the PKM rule).
117+
118+
The module rides on the model's input_ids: it re-encodes them from
119+
the ids tensor each forward, so no dataloader change is needed."""
120+
121+
eng = Engram(model.config.d_model, entries=entries, dim=dim)
122+
block = model.encoder.block[layer]
123+
124+
def hook(_module, args, output):
125+
input_ids = model._engram_ids
126+
if input_ids is None:
127+
return output
128+
memory = eng(input_ids)
129+
hidden = output[0] if isinstance(output, tuple) else output
130+
return (hidden + memory, *output[1:]) if isinstance(output, tuple) else hidden + memory
131+
132+
# capture input_ids per forward: the encoder sees them first
133+
orig_forward = model.encoder.forward
134+
135+
def encoder_forward(input_ids=None, **kw):
136+
model._engram_ids = input_ids
137+
return orig_forward(input_ids=input_ids, **kw)
138+
139+
model.encoder.forward = encoder_forward
140+
block.register_forward_hook(hook)
141+
model._engram = eng
142+
base = sum(p.numel() for p in model.parameters())
143+
print(
144+
f"[engram] attached at encoder block {layer}: "
145+
f"+{sum(p.numel() for p in eng.parameters()) / 1e6:.1f}M params "
146+
f"on a {base / 1e6:.0f}M model",
147+
flush=True,
148+
)
149+
return model
150+
151+
152+
def engram_param_split(model):
153+
"""(table_params, other_params) of the attached module — the table
154+
pairs with the Sinkhorn-balanced update, the projection stays on
155+
its optimizer."""
156+
eng = getattr(model, "_engram", None)
157+
if eng is None:
158+
return [], []
159+
return [eng.table.weight], [eng.proj.weight]

‎src/gpu/modal_distill.py‎

Lines changed: 32 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -326,6 +326,10 @@ def distill(spec_id: str, epochs: int = 3, alpha: float = 0.5, temperature: floa
326326
_maybe_stitch(spec_id, spec, student)
327327
else:
328328
student = AutoModelForSeq2SeqLM.from_pretrained(spec["student_init"]).to(device)
329+
if spec.get("engram"):
330+
from gpu.engram import attach_engram
331+
332+
attach_engram(student, **spec["engram"])
329333
student.train()
330334

331335
class Pairs(Dataset):
@@ -427,6 +431,11 @@ def val_loss() -> float:
427431
loss.backward()
428432
torch.nn.utils.clip_grad_norm_(student.parameters(), 1.0)
429433
optimizer.step()
434+
# the Engram table rides its own update rule (Sinkhorn-
435+
# balanced), stepped beside the main optimizer
436+
table_opt = getattr(optimizer, "_engram_table_opt", None)
437+
if table_opt is not None:
438+
table_opt.step()
430439
scheduler.step()
431440
optimizer.zero_grad()
432441
step += 1
@@ -745,6 +754,11 @@ def distill_sequence(spec_id: str, epochs: int = 3) -> dict:
745754
from gpu.pkm import inject_pkm
746755

747756
inject_pkm(student, **spec["pkm"])
757+
if spec.get("engram"):
758+
_ensure_src_path()
759+
from gpu.engram import attach_engram
760+
761+
attach_engram(student, **spec["engram"])
748762
mtp_head = None
749763
if spec.get("mtp_aux"):
750764
_ensure_src_path()
@@ -1078,6 +1092,14 @@ def __getitem__(self, i):
10781092

10791093
named += list(mtp_named(mtp_head))
10801094
muon_params, adamw_params = split_parameters(named)
1095+
sinkhorn_tables = []
1096+
if spec.get("engram"):
1097+
from gpu.engram import engram_param_split
1098+
1099+
tables, projs = engram_param_split(student)
1100+
sinkhorn_tables = tables
1101+
proj_ids = {id(p) for p in projs}
1102+
muon_params = [p for p in muon_params if id(p) not in proj_ids]
10811103
headwise = []
10821104
if spec.get("headwise_muon"):
10831105
from gpu.muon import qk_named
@@ -1092,6 +1114,11 @@ def __getitem__(self, i):
10921114
if headwise:
10931115
heads = int(spec.get("student_config", {}).get("num_heads", 6))
10941116
optimizer.add_headwise_group(headwise, heads=heads)
1117+
if sinkhorn_tables:
1118+
from gpu.sinkhorn_update import SinkhornUpdate
1119+
1120+
table_opt = SinkhornUpdate(sinkhorn_tables, lr=float(spec.get("engram_lr", 5e-4)))
1121+
optimizer._engram_table_opt = table_opt # stepped alongside
10951122
optimizer.add_adamw_group(adamw_params, lr=1e-4, weight_decay=0.0)
10961123
print(
10971124
f"[{spec_id}] muon: {len(muon_params)} matrix / "
@@ -1189,6 +1216,11 @@ def _usable(ck: Path) -> bool:
11891216
loss.backward()
11901217
torch.nn.utils.clip_grad_norm_(student.parameters(), 1.0)
11911218
optimizer.step()
1219+
# the Engram table rides its own update rule (Sinkhorn-
1220+
# balanced), stepped beside the main optimizer
1221+
table_opt = getattr(optimizer, "_engram_table_opt", None)
1222+
if table_opt is not None:
1223+
table_opt.step()
11921224
scheduler.step()
11931225
optimizer.zero_grad()
11941226
step += 1

‎tests/test_engram_attach.py‎

Lines changed: 69 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,69 @@
1+
"""Engram attachment specs (TODO.impl/10's run): zero-init identity,
2+
hook plumbing, param split for optimizer routing."""
3+
4+
from __future__ import annotations
5+
6+
import sys
7+
from pathlib import Path
8+
9+
sys.path.insert(0, str(Path(__file__).resolve().parent.parent / "src"))
10+
11+
import pytest
12+
13+
torch = pytest.importorskip("torch")
14+
transformers = pytest.importorskip("transformers")
15+
16+
from gpu.engram import attach_engram, engram_param_split # noqa: E402
17+
18+
19+
def _tiny_t5():
20+
from transformers import T5Config, T5ForConditionalGeneration
21+
22+
return T5ForConditionalGeneration(
23+
T5Config(
24+
vocab_size=259, d_model=32, d_ff=64, d_kv=16,
25+
num_layers=4, num_decoder_layers=4, num_heads=2,
26+
feed_forward_proj="relu", decoder_start_token_id=0,
27+
)
28+
)
29+
30+
31+
def test_attached_model_is_identity_at_step0() -> None:
32+
"""The PKM rule: zero-init projection means the attached model's
33+
forward equals the backbone's, bit-for-bit."""
34+
torch.manual_seed(0)
35+
base = _tiny_t5().eval()
36+
with torch.no_grad():
37+
for p in base.parameters():
38+
p.copy_(torch.randn_like(p).mul(0.1))
39+
attached = _tiny_t5().eval()
40+
attached.load_state_dict(base.state_dict())
41+
attach_engram(attached, layer=2, entries=1024, dim=8)
42+
43+
ids = torch.tensor([[5, 6, 7, 8, 1]])
44+
with torch.no_grad():
45+
out_base = base(input_ids=ids, decoder_input_ids=torch.tensor([[0]]))
46+
out_att = attached(input_ids=ids, decoder_input_ids=torch.tensor([[0]]))
47+
torch.testing.assert_close(out_att.logits, out_base.logits)
48+
49+
50+
def test_memory_flows_once_projection_trains() -> None:
51+
attached = _tiny_t5()
52+
attach_engram(attached, layer=1, entries=4096, dim=8)
53+
with torch.no_grad():
54+
attached._engram.proj.weight.normal_(std=0.1)
55+
ids = torch.tensor([[5, 6, 7, 8, 1]])
56+
out = attached(input_ids=ids, decoder_input_ids=torch.tensor([[0]])).logits
57+
ids2 = torch.tensor([[5, 6, 9, 8, 1]])
58+
out2 = attached(input_ids=ids2, decoder_input_ids=torch.tensor([[0]])).logits
59+
assert not torch.allclose(out, out2), "memory must reach the logits"
60+
61+
62+
def test_param_split_routes_table_and_projection() -> None:
63+
model = _tiny_t5()
64+
attach_engram(model, layer=0, entries=512, dim=8)
65+
tables, others = engram_param_split(model)
66+
assert len(tables) == 1 and len(others) == 1
67+
assert tables[0].shape == (512, 8)
68+
assert tuple(others[0].shape) == (model.config.d_model, 8)
69+
assert not any(p.requires_grad is False for p in tables + others)

0 commit comments

Comments
 (0)