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
28 changes: 28 additions & 0 deletions src/gpu/distill_specs.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
51 changes: 51 additions & 0 deletions src/gpu/engram.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]
32 changes: 32 additions & 0 deletions src/gpu/modal_distill.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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()
Expand Down Expand Up @@ -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
Expand All @@ -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 / "
Expand Down Expand Up @@ -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
Expand Down
69 changes: 69 additions & 0 deletions tests/test_engram_attach.py
Original file line number Diff line number Diff line change
@@ -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)
Loading