Bump sgl-deep-gemm to 0.1.7 (#37279)
This commit is contained in:
+1
-1
@@ -6,7 +6,7 @@ ARG BUILD_TYPE=all
|
|||||||
ARG BRANCH_TYPE=remote
|
ARG BRANCH_TYPE=remote
|
||||||
ARG SGL_KERNEL_VERSION=0.4.6.post1
|
ARG SGL_KERNEL_VERSION=0.4.6.post1
|
||||||
ARG SGL_VERSION
|
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 USE_LATEST_SGLANG=0
|
||||||
ARG GDRCOPY_VERSION=2.5.1
|
ARG GDRCOPY_VERSION=2.5.1
|
||||||
ARG SGL_NCCL_VERSION=2.30.7
|
ARG SGL_NCCL_VERSION=2.30.7
|
||||||
|
|||||||
@@ -75,7 +75,7 @@ dependencies = [
|
|||||||
"sentencepiece",
|
"sentencepiece",
|
||||||
"setproctitle",
|
"setproctitle",
|
||||||
"sgl-deep-ep==0.1.2",
|
"sgl-deep-ep==0.1.2",
|
||||||
"sgl-deep-gemm==0.1.6",
|
"sgl-deep-gemm==0.1.7",
|
||||||
"sglang-kernel==0.4.6.post1",
|
"sglang-kernel==0.4.6.post1",
|
||||||
"smg-grpc-servicer>=0.5.0",
|
"smg-grpc-servicer>=0.5.0",
|
||||||
"soundfile==0.13.1",
|
"soundfile==0.13.1",
|
||||||
|
|||||||
@@ -1,7 +1,6 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import logging
|
import logging
|
||||||
import os
|
|
||||||
from typing import TYPE_CHECKING
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
@@ -17,7 +16,6 @@ logger = logging.getLogger(__name__)
|
|||||||
|
|
||||||
def handle_mega_moe(server_args: ServerArgs) -> None:
|
def handle_mega_moe(server_args: ServerArgs) -> None:
|
||||||
handle_moe_runner_backend_alias(server_args)
|
handle_moe_runner_backend_alias(server_args)
|
||||||
handle_w4a4_mxfp4_megamoe_env(server_args)
|
|
||||||
|
|
||||||
|
|
||||||
def handle_moe_runner_backend_alias(server_args: ServerArgs) -> None:
|
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_runner_backend="auto",
|
||||||
moe_a2a_backend="megamoe",
|
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"
|
|
||||||
|
|||||||
@@ -16,7 +16,6 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import functools
|
import functools
|
||||||
import os
|
|
||||||
from contextlib import contextmanager, nullcontext
|
from contextlib import contextmanager, nullcontext
|
||||||
from typing import TYPE_CHECKING, Optional
|
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.layers.moe.utils import get_moe_a2a_backend
|
||||||
from sglang.srt.model_executor.runner import get_is_capture_mode
|
from sglang.srt.model_executor.runner import get_is_capture_mode
|
||||||
from sglang.srt.models.deepseek_common.utils import _device_sm
|
from sglang.srt.models.deepseek_common.utils import _device_sm
|
||||||
|
from sglang.srt.runtime_context import get_exec
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from deep_gemm import SymmBuffer
|
from deep_gemm import SymmBuffer
|
||||||
@@ -45,6 +45,10 @@ if TYPE_CHECKING:
|
|||||||
_MEGA_MOE_SYMM_BUFFER: dict = {}
|
_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)
|
@functools.lru_cache(maxsize=1)
|
||||||
def _mega_moe_max_num_sms() -> Optional[int]:
|
def _mega_moe_max_num_sms() -> Optional[int]:
|
||||||
if _device_sm < 100:
|
if _device_sm < 100:
|
||||||
@@ -92,6 +96,7 @@ def _get_mega_moe_symm_buffer(
|
|||||||
) -> SymmBuffer:
|
) -> SymmBuffer:
|
||||||
import deep_gemm
|
import deep_gemm
|
||||||
|
|
||||||
|
mma_type = _mega_moe_mma_type()
|
||||||
key = (
|
key = (
|
||||||
id(group),
|
id(group),
|
||||||
num_max_tokens_per_rank,
|
num_max_tokens_per_rank,
|
||||||
@@ -99,6 +104,7 @@ def _get_mega_moe_symm_buffer(
|
|||||||
num_topk,
|
num_topk,
|
||||||
hidden,
|
hidden,
|
||||||
intermediate_hidden,
|
intermediate_hidden,
|
||||||
|
mma_type,
|
||||||
)
|
)
|
||||||
buf = _MEGA_MOE_SYMM_BUFFER.get(key)
|
buf = _MEGA_MOE_SYMM_BUFFER.get(key)
|
||||||
if buf is None:
|
if buf is None:
|
||||||
@@ -109,7 +115,7 @@ def _get_mega_moe_symm_buffer(
|
|||||||
num_topk,
|
num_topk,
|
||||||
hidden,
|
hidden,
|
||||||
intermediate_hidden,
|
intermediate_hidden,
|
||||||
use_fp8_dispatch=True,
|
mma_type=mma_type,
|
||||||
activation="swiglu",
|
activation="swiglu",
|
||||||
)
|
)
|
||||||
_MEGA_MOE_SYMM_BUFFER[key] = buf
|
_MEGA_MOE_SYMM_BUFFER[key] = buf
|
||||||
@@ -248,8 +254,8 @@ def _run_mega_routed(
|
|||||||
num_tokens,
|
num_tokens,
|
||||||
)
|
)
|
||||||
|
|
||||||
use_fp4_acts = os.getenv("DG_USE_FP4_ACTS") == "1"
|
mma_type = _mega_moe_mma_type()
|
||||||
if use_fp4_acts:
|
if mma_type == "mxf4xmxf4":
|
||||||
# FP4 path goes through DeepGEMM's mega_moe_pre_dispatch which
|
# FP4 path goes through DeepGEMM's mega_moe_pre_dispatch which
|
||||||
# handles the E2M1 packing variant. The jit implementation
|
# handles the E2M1 packing variant. The jit implementation
|
||||||
# only emits FP8.
|
# only emits FP8.
|
||||||
@@ -263,7 +269,7 @@ def _run_mega_routed(
|
|||||||
buf.topk_weights,
|
buf.topk_weights,
|
||||||
num_tokens=num_tokens,
|
num_tokens=num_tokens,
|
||||||
group_size=32,
|
group_size=32,
|
||||||
use_fp4_acts=True,
|
mma_type=mma_type,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
mega_moe_pre_dispatch(
|
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:
|
def _interleave_mega_moe_gate_up(t: torch.Tensor, gran: int = 8) -> torch.Tensor:
|
||||||
# Match DeepGEMM's L1 gate/up layout:
|
# Match DeepGEMM's L1 gate/up layouts. FP8 activations use contiguous
|
||||||
# [gate: 0..7, up: 0..7, gate: 8..15, up: 8..15, ...].
|
# gran-8 chunks; packed MXFP4 activations use even/odd gran-16 chunks.
|
||||||
num_groups, n, *rest = t.shape
|
num_groups, n, *rest = t.shape
|
||||||
half = n // 2
|
half = n // 2
|
||||||
gate = t[:, :half].reshape(num_groups, half // gran, gran, *rest)
|
gate = t[:, :half].reshape(num_groups, half // gran, gran, *rest)
|
||||||
up = t[:, half:].reshape(num_groups, half // gran, gran, *rest)
|
up = t[:, half:].reshape(num_groups, half // gran, gran, *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)
|
result = torch.stack([gate, up], dim=2).reshape(num_groups, n, *rest)
|
||||||
return torch.empty_like(t).copy_(result)
|
return torch.empty_like(t).copy_(result)
|
||||||
|
|
||||||
|
|
||||||
def _interleave_mega_moe_l1_weights(
|
def _interleave_mega_moe_l1_weights(
|
||||||
l1_weights: tuple[torch.Tensor, torch.Tensor],
|
l1_weights: tuple[torch.Tensor, torch.Tensor],
|
||||||
|
mma_type: str,
|
||||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||||
|
gran = 16 if mma_type == "mxf4xmxf4" else 8
|
||||||
return (
|
return (
|
||||||
_interleave_mega_moe_gate_up(l1_weights[0]),
|
_interleave_mega_moe_gate_up(l1_weights[0], gran=gran),
|
||||||
_interleave_mega_moe_gate_up(l1_weights[1]),
|
_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):
|
if getattr(experts, "_mega_moe_weights_built", False):
|
||||||
return
|
return
|
||||||
|
|
||||||
|
mma_type = _mega_moe_mma_type()
|
||||||
w13 = experts.w13_weight.data
|
w13 = experts.w13_weight.data
|
||||||
w13_sf_fp32 = experts.w13_weight_scale_inv.data
|
w13_sf_fp32 = experts.w13_weight_scale_inv.data
|
||||||
w2 = experts.w2_weight.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
|
# the deep-ep path consumes the non-transposed interleaved scale and a
|
||||||
# swizzle-aware activation kernel. L2 weight is untouched by the mega
|
# swizzle-aware activation kernel. L2 weight is untouched by the mega
|
||||||
# transform, so the existing `w2_weight.data` is shared directly.
|
# 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)
|
w13_sf_utccp = _transpose_mega_moe_sf_for_utccp(w13_sf_interleaved)
|
||||||
w2_sf_utccp = _transpose_mega_moe_sf_for_utccp(w2_sf)
|
w2_sf_utccp = _transpose_mega_moe_sf_for_utccp(w2_sf)
|
||||||
|
|
||||||
|
|||||||
@@ -666,9 +666,13 @@ class Mxfp4MoEMethod(FusedMoEMethodBase):
|
|||||||
# (keeping both layouts OOMs: 92 layers double the experts).
|
# (keeping both layouts OOMs: 92 layers double the experts).
|
||||||
from deep_gemm import transform_weights_for_mega_moe
|
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(
|
l1_pair, l2_pair = transform_weights_for_mega_moe(
|
||||||
(layer.w13_weight.data, layer.w13_weight_scale.data),
|
(layer.w13_weight.data, layer.w13_weight_scale.data),
|
||||||
(layer.w2_weight.data, layer.w2_weight_scale.data),
|
(layer.w2_weight.data, layer.w2_weight_scale.data),
|
||||||
|
mma_type=mma_type,
|
||||||
)
|
)
|
||||||
layer.mega_l1_weights = l1_pair
|
layer.mega_l1_weights = l1_pair
|
||||||
layer.mega_l2_weights = l2_pair
|
layer.mega_l2_weights = l2_pair
|
||||||
|
|||||||
@@ -2389,8 +2389,8 @@ class ServerArgs:
|
|||||||
] = "none"
|
] = "none"
|
||||||
enable_w4a4_mxfp4_megamoe: A[
|
enable_w4a4_mxfp4_megamoe: A[
|
||||||
bool,
|
bool,
|
||||||
"Enable the W4A4 MXFP4 MegaMoE path by setting DeepGEMM's "
|
"Enable the W4A4 MXFP4 MegaMoE path with DeepGEMM's "
|
||||||
"DG_USE_FP4_ACTS=1 and DG_USE_MXF4_KIND=1. Use with "
|
"mxf4xmxf4 MMA type. Use with "
|
||||||
"--moe-a2a-backend megamoe.",
|
"--moe-a2a-backend megamoe.",
|
||||||
NS("exec.moe"),
|
NS("exec.moe"),
|
||||||
] = False
|
] = False
|
||||||
|
|||||||
@@ -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()
|
||||||
@@ -115,7 +115,7 @@ class TestPrepareServerArgs(CustomTestCase):
|
|||||||
# daemon to build the same static EPLB layout as the engine.
|
# daemon to build the same static EPLB layout as the engine.
|
||||||
handle_load_format(args)
|
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 = {
|
deepgemm_env = {
|
||||||
"DG_USE_FP4_ACTS": "0",
|
"DG_USE_FP4_ACTS": "0",
|
||||||
"DG_USE_MXF4_KIND": "0",
|
"DG_USE_MXF4_KIND": "0",
|
||||||
@@ -134,8 +134,8 @@ class TestPrepareServerArgs(CustomTestCase):
|
|||||||
args.resolve_once()
|
args.resolve_once()
|
||||||
|
|
||||||
self.assertTrue(resolution_result(args, "enable_w4a4_mxfp4_megamoe"))
|
self.assertTrue(resolution_result(args, "enable_w4a4_mxfp4_megamoe"))
|
||||||
self.assertEqual(os.environ["DG_USE_FP4_ACTS"], "1")
|
self.assertEqual(os.environ["DG_USE_FP4_ACTS"], "0")
|
||||||
self.assertEqual(os.environ["DG_USE_MXF4_KIND"], "1")
|
self.assertEqual(os.environ["DG_USE_MXF4_KIND"], "0")
|
||||||
|
|
||||||
def test_w4a4_mxfp4_megamoe_disabled_preserves_deepgemm_env(self):
|
def test_w4a4_mxfp4_megamoe_disabled_preserves_deepgemm_env(self):
|
||||||
deepgemm_env = {
|
deepgemm_env = {
|
||||||
|
|||||||
Reference in New Issue
Block a user