Refactor Marlin MoeRunner (#14554)
Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com> Co-authored-by: Cheng Wan <54331508+ch-wan@users.noreply.github.com>
This commit is contained in:
125
python/sglang/srt/layers/moe/moe_runner/marlin.py
Normal file
125
python/sglang/srt/layers/moe/moe_runner/marlin.py
Normal file
@@ -0,0 +1,125 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import TYPE_CHECKING, Optional
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.layers.moe.moe_runner.base import (
|
||||
MoeQuantInfo,
|
||||
MoeRunnerConfig,
|
||||
RunnerInput,
|
||||
RunnerOutput,
|
||||
register_fused_func,
|
||||
)
|
||||
from sglang.srt.layers.moe.utils import MoeRunnerBackend
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.layers.moe.token_dispatcher import (
|
||||
StandardCombineInput,
|
||||
StandardDispatchOutput,
|
||||
)
|
||||
|
||||
MARLIN_MOE_WORKSPACE: Optional[torch.Tensor] = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class MarlinRunnerInput(RunnerInput):
|
||||
"""Input bundle passed to the Marlin runner core."""
|
||||
|
||||
hidden_states: torch.Tensor
|
||||
topk_weights: torch.Tensor
|
||||
topk_ids: torch.Tensor
|
||||
router_logits: torch.Tensor
|
||||
|
||||
@property
|
||||
def runner_backend(self) -> MoeRunnerBackend:
|
||||
return MoeRunnerBackend.MARLIN
|
||||
|
||||
|
||||
@dataclass
|
||||
class MarlinRunnerOutput(RunnerOutput):
|
||||
"""Output bundle returned from the Marlin runner core."""
|
||||
|
||||
hidden_states: torch.Tensor
|
||||
|
||||
@property
|
||||
def runner_backend(self) -> MoeRunnerBackend:
|
||||
return MoeRunnerBackend.MARLIN
|
||||
|
||||
|
||||
@dataclass
|
||||
class MarlinMoeQuantInfo(MoeQuantInfo):
|
||||
"""Quantization payload consumed by the Marlin backend."""
|
||||
|
||||
w13_qweight: torch.Tensor
|
||||
w2_qweight: torch.Tensor
|
||||
w13_scales: torch.Tensor
|
||||
w2_scales: torch.Tensor
|
||||
w13_g_idx_sort_indices: Optional[torch.Tensor]
|
||||
w2_g_idx_sort_indices: Optional[torch.Tensor]
|
||||
weight_bits: int
|
||||
|
||||
# GPTQ specific (Optional)
|
||||
w13_g_idx: Optional[torch.Tensor] = None
|
||||
w2_g_idx: Optional[torch.Tensor] = None
|
||||
is_k_full: bool = True
|
||||
|
||||
# AWQ specific (Optional)
|
||||
w13_qzeros: Optional[torch.Tensor] = None
|
||||
w2_qzeros: Optional[torch.Tensor] = None
|
||||
|
||||
# Optional
|
||||
expert_map: Optional[torch.Tensor] = None
|
||||
|
||||
|
||||
@register_fused_func("none", "marlin")
|
||||
def fused_experts_none_to_marlin(
|
||||
dispatch_output: StandardDispatchOutput,
|
||||
quant_info: MarlinMoeQuantInfo,
|
||||
runner_config: MoeRunnerConfig,
|
||||
) -> StandardCombineInput:
|
||||
global MARLIN_MOE_WORKSPACE
|
||||
from sglang.srt.layers.moe.fused_moe_triton.fused_marlin_moe import fused_marlin_moe
|
||||
from sglang.srt.layers.moe.token_dispatcher.standard import StandardCombineInput
|
||||
from sglang.srt.layers.quantization.marlin_utils import marlin_make_workspace
|
||||
|
||||
hidden_states = dispatch_output.hidden_states
|
||||
topk_output = dispatch_output.topk_output
|
||||
|
||||
assert runner_config.activation == "silu", "Only SiLU activation is supported."
|
||||
|
||||
if (
|
||||
MARLIN_MOE_WORKSPACE is None
|
||||
or MARLIN_MOE_WORKSPACE.device != hidden_states.device
|
||||
):
|
||||
MARLIN_MOE_WORKSPACE = marlin_make_workspace(
|
||||
hidden_states.device, max_blocks_per_sm=4
|
||||
)
|
||||
|
||||
output = fused_marlin_moe(
|
||||
hidden_states=hidden_states,
|
||||
w1=quant_info.w13_qweight,
|
||||
w2=quant_info.w2_qweight,
|
||||
w1_scale=quant_info.w13_scales,
|
||||
w2_scale=quant_info.w2_scales,
|
||||
gating_output=topk_output.router_logits,
|
||||
topk_weights=topk_output.topk_weights,
|
||||
topk_ids=topk_output.topk_ids,
|
||||
expert_map=quant_info.expert_map,
|
||||
g_idx1=quant_info.w13_g_idx,
|
||||
g_idx2=quant_info.w2_g_idx,
|
||||
sort_indices1=quant_info.w13_g_idx_sort_indices,
|
||||
sort_indices2=quant_info.w2_g_idx_sort_indices,
|
||||
w1_zeros=quant_info.w13_qzeros,
|
||||
w2_zeros=quant_info.w2_qzeros,
|
||||
workspace=MARLIN_MOE_WORKSPACE,
|
||||
num_bits=quant_info.weight_bits,
|
||||
is_k_full=quant_info.is_k_full,
|
||||
inplace=runner_config.inplace,
|
||||
routed_scaling_factor=runner_config.routed_scaling_factor,
|
||||
).to(hidden_states.dtype)
|
||||
|
||||
return StandardCombineInput(
|
||||
hidden_states=output,
|
||||
)
|
||||
@@ -37,6 +37,8 @@ class MoeRunner:
|
||||
self.runner_core = TritonKernelsRunnerCore(config)
|
||||
elif runner_backend.is_deep_gemm():
|
||||
self.runner_core = DeepGemmRunnerCore(config)
|
||||
elif runner_backend.is_marlin():
|
||||
self.runner_core = None # Marlin only supports fused path
|
||||
else:
|
||||
raise NotImplementedError(f"Unsupported runner backend: {runner_backend}")
|
||||
|
||||
|
||||
@@ -59,6 +59,7 @@ class MoeRunnerBackend(Enum):
|
||||
FLASHINFER_MXFP4 = "flashinfer_mxfp4"
|
||||
FLASHINFER_CUTEDSL = "flashinfer_cutedsl"
|
||||
CUTLASS = "cutlass"
|
||||
MARLIN = "marlin"
|
||||
|
||||
def is_auto(self):
|
||||
return self == MoeRunnerBackend.AUTO
|
||||
@@ -87,6 +88,9 @@ class MoeRunnerBackend(Enum):
|
||||
def is_cutlass(self):
|
||||
return self == MoeRunnerBackend.CUTLASS
|
||||
|
||||
def is_marlin(self):
|
||||
return self == MoeRunnerBackend.MARLIN
|
||||
|
||||
|
||||
class DeepEPMode(Enum):
|
||||
|
||||
|
||||
@@ -11,6 +11,13 @@ from sglang.srt.hardware_backend.npu.quantization.fused_moe_method_npu import (
|
||||
npu_fused_experts,
|
||||
)
|
||||
from sglang.srt.layers.linear import LinearBase, set_weight_attrs
|
||||
from sglang.srt.layers.moe import (
|
||||
MoeRunner,
|
||||
MoeRunnerBackend,
|
||||
MoeRunnerConfig,
|
||||
get_moe_runner_backend,
|
||||
)
|
||||
from sglang.srt.layers.moe.moe_runner.marlin import MarlinMoeQuantInfo
|
||||
from sglang.srt.layers.parameter import GroupQuantScaleParameter, PackedvLLMParameter
|
||||
from sglang.srt.layers.quantization.base_config import (
|
||||
FusedMoEMethodBase,
|
||||
@@ -37,10 +44,9 @@ from sglang.srt.layers.quantization.utils import get_scalar_types, replace_param
|
||||
from sglang.srt.utils.patch_torch import register_fake_if_exists
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.layers.moe.moe_runner import MoeRunnerConfig
|
||||
from sglang.srt.layers.moe.token_dispatcher import (
|
||||
StandardDispatchOutput,
|
||||
CombineInput,
|
||||
StandardDispatchOutput,
|
||||
)
|
||||
|
||||
from sglang.srt.utils import is_cuda, is_hip, is_npu, is_xpu
|
||||
@@ -753,10 +759,6 @@ class AWQMoEMethod(FusedMoEMethodBase):
|
||||
layer.register_parameter("w2_qzeros", w2_qzeros)
|
||||
set_weight_attrs(w2_qzeros, extra_weight_attrs)
|
||||
|
||||
device = layer.w13_qweight.device
|
||||
if not _is_npu:
|
||||
layer.workspace = marlin_make_workspace(device, 4)
|
||||
|
||||
def process_weights_after_loading(self, layer: torch.nn.Module) -> None:
|
||||
num_experts = layer.w13_qweight.shape[0]
|
||||
device = layer.w13_qweight.device
|
||||
@@ -825,44 +827,29 @@ class AWQMoEMethod(FusedMoEMethodBase):
|
||||
def create_moe_runner(
|
||||
self, layer: torch.nn.Module, moe_runner_config: MoeRunnerConfig
|
||||
):
|
||||
assert get_moe_runner_backend().is_auto()
|
||||
self.moe_runner_config = moe_runner_config
|
||||
self.runner = MoeRunner(MoeRunnerBackend.MARLIN, moe_runner_config)
|
||||
|
||||
def apply(
|
||||
self,
|
||||
layer: torch.nn.Module,
|
||||
dispatch_output: StandardDispatchOutput,
|
||||
) -> CombineInput:
|
||||
from sglang.srt.layers.moe.fused_moe_triton.fused_marlin_moe import (
|
||||
fused_marlin_moe,
|
||||
|
||||
quant_info = MarlinMoeQuantInfo(
|
||||
w13_qweight=layer.w13_qweight,
|
||||
w2_qweight=layer.w2_qweight,
|
||||
w13_scales=layer.w13_scales,
|
||||
w2_scales=layer.w2_scales,
|
||||
w13_g_idx_sort_indices=layer.w13_g_idx_sort_indices,
|
||||
w2_g_idx_sort_indices=layer.w2_g_idx_sort_indices,
|
||||
w13_qzeros=layer.w13_qzeros,
|
||||
w2_qzeros=layer.w2_qzeros,
|
||||
weight_bits=self.quant_config.weight_bits,
|
||||
)
|
||||
from sglang.srt.layers.moe.token_dispatcher import StandardCombineInput
|
||||
|
||||
assert (
|
||||
self.moe_runner_config.activation == "silu"
|
||||
), "Only SiLU activation is supported."
|
||||
|
||||
x = dispatch_output.hidden_states
|
||||
topk_output = dispatch_output.topk_output
|
||||
orig_dtype = x.dtype
|
||||
|
||||
topk_weights, topk_ids, router_logits = topk_output
|
||||
|
||||
output = fused_marlin_moe(
|
||||
x,
|
||||
layer.w13_qweight,
|
||||
layer.w2_qweight,
|
||||
layer.w13_scales,
|
||||
layer.w2_scales,
|
||||
router_logits,
|
||||
topk_weights,
|
||||
topk_ids,
|
||||
sort_indices1=layer.w13_g_idx_sort_indices,
|
||||
sort_indices2=layer.w2_g_idx_sort_indices,
|
||||
w1_zeros=layer.w13_qzeros,
|
||||
w2_zeros=layer.w2_qzeros,
|
||||
num_bits=self.quant_config.weight_bits,
|
||||
).to(orig_dtype)
|
||||
return StandardCombineInput(hidden_states=output)
|
||||
return self.runner.run(dispatch_output, quant_info)
|
||||
|
||||
|
||||
class AWQMoEAscendMethod(AWQMoEMethod):
|
||||
|
||||
@@ -7,6 +7,13 @@ from typing import TYPE_CHECKING, Any, Callable, Dict, List, Optional, Union
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.layers.moe import (
|
||||
MoeRunner,
|
||||
MoeRunnerBackend,
|
||||
MoeRunnerConfig,
|
||||
get_moe_runner_backend,
|
||||
)
|
||||
from sglang.srt.layers.moe.moe_runner.marlin import MarlinMoeQuantInfo
|
||||
from sglang.srt.layers.parameter import (
|
||||
BasevLLMParameter,
|
||||
ChannelQuantScaleParameter,
|
||||
@@ -46,7 +53,6 @@ from sglang.srt.utils import is_cuda
|
||||
from sglang.srt.utils.patch_torch import register_fake_if_exists
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.layers.moe.moe_runner import MoeRunnerConfig
|
||||
from sglang.srt.layers.moe.token_dispatcher import (
|
||||
CombineInput,
|
||||
StandardDispatchOutput,
|
||||
@@ -1052,48 +1058,29 @@ class GPTQMarlinMoEMethod(FusedMoEMethodBase):
|
||||
def create_moe_runner(
|
||||
self, layer: torch.nn.Module, moe_runner_config: MoeRunnerConfig
|
||||
):
|
||||
assert get_moe_runner_backend().is_auto()
|
||||
self.moe_runner_config = moe_runner_config
|
||||
self.runner = MoeRunner(MoeRunnerBackend.MARLIN, moe_runner_config)
|
||||
|
||||
def apply(
|
||||
self,
|
||||
layer: torch.nn.Module,
|
||||
dispatch_output: StandardDispatchOutput,
|
||||
) -> CombineInput:
|
||||
from sglang.srt.layers.moe.fused_moe_triton.fused_marlin_moe import (
|
||||
fused_marlin_moe,
|
||||
)
|
||||
from sglang.srt.layers.moe.token_dispatcher import StandardCombineInput
|
||||
|
||||
x = dispatch_output.hidden_states
|
||||
topk_output = dispatch_output.topk_output
|
||||
|
||||
assert (
|
||||
self.moe_runner_config.activation == "silu"
|
||||
), "Only SiLU activation is supported."
|
||||
|
||||
# The input must currently be float16
|
||||
orig_dtype = x.dtype
|
||||
x = x.half()
|
||||
|
||||
topk_weights, topk_ids, router_logits = topk_output
|
||||
|
||||
output = fused_marlin_moe(
|
||||
x,
|
||||
layer.w13_qweight,
|
||||
layer.w2_qweight,
|
||||
layer.w13_scales,
|
||||
layer.w2_scales,
|
||||
router_logits,
|
||||
topk_weights,
|
||||
topk_ids,
|
||||
g_idx1=layer.w13_g_idx,
|
||||
g_idx2=layer.w2_g_idx,
|
||||
sort_indices1=layer.w13_g_idx_sort_indices,
|
||||
sort_indices2=layer.w2_g_idx_sort_indices,
|
||||
num_bits=self.quant_config.weight_bits,
|
||||
quant_info = MarlinMoeQuantInfo(
|
||||
w13_qweight=layer.w13_qweight,
|
||||
w2_qweight=layer.w2_qweight,
|
||||
w13_scales=layer.w13_scales,
|
||||
w2_scales=layer.w2_scales,
|
||||
w13_g_idx=layer.w13_g_idx,
|
||||
w2_g_idx=layer.w2_g_idx,
|
||||
w13_g_idx_sort_indices=layer.w13_g_idx_sort_indices,
|
||||
w2_g_idx_sort_indices=layer.w2_g_idx_sort_indices,
|
||||
weight_bits=self.quant_config.weight_bits,
|
||||
is_k_full=self.is_k_full,
|
||||
).to(orig_dtype)
|
||||
return StandardCombineInput(hidden_states=output)
|
||||
)
|
||||
|
||||
return self.runner.run(dispatch_output, quant_info)
|
||||
|
||||
|
||||
# Register fake implementations for torch.compile support
|
||||
|
||||
Reference in New Issue
Block a user