[MoE Refactor] Centralize FlashInfer CUTLASS MoE runner (#28211)
This commit is contained in:
@@ -0,0 +1,372 @@
|
|||||||
|
"""FlashInfer CUTLASS MoE fused funcs.
|
||||||
|
|
||||||
|
This module owns the FlashInfer ``cutlass_fused_moe`` calls used by the
|
||||||
|
unquantized, ModelOpt FP8, ModelOpt NVFP4, and SM90 MXFP4 MoE paths.
|
||||||
|
Quantization methods prepare a small quant_info payload and route through
|
||||||
|
``MoeRunner``.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from typing import TYPE_CHECKING, Optional
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.distributed import get_tp_group
|
||||||
|
from sglang.srt.distributed.device_communicators.pynccl_allocator import (
|
||||||
|
use_symmetric_memory,
|
||||||
|
)
|
||||||
|
from sglang.srt.layers.dp_attention import is_allocation_symmetric
|
||||||
|
from sglang.srt.layers.moe.moe_runner.base import (
|
||||||
|
MoeQuantInfo,
|
||||||
|
MoeRunnerConfig,
|
||||||
|
register_fused_func,
|
||||||
|
)
|
||||||
|
from sglang.srt.layers.quantization.fp8_kernel import scaled_fp8_quant
|
||||||
|
from sglang.srt.utils import is_flashinfer_available
|
||||||
|
from sglang.srt.utils.common import next_power_of_2
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from sglang.srt.layers.moe.token_dispatcher.flashinfer import (
|
||||||
|
FlashinferCombineInput,
|
||||||
|
FlashinferDispatchOutput,
|
||||||
|
)
|
||||||
|
from sglang.srt.layers.moe.token_dispatcher.standard import (
|
||||||
|
StandardCombineInput,
|
||||||
|
StandardDispatchOutput,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class FlashInferCutlassMoeQuantInfo(MoeQuantInfo):
|
||||||
|
"""Payload for FlashInfer CUTLASS fused MoE.
|
||||||
|
|
||||||
|
``quant_type`` selects the input/weight conventions:
|
||||||
|
- ``"bf16"``: unquantized weights, BF16/FP16 input, no quant scales.
|
||||||
|
- ``"fp8"``: FP8 weights, FP8-quantized input, per-tensor scales.
|
||||||
|
- ``"fp4"``: NVFP4 packed weights and optional NVFP4 packed input.
|
||||||
|
"""
|
||||||
|
|
||||||
|
quant_type: str
|
||||||
|
w13_weight: torch.Tensor
|
||||||
|
w2_weight: torch.Tensor
|
||||||
|
quant_scales: Optional[list[torch.Tensor]] = None
|
||||||
|
output_dtype: Optional[torch.dtype] = None
|
||||||
|
moe_tp_size: int = 1
|
||||||
|
moe_tp_rank: int = 0
|
||||||
|
moe_ep_size: int = 1
|
||||||
|
moe_ep_rank: int = 0
|
||||||
|
apply_routed_scaling_factor: bool = True
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class FlashInferCutlassMxfp4MoeQuantInfo(MoeQuantInfo):
|
||||||
|
"""Quantization payload for the SM90 CUTLASS W4A16 MXFP4 MoE path.
|
||||||
|
|
||||||
|
Weights and scales are pre-interleaved at load time via
|
||||||
|
``interleave_moe_{weights,scales}_for_sm90_mixed_gemm``; this dataclass
|
||||||
|
only carries references plus the per-call routing/topology fields.
|
||||||
|
"""
|
||||||
|
|
||||||
|
# Pre-interleaved weights (uint8, packed FP4)
|
||||||
|
w13_weight: torch.Tensor # [E, 2*N, K/2]
|
||||||
|
w2_weight: torch.Tensor # [E, K, N/2]
|
||||||
|
|
||||||
|
# Pre-interleaved E8M0 block scales (uint8; viewed as int32 at call time)
|
||||||
|
w13_weight_scale: torch.Tensor # [E, 2*N, K/32]
|
||||||
|
w2_weight_scale: torch.Tensor # [E, K, N/32]
|
||||||
|
|
||||||
|
# Per-expert bias. GPT-OSS has both; DSv4 leaves both None.
|
||||||
|
w13_bias: Optional[torch.Tensor] = None # bf16 [E, 2*N]
|
||||||
|
w2_bias: Optional[torch.Tensor] = None # bf16 [E, K]
|
||||||
|
|
||||||
|
# Per-expert SwiGLU scalars (fp32 [E]). Either all three are present
|
||||||
|
# (clamped SwiGLU) or all three are None (kernel default SwiGLU).
|
||||||
|
swiglu_alpha: Optional[torch.Tensor] = None
|
||||||
|
swiglu_beta: Optional[torch.Tensor] = None
|
||||||
|
swiglu_limit: Optional[torch.Tensor] = None
|
||||||
|
|
||||||
|
# TP/EP topology (forwarded to the FlashInfer kernel)
|
||||||
|
moe_tp_size: int = 1
|
||||||
|
moe_tp_rank: int = 0
|
||||||
|
moe_ep_size: int = 1
|
||||||
|
moe_ep_rank: int = 0
|
||||||
|
|
||||||
|
# GPT-OSS pads its input hidden dim up to the (pre-padded) loaded weight
|
||||||
|
# width and trims the output back. DSv4 leaves this as ``None`` (no pad).
|
||||||
|
padded_hidden: Optional[int] = None
|
||||||
|
|
||||||
|
|
||||||
|
def _flashinfer_cutlass_fused_moe():
|
||||||
|
if not is_flashinfer_available():
|
||||||
|
raise RuntimeError(
|
||||||
|
"flashinfer_cutlass MoE runner backend requires flashinfer to be installed."
|
||||||
|
)
|
||||||
|
from flashinfer.fused_moe import cutlass_fused_moe
|
||||||
|
from flashinfer.fused_moe.core import ActivationType
|
||||||
|
|
||||||
|
return cutlass_fused_moe, ActivationType
|
||||||
|
|
||||||
|
|
||||||
|
def _activation_type(runner_config: MoeRunnerConfig):
|
||||||
|
from sglang.srt.layers.moe.moe_runner.flashinfer_trtllm import get_activation_type
|
||||||
|
|
||||||
|
_, ActivationType = _flashinfer_cutlass_fused_moe()
|
||||||
|
activation = ActivationType(
|
||||||
|
get_activation_type(
|
||||||
|
runner_config.activation,
|
||||||
|
is_gated=runner_config.is_gated,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
supported = {
|
||||||
|
ActivationType.Swiglu,
|
||||||
|
ActivationType.Geglu,
|
||||||
|
ActivationType.Relu2,
|
||||||
|
ActivationType.Identity,
|
||||||
|
}
|
||||||
|
assert activation in supported, (
|
||||||
|
f"Activation {runner_config.activation!r} "
|
||||||
|
f"(is_gated={runner_config.is_gated}) maps to {activation.name}, "
|
||||||
|
"which is not supported by flashinfer cutlass moe."
|
||||||
|
)
|
||||||
|
return activation
|
||||||
|
|
||||||
|
|
||||||
|
def _maybe_apply_routed_scaling_factor(
|
||||||
|
output: torch.Tensor,
|
||||||
|
quant_info: FlashInferCutlassMoeQuantInfo,
|
||||||
|
runner_config: MoeRunnerConfig,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
if (
|
||||||
|
quant_info.apply_routed_scaling_factor
|
||||||
|
and runner_config.routed_scaling_factor is not None
|
||||||
|
):
|
||||||
|
output.mul_(runner_config.routed_scaling_factor)
|
||||||
|
return output
|
||||||
|
|
||||||
|
|
||||||
|
def _prepare_input(
|
||||||
|
dispatch_output,
|
||||||
|
quant_info: FlashInferCutlassMoeQuantInfo,
|
||||||
|
runner_config: MoeRunnerConfig,
|
||||||
|
) -> tuple[torch.Tensor, Optional[torch.Tensor], torch.dtype, int]:
|
||||||
|
x = dispatch_output.hidden_states
|
||||||
|
x_sf = dispatch_output.hidden_states_scale
|
||||||
|
|
||||||
|
if quant_info.quant_type == "fp8":
|
||||||
|
assert quant_info.quant_scales is not None and len(quant_info.quant_scales) == 4
|
||||||
|
x, _ = scaled_fp8_quant(x, quant_info.quant_scales[3])
|
||||||
|
x_sf = None
|
||||||
|
output_dtype = quant_info.output_dtype or dispatch_output.hidden_states.dtype
|
||||||
|
output_col = dispatch_output.hidden_states.shape[1]
|
||||||
|
elif quant_info.quant_type == "fp4":
|
||||||
|
output_dtype = quant_info.output_dtype or torch.bfloat16
|
||||||
|
output_col = x.shape[1]
|
||||||
|
if x_sf is not None and runner_config.is_gated:
|
||||||
|
output_col *= 2
|
||||||
|
else:
|
||||||
|
assert quant_info.quant_type == "bf16"
|
||||||
|
output_dtype = quant_info.output_dtype or x.dtype
|
||||||
|
output_col = x.shape[1]
|
||||||
|
|
||||||
|
return x, x_sf, output_dtype, output_col
|
||||||
|
|
||||||
|
|
||||||
|
def _run_flashinfer_cutlass(
|
||||||
|
*,
|
||||||
|
dispatch_output,
|
||||||
|
quant_info: FlashInferCutlassMoeQuantInfo,
|
||||||
|
runner_config: MoeRunnerConfig,
|
||||||
|
output: Optional[torch.Tensor] = None,
|
||||||
|
enable_alltoall: bool = False,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
flashinfer_cutlass_fused_moe, _ = _flashinfer_cutlass_fused_moe()
|
||||||
|
|
||||||
|
topk_output = dispatch_output.topk_output
|
||||||
|
topk_weights = topk_output.topk_weights
|
||||||
|
topk_ids = topk_output.topk_ids
|
||||||
|
x, x_sf, output_dtype, output_col = _prepare_input(
|
||||||
|
dispatch_output, quant_info, runner_config
|
||||||
|
)
|
||||||
|
|
||||||
|
if output is None:
|
||||||
|
with use_symmetric_memory(
|
||||||
|
get_tp_group(), disabled=not is_allocation_symmetric()
|
||||||
|
):
|
||||||
|
output = torch.empty(
|
||||||
|
x.shape[0],
|
||||||
|
output_col,
|
||||||
|
dtype=output_dtype,
|
||||||
|
device=x.device,
|
||||||
|
)
|
||||||
|
|
||||||
|
w13_weight = quant_info.w13_weight
|
||||||
|
w2_weight = quant_info.w2_weight
|
||||||
|
quant_scales = quant_info.quant_scales
|
||||||
|
if quant_info.quant_type == "fp4":
|
||||||
|
w13_weight = w13_weight.view(torch.long)
|
||||||
|
w2_weight = w2_weight.view(torch.long)
|
||||||
|
assert quant_scales is not None and len(quant_scales) == 6
|
||||||
|
quant_scales = [
|
||||||
|
quant_scales[0],
|
||||||
|
quant_scales[1].view(torch.int32),
|
||||||
|
quant_scales[2],
|
||||||
|
quant_scales[3],
|
||||||
|
quant_scales[4].view(torch.int32),
|
||||||
|
quant_scales[5],
|
||||||
|
]
|
||||||
|
|
||||||
|
output = flashinfer_cutlass_fused_moe(
|
||||||
|
output=output,
|
||||||
|
input=x,
|
||||||
|
token_selected_experts=topk_ids.to(torch.int),
|
||||||
|
token_final_scales=topk_weights,
|
||||||
|
fc1_expert_weights=w13_weight,
|
||||||
|
fc2_expert_weights=w2_weight,
|
||||||
|
output_dtype=output_dtype,
|
||||||
|
input_sf=x_sf,
|
||||||
|
quant_scales=quant_scales,
|
||||||
|
ep_size=quant_info.moe_ep_size,
|
||||||
|
ep_rank=quant_info.moe_ep_rank,
|
||||||
|
tp_size=quant_info.moe_tp_size,
|
||||||
|
tp_rank=quant_info.moe_tp_rank,
|
||||||
|
tune_max_num_tokens=next_power_of_2(x.shape[0]),
|
||||||
|
activation_type=_activation_type(runner_config),
|
||||||
|
enable_alltoall=enable_alltoall,
|
||||||
|
)[0]
|
||||||
|
|
||||||
|
if quant_info.quant_type in ("bf16", "fp8"):
|
||||||
|
_maybe_apply_routed_scaling_factor(output, quant_info, runner_config)
|
||||||
|
return output
|
||||||
|
|
||||||
|
|
||||||
|
@register_fused_func("none", "flashinfer_cutlass")
|
||||||
|
def fused_experts_none_to_flashinfer_cutlass(
|
||||||
|
dispatch_output: StandardDispatchOutput,
|
||||||
|
quant_info: MoeQuantInfo,
|
||||||
|
runner_config: MoeRunnerConfig,
|
||||||
|
) -> StandardCombineInput:
|
||||||
|
from sglang.srt.layers.moe.token_dispatcher.standard import StandardCombineInput
|
||||||
|
|
||||||
|
assert isinstance(
|
||||||
|
quant_info, FlashInferCutlassMoeQuantInfo
|
||||||
|
), f"Unexpected quant_info type for flashinfer_cutlass: {type(quant_info)}"
|
||||||
|
assert (
|
||||||
|
not runner_config.apply_router_weight_on_input
|
||||||
|
), "apply_router_weight_on_input is not supported for FlashInfer CUTLASS"
|
||||||
|
|
||||||
|
output = _run_flashinfer_cutlass(
|
||||||
|
dispatch_output=dispatch_output,
|
||||||
|
quant_info=quant_info,
|
||||||
|
runner_config=runner_config,
|
||||||
|
)
|
||||||
|
return StandardCombineInput(hidden_states=output)
|
||||||
|
|
||||||
|
|
||||||
|
@register_fused_func("flashinfer", "flashinfer_cutlass")
|
||||||
|
def fused_experts_flashinfer_to_flashinfer_cutlass(
|
||||||
|
dispatch_output: FlashinferDispatchOutput,
|
||||||
|
quant_info: MoeQuantInfo,
|
||||||
|
runner_config: MoeRunnerConfig,
|
||||||
|
) -> FlashinferCombineInput:
|
||||||
|
from sglang.srt.layers.moe.token_dispatcher.flashinfer import (
|
||||||
|
FlashinferCombineInput,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert isinstance(
|
||||||
|
quant_info, FlashInferCutlassMoeQuantInfo
|
||||||
|
), f"Unexpected quant_info type for flashinfer_cutlass: {type(quant_info)}"
|
||||||
|
assert (
|
||||||
|
not runner_config.apply_router_weight_on_input
|
||||||
|
), "apply_router_weight_on_input is not supported for FlashInfer CUTLASS"
|
||||||
|
|
||||||
|
output = _run_flashinfer_cutlass(
|
||||||
|
dispatch_output=dispatch_output,
|
||||||
|
quant_info=quant_info,
|
||||||
|
runner_config=runner_config,
|
||||||
|
output=dispatch_output.moe_output,
|
||||||
|
enable_alltoall=True,
|
||||||
|
)
|
||||||
|
return FlashinferCombineInput(hidden_states=output)
|
||||||
|
|
||||||
|
|
||||||
|
@register_fused_func("none", "flashinfer_mxfp4")
|
||||||
|
def fused_experts_none_to_flashinfer_mxfp4(
|
||||||
|
dispatch_output: StandardDispatchOutput,
|
||||||
|
quant_info: MoeQuantInfo,
|
||||||
|
runner_config: MoeRunnerConfig,
|
||||||
|
) -> StandardCombineInput:
|
||||||
|
"""SM90 W4A16 MXFP4 fused expert forward pass.
|
||||||
|
|
||||||
|
This preserves the ``flashinfer_mxfp4`` runner backend registration while
|
||||||
|
centralizing the CUTLASS execution in this module.
|
||||||
|
"""
|
||||||
|
from sglang.srt.layers.moe.token_dispatcher.standard import StandardCombineInput
|
||||||
|
from sglang.srt.layers.moe.topk import TopKOutputChecker
|
||||||
|
|
||||||
|
assert isinstance(
|
||||||
|
quant_info, FlashInferCutlassMxfp4MoeQuantInfo
|
||||||
|
), f"Unexpected quant_info type for flashinfer_mxfp4: {type(quant_info)}"
|
||||||
|
|
||||||
|
flashinfer_cutlass_fused_moe, ActivationType = _flashinfer_cutlass_fused_moe()
|
||||||
|
|
||||||
|
x = dispatch_output.hidden_states
|
||||||
|
topk_output = dispatch_output.topk_output
|
||||||
|
|
||||||
|
# Under ``--moe-runner-backend flashinfer_mxfp4`` topk may be in bypassed
|
||||||
|
# form (the SM100 trtllm-gen path does routing internally). The CUTLASS
|
||||||
|
# SM90 path needs explicit topk_ids / topk_weights; materialize here.
|
||||||
|
if TopKOutputChecker.format_is_bypassed(topk_output):
|
||||||
|
topk_output = topk_output.to_standard()
|
||||||
|
topk_ids = topk_output.topk_ids
|
||||||
|
topk_weights = topk_output.topk_weights
|
||||||
|
|
||||||
|
# GPT-OSS: pad input hidden dim up to the loaded weight width. DSv4
|
||||||
|
# leaves padded_hidden as None (or equal to origin_hidden), no pad.
|
||||||
|
origin_hidden = x.shape[-1]
|
||||||
|
padded_hidden = quant_info.padded_hidden
|
||||||
|
do_pad = padded_hidden is not None and padded_hidden != origin_hidden
|
||||||
|
if do_pad:
|
||||||
|
x = torch.nn.functional.pad(
|
||||||
|
x,
|
||||||
|
(0, padded_hidden - origin_hidden),
|
||||||
|
mode="constant",
|
||||||
|
value=0.0,
|
||||||
|
)
|
||||||
|
|
||||||
|
out_hidden = padded_hidden if do_pad else origin_hidden
|
||||||
|
output_dtype = torch.bfloat16
|
||||||
|
with use_symmetric_memory(get_tp_group(), disabled=not is_allocation_symmetric()):
|
||||||
|
out = torch.empty(x.shape[0], out_hidden, dtype=output_dtype, device=x.device)
|
||||||
|
|
||||||
|
flashinfer_cutlass_fused_moe(
|
||||||
|
input=x,
|
||||||
|
token_selected_experts=topk_ids.to(torch.int),
|
||||||
|
token_final_scales=topk_weights,
|
||||||
|
fc1_expert_weights=quant_info.w13_weight,
|
||||||
|
fc2_expert_weights=quant_info.w2_weight,
|
||||||
|
output_dtype=output_dtype,
|
||||||
|
quant_scales=[
|
||||||
|
quant_info.w13_weight_scale.view(torch.int32),
|
||||||
|
quant_info.w2_weight_scale.view(torch.int32),
|
||||||
|
],
|
||||||
|
fc1_expert_biases=quant_info.w13_bias,
|
||||||
|
fc2_expert_biases=quant_info.w2_bias,
|
||||||
|
swiglu_alpha=quant_info.swiglu_alpha,
|
||||||
|
swiglu_beta=quant_info.swiglu_beta,
|
||||||
|
swiglu_limit=quant_info.swiglu_limit,
|
||||||
|
tp_size=quant_info.moe_tp_size,
|
||||||
|
tp_rank=quant_info.moe_tp_rank,
|
||||||
|
ep_size=quant_info.moe_ep_size,
|
||||||
|
ep_rank=quant_info.moe_ep_rank,
|
||||||
|
use_w4_group_scaling=True,
|
||||||
|
activation_type=ActivationType.Swiglu,
|
||||||
|
tune_max_num_tokens=next_power_of_2(x.shape[0]),
|
||||||
|
output=out,
|
||||||
|
)
|
||||||
|
|
||||||
|
if do_pad:
|
||||||
|
out = out[:, :origin_hidden].contiguous()
|
||||||
|
|
||||||
|
return StandardCombineInput(hidden_states=out)
|
||||||
@@ -1,174 +0,0 @@
|
|||||||
"""FlashInfer SM90 cutlass mixed-input W4A16 MXFP4 MoE fused func.
|
|
||||||
|
|
||||||
Registered for ``("none", "flashinfer_mxfp4")``. Drives FlashInfer's
|
|
||||||
``cutlass_fused_moe(use_w4_group_scaling=True)`` (PR #3084 in flashinfer,
|
|
||||||
SM90 only). Quant methods build the quant_info each forward and call
|
|
||||||
``MoeRunner.run(dispatch_output, quant_info)``.
|
|
||||||
|
|
||||||
Two production call sites share this fused func:
|
|
||||||
- GPT-OSS via :class:`Mxfp4MoEMethod` (input pad/output trim + per-expert
|
|
||||||
SwiGLU scalars + per-expert bias)
|
|
||||||
- DSv4 via :class:`Mxfp4FlashinferCutlassMoEMethod` (no bias, optional
|
|
||||||
SwiGLU scalars, no padding)
|
|
||||||
|
|
||||||
The SM100 trtllm-gen path also lives under ``MoeRunnerBackend.FLASHINFER_MXFP4``
|
|
||||||
but is intentionally left in the legacy bypass path for now; migrating it is a
|
|
||||||
follow-up.
|
|
||||||
"""
|
|
||||||
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
from dataclasses import dataclass
|
|
||||||
from typing import TYPE_CHECKING, Optional
|
|
||||||
|
|
||||||
import torch
|
|
||||||
|
|
||||||
from sglang.srt.distributed import get_tp_group
|
|
||||||
from sglang.srt.distributed.device_communicators.pynccl_allocator import (
|
|
||||||
use_symmetric_memory,
|
|
||||||
)
|
|
||||||
from sglang.srt.layers.dp_attention import is_allocation_symmetric
|
|
||||||
from sglang.srt.layers.moe.moe_runner.base import (
|
|
||||||
MoeQuantInfo,
|
|
||||||
MoeRunnerConfig,
|
|
||||||
register_fused_func,
|
|
||||||
)
|
|
||||||
from sglang.srt.utils import is_flashinfer_available
|
|
||||||
from sglang.srt.utils.common import next_power_of_2
|
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
|
||||||
from sglang.srt.layers.moe.token_dispatcher import StandardDispatchOutput
|
|
||||||
from sglang.srt.layers.moe.token_dispatcher.standard import StandardCombineInput
|
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
|
||||||
class FlashInferMxfp4CutlassMoeQuantInfo(MoeQuantInfo):
|
|
||||||
"""Quantization payload for the SM90 cutlass W4A16 MXFP4 MoE path.
|
|
||||||
|
|
||||||
Weights and scales are pre-interleaved at load time via
|
|
||||||
``interleave_moe_{weights,scales}_for_sm90_mixed_gemm``; this dataclass
|
|
||||||
only carries references plus the per-call routing/topology fields.
|
|
||||||
"""
|
|
||||||
|
|
||||||
# Pre-interleaved weights (uint8, packed FP4)
|
|
||||||
w13_weight: torch.Tensor # [E, 2*N, K/2]
|
|
||||||
w2_weight: torch.Tensor # [E, K, N/2]
|
|
||||||
|
|
||||||
# Pre-interleaved E8M0 block scales (uint8; viewed as int32 at call time)
|
|
||||||
w13_weight_scale: torch.Tensor # [E, 2*N, K/32]
|
|
||||||
w2_weight_scale: torch.Tensor # [E, K, N/32]
|
|
||||||
|
|
||||||
# Per-expert bias. GPT-OSS has both; DSv4 leaves both None.
|
|
||||||
w13_bias: Optional[torch.Tensor] = None # bf16 [E, 2*N]
|
|
||||||
w2_bias: Optional[torch.Tensor] = None # bf16 [E, K]
|
|
||||||
|
|
||||||
# Per-expert SwiGLU scalars (fp32 [E]). Either all three are present
|
|
||||||
# (clamped SwiGLU) or all three are None (kernel default SwiGLU).
|
|
||||||
swiglu_alpha: Optional[torch.Tensor] = None
|
|
||||||
swiglu_beta: Optional[torch.Tensor] = None
|
|
||||||
swiglu_limit: Optional[torch.Tensor] = None
|
|
||||||
|
|
||||||
# TP/EP topology (forwarded to the FlashInfer kernel)
|
|
||||||
moe_tp_size: int = 1
|
|
||||||
moe_tp_rank: int = 0
|
|
||||||
moe_ep_size: int = 1
|
|
||||||
moe_ep_rank: int = 0
|
|
||||||
|
|
||||||
# GPT-OSS pads its input hidden dim up to the (pre-padded) loaded weight
|
|
||||||
# width and trims the output back. DSv4 leaves this as ``None`` (no pad).
|
|
||||||
padded_hidden: Optional[int] = None
|
|
||||||
|
|
||||||
|
|
||||||
def _flashinfer_cutlass_fused_moe():
|
|
||||||
"""Lazy import — keeps non-flashinfer wheels importable."""
|
|
||||||
if not is_flashinfer_available():
|
|
||||||
raise RuntimeError(
|
|
||||||
"flashinfer_mxfp4 runner backend requires flashinfer to be installed."
|
|
||||||
)
|
|
||||||
from flashinfer.fused_moe import cutlass_fused_moe
|
|
||||||
from flashinfer.fused_moe.core import ActivationType
|
|
||||||
|
|
||||||
return cutlass_fused_moe, ActivationType
|
|
||||||
|
|
||||||
|
|
||||||
@register_fused_func("none", "flashinfer_mxfp4")
|
|
||||||
def fused_experts_none_to_flashinfer_mxfp4(
|
|
||||||
dispatch_output: StandardDispatchOutput,
|
|
||||||
quant_info: MoeQuantInfo,
|
|
||||||
runner_config: MoeRunnerConfig,
|
|
||||||
) -> StandardCombineInput:
|
|
||||||
"""SM90 W4A16 MXFP4 fused expert forward pass.
|
|
||||||
|
|
||||||
Mirrors the legacy ``Mxfp4MoEMethod._apply_sm90_cutlass`` and DSv4's
|
|
||||||
``Mxfp4FlashinferCutlassMoEMethod.apply`` exactly; difference vs those is
|
|
||||||
that all per-layer state arrives via ``quant_info`` rather than via the
|
|
||||||
layer module, so this function is layer-agnostic.
|
|
||||||
"""
|
|
||||||
from sglang.srt.layers.moe.token_dispatcher.standard import StandardCombineInput
|
|
||||||
from sglang.srt.layers.moe.topk import TopKOutputChecker
|
|
||||||
|
|
||||||
assert isinstance(
|
|
||||||
quant_info, FlashInferMxfp4CutlassMoeQuantInfo
|
|
||||||
), f"Unexpected quant_info type for flashinfer_mxfp4: {type(quant_info)}"
|
|
||||||
|
|
||||||
flashinfer_cutlass_fused_moe, ActivationType = _flashinfer_cutlass_fused_moe()
|
|
||||||
|
|
||||||
x = dispatch_output.hidden_states
|
|
||||||
topk_output = dispatch_output.topk_output
|
|
||||||
|
|
||||||
# Under ``--moe-runner-backend flashinfer_mxfp4`` topk may be in bypassed
|
|
||||||
# form (the SM100 trtllm-gen path does routing internally). The cutlass
|
|
||||||
# SM90 path needs explicit topk_ids / topk_weights; materialize here.
|
|
||||||
if TopKOutputChecker.format_is_bypassed(topk_output):
|
|
||||||
topk_output = topk_output.to_standard()
|
|
||||||
topk_ids = topk_output.topk_ids
|
|
||||||
topk_weights = topk_output.topk_weights
|
|
||||||
|
|
||||||
# GPT-OSS: pad input hidden dim up to the loaded weight width. DSv4
|
|
||||||
# leaves padded_hidden as None (or equal to origin_hidden), no pad.
|
|
||||||
origin_hidden = x.shape[-1]
|
|
||||||
padded_hidden = quant_info.padded_hidden
|
|
||||||
do_pad = padded_hidden is not None and padded_hidden != origin_hidden
|
|
||||||
if do_pad:
|
|
||||||
x = torch.nn.functional.pad(
|
|
||||||
x,
|
|
||||||
(0, padded_hidden - origin_hidden),
|
|
||||||
mode="constant",
|
|
||||||
value=0.0,
|
|
||||||
)
|
|
||||||
|
|
||||||
out_hidden = padded_hidden if do_pad else origin_hidden
|
|
||||||
output_dtype = torch.bfloat16
|
|
||||||
with use_symmetric_memory(get_tp_group(), disabled=not is_allocation_symmetric()):
|
|
||||||
out = torch.empty(x.shape[0], out_hidden, dtype=output_dtype, device=x.device)
|
|
||||||
|
|
||||||
flashinfer_cutlass_fused_moe(
|
|
||||||
input=x,
|
|
||||||
token_selected_experts=topk_ids.to(torch.int),
|
|
||||||
token_final_scales=topk_weights,
|
|
||||||
fc1_expert_weights=quant_info.w13_weight,
|
|
||||||
fc2_expert_weights=quant_info.w2_weight,
|
|
||||||
output_dtype=output_dtype,
|
|
||||||
quant_scales=[
|
|
||||||
quant_info.w13_weight_scale.view(torch.int32),
|
|
||||||
quant_info.w2_weight_scale.view(torch.int32),
|
|
||||||
],
|
|
||||||
fc1_expert_biases=quant_info.w13_bias,
|
|
||||||
fc2_expert_biases=quant_info.w2_bias,
|
|
||||||
swiglu_alpha=quant_info.swiglu_alpha,
|
|
||||||
swiglu_beta=quant_info.swiglu_beta,
|
|
||||||
swiglu_limit=quant_info.swiglu_limit,
|
|
||||||
tp_size=quant_info.moe_tp_size,
|
|
||||||
tp_rank=quant_info.moe_tp_rank,
|
|
||||||
ep_size=quant_info.moe_ep_size,
|
|
||||||
ep_rank=quant_info.moe_ep_rank,
|
|
||||||
use_w4_group_scaling=True,
|
|
||||||
activation_type=ActivationType.Swiglu,
|
|
||||||
tune_max_num_tokens=next_power_of_2(x.shape[0]),
|
|
||||||
output=out,
|
|
||||||
)
|
|
||||||
|
|
||||||
if do_pad:
|
|
||||||
out = out[:, :origin_hidden].contiguous()
|
|
||||||
|
|
||||||
return StandardCombineInput(hidden_states=out)
|
|
||||||
@@ -61,6 +61,8 @@ class MoeRunner:
|
|||||||
self.runner_core = None # FlashInfer TRT-LLM only supports fused path
|
self.runner_core = None # FlashInfer TRT-LLM only supports fused path
|
||||||
elif runner_backend.is_flashinfer_cutedsl():
|
elif runner_backend.is_flashinfer_cutedsl():
|
||||||
self.runner_core = None # FlashInfer CuteDSL only supports fused path
|
self.runner_core = None # FlashInfer CuteDSL only supports fused path
|
||||||
|
elif runner_backend.is_flashinfer_cutlass():
|
||||||
|
self.runner_core = None # FlashInfer CUTLASS only supports fused path
|
||||||
elif runner_backend.is_flashinfer_mxfp4():
|
elif runner_backend.is_flashinfer_mxfp4():
|
||||||
self.runner_core = None # FlashInfer MXFP4 only supports fused path
|
self.runner_core = None # FlashInfer MXFP4 only supports fused path
|
||||||
else:
|
else:
|
||||||
|
|||||||
@@ -4,24 +4,17 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import logging
|
import logging
|
||||||
from enum import IntEnum
|
|
||||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional
|
from typing import TYPE_CHECKING, Any, Dict, List, Optional
|
||||||
|
|
||||||
import regex as re
|
import regex as re
|
||||||
import torch
|
import torch
|
||||||
from torch.nn.parameter import Parameter
|
from torch.nn.parameter import Parameter
|
||||||
|
|
||||||
from sglang.srt.distributed import get_tp_group
|
|
||||||
from sglang.srt.distributed.device_communicators.pynccl_allocator import (
|
|
||||||
use_symmetric_memory,
|
|
||||||
)
|
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.layers.dp_attention import is_allocation_symmetric
|
|
||||||
from sglang.srt.layers.moe import (
|
from sglang.srt.layers.moe import (
|
||||||
MoeRunner,
|
MoeRunner,
|
||||||
MoeRunnerBackend,
|
MoeRunnerBackend,
|
||||||
MoeRunnerConfig,
|
MoeRunnerConfig,
|
||||||
get_moe_a2a_backend,
|
|
||||||
get_moe_runner_backend,
|
get_moe_runner_backend,
|
||||||
)
|
)
|
||||||
from sglang.srt.layers.moe.cutlass_moe_params import CutlassMoEParams, CutlassMoEType
|
from sglang.srt.layers.moe.cutlass_moe_params import CutlassMoEParams, CutlassMoEType
|
||||||
@@ -71,7 +64,6 @@ from sglang.srt.utils.common import (
|
|||||||
is_flashinfer_available,
|
is_flashinfer_available,
|
||||||
is_sm100_supported,
|
is_sm100_supported,
|
||||||
is_sm120_supported,
|
is_sm120_supported,
|
||||||
next_power_of_2,
|
|
||||||
round_up,
|
round_up,
|
||||||
)
|
)
|
||||||
from sglang.srt.utils.custom_op import register_custom_op
|
from sglang.srt.utils.custom_op import register_custom_op
|
||||||
@@ -104,20 +96,6 @@ if is_cuda():
|
|||||||
else:
|
else:
|
||||||
cutlass_fp4_gemm = None
|
cutlass_fp4_gemm = None
|
||||||
|
|
||||||
try:
|
|
||||||
from flashinfer.fused_moe import cutlass_fused_moe as flashinfer_cutlass_fused_moe
|
|
||||||
from flashinfer.fused_moe.core import ActivationType
|
|
||||||
except ImportError:
|
|
||||||
flashinfer_cutlass_fused_moe = None
|
|
||||||
|
|
||||||
# Define a minimal ActivationType enum if flashinfer is not available
|
|
||||||
class ActivationType(IntEnum):
|
|
||||||
Swiglu = 3
|
|
||||||
Geglu = 4
|
|
||||||
Relu2 = 6
|
|
||||||
Identity = 7
|
|
||||||
|
|
||||||
|
|
||||||
# Initialize logger for the module
|
# Initialize logger for the module
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -1031,6 +1009,14 @@ class ModelOptFp8MoEMethod(FusedMoEMethodBase):
|
|||||||
self, layer: torch.nn.Module, moe_runner_config: MoeRunnerConfig
|
self, layer: torch.nn.Module, moe_runner_config: MoeRunnerConfig
|
||||||
):
|
):
|
||||||
self.moe_runner_config = moe_runner_config
|
self.moe_runner_config = moe_runner_config
|
||||||
|
moe_runner_backend = get_moe_runner_backend()
|
||||||
|
if moe_runner_backend.is_flashinfer_cutlass():
|
||||||
|
import sglang.srt.layers.moe.moe_runner.flashinfer_cutlass # noqa: F401
|
||||||
|
|
||||||
|
self.runner = MoeRunner(
|
||||||
|
MoeRunnerBackend.FLASHINFER_CUTLASS, moe_runner_config
|
||||||
|
)
|
||||||
|
else:
|
||||||
self.runner = MoeRunner(MoeRunnerBackend.TRITON, moe_runner_config)
|
self.runner = MoeRunner(MoeRunnerBackend.TRITON, moe_runner_config)
|
||||||
|
|
||||||
def apply(
|
def apply(
|
||||||
@@ -1040,7 +1026,6 @@ class ModelOptFp8MoEMethod(FusedMoEMethodBase):
|
|||||||
) -> CombineInput:
|
) -> CombineInput:
|
||||||
x = dispatch_output.hidden_states
|
x = dispatch_output.hidden_states
|
||||||
topk_output = dispatch_output.topk_output
|
topk_output = dispatch_output.topk_output
|
||||||
from sglang.srt.layers.moe.token_dispatcher import StandardCombineInput
|
|
||||||
from sglang.srt.layers.moe.topk import TopKOutputChecker
|
from sglang.srt.layers.moe.topk import TopKOutputChecker
|
||||||
|
|
||||||
# Fast path: TRT-LLM FP8 per-tensor MoE using BYPASSED TopK routing
|
# Fast path: TRT-LLM FP8 per-tensor MoE using BYPASSED TopK routing
|
||||||
@@ -1052,14 +1037,9 @@ class ModelOptFp8MoEMethod(FusedMoEMethodBase):
|
|||||||
from sglang.srt.layers.moe.moe_runner.flashinfer_trtllm import (
|
from sglang.srt.layers.moe.moe_runner.flashinfer_trtllm import (
|
||||||
FlashInferTrtllmFp8MoeQuantInfo,
|
FlashInferTrtllmFp8MoeQuantInfo,
|
||||||
fused_experts_none_to_flashinfer_trtllm_fp8,
|
fused_experts_none_to_flashinfer_trtllm_fp8,
|
||||||
)
|
|
||||||
from sglang.srt.layers.moe.utils import RoutingMethodType
|
|
||||||
|
|
||||||
topk_config = topk_output.topk_config
|
|
||||||
|
|
||||||
from sglang.srt.layers.moe.moe_runner.flashinfer_trtllm import (
|
|
||||||
get_activation_type,
|
get_activation_type,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.layers.moe.utils import RoutingMethodType
|
||||||
|
|
||||||
_SUPPORTED_FP8_ACTIVATIONS = {"silu", "relu2"}
|
_SUPPORTED_FP8_ACTIVATIONS = {"silu", "relu2"}
|
||||||
assert self.moe_runner_config.activation in _SUPPORTED_FP8_ACTIVATIONS, (
|
assert self.moe_runner_config.activation in _SUPPORTED_FP8_ACTIVATIONS, (
|
||||||
@@ -1096,75 +1076,33 @@ class ModelOptFp8MoEMethod(FusedMoEMethodBase):
|
|||||||
)
|
)
|
||||||
|
|
||||||
if get_moe_runner_backend().is_flashinfer_cutlass():
|
if get_moe_runner_backend().is_flashinfer_cutlass():
|
||||||
from sglang.srt.layers.moe.moe_runner.flashinfer_trtllm import (
|
|
||||||
get_activation_type,
|
|
||||||
)
|
|
||||||
|
|
||||||
activation_str = self.moe_runner_config.activation
|
activation_str = self.moe_runner_config.activation
|
||||||
assert activation_str in _SUPPORTED_ACT_STRS, (
|
assert activation_str in _SUPPORTED_ACT_STRS, (
|
||||||
f"Activation {activation_str!r} is not supported for "
|
f"Activation {activation_str!r} is not supported for "
|
||||||
f"flashinfer cutlass fp8 moe (supported: {_SUPPORTED_ACT_STRS})."
|
f"flashinfer cutlass fp8 moe (supported: {_SUPPORTED_ACT_STRS})."
|
||||||
)
|
)
|
||||||
activation = ActivationType(
|
from sglang.srt.layers.moe.moe_runner.flashinfer_cutlass import (
|
||||||
get_activation_type(
|
FlashInferCutlassMoeQuantInfo,
|
||||||
activation_str, is_gated=self.moe_runner_config.is_gated
|
|
||||||
)
|
)
|
||||||
)
|
|
||||||
# FlashInfer CUTLASS MoE supports gated Swiglu/Geglu and non-gated
|
|
||||||
# Relu2/Identity. Non-gated Silu/Gelu are not implemented.
|
|
||||||
_CUTLASS_SUPPORTED = {
|
|
||||||
ActivationType.Swiglu,
|
|
||||||
ActivationType.Geglu,
|
|
||||||
ActivationType.Relu2,
|
|
||||||
ActivationType.Identity,
|
|
||||||
}
|
|
||||||
assert activation in _CUTLASS_SUPPORTED, (
|
|
||||||
f"Activation {activation_str!r} (is_gated="
|
|
||||||
f"{self.moe_runner_config.is_gated}) maps to {activation.name}, "
|
|
||||||
"which is not supported by flashinfer cutlass fp8 moe."
|
|
||||||
)
|
|
||||||
topk_weights, topk_ids = topk_output.topk_weights, topk_output.topk_ids
|
|
||||||
x_fp8, _ = scaled_fp8_quant(x, layer.w13_input_scale)
|
|
||||||
output_dtype = x.dtype
|
|
||||||
original_col = x.shape[1]
|
|
||||||
x_sf = None
|
|
||||||
|
|
||||||
with use_symmetric_memory(
|
quant_info = FlashInferCutlassMoeQuantInfo(
|
||||||
get_tp_group(), disabled=not is_allocation_symmetric()
|
quant_type="fp8",
|
||||||
):
|
w13_weight=layer.w13_weight,
|
||||||
symm_output = torch.empty(
|
w2_weight=layer.w2_weight,
|
||||||
x.shape[0], original_col, dtype=output_dtype, device=x.device
|
|
||||||
)
|
|
||||||
output = flashinfer_cutlass_fused_moe(
|
|
||||||
output=symm_output,
|
|
||||||
input=x_fp8,
|
|
||||||
token_selected_experts=topk_ids.to(torch.int),
|
|
||||||
token_final_scales=topk_weights,
|
|
||||||
fc1_expert_weights=layer.w13_weight,
|
|
||||||
fc2_expert_weights=layer.w2_weight,
|
|
||||||
output_dtype=output_dtype,
|
|
||||||
input_sf=x_sf,
|
|
||||||
quant_scales=[
|
quant_scales=[
|
||||||
layer.fc1_dequant,
|
layer.fc1_dequant,
|
||||||
layer.fc2_quant,
|
layer.fc2_quant,
|
||||||
layer.fc2_dequant,
|
layer.fc2_dequant,
|
||||||
layer.fc1_input_dequant,
|
layer.fc1_input_dequant,
|
||||||
],
|
],
|
||||||
ep_size=layer.moe_ep_size,
|
output_dtype=x.dtype,
|
||||||
ep_rank=layer.moe_ep_rank,
|
moe_ep_size=layer.moe_ep_size,
|
||||||
tp_size=layer.moe_tp_size,
|
moe_ep_rank=layer.moe_ep_rank,
|
||||||
tp_rank=layer.moe_tp_rank,
|
moe_tp_size=layer.moe_tp_size,
|
||||||
tune_max_num_tokens=next_power_of_2(x.shape[0]),
|
moe_tp_rank=layer.moe_tp_rank,
|
||||||
activation_type=activation,
|
apply_routed_scaling_factor=not layer.should_fuse_routed_scaling_factor_in_topk,
|
||||||
)[0]
|
)
|
||||||
|
return self.runner.run(dispatch_output, quant_info)
|
||||||
if (
|
|
||||||
not layer.should_fuse_routed_scaling_factor_in_topk
|
|
||||||
and self.moe_runner_config.routed_scaling_factor is not None
|
|
||||||
):
|
|
||||||
output.mul_(self.moe_runner_config.routed_scaling_factor)
|
|
||||||
|
|
||||||
return StandardCombineInput(hidden_states=output)
|
|
||||||
|
|
||||||
quant_info = TritonMoeQuantInfo(
|
quant_info = TritonMoeQuantInfo(
|
||||||
w13_weight=layer.w13_weight,
|
w13_weight=layer.w13_weight,
|
||||||
@@ -1320,9 +1258,9 @@ class ModelOptFp4Config(ModelOptQuantConfig):
|
|||||||
"Expected either flat format (config.json) or nested format (hf_quant_config.json)."
|
"Expected either flat format (config.json) or nested format (hf_quant_config.json)."
|
||||||
)
|
)
|
||||||
|
|
||||||
if not quant_method in ["FP8", "NVFP4"]:
|
if quant_method not in ["FP8", "NVFP4"]:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"ModelOpt currently only supports: FP8, NVFP4"
|
"ModelOpt currently only supports: FP8, NVFP4"
|
||||||
" quantizations in sglang. Please check the "
|
" quantizations in sglang. Please check the "
|
||||||
"quantization config for your model's configuration."
|
"quantization config for your model's configuration."
|
||||||
)
|
)
|
||||||
@@ -2327,7 +2265,9 @@ class ModelOptNvFp4FusedMoEMethod(FusedMoEMethodBase):
|
|||||||
if moe_runner_backend.is_flashinfer_cutedsl():
|
if moe_runner_backend.is_flashinfer_cutedsl():
|
||||||
import sglang.srt.layers.moe.moe_runner.flashinfer_cutedsl # noqa: F401 – triggers @register_fused_func
|
import sglang.srt.layers.moe.moe_runner.flashinfer_cutedsl # noqa: F401 – triggers @register_fused_func
|
||||||
|
|
||||||
if not moe_runner_backend.is_flashinfer_cutlass():
|
if moe_runner_backend.is_flashinfer_cutlass():
|
||||||
|
import sglang.srt.layers.moe.moe_runner.flashinfer_cutlass # noqa: F401
|
||||||
|
|
||||||
self.runner = MoeRunner(moe_runner_backend, moe_runner_config)
|
self.runner = MoeRunner(moe_runner_backend, moe_runner_config)
|
||||||
|
|
||||||
def apply(
|
def apply(
|
||||||
@@ -2458,85 +2398,33 @@ class ModelOptNvFp4FusedMoEMethod(FusedMoEMethodBase):
|
|||||||
return self.runner.run(dispatch_output, quant_info)
|
return self.runner.run(dispatch_output, quant_info)
|
||||||
|
|
||||||
if self.enable_flashinfer_cutlass_moe:
|
if self.enable_flashinfer_cutlass_moe:
|
||||||
from sglang.srt.layers.moe.moe_runner.flashinfer_trtllm import (
|
from sglang.srt.layers.moe.moe_runner.flashinfer_cutlass import (
|
||||||
get_activation_type,
|
FlashInferCutlassMoeQuantInfo,
|
||||||
)
|
)
|
||||||
from sglang.srt.layers.moe.token_dispatcher import DispatchOutputChecker
|
|
||||||
|
|
||||||
assert (
|
assert (
|
||||||
not moe_runner_config.apply_router_weight_on_input
|
not moe_runner_config.apply_router_weight_on_input
|
||||||
), "apply_router_weight_on_input is not supported for Flashinfer"
|
), "apply_router_weight_on_input is not supported for Flashinfer"
|
||||||
# Resolve the FlashInfer ActivationType honoring the gated flag,
|
quant_info = FlashInferCutlassMoeQuantInfo(
|
||||||
# then verify the CUTLASS FP4 kernel supports it.
|
quant_type="fp4",
|
||||||
fi_activation = ActivationType(
|
w13_weight=layer.w13_weight,
|
||||||
get_activation_type(activation, is_gated=moe_runner_config.is_gated)
|
w2_weight=layer.w2_weight,
|
||||||
)
|
output_dtype=torch.bfloat16,
|
||||||
_CUTLASS_FP4_SUPPORTED = {
|
|
||||||
ActivationType.Swiglu,
|
|
||||||
ActivationType.Geglu,
|
|
||||||
ActivationType.Relu2,
|
|
||||||
ActivationType.Identity,
|
|
||||||
}
|
|
||||||
assert fi_activation in _CUTLASS_FP4_SUPPORTED, (
|
|
||||||
f"Activation {activation!r} (is_gated={moe_runner_config.is_gated}) "
|
|
||||||
f"maps to {fi_activation.name}, which is not supported by the "
|
|
||||||
"flashinfer cutlass fp4 moe kernel."
|
|
||||||
)
|
|
||||||
# TRTLLM Cutlass moe takes in activations in BF16/Half/nvfp4 precision
|
|
||||||
# and fp4 quantized weights loaded from the checkpoint
|
|
||||||
x = dispatch_output.hidden_states
|
|
||||||
x_sf = dispatch_output.hidden_states_scale
|
|
||||||
topk_output = dispatch_output.topk_output
|
|
||||||
topk_weights, topk_ids = topk_output.topk_weights, topk_output.topk_ids
|
|
||||||
|
|
||||||
output_dtype = torch.bfloat16
|
|
||||||
|
|
||||||
if DispatchOutputChecker.format_is_flashinfer(dispatch_output):
|
|
||||||
symm_output = dispatch_output.moe_output
|
|
||||||
else:
|
|
||||||
# If x_sf is not None, x is FP4 packed (half size), so we need * 2
|
|
||||||
# If x_sf is None, x is not packed, so output_col = x.shape[1]
|
|
||||||
output_col = x.shape[1]
|
|
||||||
if x_sf is not None and layer.moe_runner_config.is_gated:
|
|
||||||
output_col *= 2
|
|
||||||
with use_symmetric_memory(
|
|
||||||
get_tp_group(), disabled=not is_allocation_symmetric()
|
|
||||||
):
|
|
||||||
symm_output = torch.empty(
|
|
||||||
x.shape[0],
|
|
||||||
output_col,
|
|
||||||
dtype=output_dtype,
|
|
||||||
device=x.device,
|
|
||||||
)
|
|
||||||
|
|
||||||
output = flashinfer_cutlass_fused_moe(
|
|
||||||
output=symm_output,
|
|
||||||
input=x,
|
|
||||||
token_selected_experts=topk_ids.to(torch.int),
|
|
||||||
token_final_scales=topk_weights,
|
|
||||||
fc1_expert_weights=layer.w13_weight.view(torch.long),
|
|
||||||
fc2_expert_weights=layer.w2_weight.view(torch.long),
|
|
||||||
output_dtype=output_dtype,
|
|
||||||
input_sf=x_sf,
|
|
||||||
# swizzled_input_sf intentionally omitted; not used for this path.
|
|
||||||
quant_scales=[
|
quant_scales=[
|
||||||
layer.w13_input_scale_quant,
|
layer.w13_input_scale_quant,
|
||||||
layer.w13_blockscale_swizzled.view(torch.int32),
|
layer.w13_blockscale_swizzled,
|
||||||
layer.g1_alphas,
|
layer.g1_alphas,
|
||||||
layer.w2_input_scale_quant,
|
layer.w2_input_scale_quant,
|
||||||
layer.w2_blockscale_swizzled.view(torch.int32),
|
layer.w2_blockscale_swizzled,
|
||||||
layer.g2_alphas,
|
layer.g2_alphas,
|
||||||
],
|
],
|
||||||
ep_size=layer.moe_ep_size,
|
moe_ep_size=layer.moe_ep_size,
|
||||||
ep_rank=layer.moe_ep_rank,
|
moe_ep_rank=layer.moe_ep_rank,
|
||||||
tp_size=layer.moe_tp_size,
|
moe_tp_size=layer.moe_tp_size,
|
||||||
tp_rank=layer.moe_tp_rank,
|
moe_tp_rank=layer.moe_tp_rank,
|
||||||
tune_max_num_tokens=next_power_of_2(x.shape[0]),
|
apply_routed_scaling_factor=False,
|
||||||
activation_type=fi_activation,
|
)
|
||||||
enable_alltoall=get_moe_a2a_backend().is_flashinfer(),
|
return self.runner.run(dispatch_output, quant_info)
|
||||||
)[0]
|
|
||||||
|
|
||||||
return StandardCombineInput(hidden_states=output)
|
|
||||||
|
|
||||||
from sglang.srt.layers.moe.cutlass_moe import cutlass_moe_fp4
|
from sglang.srt.layers.moe.cutlass_moe import cutlass_moe_fp4
|
||||||
|
|
||||||
|
|||||||
@@ -1039,7 +1039,7 @@ class Mxfp4MoEMethod(FusedMoEMethodBase):
|
|||||||
):
|
):
|
||||||
# Register the fused func at runner construction so the FusedOpPool
|
# Register the fused func at runner construction so the FusedOpPool
|
||||||
# lookup at `MoeRunner.__init__` finds it.
|
# lookup at `MoeRunner.__init__` finds it.
|
||||||
import sglang.srt.layers.moe.moe_runner.flashinfer_mxfp4 # noqa: F401
|
import sglang.srt.layers.moe.moe_runner.flashinfer_cutlass # noqa: F401
|
||||||
|
|
||||||
self.runner = MoeRunner(moe_runner_backend, moe_runner_config)
|
self.runner = MoeRunner(moe_runner_backend, moe_runner_config)
|
||||||
else:
|
else:
|
||||||
@@ -1051,12 +1051,12 @@ class Mxfp4MoEMethod(FusedMoEMethodBase):
|
|||||||
"""SM90 (Hopper) MXFP4 x BF16 MoE via FlashInfer's cutlass mixed-input
|
"""SM90 (Hopper) MXFP4 x BF16 MoE via FlashInfer's cutlass mixed-input
|
||||||
path (PR #3084). Routed through the unified ``MoeRunner`` -- this
|
path (PR #3084). Routed through the unified ``MoeRunner`` -- this
|
||||||
helper only builds the quant_info; the actual kernel call lives in
|
helper only builds the quant_info; the actual kernel call lives in
|
||||||
:mod:`sglang.srt.layers.moe.moe_runner.flashinfer_mxfp4`."""
|
:mod:`sglang.srt.layers.moe.moe_runner.flashinfer_cutlass`."""
|
||||||
from sglang.srt.layers.moe.moe_runner.flashinfer_mxfp4 import (
|
from sglang.srt.layers.moe.moe_runner.flashinfer_cutlass import (
|
||||||
FlashInferMxfp4CutlassMoeQuantInfo,
|
FlashInferCutlassMxfp4MoeQuantInfo,
|
||||||
)
|
)
|
||||||
|
|
||||||
quant_info = FlashInferMxfp4CutlassMoeQuantInfo(
|
quant_info = FlashInferCutlassMxfp4MoeQuantInfo(
|
||||||
w13_weight=layer.w13_weight,
|
w13_weight=layer.w13_weight,
|
||||||
w2_weight=layer.w2_weight,
|
w2_weight=layer.w2_weight,
|
||||||
w13_weight_scale=layer.w13_weight_scale,
|
w13_weight_scale=layer.w13_weight_scale,
|
||||||
|
|||||||
@@ -145,7 +145,7 @@ class Mxfp4FlashinferCutlassMoEMethod:
|
|||||||
|
|
||||||
# Register the fused func at runner construction so the FusedOpPool
|
# Register the fused func at runner construction so the FusedOpPool
|
||||||
# lookup at `MoeRunner.__init__` finds it.
|
# lookup at `MoeRunner.__init__` finds it.
|
||||||
import sglang.srt.layers.moe.moe_runner.flashinfer_mxfp4 # noqa: F401
|
import sglang.srt.layers.moe.moe_runner.flashinfer_cutlass # noqa: F401
|
||||||
|
|
||||||
self.runner = MoeRunner(MoeRunnerBackend.FLASHINFER_MXFP4, moe_runner_config)
|
self.runner = MoeRunner(MoeRunnerBackend.FLASHINFER_MXFP4, moe_runner_config)
|
||||||
|
|
||||||
@@ -217,8 +217,8 @@ class Mxfp4FlashinferCutlassMoEMethod:
|
|||||||
layer: Module,
|
layer: Module,
|
||||||
dispatch_output: DispatchOutput,
|
dispatch_output: DispatchOutput,
|
||||||
) -> CombineInput:
|
) -> CombineInput:
|
||||||
from sglang.srt.layers.moe.moe_runner.flashinfer_mxfp4 import (
|
from sglang.srt.layers.moe.moe_runner.flashinfer_cutlass import (
|
||||||
FlashInferMxfp4CutlassMoeQuantInfo,
|
FlashInferCutlassMxfp4MoeQuantInfo,
|
||||||
)
|
)
|
||||||
|
|
||||||
# DSv4 always feeds StandardDispatchOutput; the fused func tolerates
|
# DSv4 always feeds StandardDispatchOutput; the fused func tolerates
|
||||||
@@ -227,7 +227,7 @@ class Mxfp4FlashinferCutlassMoEMethod:
|
|||||||
if not TopKOutputChecker.format_is_standard(topk_output):
|
if not TopKOutputChecker.format_is_standard(topk_output):
|
||||||
raise ValueError(f"Unsupported topk output format: {topk_output.format}")
|
raise ValueError(f"Unsupported topk output format: {topk_output.format}")
|
||||||
|
|
||||||
quant_info = FlashInferMxfp4CutlassMoeQuantInfo(
|
quant_info = FlashInferCutlassMxfp4MoeQuantInfo(
|
||||||
w13_weight=layer.w13_weight,
|
w13_weight=layer.w13_weight,
|
||||||
w2_weight=layer.w2_weight,
|
w2_weight=layer.w2_weight,
|
||||||
w13_weight_scale=layer.w13_weight_scale_inv,
|
w13_weight_scale=layer.w13_weight_scale_inv,
|
||||||
|
|||||||
@@ -35,7 +35,6 @@ from sglang.srt.utils import (
|
|||||||
is_cpu,
|
is_cpu,
|
||||||
is_hip,
|
is_hip,
|
||||||
is_npu,
|
is_npu,
|
||||||
next_power_of_2,
|
|
||||||
set_weight_attrs,
|
set_weight_attrs,
|
||||||
use_intel_amx_backend,
|
use_intel_amx_backend,
|
||||||
use_intel_xpu_backend,
|
use_intel_xpu_backend,
|
||||||
@@ -62,12 +61,6 @@ if _use_aiter:
|
|||||||
if _is_npu:
|
if _is_npu:
|
||||||
from sglang.srt.hardware_backend.npu.utils import npu_format_cast
|
from sglang.srt.hardware_backend.npu.utils import npu_format_cast
|
||||||
|
|
||||||
try:
|
|
||||||
from flashinfer.fused_moe import cutlass_fused_moe as flashinfer_cutlass_fused_moe
|
|
||||||
from flashinfer.fused_moe.core import ActivationType
|
|
||||||
except ImportError:
|
|
||||||
flashinfer_cutlass_fused_moe = None
|
|
||||||
|
|
||||||
|
|
||||||
class UnquantizedEmbeddingMethod(QuantizeMethodBase):
|
class UnquantizedEmbeddingMethod(QuantizeMethodBase):
|
||||||
"""Unquantized method for embeddings."""
|
"""Unquantized method for embeddings."""
|
||||||
@@ -412,6 +405,10 @@ class UnquantizedFusedMoEMethod(FusedMoEMethodBase, MultiPlatformOp):
|
|||||||
if get_moe_runner_backend().is_flashinfer_trtllm_routed()
|
if get_moe_runner_backend().is_flashinfer_trtllm_routed()
|
||||||
else MoeRunnerBackend.FLASHINFER_TRTLLM
|
else MoeRunnerBackend.FLASHINFER_TRTLLM
|
||||||
)
|
)
|
||||||
|
elif self.use_flashinfer_cutlass:
|
||||||
|
import sglang.srt.layers.moe.moe_runner.flashinfer_cutlass # noqa: F401
|
||||||
|
|
||||||
|
backend = MoeRunnerBackend.FLASHINFER_CUTLASS
|
||||||
elif self.use_deep_gemm:
|
elif self.use_deep_gemm:
|
||||||
backend = MoeRunnerBackend.DEEP_GEMM
|
backend = MoeRunnerBackend.DEEP_GEMM
|
||||||
elif self.use_triton_kernels:
|
elif self.use_triton_kernels:
|
||||||
@@ -467,12 +464,8 @@ class UnquantizedFusedMoEMethod(FusedMoEMethodBase, MultiPlatformOp):
|
|||||||
layer: torch.nn.Module,
|
layer: torch.nn.Module,
|
||||||
dispatch_output: StandardDispatchOutput,
|
dispatch_output: StandardDispatchOutput,
|
||||||
) -> CombineInput:
|
) -> CombineInput:
|
||||||
from sglang.srt.layers.moe.token_dispatcher import StandardCombineInput
|
|
||||||
|
|
||||||
x = dispatch_output.hidden_states
|
x = dispatch_output.hidden_states
|
||||||
|
|
||||||
moe_runner_config = self.moe_runner_config
|
|
||||||
|
|
||||||
backend = self.runner.runner_backend
|
backend = self.runner.runner_backend
|
||||||
if backend.is_triton_kernels():
|
if backend.is_triton_kernels():
|
||||||
from sglang.srt.layers.moe.moe_runner.triton_kernels import (
|
from sglang.srt.layers.moe.moe_runner.triton_kernels import (
|
||||||
@@ -501,34 +494,22 @@ class UnquantizedFusedMoEMethod(FusedMoEMethodBase, MultiPlatformOp):
|
|||||||
)
|
)
|
||||||
return self.runner.run(dispatch_output, quant_info)
|
return self.runner.run(dispatch_output, quant_info)
|
||||||
elif self.use_flashinfer_cutlass:
|
elif self.use_flashinfer_cutlass:
|
||||||
topk_output = dispatch_output.topk_output
|
from sglang.srt.layers.moe.moe_runner.flashinfer_cutlass import (
|
||||||
output = flashinfer_cutlass_fused_moe(
|
FlashInferCutlassMoeQuantInfo,
|
||||||
input=x,
|
)
|
||||||
token_selected_experts=topk_output.topk_ids,
|
|
||||||
token_final_scales=topk_output.topk_weights,
|
quant_info = FlashInferCutlassMoeQuantInfo(
|
||||||
fc1_expert_weights=layer.w13_weight,
|
quant_type="bf16",
|
||||||
fc2_expert_weights=layer.w2_weight,
|
w13_weight=layer.w13_weight,
|
||||||
|
w2_weight=layer.w2_weight,
|
||||||
output_dtype=x.dtype,
|
output_dtype=x.dtype,
|
||||||
quant_scales=None,
|
moe_ep_size=layer.moe_ep_size,
|
||||||
ep_size=layer.moe_ep_size,
|
moe_ep_rank=layer.moe_ep_rank,
|
||||||
ep_rank=layer.moe_ep_rank,
|
moe_tp_size=layer.moe_tp_size,
|
||||||
tp_size=layer.moe_tp_size,
|
moe_tp_rank=layer.moe_tp_rank,
|
||||||
tp_rank=layer.moe_tp_rank,
|
apply_routed_scaling_factor=not layer.should_fuse_routed_scaling_factor_in_topk,
|
||||||
tune_max_num_tokens=next_power_of_2(x.shape[0]),
|
)
|
||||||
activation_type=(
|
return self.runner.run(dispatch_output, quant_info)
|
||||||
ActivationType.Relu2
|
|
||||||
if moe_runner_config.activation == "relu2"
|
|
||||||
else ActivationType.Swiglu
|
|
||||||
),
|
|
||||||
)[0]
|
|
||||||
|
|
||||||
if (
|
|
||||||
not layer.should_fuse_routed_scaling_factor_in_topk
|
|
||||||
and moe_runner_config.routed_scaling_factor is not None
|
|
||||||
):
|
|
||||||
output.mul_(moe_runner_config.routed_scaling_factor)
|
|
||||||
|
|
||||||
return StandardCombineInput(hidden_states=output)
|
|
||||||
elif self.use_flashinfer_trtllm_moe:
|
elif self.use_flashinfer_trtllm_moe:
|
||||||
from sglang.srt.layers.moe.moe_runner.flashinfer_trtllm import (
|
from sglang.srt.layers.moe.moe_runner.flashinfer_trtllm import (
|
||||||
FlashInferTrtllmBf16MoeQuantInfo,
|
FlashInferTrtllmBf16MoeQuantInfo,
|
||||||
|
|||||||
@@ -171,7 +171,7 @@ def _build_flashinfer_mxfp4_runner(num_experts, hidden, inter):
|
|||||||
and wires the runner with a minimal MoeRunnerConfig sufficient for the
|
and wires the runner with a minimal MoeRunnerConfig sufficient for the
|
||||||
cutlass SM90 fused func, which only reads dispatch_output / quant_info.
|
cutlass SM90 fused func, which only reads dispatch_output / quant_info.
|
||||||
"""
|
"""
|
||||||
import sglang.srt.layers.moe.moe_runner.flashinfer_mxfp4 # noqa: F401
|
import sglang.srt.layers.moe.moe_runner.flashinfer_cutlass # noqa: F401
|
||||||
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
|
||||||
from sglang.srt.layers.moe.utils import MoeRunnerBackend
|
from sglang.srt.layers.moe.utils import MoeRunnerBackend
|
||||||
@@ -325,7 +325,7 @@ def test_apply_sm90_cutlass_matches_flashinfer_direct(
|
|||||||
covered separately by ``test_process_weights_matches_direct_interleave``;
|
covered separately by ``test_process_weights_matches_direct_interleave``;
|
||||||
here we just verify that ``apply`` calls the kernel with the right
|
here we just verify that ``apply`` calls the kernel with the right
|
||||||
arguments (incl. input padding + output trim)."""
|
arguments (incl. input padding + output trim)."""
|
||||||
import sglang.srt.layers.moe.moe_runner.flashinfer_mxfp4 as fi_mxfp4_mod
|
import sglang.srt.layers.moe.moe_runner.flashinfer_cutlass as fi_cutlass_mod
|
||||||
import sglang.srt.layers.quantization.mxfp4 as mxfp4_mod
|
import sglang.srt.layers.quantization.mxfp4 as mxfp4_mod
|
||||||
|
|
||||||
# Bypass symmetric-memory / TP-group in both the legacy quant_method and
|
# Bypass symmetric-memory / TP-group in both the legacy quant_method and
|
||||||
@@ -336,10 +336,10 @@ def test_apply_sm90_cutlass_matches_flashinfer_direct(
|
|||||||
monkeypatch.setattr(mxfp4_mod, "is_allocation_symmetric", lambda: False)
|
monkeypatch.setattr(mxfp4_mod, "is_allocation_symmetric", lambda: False)
|
||||||
monkeypatch.setattr(mxfp4_mod, "get_tp_group", lambda: None)
|
monkeypatch.setattr(mxfp4_mod, "get_tp_group", lambda: None)
|
||||||
monkeypatch.setattr(
|
monkeypatch.setattr(
|
||||||
fi_mxfp4_mod, "use_symmetric_memory", lambda *a, **kw: nullcontext()
|
fi_cutlass_mod, "use_symmetric_memory", lambda *a, **kw: nullcontext()
|
||||||
)
|
)
|
||||||
monkeypatch.setattr(fi_mxfp4_mod, "is_allocation_symmetric", lambda: False)
|
monkeypatch.setattr(fi_cutlass_mod, "is_allocation_symmetric", lambda: False)
|
||||||
monkeypatch.setattr(fi_mxfp4_mod, "get_tp_group", lambda: None)
|
monkeypatch.setattr(fi_cutlass_mod, "get_tp_group", lambda: None)
|
||||||
|
|
||||||
w13, w2, w13_s, w2_s, w13_b, w2_b = _make_random_mxfp4(num_experts, hidden, inter)
|
w13, w2, w13_s, w2_s, w13_b, w2_b = _make_random_mxfp4(num_experts, hidden, inter)
|
||||||
x = torch.randn(tokens, hidden, dtype=torch.bfloat16, device="cuda") * 0.1
|
x = torch.randn(tokens, hidden, dtype=torch.bfloat16, device="cuda") * 0.1
|
||||||
@@ -463,17 +463,17 @@ def test_dsv4_apply_matches_flashinfer_direct(
|
|||||||
the equivalent reorder + scale-cast + interleave applied manually."""
|
the equivalent reorder + scale-cast + interleave applied manually."""
|
||||||
from types import SimpleNamespace
|
from types import SimpleNamespace
|
||||||
|
|
||||||
import sglang.srt.layers.moe.moe_runner.flashinfer_mxfp4 as fi_mxfp4_mod
|
import sglang.srt.layers.moe.moe_runner.flashinfer_cutlass as fi_cutlass_mod
|
||||||
import sglang.srt.layers.quantization.mxfp4_flashinfer_cutlass_moe as ds_mod
|
import sglang.srt.layers.quantization.mxfp4_flashinfer_cutlass_moe as ds_mod
|
||||||
from sglang.srt.layers.quantization.utils import reorder_w1w3_to_w3w1
|
from sglang.srt.layers.quantization.utils import reorder_w1w3_to_w3w1
|
||||||
|
|
||||||
# Bypass symmetric-memory / TP-group stack in the new fused-func module
|
# Bypass symmetric-memory / TP-group stack in the new fused-func module
|
||||||
# (where DSv4 ``apply`` now dispatches the kernel call through).
|
# (where DSv4 ``apply`` now dispatches the kernel call through).
|
||||||
monkeypatch.setattr(
|
monkeypatch.setattr(
|
||||||
fi_mxfp4_mod, "use_symmetric_memory", lambda *a, **kw: nullcontext()
|
fi_cutlass_mod, "use_symmetric_memory", lambda *a, **kw: nullcontext()
|
||||||
)
|
)
|
||||||
monkeypatch.setattr(fi_mxfp4_mod, "is_allocation_symmetric", lambda: False)
|
monkeypatch.setattr(fi_cutlass_mod, "is_allocation_symmetric", lambda: False)
|
||||||
monkeypatch.setattr(fi_mxfp4_mod, "get_tp_group", lambda: None)
|
monkeypatch.setattr(fi_cutlass_mod, "get_tp_group", lambda: None)
|
||||||
|
|
||||||
w13, w2, w13_s, w2_s = _make_random_dsv4_mxfp4(num_experts, hidden, inter)
|
w13, w2, w13_s, w2_s = _make_random_dsv4_mxfp4(num_experts, hidden, inter)
|
||||||
x = torch.randn(tokens, hidden, dtype=torch.bfloat16, device="cuda") * 0.1
|
x = torch.randn(tokens, hidden, dtype=torch.bfloat16, device="cuda") * 0.1
|
||||||
|
|||||||
Reference in New Issue
Block a user