[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:
co-authored by
Alex Nails
Claude Opus 5
parent
ca8ff035c3
commit
ed39568e79
@@ -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
|
||||
|
||||
@@ -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 = (
|
||||
|
||||
@@ -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",
|
||||
),
|
||||
|
||||
Reference in New Issue
Block a user