[MoE] Add extension points for custom runner backends (#32665)

Co-authored-by: Alex Nails <alex.nails@radixark.ai>
Co-authored-by: Claude Opus 5 (1M context) <noreply@anthropic.com>
This commit is contained in:
Kurt Shuster
2026-08-29 19:03:41 -07:00
committed by GitHub
co-authored by Alex Nails Claude Opus 5
parent ca8ff035c3
commit ed39568e79
12 changed files with 613 additions and 104 deletions
@@ -283,6 +283,7 @@ class FusedMoE(torch.nn.Module):
params_dtype: Data type for the parameters.
reduce_results: Whether to apply all_reduce on the output of the layer
quant_config: Quantization configuration.
quant_method: Explicit quant method, overriding selection from quant_config.
inplace: suggestion to compute inplace (modify input activation).
enable_qwen35_fp8_deferred_finalize: Whether this concrete Qwen3.5
layer may expose FlashInfer's block-FP8 deferred MoE output.
@@ -325,6 +326,7 @@ class FusedMoE(torch.nn.Module):
is_gated: bool = True,
gate_up_interleaved: bool = True,
enable_qwen35_fp8_deferred_finalize: bool = False,
quant_method: Optional[FusedMoEMethodBase] = None,
):
super().__init__()
if params_dtype is None:
@@ -430,17 +432,19 @@ class FusedMoE(torch.nn.Module):
gate_up_interleaved=gate_up_interleaved,
)
self.quant_method: Optional[FusedMoEMethodBase] = None
self.quant_method = quant_method
server_args = get_server_args()
kt_config = create_kt_config_from_server_args(server_args, layer_id)
if kt_config is not None:
if quant_config is not None:
if self.quant_method is not None:
gpu_method = self.quant_method
elif quant_config is not None:
gpu_method = quant_config.get_quant_method(self, prefix)
else:
gpu_method = UnquantizedFusedMoEMethod(self.use_triton_kernels)
self.quant_method = KTEPWrapperMethod(gpu_method, kt_config)
else:
if quant_config is not None:
if self.quant_method is None and quant_config is not None:
self.quant_method = quant_config.get_quant_method(self, prefix)
if self.quant_method is None:
self.quant_method = UnquantizedFusedMoEMethod(
@@ -525,8 +529,7 @@ class FusedMoE(torch.nn.Module):
self._dwdp_bound = False
if self.quant_method is not None and hasattr(self.quant_method, "runner"):
self.runner = self.quant_method.runner
self.runner = self.quant_method.runner
@property
def num_global_routed_experts(self) -> int:
@@ -1685,7 +1688,7 @@ class FusedMoE(torch.nn.Module):
def set_overlap_args(
self, down_gemm_overlap_args: DownGemmOverlapArgs, meta_overlap_args: dict
):
if hasattr(self, "runner"):
if self.runner is not None:
self.runner.set_overlap_args(down_gemm_overlap_args, meta_overlap_args)
else:
# TODO: remove this branch after MoE refactor
@@ -1693,7 +1696,7 @@ class FusedMoE(torch.nn.Module):
self.meta_overlap_args = meta_overlap_args
def clear_overlap_args(self) -> None:
if hasattr(self, "runner"):
if self.runner is not None:
self.runner.clear_overlap_args()
else:
# TODO: remove this branch after MoE refactor
@@ -1,4 +1,4 @@
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.moe_runner.runner import MoeRunner, register_moe_runner_core
__all__ = ["MoeRunnerConfig", "MoeRunner"]
__all__ = ["MoeRunnerConfig", "MoeRunner", "register_moe_runner_core"]
@@ -9,6 +9,7 @@ import torch
from sglang.srt.layers.moe.utils import (
MoeA2ABackend,
MoeRunnerBackend,
MoeRunnerBackendLike,
RoutingMethodType,
)
@@ -113,6 +114,26 @@ class MoeRunnerCore(ABC):
return self.runner_backend == MoeRunnerBackend.TRITON
class DispatchMoeRunnerCore(ABC):
"""Runner core that consumes the standard dispatch representation directly."""
def __init__(self, config: MoeRunnerConfig):
self.config = config
@property
@abstractmethod
def runner_backend(self) -> MoeRunnerBackendLike: ...
@abstractmethod
def run_from_dispatch(
self,
dispatch_output: DispatchOutput,
quant_info: MoeQuantInfo,
runner_config: MoeRunnerConfig,
hooks: Any = None,
) -> CombineInput: ...
class FusedOpPool:
_fused_funcs: dict[str, Callable] = {}
@@ -2,32 +2,58 @@ from __future__ import annotations
import logging
import os
from typing import TYPE_CHECKING, Any, Optional
from collections.abc import Callable
from typing import TYPE_CHECKING, Optional
from sglang.srt.layers.moe.moe_runner.base import (
DispatchMoeRunnerCore,
FusedOpPool,
MoeRunnerConfig,
PermuteMethodPool,
)
from sglang.srt.layers.moe.moe_runner.deep_gemm import DeepGemmRunnerCore
from sglang.srt.layers.moe.moe_runner.triton import TritonRunnerCore
from sglang.srt.layers.moe.moe_runner.triton import TritonRunnerCore, TritonRunnerInput
from sglang.srt.layers.moe.moe_runner.triton_kernels import TritonKernelsRunnerCore
from sglang.srt.layers.moe.utils import get_moe_a2a_backend, get_moe_runner_backend
from sglang.srt.layers.moe.utils import (
MoeRunnerBackendLike,
get_moe_a2a_backend,
get_moe_runner_backend,
register_moe_runner_backend_name,
resolve_moe_runner_backend,
)
if TYPE_CHECKING:
from sglang.srt.batch_overlap.single_batch_overlap import DownGemmOverlapArgs
from sglang.srt.layers.moe.moe_runner.base import MoeQuantInfo
from sglang.srt.layers.moe.token_dispatcher.base import CombineInput, DispatchOutput
from sglang.srt.layers.moe.utils import MoeRunnerBackend
from sglang.srt.lora.lora_moe_runners import LoRAHooks
logger = logging.getLogger(__name__)
_CUSTOM_RUNNER_CORE_FACTORIES: dict[
str, Callable[[MoeRunnerConfig], DispatchMoeRunnerCore]
] = {}
def register_moe_runner_core(
backend_name: str,
factory: Callable[[MoeRunnerConfig], DispatchMoeRunnerCore],
) -> None:
"""Register a runner-core factory for a new or built-in backend name."""
if backend_name in _CUSTOM_RUNNER_CORE_FACTORIES:
raise ValueError(f"Runner core for {backend_name!r} is already registered")
try:
resolve_moe_runner_backend(backend_name)
except ValueError:
register_moe_runner_backend_name(backend_name)
_CUSTOM_RUNNER_CORE_FACTORIES[backend_name] = factory
class MoeRunner:
def __init__(
self,
runner_backend: MoeRunnerBackend,
runner_backend: MoeRunnerBackendLike,
config: MoeRunnerConfig,
lora_enabled: bool = False,
):
@@ -61,7 +87,9 @@ class MoeRunner:
self.fused_func = None
if runner_backend.is_triton():
if custom_factory := _CUSTOM_RUNNER_CORE_FACTORIES.get(runner_backend.value):
self.runner_core = custom_factory(config)
elif runner_backend.is_triton():
self.runner_core = TritonRunnerCore(config)
elif runner_backend.is_ascend():
from sglang.srt.layers.moe.moe_runner.ascend import AscendRunnerCore
@@ -157,7 +185,16 @@ class MoeRunner:
assert self.runner_core is not None
def _maybe_build_lora_hooks(_runner_input: Any) -> LoRAHooks:
def _maybe_build_lora_hooks(
_runner_input: DispatchOutput | TritonRunnerInput,
) -> Optional[LoRAHooks]:
# Bail out before touching the runner input: LoRA is only wired up
# for the Triton runner, so every other backend (deep_gemm,
# triton_kernels, aiter, ascend, ...) gets here with LoRA disabled
# and its runner input carries no topk_ids to read.
if not self.lora_enabled or lora_info is None:
return None
from sglang.srt.layers.moe.token_dispatcher.base import DispatchOutput
from sglang.srt.lora.lora_moe_runners import build_lora_hooks
@@ -167,19 +204,18 @@ class MoeRunner:
_runner_input.topk_output.topk_ids,
)
else:
assert isinstance(_runner_input, TritonRunnerInput), type(_runner_input)
hidden_states = _runner_input.hidden_states
topk_ids = getattr(_runner_input, "topk_ids", None)
if self.lora_enabled and lora_info is not None:
return build_lora_hooks(
hidden_states,
lora_info,
topk_ids,
)
return None
topk_ids = _runner_input.topk_ids
return build_lora_hooks(
hidden_states,
lora_info,
topk_ids,
)
# Runners that handle dispatch_output directly (e.g., MarlinRunnerCore)
# bypass the pre-permute step and do their own alignment internally.
if hasattr(self.runner_core, "run_from_dispatch"):
if isinstance(self.runner_core, DispatchMoeRunnerCore):
hooks = _maybe_build_lora_hooks(dispatch_output)
return self.runner_core.run_from_dispatch(
dispatch_output, quant_info, self.config, hooks=hooks
+102 -56
View File
@@ -3,6 +3,7 @@ from __future__ import annotations
import logging
import os
from contextlib import contextmanager
from dataclasses import dataclass
from enum import Enum, IntEnum
from typing import NamedTuple
@@ -100,7 +101,73 @@ class MoeA2ABackend(Enum):
)
class MoeRunnerBackend(Enum):
class _MoeRunnerBackendPredicates:
value: str
def is_auto(self):
return self.value == MoeRunnerBackend.AUTO.value
def is_hpc_ops(self):
return self.value == MoeRunnerBackend.HPC_OPS.value
def is_deep_gemm(self):
return self.value == MoeRunnerBackend.DEEP_GEMM.value
def is_triton(self):
return self.value == MoeRunnerBackend.TRITON.value
def is_ascend(self):
return self.value == MoeRunnerBackend.ASCEND.value
def is_triton_kernels(self):
return self.value == MoeRunnerBackend.TRITON_KERNELS.value
def is_flashinfer_trtllm(self):
# experimental_sgl_trtllm shares the TRT-LLM FP8 kernels + layout, so it inherits
# trtllm weight-prep here; divergent sites check is_experimental_sgl_trtllm() first.
return self.value in (
MoeRunnerBackend.FLASHINFER_TRTLLM.value,
MoeRunnerBackend.EXPERIMENTAL_SGL_TRTLLM.value,
)
def is_experimental_sgl_trtllm(self):
return self.value == MoeRunnerBackend.EXPERIMENTAL_SGL_TRTLLM.value
def is_flashinfer_trtllm_routed(self):
return self.value == MoeRunnerBackend.FLASHINFER_TRTLLM_ROUTED.value
def is_flashinfer_cutlass(self):
return self.value == MoeRunnerBackend.FLASHINFER_CUTLASS.value
def is_flashinfer_cutedsl(self):
return self.value == MoeRunnerBackend.FLASHINFER_CUTEDSL.value
def is_flashinfer_mxfp4(self):
return self.value == MoeRunnerBackend.FLASHINFER_MXFP4.value
def is_cutlass(self):
return self.value == MoeRunnerBackend.CUTLASS.value
def is_marlin(self):
# experimental_sgl_marlin shares the marlin weight repack, quant-method
# selection, and base fused path; divergent sites (the LoRA MoE dispatch)
# check is_experimental_sgl_marlin() first.
return self.value in (
MoeRunnerBackend.MARLIN.value,
MoeRunnerBackend.EXPERIMENTAL_SGL_MARLIN.value,
)
def is_experimental_sgl_marlin(self):
return self.value == MoeRunnerBackend.EXPERIMENTAL_SGL_MARLIN.value
def is_humming(self):
return self.value == MoeRunnerBackend.HUMMING.value
def is_aiter(self):
return self.value == MoeRunnerBackend.AITER.value
class MoeRunnerBackend(_MoeRunnerBackendPredicates, Enum):
AUTO = "auto"
DEEP_GEMM = "deep_gemm"
@@ -121,67 +188,46 @@ class MoeRunnerBackend(Enum):
HPC_OPS = "hpc_ops"
INTEL_XPU = "intel_xpu"
def is_auto(self):
return self == MoeRunnerBackend.AUTO
def is_hpc_ops(self):
return self == MoeRunnerBackend.HPC_OPS
@dataclass(frozen=True)
class RegisteredMoeRunnerBackend(_MoeRunnerBackendPredicates):
"""Identifier for an MoE runner backend supplied by an extension."""
def is_deep_gemm(self):
return self == MoeRunnerBackend.DEEP_GEMM
value: str
def is_triton(self):
return self == MoeRunnerBackend.TRITON
def is_ascend(self):
return self == MoeRunnerBackend.ASCEND
MoeRunnerBackendLike = MoeRunnerBackend | RegisteredMoeRunnerBackend
_REGISTERED_MOE_RUNNER_BACKEND_NAMES: set[str] = set()
def is_triton_kernels(self):
return self == MoeRunnerBackend.TRITON_KERNELS
def is_flashinfer_trtllm(self):
# experimental_sgl_trtllm shares the TRT-LLM FP8 kernels + layout, so it inherits
# trtllm weight-prep here; divergent sites check is_experimental_sgl_trtllm() first.
return self in (
MoeRunnerBackend.FLASHINFER_TRTLLM,
MoeRunnerBackend.EXPERIMENTAL_SGL_TRTLLM,
)
def register_moe_runner_backend_name(name: str) -> None:
"""Register a backend name supplied by an out-of-tree extension."""
def is_experimental_sgl_trtllm(self):
return self == MoeRunnerBackend.EXPERIMENTAL_SGL_TRTLLM
if not name:
raise ValueError("MoE runner backend name must not be empty")
try:
MoeRunnerBackend(name)
except ValueError:
_REGISTERED_MOE_RUNNER_BACKEND_NAMES.add(name)
else:
raise ValueError(f"MoE runner backend {name!r} is already built in")
def is_flashinfer_trtllm_routed(self):
return self == MoeRunnerBackend.FLASHINFER_TRTLLM_ROUTED
def is_flashinfer_cutlass(self):
return self == MoeRunnerBackend.FLASHINFER_CUTLASS
def resolve_moe_runner_backend(
backend: str | MoeRunnerBackendLike,
) -> MoeRunnerBackendLike:
"""Resolve a built-in or registered backend identifier."""
def is_flashinfer_cutedsl(self):
return self == MoeRunnerBackend.FLASHINFER_CUTEDSL
def is_flashinfer_mxfp4(self):
return self == MoeRunnerBackend.FLASHINFER_MXFP4
def is_cutlass(self):
return self == MoeRunnerBackend.CUTLASS
def is_marlin(self):
# experimental_sgl_marlin shares the marlin weight repack, quant-method
# selection, and base fused path; divergent sites (the LoRA MoE dispatch)
# check is_experimental_sgl_marlin() first.
return self in (
MoeRunnerBackend.MARLIN,
MoeRunnerBackend.EXPERIMENTAL_SGL_MARLIN,
)
def is_experimental_sgl_marlin(self):
return self == MoeRunnerBackend.EXPERIMENTAL_SGL_MARLIN
def is_humming(self):
return self == MoeRunnerBackend.HUMMING
def is_aiter(self):
return self == MoeRunnerBackend.AITER
if isinstance(backend, (MoeRunnerBackend, RegisteredMoeRunnerBackend)):
return backend
try:
return MoeRunnerBackend(backend)
except ValueError:
if backend in _REGISTERED_MOE_RUNNER_BACKEND_NAMES:
return RegisteredMoeRunnerBackend(backend)
raise ValueError(
f"MoE runner backend {backend!r} is neither built in nor registered"
) from None
def is_intel_xpu(self):
return self == MoeRunnerBackend.INTEL_XPU
@@ -353,9 +399,9 @@ def initialize_moe_config():
spec = get_spec()
moe = get_flags().moe
moe.a2a_backend = MoeA2ABackend(exec_moe.moe_a2a_backend)
moe.runner_backend = MoeRunnerBackend(exec_moe.moe_runner_backend)
moe.runner_backend = resolve_moe_runner_backend(exec_moe.moe_runner_backend)
moe.speculative_runner_backend = (
MoeRunnerBackend(spec.speculative_moe_runner_backend)
resolve_moe_runner_backend(spec.speculative_moe_runner_backend)
if spec.speculative_moe_runner_backend is not None
else moe.runner_backend
)
@@ -391,14 +437,14 @@ def get_moe_a2a_backend() -> MoeA2ABackend:
return moe.a2a_backend
def get_moe_runner_backend() -> MoeRunnerBackend:
def get_moe_runner_backend() -> MoeRunnerBackendLike:
moe = get_flags().moe
if moe.runner_backend is None:
moe.runner_backend = MoeRunnerBackend.AUTO
return moe.runner_backend
def get_speculative_moe_runner_backend() -> MoeRunnerBackend:
def get_speculative_moe_runner_backend() -> MoeRunnerBackendLike:
moe = get_flags().moe
if moe.speculative_runner_backend is None:
logger.warning(
@@ -13,9 +13,11 @@ from torch import nn
from sglang.srt.layers.modelopt_utils import canonicalize_modelopt_quant_algo
if TYPE_CHECKING:
from sglang.srt.layers.moe.moe_runner import MoeRunnerConfig
from sglang.srt.layers.moe.moe_runner import MoeRunner, MoeRunnerConfig
from sglang.srt.layers.moe.moe_runner.base import MoeQuantInfo
from sglang.srt.layers.moe.moe_runner.triton import TritonMoeQuantInfo
from sglang.srt.layers.moe.token_dispatcher import CombineInput, DispatchOutput
from sglang.srt.layers.moe.utils import MoeRunnerBackendLike
from sglang.srt.models.utils import WeightsMapper
@@ -86,6 +88,7 @@ class LinearMethodBase(QuantizeMethodBase):
class FusedMoEMethodBase(QuantizeMethodBase):
runner: MoeRunner | None = None
def create_weights(
self,
@@ -124,6 +127,15 @@ class FusedMoEMethodBase(QuantizeMethodBase):
f"{type(self).__name__} must implement get_triton_quant_info()"
)
def get_moe_quant_info(
self, layer: torch.nn.Module, runner_backend: MoeRunnerBackendLike
) -> MoeQuantInfo:
if runner_backend.is_triton():
return self.get_triton_quant_info(layer)
raise NotImplementedError(
f"{type(self).__name__} does not expose quant info for {runner_backend.value!r}"
)
class QuantizationConfig(ABC):
"""Base class for quantization configs."""
@@ -57,6 +57,10 @@ class Mxfp4FlashinferTrtllmMoEMethod:
def create_moe_runner(self, layer, moe_runner_config):
self.moe_runner_config = moe_runner_config
# Applies flashinfer trtllm directly instead of going through a
# MoeRunner; FusedMoE still reads `.runner`, and this class is not a
# FusedMoEMethodBase subclass so it inherits no default.
self.runner = None
swiglu_limit = moe_runner_config.swiglu_limit
self._gemm1_clamp_limit_tensor = (
+14 -14
View File
@@ -970,13 +970,14 @@ class FusedMoEWithLoRA(BaseLayerWithLoRA):
base_layer.should_fuse_routed_scaling_factor_in_topk
)
self.tp_size = getattr(base_layer, "moe_tp_size", 1)
self.tp_rank = getattr(base_layer, "moe_tp_rank", 0)
self.intermediate_size_per_partition = getattr(
base_layer, "intermediate_size_per_partition", None
self.tp_size = base_layer.moe_tp_size
self.tp_rank = base_layer.moe_tp_rank
self.intermediate_size_per_partition = (
base_layer.intermediate_size_per_partition
)
# Stock MoE LoRA buffers are split gate/up except for GPT-OSS-style weights.
self._uses_interleaved_gate_up = (
getattr(base_layer.moe_runner_config, "gemm1_alpha", None) is not None
base_layer.moe_runner_config.gemm1_alpha is not None
)
# Initialize triton_lora moe runner for batches with lora enabled
@@ -984,15 +985,13 @@ class FusedMoEWithLoRA(BaseLayerWithLoRA):
from sglang.srt.layers.moe.moe_runner.runner import MoeRunner
from sglang.srt.layers.moe.utils import get_moe_runner_backend
# Determine runner backend: prefer server arg, fall back to quant method's runner
# Use the runner selected by the quant method so per-format backend resolution
# stays identical between base and LoRA forwards.
global_backend = get_moe_runner_backend()
if not global_backend.is_auto():
if base_layer.runner is not None:
runner_backend = base_layer.runner.runner_backend
elif not global_backend.is_auto():
runner_backend = global_backend
elif (
hasattr(base_layer.quant_method, "runner")
and base_layer.quant_method.runner is not None
):
runner_backend = base_layer.quant_method.runner.runner_backend
else:
runner_backend = MoeRunnerBackend.TRITON
@@ -1051,8 +1050,9 @@ class FusedMoEWithLoRA(BaseLayerWithLoRA):
assert base_layer.quant_method is not None, "Quant method must be set"
self._quant_info = base_layer.quant_method.get_triton_quant_info(base_layer)
else:
raise NotImplementedError(
f"LoRA MoE not supported for backend {runner_backend}"
assert base_layer.quant_method is not None, "Quant method must be set"
self._quant_info = base_layer.quant_method.get_moe_quant_info(
base_layer, runner_backend
)
def set_lora_info(
@@ -10,8 +10,9 @@ from typing import TYPE_CHECKING
import torch
from sglang.srt.layers.moe.moe_runner.base import MoeRunnerConfig
from sglang.srt.layers.moe.moe_runner.base import DispatchMoeRunnerCore, MoeRunnerConfig
from sglang.srt.layers.moe.moe_runner.marlin import MarlinMoeQuantInfo
from sglang.srt.layers.moe.utils import MoeRunnerBackend
from sglang.srt.utils import is_cuda
if TYPE_CHECKING:
@@ -37,7 +38,7 @@ if _is_cuda:
from sglang.srt.layers.quantization.marlin_utils import marlin_make_workspace
class MarlinLoraRunnerCore:
class MarlinLoraRunnerCore(DispatchMoeRunnerCore):
"""
MoE runner using Marlin kernels for base projections, with hooks for LoRA.
@@ -53,6 +54,10 @@ class MarlinLoraRunnerCore:
def __init__(self, config: MoeRunnerConfig):
self.config = config
@property
def runner_backend(self) -> MoeRunnerBackend:
return MoeRunnerBackend.MARLIN
def run_from_dispatch(
self,
dispatch_output: StandardDispatchOutput,
@@ -0,0 +1,378 @@
import sys
from types import SimpleNamespace
import pytest
import torch
from sglang.srt.layers.moe import MoeA2ABackend, MoeRunnerBackend
from sglang.srt.layers.moe.fused_moe_triton import layer as fused_moe_layer_module
from sglang.srt.layers.moe.fused_moe_triton.layer import FusedMoE
from sglang.srt.layers.moe.moe_runner import runner as runner_module
from sglang.srt.layers.moe.moe_runner.base import (
DispatchMoeRunnerCore,
MoeQuantInfo,
MoeRunnerConfig,
MoeRunnerCore,
PermuteMethodPool,
RunnerInput,
RunnerOutput,
)
from sglang.srt.layers.moe.token_dispatcher.standard import (
StandardCombineInput,
StandardDispatchOutput,
)
from sglang.srt.layers.moe.topk import StandardTopKOutput
from sglang.srt.layers.moe.utils import (
RegisteredMoeRunnerBackend,
register_moe_runner_backend_name,
resolve_moe_runner_backend,
)
from sglang.srt.layers.quantization.unquant import UnquantizedFusedMoEMethod
from sglang.srt.lora.layers import FusedMoEWithLoRA
from sglang.srt.runtime_context import get_context, get_flags, get_parallel
from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=15, suite="base-c-test-cpu")
class _TestDispatchRunnerCore(DispatchMoeRunnerCore):
def __init__(self, config: MoeRunnerConfig, backend):
super().__init__(config)
self._backend = backend
self.calls = []
@property
def runner_backend(self):
return self._backend
def run_from_dispatch(
self,
dispatch_output,
quant_info,
runner_config,
hooks=None,
):
self.calls.append((dispatch_output, quant_info, runner_config, hooks))
return StandardCombineInput(dispatch_output.hidden_states + 1)
@pytest.fixture
def isolated_runner_registries(monkeypatch):
from sglang.srt.layers.moe import utils as moe_utils
monkeypatch.setattr(moe_utils, "_REGISTERED_MOE_RUNNER_BACKEND_NAMES", set())
monkeypatch.setattr(runner_module, "_CUSTOM_RUNNER_CORE_FACTORIES", {})
with get_flags().moe.override(a2a_backend=MoeA2ABackend.NONE):
yield
def test_registered_runner_core_uses_standard_dispatch(
isolated_runner_registries,
) -> None:
backend_name = "test_dispatch_extension"
runner_module.register_moe_runner_core(
backend_name,
lambda config: _TestDispatchRunnerCore(
config, resolve_moe_runner_backend(backend_name)
),
)
backend = resolve_moe_runner_backend(backend_name)
assert isinstance(backend, RegisteredMoeRunnerBackend)
assert backend.value == backend_name
runner = runner_module.MoeRunner(backend, MoeRunnerConfig())
dispatch_output = StandardDispatchOutput(
hidden_states=torch.zeros(1, 2),
hidden_states_scale=None,
topk_output=StandardTopKOutput(
topk_weights=torch.ones(1, 1),
topk_ids=torch.zeros(1, 1, dtype=torch.int64),
router_logits=torch.zeros(1, 1),
),
)
quant_info = MoeQuantInfo()
result = runner.run(dispatch_output, quant_info)
assert torch.equal(result.hidden_states, torch.ones(1, 2))
assert runner.runner_core.calls == [
(dispatch_output, quant_info, runner.config, None)
]
with pytest.raises(ValueError, match="already registered"):
runner_module.register_moe_runner_core(backend_name, _TestDispatchRunnerCore)
def test_runner_core_registration_can_override_builtin_backend(
isolated_runner_registries,
) -> None:
backend = MoeRunnerBackend.FLASHINFER_CUTLASS
runner_module.register_moe_runner_core(
backend.value,
lambda config: _TestDispatchRunnerCore(config, backend),
)
runner = runner_module.MoeRunner(backend, MoeRunnerConfig())
assert isinstance(runner.runner_core, _TestDispatchRunnerCore)
assert runner.runner_core.runner_backend is backend
def test_runner_backend_names_must_be_builtin_or_registered(
isolated_runner_registries,
) -> None:
assert resolve_moe_runner_backend("triton") is MoeRunnerBackend.TRITON
assert MoeRunnerBackend.EXPERIMENTAL_SGL_TRTLLM.is_flashinfer_trtllm()
assert MoeRunnerBackend.EXPERIMENTAL_SGL_MARLIN.is_marlin()
with pytest.raises(ValueError, match="neither built in nor registered"):
resolve_moe_runner_backend("unknown_backend")
with pytest.raises(ValueError, match="must not be empty"):
register_moe_runner_backend_name("")
with pytest.raises(ValueError, match="already built in"):
register_moe_runner_backend_name("triton")
def test_fused_moe_uses_explicit_quant_method_for_full_lifecycle(monkeypatch) -> None:
calls = []
method = UnquantizedFusedMoEMethod()
monkeypatch.setattr(
method, "create_weights", lambda **kwargs: calls.append("weights")
)
def create_runner(layer, config) -> None:
calls.append("runner")
method.runner = SimpleNamespace()
monkeypatch.setattr(method, "create_moe_runner", create_runner)
monkeypatch.setattr(
fused_moe_layer_module,
"create_moe_dispatcher",
lambda config: SimpleNamespace(),
)
with get_context().override_server_args(
model_path="dummy"
), get_flags().moe.override(
runner_backend=MoeRunnerBackend.AUTO,
a2a_backend=MoeA2ABackend.NONE,
), get_parallel().override(
moe_ep_size=1,
moe_ep_rank=0,
moe_tp_size=1,
moe_tp_rank=0,
tp_size=1,
tp_rank=0,
):
layer = FusedMoE(
num_experts=2,
hidden_size=4,
intermediate_size=8,
layer_id=0,
quant_method=method,
)
assert layer.quant_method is method
assert layer.runner is method.runner
assert calls == ["weights", "runner"]
def test_lora_uses_quant_method_contract_for_registered_backend(
monkeypatch, isolated_runner_registries
) -> None:
backend_name = "test_lora_extension"
register_moe_runner_backend_name(backend_name)
backend = resolve_moe_runner_backend(backend_name)
quant_info = object()
quant_calls = []
class FakeQuantMethod:
def get_moe_quant_info(self, layer, runner_backend):
quant_calls.append((layer, runner_backend))
return quant_info
base_layer = FusedMoE.__new__(FusedMoE)
torch.nn.Module.__init__(base_layer)
base_layer.quant_method = FakeQuantMethod()
base_layer.moe_runner_config = MoeRunnerConfig()
base_layer.dispatcher = object()
base_layer.num_local_experts = 2
base_layer.should_fuse_routed_scaling_factor_in_topk = False
base_layer.moe_tp_size = 1
base_layer.moe_tp_rank = 0
base_layer.intermediate_size_per_partition = 8
base_layer.runner = SimpleNamespace(runner_backend=backend)
lora_backend = SimpleNamespace(is_moe_lora=False)
created_runners = []
monkeypatch.setattr(
runner_module,
"MoeRunner",
lambda selected_backend, config, lora_enabled: created_runners.append(
(selected_backend, config, lora_enabled)
)
or object(),
)
wrapper = FusedMoEWithLoRA(base_layer, lora_backend)
assert wrapper._quant_info is quant_info
assert quant_calls == [(base_layer, backend)]
assert created_runners == [(backend, base_layer.moe_runner_config, True)]
class _NonTritonRunnerInput(RunnerInput):
"""Stands in for deep_gemm/aiter/ascend inputs: no ``topk_ids`` field."""
def __init__(self, backend, hidden_states):
self._backend = backend
self.hidden_states = hidden_states
@property
def runner_backend(self):
return self._backend
class _NonTritonRunnerOutput(RunnerOutput):
def __init__(self, backend, hidden_states):
self._backend = backend
self.hidden_states = hidden_states
@property
def runner_backend(self):
return self._backend
class _TestPermuteRunnerCore(MoeRunnerCore):
def __init__(self, config: MoeRunnerConfig, backend):
super().__init__(config)
self._backend = backend
self.hooks_seen = []
@property
def runner_backend(self):
return self._backend
def run(self, runner_input, quant_info, running_state, hooks=None):
self.hooks_seen.append(hooks)
return _NonTritonRunnerOutput(self._backend, runner_input.hidden_states + 1)
def test_non_triton_runner_input_skips_lora_hooks(
monkeypatch, isolated_runner_registries
) -> None:
"""LoRA-disabled runs must not inspect LoRA fields on the runner input.
Every non-Triton backend (deep_gemm, triton_kernels, aiter, ascend, ...)
produces a runner input without ``topk_ids``, so building hooks eagerly
would break each of them.
"""
monkeypatch.setattr(
PermuteMethodPool,
"_pre_permute_methods",
dict(PermuteMethodPool._pre_permute_methods),
)
monkeypatch.setattr(
PermuteMethodPool,
"_post_permute_methods",
dict(PermuteMethodPool._post_permute_methods),
)
backend_name = "test_permute_extension"
runner_module.register_moe_runner_core(
backend_name,
lambda config: _TestPermuteRunnerCore(
config, resolve_moe_runner_backend(backend_name)
),
)
backend = resolve_moe_runner_backend(backend_name)
PermuteMethodPool.register_pre_permute(
"standard",
backend_name,
lambda dispatch_output, quant_info, config, state: _NonTritonRunnerInput(
backend, dispatch_output.hidden_states
),
)
PermuteMethodPool.register_post_permute(
backend_name,
"standard",
lambda runner_output, quant_info, config, state: StandardCombineInput(
runner_output.hidden_states
),
)
runner = runner_module.MoeRunner(backend, MoeRunnerConfig())
dispatch_output = StandardDispatchOutput(
hidden_states=torch.zeros(1, 2),
hidden_states_scale=None,
topk_output=StandardTopKOutput(
topk_weights=torch.ones(1, 1),
topk_ids=torch.zeros(1, 1, dtype=torch.int64),
router_logits=torch.zeros(1, 1),
),
)
result = runner.run(dispatch_output, MoeQuantInfo())
assert torch.equal(result.hidden_states, torch.ones(1, 2))
assert runner.runner_core.hooks_seen == [None]
def test_trtllm_quant_method_defines_runner_after_create_moe_runner() -> None:
"""FusedMoE reads `quant_method.runner` right after `create_moe_runner`, so
a method that never builds a MoeRunner must still define the attribute."""
from sglang.srt.layers.quantization.mxfp4_flashinfer_trtllm_moe import (
Mxfp4FlashinferTrtllmMoEMethod,
)
method = Mxfp4FlashinferTrtllmMoEMethod.__new__(Mxfp4FlashinferTrtllmMoEMethod)
assert not hasattr(method, "runner")
method.create_moe_runner(
SimpleNamespace(num_local_experts=2), MoeRunnerConfig(swiglu_limit=None)
)
assert method.runner is None
def test_fused_moe_layer_runner_is_none_when_method_builds_no_runner(
monkeypatch,
) -> None:
"""The overlap-args helpers key off `runner is not None`, so a layer whose
quant method builds no MoeRunner must fall back instead of raising."""
method = UnquantizedFusedMoEMethod()
monkeypatch.setattr(method, "create_weights", lambda **kwargs: None)
monkeypatch.setattr(method, "create_moe_runner", lambda layer, config: None)
monkeypatch.setattr(
fused_moe_layer_module,
"create_moe_dispatcher",
lambda config: SimpleNamespace(),
)
with get_context().override_server_args(
model_path="dummy"
), get_flags().moe.override(
runner_backend=MoeRunnerBackend.AUTO,
a2a_backend=MoeA2ABackend.NONE,
), get_parallel().override(
moe_ep_size=1,
moe_ep_rank=0,
moe_tp_size=1,
moe_tp_rank=0,
tp_size=1,
tp_rank=0,
):
layer = FusedMoE(
num_experts=2,
hidden_size=4,
intermediate_size=8,
layer_id=0,
quant_method=method,
)
assert layer.runner is None
layer.clear_overlap_args()
assert layer.down_gemm_overlap_args is None
if __name__ == "__main__":
sys.exit(pytest.main([__file__, "-v"]))
@@ -65,6 +65,7 @@ def _make_base_layer(quant_method=None) -> types.SimpleNamespace:
moe_tp_size=1,
moe_tp_rank=0,
intermediate_size_per_partition=32,
runner=None,
)
@@ -114,15 +114,18 @@ def _load_lora_weight_to_buffer(pool, **kwargs):
def _load_moe_backend_enum():
tree = ast.parse(MOE_UTILS_PATH.read_text())
backend = next(
node
for node in tree.body
if isinstance(node, ast.ClassDef) and node.name == "MoeRunnerBackend"
)
classes = {node.name: node for node in tree.body if isinstance(node, ast.ClassDef)}
backend = classes["MoeRunnerBackend"]
body = [
classes[base.id]
for base in backend.bases
if isinstance(base, ast.Name) and base.id in classes
]
body.append(backend)
namespace = {"Enum": Enum}
exec(
compile(
ast.fix_missing_locations(ast.Module(body=[backend], type_ignores=[])),
ast.fix_missing_locations(ast.Module(body=body, type_ignores=[])),
str(MOE_UTILS_PATH),
"exec",
),