Skip to content

Commit 464ad76

Browse files
author
Ronald Tse
committed
feat(e5): MTP-aux distillation rung — registered + implemented
EXPERIMENTS.md E5: control = run-006-r7-muon verbatim, single delta = a 3-step multi-token-prediction head (~1.1M params, 0.4%) on the student decoder, aux CE at beta 0.15, head discarded at inference so the shipped artifact stays vanilla and size-identical. Source: Hy4's native MTP mapped as a TRAINING auxiliary (TODO 07-hy4) — serving-side speculation stays parked, decode measured non-binding. gpu/mtp.py (MTPHead + build/mtp_named), distill_sequence wiring (head before optimizer so Muon sees its matrices; mtp_head.pt saved beside student.pt for resume; provenance-only copy under best/), spec ara-diac-small-2-mtp -> run-007-r7-muon-mtp. 3 unit tests (shift/mask semantics, uniform-logit baseline, naming). Gate (pre-agreed): adopt at <=4.5218 full-set windowed DER; honest report in [4.5218, 4.8218). Prediction: 4.5-4.75.
1 parent 70b4d38 commit 464ad76

5 files changed

Lines changed: 190 additions & 2 deletions

File tree

‎docs/EXPERIMENTS.md‎

Lines changed: 25 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -195,6 +195,31 @@ All rows passed the CER parity gate at release. Readings:
195195
teacher+0.5pp budget remains the disclosed north star.
196196
- **Prediction (registered):** 4.3–5.0, by E3's 4.829 on weaker labels.
197197

198+
## E5 — MTP-aux distillation rung (run-007-r7-muon-mtp)
199+
200+
- **Status:** REGISTERED 2026-09-01, launching.
201+
- **Source:** Tencent Hy4-preview carries a native multi-token
202+
prediction layer; mapped to our stack as a TRAINING auxiliary (TODO
203+
07-hy4) — per-position multi-step heads densify supervision on the
204+
decode path; serving-side speculation stays parked (decode measured
205+
non-binding at our sizes).
206+
- **Hypothesis:** the student's residual errors concentrate where the
207+
single-step objective leaves the byte decision underconstrained;
208+
forcing each decoder position to also predict t+1..t+3 regularizes
209+
the hidden state toward the local sequence structure that
210+
diacritization output exhibits (letter + haraqat pattern).
211+
- **Design:** control = run-006-r7-muon verbatim (same teacher labels,
212+
corpus, limits, schedule, Muon groups); single delta = MTPHead
213+
attached to the student decoder (3 steps, byte-vocab 259, ~1.1M
214+
params ≈ 0.4%), auxiliary CE weighted β=0.15, head DISCARDED at
215+
inference (zero serving cost/size delta in the shipped artifact).
216+
- **Pre-agreed gate:** adopt if full-set windowed DER ≤ 4.5218
217+
(≥0.3pp over 4.8218, the E3-style bar); report honestly in
218+
[4.5218, 4.8218); investigate if worse.
219+
- **Prediction (registered):** 4.5–4.75 — denser supervision helps
220+
the tail, but the factorial attributes most of the remaining gap to
221+
domain coverage, so the effect should be second-order.
222+
198223
## Parked
199224

200225
- **Speculative decoding** (LongCat converts sparsity→speed): revisit

‎src/gpu/distill_specs.yaml‎

Lines changed: 23 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -116,6 +116,29 @@ ara-diac-small-muon:
116116
labels_complete: 'true'
117117
mode: sequence
118118
note: vanilla ByT5-small + Muon (factorial cell 4)
119+
ara-diac-small-2-mtp:
120+
# E5 (EXPERIMENTS.md): control = ara-diac-small-2 verbatim; single
121+
# delta = MTP-aux head (3 steps, beta 0.15), discarded at inference
122+
teacher: rababa_arabic_byt5/run-007-news/best
123+
teacher_volume: rababa
124+
student_init: google/byt5-small
125+
optimizer: muon
126+
muon_lr: '0.01'
127+
train: r5-units/domain.txt
128+
train_extra:
129+
- r5-units/replay.txt
130+
unit_limits:
131+
- 24000
132+
- 6000
133+
max_len: 1450
134+
label_beams: '1'
135+
out: rababa_arabic_distill_small/run-007-r7-muon-mtp
136+
labels_file: teacher_labels_r7.jsonl
137+
mtp_aux:
138+
steps: 3
139+
beta: 0.15
140+
mode: sequence
141+
119142
ara-diac-small-2:
120143
teacher: rababa_arabic_byt5/run-007-news/best
121144
teacher_volume: rababa

‎src/gpu/modal_distill.py‎

Lines changed: 40 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -429,6 +429,8 @@ def val_loss() -> float:
429429
ck.mkdir(exist_ok=True)
430430
torch.save(student.state_dict(), ck / "student.pt")
431431
torch.save(optimizer.state_dict(), ck / "optim.pt")
432+
if mtp_head is not None:
433+
torch.save(mtp_head.state_dict(), ck / "mtp_head.pt")
432434
CHECKPOINTS.commit()
433435

434436
vl = val_loss()
@@ -742,6 +744,14 @@ def distill_sequence(spec_id: str, epochs: int = 3) -> dict:
742744
from gpu.pkm import inject_pkm
743745

744746
inject_pkm(student, **spec["pkm"])
747+
mtp_head = None
748+
if spec.get("mtp_aux"):
749+
_ensure_src_path()
750+
from gpu.mtp import build_mtp
751+
752+
mtp_head = build_mtp(student, **spec["mtp_aux"])
753+
n = sum(q.numel() for q in mtp_head.parameters()) / 1e6
754+
print(f"[{spec_id}] mtp_aux head: {n:.2f}M params", flush=True)
745755
student.train()
746756

747757
class Pairs(Dataset):
@@ -1043,7 +1053,12 @@ def __getitem__(self, i):
10431053
_ensure_src_path()
10441054
from gpu.muon import Muon, split_parameters
10451055

1046-
muon_params, adamw_params = split_parameters(student.named_parameters())
1056+
named = list(student.named_parameters())
1057+
if mtp_head is not None:
1058+
from gpu.mtp import mtp_named
1059+
1060+
named += list(mtp_named(mtp_head))
1061+
muon_params, adamw_params = split_parameters(named)
10471062
optimizer = Muon(
10481063
muon_params, lr=float(spec.get("muon_lr", 0.01)),
10491064
momentum=0.95, weight_decay=0.01,
@@ -1083,6 +1098,16 @@ def _usable(ck: Path) -> bool:
10831098
optimizer.load_state_dict(
10841099
torch.load(ckpts[-1] / "optim.pt", map_location="cpu", weights_only=True)
10851100
)
1101+
if mtp_head is not None and (ckpts[-1] / "mtp_head.pt").exists():
1102+
mtp_head.load_state_dict(
1103+
torch.load(ckpts[-1] / "mtp_head.pt", map_location="cpu", weights_only=True)
1104+
)
1105+
elif mtp_head is not None:
1106+
print(
1107+
f"[{spec_id}] WARNING: no mtp_head.pt at resume — "
1108+
"fresh head, aux dynamics reset",
1109+
flush=True,
1110+
)
10861111
step = int(ckpts[-1].name.split("-")[1])
10871112
for _ in range(step):
10881113
scheduler.step()
@@ -1093,7 +1118,18 @@ def _usable(ck: Path) -> bool:
10931118
if step >= total_steps:
10941119
break
10951120
ids, am, labels = ids.to("cuda"), am.to("cuda"), labels.to("cuda")
1096-
loss = student(input_ids=ids, attention_mask=am, labels=labels).loss
1121+
if mtp_head is not None:
1122+
beta = float(spec["mtp_aux"].get("beta", 0.15))
1123+
s_out = student(
1124+
input_ids=ids, attention_mask=am, labels=labels,
1125+
output_hidden_states=True,
1126+
)
1127+
aux = mtp_head.aux_loss(
1128+
s_out.decoder_hidden_states[-1], labels
1129+
)
1130+
loss = s_out.loss + beta * aux
1131+
else:
1132+
loss = student(input_ids=ids, attention_mask=am, labels=labels).loss
10971133
loss.backward()
10981134
torch.nn.utils.clip_grad_norm_(student.parameters(), 1.0)
10991135
optimizer.step()
@@ -1119,6 +1155,8 @@ def _usable(ck: Path) -> bool:
11191155
best.mkdir(exist_ok=True)
11201156
student.save_pretrained(str(best))
11211157
student_tok.save_pretrained(str(best))
1158+
if mtp_head is not None: # provenance only; never in the shipped artifact
1159+
torch.save(mtp_head.state_dict(), best / "mtp_head.pt")
11221160
CHECKPOINTS.commit()
11231161
SECRYST_CHECKPOINTS.commit()
11241162
PERSIAN_CHECKPOINTS.commit()

‎src/gpu/mtp.py‎

Lines changed: 51 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,51 @@
1+
"""Multi-token-prediction auxiliary head (E5, TODO 07-hy4).
2+
3+
Per-position multi-step prediction as a TRAINING auxiliary: each
4+
decoder position's hidden state additionally predicts the target
5+
tokens at t+1..t+k, densifying supervision on the decode path. The
6+
head is discarded at inference — the shipped student stays vanilla
7+
and size-identical; only the training run carries it (saved as
8+
mtp_head.pt beside student.pt for resume, never exported).
9+
"""
10+
11+
from __future__ import annotations
12+
13+
import torch
14+
from torch import nn
15+
import torch.nn.functional as F
16+
17+
18+
class MTPHead(nn.Module):
19+
def __init__(self, d_model: int, vocab: int, steps: int = 3):
20+
super().__init__()
21+
self.steps = steps
22+
self.heads = nn.ModuleList(
23+
nn.Linear(d_model, vocab) for _ in range(steps)
24+
)
25+
26+
def aux_loss(self, hidden: torch.Tensor, labels: torch.Tensor) -> torch.Tensor:
27+
"""hidden [B, T, d] is the decoder's final hidden state whose
28+
position t already predicts labels[t] through the main head;
29+
step-k heads predict labels[t+k] from the same position."""
30+
total = 0.0
31+
n = 0
32+
for k, head in enumerate(self.heads, start=1):
33+
tgt = labels[:, k:]
34+
m = tgt != -100
35+
if m.any():
36+
total = total + F.cross_entropy(
37+
head(hidden[:, :-k])[m].float(), tgt[m]
38+
)
39+
n += 1
40+
return total / max(n, 1)
41+
42+
43+
def build_mtp(student, steps: int = 3) -> MTPHead:
44+
return MTPHead(
45+
student.config.d_model, student.config.vocab_size, steps=steps
46+
).to(student.device)
47+
48+
49+
def mtp_named(head: MTPHead):
50+
for name, p in head.named_parameters():
51+
yield f"mtp.{name}", p

‎tests/test_gpu_mtp.py‎

Lines changed: 51 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,51 @@
1+
import torch
2+
3+
from gpu.mtp import MTPHead, mtp_named
4+
5+
6+
def _labels():
7+
t = torch.full((2, 6), -100)
8+
t[0, :5] = torch.tensor([10, 11, 12, 13, 14])
9+
t[1, :3] = torch.tensor([4, 5, 6])
10+
return t
11+
12+
13+
def test_mtp_head_shapes_and_shift():
14+
torch.manual_seed(0)
15+
head = MTPHead(d_model=8, vocab=16, steps=3)
16+
hidden = torch.randn(2, 6, 8)
17+
loss = head.aux_loss(hidden, _labels())
18+
assert loss.ndim == 0 and float(loss) > 0
19+
20+
21+
def test_mtp_shift_targets_not_inputs():
22+
torch.manual_seed(0)
23+
head = MTPHead(d_model=8, vocab=16, steps=1)
24+
head.heads[0].weight.data.zero_()
25+
head.heads[0].bias.data.zero_()
26+
27+
labels = _labels()
28+
hidden = torch.randn(2, 6, 8)
29+
# uniform-logit head: CE is log(vocab) wherever any target is valid
30+
loss = head.aux_loss(hidden, labels)
31+
import math
32+
33+
assert math.isclose(float(loss), math.log(16), rel_tol=1e-4)
34+
35+
# a step-1 head must see labels[:, 1:] as targets: with all -100
36+
# beyond position 0 there is nothing to predict
37+
only_first = torch.full((1, 4), -100)
38+
only_first[0, 0] = 7.0
39+
assert float(head.aux_loss(torch.randn(1, 4, 8), only_first)) == 0.0
40+
41+
42+
def test_mtp_learnable_and_named():
43+
head = MTPHead(d_model=8, vocab=16, steps=2)
44+
named = dict(mtp_named(head))
45+
assert set(named) == {
46+
"mtp.heads.0.weight", "mtp.heads.0.bias",
47+
"mtp.heads.1.weight", "mtp.heads.1.bias",
48+
}
49+
out = head.aux_loss(torch.randn(2, 6, 8), _labels())
50+
out.backward()
51+
assert head.heads[0].weight.grad is not None

0 commit comments

Comments
 (0)