feat(moe): add fp8 megamoe fused path and nvfp4 layout guards

Wire MegaMoE FP8 through DeepGEMM's fused fp8_mega_moe path, preserve fallback runner layouts, and add explicit NVFP4 group-size guardrails for unsupported DeepGEMM scale transforms.

Tested: PYTHONPYCACHEPREFIX=/private/tmp/sglang_pycache python3 -m py_compile python/sglang/srt/layers/moe/mega_moe.py python/sglang/srt/layers/moe/moe_runner/flashinfer_trtllm.py python/sglang/srt/layers/quantization/fp8.py python/sglang/srt/layers/quantization/modelopt_quant.py test/registered/unit/moe/test_glm_megamoe.py

Tested: git diff --check

Not-tested: python3 -m unittest test/registered/unit/moe/test_glm_megamoe.py (local Python environments do not have torch installed).
This commit is contained in:
LuminolT
2026-07-08 11:42:51 +08:00
parent eb7f44a8ee
commit 89ba17ad05
5 changed files with 465 additions and 15 deletions
+77 -1
View File
@@ -33,7 +33,10 @@ class TestGLMMegaMoE(unittest.TestCase):
def test_should_use_mega_moe_respects_env_and_token_cap(self):
hidden_states = torch.empty((2, 32))
moe = SimpleNamespace(
experts=SimpleNamespace(_mega_moe_weights_built=True),
experts=SimpleNamespace(
_mega_moe_weights_built=True,
_mega_moe_weight_format="fp4",
),
)
with patch.object(
@@ -67,6 +70,22 @@ class TestGLMMegaMoE(unittest.TestCase):
), envs.SGLANG_OPT_DEEPGEMM_MEGA_MOE_NUM_MAX_TOKENS_PER_RANK.override(2):
self.assertTrue(mega_moe.should_use_mega_moe(moe, hidden_states))
def test_should_use_mega_moe_fp8_does_not_require_fp4_act_env(self):
hidden_states = torch.empty((2, 32))
moe = SimpleNamespace(
experts=SimpleNamespace(
_mega_moe_weights_built=True,
_mega_moe_weight_format="fp8",
),
)
with patch.object(
mega_moe, "get_moe_a2a_backend", return_value=_MegaBackend()
), patch.object(
mega_moe.deep_gemm_wrapper, "ENABLE_JIT_DEEPGEMM", True
), envs.SGLANG_OPT_DEEPGEMM_MEGA_MOE_USE_FP4_ACTS.override(False):
self.assertTrue(mega_moe.should_use_mega_moe(moe, hidden_states))
def test_glm_forward_uses_megamoe_fast_path(self):
block = object.__new__(glm4_moe.Glm4MoeSparseMoeBlock)
hidden_states = torch.empty((1, 32))
@@ -125,6 +144,7 @@ class TestGLMMegaMoE(unittest.TestCase):
self.assertTrue(experts._mega_moe_weights_built)
self.assertTrue(experts._mega_moe_preserve_runner_layout)
self.assertEqual(experts._mega_moe_weight_group_size, 16)
self.assertEqual(experts._mega_moe_weight_format, "fp4")
self.assertTrue(torch.equal(experts.w13_weight.data, original_w13))
self.assertNotEqual(
experts.mega_l1_weights[0].data_ptr(),
@@ -155,6 +175,62 @@ class TestGLMMegaMoE(unittest.TestCase):
self.assertIs(same_data, int8_data)
self.assertIs(same_scale, scale)
def test_build_fp8_mega_moe_weights_preserves_runner_layout(self):
experts = SimpleNamespace(
w13_weight=_Param(torch.arange(2 * 8 * 4).reshape(2, 8, 4)),
w2_weight=_Param(torch.arange(2 * 4 * 4).reshape(2, 4, 4)),
w13_weight_scale=_Param(torch.arange(2, dtype=torch.float32) + 1),
w2_weight_scale=_Param(torch.arange(2, dtype=torch.float32) + 3),
quant_method=SimpleNamespace(
quant_config=SimpleNamespace(weight_block_size=None)
),
)
original_w13 = experts.w13_weight.data.clone()
fake_deep_gemm = types.SimpleNamespace(
transform_sf_into_required_layout=(
lambda sf, *, mn, k, recipe, num_groups, disable_ue8m0_cast: torch.zeros(
(num_groups, mn, max(k // recipe[1], 1)), dtype=torch.int32
)
),
transform_weights_for_mega_moe=lambda l1, l2: (l1, l2),
)
with patch.dict(sys.modules, {"deep_gemm": fake_deep_gemm}):
mega_moe.build_mega_moe_fp8_experts_weights(
experts, preserve_runner_layout=True, swap_w13_halves=True
)
self.assertTrue(experts._mega_moe_weights_built)
self.assertEqual(experts._mega_moe_weight_format, "fp8")
self.assertFalse(experts._mega_moe_fp8_block_quant)
self.assertTrue(torch.equal(experts.w13_weight.data, original_w13))
self.assertTrue(
torch.equal(
experts.mega_fp8_w13_weight,
torch.cat((original_w13[:, 4:], original_w13[:, :4]), dim=1),
)
)
self.assertNotEqual(
experts.mega_fp8_w13_weight.data_ptr(),
experts.w13_weight.data.data_ptr(),
)
def test_transform_sf_reports_group16_dependency_gap(self):
def fake_transform_sf(*args, **kwargs):
del args, kwargs
raise RuntimeError("Unknown SF transformation")
with self.assertRaisesRegex(RuntimeError, "group16"):
mega_moe._transform_mega_moe_sf(
fake_transform_sf,
torch.ones((2, 256, 2), dtype=torch.float32),
mn=256,
k=32,
group_size=16,
num_groups=2,
name="w13",
)
if __name__ == "__main__":
unittest.main()