[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. params_dtype: Data type for the parameters.
reduce_results: Whether to apply all_reduce on the output of the layer reduce_results: Whether to apply all_reduce on the output of the layer
quant_config: Quantization configuration. quant_config: Quantization configuration.
quant_method: Explicit quant method, overriding selection from quant_config.
inplace: suggestion to compute inplace (modify input activation). inplace: suggestion to compute inplace (modify input activation).
enable_qwen35_fp8_deferred_finalize: Whether this concrete Qwen3.5 enable_qwen35_fp8_deferred_finalize: Whether this concrete Qwen3.5
layer may expose FlashInfer's block-FP8 deferred MoE output. layer may expose FlashInfer's block-FP8 deferred MoE output.
@@ -325,6 +326,7 @@ class FusedMoE(torch.nn.Module):
is_gated: bool = True, is_gated: bool = True,
gate_up_interleaved: bool = True, gate_up_interleaved: bool = True,
enable_qwen35_fp8_deferred_finalize: bool = False, enable_qwen35_fp8_deferred_finalize: bool = False,
quant_method: Optional[FusedMoEMethodBase] = None,
): ):
super().__init__() super().__init__()
if params_dtype is None: if params_dtype is None:
@@ -430,17 +432,19 @@ class FusedMoE(torch.nn.Module):
gate_up_interleaved=gate_up_interleaved, gate_up_interleaved=gate_up_interleaved,
) )
self.quant_method: Optional[FusedMoEMethodBase] = None self.quant_method = quant_method
server_args = get_server_args() server_args = get_server_args()
kt_config = create_kt_config_from_server_args(server_args, layer_id) kt_config = create_kt_config_from_server_args(server_args, layer_id)
if kt_config is not None: 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) gpu_method = quant_config.get_quant_method(self, prefix)
else: else:
gpu_method = UnquantizedFusedMoEMethod(self.use_triton_kernels) gpu_method = UnquantizedFusedMoEMethod(self.use_triton_kernels)
self.quant_method = KTEPWrapperMethod(gpu_method, kt_config) self.quant_method = KTEPWrapperMethod(gpu_method, kt_config)
else: 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) self.quant_method = quant_config.get_quant_method(self, prefix)
if self.quant_method is None: if self.quant_method is None:
self.quant_method = UnquantizedFusedMoEMethod( self.quant_method = UnquantizedFusedMoEMethod(
@@ -525,8 +529,7 @@ class FusedMoE(torch.nn.Module):
self._dwdp_bound = False 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 @property
def num_global_routed_experts(self) -> int: def num_global_routed_experts(self) -> int:
@@ -1685,7 +1688,7 @@ class FusedMoE(torch.nn.Module):
def set_overlap_args( def set_overlap_args(
self, down_gemm_overlap_args: DownGemmOverlapArgs, meta_overlap_args: dict 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) self.runner.set_overlap_args(down_gemm_overlap_args, meta_overlap_args)
else: else:
# TODO: remove this branch after MoE refactor # TODO: remove this branch after MoE refactor
@@ -1693,7 +1696,7 @@ class FusedMoE(torch.nn.Module):
self.meta_overlap_args = meta_overlap_args self.meta_overlap_args = meta_overlap_args
def clear_overlap_args(self) -> None: def clear_overlap_args(self) -> None:
if hasattr(self, "runner"): if self.runner is not None:
self.runner.clear_overlap_args() self.runner.clear_overlap_args()
else: else:
# TODO: remove this branch after MoE refactor # 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.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 ( from sglang.srt.layers.moe.utils import (
MoeA2ABackend, MoeA2ABackend,
MoeRunnerBackend, MoeRunnerBackend,
MoeRunnerBackendLike,
RoutingMethodType, RoutingMethodType,
) )
@@ -113,6 +114,26 @@ class MoeRunnerCore(ABC):
return self.runner_backend == MoeRunnerBackend.TRITON 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: class FusedOpPool:
_fused_funcs: dict[str, Callable] = {} _fused_funcs: dict[str, Callable] = {}
@@ -2,32 +2,58 @@ from __future__ import annotations
import logging import logging
import os 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 ( from sglang.srt.layers.moe.moe_runner.base import (
DispatchMoeRunnerCore,
FusedOpPool, FusedOpPool,
MoeRunnerConfig, MoeRunnerConfig,
PermuteMethodPool, PermuteMethodPool,
) )
from sglang.srt.layers.moe.moe_runner.deep_gemm import DeepGemmRunnerCore 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.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: if TYPE_CHECKING:
from sglang.srt.batch_overlap.single_batch_overlap import DownGemmOverlapArgs 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.moe_runner.base import MoeQuantInfo
from sglang.srt.layers.moe.token_dispatcher.base import CombineInput, DispatchOutput 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 from sglang.srt.lora.lora_moe_runners import LoRAHooks
logger = logging.getLogger(__name__) 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: class MoeRunner:
def __init__( def __init__(
self, self,
runner_backend: MoeRunnerBackend, runner_backend: MoeRunnerBackendLike,
config: MoeRunnerConfig, config: MoeRunnerConfig,
lora_enabled: bool = False, lora_enabled: bool = False,
): ):
@@ -61,7 +87,9 @@ class MoeRunner:
self.fused_func = None 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) self.runner_core = TritonRunnerCore(config)
elif runner_backend.is_ascend(): elif runner_backend.is_ascend():
from sglang.srt.layers.moe.moe_runner.ascend import AscendRunnerCore from sglang.srt.layers.moe.moe_runner.ascend import AscendRunnerCore
@@ -157,7 +185,16 @@ class MoeRunner:
assert self.runner_core is not None 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.layers.moe.token_dispatcher.base import DispatchOutput
from sglang.srt.lora.lora_moe_runners import build_lora_hooks from sglang.srt.lora.lora_moe_runners import build_lora_hooks
@@ -167,19 +204,18 @@ class MoeRunner:
_runner_input.topk_output.topk_ids, _runner_input.topk_output.topk_ids,
) )
else: else:
assert isinstance(_runner_input, TritonRunnerInput), type(_runner_input)
hidden_states = _runner_input.hidden_states hidden_states = _runner_input.hidden_states
topk_ids = getattr(_runner_input, "topk_ids", None) topk_ids = _runner_input.topk_ids
if self.lora_enabled and lora_info is not None: return build_lora_hooks(
return build_lora_hooks( hidden_states,
hidden_states, lora_info,
lora_info, topk_ids,
topk_ids, )
)
return None
# Runners that handle dispatch_output directly (e.g., MarlinRunnerCore) # Runners that handle dispatch_output directly (e.g., MarlinRunnerCore)
# bypass the pre-permute step and do their own alignment internally. # 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) hooks = _maybe_build_lora_hooks(dispatch_output)
return self.runner_core.run_from_dispatch( return self.runner_core.run_from_dispatch(
dispatch_output, quant_info, self.config, hooks=hooks dispatch_output, quant_info, self.config, hooks=hooks
+102 -56
View File
@@ -3,6 +3,7 @@ from __future__ import annotations
import logging import logging
import os import os
from contextlib import contextmanager from contextlib import contextmanager
from dataclasses import dataclass
from enum import Enum, IntEnum from enum import Enum, IntEnum
from typing import NamedTuple 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" AUTO = "auto"
DEEP_GEMM = "deep_gemm" DEEP_GEMM = "deep_gemm"
@@ -121,67 +188,46 @@ class MoeRunnerBackend(Enum):
HPC_OPS = "hpc_ops" HPC_OPS = "hpc_ops"
INTEL_XPU = "intel_xpu" INTEL_XPU = "intel_xpu"
def is_auto(self):
return self == MoeRunnerBackend.AUTO
def is_hpc_ops(self): @dataclass(frozen=True)
return self == MoeRunnerBackend.HPC_OPS class RegisteredMoeRunnerBackend(_MoeRunnerBackendPredicates):
"""Identifier for an MoE runner backend supplied by an extension."""
def is_deep_gemm(self): value: str
return self == MoeRunnerBackend.DEEP_GEMM
def is_triton(self):
return self == MoeRunnerBackend.TRITON
def is_ascend(self): MoeRunnerBackendLike = MoeRunnerBackend | RegisteredMoeRunnerBackend
return self == MoeRunnerBackend.ASCEND _REGISTERED_MOE_RUNNER_BACKEND_NAMES: set[str] = set()
def is_triton_kernels(self):
return self == MoeRunnerBackend.TRITON_KERNELS
def is_flashinfer_trtllm(self): def register_moe_runner_backend_name(name: str) -> None:
# experimental_sgl_trtllm shares the TRT-LLM FP8 kernels + layout, so it inherits """Register a backend name supplied by an out-of-tree extension."""
# trtllm weight-prep here; divergent sites check is_experimental_sgl_trtllm() first.
return self in (
MoeRunnerBackend.FLASHINFER_TRTLLM,
MoeRunnerBackend.EXPERIMENTAL_SGL_TRTLLM,
)
def is_experimental_sgl_trtllm(self): if not name:
return self == MoeRunnerBackend.EXPERIMENTAL_SGL_TRTLLM 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): def resolve_moe_runner_backend(
return self == MoeRunnerBackend.FLASHINFER_CUTLASS backend: str | MoeRunnerBackendLike,
) -> MoeRunnerBackendLike:
"""Resolve a built-in or registered backend identifier."""
def is_flashinfer_cutedsl(self): if isinstance(backend, (MoeRunnerBackend, RegisteredMoeRunnerBackend)):
return self == MoeRunnerBackend.FLASHINFER_CUTEDSL return backend
try:
def is_flashinfer_mxfp4(self): return MoeRunnerBackend(backend)
return self == MoeRunnerBackend.FLASHINFER_MXFP4 except ValueError:
if backend in _REGISTERED_MOE_RUNNER_BACKEND_NAMES:
def is_cutlass(self): return RegisteredMoeRunnerBackend(backend)
return self == MoeRunnerBackend.CUTLASS raise ValueError(
f"MoE runner backend {backend!r} is neither built in nor registered"
def is_marlin(self): ) from None
# 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
def is_intel_xpu(self): def is_intel_xpu(self):
return self == MoeRunnerBackend.INTEL_XPU return self == MoeRunnerBackend.INTEL_XPU
@@ -353,9 +399,9 @@ def initialize_moe_config():
spec = get_spec() spec = get_spec()
moe = get_flags().moe moe = get_flags().moe
moe.a2a_backend = MoeA2ABackend(exec_moe.moe_a2a_backend) 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 = ( 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 if spec.speculative_moe_runner_backend is not None
else moe.runner_backend else moe.runner_backend
) )
@@ -391,14 +437,14 @@ def get_moe_a2a_backend() -> MoeA2ABackend:
return moe.a2a_backend return moe.a2a_backend
def get_moe_runner_backend() -> MoeRunnerBackend: def get_moe_runner_backend() -> MoeRunnerBackendLike:
moe = get_flags().moe moe = get_flags().moe
if moe.runner_backend is None: if moe.runner_backend is None:
moe.runner_backend = MoeRunnerBackend.AUTO moe.runner_backend = MoeRunnerBackend.AUTO
return moe.runner_backend return moe.runner_backend
def get_speculative_moe_runner_backend() -> MoeRunnerBackend: def get_speculative_moe_runner_backend() -> MoeRunnerBackendLike:
moe = get_flags().moe moe = get_flags().moe
if moe.speculative_runner_backend is None: if moe.speculative_runner_backend is None:
logger.warning( logger.warning(
@@ -13,9 +13,11 @@ from torch import nn
from sglang.srt.layers.modelopt_utils import canonicalize_modelopt_quant_algo from sglang.srt.layers.modelopt_utils import canonicalize_modelopt_quant_algo
if TYPE_CHECKING: 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.moe_runner.triton import TritonMoeQuantInfo
from sglang.srt.layers.moe.token_dispatcher import CombineInput, DispatchOutput 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 from sglang.srt.models.utils import WeightsMapper
@@ -86,6 +88,7 @@ class LinearMethodBase(QuantizeMethodBase):
class FusedMoEMethodBase(QuantizeMethodBase): class FusedMoEMethodBase(QuantizeMethodBase):
runner: MoeRunner | None = None
def create_weights( def create_weights(
self, self,
@@ -124,6 +127,15 @@ class FusedMoEMethodBase(QuantizeMethodBase):
f"{type(self).__name__} must implement get_triton_quant_info()" 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): class QuantizationConfig(ABC):
"""Base class for quantization configs.""" """Base class for quantization configs."""
@@ -57,6 +57,10 @@ class Mxfp4FlashinferTrtllmMoEMethod:
def create_moe_runner(self, layer, moe_runner_config): def create_moe_runner(self, layer, moe_runner_config):
self.moe_runner_config = 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 swiglu_limit = moe_runner_config.swiglu_limit
self._gemm1_clamp_limit_tensor = ( 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 base_layer.should_fuse_routed_scaling_factor_in_topk
) )
self.tp_size = getattr(base_layer, "moe_tp_size", 1) self.tp_size = base_layer.moe_tp_size
self.tp_rank = getattr(base_layer, "moe_tp_rank", 0) self.tp_rank = base_layer.moe_tp_rank
self.intermediate_size_per_partition = getattr( self.intermediate_size_per_partition = (
base_layer, "intermediate_size_per_partition", None 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 = ( 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 # 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.moe_runner.runner import MoeRunner
from sglang.srt.layers.moe.utils import get_moe_runner_backend 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() 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 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: else:
runner_backend = MoeRunnerBackend.TRITON runner_backend = MoeRunnerBackend.TRITON
@@ -1051,8 +1050,9 @@ class FusedMoEWithLoRA(BaseLayerWithLoRA):
assert base_layer.quant_method is not None, "Quant method must be set" 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) self._quant_info = base_layer.quant_method.get_triton_quant_info(base_layer)
else: else:
raise NotImplementedError( assert base_layer.quant_method is not None, "Quant method must be set"
f"LoRA MoE not supported for backend {runner_backend}" self._quant_info = base_layer.quant_method.get_moe_quant_info(
base_layer, runner_backend
) )
def set_lora_info( def set_lora_info(
@@ -10,8 +10,9 @@ from typing import TYPE_CHECKING
import torch 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.moe_runner.marlin import MarlinMoeQuantInfo
from sglang.srt.layers.moe.utils import MoeRunnerBackend
from sglang.srt.utils import is_cuda from sglang.srt.utils import is_cuda
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -37,7 +38,7 @@ if _is_cuda:
from sglang.srt.layers.quantization.marlin_utils import marlin_make_workspace 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. MoE runner using Marlin kernels for base projections, with hooks for LoRA.
@@ -53,6 +54,10 @@ class MarlinLoraRunnerCore:
def __init__(self, config: MoeRunnerConfig): def __init__(self, config: MoeRunnerConfig):
self.config = config self.config = config
@property
def runner_backend(self) -> MoeRunnerBackend:
return MoeRunnerBackend.MARLIN
def run_from_dispatch( def run_from_dispatch(
self, self,
dispatch_output: StandardDispatchOutput, 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_size=1,
moe_tp_rank=0, moe_tp_rank=0,
intermediate_size_per_partition=32, intermediate_size_per_partition=32,
runner=None,
) )
@@ -114,15 +114,18 @@ def _load_lora_weight_to_buffer(pool, **kwargs):
def _load_moe_backend_enum(): def _load_moe_backend_enum():
tree = ast.parse(MOE_UTILS_PATH.read_text()) tree = ast.parse(MOE_UTILS_PATH.read_text())
backend = next( classes = {node.name: node for node in tree.body if isinstance(node, ast.ClassDef)}
node backend = classes["MoeRunnerBackend"]
for node in tree.body body = [
if isinstance(node, ast.ClassDef) and node.name == "MoeRunnerBackend" classes[base.id]
) for base in backend.bases
if isinstance(base, ast.Name) and base.id in classes
]
body.append(backend)
namespace = {"Enum": Enum} namespace = {"Enum": Enum}
exec( exec(
compile( 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), str(MOE_UTILS_PATH),
"exec", "exec",
), ),