From fbea2d717fc3a305c40c465941c1be8d8265fc7e Mon Sep 17 00:00:00 2001 From: Ronald Tse Date: Tue, 29 Sep 2026 23:25:43 +0800 Subject: [PATCH] =?UTF-8?q?feat:=20recipe=20arms=20=E2=80=94=20headwise=20?= =?UTF-8?q?Muon=20+=20Sinkhorn=20embedding=20update?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Two single-variable arms off the 2.1 recipe (run-007, 4.5701), DeepSeek-V4.1-Flash sec 2.5 / Alg 1: - ara-diac-small-2-1-hwmuon: head-wise Muon on Q/K (existing flag, first recipe run) -> run-014-hwmuon - ara-diac-small-2-1-skembed: sinkhorn_embed routes the 2D embedding-like tensors (tied table, prediction heads) from AdamW to SinkhornUpdate via the existing side-opt hook -> run-015-skembed muon.embed_named extracts the 2D embedding-like predicate so the selector and split_parameters share one classification (specs cover the MECE contract: exactly split's 2D AdamW set, never Muon's). Also fixes an engram-side routing bug: split_parameters handed _engram.table.weight to Muon while the side-opt stepped it too — double-stepped in run-013-engram. Tables are now excluded from muon_params (that run's FLAT verdict stands as measured; this corrects future engram arms). --- src/gpu/distill_specs.yaml | 47 +++++++++++++++++++++++++++++++++++++ src/gpu/modal_distill.py | 22 +++++++++++++---- src/gpu/muon.py | 36 +++++++++++++++++++--------- tests/test_muon_headwise.py | 44 +++++++++++++++++++++++++++++++++- 4 files changed, 133 insertions(+), 16 deletions(-) diff --git a/src/gpu/distill_specs.yaml b/src/gpu/distill_specs.yaml index d6d71ef..76ac5f0 100644 --- a/src/gpu/distill_specs.yaml +++ b/src/gpu/distill_specs.yaml @@ -380,6 +380,53 @@ ara-diac-small-2-1-engram: mode: sequence optimizer: muon note: TODO.impl/10 lexical-memory rung; gate vs 2.1 (4.5701) +ara-diac-small-2-1-hwmuon: + # TODO.impl/03 recipe arm: the 2.1 recipe (run-007, 4.5701) plus + # head-wise Muon on the Q/K projections (DeepSeek V4.1 Flash sec + # 2.5; GLM-5 and Kimi-K3 validate the same split). Single variable. + teacher: rababa_arabic_byt5/run-007-news/best + teacher_volume: rababa + out_volume: rababa + student_init: google/byt5-small + 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-014-hwmuon + labels_file: teacher_labels_r7.jsonl + labels_complete: 'true' + mode: sequence + optimizer: muon + headwise_muon: 'true' + note: TODO.impl/03 recipe arm; gate vs 2.1 (4.5701) +ara-diac-small-2-1-skembed: + # TODO.impl/06 recipe arm: the 2.1 recipe plus the Sinkhorn-balanced + # embedding update (Alg. 1) on the tied table / prediction head in + # place of AdamW. Single variable vs run-007. + teacher: rababa_arabic_byt5/run-007-news/best + teacher_volume: rababa + out_volume: rababa + student_init: google/byt5-small + 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-015-skembed + labels_file: teacher_labels_r7.jsonl + labels_complete: 'true' + mode: sequence + optimizer: muon + sinkhorn_embed: 'true' + sinkhorn_lr: '0.00026' + note: TODO.impl/06 recipe arm; 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/modal_distill.py b/src/gpu/modal_distill.py index 401c24f..543e977 100644 --- a/src/gpu/modal_distill.py +++ b/src/gpu/modal_distill.py @@ -1110,8 +1110,11 @@ def __getitem__(self, i): 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] + # the table rides SinkhornUpdate only — split_parameters + # would otherwise also hand it to Muon (double-stepped; + # observed in the run-013-engram wiring) + side_ids = {id(p) for p in tables + projs} + muon_params = [p for p in muon_params if id(p) not in side_ids] headwise = [] if spec.get("headwise_muon"): from gpu.muon import qk_named @@ -1119,6 +1122,13 @@ def __getitem__(self, i): headwise = [p for _, p in qk_named(named)] headwise_ids = {id(p) for p in headwise} muon_params = [p for p in muon_params if id(p) not in headwise_ids] + sinkhorn_embed = [] + if spec.get("sinkhorn_embed"): + from gpu.muon import embed_named + + sinkhorn_embed = [p for _, p in embed_named(named)] + embed_ids = {id(p) for p in sinkhorn_embed} + adamw_params = [p for p in adamw_params if id(p) not in embed_ids] optimizer = Muon( muon_params, lr=float(spec.get("muon_lr", 0.01)), momentum=0.95, weight_decay=0.01, @@ -1126,15 +1136,19 @@ 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: + if sinkhorn_tables or sinkhorn_embed: from gpu.sinkhorn_update import SinkhornUpdate - table_opt = SinkhornUpdate(sinkhorn_tables, lr=float(spec.get("engram_lr", 5e-4))) + table_opt = SinkhornUpdate( + sinkhorn_tables + sinkhorn_embed, + lr=float(spec.get("sinkhorn_lr", spec.get("engram_lr", 2.6e-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 / " f"{len(headwise)} headwise q/k / " + f"{len(sinkhorn_embed)} sinkhorn-embed / " f"{len(adamw_params)} embedding-like params", flush=True, ) diff --git a/src/gpu/muon.py b/src/gpu/muon.py index cb1ae1a..f8aac00 100644 --- a/src/gpu/muon.py +++ b/src/gpu/muon.py @@ -153,6 +153,30 @@ def _adamw_step(self, group) -> None: p.addcdiv_(exp_avg / bias_c1, denom, value=-group["lr"]) +def _is_embedding_like(name: str, p) -> bool: + return ( + p.ndim < 2 + or "embed_tokens" in name + or name == "shared.weight" # T5 tied byte embedding (transformers 5.x) + or "lm_head" in name + or "relative_attention" in name + or "memory.values" in name + or "memory.k1" in name + or "memory.k2" in name + ) + + +def embed_named(named_params): + """The 2D embedding-like tensors (tied embedding, prediction heads, + memory tables) — SinkhornUpdate's surface: rows are token identity, + columns features. 1D params (layer norms) are excluded; they have + no row/column structure to balance.""" + return [ + (name, p) for name, p in named_params + if p.requires_grad and p.ndim >= 2 and _is_embedding_like(name, p) + ] + + def split_parameters(named_params): """The standard split: orthogonalizable 2D hidden weights vs embedding-like tensors (1D params, embeddings, tied head, relative @@ -161,17 +185,7 @@ def split_parameters(named_params): for name, p in named_params: if not p.requires_grad: continue - embedding_like = ( - p.ndim < 2 - or "embed_tokens" in name - or name == "shared.weight" # T5 tied byte embedding (transformers 5.x) - or "lm_head" in name - or "relative_attention" in name - or "memory.values" in name - or "memory.k1" in name - or "memory.k2" in name - ) - (adamw if embedding_like else muon).append(p) + (adamw if _is_embedding_like(name, p) else muon).append(p) return muon, adamw diff --git a/tests/test_muon_headwise.py b/tests/test_muon_headwise.py index a00193c..6929a9e 100644 --- a/tests/test_muon_headwise.py +++ b/tests/test_muon_headwise.py @@ -14,7 +14,7 @@ torch = pytest.importorskip("torch") -from gpu.muon import Muon, qk_named, split_parameters # noqa: E402 +from gpu.muon import Muon, embed_named, qk_named, split_parameters # noqa: E402 def _grad_like(shape, seed): @@ -121,3 +121,45 @@ def test_split_parameters_unchanged_when_headwise_unused() -> None: muon, adamw = split_parameters(named) assert [p.shape for p in muon] == [torch.Size([4, 3])] assert [p.shape for p in adamw] == [torch.Size([3, 4]), torch.Size([3])] + + +def test_embed_named_selects_2d_embedding_like_only() -> None: + """The Sinkhorn surface: 2D tensors with token-row structure (tied + embedding, prediction heads). 1D layer norms stay with AdamW; 2D + hidden matrices are not embedding-like.""" + named = [ + ("shared.weight", torch.zeros(259, 4, requires_grad=True)), + ("lm_head.weight", torch.zeros(259, 4, requires_grad=True)), + ("encoder.block.0.layer.0.layer_norm.weight", + torch.zeros(4, requires_grad=True)), + ("encoder.block.0.layer.1.DenseReluDense.wi_0.weight", + torch.zeros(8, 4, requires_grad=True)), + ("encoder.embed_tokens.weight", torch.zeros(259, 4, requires_grad=True)), + ] + assert {n for n, _ in embed_named(named)} == { + "shared.weight", "lm_head.weight", "encoder.embed_tokens.weight", + } + + +def test_embed_named_agrees_with_split_parameters() -> None: + """MECE: every 2D param split_parameters routes to AdamW that + embed_named claims must be exactly embed_named's set — rerouting + them cannot strand or double-count a tensor.""" + named = [ + ("shared.weight", torch.zeros(259, 4, requires_grad=True)), + ("decoder.lm_head.weight", torch.zeros(259, 4, requires_grad=True)), + ("encoder.block.0.layer.0.layer_norm.weight", + torch.zeros(4, requires_grad=True)), + ("encoder.block.0.layer.1.DenseReluDense.wi_0.weight", + torch.zeros(8, 4, requires_grad=True)), + ("encoder.block.0.layer.0.SelfAttention.q.weight", + torch.zeros(4, 4, requires_grad=True)), + ("memory.values", torch.zeros(8, 4, requires_grad=True)), + ("encoder.block.0.layer.0.EncDecAttention.relative_attention_bias.weight", + torch.zeros(8, 4, requires_grad=True)), + ] + muon, adamw = split_parameters(named) + claimed = {id(p) for _, p in embed_named(named)} + adamw_2d = [p for p in adamw if p.ndim == 2] + assert {id(p) for p in adamw_2d} == claimed + assert not (claimed & {id(p) for p in muon})