diff --git a/src/maxtext/utils/muon_utils.py b/src/maxtext/utils/muon_utils.py index ff77c57807..50d5fab10b 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 @@ -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 58bfadf29a..61f32549a1 100644 --- a/tests/unit/muon_utils_test.py +++ b/tests/unit/muon_utils_test.py @@ -53,6 +53,14 @@ 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): + # Bias names (e.g., in MoE with mlp_bias=True) are excluded across all expert projections. + self.assertIsNone(muon_utils.transform_logic(("decoder", "moe_block", "wi_0_bias"))) + self.assertIsNone(muon_utils.transform_logic(("decoder", "GptOssMlp", "wi_0_bias"))) + self.assertIsNone(muon_utils.transform_logic(("decoder", "GptOssMlp", "wi_1_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"))) @@ -69,9 +77,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):