[MegaMoE] Wire Qwen MoE blocks to DeepGEMM MegaMoE (MXFP4 and NVFP4 experts) (#38080)

This commit is contained in:
YAMY
2026-09-19 16:00:38 -07:00
committed by GitHub
parent 3a64faa1f2
commit 9cc7da2ab0
15 changed files with 686 additions and 127 deletions
@@ -203,6 +203,8 @@ class TestDeepGemmMegaMoeApi(CustomTestCase):
with (
patch.dict(sys.modules, {"deep_gemm": deep_gemm}),
patch.object(mega_moe, "_mega_moe_mma_type", return_value="mxf4xmxf4"),
# Toy shapes.
patch.object(mega_moe, "check_mega_moe_shapes"),
):
method.process_weights_after_loading(layer)
@@ -224,6 +226,9 @@ class TestDeepGemmMegaMoeApi(CustomTestCase):
mega_l1_weights=object(),
mega_l2_weights=object(),
should_fuse_routed_scaling_factor_in_topk=True,
moe_runner_config=SimpleNamespace(swiglu_limit=None),
_mega_moe_weights_built=True,
_mega_moe_nvfp4=False,
)
topk_output = SimpleNamespace(
topk_ids=torch.tensor([[0, 1]]),
@@ -259,14 +264,18 @@ class TestDeepGemmMegaMoeApi(CustomTestCase):
"_configure_mega_moe_deep_gemm_num_sms",
return_value=nullcontext(),
),
# Toy shapes.
patch.object(mega_moe, "check_mega_moe_shapes"),
patch.object(
mega_moe.ExpertLocationDispatchInfo,
"init_new",
return_value=object(),
),
patch(
"sglang.srt.distributed.parallel_state.get_moe_ep_group",
return_value=SimpleNamespace(device_group=object()),
"sglang.srt.runtime_context.get_parallel",
return_value=SimpleNamespace(
moe_ep_group=SimpleNamespace(device_group=object())
),
),
):
mega_moe._run_mega_routed(
@@ -282,6 +291,87 @@ class TestDeepGemmMegaMoeApi(CustomTestCase):
self.assertEqual(call.kwargs.get("mma_type"), "mxf4xmxf4")
self.assertNotIn("use_fp4_acts", call.kwargs)
def test_run_mega_routed_experts_generic_entry(self):
deep_gemm = ModuleType("deep_gemm")
deep_gemm.fp8_fp4_mega_moe = MagicMock()
buffer = SimpleNamespace(
x=object(),
x_sf=object(),
topk_idx=object(),
topk_weights=object(),
)
experts = SimpleNamespace(
num_experts=8,
mega_l1_weights=object(),
mega_l2_weights=object(),
_mega_moe_weights_built=True,
_mega_moe_nvfp4=False,
)
hidden_states = torch.zeros((3, 4), dtype=torch.bfloat16)
topk_ids = torch.tensor([[0, 1], [2, 3], [4, 5]], dtype=torch.int64)
topk_weights = torch.full((3, 2), 0.5, dtype=torch.bfloat16)
with (
patch.dict(sys.modules, {"deep_gemm": deep_gemm}),
patch.object(mega_moe, "_device_sm", 100),
patch.object(mega_moe, "_mega_moe_mma_type", return_value="fp8xfp4"),
patch.object(mega_moe, "mega_moe_pre_dispatch") as pre_dispatch,
patch.object(
mega_moe, "_get_mega_moe_symm_buffer", return_value=buffer
) as get_buffer,
patch.object(
mega_moe,
"_configure_mega_moe_deep_gemm_num_sms",
return_value=nullcontext(),
),
# Toy shapes.
patch.object(mega_moe, "check_mega_moe_shapes"),
patch(
"sglang.srt.runtime_context.get_parallel",
return_value=SimpleNamespace(
moe_ep_group=SimpleNamespace(device_group=object())
),
),
):
out = mega_moe.run_mega_routed_experts(
experts,
hidden_states,
topk_ids,
topk_weights,
hidden_size=4,
intermediate_size=8,
top_k=2,
num_tokens=3,
activation_clamp=7.0,
routed_scaling_factor=1.0,
)
self.assertEqual(out.shape, (3, 4))
self.assertEqual(out.dtype, torch.bfloat16)
buf_call = get_buffer.call_args
self.assertEqual(buf_call.kwargs.get("num_topk"), 2)
self.assertEqual(buf_call.kwargs.get("hidden"), 4)
self.assertEqual(buf_call.kwargs.get("intermediate_hidden"), 8)
# The kernel wants int32 ids and fp32 weights regardless of the router dtype.
ids_arg, weights_arg = pre_dispatch.call_args.args[1:3]
self.assertEqual(ids_arg.dtype, torch.int32)
self.assertEqual(weights_arg.dtype, torch.float32)
mega_call = deep_gemm.fp8_fp4_mega_moe.call_args
self.assertIs(mega_call.args[1], experts.mega_l1_weights)
self.assertIs(mega_call.args[2], experts.mega_l2_weights)
self.assertEqual(mega_call.kwargs.get("activation_clamp"), 7.0)
def test_shape_check_rejects_unaligned_intermediate(self):
# Qwen3-30B-A3B: intermediate 768 leaves a 24-byte scale row.
with self.assertRaisesRegex(ValueError, "multiples of 512"):
mega_moe.check_mega_moe_shapes(2048, 768, "fp8xfp4")
mega_moe.check_mega_moe_shapes(4096, 1536, "fp8xfp4")
mega_moe.check_mega_moe_shapes(4096, 1024, "mxf4xmxf4")
# 768 is a multiple of 256, so the NVFP4 (g16) rule accepts it.
mega_moe.check_mega_moe_shapes(2048, 768, "nvfp4xnvfp4")
with self.assertRaisesRegex(ValueError, "multiples of 256"):
mega_moe.check_mega_moe_shapes(2048, 384, "nvfp4xnvfp4")
def test_mxf4_l1_uses_packed_gate_up_interleave(self):
source = torch.arange(32).reshape(1, 32)
expected = torch.tensor(
@@ -189,6 +189,94 @@ class TestPrepareServerArgs(CustomTestCase):
self.assertEqual(os.environ["DG_USE_FP4_ACTS"], "0")
self.assertEqual(os.environ["DG_USE_MXF4_KIND"], "0")
def test_megamoe_rejects_two_batch_overlap(self):
# The fused kernel has no dispatch/combine split for the TBO ops to call.
with override_platform(is_cuda=True, is_sm90=False, is_sm100=True):
args = ServerArgs(
model_path="dummy",
moe_a2a_backend="megamoe",
enable_two_batch_overlap=True,
)
with self.assertRaisesRegex(ValueError, "overlap"):
args.resolve_once()
def test_megamoe_requires_sm90_or_sm100(self):
with override_platform(is_cuda=True, is_sm90=False, is_sm100=False):
args = ServerArgs(model_path="dummy", moe_a2a_backend="megamoe")
with self.assertRaisesRegex(ValueError, "SM90"):
args.resolve_once()
with override_platform(is_cuda=False, is_sm90=False, is_sm100=False):
args = ServerArgs(model_path="dummy", moe_a2a_backend="megamoe")
with self.assertRaisesRegex(ValueError, "CUDA"):
args.resolve_once()
with override_platform(is_cuda=True, is_sm90=False, is_sm100=True):
ServerArgs(model_path="dummy", moe_a2a_backend="megamoe").resolve_once()
def test_megamoe_token_budget_must_cover_chunked_prefill(self):
from sglang.srt.arg_groups.mega_moe_hook import validate_mega_moe_token_budget
from sglang.srt.environ import envs
with override_platform(is_cuda=True, is_sm90=False, is_sm100=True):
args = ServerArgs(
model_path="dummy",
moe_a2a_backend="megamoe",
chunked_prefill_size=16384,
)
args.resolve_once()
with envs.SGLANG_OPT_DEEPGEMM_MEGA_MOE_NUM_MAX_TOKENS_PER_RANK.override(
8192
):
with self.assertRaisesRegex(ValueError, "required_per_rank=16384"):
validate_mega_moe_token_budget(args, "Qwen3MoeForCausalLM")
with envs.SGLANG_OPT_DEEPGEMM_MEGA_MOE_NUM_MAX_TOKENS_PER_RANK.override(
16384
):
validate_mega_moe_token_budget(args, "Qwen3MoeForCausalLM")
def test_megamoe_token_budget_gate_by_arch_and_nvfp4(self):
from sglang.srt.arg_groups.mega_moe_hook import mega_moe_needs_token_budget
for arch in (
"InternS2PreviewForConditionalGeneration",
"MellumForCausalLM",
"Qwen3MoeForCausalLM",
"DeepseekV4ForCausalLM",
):
self.assertTrue(mega_moe_needs_token_budget(arch, {}, None), arch)
# MXFP4 DeepSeek-family models keep their runtime fallback.
self.assertFalse(mega_moe_needs_token_budget("DeepseekV3ForCausalLM", {}, None))
# NVFP4 experts are repacked at load: every model is checked.
self.assertTrue(
mega_moe_needs_token_budget(
"DeepseekV3ForCausalLM", {"quant_algo": "NVFP4"}, None
)
)
self.assertTrue(
mega_moe_needs_token_budget("DeepseekV3ForCausalLM", {}, "modelopt_fp4")
)
def test_megamoe_decode_tokens_per_rank_follows_graph_bs_and_draft_tokens(self):
# cuda_graph_config only exists on a GPU host; use a stand-in view.
from sglang.srt.arg_groups.mega_moe_hook import mega_moe_decode_tokens_per_rank
def view(max_bs, algorithm=None, draft_tokens=None, cg=True):
return SimpleNamespace(
cuda_graph_config=(
SimpleNamespace(decode=SimpleNamespace(max_bs=max_bs))
if cg
else None
),
speculative_algorithm=algorithm,
speculative_num_draft_tokens=draft_tokens,
)
self.assertEqual(mega_moe_decode_tokens_per_rank(view(16384)), 16384)
self.assertEqual(mega_moe_decode_tokens_per_rank(view(256, "EAGLE", 4)), 1024)
# Draft tokens only count under a speculative algorithm.
self.assertEqual(mega_moe_decode_tokens_per_rank(view(256, None, 4)), 256)
self.assertEqual(mega_moe_decode_tokens_per_rank(view(None)), 0)
self.assertEqual(mega_moe_decode_tokens_per_rank(view(0, cg=False)), 0)
def test_w4a4_mxfp4_megamoe_disabled_preserves_deepgemm_env(self):
deepgemm_env = {
"DG_USE_FP4_ACTS": "0",