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
47 changes: 47 additions & 0 deletions src/gpu/distill_specs.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
22 changes: 18 additions & 4 deletions src/gpu/modal_distill.py
Original file line number Diff line number Diff line change
Expand Up @@ -1110,31 +1110,45 @@ 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

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,
)
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,
)
Expand Down
36 changes: 25 additions & 11 deletions src/gpu/muon.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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


Expand Down
44 changes: 43 additions & 1 deletion tests/test_muon_headwise.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down Expand Up @@ -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})
Loading