[MoE Refactor] Migrate SM90 Cutlass W4A16 to MoeRunner (#26489)

Co-authored-by: luoyuan.luo <luoyuan.luo@antgroup.com>
This commit is contained in:
Yuan Luo
2026-05-30 02:02:56 -07:00
committed by GitHub
co-authored by luoyuan.luo
parent c93a559e5a
commit 0d9a2a9de3
5 changed files with 281 additions and 116 deletions
@@ -160,9 +160,34 @@ def _build_method(num_experts, hidden, inter):
method._padded_hidden = _round_up(hidden, 128)
method._padded_intermediate = _round_up(inter, 128)
method.use_flashinfer = True
method.runner = _build_flashinfer_mxfp4_runner(num_experts, hidden, inter)
return method
def _build_flashinfer_mxfp4_runner(num_experts, hidden, inter):
"""Construct a real MoeRunner bound to the flashinfer_mxfp4 fused func.
Bypasses ``create_moe_runner`` (which needs a live server arg context)
and wires the runner with a minimal MoeRunnerConfig sufficient for the
cutlass SM90 fused func, which only reads dispatch_output / quant_info.
"""
import sglang.srt.layers.moe.moe_runner.flashinfer_mxfp4 # noqa: F401
from sglang.srt.layers.moe.moe_runner.base import MoeRunnerConfig
from sglang.srt.layers.moe.moe_runner.runner import MoeRunner
from sglang.srt.layers.moe.utils import MoeRunnerBackend
cfg = MoeRunnerConfig(
num_experts=num_experts,
num_local_experts=num_experts,
hidden_size=hidden,
intermediate_size_per_partition=inter,
top_k=None,
activation="silu",
is_gated=True,
)
return MoeRunner(MoeRunnerBackend.FLASHINFER_MXFP4, cfg)
def _expected_w13_processed(w13_un, w13_s_un, w13_b_un, N_pad, K_pad, group_size):
"""Replicate ``_process_weights_for_sm90_cutlass`` for w13: de-interleave
HF's pair-wise ``[g_0, u_0, g_1, u_1, ...]`` layout into halved
@@ -300,14 +325,21 @@ def test_apply_sm90_cutlass_matches_flashinfer_direct(
covered separately by ``test_process_weights_matches_direct_interleave``;
here we just verify that ``apply`` calls the kernel with the right
arguments (incl. input padding + output trim)."""
import sglang.srt.layers.moe.moe_runner.flashinfer_mxfp4 as fi_mxfp4_mod
import sglang.srt.layers.quantization.mxfp4 as mxfp4_mod
# Bypass symmetric-memory / TP-group: not relevant to numerics.
# Bypass symmetric-memory / TP-group in both the legacy quant_method and
# the new fused-func module (where the kernel call now lives).
monkeypatch.setattr(
mxfp4_mod, "use_symmetric_memory", lambda *a, **kw: nullcontext()
)
monkeypatch.setattr(mxfp4_mod, "is_allocation_symmetric", lambda: False)
monkeypatch.setattr(mxfp4_mod, "get_tp_group", lambda: None)
monkeypatch.setattr(
fi_mxfp4_mod, "use_symmetric_memory", lambda *a, **kw: nullcontext()
)
monkeypatch.setattr(fi_mxfp4_mod, "is_allocation_symmetric", lambda: False)
monkeypatch.setattr(fi_mxfp4_mod, "get_tp_group", lambda: None)
w13, w2, w13_s, w2_s, w13_b, w2_b = _make_random_mxfp4(num_experts, hidden, inter)
x = torch.randn(tokens, hidden, dtype=torch.bfloat16, device="cuda") * 0.1
@@ -321,7 +353,7 @@ def test_apply_sm90_cutlass_matches_flashinfer_direct(
method._process_weights_for_sm90_cutlass(layer)
out_sglang = method._apply_sm90_cutlass(
layer, x.clone(), _MockTopKOutput(topk_w, topk_i)
layer, _MockDispatchOutput(x.clone(), topk_w, topk_i)
).hidden_states
# ---- FlashInfer-direct reference using the same processed weights ----
@@ -431,13 +463,17 @@ def test_dsv4_apply_matches_flashinfer_direct(
the equivalent reorder + scale-cast + interleave applied manually."""
from types import SimpleNamespace
import sglang.srt.layers.moe.moe_runner.flashinfer_mxfp4 as fi_mxfp4_mod
import sglang.srt.layers.quantization.mxfp4_flashinfer_cutlass_moe as ds_mod
from sglang.srt.layers.quantization.utils import reorder_w1w3_to_w3w1
# Bypass symmetric-memory / TP-group stack -- not relevant to numerics.
monkeypatch.setattr(ds_mod, "use_symmetric_memory", lambda *a, **kw: nullcontext())
monkeypatch.setattr(ds_mod, "is_allocation_symmetric", lambda: False)
monkeypatch.setattr(ds_mod, "get_tp_group", lambda: None)
# Bypass symmetric-memory / TP-group stack in the new fused-func module
# (where DSv4 ``apply`` now dispatches the kernel call through).
monkeypatch.setattr(
fi_mxfp4_mod, "use_symmetric_memory", lambda *a, **kw: nullcontext()
)
monkeypatch.setattr(fi_mxfp4_mod, "is_allocation_symmetric", lambda: False)
monkeypatch.setattr(fi_mxfp4_mod, "get_tp_group", lambda: None)
w13, w2, w13_s, w2_s = _make_random_dsv4_mxfp4(num_experts, hidden, inter)
x = torch.randn(tokens, hidden, dtype=torch.bfloat16, device="cuda") * 0.1
@@ -455,6 +491,9 @@ def test_dsv4_apply_matches_flashinfer_direct(
method._swiglu_alpha_tensor = None
method._swiglu_beta_tensor = None
method._swiglu_limit_tensor = None
# Wire the unified MoeRunner -> flashinfer_mxfp4 fused func that
# ``apply`` now dispatches through.
method.runner = _build_flashinfer_mxfp4_runner(num_experts, hidden, inter)
layer = _MockLayer()
layer.w13_weight = torch.nn.Parameter(w13.clone(), requires_grad=False)