From a0cf62bbac05bed7ffa83ccb08292b16b34f986e Mon Sep 17 00:00:00 2001 From: Parsa Assadi Date: Tue, 11 Aug 2026 21:05:16 +0000 Subject: [PATCH 1/2] Exclude bias from Muon orthogonalization --- src/maxtext/utils/muon_utils.py | 2 +- tests/unit/muon_utils_test.py | 5 +++++ 2 files changed, 6 insertions(+), 1 deletion(-) diff --git a/src/maxtext/utils/muon_utils.py b/src/maxtext/utils/muon_utils.py index ff77c57807..30bb1bd10c 100644 --- a/src/maxtext/utils/muon_utils.py +++ b/src/maxtext/utils/muon_utils.py @@ -91,9 +91,9 @@ def transform_logic(path: Tuple[str, ...]) -> Optional[mdn]: "hc_base", "sinks", "tid2eid", + "bias", ) ) - or segment == "bias" for segment in path ): return None diff --git a/tests/unit/muon_utils_test.py b/tests/unit/muon_utils_test.py index 58bfadf29a..18b8967908 100644 --- a/tests/unit/muon_utils_test.py +++ b/tests/unit/muon_utils_test.py @@ -53,6 +53,11 @@ def test_scale_is_excluded(self): def test_bias_is_excluded(self): self.assertIsNone(muon_utils.transform_logic(("decoder", "dense", "bias"))) + def test_bias_variant_is_excluded(self): + self.assertIsNone(muon_utils.transform_logic(("decoder", "moe_block", "wi_0_bias"))) + self.assertIsNone(muon_utils.transform_logic(("decoder", "GptOssMlp", "wo_bias"))) + self.assertIsNone(muon_utils.transform_logic(("decoder", "mlp", "mlp_bias"))) + def test_embedding_is_excluded(self): self.assertIsNone(muon_utils.transform_logic(("token_embedder", "embedding"))) From e9f9b546e71096fd09513b1c7ee2ff9fa4ece743 Mon Sep 17 00:00:00 2001 From: Parsa Assadi Date: Tue, 11 Aug 2026 21:12:26 +0000 Subject: [PATCH 2/2] Support Qwen3, GPT-OSS MoE blocks in Muon --- src/maxtext/utils/muon_utils.py | 4 ++-- tests/unit/muon_utils_test.py | 20 +++++++++++++++++++- 2 files changed, 21 insertions(+), 3 deletions(-) diff --git a/src/maxtext/utils/muon_utils.py b/src/maxtext/utils/muon_utils.py index 30bb1bd10c..50d5fab10b 100644 --- a/src/maxtext/utils/muon_utils.py +++ b/src/maxtext/utils/muon_utils.py @@ -101,9 +101,9 @@ def transform_logic(path: Tuple[str, ...]) -> Optional[mdn]: # 2 Special weights # 2.1 Special weights: MoE, [0, L, -2, -1] # L (optional) stands for layer when scan_layers=True - if "MoeBlock_0" in path: + if _is_path_contain_any(("MoeBlock_0", "routed_experts", "moe_block", "GptOssMlp"), path): # exclude gate - if _is_path_contain_any(("wi_0", "wi_1", "wo"), path): + if _is_path_contain_any(("wi", "wi_0", "wi_1", "wo"), path): return mdn((-2,), (-1,)) # 2.2 Special weights: Self attention diff --git a/tests/unit/muon_utils_test.py b/tests/unit/muon_utils_test.py index 18b8967908..4dfff9bb34 100644 --- a/tests/unit/muon_utils_test.py +++ b/tests/unit/muon_utils_test.py @@ -74,9 +74,27 @@ def test_moe_wi_1_uses_last_two_axes(self): def test_moe_wo_uses_last_two_axes(self): self.assertEqual(muon_utils.transform_logic(("decoder", "MoeBlock_0", "wo")), mdn((-2,), (-1,))) + def test_moe_prefused_wi_uses_last_two_axes(self): + self.assertEqual(muon_utils.transform_logic(("decoder", "MoeBlock_0", "wi")), mdn((-2,), (-1,))) + self.assertEqual(muon_utils.transform_logic(("decoder", "routed_experts", "wi")), mdn((-2,), (-1,))) + + def test_qwen3_next_moe_routed_experts(self): + self.assertEqual(muon_utils.transform_logic(("decoder", "mlp", "routed_experts", "wi_0")), mdn((-2,), (-1,))) + self.assertEqual(muon_utils.transform_logic(("decoder", "mlp", "routed_experts", "wi_1")), mdn((-2,), (-1,))) + self.assertEqual(muon_utils.transform_logic(("decoder", "mlp", "routed_experts", "wo")), mdn((-2,), (-1,))) + + def test_qwen3_moe_block(self): + self.assertEqual(muon_utils.transform_logic(("decoder", "moe_block", "wi_0")), mdn((-2,), (-1,))) + self.assertEqual(muon_utils.transform_logic(("decoder", "moe_block", "wo")), mdn((-2,), (-1,))) + + def test_gpt_oss_mlp_moe(self): + self.assertEqual(muon_utils.transform_logic(("decoder", "GptOssMlp", "wi_0")), mdn((-2,), (-1,))) + self.assertEqual(muon_utils.transform_logic(("decoder", "GptOssMlp", "wo")), mdn((-2,), (-1,))) + def test_moe_gate_falls_through_to_standard(self): - # 'gate' is inside MoeBlock_0 but not one of (wi_0, wi_1, wo) → standard. + # 'gate' is inside MoeBlock_0 but not one of (wi, wi_0, wi_1, wo) → standard. self.assertEqual(muon_utils.transform_logic(("decoder", "MoeBlock_0", "gate", "kernel")), mdn((0,), (-1,))) + self.assertEqual(muon_utils.transform_logic(("decoder", "routed_experts", "gate", "kernel")), mdn((0,), (-1,))) # --- 2.2 Self-attention --- def test_self_attention_out_projection(self):