From 687c525384c5ce0424c929ce5cd54a6524181dba Mon Sep 17 00:00:00 2001 From: thad0ctor Date: Thu, 9 Jul 2026 22:01:29 -0700 Subject: [PATCH] fix: address Codex review on Gefen-X parameter routing - make_gefenx_param_groups emits (name, param) pairs so gefen keeps the real parameter names instead of synthesizing group_i_param_j; without this GefenXConfig.period_one_substrings (matched against names) never fired. - Only non-empty buckets become param groups: gefen rejects an empty parameter group, so an embeddings-only run (all trainable params are modules_to_save) with embedding_lr set crashed on the empty non_embeddings group. - tests: (name, param) emission, embedding-only omits the empty group, and real gefen regression tests for name preservation + embedding-only construction. 35 tests pass. --- tests/python/test_gefenx_optimizer.py | 74 +++++++++++++++++++++++++++ unsloth/optimizers/gefenx.py | 21 ++++++-- 2 files changed, 90 insertions(+), 5 deletions(-) diff --git a/tests/python/test_gefenx_optimizer.py b/tests/python/test_gefenx_optimizer.py index 89fce7a4d3..6a7531cbd5 100644 --- a/tests/python/test_gefenx_optimizer.py +++ b/tests/python/test_gefenx_optimizer.py @@ -169,6 +169,32 @@ def test_param_groups_splits_embedding_lr(): assert embed["lr"] == 5e-6 and len(embed["params"]) == 1 +def test_param_groups_emit_named_pairs(): + # Params must be (name, param) pairs so gefen keeps real names (period_one_substrings). + p = _FakeParam() + model = _FakeModel([("model.layers.0.self_attn.q_proj.weight", p)]) + groups = gefenx.make_gefenx_param_groups(model, lr = 1e-4, weight_decay = 0.0) + entry = groups[0]["params"][0] + assert isinstance(entry, tuple) and entry[0] == "model.layers.0.self_attn.q_proj.weight" + assert entry[1] is p + + +def test_param_groups_embedding_only_omits_empty_group(): + # All trainable params are PEFT modules_to_save + embedding_lr => a single + # embeddings group, NOT an empty non_embeddings group (gefen rejects empties). + model = _FakeModel( + [ + ("model.embed_tokens.modules_to_save.default.weight", _FakeParam()), + ("lm_head.modules_to_save.default.weight", _FakeParam()), + ] + ) + groups = gefenx.make_gefenx_param_groups(model, lr = 1e-4, weight_decay = 0.0, embedding_lr = 5e-6) + assert len(groups) == 1 + assert groups[0]["lr"] == 5e-6 + assert len(groups[0]["params"]) == 2 + assert all(len(g["params"]) > 0 for g in groups) + + # --------------------------------------------------------------------------- # # build_gefenx_optimizer # --------------------------------------------------------------------------- # @@ -504,6 +530,54 @@ def test_real_gefenx_muon_updates_all_params_cpu(): assert _num_changed(torch, before, model) == len(before) +def test_real_gefenx_preserves_param_names(): + # gefen must receive real names (not synthesized group_i_param_j), otherwise + # GefenXConfig.period_one_substrings can never match. + torch = pytest.importorskip("torch") + pytest.importorskip("gefen") + + model = torch.nn.Sequential(torch.nn.Embedding(16, 8), torch.nn.Linear(8, 8)) + opt = gefenx.build_gefenx_optimizer( + model, + _GefenXConfig(fused = False, period_one_substrings = ("embed",)), + lr = 1e-3, + weight_decay = 0.0, + betas = (0.9, 0.999), + eps = 1e-8, + ) + names = [opt.state[p].get("name") for g in opt.param_groups for p in g["params"]] + assert names and not any(str(n).startswith("group_") for n in names) + + +def test_real_gefenx_embedding_only_builds(): + # Regression: an embeddings-only run (all trainable params are modules_to_save) + # with embedding_lr set must not crash on an empty non_embeddings group. + torch = pytest.importorskip("torch") + pytest.importorskip("gefen") + from gefen import Gefen + + class _EmbedOnly(torch.nn.Module): + def __init__(self): + super().__init__() + self.w = torch.nn.Parameter(torch.randn(8, 4)) + + def named_parameters(self, *a, **k): + yield "base.embed_tokens.modules_to_save.default.weight", self.w + + model = _EmbedOnly() + opt = gefenx.build_gefenx_optimizer( + model, + _GefenXConfig(fused = False), + lr = 1e-3, + weight_decay = 0.0, + betas = (0.9, 0.999), + eps = 1e-8, + embedding_lr = 5e-6, + ) + assert isinstance(opt, Gefen) + assert all(len(g["params"]) > 0 for g in opt.param_groups) + + @pytest.mark.skipif(not _CUDA, reason = "requires NVIDIA CUDA for the fused gefen kernels") def test_real_gefenx_cuda_fused_updates_all_params(): import torch diff --git a/unsloth/optimizers/gefenx.py b/unsloth/optimizers/gefenx.py index 0cc284e126..434ac299d1 100644 --- a/unsloth/optimizers/gefenx.py +++ b/unsloth/optimizers/gefenx.py @@ -182,6 +182,13 @@ def make_gefenx_param_groups( Matches Unsloth's ``_create_unsloth_optimizer`` embedding split: params saved via PEFT ``modules_to_save`` (embeddings / heads trained at full rank) get the dedicated ``embedding_lr`` when one is provided; everything else shares ``lr``. + + Params are emitted as ``(name, param)`` pairs so gefen keeps the real + parameter names — otherwise it synthesises ``group_i_param_j`` names and + ``GefenXConfig.period_one_substrings`` (matched against names) never fires. + Only non-empty buckets become groups: gefen rejects an empty parameter group, + so an embeddings-only run (all trainable params are ``modules_to_save``) would + otherwise crash on the empty ``non_embeddings`` group. """ non_embeddings: List[Any] = [] embeddings: List[Any] = [] @@ -195,17 +202,21 @@ def make_gefenx_param_groups( print( f"Unsloth: Setting lr = {embedding_lr:.2e} instead of {lr:.2e} for {partial_name}." ) - embeddings.append(param) + embeddings.append((name, param)) else: - non_embeddings.append(param) + non_embeddings.append((name, param)) - param_groups: List[Dict[str, Any]] = [ - {"params": non_embeddings, "weight_decay": weight_decay, "lr": lr}, - ] + param_groups: List[Dict[str, Any]] = [] + if non_embeddings: + param_groups.append({"params": non_embeddings, "weight_decay": weight_decay, "lr": lr}) if embeddings: param_groups.append( {"params": embeddings, "weight_decay": weight_decay, "lr": embedding_lr} ) + if not param_groups: + # No trainable params at all — hand gefen a single (empty) group so it + # raises its own clear "empty parameter list" error. + param_groups.append({"params": non_embeddings, "weight_decay": weight_decay, "lr": lr}) return param_groups