[NVIDIA] Support flashinfer Mega Moe (#31470)
Co-authored-by: djns99 <40156487+djns99@users.noreply.github.com> Co-authored-by: 云挚 <ningyunxiao.nyx@antgroup.com> Co-authored-by: Yangmin Li <yangminl@nvidia.com> Co-authored-by: Po-Han Huang (NVIDIA) <53919306+nvpohanh@users.noreply.github.com>
This commit is contained in:
co-authored by
djns99
云挚
Yangmin Li
Po-Han Huang
parent
c0b790cf7f
commit
1b77f498a0
@@ -1818,7 +1818,13 @@ Please consult the documentation below and [server_args.py](https://github.com/s
|
||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>`auto`</td>
|
||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}><code>auto</code>, <code>bf16</code>, <code>fp8</code>, <code>int8</code>, <code>nvfp4</code></td>
|
||||
</tr>
|
||||
<tr>
|
||||
<tr>
|
||||
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><code>--flashinfer-a2a-dispatch-type</code></td>
|
||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Select FlashInfer A2A dispatcher activation dtype. When omitted, it falls back to the <code>SGLANG_MOE_NVFP4_DISPATCH</code> environment variable. Explicit <code>auto</code> selects <code>mxfp8</code> for <code>--quantization mxfp8</code>, <code>nvfp4</code> for modelopt FP4 or hybrid NVFP4 MoE checkpoints, and <code>bf16</code> otherwise. The <code>SGLANG_MOE_NVFP4_DISPATCH</code> environment variable cannot be set with this argument.</td>
|
||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}><code>None</code></td>
|
||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}><code>auto</code>, <code>bf16</code>, <code>nvfp4</code>, <code>mxfp8</code></td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}>`--ep-num-redundant-experts`</td>
|
||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Allocate this number of redundant experts in expert parallel.</td>
|
||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>`0`</td>
|
||||
|
||||
@@ -326,6 +326,21 @@ SGLang supports various environment variables that can be used to configure its
|
||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>The maximum number of dispatched tokens on each GPU for --moe-a2a-backend=flashinfer</td>
|
||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}>`"1024"`</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><code>SGLANG_FLASHINFER_MEGAMOE_COMBINE_DTYPE</code></td>
|
||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>FlashInfer NVFP4 MegaMOE cross-rank combine wire format. Supported values are <code>bf16</code>, <code>mxfp8</code>, and <code>nvfp4</code>. Quantized formats are incompatible with <code>SGLANG_FLASHINFER_MEGAMOE_IN_KERNEL_FC2_REDUCE=1</code>.</td>
|
||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}><code>"bf16"</code></td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><code>SGLANG_FLASHINFER_MEGAMOE_MAX_TOKENS_PER_RANK</code></td>
|
||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Per-rank FlashInfer MegaMOE symmetric-workspace token capacity. A value of <code>0</code> derives the capacity from the runtime token limits.</td>
|
||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}><code>0</code></td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}><code>SGLANG_FLASHINFER_MEGAMOE_IN_KERNEL_FC2_REDUCE</code></td>
|
||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Use the CuTe DSL MegaMOE in-kernel FC2 reduction. This can reduce workspace size and improve large-batch performance, but BF16 atomic accumulation is nondeterministic. It is incompatible with quantized combine dtypes.</td>
|
||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.02)"}}><code>"false"</code></td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td style={{padding: "9px 12px", fontWeight: 500, backgroundColor: "rgba(255,255,255,0.02)"}}>`SGLANG_DEEPEP_LL_COMBINE_SEND_NUM_SMS`</td>
|
||||
<td style={{padding: "9px 12px", backgroundColor: "rgba(255,255,255,0.05)"}}>Number of SMs used for DeepEP combine when single batch overlap is enabled</td>
|
||||
|
||||
@@ -55,6 +55,7 @@ dependencies = [
|
||||
"nvidia-cutlass-dsl[cu13]==4.6.2",
|
||||
"nvidia-mathdx==25.6.0",
|
||||
"nvidia-ml-py",
|
||||
"nvshmem4py-cu13",
|
||||
"openai==2.6.1",
|
||||
"openai-harmony==0.0.4",
|
||||
"orjson",
|
||||
|
||||
@@ -146,6 +146,7 @@ MOE_RUNNER_BACKEND_CHOICES = [
|
||||
"flashinfer_cutlass",
|
||||
"flashinfer_mxfp4",
|
||||
"flashinfer_cutedsl",
|
||||
"flashinfer_megamoe",
|
||||
"cutlass",
|
||||
"aiter",
|
||||
"marlin",
|
||||
@@ -159,6 +160,7 @@ MOE_RUNNER_BACKEND_CHOICES = [
|
||||
MXFP8_MOE_RUNNER_BACKEND_CHOICES = [
|
||||
"cutlass",
|
||||
"deep_gemm",
|
||||
"flashinfer_megamoe",
|
||||
"flashinfer_trtllm",
|
||||
"flashinfer_trtllm_routed",
|
||||
]
|
||||
|
||||
@@ -636,6 +636,7 @@ class ExecMoe(msgspec.Struct):
|
||||
"deepep_v2",
|
||||
"ascend_tp",
|
||||
"pplx",
|
||||
"flashinfer_megamoe",
|
||||
],
|
||||
Arg(
|
||||
help="Choose the backend for MoE A2A.",
|
||||
@@ -651,6 +652,7 @@ class ExecMoe(msgspec.Struct):
|
||||
"deepep_v2",
|
||||
"pplx",
|
||||
"ascend_tp",
|
||||
"flashinfer_megamoe",
|
||||
],
|
||||
resolvable=True,
|
||||
),
|
||||
@@ -695,6 +697,10 @@ class ExecMoe(msgspec.Struct):
|
||||
Literal["auto", "bf16", "fp8", "int8", "nvfp4"],
|
||||
"Select DeepEP dispatcher output dtype",
|
||||
] = "auto"
|
||||
flashinfer_a2a_dispatch_type: A[
|
||||
Optional[Literal["auto", "bf16", "nvfp4", "mxfp8"]],
|
||||
"Select FlashInfer A2A dispatcher activation dtype.",
|
||||
] = None
|
||||
ep_num_redundant_experts: A[
|
||||
int, "Allocate this number of redundant experts in expert parallel."
|
||||
] = 0
|
||||
|
||||
@@ -26,7 +26,7 @@ from sglang.srt.connector import ConnectorType
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.model_executor.cuda_graph_config import Backend, Phase, with_phase
|
||||
from sglang.srt.runtime_context import get_platform
|
||||
from sglang.srt.utils.common import parse_connector_type
|
||||
from sglang.srt.utils.common import is_sm100_supported, parse_connector_type
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -126,6 +126,110 @@ def handle_moe_kernel_config(server_args: Any):
|
||||
)
|
||||
|
||||
|
||||
def handle_flashinfer_a2a_dispatch_type(server_args: Any):
|
||||
cfg = resolving_view(server_args)
|
||||
cli_dispatch_type = cfg.flashinfer_a2a_dispatch_type
|
||||
nvfp4_dispatch_env_is_set = envs.SGLANG_MOE_NVFP4_DISPATCH.is_set()
|
||||
|
||||
if nvfp4_dispatch_env_is_set:
|
||||
raise ValueError(
|
||||
"SGLANG_MOE_NVFP4_DISPATCH cannot be set together with "
|
||||
"--flashinfer-a2a-dispatch-type."
|
||||
)
|
||||
|
||||
dispatch_type = cli_dispatch_type or "auto"
|
||||
|
||||
supports_nvfp4_dispatch = (
|
||||
cfg.quantization == "modelopt_fp4"
|
||||
or model_config_of(server_args).nvfp4_moe_meta is not None
|
||||
)
|
||||
if dispatch_type == "auto":
|
||||
if cfg.quantization == "mxfp8":
|
||||
dispatch_type = "mxfp8"
|
||||
elif supports_nvfp4_dispatch:
|
||||
dispatch_type = "nvfp4"
|
||||
else:
|
||||
dispatch_type = "bf16"
|
||||
|
||||
if dispatch_type == "mxfp8":
|
||||
if cfg.quantization != "mxfp8":
|
||||
raise ValueError(
|
||||
"--flashinfer-a2a-dispatch-type mxfp8 requires --quantization mxfp8."
|
||||
)
|
||||
if cfg.moe_runner_backend != "flashinfer_trtllm_routed":
|
||||
raise ValueError(
|
||||
"--flashinfer-a2a-dispatch-type mxfp8 requires "
|
||||
"--moe-runner-backend flashinfer_trtllm_routed."
|
||||
)
|
||||
elif dispatch_type == "nvfp4" and not supports_nvfp4_dispatch:
|
||||
raise ValueError(
|
||||
"--flashinfer-a2a-dispatch-type nvfp4 requires NVFP4/"
|
||||
"modelopt-FP4 quantization or hybrid NVFP4 MoE metadata."
|
||||
)
|
||||
|
||||
declare_resolution(
|
||||
server_args,
|
||||
"_handle_flashinfer_a2a_dispatch_type",
|
||||
flashinfer_a2a_dispatch_type=dispatch_type,
|
||||
)
|
||||
|
||||
|
||||
def validate_flashinfer_megamoe_envs() -> None:
|
||||
combine_dtype = envs.SGLANG_FLASHINFER_MEGAMOE_COMBINE_DTYPE.get().strip().lower()
|
||||
if combine_dtype not in ("bf16", "mxfp8", "nvfp4"):
|
||||
raise ValueError(
|
||||
"SGLANG_FLASHINFER_MEGAMOE_COMBINE_DTYPE must be one of "
|
||||
f"'bf16', 'mxfp8', or 'nvfp4', got {combine_dtype!r}."
|
||||
)
|
||||
if (
|
||||
combine_dtype != "bf16"
|
||||
and envs.SGLANG_FLASHINFER_MEGAMOE_IN_KERNEL_FC2_REDUCE.get()
|
||||
):
|
||||
raise ValueError(
|
||||
"SGLANG_FLASHINFER_MEGAMOE_COMBINE_DTYPE="
|
||||
f"{combine_dtype!r} is incompatible with "
|
||||
"SGLANG_FLASHINFER_MEGAMOE_IN_KERNEL_FC2_REDUCE=1."
|
||||
)
|
||||
|
||||
|
||||
def validate_flashinfer_megamoe_model(server_args: Any) -> None:
|
||||
model_config = model_config_of(server_args)
|
||||
architectures = model_config.hf_config.architectures or []
|
||||
validated_architectures = (
|
||||
"DeepseekV2ForCausalLM",
|
||||
"DeepseekV3ForCausalLM",
|
||||
"DeepseekV32ForCausalLM",
|
||||
"DeepseekV4ForCausalLM",
|
||||
"Glm4MoeForCausalLM",
|
||||
"NemotronHForCausalLM",
|
||||
"NemotronHPuzzleForCausalLM",
|
||||
"Qwen2MoeForCausalLM",
|
||||
"Qwen3MoeForCausalLM",
|
||||
)
|
||||
if not any(
|
||||
architecture in validated_architectures for architecture in architectures
|
||||
):
|
||||
raise ValueError(
|
||||
"FlashInfer MegaMOE is not validated for model architectures "
|
||||
f"{architectures}. Supported architectures: "
|
||||
f"{sorted(validated_architectures)}."
|
||||
)
|
||||
|
||||
quantization = resolved_view(server_args).quantization
|
||||
supports_megamoe_quantization = (
|
||||
quantization in ("mxfp8", "modelopt_fp4")
|
||||
or model_config.is_fp4_experts
|
||||
or model_config.nvfp4_moe_meta is not None
|
||||
)
|
||||
if not supports_megamoe_quantization:
|
||||
raise ValueError(
|
||||
"FlashInfer MegaMOE currently supports only MXFP8, ModelOpt "
|
||||
"NVFP4, FP4-expert, or hybrid NVFP4 MoE checkpoints; got "
|
||||
f"quantization={quantization!r}. Standard FP8 MoE checkpoints "
|
||||
"are not supported."
|
||||
)
|
||||
|
||||
|
||||
def handle_a2a_moe(server_args: Any):
|
||||
# The backend overrides and the ep_size=tp_size adjustments moved to
|
||||
# the resolution pipeline (arg_groups/overrides.py:
|
||||
@@ -149,6 +253,38 @@ def handle_a2a_moe(server_args: Any):
|
||||
)
|
||||
logger.info(f"Waterfill is enabled with moe_a2a_backend='{a2a_backend}'.")
|
||||
|
||||
if a2a_backend != "flashinfer" and cfg.flashinfer_a2a_dispatch_type not in (
|
||||
None,
|
||||
"auto",
|
||||
):
|
||||
raise ValueError(
|
||||
"--flashinfer-a2a-dispatch-type requires --moe-a2a-backend flashinfer."
|
||||
)
|
||||
|
||||
if a2a_backend == "flashinfer_megamoe":
|
||||
validate_flashinfer_megamoe_model(server_args)
|
||||
validate_flashinfer_megamoe_envs()
|
||||
assert cfg.enable_dp_attention and cfg.dp_size == cfg.tp_size, (
|
||||
"FlashInfer MegaMOE is only supported with dp_size == tp_size and --enable-dp-attention"
|
||||
)
|
||||
if resolved_view(server_args).moe_runner_backend == "auto":
|
||||
declare_resolution(
|
||||
server_args, "_handle_a2a_moe", moe_runner_backend="flashinfer_megamoe"
|
||||
)
|
||||
assert resolved_view(server_args).moe_runner_backend == "flashinfer_megamoe", (
|
||||
"FlashInfer MegaMOE a2a backend requires --moe-runner-backend flashinfer_megamoe"
|
||||
)
|
||||
if not is_sm100_supported():
|
||||
raise ValueError(
|
||||
"FlashInfer MegaMOE currently requires an SM100-family "
|
||||
"CUDA device for all supported quantization formats."
|
||||
)
|
||||
logger.info(
|
||||
"FlashInfer MegaMOE is enabled. The expert parallel size is "
|
||||
"adjusted to be the same as the tensor parallel size[%s].",
|
||||
cfg.tp_size,
|
||||
)
|
||||
|
||||
if a2a_backend == "deepep":
|
||||
if cfg.moe_runner_backend == "flashinfer_cutedsl":
|
||||
if cfg.deepep_mode == "auto":
|
||||
@@ -260,7 +396,7 @@ def handle_a2a_moe(server_args: Any):
|
||||
moe_a2a_backend="none",
|
||||
)
|
||||
|
||||
if cfg.moe_a2a_backend == "flashinfer":
|
||||
if a2a_now == "flashinfer":
|
||||
assert (
|
||||
resolved_view(server_args).enable_dp_attention
|
||||
and cfg.dp_size == cfg.tp_size
|
||||
@@ -273,20 +409,44 @@ def handle_a2a_moe(server_args: Any):
|
||||
resolved_view(server_args).moe_runner_backend == "flashinfer_cutedsl"
|
||||
and envs.SGLANG_FLASHINFER_CUTEDSL_NVFP4_W4A16.get()
|
||||
)
|
||||
if use_cutedsl_w4a16:
|
||||
if envs.SGLANG_MOE_NVFP4_DISPATCH.get():
|
||||
raise ValueError(
|
||||
"CuTe DSL NVFP4 W4A16 requires BF16 FlashInfer MoE "
|
||||
"dispatch; unset SGLANG_MOE_NVFP4_DISPATCH."
|
||||
)
|
||||
elif not envs.SGLANG_MOE_NVFP4_DISPATCH.is_set() and (
|
||||
resolved_view(server_args).quantization == "modelopt_fp4"
|
||||
or model_config_of(server_args).nvfp4_moe_meta is not None
|
||||
):
|
||||
envs.SGLANG_MOE_NVFP4_DISPATCH.set(True)
|
||||
logger.warning(
|
||||
"SGLANG_MOE_NVFP4_DISPATCH is set to True for Flashinfer MoE A2A"
|
||||
if use_cutedsl_w4a16 and envs.SGLANG_MOE_NVFP4_DISPATCH.get():
|
||||
raise ValueError(
|
||||
"CuTe DSL NVFP4 W4A16 requires BF16 FlashInfer MoE "
|
||||
"dispatch; unset SGLANG_MOE_NVFP4_DISPATCH."
|
||||
)
|
||||
if cfg.flashinfer_a2a_dispatch_type is None:
|
||||
if (
|
||||
not use_cutedsl_w4a16
|
||||
and not envs.SGLANG_MOE_NVFP4_DISPATCH.is_set()
|
||||
and (
|
||||
resolved_view(server_args).quantization == "modelopt_fp4"
|
||||
or model_config_of(server_args).nvfp4_moe_meta is not None
|
||||
)
|
||||
):
|
||||
envs.SGLANG_MOE_NVFP4_DISPATCH.set(True)
|
||||
logger.warning(
|
||||
"SGLANG_MOE_NVFP4_DISPATCH is set to True for Flashinfer MoE A2A"
|
||||
)
|
||||
else:
|
||||
if resolved_view(server_args).moe_runner_backend == "flashinfer_trtllm":
|
||||
declare_resolution(
|
||||
server_args,
|
||||
"_handle_a2a_moe",
|
||||
moe_runner_backend="flashinfer_trtllm_routed",
|
||||
)
|
||||
logger.warning(
|
||||
"Flashinfer MoE A2A is enabled with flashinfer_trtllm. "
|
||||
"Using flashinfer_trtllm_routed because A2A dispatch "
|
||||
"provides top-k ids and weights."
|
||||
)
|
||||
if use_cutedsl_w4a16 and cfg.flashinfer_a2a_dispatch_type in (
|
||||
"auto",
|
||||
"nvfp4",
|
||||
):
|
||||
raise ValueError(
|
||||
"CuTe DSL NVFP4 W4A16 requires --flashinfer-a2a-dispatch-type bf16."
|
||||
)
|
||||
handle_flashinfer_a2a_dispatch_type(server_args)
|
||||
assert resolved_view(server_args).moe_runner_backend in [
|
||||
"flashinfer_cutlass",
|
||||
"flashinfer_cutedsl",
|
||||
|
||||
@@ -1483,7 +1483,12 @@ def _moe_runner_backend_quant_constraints(view: Any) -> dict:
|
||||
allowed = list(MXFP8_MOE_RUNNER_BACKEND_CHOICES)
|
||||
if is_gfx95_mxfp8:
|
||||
allowed.append("triton")
|
||||
mxfp8_default = "triton" if is_gfx95_mxfp8 else "flashinfer_trtllm"
|
||||
|
||||
if view.moe_a2a_backend == "flashinfer_megamoe":
|
||||
mxfp8_default = "flashinfer_megamoe"
|
||||
else:
|
||||
mxfp8_default = "triton" if is_gfx95_mxfp8 else "flashinfer_trtllm"
|
||||
|
||||
if moe_runner_backend == "auto":
|
||||
moe_runner_backend = mxfp8_default
|
||||
elif moe_runner_backend not in allowed:
|
||||
@@ -1537,7 +1542,8 @@ def _moe_runner_fusion_disable(view: Any) -> dict:
|
||||
def _a2a_fusion_adjustments(view: Any) -> dict:
|
||||
"""A2A-backend-driven shared-experts fusion adjustments, declared at the
|
||||
legacy write slots in _handle_a2a_moe: Waterfill requires the
|
||||
fusion enabled; FlashInfer and DeepEP v2 A2A require it disabled."""
|
||||
fusion enabled; FlashInfer, FlashInfer MegaMOE, and DeepEP v2 A2A require it disabled.
|
||||
"""
|
||||
if view.moe_a2a_backend in ("deepep", "megamoe") and view.enable_waterfill:
|
||||
if view.disable_shared_experts_fusion:
|
||||
logger.warning(
|
||||
@@ -1545,7 +1551,7 @@ def _a2a_fusion_adjustments(view: Any) -> dict:
|
||||
)
|
||||
return {"disable_shared_experts_fusion": False}
|
||||
return {}
|
||||
if view.moe_a2a_backend == "flashinfer":
|
||||
if view.moe_a2a_backend in ("flashinfer", "flashinfer_megamoe"):
|
||||
logger.warning(
|
||||
"Flashinfer MoE A2A is enabled. --disable-shared-experts-fusion is automatically set."
|
||||
)
|
||||
@@ -1566,6 +1572,7 @@ _A2A_EP_SPANNING_BACKENDS = frozenset(
|
||||
"nixl",
|
||||
"ascend_fuseep",
|
||||
"flashinfer",
|
||||
"flashinfer_megamoe",
|
||||
"mori",
|
||||
"pplx",
|
||||
"deepep_v2",
|
||||
|
||||
@@ -989,6 +989,21 @@ class Envs:
|
||||
# Per-rank dispatch capacity of the FlashInfer MoE A2A dispatcher. Unset
|
||||
# means each call site keeps its own default.
|
||||
SGLANG_FLASHINFER_NUM_MAX_DISPATCH_TOKENS_PER_RANK = EnvInt(None)
|
||||
# FlashInfer MegaMOE (generic moe_ep.MoEEpMegaLayer backend). Sizes the
|
||||
# per-rank symmetric workspace; must be >= the largest padded per-rank batch
|
||||
# (derived from cuda_graph_max_bs / chunked_prefill_size when unset).
|
||||
SGLANG_FLASHINFER_MEGAMOE_MAX_TOKENS_PER_RANK = EnvInt(0)
|
||||
# Opt-in in-kernel FC2 top-k reduce (cross-rank REDG atomic-add) for the
|
||||
# cutedsl mega kernels (NVFP4 / MXFP8). Deletes the multi-GB combine staging
|
||||
# region and can win at large batch, but makes the output accumulation order
|
||||
# nondeterministic (bf16 unordered sum) -- keep off for bit-reproducibility.
|
||||
# No effect on the DeepGEMM (block-FP8) mega path, which lacks the knob.
|
||||
SGLANG_FLASHINFER_MEGAMOE_IN_KERNEL_FC2_REDUCE = EnvBool(False)
|
||||
# Cross-rank combine wire format for the FlashInfer NVFP4 cutedsl MegaMOE
|
||||
# kernel. "bf16" is exact/default; "mxfp8" and "nvfp4" reduce combine
|
||||
# traffic with a small accuracy tradeoff and require FC2 reduce outside the
|
||||
# kernel.
|
||||
SGLANG_FLASHINFER_MEGAMOE_COMBINE_DTYPE = EnvStr("bf16")
|
||||
# Enable per-token FP32 activation scaling for serialized ModelOpt FP4 with
|
||||
# FlashInfer TRT-LLM or CuTe DSL v2 MoE.
|
||||
SGLANG_FLASHINFER_NVFP4_PER_TOKEN_ACTIVATION = EnvBool(False)
|
||||
|
||||
@@ -0,0 +1,672 @@
|
||||
# Copyright 2023-2024 SGLang Team
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
# ==============================================================================
|
||||
"""Generic FlashInfer MegaMOE backend (moe_ep.MoEEpMegaLayer).
|
||||
|
||||
Wraps FlashInfer's fused EP all-to-all + expert-compute mega kernel so it can
|
||||
be selected as a model-agnostic MoE runner backend through the standard
|
||||
FusedMoE dispatch -> run_moe_core -> combine flow. The mega kernel does its EP
|
||||
communication internally via the deep_gemm symmetric buffer, so the dispatcher
|
||||
and combine stay pure no-ops; this module owns the layer build + forward.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from collections.abc import Callable, Generator
|
||||
from contextlib import contextmanager
|
||||
from dataclasses import dataclass
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.layers.moe.moe_runner.base import (
|
||||
MoeQuantInfo,
|
||||
MoeRunnerConfig,
|
||||
register_fused_func,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _format_megakernel_config(config: Any) -> str:
|
||||
"""Readable one-line repr of a mega kernel config.
|
||||
|
||||
The config dataclasses carry per-expert tensor fields (e.g. fc1_alpha /
|
||||
fc2_alpha / fc1_norm_const); their default repr dumps every element, so
|
||||
abbreviate tensors to shape/dtype/device instead.
|
||||
"""
|
||||
import dataclasses
|
||||
|
||||
if not dataclasses.is_dataclass(config):
|
||||
return repr(config)
|
||||
|
||||
parts = []
|
||||
for field, value in zip(dataclasses.fields(config), dataclasses.astuple(config)):
|
||||
if isinstance(value, torch.Tensor):
|
||||
value = (
|
||||
f"Tensor(shape={tuple(value.shape)}, dtype={value.dtype}, "
|
||||
f"device={value.device})"
|
||||
)
|
||||
else:
|
||||
value = repr(value)
|
||||
parts.append(f"{field.name}={value}")
|
||||
return f"{type(config).__name__}({', '.join(parts)})"
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.layers.moe.fused_moe_triton.layer import FusedMoE
|
||||
from sglang.srt.layers.moe.token_dispatcher import (
|
||||
DispatchOutput,
|
||||
StandardCombineInput,
|
||||
)
|
||||
|
||||
|
||||
@contextmanager
|
||||
def _capture_safe_ue8m0_pack() -> Generator[None, None, None]:
|
||||
"""Make deep_gemm's UE8M0 scale packing safe under CUDA graph capture.
|
||||
|
||||
The deep_gemm mega staging path (block-FP8 models such as DeepSeek-V4-Flash)
|
||||
runs ``per_token_cast_to_fp8(..., use_packed_ue8m0=True)`` on every forward,
|
||||
which calls ``deep_gemm.utils.math.pack_ue8m0_to_int``. Its upstream
|
||||
implementation carries two debug assertions::
|
||||
|
||||
assert (x_int >= 0).all() and (x_int & 0x7fffff == 0).all()
|
||||
|
||||
``.all()`` forces a device->host sync, which is illegal while a CUDA graph is
|
||||
capturing. Replace the helper only around the mega forward that runs during
|
||||
capture, then restore the exact upstream function.
|
||||
|
||||
TODO(deepseek-ai/DeepGEMM#414): remove once upstream provides a capture-safe
|
||||
helper: https://github.com/deepseek-ai/DeepGEMM/issues/414
|
||||
"""
|
||||
if not torch.cuda.is_available() or not torch.cuda.is_current_stream_capturing():
|
||||
yield
|
||||
return
|
||||
|
||||
try:
|
||||
import deep_gemm.utils.math as _dgm
|
||||
except ImportError:
|
||||
# deep_gemm is only needed by the block-FP8 mega path; NVFP4/MXFP8 mega
|
||||
# runs on cutedsl and does not import it. Nothing to patch here.
|
||||
yield
|
||||
return
|
||||
|
||||
def _pack_ue8m0_to_int(x: torch.Tensor) -> torch.Tensor:
|
||||
x_int = x.view(torch.int)
|
||||
return (x_int >> 23).to(torch.uint8).view(torch.int)
|
||||
|
||||
original = _dgm.pack_ue8m0_to_int
|
||||
_dgm.pack_ue8m0_to_int = _pack_ue8m0_to_int
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
_dgm.pack_ue8m0_to_int = original
|
||||
|
||||
|
||||
@dataclass
|
||||
class FlashInferMegaMoeQuantInfo(MoeQuantInfo):
|
||||
mega: Any
|
||||
mega_forward: Callable[[Any, Any], torch.Tensor] | None = None
|
||||
fc1_alpha: torch.Tensor | None = None
|
||||
fc2_alpha: torch.Tensor | None = None
|
||||
fc1_norm_const: torch.Tensor | None = None
|
||||
apply_routed_scaling_factor: bool = False
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
if self.mega_forward is None:
|
||||
self.mega_forward = _select_megamoe_forward(self.mega)
|
||||
|
||||
|
||||
def _forward_megamoe_with_workspace_view(mega: Any, tensors: Any) -> torch.Tensor:
|
||||
return mega.forward(tensors, return_workspace_view=True)
|
||||
|
||||
|
||||
def _forward_megamoe_legacy(mega: Any, tensors: Any) -> torch.Tensor:
|
||||
return mega.forward(tensors)
|
||||
|
||||
|
||||
def _select_megamoe_forward(mega: Any) -> Callable[[Any, Any], torch.Tensor]:
|
||||
import inspect
|
||||
|
||||
if "return_workspace_view" in inspect.signature(mega.forward).parameters:
|
||||
return _forward_megamoe_with_workspace_view
|
||||
return _forward_megamoe_legacy
|
||||
|
||||
|
||||
def _resolve_max_tokens_per_rank() -> int:
|
||||
"""Per-rank symmetric-buffer sizing for the mega kernel.
|
||||
|
||||
Honors the explicit env override; otherwise derives the largest per-(DP)rank
|
||||
token count a single MoE forward can route (same bound the cutedsl A2A path
|
||||
uses), falling back to 1024 if it cannot be determined.
|
||||
"""
|
||||
configured = envs.SGLANG_FLASHINFER_MEGAMOE_MAX_TOKENS_PER_RANK.get()
|
||||
if configured > 0:
|
||||
return configured
|
||||
|
||||
from sglang.srt.runtime_context import cutedsl_moe_max_num_tokens
|
||||
|
||||
derived = cutedsl_moe_max_num_tokens()
|
||||
return derived if derived > 0 else 1024
|
||||
|
||||
|
||||
def resolve_flashinfer_megamoe_combine_dtype() -> str:
|
||||
combine_dtype = envs.SGLANG_FLASHINFER_MEGAMOE_COMBINE_DTYPE.get().strip().lower()
|
||||
if combine_dtype not in ("bf16", "mxfp8", "nvfp4"):
|
||||
raise ValueError(
|
||||
"SGLANG_FLASHINFER_MEGAMOE_COMBINE_DTYPE must be one of "
|
||||
f"'bf16', 'mxfp8', or 'nvfp4', got {combine_dtype!r}."
|
||||
)
|
||||
if (
|
||||
combine_dtype != "bf16"
|
||||
and envs.SGLANG_FLASHINFER_MEGAMOE_IN_KERNEL_FC2_REDUCE.get()
|
||||
):
|
||||
raise ValueError(
|
||||
"SGLANG_FLASHINFER_MEGAMOE_COMBINE_DTYPE="
|
||||
f"{combine_dtype!r} is incompatible with "
|
||||
"SGLANG_FLASHINFER_MEGAMOE_IN_KERNEL_FC2_REDUCE=1."
|
||||
)
|
||||
return combine_dtype
|
||||
|
||||
|
||||
def _layer_ep_world_rank(layer: FusedMoE) -> tuple[int, int]:
|
||||
world_size = int(layer.moe_ep_size)
|
||||
rank = int(layer.moe_ep_rank)
|
||||
if world_size <= 0:
|
||||
raise ValueError(f"moe_ep_size must be positive, got {world_size}.")
|
||||
if rank < 0 or rank >= world_size:
|
||||
raise ValueError(f"moe_ep_rank must be in [0, {world_size}), got {rank}.")
|
||||
return world_size, rank
|
||||
|
||||
|
||||
def _scalar_float(value: Any) -> float:
|
||||
if isinstance(value, torch.Tensor):
|
||||
return float(value.detach().to(torch.float32).max())
|
||||
return float(value)
|
||||
|
||||
|
||||
def _local_expert_vector(value: torch.Tensor, num_local_experts: int) -> torch.Tensor:
|
||||
value = value.detach().to(torch.float32)
|
||||
if value.dim() == 0:
|
||||
return value.expand(num_local_experts).contiguous()
|
||||
if value.shape != (num_local_experts,):
|
||||
raise ValueError(
|
||||
f"expected per-local-expert vector of shape ({num_local_experts},), "
|
||||
f"got {tuple(value.shape)}"
|
||||
)
|
||||
return value.contiguous()
|
||||
|
||||
|
||||
def _validate_nvfp4_fc1_alpha(layer: FusedMoE) -> None:
|
||||
"""MegaMOE reuses ``g1_alphas`` as fc1_alpha; the kernel takes one alpha per
|
||||
expert, so the gate and up FC1 alphas must agree."""
|
||||
if not layer.moe_runner_config.is_gated:
|
||||
return
|
||||
gate_alpha = _local_expert_vector(layer.g1_alphas, layer.num_local_experts)
|
||||
up_alpha = _local_expert_vector(layer.g1_alphas_up, layer.num_local_experts)
|
||||
if not torch.allclose(gate_alpha, up_alpha):
|
||||
raise ValueError(
|
||||
"FlashInfer NVFP4 MegaMOE requires matching gate/up FC1 alpha "
|
||||
"values because the kernel accepts one alpha per expert."
|
||||
)
|
||||
|
||||
|
||||
def _bind_transformed_weights(
|
||||
layer: FusedMoE,
|
||||
transformed_weights: Any,
|
||||
*,
|
||||
w13_scale_name: str,
|
||||
w2_scale_name: str,
|
||||
) -> None:
|
||||
from sglang.srt.layers.utils.common import copy_or_rebind_param
|
||||
|
||||
(w13_weight, w13_scale), (w2_weight, w2_scale) = transformed_weights
|
||||
copy_or_rebind_param(layer, "w13_weight", w13_weight)
|
||||
copy_or_rebind_param(layer, w13_scale_name, w13_scale)
|
||||
copy_or_rebind_param(layer, "w2_weight", w2_weight)
|
||||
copy_or_rebind_param(layer, w2_scale_name, w2_scale)
|
||||
|
||||
|
||||
def _init_flashinfer_megamoe_layer_state(layer: FusedMoE) -> None:
|
||||
layer._flashinfer_megamoe_layer = None
|
||||
layer._flashinfer_megamoe_forward = None
|
||||
layer._flashinfer_megamoe_input_norm_const = None
|
||||
|
||||
|
||||
def _get_or_init_flashinfer_megamoe_layer_state(layer: FusedMoE) -> Any:
|
||||
if not hasattr(layer, "_flashinfer_megamoe_layer"):
|
||||
_init_flashinfer_megamoe_layer_state(layer)
|
||||
return layer._flashinfer_megamoe_layer
|
||||
|
||||
|
||||
def _ensure_flashinfer_megamoe_layer(
|
||||
layer: FusedMoE,
|
||||
*,
|
||||
megakernel_config: Any,
|
||||
w13_scale: torch.Tensor,
|
||||
w2_scale: torch.Tensor,
|
||||
) -> Any:
|
||||
mega = _get_or_init_flashinfer_megamoe_layer_state(layer)
|
||||
if mega is not None:
|
||||
return mega
|
||||
|
||||
from flashinfer.moe_ep import (
|
||||
BootstrapConfig,
|
||||
FleetParams,
|
||||
MegaConfig,
|
||||
MoEEpMegaLayer,
|
||||
)
|
||||
|
||||
transformed_weights = (
|
||||
(layer.w13_weight.data, w13_scale.data),
|
||||
(layer.w2_weight.data, w2_scale.data),
|
||||
)
|
||||
world_size, rank = _layer_ep_world_rank(layer)
|
||||
|
||||
max_tokens_per_rank = _resolve_max_tokens_per_rank()
|
||||
logger.debug(
|
||||
"FlashInfer MegaMOE layer[%s] build: megakernel_config=%s "
|
||||
"(world_size=%d, num_experts=%d, max_tokens_per_rank=%d, hidden_size=%d)",
|
||||
layer.layer_id,
|
||||
_format_megakernel_config(megakernel_config),
|
||||
world_size,
|
||||
layer.num_experts,
|
||||
max_tokens_per_rank,
|
||||
layer.hidden_size,
|
||||
)
|
||||
|
||||
mega = MoEEpMegaLayer(
|
||||
bootstrap=BootstrapConfig(
|
||||
world_size=world_size, rank=rank, device=torch.cuda.current_device()
|
||||
),
|
||||
fleet_params=FleetParams(
|
||||
num_experts=layer.num_experts,
|
||||
max_tokens_per_rank=max_tokens_per_rank,
|
||||
token_hidden_size=layer.hidden_size,
|
||||
),
|
||||
# weights already preprocessed in prepare_*; with transformed_weights set
|
||||
# the kernel never reads `weights` (see MoEEpMegaLayer), so pass None.
|
||||
weights=None,
|
||||
backend=MegaConfig(
|
||||
megakernel=megakernel_config,
|
||||
preprocess_weights=False,
|
||||
transformed_weights=transformed_weights,
|
||||
),
|
||||
)
|
||||
layer._flashinfer_megamoe_layer = mega
|
||||
layer._flashinfer_megamoe_forward = _select_megamoe_forward(mega)
|
||||
return mega
|
||||
|
||||
|
||||
def ensure_fp4_moe_layer_for_flashinfer_megamoe(layer: FusedMoE) -> Any:
|
||||
mega = _get_or_init_flashinfer_megamoe_layer_state(layer)
|
||||
if mega is not None:
|
||||
return mega
|
||||
|
||||
from flashinfer.moe_ep import DeepGemmMegaMoeConfig
|
||||
|
||||
return _ensure_flashinfer_megamoe_layer(
|
||||
layer,
|
||||
megakernel_config=DeepGemmMegaMoeConfig(
|
||||
intermediate_size=layer.intermediate_size_per_partition,
|
||||
top_k=layer.top_k,
|
||||
activation_clamp=layer.moe_runner_config.swiglu_limit,
|
||||
),
|
||||
w13_scale=layer.w13_weight_scale_inv,
|
||||
w2_scale=layer.w2_weight_scale_inv,
|
||||
)
|
||||
|
||||
|
||||
def ensure_nvfp4_moe_layer_for_flashinfer_megamoe(layer: FusedMoE) -> Any:
|
||||
mega = _get_or_init_flashinfer_megamoe_layer_state(layer)
|
||||
if mega is not None:
|
||||
return mega
|
||||
|
||||
from flashinfer.moe_ep import Nvfp4CutedslMegaMoeConfig
|
||||
|
||||
input_norm_const = layer._flashinfer_megamoe_input_norm_const
|
||||
if input_norm_const is None:
|
||||
if torch.cuda.is_available() and torch.cuda.is_current_stream_capturing():
|
||||
raise RuntimeError(
|
||||
"FlashInfer NVFP4 MegaMOE layer must be initialized before "
|
||||
"CUDA graph capture."
|
||||
)
|
||||
logger.warning(
|
||||
"FlashInfer NVFP4 MegaMOE layer[%s]: input_norm_const was not "
|
||||
"precomputed at weight-load time; computing it lazily now via a "
|
||||
"blocking device sync. This should only happen once per layer, "
|
||||
"but if it happens during warmup it will look like a stall.",
|
||||
layer.layer_id,
|
||||
)
|
||||
input_norm_const = _scalar_float(layer.w13_input_scale_quant)
|
||||
layer._flashinfer_megamoe_input_norm_const = input_norm_const
|
||||
|
||||
return _ensure_flashinfer_megamoe_layer(
|
||||
layer,
|
||||
megakernel_config=Nvfp4CutedslMegaMoeConfig(
|
||||
intermediate_size=layer.intermediate_size_per_partition,
|
||||
top_k=layer.top_k,
|
||||
gate_up_clamp=layer.moe_runner_config.swiglu_limit,
|
||||
apply_topk_in_fc1=True,
|
||||
in_kernel_fc2_reduce=envs.SGLANG_FLASHINFER_MEGAMOE_IN_KERNEL_FC2_REDUCE.get(),
|
||||
combine_dtype=resolve_flashinfer_megamoe_combine_dtype(),
|
||||
input_norm_const=input_norm_const,
|
||||
fc1_alpha=layer.g1_alphas,
|
||||
fc2_alpha=layer.g2_alphas,
|
||||
fc1_norm_const=layer.w2_input_scale_quant,
|
||||
),
|
||||
w13_scale=layer.w13_weight_scale,
|
||||
w2_scale=layer.w2_weight_scale,
|
||||
)
|
||||
|
||||
|
||||
def ensure_mxfp8_moe_layer_for_flashinfer_megamoe(layer: FusedMoE) -> Any:
|
||||
mega = _get_or_init_flashinfer_megamoe_layer_state(layer)
|
||||
if mega is not None:
|
||||
return mega
|
||||
|
||||
from flashinfer.moe_ep import Mxfp8CutedslMegaMoeConfig
|
||||
|
||||
return _ensure_flashinfer_megamoe_layer(
|
||||
layer,
|
||||
megakernel_config=Mxfp8CutedslMegaMoeConfig(
|
||||
intermediate_size=layer.intermediate_size_per_partition,
|
||||
top_k=layer.top_k,
|
||||
kind="mxfp8_e4m3",
|
||||
gate_up_clamp=layer.moe_runner_config.swiglu_limit,
|
||||
in_kernel_fc2_reduce=envs.SGLANG_FLASHINFER_MEGAMOE_IN_KERNEL_FC2_REDUCE.get(),
|
||||
),
|
||||
w13_scale=layer.w13_weight_scale_inv,
|
||||
w2_scale=layer.w2_weight_scale_inv,
|
||||
)
|
||||
|
||||
|
||||
def prepare_fp4_moe_weights_for_flashinfer_megamoe(
|
||||
layer: FusedMoE,
|
||||
) -> None:
|
||||
"""Prepare loaded FP4 weights for MegaMOE.
|
||||
|
||||
SGLang loads FP4-packed expert weights plus raw block scales. FlashInfer's
|
||||
current moe_ep API owns backend-specific weight preprocessing, including
|
||||
DeepGEMM scale layout transforms.
|
||||
"""
|
||||
_init_flashinfer_megamoe_layer_state(layer)
|
||||
|
||||
from flashinfer.moe_ep import (
|
||||
MoEWeightPack,
|
||||
preprocess_mega_weights,
|
||||
)
|
||||
|
||||
weights = MoEWeightPack(
|
||||
w13=layer.w13_weight.data,
|
||||
w2=layer.w2_weight.data,
|
||||
w13_scale=layer.w13_weight_scale_inv.data,
|
||||
w2_scale=layer.w2_weight_scale_inv.data,
|
||||
)
|
||||
transformed_weights = preprocess_mega_weights(
|
||||
weights,
|
||||
intermediate_size=layer.intermediate_size_per_partition,
|
||||
hidden_size=layer.hidden_size,
|
||||
)
|
||||
_bind_transformed_weights(
|
||||
layer,
|
||||
transformed_weights,
|
||||
w13_scale_name="w13_weight_scale_inv",
|
||||
w2_scale_name="w2_weight_scale_inv",
|
||||
)
|
||||
|
||||
|
||||
def prepare_nvfp4_moe_weights_for_flashinfer_megamoe(
|
||||
layer: FusedMoE,
|
||||
) -> None:
|
||||
_init_flashinfer_megamoe_layer_state(layer)
|
||||
|
||||
from flashinfer.moe_ep import (
|
||||
MoEWeightPack,
|
||||
preprocess_nvfp4_cutedsl_mega_weights,
|
||||
)
|
||||
|
||||
if layer.hidden_size % 128 != 0:
|
||||
raise ValueError(
|
||||
"FlashInfer NVFP4 MegaMOE requires hidden_size to be a multiple "
|
||||
f"of 128, got {layer.hidden_size}."
|
||||
)
|
||||
if layer.quant_config.use_per_token_activation:
|
||||
raise ValueError(
|
||||
"FlashInfer NVFP4 MegaMOE does not support per-token activation "
|
||||
"scaling. Use flashinfer_trtllm/flashinfer_trtllm_routed for "
|
||||
"ModelOpt NVFP4 per-token activation."
|
||||
)
|
||||
if layer.intermediate_size_per_partition % 128 != 0:
|
||||
raise ValueError(
|
||||
"FlashInfer NVFP4 MegaMOE requires intermediate_size_per_partition "
|
||||
f"to be a multiple of 128, got {layer.intermediate_size_per_partition}."
|
||||
)
|
||||
if layer.num_experts % layer.moe_ep_size != 0:
|
||||
raise ValueError(
|
||||
"FlashInfer NVFP4 MegaMOE requires num_experts to be divisible by "
|
||||
f"ep_size, got {layer.num_experts=} and {layer.moe_ep_size=}."
|
||||
)
|
||||
|
||||
_validate_nvfp4_fc1_alpha(layer)
|
||||
layer._flashinfer_megamoe_input_norm_const = _scalar_float(
|
||||
layer.w13_input_scale_quant
|
||||
)
|
||||
|
||||
gate_up_clamp = layer.moe_runner_config.swiglu_limit
|
||||
|
||||
weights = MoEWeightPack(
|
||||
w13=layer.w13_weight.data,
|
||||
w2=layer.w2_weight.data,
|
||||
w13_scale=layer.w13_weight_scale.data,
|
||||
w2_scale=layer.w2_weight_scale.data,
|
||||
)
|
||||
transformed_weights = preprocess_nvfp4_cutedsl_mega_weights(
|
||||
weights,
|
||||
intermediate_size=layer.intermediate_size_per_partition,
|
||||
hidden_size=layer.hidden_size,
|
||||
gate_up_clamp=gate_up_clamp,
|
||||
activation_clamp=None,
|
||||
)
|
||||
_bind_transformed_weights(
|
||||
layer,
|
||||
transformed_weights,
|
||||
w13_scale_name="w13_weight_scale",
|
||||
w2_scale_name="w2_weight_scale",
|
||||
)
|
||||
|
||||
|
||||
def prepare_mxfp8_moe_weights_for_flashinfer_megamoe(
|
||||
layer: FusedMoE,
|
||||
) -> None:
|
||||
_init_flashinfer_megamoe_layer_state(layer)
|
||||
|
||||
from flashinfer.moe_ep import (
|
||||
MoEWeightPack,
|
||||
preprocess_mxfp8_cutedsl_mega_weights,
|
||||
)
|
||||
|
||||
if layer.hidden_size % 128 != 0:
|
||||
raise ValueError(
|
||||
"FlashInfer MXFP8 MegaMOE requires hidden_size to be a multiple "
|
||||
f"of 128, got {layer.hidden_size}."
|
||||
)
|
||||
if layer.intermediate_size_per_partition % 128 != 0:
|
||||
raise ValueError(
|
||||
"FlashInfer MXFP8 MegaMOE requires intermediate_size_per_partition "
|
||||
f"to be a multiple of 128, got {layer.intermediate_size_per_partition}."
|
||||
)
|
||||
if layer.num_experts % layer.moe_ep_size != 0:
|
||||
raise ValueError(
|
||||
"FlashInfer MXFP8 MegaMOE requires num_experts to be divisible by "
|
||||
f"ep_size, got {layer.num_experts=} and {layer.moe_ep_size=}."
|
||||
)
|
||||
|
||||
weights = MoEWeightPack(
|
||||
w13=layer.w13_weight.data,
|
||||
w2=layer.w2_weight.data,
|
||||
w13_scale=layer.w13_weight_scale_inv.data,
|
||||
w2_scale=layer.w2_weight_scale_inv.data,
|
||||
)
|
||||
transformed_weights = preprocess_mxfp8_cutedsl_mega_weights(
|
||||
weights,
|
||||
intermediate_size=layer.intermediate_size_per_partition,
|
||||
hidden_size=layer.hidden_size,
|
||||
kind="mxfp8_e4m3",
|
||||
gate_up_clamp=layer.moe_runner_config.swiglu_limit,
|
||||
activation_clamp=None,
|
||||
)
|
||||
_bind_transformed_weights(
|
||||
layer,
|
||||
transformed_weights,
|
||||
w13_scale_name="w13_weight_scale_inv",
|
||||
w2_scale_name="w2_weight_scale_inv",
|
||||
)
|
||||
|
||||
|
||||
def _ensure_shared_workspace(mega: Any) -> None:
|
||||
"""Share this layer's workspace across MegaMOE layers with identical
|
||||
fleet/kernel geometry.
|
||||
FlashInfer's own workspace pool keys by fc1_alpha/fc2_alpha/fc1_norm_const
|
||||
tensor identity, but sglang binds distinct tensor objects per layer, so
|
||||
its pool never hits across layers; key on geometry instead, since those
|
||||
values are re-staged into the workspace per forward rather than baked
|
||||
into the compiled kernel. Creation is collective, so it must only happen
|
||||
on a layer's first forward, under warmup's cross-rank lockstep.
|
||||
"""
|
||||
if mega._workspace is not None:
|
||||
return
|
||||
fp = mega._fleet_params
|
||||
kc = mega._megakernel_config
|
||||
mc = mega._mega_config
|
||||
from sglang.srt.runtime_context import get_resources
|
||||
|
||||
key = (
|
||||
getattr(kc, "kernel_name", kc.__class__.__name__),
|
||||
mega._bootstrap.world_size,
|
||||
fp.num_experts,
|
||||
fp.max_tokens_per_rank,
|
||||
fp.token_hidden_size,
|
||||
kc.top_k,
|
||||
kc.intermediate_size,
|
||||
getattr(kc, "gate_up_clamp", None),
|
||||
getattr(kc, "activation_clamp", None),
|
||||
getattr(kc, "apply_topk_in_fc1", None),
|
||||
getattr(kc, "kind", None),
|
||||
getattr(kc, "in_kernel_fc2_reduce", None),
|
||||
getattr(kc, "combine_dtype", None),
|
||||
getattr(kc, "token_back_by_dispatch", None),
|
||||
getattr(kc, "fast_math", None),
|
||||
mc.quantize_input,
|
||||
)
|
||||
workspaces = get_resources().flashinfer_megamoe_workspaces
|
||||
shared = workspaces.get(key)
|
||||
if shared is None:
|
||||
workspaces[key] = mega._ensure_workspace()
|
||||
else:
|
||||
mega._workspace = shared
|
||||
|
||||
|
||||
@register_fused_func("flashinfer_megamoe", "flashinfer_megamoe")
|
||||
def run_flashinfer_megamoe(
|
||||
dispatch_output: DispatchOutput,
|
||||
quant_info: MoeQuantInfo,
|
||||
runner_config: MoeRunnerConfig,
|
||||
) -> StandardCombineInput:
|
||||
"""Run the fused mega kernel and return per-rank outputs (no combine)."""
|
||||
from flashinfer.moe_ep import MoEEpTensors
|
||||
|
||||
from sglang.srt.layers.moe.token_dispatcher import StandardCombineInput
|
||||
|
||||
assert isinstance(quant_info, FlashInferMegaMoeQuantInfo), (
|
||||
f"Unexpected quant_info type for flashinfer_megamoe: {type(quant_info)}"
|
||||
)
|
||||
|
||||
x = dispatch_output.hidden_states
|
||||
topk_output = dispatch_output.topk_output
|
||||
topk_weights = topk_output.topk_weights
|
||||
topk_ids = topk_output.topk_ids
|
||||
mega = quant_info.mega
|
||||
_ensure_shared_workspace(mega)
|
||||
|
||||
t = MoEEpTensors(
|
||||
hidden_states=x.to(torch.bfloat16),
|
||||
# FlashInfer's fused staging accepts the int32 router output and widens
|
||||
# directly into its final int64 workspace buffer. Keep this path copy-free.
|
||||
topk_ids=topk_ids,
|
||||
topk_weights=topk_weights.to(torch.float32),
|
||||
fc1_alpha=quant_info.fc1_alpha,
|
||||
fc2_alpha=quant_info.fc2_alpha,
|
||||
fc1_norm_const=quant_info.fc1_norm_const,
|
||||
)
|
||||
with _capture_safe_ue8m0_pack():
|
||||
assert quant_info.mega_forward is not None
|
||||
y = quant_info.mega_forward(mega, t)
|
||||
|
||||
if quant_info.apply_routed_scaling_factor:
|
||||
rsf = runner_config.routed_scaling_factor
|
||||
if rsf is not None and rsf != 1.0:
|
||||
y.mul_(rsf)
|
||||
|
||||
return StandardCombineInput(hidden_states=y)
|
||||
|
||||
|
||||
def warmup_all_flashinfer_megamoe_layers(model: torch.nn.Module) -> None:
|
||||
"""Force every FlashInfer MegaMOE layer to build before CUDA graph capture.
|
||||
``ensure_*_moe_layer_for_flashinfer_megamoe`` builds a layer's state
|
||||
lazily on its first forward; if that first forward instead happens inside
|
||||
a CUDA graph capture, the lazy build's blocking device sync isn't allowed
|
||||
and it raises rather than silently recovering. Call this once, explicitly,
|
||||
right before capture begins (see ``ModelRunner.init_cuda_graphs``) so
|
||||
every layer is built eagerly outside any graph.
|
||||
Only the nvfp4 path is wired up below; extend the dispatch if fp4/mxfp8
|
||||
MegaMOE hits the same gap.
|
||||
"""
|
||||
from sglang.srt.layers.moe.fused_moe_triton.layer import FusedMoE
|
||||
|
||||
n_built = 0
|
||||
for module in model.modules():
|
||||
if not isinstance(module, FusedMoE):
|
||||
continue
|
||||
if not hasattr(module, "_flashinfer_megamoe_layer"):
|
||||
# This layer's quant method never went through one of the
|
||||
# prepare_*_moe_weights_for_flashinfer_megamoe hooks -- not a
|
||||
# MegaMOE layer (or not on the flashinfer_megamoe backend).
|
||||
continue
|
||||
if getattr(module, "_flashinfer_megamoe_layer", None) is not None:
|
||||
continue # already built (e.g. warmup's dummy batch hit it)
|
||||
|
||||
# Dispatch mirrors modelopt_quant.py's apply(): only the nvfp4
|
||||
# method is wired up here today.
|
||||
if type(module.quant_method).__name__ == "ModelOptNvFp4FusedMoEMethod":
|
||||
ensure_nvfp4_moe_layer_for_flashinfer_megamoe(module)
|
||||
n_built += 1
|
||||
else:
|
||||
logger.warning(
|
||||
"warmup_all_flashinfer_megamoe_layers: layer[%s] uses "
|
||||
"quant_method=%s, which this eager pre-capture warmup does "
|
||||
"not know how to build. If capture then fails with "
|
||||
"'must be initialized before CUDA graph capture', add a "
|
||||
"branch here for that quant method.",
|
||||
getattr(module, "layer_id", "?"),
|
||||
type(module.quant_method).__name__,
|
||||
)
|
||||
|
||||
if n_built:
|
||||
logger.info(
|
||||
"warmup_all_flashinfer_megamoe_layers: eagerly built %d "
|
||||
"FlashInfer MegaMOE layer(s) before CUDA graph capture.",
|
||||
n_built,
|
||||
)
|
||||
@@ -164,11 +164,14 @@ def create_moe_dispatcher(moe_runner_config: MoeRunnerConfig) -> BaseDispatcher:
|
||||
elif (
|
||||
a2a_backend.is_none()
|
||||
or a2a_backend.is_megamoe()
|
||||
or a2a_backend.is_flashinfer_megamoe()
|
||||
or a2a_backend.is_ascend_fuseep()
|
||||
):
|
||||
# ascend_fuseep bypasses the dispatcher abstraction (see
|
||||
# forward_fuseep in hardware_backend/npu/moe/fuseep.py); a
|
||||
# StandardDispatcher is created but never invoked.
|
||||
# flashinfer_megamoe does its EP all-to-all inside the kernel, so the
|
||||
# dispatcher stays a pure noop passthrough.
|
||||
return StandardDispatcher(moe_runner_config)
|
||||
elif (
|
||||
a2a_backend.is_deepep()
|
||||
|
||||
@@ -133,6 +133,8 @@ def round_up_to_multiple(x: int, m: int) -> int:
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.layers.moe.token_dispatcher import (
|
||||
FlashinferCombineInput,
|
||||
FlashinferDispatchOutput,
|
||||
StandardCombineInput,
|
||||
StandardDispatchOutput,
|
||||
)
|
||||
@@ -721,13 +723,16 @@ class FlashInferTrtllmFp8MoeQuantInfo(MoeQuantInfo):
|
||||
|
||||
|
||||
def fused_experts_none_to_flashinfer_trtllm_fp8(
|
||||
dispatch_output: StandardDispatchOutput,
|
||||
dispatch_output: StandardDispatchOutput | FlashinferDispatchOutput,
|
||||
quant_info: FlashInferTrtllmFp8MoeQuantInfo,
|
||||
runner_config: MoeRunnerConfig,
|
||||
use_routed_topk: bool = False,
|
||||
) -> StandardCombineInput:
|
||||
from flashinfer.fused_moe import Fp8QuantizationType
|
||||
|
||||
from sglang.srt.layers.moe.token_dispatcher.flashinfer import (
|
||||
FlashinferDispatchOutput,
|
||||
)
|
||||
from sglang.srt.layers.moe.token_dispatcher.standard import StandardCombineInput
|
||||
from sglang.srt.layers.moe.topk import TopKOutputChecker
|
||||
from sglang.srt.layers.moe.utils import RoutingMethodType
|
||||
@@ -740,6 +745,11 @@ def fused_experts_none_to_flashinfer_trtllm_fp8(
|
||||
assert not runner_config.no_combine, "no_combine is not supported for flashinfer."
|
||||
|
||||
hidden_states = dispatch_output.hidden_states
|
||||
output_dtype = (
|
||||
dispatch_output.output_dtype
|
||||
if isinstance(dispatch_output, FlashinferDispatchOutput)
|
||||
else hidden_states.dtype
|
||||
)
|
||||
topk_output = dispatch_output.topk_output
|
||||
if TopKOutputChecker.format_is_bypassed(topk_output):
|
||||
router_logits = topk_output.router_logits
|
||||
@@ -775,14 +785,26 @@ def fused_experts_none_to_flashinfer_trtllm_fp8(
|
||||
|
||||
if quant_info.use_mxfp8:
|
||||
assert quant_info.weight_block_k == 32
|
||||
from sglang.srt.layers.quantization.fp8_utils import (
|
||||
flashinfer_mxfp8_quantize,
|
||||
)
|
||||
if dispatch_output.hidden_states_scale is not None:
|
||||
a_q = hidden_states
|
||||
a_sf_t = dispatch_output.hidden_states_scale
|
||||
a_sf_t = a_sf_t.reshape(hidden_states.shape[0], -1).contiguous()
|
||||
else:
|
||||
from sglang.srt.layers.quantization.fp8_utils import (
|
||||
flashinfer_mxfp8_quantize,
|
||||
)
|
||||
|
||||
a_q, a_sf = flashinfer_mxfp8_quantize(hidden_states, False)
|
||||
# FlashInfer TRT-LLM MxFP8 expects token-major activation scales:
|
||||
# [num_tokens, hidden_size // 32] (no transpose).
|
||||
a_sf_t = a_sf.view(torch.uint8).reshape(hidden_states.shape[0], -1)
|
||||
a_q, a_sf = flashinfer_mxfp8_quantize(hidden_states, False)
|
||||
# FlashInfer TRT-LLM MxFP8 expects token-major activation scales:
|
||||
# [num_tokens, hidden_size // 32] (no transpose).
|
||||
a_sf_t = (
|
||||
a_sf.view(torch.uint8)
|
||||
.reshape(hidden_states.shape[0], -1)
|
||||
.contiguous()
|
||||
)
|
||||
assert a_q.dtype == torch.float8_e4m3fn
|
||||
assert a_sf_t.dtype == torch.uint8
|
||||
assert a_sf_t.shape[1] == hidden_states.shape[1] // 32
|
||||
else:
|
||||
a_q, a_sf = per_token_group_quant_fp8(
|
||||
hidden_states, quant_info.weight_block_k, column_major_scales=True
|
||||
@@ -799,7 +821,7 @@ def fused_experts_none_to_flashinfer_trtllm_fp8(
|
||||
symm_output = torch.empty(
|
||||
hidden_states.shape[0],
|
||||
hidden_states.shape[1],
|
||||
dtype=hidden_states.dtype,
|
||||
dtype=output_dtype,
|
||||
device=hidden_states.device,
|
||||
)
|
||||
|
||||
|
||||
@@ -130,6 +130,13 @@ class MoeRunner:
|
||||
from sglang.srt.layers.moe.moe_runner import ( # noqa: F401
|
||||
flashinfer_cutlass,
|
||||
)
|
||||
elif runner_backend.is_flashinfer_megamoe():
|
||||
if lora_enabled:
|
||||
raise NotImplementedError(
|
||||
"FlashInfer MegaMOE does not support LoRA because it requires a fused path."
|
||||
)
|
||||
self.runner_core = None # FlashInfer MegaMOE only supports fused path
|
||||
import sglang.srt.layers.moe.flashinfer_megamoe # noqa: F401
|
||||
elif runner_backend.is_cutlass():
|
||||
self.runner_core = None # CUTLASS uses the direct cutlass_moe_fp4 path
|
||||
elif runner_backend.is_hpc_ops():
|
||||
|
||||
@@ -32,7 +32,11 @@ from sglang.srt.layers.moe.topk import (
|
||||
TopKOutput,
|
||||
TopKOutputChecker,
|
||||
)
|
||||
from sglang.srt.layers.moe.utils import get_moe_runner_backend
|
||||
from sglang.srt.layers.moe.utils import (
|
||||
FlashinferA2ADispatchType,
|
||||
get_flashinfer_a2a_dispatch_type,
|
||||
get_moe_runner_backend,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_flags, get_parallel, get_schedule, get_spec
|
||||
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
||||
|
||||
@@ -50,8 +54,6 @@ except ImportError:
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
MOE_NVFP4_DISPATCH = envs.SGLANG_MOE_NVFP4_DISPATCH.get()
|
||||
|
||||
# FlashInfer keys MNNVL allocations by workspace size; aligned tail padding gives
|
||||
# concurrently live paths distinct persistent workspaces without extra token work.
|
||||
_WORKSPACE_NAMESPACE_ALIGNMENT = 128
|
||||
@@ -91,6 +93,7 @@ class FlashinferDispatchOutput(NamedTuple):
|
||||
topk_output: StandardTopKOutput
|
||||
# Provide an output tensor to fused_moe so it writes directly to our buffer
|
||||
moe_output: Optional[torch.Tensor] = None
|
||||
output_dtype: Optional[torch.dtype] = None
|
||||
|
||||
@property
|
||||
def format(self) -> DispatchOutputFormat:
|
||||
@@ -151,6 +154,7 @@ class FlashinferDispatcher(BaseDispatcher):
|
||||
)
|
||||
# TODO: Can other moe runners use payload_in_workspace too?
|
||||
self.payload_in_workspace = get_moe_runner_backend().is_flashinfer_cutlass()
|
||||
self.dispatch_type = get_flashinfer_a2a_dispatch_type()
|
||||
if moe_runner_config is None:
|
||||
from sglang.srt.layers.moe.moe_runner.base import MoeRunnerConfig
|
||||
|
||||
@@ -184,24 +188,42 @@ class FlashinferDispatcher(BaseDispatcher):
|
||||
if configured_max_tokens is not None
|
||||
else default_max_tokens
|
||||
)
|
||||
|
||||
# Calculate workspace size. For eagle mode, use the larger workspace size since nextn layer will be unquantized.
|
||||
speculative_algo = SpeculativeAlgorithm.from_string(
|
||||
get_spec().speculative_algorithm
|
||||
)
|
||||
if MOE_NVFP4_DISPATCH and not speculative_algo.is_eagle():
|
||||
total_dispatch_payload_size_per_token = (
|
||||
can_use_quantized_dispatch = not speculative_algo.is_eagle()
|
||||
topk_id_and_weight_bytes = self.router_topk * 4 + self.router_topk * 4
|
||||
bf16_dispatch_payload_size_per_token = (
|
||||
hidden_size * 2 + topk_id_and_weight_bytes # bf16 hidden states
|
||||
)
|
||||
if (
|
||||
self.dispatch_type == FlashinferA2ADispatchType.NVFP4
|
||||
and can_use_quantized_dispatch
|
||||
):
|
||||
quantized_dispatch_payload_size_per_token = (
|
||||
hidden_size // 2 # nvfp4 hidden states
|
||||
+ hidden_size // 16 # fp8 scaling factors
|
||||
+ self.router_topk * 4 # int32 topks ids
|
||||
+ self.router_topk * 4 # float32 topk weights
|
||||
+ hidden_size // 16 # uint8 scaling factors
|
||||
+ topk_id_and_weight_bytes
|
||||
)
|
||||
total_dispatch_payload_size_per_token = max(
|
||||
quantized_dispatch_payload_size_per_token,
|
||||
bf16_dispatch_payload_size_per_token,
|
||||
)
|
||||
elif (
|
||||
self.dispatch_type == FlashinferA2ADispatchType.MXFP8
|
||||
and can_use_quantized_dispatch
|
||||
):
|
||||
quantized_dispatch_payload_size_per_token = (
|
||||
hidden_size # fp8 hidden states
|
||||
+ hidden_size // 32 # ue8m0 scaling factors
|
||||
+ topk_id_and_weight_bytes
|
||||
)
|
||||
total_dispatch_payload_size_per_token = max(
|
||||
quantized_dispatch_payload_size_per_token,
|
||||
bf16_dispatch_payload_size_per_token,
|
||||
)
|
||||
else:
|
||||
total_dispatch_payload_size_per_token = (
|
||||
hidden_size * 2 # bf16 hidden states
|
||||
+ self.router_topk * 4 # int32 topks ids
|
||||
+ self.router_topk * 4 # float32 topk weights
|
||||
)
|
||||
total_dispatch_payload_size_per_token = bf16_dispatch_payload_size_per_token
|
||||
combine_payload_size_per_token = hidden_size * 2 # bf16 hidden states
|
||||
self.workspace_size = moe_a2a_get_workspace_size_per_rank(
|
||||
ep_size=self.ep_size,
|
||||
@@ -293,6 +315,18 @@ class FlashinferDispatcher(BaseDispatcher):
|
||||
StandardTopKOutput(topk_weights, topk_ids, topk_output.router_logits),
|
||||
)
|
||||
|
||||
def _effective_dispatch_type(self) -> FlashinferA2ADispatchType:
|
||||
if self.dispatch_type == FlashinferA2ADispatchType.NVFP4:
|
||||
global_scale = (self.quant_config or {}).get("input_global_scale", None)
|
||||
if global_scale is None:
|
||||
return FlashinferA2ADispatchType.BF16
|
||||
elif self.dispatch_type == FlashinferA2ADispatchType.MXFP8:
|
||||
# Draft/NextN or mixed layers may not be MXFP8 even when the
|
||||
# process-wide default is MXFP8.
|
||||
if not (self.quant_config or {}).get("use_mxfp8", False):
|
||||
return FlashinferA2ADispatchType.BF16
|
||||
return self.dispatch_type
|
||||
|
||||
@debug_kernel_api
|
||||
def dispatch(
|
||||
self, hidden_states: torch.Tensor, topk_output: TopKOutput
|
||||
@@ -320,6 +354,7 @@ class FlashinferDispatcher(BaseDispatcher):
|
||||
)
|
||||
|
||||
output_dtype = hidden_states.dtype
|
||||
dispatch_type = self._effective_dispatch_type()
|
||||
x = hidden_states
|
||||
x_sf = None
|
||||
# FlashInfer dispatch requires materialized top-k IDs and weights.
|
||||
@@ -330,14 +365,36 @@ class FlashinferDispatcher(BaseDispatcher):
|
||||
topk_ids = topk_output.topk_ids.to(torch.int32)
|
||||
topk_weights = topk_output.topk_weights
|
||||
|
||||
global_scale = self.quant_config.get("input_global_scale", None)
|
||||
if global_scale is not None:
|
||||
if dispatch_type == FlashinferA2ADispatchType.NVFP4:
|
||||
global_scale = (self.quant_config or {}).get("input_global_scale", None)
|
||||
assert global_scale is not None
|
||||
if x.shape[0] > 0:
|
||||
x, x_sf = fp4_quantize(x, global_scale, is_sf_swizzled_layout=False)
|
||||
else:
|
||||
x_col = x.shape[1]
|
||||
x = torch.zeros(0, x_col // 2, dtype=torch.uint8, device=x.device)
|
||||
x_sf = torch.zeros(0, x_col // 16, dtype=torch.uint8, device=x.device)
|
||||
x = torch.zeros(
|
||||
0, self.hidden_size // 2, dtype=torch.uint8, device=x.device
|
||||
)
|
||||
x_sf = torch.zeros(
|
||||
0, self.hidden_size // 16, dtype=torch.uint8, device=x.device
|
||||
)
|
||||
elif dispatch_type == FlashinferA2ADispatchType.MXFP8:
|
||||
if x.shape[0] > 0:
|
||||
from flashinfer import mxfp8_quantize
|
||||
|
||||
x, x_sf = mxfp8_quantize(x, False)
|
||||
x_sf = x_sf.view(torch.uint8).reshape(
|
||||
x.shape[0], self.hidden_size // 32
|
||||
)
|
||||
else:
|
||||
x = torch.zeros(
|
||||
0,
|
||||
self.hidden_size,
|
||||
dtype=torch.float8_e4m3fn,
|
||||
device=x.device,
|
||||
)
|
||||
x_sf = torch.zeros(
|
||||
0, self.hidden_size // 32, dtype=torch.uint8, device=x.device
|
||||
)
|
||||
|
||||
payloads = []
|
||||
payloads.append(x)
|
||||
@@ -423,12 +480,17 @@ class FlashinferDispatcher(BaseDispatcher):
|
||||
if x_sf is not None:
|
||||
x_recv, x_sf_recv, topk_ids_recv, topk_weights_recv = recv_tensors
|
||||
x_sf = x_sf_recv.view(-1, x_sf_recv.shape[-1])
|
||||
# TODO: fuse interleave into cutlass moe
|
||||
if get_moe_runner_backend().is_flashinfer_cutlass():
|
||||
# TODO: Fuse interleave into cutlass moe when FlashInfer supports it.
|
||||
if (
|
||||
dispatch_type == FlashinferA2ADispatchType.NVFP4
|
||||
and get_moe_runner_backend().is_flashinfer_cutlass()
|
||||
):
|
||||
x_sf = nvfp4_block_scale_interleave(x_sf)
|
||||
else:
|
||||
x_recv, topk_ids_recv, topk_weights_recv = recv_tensors
|
||||
x = x_recv.view(-1, x_recv.shape[-1])
|
||||
if dispatch_type == FlashinferA2ADispatchType.MXFP8:
|
||||
x = x.view(torch.float8_e4m3fn)
|
||||
topk_ids = topk_ids_recv.view(-1, topk_ids_recv.shape[-1])
|
||||
topk_weights = topk_weights_recv.view(-1, topk_weights_recv.shape[-1])
|
||||
|
||||
@@ -443,6 +505,7 @@ class FlashinferDispatcher(BaseDispatcher):
|
||||
x_sf,
|
||||
StandardTopKOutput(topk_weights, topk_ids, topk_output.router_logits),
|
||||
moe_output,
|
||||
output_dtype,
|
||||
)
|
||||
|
||||
@debug_kernel_api
|
||||
|
||||
@@ -112,6 +112,7 @@ class StandardDispatcher(BaseDispatcher):
|
||||
# - cutlass / cutedsl / trtllm_routed handle EP internally
|
||||
# - mxfp4 dispatcher mapping is already global
|
||||
# - hpc_ops consumes global ids together with rank_ep / num_expert_total
|
||||
# - flashinfer_megamoe routes by global expert ID inside the mega kernel
|
||||
self.skip_local_expert_mapping = (
|
||||
backend.is_flashinfer_cutlass()
|
||||
or backend.is_flashinfer_cutedsl()
|
||||
@@ -119,6 +120,7 @@ class StandardDispatcher(BaseDispatcher):
|
||||
or backend.is_experimental_sgl_trtllm()
|
||||
or backend.is_flashinfer_trtllm_routed()
|
||||
or backend.is_hpc_ops()
|
||||
or backend.is_flashinfer_megamoe()
|
||||
or self.enable_flashinfer_mxfp4_moe
|
||||
)
|
||||
self.num_experts = moe_runner_config.num_experts
|
||||
|
||||
@@ -43,6 +43,7 @@ class MoeA2ABackend(Enum):
|
||||
MEGAMOE = "megamoe"
|
||||
DEEPEP_V2 = "deepep_v2"
|
||||
PPLX = "pplx"
|
||||
FLASHINFER_MEGAMOE = "flashinfer_megamoe"
|
||||
CUSTOMIZED = "customized"
|
||||
|
||||
@classmethod
|
||||
@@ -87,6 +88,9 @@ class MoeA2ABackend(Enum):
|
||||
def is_pplx(self):
|
||||
return self == MoeA2ABackend.PPLX
|
||||
|
||||
def is_flashinfer_megamoe(self):
|
||||
return self == MoeA2ABackend.FLASHINFER_MEGAMOE
|
||||
|
||||
def is_customized(self):
|
||||
return self == MoeA2ABackend.CUSTOMIZED
|
||||
|
||||
@@ -141,6 +145,9 @@ class _MoeRunnerBackendPredicates:
|
||||
def is_flashinfer_cutedsl(self):
|
||||
return self.value == MoeRunnerBackend.FLASHINFER_CUTEDSL.value
|
||||
|
||||
def is_flashinfer_megamoe(self):
|
||||
return self.value == MoeRunnerBackend.FLASHINFER_MEGAMOE.value
|
||||
|
||||
def is_flashinfer_mxfp4(self):
|
||||
return self.value == MoeRunnerBackend.FLASHINFER_MXFP4.value
|
||||
|
||||
@@ -165,6 +172,9 @@ class _MoeRunnerBackendPredicates:
|
||||
def is_aiter(self):
|
||||
return self.value == MoeRunnerBackend.AITER.value
|
||||
|
||||
def is_intel_xpu(self):
|
||||
return self.value == MoeRunnerBackend.INTEL_XPU.value
|
||||
|
||||
|
||||
class MoeRunnerBackend(_MoeRunnerBackendPredicates, Enum):
|
||||
AUTO = "auto"
|
||||
@@ -178,6 +188,7 @@ class MoeRunnerBackend(_MoeRunnerBackendPredicates, Enum):
|
||||
FLASHINFER_CUTLASS = "flashinfer_cutlass"
|
||||
FLASHINFER_MXFP4 = "flashinfer_mxfp4"
|
||||
FLASHINFER_CUTEDSL = "flashinfer_cutedsl"
|
||||
FLASHINFER_MEGAMOE = "flashinfer_megamoe"
|
||||
CUTLASS = "cutlass"
|
||||
MARLIN = "marlin"
|
||||
HUMMING = "humming"
|
||||
@@ -227,9 +238,6 @@ def resolve_moe_runner_backend(
|
||||
f"MoE runner backend {backend!r} is neither built in nor registered"
|
||||
) from None
|
||||
|
||||
def is_intel_xpu(self):
|
||||
return self == MoeRunnerBackend.INTEL_XPU
|
||||
|
||||
|
||||
class DeepEPv2Fp8ScaleFormat(NamedTuple):
|
||||
"""DeepGEMM FP8 activation-scale layout expected from DeepEP v2."""
|
||||
@@ -286,6 +294,33 @@ class DispatcherOutputDtype(Enum):
|
||||
MXFP8 = "mxfp8"
|
||||
|
||||
|
||||
class FlashinferA2ADispatchType(Enum):
|
||||
BF16 = "bf16"
|
||||
NVFP4 = "nvfp4"
|
||||
MXFP8 = "mxfp8"
|
||||
|
||||
|
||||
def get_flashinfer_a2a_dispatch_type() -> FlashinferA2ADispatchType:
|
||||
dispatch_type = get_exec().moe.flashinfer_a2a_dispatch_type
|
||||
|
||||
if dispatch_type is None:
|
||||
if envs.SGLANG_MOE_NVFP4_DISPATCH.is_set():
|
||||
return (
|
||||
FlashinferA2ADispatchType.NVFP4
|
||||
if envs.SGLANG_MOE_NVFP4_DISPATCH.get()
|
||||
else FlashinferA2ADispatchType.BF16
|
||||
)
|
||||
return FlashinferA2ADispatchType.BF16
|
||||
|
||||
if dispatch_type != "auto":
|
||||
return FlashinferA2ADispatchType(dispatch_type)
|
||||
|
||||
raise RuntimeError(
|
||||
"flashinfer_a2a_dispatch_type='auto' reached the published runtime "
|
||||
"configuration; ServerArgs must resolve it before publication"
|
||||
)
|
||||
|
||||
|
||||
def get_deepep_output_dtype(self) -> DispatcherOutputDtype:
|
||||
"""
|
||||
Automatically choose the dispatch output dtype for DeepEP.
|
||||
@@ -716,6 +751,11 @@ def should_skip_post_experts_all_reduce(*, is_tp_path: bool) -> bool:
|
||||
# pplx's AllToAll.combine already sums each token's expert outputs back
|
||||
# to the source rank
|
||||
return True
|
||||
if get_moe_a2a_backend().is_flashinfer_megamoe():
|
||||
# The mega kernel does its EP all-to-all + combine internally and
|
||||
# returns per-rank outputs, so any further EP/TP all-reduce would
|
||||
# double-count. Same opt-in as the flashinfer a2a dispatcher.
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
|
||||
@@ -1667,6 +1667,13 @@ class Fp8MoEMethod(FusedMoEMethodBase):
|
||||
layer.w2_weight.contiguous(), (16, 16)
|
||||
)
|
||||
return
|
||||
elif self.use_mxfp8 and get_moe_a2a_backend().is_flashinfer_megamoe():
|
||||
from sglang.srt.layers.moe.flashinfer_megamoe import (
|
||||
prepare_mxfp8_moe_weights_for_flashinfer_megamoe,
|
||||
)
|
||||
|
||||
prepare_mxfp8_moe_weights_for_flashinfer_megamoe(layer)
|
||||
return
|
||||
elif self.use_mxfp8:
|
||||
self._process_mxfp8_moe_weights(
|
||||
layer, quantize=not self.quant_config.is_checkpoint_fp8_serialized
|
||||
@@ -1762,6 +1769,14 @@ class Fp8MoEMethod(FusedMoEMethodBase):
|
||||
build_mega_moe_experts_weights(layer)
|
||||
return
|
||||
|
||||
if get_moe_a2a_backend().is_flashinfer_megamoe():
|
||||
from sglang.srt.layers.moe.flashinfer_megamoe import (
|
||||
prepare_fp4_moe_weights_for_flashinfer_megamoe,
|
||||
)
|
||||
|
||||
prepare_fp4_moe_weights_for_flashinfer_megamoe(layer)
|
||||
return
|
||||
|
||||
if deep_gemm_wrapper.DEEPGEMM_SCALE_UE8M0 and will_use_deepgemm:
|
||||
from deep_gemm import transform_sf_into_required_layout
|
||||
|
||||
@@ -2246,7 +2261,12 @@ class Fp8MoEMethod(FusedMoEMethodBase):
|
||||
self._prepare_hpc_ops_weights(layer)
|
||||
|
||||
if hasattr(layer, "dispatcher"):
|
||||
layer.dispatcher.set_quant_config({"weight_dtype": layer.w13_weight.dtype})
|
||||
layer.dispatcher.set_quant_config(
|
||||
{
|
||||
"weight_dtype": layer.w13_weight.dtype,
|
||||
"use_mxfp8": self.use_mxfp8,
|
||||
}
|
||||
)
|
||||
|
||||
def _prepare_flashinfer_trtllm_activation_params(self, layer: Module) -> None:
|
||||
"""Materialize optional TRT-LLM SwiGLU parameters once per expert."""
|
||||
@@ -2440,6 +2460,7 @@ class Fp8MoEMethod(FusedMoEMethodBase):
|
||||
or moe_runner_backend.is_flashinfer_trtllm()
|
||||
or moe_runner_backend.is_flashinfer_trtllm_routed()
|
||||
or moe_runner_backend.is_hpc_ops()
|
||||
or moe_runner_backend.is_flashinfer_megamoe()
|
||||
):
|
||||
self.runner = MoeRunner(moe_runner_backend, moe_runner_config)
|
||||
self._owns_moe_runner = True
|
||||
@@ -2598,7 +2619,31 @@ class Fp8MoEMethod(FusedMoEMethodBase):
|
||||
)
|
||||
return StandardCombineInput(hidden_states=output)
|
||||
|
||||
if self.runner.runner_backend.is_deep_gemm():
|
||||
if self.runner.runner_backend.is_flashinfer_megamoe():
|
||||
from sglang.srt.layers.moe.flashinfer_megamoe import (
|
||||
FlashInferMegaMoeQuantInfo,
|
||||
ensure_fp4_moe_layer_for_flashinfer_megamoe,
|
||||
ensure_mxfp8_moe_layer_for_flashinfer_megamoe,
|
||||
)
|
||||
|
||||
if self.use_mxfp8:
|
||||
ensure_megamoe_layer = ensure_mxfp8_moe_layer_for_flashinfer_megamoe
|
||||
elif self.is_fp4_expert:
|
||||
ensure_megamoe_layer = ensure_fp4_moe_layer_for_flashinfer_megamoe
|
||||
else:
|
||||
raise ValueError(
|
||||
"FlashInfer MegaMOE does not support standard FP8 MoE "
|
||||
"weights; use MXFP8, NVFP4, or an FP4-expert checkpoint."
|
||||
)
|
||||
mega = ensure_megamoe_layer(layer)
|
||||
quant_info = FlashInferMegaMoeQuantInfo(
|
||||
mega=mega,
|
||||
mega_forward=layer._flashinfer_megamoe_forward,
|
||||
apply_routed_scaling_factor=(
|
||||
not layer.should_fuse_routed_scaling_factor_in_topk
|
||||
),
|
||||
)
|
||||
elif self.runner.runner_backend.is_deep_gemm():
|
||||
w13_weight = layer.w13_weight
|
||||
w2_weight = layer.w2_weight
|
||||
|
||||
|
||||
@@ -16,10 +16,13 @@ from sglang.srt.layers.moe import (
|
||||
MoeRunner,
|
||||
MoeRunnerBackend,
|
||||
MoeRunnerConfig,
|
||||
get_moe_a2a_backend,
|
||||
get_moe_runner_backend,
|
||||
)
|
||||
from sglang.srt.layers.moe.moe_runner.triton import TritonMoeQuantInfo
|
||||
from sglang.srt.layers.moe.utils import (
|
||||
FlashinferA2ADispatchType,
|
||||
get_flashinfer_a2a_dispatch_type,
|
||||
is_flashinfer_cutedsl_v1_path,
|
||||
should_use_flashinfer_cutlass_moe_fp4_allgather,
|
||||
)
|
||||
@@ -291,6 +294,12 @@ MOE_NVFP4_DISPATCH = envs.SGLANG_MOE_NVFP4_DISPATCH.get()
|
||||
ACTIVATION_SCHEMES = ["static"]
|
||||
|
||||
|
||||
def _use_nvfp4_dispatch() -> bool:
|
||||
if not get_moe_a2a_backend().is_flashinfer():
|
||||
return MOE_NVFP4_DISPATCH
|
||||
return get_flashinfer_a2a_dispatch_type() == FlashinferA2ADispatchType.NVFP4
|
||||
|
||||
|
||||
_SUPPORTED_ACT_STRS = ("silu", "relu2", "gelu")
|
||||
|
||||
|
||||
@@ -2210,6 +2219,31 @@ class ModelOptNvFp4A16LinearMethod(LinearMethodBase):
|
||||
)
|
||||
|
||||
|
||||
def _input_scale_to_local_experts(
|
||||
input_scale: torch.Tensor,
|
||||
num_local_experts: int,
|
||||
num_experts: int,
|
||||
moe_ep_rank: int,
|
||||
) -> torch.Tensor:
|
||||
"""Normalize a checkpoint input scale to this rank's local experts.
|
||||
|
||||
Checkpoints may store the activation scale as a scalar, a per-local-expert
|
||||
vector, or a global per-expert vector; return a (num_local_experts,) vector.
|
||||
"""
|
||||
input_scale = input_scale.detach().to(torch.float32)
|
||||
if input_scale.dim() == 0:
|
||||
return input_scale.expand(num_local_experts).contiguous()
|
||||
if input_scale.shape == (num_local_experts,):
|
||||
return input_scale.contiguous()
|
||||
if input_scale.shape == (num_experts,):
|
||||
start = moe_ep_rank * num_local_experts
|
||||
return input_scale[start : start + num_local_experts].contiguous()
|
||||
raise ValueError(
|
||||
f"input scale must be scalar, ({num_local_experts},), or "
|
||||
f"({num_experts},); got {tuple(input_scale.shape)}"
|
||||
)
|
||||
|
||||
|
||||
def _compute_gemm1_alphas(
|
||||
w13_weight_scale_2: torch.Tensor,
|
||||
w13_input_scale: torch.Tensor,
|
||||
@@ -2383,7 +2417,10 @@ class ModelOptNvFp4FusedMoEMethod(FusedMoEMethodBase):
|
||||
# TRTLLM replaces blockscale_swizzled with an alias to weight_scale
|
||||
# during process_weights_after_loading, so skip the expensive
|
||||
# swizzle+allocate here to avoid GPU memory fragmentation
|
||||
if self.enable_flashinfer_trtllm_moe:
|
||||
if (
|
||||
self.enable_flashinfer_trtllm_moe
|
||||
or get_moe_runner_backend().is_flashinfer_megamoe()
|
||||
):
|
||||
layer.w13_blockscale_swizzled = None
|
||||
else:
|
||||
layer.w13_blockscale_swizzled = Parameter(
|
||||
@@ -2403,7 +2440,10 @@ class ModelOptNvFp4FusedMoEMethod(FusedMoEMethodBase):
|
||||
)
|
||||
layer.register_parameter("w2_weight_scale", w2_weight_scale)
|
||||
|
||||
if self.enable_flashinfer_trtllm_moe:
|
||||
if (
|
||||
self.enable_flashinfer_trtllm_moe
|
||||
or get_moe_runner_backend().is_flashinfer_megamoe()
|
||||
):
|
||||
layer.w2_blockscale_swizzled = None
|
||||
else:
|
||||
layer.w2_blockscale_swizzled = Parameter(
|
||||
@@ -2493,6 +2533,7 @@ class ModelOptNvFp4FusedMoEMethod(FusedMoEMethodBase):
|
||||
moe_runner_backend = getattr(
|
||||
self, "_moe_runner_backend", get_moe_runner_backend()
|
||||
)
|
||||
use_nvfp4_dispatch = _use_nvfp4_dispatch()
|
||||
if moe_runner_backend.is_marlin():
|
||||
# Marlin supports only a single shared w1/w3 weight scale, so collapse
|
||||
# the gate/up columns to the gate scale here. Other backends keep the
|
||||
@@ -2526,6 +2567,18 @@ class ModelOptNvFp4FusedMoEMethod(FusedMoEMethodBase):
|
||||
if self.enable_flashinfer_cutlass_moe or self.enable_flashinfer_trtllm_moe:
|
||||
w13_input_scale = layer.w13_input_scale.max().to(torch.float32)
|
||||
w2_input_scale = layer.w2_input_scale.max().to(torch.float32)
|
||||
elif moe_runner_backend.is_flashinfer_megamoe():
|
||||
# MegaMOE folds a scalar w13 input scale into input_norm_const but keeps
|
||||
# per-expert w2 scales, so g2_alphas / w2_input_scale_quant stay
|
||||
# per-expert to feed the mega kernel's fc2_alpha / fc1_norm_const (keeps
|
||||
# FC1-output renorm and FC2 dequant on the same per-expert scale).
|
||||
w13_input_scale = layer.w13_input_scale.max().to(torch.float32)
|
||||
w2_input_scale = _input_scale_to_local_experts(
|
||||
layer.w2_input_scale,
|
||||
layer.num_local_experts,
|
||||
layer.num_experts,
|
||||
layer.moe_ep_rank,
|
||||
)
|
||||
elif self.enable_flashinfer_cutedsl_moe:
|
||||
# CuteDSL standard path uses a single scalar input scale (all experts).
|
||||
w13_input_scale = (
|
||||
@@ -2548,7 +2601,7 @@ class ModelOptNvFp4FusedMoEMethod(FusedMoEMethodBase):
|
||||
w13_input_scale = _slice_scale(w13_input_scale)
|
||||
w2_input_scale = _slice_scale(w2_input_scale)
|
||||
|
||||
if MOE_NVFP4_DISPATCH:
|
||||
if use_nvfp4_dispatch:
|
||||
assert torch.all(w13_input_scale == w13_input_scale[0])
|
||||
w13_input_scale = w13_input_scale[0]
|
||||
else:
|
||||
@@ -2633,7 +2686,7 @@ class ModelOptNvFp4FusedMoEMethod(FusedMoEMethodBase):
|
||||
not self.quant_config.use_per_token_activation
|
||||
and not use_cutedsl_w4a16
|
||||
and (
|
||||
MOE_NVFP4_DISPATCH or should_use_flashinfer_cutlass_moe_fp4_allgather()
|
||||
use_nvfp4_dispatch or should_use_flashinfer_cutlass_moe_fp4_allgather()
|
||||
)
|
||||
)
|
||||
|
||||
@@ -2672,6 +2725,14 @@ class ModelOptNvFp4FusedMoEMethod(FusedMoEMethodBase):
|
||||
f"{name} Weight Blockscale must be represented as FP8-E4M3"
|
||||
)
|
||||
|
||||
if moe_runner_backend.is_flashinfer_megamoe():
|
||||
from sglang.srt.layers.moe.flashinfer_megamoe import (
|
||||
prepare_nvfp4_moe_weights_for_flashinfer_megamoe,
|
||||
)
|
||||
|
||||
prepare_nvfp4_moe_weights_for_flashinfer_megamoe(layer)
|
||||
return
|
||||
|
||||
# Weight processing based on strategy
|
||||
if (
|
||||
self.enable_flashinfer_trtllm_moe
|
||||
@@ -2900,6 +2961,25 @@ class ModelOptNvFp4FusedMoEMethod(FusedMoEMethodBase):
|
||||
), f"{activation=} is unsupported by {moe_runner_backend}"
|
||||
moe_runner_config = self.moe_runner_config
|
||||
|
||||
if moe_runner_backend.is_flashinfer_megamoe():
|
||||
from sglang.srt.layers.moe.flashinfer_megamoe import (
|
||||
FlashInferMegaMoeQuantInfo,
|
||||
ensure_nvfp4_moe_layer_for_flashinfer_megamoe,
|
||||
)
|
||||
|
||||
mega = ensure_nvfp4_moe_layer_for_flashinfer_megamoe(layer)
|
||||
quant_info = FlashInferMegaMoeQuantInfo(
|
||||
mega=mega,
|
||||
mega_forward=layer._flashinfer_megamoe_forward,
|
||||
fc1_alpha=layer.g1_alphas,
|
||||
fc2_alpha=layer.g2_alphas,
|
||||
fc1_norm_const=layer.w2_input_scale_quant,
|
||||
apply_routed_scaling_factor=(
|
||||
not layer.should_fuse_routed_scaling_factor_in_topk
|
||||
),
|
||||
)
|
||||
return self.runner.run(dispatch_output, quant_info)
|
||||
|
||||
if moe_runner_backend.is_marlin():
|
||||
quant_info = self.get_marlin_quant_info(layer)
|
||||
return self.runner.run(dispatch_output, quant_info)
|
||||
|
||||
@@ -1081,6 +1081,19 @@ class ModelRunner:
|
||||
return self.sampling_prewarm_result
|
||||
|
||||
def init_cuda_graphs(self, capture_decode_cuda_graph: bool = True):
|
||||
# from sglang.srt.layers.moe.utils import get_moe_runner_backend
|
||||
|
||||
# if get_moe_runner_backend().is_flashinfer_megamoe():
|
||||
# # Warmup's dummy batches aren't guaranteed to route through every
|
||||
# # MoE layer; a layer that first builds mid-capture instead of
|
||||
# # during warmup hits a hard RuntimeError (capture forbids the
|
||||
# # lazy build's blocking device sync). Force every layer to build
|
||||
# # here, eagerly, outside any graph.
|
||||
# from sglang.srt.layers.moe.flashinfer_megamoe import (
|
||||
# warmup_all_flashinfer_megamoe_layers,
|
||||
# )
|
||||
|
||||
# warmup_all_flashinfer_megamoe_layers(self.model)
|
||||
capture = capture_cuda_graphs(
|
||||
model_runner=self, capture_decode_cuda_graph=capture_decode_cuda_graph
|
||||
)
|
||||
|
||||
@@ -717,6 +717,7 @@ class DeepseekV2MoE(nn.Module):
|
||||
or get_moe_a2a_backend().is_ascend_fuseep()
|
||||
or get_moe_a2a_backend().is_flashinfer()
|
||||
or get_moe_a2a_backend().is_megamoe()
|
||||
or get_moe_a2a_backend().is_flashinfer_megamoe()
|
||||
or get_moe_a2a_backend().is_deepep_v2()
|
||||
or should_use_flashinfer_cutlass_moe_fp4_allgather()
|
||||
or envs.SGLANG_SHARED_EXPERT_TP1.get()
|
||||
|
||||
@@ -468,6 +468,7 @@ class Glm4MoeSparseMoeBlock(nn.Module):
|
||||
or get_moe_a2a_backend().is_mori()
|
||||
or get_moe_a2a_backend().is_ascend_fuseep()
|
||||
or get_moe_a2a_backend().is_flashinfer()
|
||||
or get_moe_a2a_backend().is_flashinfer_megamoe()
|
||||
or should_use_flashinfer_cutlass_moe_fp4_allgather()
|
||||
else {}
|
||||
),
|
||||
|
||||
@@ -247,6 +247,7 @@ class NemotronHMoE(nn.Module):
|
||||
dict(tp_rank=0, tp_size=1)
|
||||
if get_moe_a2a_backend().is_deepep()
|
||||
or get_moe_a2a_backend().is_flashinfer()
|
||||
or get_moe_a2a_backend().is_flashinfer_megamoe()
|
||||
else {}
|
||||
),
|
||||
prefix=f"{prefix}.shared_experts",
|
||||
|
||||
@@ -376,6 +376,7 @@ class Qwen2MoeSparseMoeBlock(nn.Module):
|
||||
or get_moe_a2a_backend().is_mori()
|
||||
or get_moe_a2a_backend().is_deepep_v2()
|
||||
or get_moe_a2a_backend().is_flashinfer()
|
||||
or get_moe_a2a_backend().is_flashinfer_megamoe()
|
||||
)
|
||||
else {}
|
||||
),
|
||||
|
||||
@@ -318,9 +318,12 @@ class Qwen3MoeSparseMoeBlock(nn.Module):
|
||||
num_tokens, hidden_dim = hidden_states.shape
|
||||
hidden_states = hidden_states.view(-1, hidden_dim)
|
||||
|
||||
# router_logits: (num_tokens, n_experts)
|
||||
router_logits, _ = self.gate(hidden_states)
|
||||
topk_output = self.topk(hidden_states, router_logits)
|
||||
if hidden_states.shape[0] > 0:
|
||||
# router_logits: (num_tokens, n_experts)
|
||||
router_logits, _ = self.gate(hidden_states)
|
||||
topk_output = self.topk(hidden_states, router_logits)
|
||||
else:
|
||||
topk_output = self.topk.empty_topk_output(hidden_states.device)
|
||||
final_hidden_states = self.experts(hidden_states, topk_output)
|
||||
|
||||
if self.ep_size > 1 and not should_skip_post_experts_all_reduce(
|
||||
|
||||
@@ -622,6 +622,7 @@ class Resources(_FlagGroupBase):
|
||||
# Persistent reusable CUDA events for non-EP DP TBO, keyed by
|
||||
# (kind, subbatch) — see dp_attention._tbo_event for why reuse matters.
|
||||
tbo_event_pool: dict = msgspec.field(default_factory=dict)
|
||||
flashinfer_megamoe_workspaces: dict = msgspec.field(default_factory=dict)
|
||||
# State capturers (installed by their subsystems when capture is on).
|
||||
indexer_capturer: Any = None
|
||||
experts_capturer: Any = None
|
||||
|
||||
@@ -4,13 +4,15 @@ import torch
|
||||
|
||||
from sglang.srt.distributed import init_distributed_environment
|
||||
from sglang.srt.distributed.parallel_state import (
|
||||
destroy_distributed_environment,
|
||||
destroy_model_parallel,
|
||||
get_tp_group,
|
||||
initialize_model_parallel,
|
||||
)
|
||||
from sglang.srt.layers.dp_attention import set_dp_buffer_len
|
||||
from sglang.srt.layers.moe.token_dispatcher.flashinfer import FlashinferDispatcher
|
||||
from sglang.srt.layers.moe.utils import initialize_moe_config
|
||||
from sglang.srt.runtime_context import publish
|
||||
from sglang.srt.runtime_context import get_context, publish
|
||||
from sglang.srt.server_args import ServerArgs, set_global_server_args_for_scheduler
|
||||
from sglang.test.test_utils import CustomTestCase
|
||||
|
||||
@@ -21,6 +23,7 @@ class TestFlashinferDispatcher(CustomTestCase):
|
||||
server_args = ServerArgs(model_path="dummy")
|
||||
server_args.moe_runner_backend = "flashinfer_cutlass"
|
||||
server_args.moe_a2a_backend = "flashinfer"
|
||||
cls.server_args = server_args
|
||||
set_global_server_args_for_scheduler(server_args)
|
||||
publish(server_args, role="scheduler")
|
||||
initialize_moe_config()
|
||||
@@ -41,9 +44,18 @@ class TestFlashinferDispatcher(CustomTestCase):
|
||||
|
||||
@classmethod
|
||||
def tearDownClass(cls):
|
||||
# Clean up distributed environment
|
||||
if torch.distributed.is_initialized():
|
||||
torch.distributed.destroy_process_group()
|
||||
try:
|
||||
from flashinfer.comm.trtllm_moe_alltoall import MoeAlltoAll
|
||||
|
||||
for workspace in MoeAlltoAll._WORKSPACE_CACHE.values():
|
||||
mnnvl_mem = workspace.get("mnnvl_mem")
|
||||
if mnnvl_mem is not None and "ptr" in vars(mnnvl_mem):
|
||||
del mnnvl_mem.ptr
|
||||
MoeAlltoAll._WORKSPACE_CACHE.clear()
|
||||
except ImportError:
|
||||
pass
|
||||
destroy_model_parallel()
|
||||
destroy_distributed_environment()
|
||||
|
||||
def create_dispatcher(
|
||||
self, router_topk=2, num_experts=8, num_local_experts=4, hidden_size=128
|
||||
@@ -58,8 +70,40 @@ class TestFlashinferDispatcher(CustomTestCase):
|
||||
params_dtype=torch.bfloat16,
|
||||
)
|
||||
|
||||
def set_dispatch_type(self, dispatch_type):
|
||||
get_context().override(
|
||||
"test_flashinfer_dispatcher",
|
||||
flashinfer_a2a_dispatch_type=dispatch_type,
|
||||
)
|
||||
|
||||
def _zero_moe_a2a_dispatch_payloads(self):
|
||||
# Shared MoeAlltoAll workspaces keep stale recv payloads across tests.
|
||||
# Zero only the payload region so unused-source == 0 asserts stay valid.
|
||||
try:
|
||||
from flashinfer.comm.trtllm_moe_alltoall import (
|
||||
MoeAlltoAll,
|
||||
get_moe_alltoall_module,
|
||||
)
|
||||
except ImportError:
|
||||
return
|
||||
|
||||
module = get_moe_alltoall_module()
|
||||
for ws in MoeAlltoAll._WORKSPACE_CACHE.values():
|
||||
workspace = ws["workspace"]
|
||||
aux = int(
|
||||
module.moe_a2a_get_aux_data_size(
|
||||
ws["ep_size"],
|
||||
ws["max_num_tokens"],
|
||||
ws["eplb_stats_num_experts"],
|
||||
)
|
||||
)
|
||||
aux = ((aux + 127) // 128) * 128
|
||||
if aux < workspace.shape[1]:
|
||||
workspace[:, aux:].zero_()
|
||||
|
||||
def test_dispatch_basic(self):
|
||||
"""Test basic dispatch functionality"""
|
||||
self.set_dispatch_type("bf16")
|
||||
num_tokens = 16
|
||||
hidden_size = 128
|
||||
router_topk = 1 # Single expert per token for simplicity
|
||||
@@ -143,9 +187,10 @@ class TestFlashinferDispatcher(CustomTestCase):
|
||||
|
||||
def test_dispatch_with_empty_tokens(self):
|
||||
"""Test dispatch when there are no tokens (edge case)"""
|
||||
self.set_dispatch_type("bf16")
|
||||
# This tests the dummy token handling
|
||||
num_tokens = 16
|
||||
hidden_size = 1
|
||||
hidden_size = 128
|
||||
router_topk = 1 # Single expert per token for simplicity
|
||||
world_size = torch.distributed.get_world_size()
|
||||
rank = torch.distributed.get_rank()
|
||||
@@ -195,6 +240,9 @@ class TestFlashinferDispatcher(CustomTestCase):
|
||||
topk_weights=topk_weights, topk_ids=topk_ids, router_logits=None
|
||||
)
|
||||
|
||||
self._zero_moe_a2a_dispatch_payloads()
|
||||
torch.distributed.barrier()
|
||||
|
||||
dispatcher = self.create_dispatcher(
|
||||
router_topk=router_topk,
|
||||
num_experts=num_experts,
|
||||
@@ -250,6 +298,7 @@ class TestFlashinferDispatcher(CustomTestCase):
|
||||
|
||||
def test_dispatch_with_fp4_quantization(self):
|
||||
"""Test dispatch with FP4 quantization enabled"""
|
||||
self.set_dispatch_type("nvfp4")
|
||||
num_tokens = 128
|
||||
hidden_size = 128
|
||||
router_topk = 1 # Single expert per token for simplicity
|
||||
@@ -312,6 +361,133 @@ class TestFlashinferDispatcher(CustomTestCase):
|
||||
)
|
||||
self.assertEqual(dispatch_output.hidden_states_scale.dtype, torch.uint8)
|
||||
|
||||
def test_dispatch_with_mxfp8_quantization(self):
|
||||
"""Test dispatch with MXFP8 quantization enabled"""
|
||||
self.set_dispatch_type("mxfp8")
|
||||
num_tokens = 128
|
||||
hidden_size = 128
|
||||
router_topk = 1
|
||||
world_size = torch.distributed.get_world_size()
|
||||
rank = torch.distributed.get_rank()
|
||||
num_experts = world_size
|
||||
num_local_experts = 1
|
||||
|
||||
set_dp_buffer_len(
|
||||
global_dp_buffer_len=num_tokens * world_size,
|
||||
local_dp_buffer_len=num_tokens,
|
||||
dp_max_padding=True,
|
||||
global_num_tokens=None,
|
||||
)
|
||||
|
||||
hidden_states = torch.randn(
|
||||
(num_tokens, hidden_size), dtype=torch.bfloat16, device="cuda"
|
||||
)
|
||||
|
||||
target_rank = (rank + 1) % world_size
|
||||
target_expert = target_rank
|
||||
topk_ids = torch.full(
|
||||
(num_tokens, router_topk), target_expert, dtype=torch.int32, device="cuda"
|
||||
)
|
||||
topk_weights = torch.ones(
|
||||
(num_tokens, router_topk), dtype=torch.float32, device="cuda"
|
||||
)
|
||||
|
||||
from sglang.srt.layers.moe.topk import StandardTopKOutput
|
||||
|
||||
topk_output = StandardTopKOutput(
|
||||
topk_weights=topk_weights, topk_ids=topk_ids, router_logits=None
|
||||
)
|
||||
|
||||
dispatcher = self.create_dispatcher(
|
||||
router_topk=router_topk,
|
||||
num_experts=num_experts,
|
||||
num_local_experts=num_local_experts,
|
||||
hidden_size=hidden_size,
|
||||
)
|
||||
dispatcher.set_quant_config({"input_global_scale": None, "use_mxfp8": True})
|
||||
|
||||
dispatch_output = dispatcher.dispatch(hidden_states, topk_output)
|
||||
|
||||
self.assertEqual(
|
||||
dispatch_output.hidden_states.shape,
|
||||
(num_tokens * world_size, hidden_size),
|
||||
)
|
||||
self.assertEqual(dispatch_output.hidden_states.dtype, torch.float8_e4m3fn)
|
||||
self.assertEqual(dispatch_output.output_dtype, torch.bfloat16)
|
||||
|
||||
self.assertIsNotNone(dispatch_output.hidden_states_scale)
|
||||
self.assertEqual(
|
||||
dispatch_output.hidden_states_scale.shape,
|
||||
(num_tokens * world_size, hidden_size // 32),
|
||||
)
|
||||
self.assertEqual(dispatch_output.hidden_states_scale.dtype, torch.uint8)
|
||||
self.assertEqual(
|
||||
dispatch_output.topk_output.topk_ids.shape,
|
||||
(num_tokens * world_size, router_topk),
|
||||
)
|
||||
self.assertEqual(dispatch_output.topk_output.topk_ids.dtype, torch.int32)
|
||||
|
||||
def test_dispatch_with_mxfp8_quantization_and_empty_rank(self):
|
||||
"""All ranks must contribute the same payload dtypes, including empty ranks."""
|
||||
self.set_dispatch_type("mxfp8")
|
||||
num_tokens = 16
|
||||
hidden_size = 128
|
||||
router_topk = 1
|
||||
world_size = torch.distributed.get_world_size()
|
||||
rank = torch.distributed.get_rank()
|
||||
empty_rank = 1
|
||||
|
||||
global_num_tokens = [num_tokens] * world_size
|
||||
global_num_tokens[empty_rank] = 0
|
||||
set_dp_buffer_len(
|
||||
global_dp_buffer_len=num_tokens * world_size,
|
||||
local_dp_buffer_len=num_tokens,
|
||||
dp_max_padding=False,
|
||||
global_num_tokens=global_num_tokens,
|
||||
)
|
||||
|
||||
local_tokens = 0 if rank == empty_rank else num_tokens
|
||||
hidden_states = torch.randn(
|
||||
(local_tokens, hidden_size), dtype=torch.bfloat16, device="cuda"
|
||||
)
|
||||
target_expert = (rank + 1) % world_size
|
||||
topk_ids = torch.full(
|
||||
(local_tokens, router_topk),
|
||||
target_expert,
|
||||
dtype=torch.int32,
|
||||
device="cuda",
|
||||
)
|
||||
topk_weights = torch.ones(
|
||||
(local_tokens, router_topk), dtype=torch.float32, device="cuda"
|
||||
)
|
||||
|
||||
from sglang.srt.layers.moe.topk import StandardTopKOutput
|
||||
|
||||
dispatcher = self.create_dispatcher(
|
||||
router_topk=router_topk,
|
||||
num_experts=world_size,
|
||||
num_local_experts=1,
|
||||
hidden_size=hidden_size,
|
||||
)
|
||||
dispatcher.set_quant_config({"input_global_scale": None, "use_mxfp8": True})
|
||||
self._zero_moe_a2a_dispatch_payloads()
|
||||
torch.distributed.barrier()
|
||||
dispatch_output = dispatcher.dispatch(
|
||||
hidden_states,
|
||||
StandardTopKOutput(
|
||||
topk_weights=topk_weights,
|
||||
topk_ids=topk_ids,
|
||||
router_logits=None,
|
||||
),
|
||||
)
|
||||
|
||||
self.assertEqual(dispatch_output.hidden_states.dtype, torch.float8_e4m3fn)
|
||||
self.assertEqual(dispatch_output.hidden_states_scale.dtype, torch.uint8)
|
||||
self.assertEqual(
|
||||
dispatch_output.hidden_states.shape,
|
||||
(num_tokens * world_size, hidden_size),
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
"""
|
||||
|
||||
@@ -155,6 +155,56 @@ class FlashinferTrtllmGenMoeBackendMXFP8Base:
|
||||
self.assertGreater(metrics["score"], 0.93)
|
||||
|
||||
|
||||
class FlashinferTrtllmGenMoeBackendMXFP8A2ABase:
|
||||
backend = "flashinfer_trtllm_routed"
|
||||
|
||||
@classmethod
|
||||
def setUpClass(cls):
|
||||
cls.model = "zianglih/Qwen3-30B-A3B-Instruct-2507-MXFP8"
|
||||
cls.base_url = DEFAULT_URL_FOR_TEST
|
||||
cls.process = popen_launch_server(
|
||||
cls.model,
|
||||
cls.base_url,
|
||||
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
||||
env={**os.environ, "SGLANG_ENABLE_JIT_DEEPGEMM": "False"},
|
||||
other_args=[
|
||||
"--quantization",
|
||||
"mxfp8",
|
||||
"--enable-dp-attention",
|
||||
"--dp-size",
|
||||
"4",
|
||||
"--tp-size",
|
||||
"4",
|
||||
"--moe-a2a-backend",
|
||||
"flashinfer",
|
||||
"--moe-runner-backend",
|
||||
cls.backend,
|
||||
"--flashinfer-a2a-dispatch-type",
|
||||
"mxfp8",
|
||||
"--mem-fraction-static",
|
||||
"0.7",
|
||||
],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def tearDownClass(cls):
|
||||
kill_process_tree(cls.process.pid)
|
||||
|
||||
def test_gsm8k(self):
|
||||
args = SimpleNamespace(
|
||||
base_url=self.base_url,
|
||||
model=self.model,
|
||||
eval_name="gsm8k",
|
||||
api="completion",
|
||||
max_tokens=512,
|
||||
num_examples=200,
|
||||
num_threads=128,
|
||||
)
|
||||
metrics = run_eval(args)
|
||||
print(f"{metrics=}")
|
||||
self.assertGreater(metrics["score"], 0.93)
|
||||
|
||||
|
||||
class FlashinferTrtllmGenMoeBackendMXFP8MixedBF16Base:
|
||||
backend = None
|
||||
|
||||
@@ -261,6 +311,12 @@ class TestFlashinferTrtllmGenMoeBackendMXFP8Routed(
|
||||
backend = "flashinfer_trtllm_routed"
|
||||
|
||||
|
||||
class TestFlashinferTrtllmGenMoeBackendMXFP8A2A(
|
||||
FlashinferTrtllmGenMoeBackendMXFP8A2ABase, CustomTestCase
|
||||
):
|
||||
pass
|
||||
|
||||
|
||||
class TestFlashinferTrtllmRoutedMxfp8MixedBF16(
|
||||
FlashinferTrtllmGenMoeBackendMXFP8MixedBF16Base, CustomTestCase
|
||||
):
|
||||
|
||||
@@ -0,0 +1,62 @@
|
||||
import sys
|
||||
from unittest.mock import patch
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.layers.moe.token_dispatcher.flashinfer import FlashinferDispatcher
|
||||
from sglang.srt.layers.moe.topk import StandardTopKOutput
|
||||
from sglang.srt.layers.moe.utils import FlashinferA2ADispatchType
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
|
||||
register_cpu_ci(est_time=1, suite="base-a-test-cpu")
|
||||
|
||||
|
||||
def test_empty_mxfp8_dispatch_uses_same_payload_dtype_as_nonempty_rank():
|
||||
class FakeMoeAlltoAll:
|
||||
def dispatch(self, _topk_ids, payloads, *_args, **_kwargs):
|
||||
self.payload_dtypes = [payload.dtype for payload in payloads]
|
||||
return payloads
|
||||
|
||||
dispatcher = object.__new__(FlashinferDispatcher)
|
||||
dispatcher.dispatch_type = FlashinferA2ADispatchType.MXFP8
|
||||
dispatcher.hidden_size = 128
|
||||
dispatcher.max_num_tokens = 0
|
||||
dispatcher.ep_size = 1
|
||||
dispatcher.invalid_token_expert_id = 8
|
||||
dispatcher.payload_in_workspace = False
|
||||
dispatcher.quant_config = {"use_mxfp8": True}
|
||||
dispatcher.moe_a2a = FakeMoeAlltoAll()
|
||||
|
||||
hidden_states = torch.empty((0, 128), dtype=torch.bfloat16)
|
||||
topk_output = StandardTopKOutput(
|
||||
topk_weights=torch.empty((0, 1), dtype=torch.float32),
|
||||
topk_ids=torch.empty((0, 1), dtype=torch.int32),
|
||||
router_logits=None,
|
||||
)
|
||||
|
||||
with (
|
||||
patch(
|
||||
"sglang.srt.layers.moe.token_dispatcher.flashinfer.get_dp_global_num_tokens",
|
||||
return_value=None,
|
||||
),
|
||||
patch(
|
||||
"sglang.srt.layers.moe.token_dispatcher.flashinfer.is_dp_attention_enabled",
|
||||
return_value=False,
|
||||
),
|
||||
):
|
||||
output = dispatcher.dispatch(hidden_states, topk_output)
|
||||
|
||||
assert dispatcher.moe_a2a.payload_dtypes == [
|
||||
torch.float8_e4m3fn,
|
||||
torch.uint8,
|
||||
torch.int32,
|
||||
torch.float32,
|
||||
]
|
||||
assert output.hidden_states.dtype == torch.float8_e4m3fn
|
||||
assert output.hidden_states_scale.dtype == torch.uint8
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
import pytest
|
||||
|
||||
sys.exit(pytest.main([__file__, "-v"]))
|
||||
@@ -0,0 +1,220 @@
|
||||
import importlib.util
|
||||
import sys
|
||||
import types
|
||||
from pathlib import Path
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
|
||||
register_cpu_ci(est_time=5, suite="base-a-test-cpu")
|
||||
|
||||
|
||||
def _load_megamoe_module(monkeypatch):
|
||||
"""Load the adapter with only its small import-time dependencies stubbed."""
|
||||
|
||||
class MoeQuantInfo:
|
||||
pass
|
||||
|
||||
class MoeRunnerConfig:
|
||||
pass
|
||||
|
||||
def register_fused_func(*_args, **_kwargs):
|
||||
return lambda fn: fn
|
||||
|
||||
fake_modules = {
|
||||
"sglang": types.ModuleType("sglang"),
|
||||
"sglang.srt": types.ModuleType("sglang.srt"),
|
||||
"sglang.srt.environ": types.ModuleType("sglang.srt.environ"),
|
||||
"sglang.srt.layers": types.ModuleType("sglang.srt.layers"),
|
||||
"sglang.srt.layers.moe": types.ModuleType("sglang.srt.layers.moe"),
|
||||
"sglang.srt.layers.moe.moe_runner": types.ModuleType(
|
||||
"sglang.srt.layers.moe.moe_runner"
|
||||
),
|
||||
"sglang.srt.layers.moe.moe_runner.base": types.ModuleType(
|
||||
"sglang.srt.layers.moe.moe_runner.base"
|
||||
),
|
||||
"sglang.srt.layers.moe.token_dispatcher": types.ModuleType(
|
||||
"sglang.srt.layers.moe.token_dispatcher"
|
||||
),
|
||||
"sglang.srt.runtime_context": types.ModuleType("sglang.srt.runtime_context"),
|
||||
"deep_gemm": types.ModuleType("deep_gemm"),
|
||||
"deep_gemm.utils": types.ModuleType("deep_gemm.utils"),
|
||||
"deep_gemm.utils.math": types.ModuleType("deep_gemm.utils.math"),
|
||||
}
|
||||
fake_modules["sglang.srt.environ"].envs = types.SimpleNamespace(
|
||||
SGLANG_FLASHINFER_MEGAMOE_MAX_TOKENS_PER_RANK=types.SimpleNamespace(
|
||||
get=lambda: 0
|
||||
),
|
||||
SGLANG_FLASHINFER_MEGAMOE_COMBINE_DTYPE=types.SimpleNamespace(
|
||||
get=lambda: "bf16"
|
||||
),
|
||||
SGLANG_FLASHINFER_MEGAMOE_IN_KERNEL_FC2_REDUCE=types.SimpleNamespace(
|
||||
get=lambda: False
|
||||
),
|
||||
)
|
||||
runtime_context = fake_modules["sglang.srt.runtime_context"]
|
||||
runtime_context.cutedsl_moe_max_num_tokens = lambda: 2048
|
||||
base = fake_modules["sglang.srt.layers.moe.moe_runner.base"]
|
||||
base.MoeQuantInfo = MoeQuantInfo
|
||||
base.MoeRunnerConfig = MoeRunnerConfig
|
||||
base.register_fused_func = register_fused_func
|
||||
token_dispatcher = fake_modules["sglang.srt.layers.moe.token_dispatcher"]
|
||||
|
||||
class StandardCombineInput:
|
||||
def __init__(self, *, hidden_states):
|
||||
self.hidden_states = hidden_states
|
||||
|
||||
token_dispatcher.StandardCombineInput = StandardCombineInput
|
||||
for name, module in fake_modules.items():
|
||||
monkeypatch.setitem(sys.modules, name, module)
|
||||
|
||||
module_path = (
|
||||
Path(__file__).resolve().parents[5]
|
||||
/ "python/sglang/srt/layers/moe/flashinfer_megamoe.py"
|
||||
)
|
||||
module_name = "sglang_flashinfer_megamoe_adapter_test"
|
||||
spec = importlib.util.spec_from_file_location(module_name, module_path)
|
||||
assert spec is not None and spec.loader is not None
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
monkeypatch.setitem(sys.modules, module_name, module)
|
||||
spec.loader.exec_module(module)
|
||||
return module
|
||||
|
||||
|
||||
def test_max_tokens_uses_runtime_context_accessor(monkeypatch):
|
||||
module = _load_megamoe_module(monkeypatch)
|
||||
|
||||
assert module._resolve_max_tokens_per_rank() == 2048
|
||||
|
||||
runtime_context = sys.modules["sglang.srt.runtime_context"]
|
||||
runtime_context.cutedsl_moe_max_num_tokens = lambda: 0
|
||||
assert module._resolve_max_tokens_per_rank() == 1024
|
||||
|
||||
|
||||
def test_adapter_keeps_router_ids_int32(monkeypatch):
|
||||
module = _load_megamoe_module(monkeypatch)
|
||||
|
||||
class FakeMoEEpTensors:
|
||||
def __init__(self, **kwargs):
|
||||
self.__dict__.update(kwargs)
|
||||
|
||||
fake_moe_ep = types.ModuleType("flashinfer.moe_ep")
|
||||
fake_moe_ep.MoEEpTensors = FakeMoEEpTensors
|
||||
fake_flashinfer = types.ModuleType("flashinfer")
|
||||
fake_flashinfer.moe_ep = fake_moe_ep
|
||||
monkeypatch.setitem(sys.modules, "flashinfer", fake_flashinfer)
|
||||
monkeypatch.setitem(sys.modules, "flashinfer.moe_ep", fake_moe_ep)
|
||||
|
||||
hidden_states = torch.randn((3, 4), dtype=torch.bfloat16)
|
||||
topk_ids = torch.tensor([[0, 1], [1, 0], [0, 1]], dtype=torch.int32)
|
||||
topk_weights = torch.randn((3, 2), dtype=torch.float32)
|
||||
output = torch.randn_like(hidden_states)
|
||||
|
||||
class Mega:
|
||||
_workspace = object()
|
||||
|
||||
def forward(self, tensors):
|
||||
self.tensors = tensors
|
||||
return output
|
||||
|
||||
mega = Mega()
|
||||
dispatch_output = types.SimpleNamespace(
|
||||
hidden_states=hidden_states,
|
||||
topk_output=types.SimpleNamespace(
|
||||
topk_ids=topk_ids,
|
||||
topk_weights=topk_weights,
|
||||
),
|
||||
)
|
||||
quant_info = module.FlashInferMegaMoeQuantInfo(mega=mega)
|
||||
runner_config = types.SimpleNamespace(routed_scaling_factor=1.0)
|
||||
|
||||
result = module.run_flashinfer_megamoe(
|
||||
dispatch_output,
|
||||
quant_info,
|
||||
runner_config,
|
||||
)
|
||||
|
||||
assert mega.tensors.topk_ids.data_ptr() == topk_ids.data_ptr()
|
||||
assert mega.tensors.topk_ids.dtype == torch.int32
|
||||
assert result.hidden_states is output
|
||||
|
||||
|
||||
def test_adapter_requests_workspace_output_view(monkeypatch):
|
||||
module = _load_megamoe_module(monkeypatch)
|
||||
|
||||
class FakeMoEEpTensors:
|
||||
def __init__(self, **kwargs):
|
||||
self.__dict__.update(kwargs)
|
||||
|
||||
fake_moe_ep = types.ModuleType("flashinfer.moe_ep")
|
||||
fake_moe_ep.MoEEpTensors = FakeMoEEpTensors
|
||||
fake_flashinfer = types.ModuleType("flashinfer")
|
||||
fake_flashinfer.moe_ep = fake_moe_ep
|
||||
monkeypatch.setitem(sys.modules, "flashinfer", fake_flashinfer)
|
||||
monkeypatch.setitem(sys.modules, "flashinfer.moe_ep", fake_moe_ep)
|
||||
|
||||
hidden_states = torch.randn((2, 4), dtype=torch.bfloat16)
|
||||
topk_ids = torch.tensor([[0, 1], [1, 0]], dtype=torch.int32)
|
||||
topk_weights = torch.ones((2, 2), dtype=torch.float32)
|
||||
output = torch.randn_like(hidden_states)
|
||||
|
||||
class Mega:
|
||||
supports_output_view = True
|
||||
_workspace = object()
|
||||
|
||||
def forward(self, tensors, *, return_workspace_view=False):
|
||||
self.tensors = tensors
|
||||
self.return_workspace_view = return_workspace_view
|
||||
return output
|
||||
|
||||
mega = Mega()
|
||||
dispatch_output = types.SimpleNamespace(
|
||||
hidden_states=hidden_states,
|
||||
topk_output=types.SimpleNamespace(
|
||||
topk_ids=topk_ids,
|
||||
topk_weights=topk_weights,
|
||||
),
|
||||
)
|
||||
|
||||
result = module.run_flashinfer_megamoe(
|
||||
dispatch_output,
|
||||
module.FlashInferMegaMoeQuantInfo(mega=mega),
|
||||
types.SimpleNamespace(routed_scaling_factor=1.0),
|
||||
)
|
||||
|
||||
assert result.hidden_states is output
|
||||
assert mega.tensors.topk_ids.data_ptr() == topk_ids.data_ptr()
|
||||
assert mega.tensors.topk_ids.dtype == torch.int32
|
||||
assert mega.return_workspace_view is True
|
||||
|
||||
|
||||
def test_capture_safe_ue8m0_pack_is_scoped(monkeypatch):
|
||||
module = _load_megamoe_module(monkeypatch)
|
||||
|
||||
dgm = sys.modules["deep_gemm.utils.math"]
|
||||
|
||||
def original(value):
|
||||
return value
|
||||
|
||||
dgm.pack_ue8m0_to_int = original
|
||||
monkeypatch.setattr(torch.cuda, "is_available", lambda: True)
|
||||
monkeypatch.setattr(torch.cuda, "is_current_stream_capturing", lambda: False)
|
||||
|
||||
with module._capture_safe_ue8m0_pack():
|
||||
assert dgm.pack_ue8m0_to_int is original
|
||||
|
||||
monkeypatch.setattr(torch.cuda, "is_current_stream_capturing", lambda: True)
|
||||
|
||||
with module._capture_safe_ue8m0_pack():
|
||||
assert dgm.pack_ue8m0_to_int is not original
|
||||
packed = dgm.pack_ue8m0_to_int(torch.ones(4, dtype=torch.float32))
|
||||
assert packed.dtype == torch.int32
|
||||
|
||||
assert dgm.pack_ue8m0_to_int is original
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
import pytest
|
||||
|
||||
sys.exit(pytest.main([__file__, "-v"]))
|
||||
@@ -906,17 +906,19 @@ class TestModuleLevelHelpers(unittest.TestCase):
|
||||
# Without a specific flashinfer backend selected, default is False.
|
||||
self.assertFalse(_moe_runner_keeps_global_expert_ids())
|
||||
|
||||
def test_real_backend_predicate_matches_dispatcher_and_pool(self):
|
||||
def test_real_backend_predicates_match_supported_id_contracts(self):
|
||||
backends = _load_moe_backend_enum()
|
||||
expected_global = {
|
||||
dispatcher_global_ids = {
|
||||
backends.FLASHINFER_TRTLLM,
|
||||
backends.EXPERIMENTAL_SGL_TRTLLM,
|
||||
backends.FLASHINFER_TRTLLM_ROUTED,
|
||||
backends.FLASHINFER_CUTLASS,
|
||||
backends.FLASHINFER_MXFP4,
|
||||
backends.FLASHINFER_CUTEDSL,
|
||||
backends.FLASHINFER_MEGAMOE,
|
||||
backends.HPC_OPS,
|
||||
}
|
||||
lora_global_ids = dispatcher_global_ids - {backends.FLASHINFER_MEGAMOE}
|
||||
config = types.SimpleNamespace(
|
||||
num_experts=8,
|
||||
num_local_experts=2,
|
||||
@@ -938,11 +940,11 @@ class TestModuleLevelHelpers(unittest.TestCase):
|
||||
dispatcher = standard_dispatcher(config)
|
||||
self.assertEqual(
|
||||
dispatcher.skip_local_expert_mapping,
|
||||
backend in expected_global,
|
||||
backend in dispatcher_global_ids,
|
||||
)
|
||||
self.assertEqual(
|
||||
_moe_runner_keeps_global_expert_ids(),
|
||||
backend in expected_global,
|
||||
backend in lora_global_ids,
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -76,6 +76,10 @@ from sglang.srt.entrypoints.sidecar import (
|
||||
)
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.layers.cp.base import is_cp_enabled, is_interleave
|
||||
from sglang.srt.layers.moe.utils import (
|
||||
FlashinferA2ADispatchType,
|
||||
get_flashinfer_a2a_dispatch_type,
|
||||
)
|
||||
from sglang.srt.model_executor.cuda_graph_config import (
|
||||
Backend,
|
||||
CudaGraphConfig,
|
||||
@@ -1123,6 +1127,264 @@ class TestContextParallelServerArgs(CustomTestCase):
|
||||
self.assertTrue(is_interleave())
|
||||
|
||||
|
||||
class TestFlashinferA2ADispatchType(CustomTestCase):
|
||||
def setUp(self):
|
||||
self._nvfp4_env_backup = os.environ.get("SGLANG_MOE_NVFP4_DISPATCH")
|
||||
envs.SGLANG_MOE_NVFP4_DISPATCH.clear()
|
||||
|
||||
def tearDown(self):
|
||||
if self._nvfp4_env_backup is None:
|
||||
envs.SGLANG_MOE_NVFP4_DISPATCH.clear()
|
||||
else:
|
||||
os.environ["SGLANG_MOE_NVFP4_DISPATCH"] = self._nvfp4_env_backup
|
||||
|
||||
def _make_args(
|
||||
self,
|
||||
quantization=None,
|
||||
dispatch_type=None,
|
||||
runner_backend="flashinfer_trtllm_routed",
|
||||
):
|
||||
server_args = ServerArgs(
|
||||
model_path="dummy",
|
||||
quantization=quantization,
|
||||
moe_a2a_backend="flashinfer",
|
||||
moe_runner_backend=runner_backend,
|
||||
flashinfer_a2a_dispatch_type=dispatch_type,
|
||||
enable_dp_attention=True,
|
||||
dp_size=4,
|
||||
tp_size=4,
|
||||
)
|
||||
server_args._model_config = SimpleNamespace(nvfp4_moe_meta=None)
|
||||
return server_args
|
||||
|
||||
def test_auto_resolves_mxfp8_and_normalizes_trtllm(self):
|
||||
server_args = self._make_args(
|
||||
quantization="mxfp8",
|
||||
dispatch_type="auto",
|
||||
runner_backend="flashinfer_trtllm",
|
||||
)
|
||||
handle_a2a_moe(server_args)
|
||||
|
||||
self.assertEqual(
|
||||
resolution_result(server_args, "moe_runner_backend"),
|
||||
"flashinfer_trtllm_routed",
|
||||
)
|
||||
self.assertEqual(
|
||||
resolution_result(server_args, "flashinfer_a2a_dispatch_type"), "mxfp8"
|
||||
)
|
||||
|
||||
def test_auto_resolves_modelopt_fp4_to_nvfp4(self):
|
||||
server_args = self._make_args(quantization="modelopt_fp4", dispatch_type="auto")
|
||||
handle_a2a_moe(server_args)
|
||||
|
||||
self.assertEqual(
|
||||
resolution_result(server_args, "flashinfer_a2a_dispatch_type"), "nvfp4"
|
||||
)
|
||||
|
||||
def test_auto_resolves_hybrid_nvfp4_metadata_to_nvfp4(self):
|
||||
server_args = self._make_args(quantization="fp8", dispatch_type="auto")
|
||||
server_args._model_config = SimpleNamespace(nvfp4_moe_meta={})
|
||||
handle_a2a_moe(server_args)
|
||||
|
||||
self.assertEqual(
|
||||
resolution_result(server_args, "flashinfer_a2a_dispatch_type"), "nvfp4"
|
||||
)
|
||||
|
||||
def test_unspecified_preserves_legacy_nvfp4_auto_enable(self):
|
||||
server_args = self._make_args(quantization="modelopt_fp4")
|
||||
handle_a2a_moe(server_args)
|
||||
|
||||
self.assertIsNone(server_args.flashinfer_a2a_dispatch_type)
|
||||
self.assertTrue(envs.SGLANG_MOE_NVFP4_DISPATCH.get())
|
||||
|
||||
def test_unspecified_getter_preserves_legacy_bf16_fallback(self):
|
||||
with get_context().override_server_args(
|
||||
flashinfer_a2a_dispatch_type=None,
|
||||
quantization="mxfp8",
|
||||
):
|
||||
self.assertEqual(
|
||||
get_flashinfer_a2a_dispatch_type(),
|
||||
FlashinferA2ADispatchType.BF16,
|
||||
)
|
||||
|
||||
def test_runtime_getter_rejects_unresolved_auto(self):
|
||||
with get_context().override_server_args(
|
||||
flashinfer_a2a_dispatch_type="auto",
|
||||
):
|
||||
with self.assertRaisesRegex(RuntimeError, "must resolve it"):
|
||||
get_flashinfer_a2a_dispatch_type()
|
||||
|
||||
def test_explicit_nvfp4_checks_hybrid_metadata_for_mxfp8_quantization(self):
|
||||
server_args = self._make_args(quantization="mxfp8", dispatch_type="nvfp4")
|
||||
server_args._model_config = SimpleNamespace(nvfp4_moe_meta={})
|
||||
handle_a2a_moe(server_args)
|
||||
|
||||
self.assertEqual(
|
||||
resolution_result(server_args, "flashinfer_a2a_dispatch_type"), "nvfp4"
|
||||
)
|
||||
|
||||
def test_explicit_bf16_overrides_auto(self):
|
||||
server_args = self._make_args(quantization="modelopt_fp4", dispatch_type="bf16")
|
||||
handle_a2a_moe(server_args)
|
||||
|
||||
self.assertEqual(
|
||||
resolution_result(server_args, "flashinfer_a2a_dispatch_type"), "bf16"
|
||||
)
|
||||
|
||||
def test_legacy_env_maps_to_dispatch_type(self):
|
||||
with envs.SGLANG_MOE_NVFP4_DISPATCH.override("1"):
|
||||
server_args = self._make_args(quantization="modelopt_fp4")
|
||||
handle_a2a_moe(server_args)
|
||||
self.assertIsNone(server_args.flashinfer_a2a_dispatch_type)
|
||||
|
||||
with envs.SGLANG_MOE_NVFP4_DISPATCH.override("0"):
|
||||
server_args = self._make_args(quantization="modelopt_fp4")
|
||||
handle_a2a_moe(server_args)
|
||||
self.assertIsNone(server_args.flashinfer_a2a_dispatch_type)
|
||||
|
||||
def test_legacy_env_conflicts_with_explicit_cli(self):
|
||||
with envs.SGLANG_MOE_NVFP4_DISPATCH.override("1"):
|
||||
server_args = self._make_args(
|
||||
quantization="modelopt_fp4", dispatch_type="bf16"
|
||||
)
|
||||
with self.assertRaisesRegex(
|
||||
ValueError, "SGLANG_MOE_NVFP4_DISPATCH cannot be set"
|
||||
):
|
||||
handle_a2a_moe(server_args)
|
||||
|
||||
def test_mxfp8_dispatch_requires_mxfp8_quantization(self):
|
||||
server_args = self._make_args(quantization="fp8", dispatch_type="mxfp8")
|
||||
with self.assertRaisesRegex(ValueError, "requires --quantization mxfp8"):
|
||||
handle_a2a_moe(server_args)
|
||||
|
||||
def test_explicit_dispatch_type_requires_flashinfer_a2a(self):
|
||||
server_args = ServerArgs(
|
||||
model_path="dummy",
|
||||
moe_a2a_backend="none",
|
||||
flashinfer_a2a_dispatch_type="bf16",
|
||||
)
|
||||
with self.assertRaisesRegex(ValueError, "requires --moe-a2a-backend"):
|
||||
handle_a2a_moe(server_args)
|
||||
|
||||
|
||||
class TestFlashinferMegaMoeConfig(CustomTestCase):
|
||||
def setUp(self):
|
||||
self._combine_dtype_backup = os.environ.get(
|
||||
"SGLANG_FLASHINFER_MEGAMOE_COMBINE_DTYPE"
|
||||
)
|
||||
self._ikr_backup = os.environ.get(
|
||||
"SGLANG_FLASHINFER_MEGAMOE_IN_KERNEL_FC2_REDUCE"
|
||||
)
|
||||
envs.SGLANG_FLASHINFER_MEGAMOE_COMBINE_DTYPE.clear()
|
||||
envs.SGLANG_FLASHINFER_MEGAMOE_IN_KERNEL_FC2_REDUCE.clear()
|
||||
|
||||
def tearDown(self):
|
||||
if self._combine_dtype_backup is None:
|
||||
envs.SGLANG_FLASHINFER_MEGAMOE_COMBINE_DTYPE.clear()
|
||||
else:
|
||||
os.environ["SGLANG_FLASHINFER_MEGAMOE_COMBINE_DTYPE"] = (
|
||||
self._combine_dtype_backup
|
||||
)
|
||||
if self._ikr_backup is None:
|
||||
envs.SGLANG_FLASHINFER_MEGAMOE_IN_KERNEL_FC2_REDUCE.clear()
|
||||
else:
|
||||
os.environ["SGLANG_FLASHINFER_MEGAMOE_IN_KERNEL_FC2_REDUCE"] = (
|
||||
self._ikr_backup
|
||||
)
|
||||
|
||||
def _make_args(
|
||||
self,
|
||||
architecture="DeepseekV4ForCausalLM",
|
||||
quantization="modelopt_fp4",
|
||||
*,
|
||||
is_fp4_experts=False,
|
||||
nvfp4_moe_meta=None,
|
||||
):
|
||||
server_args = ServerArgs(
|
||||
model_path="dummy",
|
||||
quantization=quantization,
|
||||
moe_a2a_backend="flashinfer_megamoe",
|
||||
moe_runner_backend="flashinfer_megamoe",
|
||||
enable_dp_attention=True,
|
||||
dp_size=4,
|
||||
tp_size=4,
|
||||
)
|
||||
server_args._model_config = SimpleNamespace(
|
||||
hf_config=SimpleNamespace(architectures=[architecture]),
|
||||
is_fp4_experts=is_fp4_experts,
|
||||
nvfp4_moe_meta=nvfp4_moe_meta,
|
||||
)
|
||||
return server_args
|
||||
|
||||
@patch("sglang.srt.arg_groups.moe_hook.is_sm100_supported", return_value=True)
|
||||
def test_megamoe_accepts_audited_model_architectures(self, _):
|
||||
supported = (
|
||||
"DeepseekV2ForCausalLM",
|
||||
"DeepseekV3ForCausalLM",
|
||||
"DeepseekV32ForCausalLM",
|
||||
"DeepseekV4ForCausalLM",
|
||||
"Glm4MoeForCausalLM",
|
||||
"NemotronHForCausalLM",
|
||||
"NemotronHPuzzleForCausalLM",
|
||||
"Qwen2MoeForCausalLM",
|
||||
"Qwen3MoeForCausalLM",
|
||||
)
|
||||
for architecture in supported:
|
||||
with self.subTest(architecture=architecture):
|
||||
handle_a2a_moe(self._make_args(architecture))
|
||||
|
||||
def test_megamoe_rejects_unaudited_model_architecture(self):
|
||||
with self.assertRaisesRegex(
|
||||
ValueError,
|
||||
"not validated for model architectures.*UnsupportedMoeForCausalLM",
|
||||
):
|
||||
handle_a2a_moe(self._make_args("UnsupportedMoeForCausalLM"))
|
||||
|
||||
@patch("sglang.srt.arg_groups.moe_hook.is_sm100_supported", return_value=True)
|
||||
def test_megamoe_accepts_supported_quantization_formats(self, _):
|
||||
supported = (
|
||||
{"quantization": "modelopt_fp4"},
|
||||
{"quantization": "mxfp8"},
|
||||
{"quantization": "fp8", "is_fp4_experts": True},
|
||||
{"quantization": "modelopt_mixed", "nvfp4_moe_meta": {}},
|
||||
)
|
||||
for config in supported:
|
||||
with self.subTest(config=config):
|
||||
handle_a2a_moe(self._make_args(**config))
|
||||
|
||||
def test_megamoe_rejects_standard_fp8(self):
|
||||
with self.assertRaisesRegex(ValueError, "Standard FP8 MoE checkpoints"):
|
||||
handle_a2a_moe(self._make_args(quantization="fp8"))
|
||||
|
||||
@patch("sglang.srt.arg_groups.moe_hook.is_sm100_supported", return_value=False)
|
||||
def test_megamoe_requires_sm100_for_all_quantization_formats(self, _):
|
||||
with self.assertRaisesRegex(ValueError, "requires an SM100-family"):
|
||||
handle_a2a_moe(self._make_args(quantization="fp8", is_fp4_experts=True))
|
||||
|
||||
@patch("sglang.srt.arg_groups.moe_hook.is_sm100_supported", return_value=True)
|
||||
def test_megamoe_combine_dtype_accepts_quantized_values(self, _):
|
||||
with envs.SGLANG_FLASHINFER_MEGAMOE_COMBINE_DTYPE.override("nvfp4"):
|
||||
handle_a2a_moe(self._make_args())
|
||||
|
||||
with envs.SGLANG_FLASHINFER_MEGAMOE_COMBINE_DTYPE.override("mxfp8"):
|
||||
handle_a2a_moe(self._make_args())
|
||||
|
||||
@patch("sglang.srt.arg_groups.moe_hook.is_sm100_supported", return_value=True)
|
||||
def test_megamoe_combine_dtype_rejects_invalid_value(self, _):
|
||||
with envs.SGLANG_FLASHINFER_MEGAMOE_COMBINE_DTYPE.override("fp8"):
|
||||
with self.assertRaisesRegex(
|
||||
ValueError, "SGLANG_FLASHINFER_MEGAMOE_COMBINE_DTYPE"
|
||||
):
|
||||
handle_a2a_moe(self._make_args())
|
||||
|
||||
@patch("sglang.srt.arg_groups.moe_hook.is_sm100_supported", return_value=True)
|
||||
def test_megamoe_combine_dtype_conflicts_with_ikr(self, _):
|
||||
with envs.SGLANG_FLASHINFER_MEGAMOE_COMBINE_DTYPE.override("nvfp4"):
|
||||
with envs.SGLANG_FLASHINFER_MEGAMOE_IN_KERNEL_FC2_REDUCE.override("1"):
|
||||
with self.assertRaisesRegex(ValueError, "incompatible"):
|
||||
handle_a2a_moe(self._make_args())
|
||||
|
||||
|
||||
class TestPortArgs(unittest.TestCase):
|
||||
@patch("sglang.srt.server_args.tempfile.NamedTemporaryFile")
|
||||
def test_init_new_standard_case(self, mock_temp_file):
|
||||
|
||||
@@ -2466,7 +2466,11 @@ class TestGoldenModelOverrides(_IsolatedPublish):
|
||||
)
|
||||
|
||||
def _view(**kw):
|
||||
defaults = dict(quantization=None, moe_runner_backend="auto")
|
||||
defaults = dict(
|
||||
quantization=None,
|
||||
moe_runner_backend="auto",
|
||||
moe_a2a_backend="none",
|
||||
)
|
||||
defaults.update(kw)
|
||||
return ResolvedView(SimpleNamespace(**defaults))
|
||||
|
||||
|
||||
Reference in New Issue
Block a user