From 07c8f7294daa12087e969c89a11f773d36b9d42d Mon Sep 17 00:00:00 2001 From: Baizhou Zhang Date: Mon, 31 Aug 2026 17:16:34 -0700 Subject: [PATCH] Bump sgl-deep-gemm to 0.1.7 (#37279) --- docker/Dockerfile | 2 +- python/pyproject.toml | 2 +- python/sglang/srt/arg_groups/mega_moe_hook.py | 11 - python/sglang/srt/layers/moe/mega_moe.py | 39 ++- .../sglang/srt/layers/quantization/mxfp4.py | 4 + python/sglang/srt/server_args.py | 4 +- .../layers/moe/test_mega_moe_deepgemm_api.py | 262 ++++++++++++++++++ .../unit/server_args/test_server_args.py | 6 +- 8 files changed, 301 insertions(+), 29 deletions(-) create mode 100644 test/registered/unit/layers/moe/test_mega_moe_deepgemm_api.py diff --git a/docker/Dockerfile b/docker/Dockerfile index ceedf8dae..78b8054bc 100644 --- a/docker/Dockerfile +++ b/docker/Dockerfile @@ -6,7 +6,7 @@ ARG BUILD_TYPE=all ARG BRANCH_TYPE=remote ARG SGL_KERNEL_VERSION=0.4.6.post1 ARG SGL_VERSION -ARG SGL_DEEP_GEMM_VERSION=0.1.6 +ARG SGL_DEEP_GEMM_VERSION=0.1.7 ARG USE_LATEST_SGLANG=0 ARG GDRCOPY_VERSION=2.5.1 ARG SGL_NCCL_VERSION=2.30.7 diff --git a/python/pyproject.toml b/python/pyproject.toml index b82df0fb2..6b6b0aa12 100755 --- a/python/pyproject.toml +++ b/python/pyproject.toml @@ -75,7 +75,7 @@ dependencies = [ "sentencepiece", "setproctitle", "sgl-deep-ep==0.1.2", - "sgl-deep-gemm==0.1.6", + "sgl-deep-gemm==0.1.7", "sglang-kernel==0.4.6.post1", "smg-grpc-servicer>=0.5.0", "soundfile==0.13.1", diff --git a/python/sglang/srt/arg_groups/mega_moe_hook.py b/python/sglang/srt/arg_groups/mega_moe_hook.py index 99777ce15..148f8a3ce 100644 --- a/python/sglang/srt/arg_groups/mega_moe_hook.py +++ b/python/sglang/srt/arg_groups/mega_moe_hook.py @@ -1,7 +1,6 @@ from __future__ import annotations import logging -import os from typing import TYPE_CHECKING if TYPE_CHECKING: @@ -17,7 +16,6 @@ logger = logging.getLogger(__name__) def handle_mega_moe(server_args: ServerArgs) -> None: handle_moe_runner_backend_alias(server_args) - handle_w4a4_mxfp4_megamoe_env(server_args) def handle_moe_runner_backend_alias(server_args: ServerArgs) -> None: @@ -38,12 +36,3 @@ def handle_moe_runner_backend_alias(server_args: ServerArgs) -> None: moe_runner_backend="auto", moe_a2a_backend="megamoe", ) - - -def handle_w4a4_mxfp4_megamoe_env(server_args: ServerArgs) -> None: - cfg = resolving_view(server_args) - if not cfg.enable_w4a4_mxfp4_megamoe: - return - - os.environ["DG_USE_FP4_ACTS"] = "1" - os.environ["DG_USE_MXF4_KIND"] = "1" diff --git a/python/sglang/srt/layers/moe/mega_moe.py b/python/sglang/srt/layers/moe/mega_moe.py index 3485fd841..0ecf88497 100644 --- a/python/sglang/srt/layers/moe/mega_moe.py +++ b/python/sglang/srt/layers/moe/mega_moe.py @@ -16,7 +16,6 @@ from __future__ import annotations import functools -import os from contextlib import contextmanager, nullcontext from typing import TYPE_CHECKING, Optional @@ -34,6 +33,7 @@ from sglang.srt.layers.moe.mega_moe_sm90 import ( from sglang.srt.layers.moe.utils import get_moe_a2a_backend from sglang.srt.model_executor.runner import get_is_capture_mode from sglang.srt.models.deepseek_common.utils import _device_sm +from sglang.srt.runtime_context import get_exec if TYPE_CHECKING: from deep_gemm import SymmBuffer @@ -45,6 +45,10 @@ if TYPE_CHECKING: _MEGA_MOE_SYMM_BUFFER: dict = {} +def _mega_moe_mma_type() -> str: + return "mxf4xmxf4" if get_exec().moe.enable_w4a4_mxfp4_megamoe else "fp8xfp4" + + @functools.lru_cache(maxsize=1) def _mega_moe_max_num_sms() -> Optional[int]: if _device_sm < 100: @@ -92,6 +96,7 @@ def _get_mega_moe_symm_buffer( ) -> SymmBuffer: import deep_gemm + mma_type = _mega_moe_mma_type() key = ( id(group), num_max_tokens_per_rank, @@ -99,6 +104,7 @@ def _get_mega_moe_symm_buffer( num_topk, hidden, intermediate_hidden, + mma_type, ) buf = _MEGA_MOE_SYMM_BUFFER.get(key) if buf is None: @@ -109,7 +115,7 @@ def _get_mega_moe_symm_buffer( num_topk, hidden, intermediate_hidden, - use_fp8_dispatch=True, + mma_type=mma_type, activation="swiglu", ) _MEGA_MOE_SYMM_BUFFER[key] = buf @@ -248,8 +254,8 @@ def _run_mega_routed( num_tokens, ) - use_fp4_acts = os.getenv("DG_USE_FP4_ACTS") == "1" - if use_fp4_acts: + mma_type = _mega_moe_mma_type() + if mma_type == "mxf4xmxf4": # FP4 path goes through DeepGEMM's mega_moe_pre_dispatch which # handles the E2M1 packing variant. The jit implementation # only emits FP8. @@ -263,7 +269,7 @@ def _run_mega_routed( buf.topk_weights, num_tokens=num_tokens, group_size=32, - use_fp4_acts=True, + mma_type=mma_type, ) else: mega_moe_pre_dispatch( @@ -304,22 +310,30 @@ def _run_mega_routed( def _interleave_mega_moe_gate_up(t: torch.Tensor, gran: int = 8) -> torch.Tensor: - # Match DeepGEMM's L1 gate/up layout: - # [gate: 0..7, up: 0..7, gate: 8..15, up: 8..15, ...]. + # Match DeepGEMM's L1 gate/up layouts. FP8 activations use contiguous + # gran-8 chunks; packed MXFP4 activations use even/odd gran-16 chunks. num_groups, n, *rest = t.shape half = n // 2 gate = t[:, :half].reshape(num_groups, half // gran, gran, *rest) up = t[:, half:].reshape(num_groups, half // gran, gran, *rest) - result = torch.stack([gate, up], dim=2).reshape(num_groups, n, *rest) + if gran == 16: + result = torch.cat( + [gate[:, :, 0::2], up[:, :, 0::2], gate[:, :, 1::2], up[:, :, 1::2]], + dim=2, + ).reshape(num_groups, n, *rest) + else: + result = torch.stack([gate, up], dim=2).reshape(num_groups, n, *rest) return torch.empty_like(t).copy_(result) def _interleave_mega_moe_l1_weights( l1_weights: tuple[torch.Tensor, torch.Tensor], + mma_type: str, ) -> tuple[torch.Tensor, torch.Tensor]: + gran = 16 if mma_type == "mxf4xmxf4" else 8 return ( - _interleave_mega_moe_gate_up(l1_weights[0]), - _interleave_mega_moe_gate_up(l1_weights[1]), + _interleave_mega_moe_gate_up(l1_weights[0], gran=gran), + _interleave_mega_moe_gate_up(l1_weights[1], gran=gran), ) @@ -342,6 +356,7 @@ def build_mega_moe_experts_weights(experts) -> None: if getattr(experts, "_mega_moe_weights_built", False): return + mma_type = _mega_moe_mma_type() w13 = experts.w13_weight.data w13_sf_fp32 = experts.w13_weight_scale_inv.data w2 = experts.w2_weight.data @@ -375,7 +390,9 @@ def build_mega_moe_experts_weights(experts) -> None: # the deep-ep path consumes the non-transposed interleaved scale and a # swizzle-aware activation kernel. L2 weight is untouched by the mega # transform, so the existing `w2_weight.data` is shared directly. - w13_interleaved, w13_sf_interleaved = _interleave_mega_moe_l1_weights((w13, w13_sf)) + w13_interleaved, w13_sf_interleaved = _interleave_mega_moe_l1_weights( + (w13, w13_sf), mma_type + ) w13_sf_utccp = _transpose_mega_moe_sf_for_utccp(w13_sf_interleaved) w2_sf_utccp = _transpose_mega_moe_sf_for_utccp(w2_sf) diff --git a/python/sglang/srt/layers/quantization/mxfp4.py b/python/sglang/srt/layers/quantization/mxfp4.py index 369be8973..27dd4d7b4 100644 --- a/python/sglang/srt/layers/quantization/mxfp4.py +++ b/python/sglang/srt/layers/quantization/mxfp4.py @@ -666,9 +666,13 @@ class Mxfp4MoEMethod(FusedMoEMethodBase): # (keeping both layouts OOMs: 92 layers double the experts). from deep_gemm import transform_weights_for_mega_moe + from sglang.srt.layers.moe.mega_moe import _mega_moe_mma_type + + mma_type = _mega_moe_mma_type() l1_pair, l2_pair = transform_weights_for_mega_moe( (layer.w13_weight.data, layer.w13_weight_scale.data), (layer.w2_weight.data, layer.w2_weight_scale.data), + mma_type=mma_type, ) layer.mega_l1_weights = l1_pair layer.mega_l2_weights = l2_pair diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index 24134256b..fb2661b57 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -2389,8 +2389,8 @@ class ServerArgs: ] = "none" enable_w4a4_mxfp4_megamoe: A[ bool, - "Enable the W4A4 MXFP4 MegaMoE path by setting DeepGEMM's " - "DG_USE_FP4_ACTS=1 and DG_USE_MXF4_KIND=1. Use with " + "Enable the W4A4 MXFP4 MegaMoE path with DeepGEMM's " + "mxf4xmxf4 MMA type. Use with " "--moe-a2a-backend megamoe.", NS("exec.moe"), ] = False diff --git a/test/registered/unit/layers/moe/test_mega_moe_deepgemm_api.py b/test/registered/unit/layers/moe/test_mega_moe_deepgemm_api.py new file mode 100644 index 000000000..f06106530 --- /dev/null +++ b/test/registered/unit/layers/moe/test_mega_moe_deepgemm_api.py @@ -0,0 +1,262 @@ +"""Unit tests for the DeepGEMM MegaMoE interface.""" + +import sys +import unittest +from contextlib import nullcontext +from types import ModuleType, SimpleNamespace +from unittest.mock import MagicMock, patch + +import torch + +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase, maybe_stub_sgl_kernel + +maybe_stub_sgl_kernel() + +from sglang.srt.layers.moe import mega_moe + +register_cpu_ci(est_time=2, suite="base-a-test-cpu") + + +class TestDeepGemmMegaMoeApi(CustomTestCase): + def setUp(self): + super().setUp() + mega_moe._MEGA_MOE_SYMM_BUFFER.clear() + + def tearDown(self): + mega_moe._MEGA_MOE_SYMM_BUFFER.clear() + super().tearDown() + + def test_mxf4_buffer_uses_typed_api(self): + deep_gemm = ModuleType("deep_gemm") + expected_buffer = object() + deep_gemm.get_symm_buffer_for_mega_moe = MagicMock(return_value=expected_buffer) + group = object() + + with ( + patch.dict(sys.modules, {"deep_gemm": deep_gemm}), + patch.object( + mega_moe, + "_mega_moe_mma_type", + return_value="mxf4xmxf4", + create=True, + ), + ): + actual_buffer = mega_moe._get_mega_moe_symm_buffer( + group, + num_experts=8, + num_max_tokens_per_rank=64, + num_topk=2, + hidden=128, + intermediate_hidden=256, + ) + + self.assertIs(actual_buffer, expected_buffer) + call = deep_gemm.get_symm_buffer_for_mega_moe.call_args + self.assertEqual(call.kwargs.get("mma_type"), "mxf4xmxf4") + self.assertNotIn("use_fp8_dispatch", call.kwargs) + + def test_server_flag_selects_mxf4_mma_type(self): + for enabled, expected in ((False, "fp8xfp4"), (True, "mxf4xmxf4")): + with self.subTest(enabled=enabled): + config = SimpleNamespace( + moe=SimpleNamespace(enable_w4a4_mxfp4_megamoe=enabled) + ) + with patch.object(mega_moe, "get_exec", return_value=config): + self.assertEqual(mega_moe._mega_moe_mma_type(), expected) + + def test_buffer_cache_separates_mma_types(self): + deep_gemm = ModuleType("deep_gemm") + expected_buffers = (object(), object()) + deep_gemm.get_symm_buffer_for_mega_moe = MagicMock(side_effect=expected_buffers) + group = object() + + with ( + patch.dict(sys.modules, {"deep_gemm": deep_gemm}), + patch.object( + mega_moe, + "_mega_moe_mma_type", + side_effect=("fp8xfp4", "mxf4xmxf4"), + ), + ): + actual_buffers = tuple( + mega_moe._get_mega_moe_symm_buffer( + group, + num_experts=8, + num_max_tokens_per_rank=64, + num_topk=2, + hidden=128, + intermediate_hidden=256, + ) + for _ in range(2) + ) + + self.assertEqual(actual_buffers, expected_buffers) + self.assertEqual(deep_gemm.get_symm_buffer_for_mega_moe.call_count, 2) + self.assertEqual( + [ + call.kwargs["mma_type"] + for call in deep_gemm.get_symm_buffer_for_mega_moe.call_args_list + ], + ["fp8xfp4", "mxf4xmxf4"], + ) + + def test_mxf4_weight_transform_uses_matching_mma_type(self): + from sglang.srt.layers.quantization.mxfp4 import Mxfp4MoEMethod + + deep_gemm = ModuleType("deep_gemm") + deep_gemm.transform_sf_into_required_layout = MagicMock( + side_effect=lambda _sf, mn, k, recipe, num_groups, disable_ue8m0_cast: torch.zeros( + (num_groups, mn, max(1, k // 32)), dtype=torch.int32 + ) + ) + deep_gemm.transform_weights_for_mega_moe = MagicMock( + side_effect=lambda l1, l2, **_kwargs: (l1, l2) + ) + method = object.__new__(Mxfp4MoEMethod) + method.use_marlin = False + method.use_deep_gemm = False + method.use_mega_moe = True + layer = SimpleNamespace( + w13_weight=torch.nn.Parameter( + torch.zeros((1, 32, 16), dtype=torch.uint8), requires_grad=False + ), + w13_weight_scale=torch.nn.Parameter( + torch.zeros((1, 32, 1), dtype=torch.uint8), requires_grad=False + ), + w2_weight=torch.nn.Parameter( + torch.zeros((1, 32, 16), dtype=torch.uint8), requires_grad=False + ), + w2_weight_scale=torch.nn.Parameter( + torch.zeros((1, 32, 1), dtype=torch.uint8), requires_grad=False + ), + ) + + with ( + patch.dict(sys.modules, {"deep_gemm": deep_gemm}), + patch.object(mega_moe, "_mega_moe_mma_type", return_value="mxf4xmxf4"), + ): + method.process_weights_after_loading(layer) + + call = deep_gemm.transform_weights_for_mega_moe.call_args + self.assertEqual(call.kwargs.get("mma_type"), "mxf4xmxf4") + + def test_mxf4_pre_dispatch_uses_typed_api(self): + deep_gemm = ModuleType("deep_gemm") + deep_gemm.mega_moe_pre_dispatch = MagicMock() + 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(), + should_fuse_routed_scaling_factor_in_topk=True, + ) + topk_output = SimpleNamespace( + topk_ids=torch.tensor([[0, 1]]), + topk_weights=torch.tensor([[0.6, 0.4]]), + ) + moe = SimpleNamespace( + config=SimpleNamespace( + hidden_size=4, + num_experts_per_tok=2, + moe_intermediate_size=8, + swiglu_limit=None, + ), + experts=experts, + gate=MagicMock(return_value=torch.empty((1, 8))), + topk=MagicMock(return_value=topk_output), + is_hash=False, + num_fused_shared_experts=0, + layer_id=0, + routed_scaling_factor=1.0, + ) + + 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="mxf4xmxf4"), + patch.object( + mega_moe, + "_get_mega_moe_symm_buffer", + return_value=buffer, + ), + patch.object( + mega_moe, + "_configure_mega_moe_deep_gemm_num_sms", + return_value=nullcontext(), + ), + 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()), + ), + ): + mega_moe._run_mega_routed( + moe, + torch.zeros((1, 4)), + forward_batch=None, + input_ids_global=None, + num_tokens=1, + ) + + self.assertTrue(deep_gemm.mega_moe_pre_dispatch.called) + call = deep_gemm.mega_moe_pre_dispatch.call_args + self.assertEqual(call.kwargs.get("mma_type"), "mxf4xmxf4") + self.assertNotIn("use_fp4_acts", call.kwargs) + + def test_mxf4_l1_uses_packed_gate_up_interleave(self): + source = torch.arange(32).reshape(1, 32) + expected = torch.tensor( + [ + 0, + 2, + 4, + 6, + 8, + 10, + 12, + 14, + 16, + 18, + 20, + 22, + 24, + 26, + 28, + 30, + 1, + 3, + 5, + 7, + 9, + 11, + 13, + 15, + 17, + 19, + 21, + 23, + 25, + 27, + 29, + 31, + ] + ).reshape(1, 32) + + actual = mega_moe._interleave_mega_moe_gate_up(source, gran=16) + + torch.testing.assert_close(actual, expected) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/server_args/test_server_args.py b/test/registered/unit/server_args/test_server_args.py index 6272676b7..b81bdd2df 100644 --- a/test/registered/unit/server_args/test_server_args.py +++ b/test/registered/unit/server_args/test_server_args.py @@ -115,7 +115,7 @@ class TestPrepareServerArgs(CustomTestCase): # daemon to build the same static EPLB layout as the engine. handle_load_format(args) - def test_enable_w4a4_mxfp4_megamoe_sets_deepgemm_env(self): + def test_enable_w4a4_mxfp4_megamoe_preserves_legacy_deepgemm_env(self): deepgemm_env = { "DG_USE_FP4_ACTS": "0", "DG_USE_MXF4_KIND": "0", @@ -134,8 +134,8 @@ class TestPrepareServerArgs(CustomTestCase): args.resolve_once() self.assertTrue(resolution_result(args, "enable_w4a4_mxfp4_megamoe")) - self.assertEqual(os.environ["DG_USE_FP4_ACTS"], "1") - self.assertEqual(os.environ["DG_USE_MXF4_KIND"], "1") + self.assertEqual(os.environ["DG_USE_FP4_ACTS"], "0") + self.assertEqual(os.environ["DG_USE_MXF4_KIND"], "0") def test_w4a4_mxfp4_megamoe_disabled_preserves_deepgemm_env(self): deepgemm_env = {