From 3d3a7ec0315924804ccb278901d6f13fdd9b0943 Mon Sep 17 00:00:00 2001 From: Rita Brugarolas Date: Thu, 25 Jun 2026 01:56:30 -0700 Subject: [PATCH] [AMD] fix(moe): correct fused shared-expert scaling on aiter/DeepEP path (mori all-to-all) (#28237) Signed-off-by: Rita Brugarolas Brufau --- python/sglang/srt/layers/moe/topk.py | 30 ++++- .../moe/test_fused_shared_expert_scaling.py | 117 ++++++++++++++++++ 2 files changed, 144 insertions(+), 3 deletions(-) create mode 100644 test/registered/unit/layers/moe/test_fused_shared_expert_scaling.py diff --git a/python/sglang/srt/layers/moe/topk.py b/python/sglang/srt/layers/moe/topk.py index a97b12bf6..e740cded7 100644 --- a/python/sglang/srt/layers/moe/topk.py +++ b/python/sglang/srt/layers/moe/topk.py @@ -1641,7 +1641,7 @@ def _remap_topk_for_deepep( Routed IDs: e -> e + e // num_local_routed Shared IDs: ep_rank * num_local_experts + num_local_routed - Shared weight: 1 / routed_scaling_factor (compensates post-MoE scaling) + Shared weight: 1.0 on the aiter path, else 1/routed_scaling_factor (see below). """ if topk_ids.shape[0] == 0: return topk_ids, topk_weights @@ -1665,9 +1665,33 @@ def _remap_topk_for_deepep( + torch.arange(num_fused_shared_experts, device=topk_ids.device) ) - # Override shared weight: 1/routed_scaling_factor so net contribution = 1.0 + # Override the fused shared expert's weight so its net contribution is 1.0x. + # + # The correct value depends on whether routed_scaling_factor is applied to + # the MoE output AFTER the experts run, or already folded into the routed + # topk weights BEFORE dispatch: + # + # * Post-MoE scaling path (default): DeepseekV2MoE.forward_deepep later + # multiplies the whole MoE output by routed_scaling_factor, so the shared + # weight must be 1/routed_scaling_factor for (1/rsf) * rsf = 1.0. + # * aiter (HIP) path: aiter_biased_grouped_topk folds routed_scaling_factor + # into each routed topk weight, and forward_deepep SKIPS the post-MoE + # multiply for _use_aiter (see its `not (... or _use_aiter)` guard). The + # shared weight must therefore be 1.0 -- applying 1/rsf here would + # under-weight the always-on shared expert by routed_scaling_factor and + # corrupt every MoE layer. + # + # NOTE: forward_deepep also skips the post-MoE multiply for the non-aiter + # families where routed_scaling_factor is pre-folded in topk + # (should_fuse_routed_scaling_factor_in_topk / apply_routed_scaling_factor_on_output: + # ModelOpt NVFP4, cutlass/trtllm-routed fp8), so those would likewise need a + # 1.0 shared weight. This fix is deliberately scoped to the aiter path (the + # one validated on AMD MI355X); those other backends are left at their + # existing behavior and can be addressed by their maintainers. routed_scaling_factor = topk_config.routed_scaling_factor - if routed_scaling_factor is not None and routed_scaling_factor != 0: + if _use_aiter: + topk_weights[:, -num_fused_shared_experts:] = 1.0 + elif routed_scaling_factor is not None and routed_scaling_factor != 0: topk_weights[:, -num_fused_shared_experts:] = 1.0 / routed_scaling_factor return topk_ids, topk_weights diff --git a/test/registered/unit/layers/moe/test_fused_shared_expert_scaling.py b/test/registered/unit/layers/moe/test_fused_shared_expert_scaling.py new file mode 100644 index 000000000..321613ebb --- /dev/null +++ b/test/registered/unit/layers/moe/test_fused_shared_expert_scaling.py @@ -0,0 +1,117 @@ +"""Unit tests for fused shared-expert weight scaling on the DeepEP layout. + +These tests pin the contract of ``_remap_topk_for_deepep`` for the fused shared +expert's topk weight on the two paths this fix covers: + + * aiter (HIP) path: routed_scaling_factor is folded into the routed weights and + forward_deepep skips the post-MoE multiply, so the shared weight must be 1.0 + for a net 1.0x contribution. + * post-MoE scaling path (default): the whole MoE output is multiplied by + routed_scaling_factor afterward, so the shared weight must be 1/rsf. +""" + +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=7, suite="base-a-test-cpu") + +import unittest +from types import SimpleNamespace +from unittest.mock import patch + +import torch + +from sglang.srt.layers.moe import topk as topk_module +from sglang.srt.layers.moe.topk import TopKConfig +from sglang.test.test_utils import CustomTestCase + + +class TestFusedSharedExpertScaling(CustomTestCase): + # 256 physical routed experts over ep_size=8 -> 32 local routed per rank. + NUM_PHYSICAL_ROUTED = 256 + EP_SIZE = 8 + EP_RANK = 0 + ROUTED_SCALING_FACTOR = 2.5 + + def _run_remap(self, *, use_aiter): + # Layout: [routed, routed, routed, shared]; the trailing column is the + # fused shared expert. The shared weight starts at a sentinel to prove + # the function overwrites it. + topk_ids = torch.tensor([[5, 40, 100, 999]], dtype=torch.int32) + routed_weights = torch.tensor([1.0, 0.5, 0.25], dtype=torch.float32) + topk_weights = torch.tensor([[1.0, 0.5, 0.25, -123.0]], dtype=torch.float32) + + topk_config = TopKConfig( + top_k=4, + num_fused_shared_experts=1, + routed_scaling_factor=self.ROUTED_SCALING_FACTOR, + ) + + with ( + patch.object(topk_module, "_use_aiter", use_aiter), + patch.object( + topk_module, + "get_parallel", + return_value=SimpleNamespace( + moe_ep_size=self.EP_SIZE, moe_ep_rank=self.EP_RANK + ), + ), + ): + _out_ids, out_weights = topk_module._remap_topk_for_deepep( + topk_ids.clone(), + topk_weights.clone(), + num_fused_shared_experts=1, + num_physical_routed_experts=self.NUM_PHYSICAL_ROUTED, + topk_config=topk_config, + ) + + # Routed weights must never be touched by the shared-weight override. + self.assertTrue(torch.equal(out_weights[0, :-1], routed_weights)) + return out_weights[0, -1].item() + + def test_aiter_path_uses_unit_shared_weight(self): + # routed_scaling_factor is folded into the routed weights and the + # post-MoE multiply is skipped -> shared weight must be 1.0, NOT 1/rsf. + shared_weight = self._run_remap(use_aiter=True) + self.assertAlmostEqual(shared_weight, 1.0) + + def test_post_moe_scaling_path_compensates_with_inverse_rsf(self): + # Default path: the whole MoE output is multiplied by rsf afterward, so + # the shared weight must be 1/rsf to net out to 1.0. + shared_weight = self._run_remap(use_aiter=False) + self.assertAlmostEqual(shared_weight, 1.0 / self.ROUTED_SCALING_FACTOR) + + def test_shared_expert_ids_route_to_home_rank(self): + # Sanity: the shared slot id is placed at this rank's interleaved + # position (ep_rank * num_local_experts + num_local_routed). + topk_ids = torch.tensor([[5, 40, 100, 999]], dtype=torch.int32) + topk_weights = torch.tensor([[1.0, 0.5, 0.25, 0.0]], dtype=torch.float32) + topk_config = TopKConfig( + top_k=4, + num_fused_shared_experts=1, + routed_scaling_factor=self.ROUTED_SCALING_FACTOR, + ) + with ( + patch.object(topk_module, "_use_aiter", True), + patch.object( + topk_module, + "get_parallel", + return_value=SimpleNamespace( + moe_ep_size=self.EP_SIZE, moe_ep_rank=self.EP_RANK + ), + ), + ): + out_ids, _ = topk_module._remap_topk_for_deepep( + topk_ids.clone(), + topk_weights.clone(), + num_fused_shared_experts=1, + num_physical_routed_experts=self.NUM_PHYSICAL_ROUTED, + topk_config=topk_config, + ) + num_local_routed = self.NUM_PHYSICAL_ROUTED // self.EP_SIZE # 32 + num_local_experts = num_local_routed + 1 # 33 + expected_shared_id = self.EP_RANK * num_local_experts + num_local_routed + self.assertEqual(out_ids[0, -1].item(), expected_shared_id) + + +if __name__ == "__main__": + unittest.main()