[MegaMoE] Wire Qwen MoE blocks to DeepGEMM MegaMoE (MXFP4 and NVFP4 experts) (#38080)
This commit is contained in:
@@ -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",
|
||||
|
||||
Reference in New Issue
Block a user