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:
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user