diff --git a/docs_new/docs/advanced_features/quantization.mdx b/docs_new/docs/advanced_features/quantization.mdx
index 0fdbfc99b..35b0e812a 100644
--- a/docs_new/docs/advanced_features/quantization.mdx
+++ b/docs_new/docs/advanced_features/quantization.mdx
@@ -53,8 +53,8 @@ The following table summarizes quantization method support across NVIDIA and AMD
mxfp8 |
No |
No |
- Yes (A5 for Diffusion and LLM Dense Linear) |
- Ascend NPU only; online MXFP8 quantization for Diffusion models (e.g., Wan2.2) and LLM Dense Linear on A5 series; uses CANN npu_dynamic_mx_quant / npu_quant_matmul kernels |
+ Yes (A5 for Diffusion, LLM Dense Linear and LLM MoE) |
+ Ascend NPU only; online + offline MXFP8 for Diffusion models (e.g., Wan2.2), LLM Dense Linear, and LLM MoE (FusedMoE, e.g. Qwen3-30B-A3B) on A5 series; uses CANN npu_dynamic_mx_quant / npu_quant_matmul (dense) and npu_grouped_matmul_swiglu_quant_v2 / npu_grouped_matmul (MoE) kernels |
mxfp_w4a8 |
diff --git a/docs_new/docs/hardware-platforms/ascend-npus/ascend_npu_quantization.mdx b/docs_new/docs/hardware-platforms/ascend-npus/ascend_npu_quantization.mdx
index 086651ae6..8a7c01667 100644
--- a/docs_new/docs/hardware-platforms/ascend-npus/ascend_npu_quantization.mdx
+++ b/docs_new/docs/hardware-platforms/ascend-npus/ascend_npu_quantization.mdx
@@ -101,11 +101,11 @@ SGLang supports **mix-bits** quantization (independently defines and loads each
x |
- | MXFP8 |
+ MXFP8 (LLM MoE) |
MoE |
x |
x |
- WIP |
+ √ |
x |
@@ -361,7 +361,40 @@ python3 -m sglang.launch_server \
> **Implementation Notes:**
> - Online path: `Fp8Config.get_quant_method()` dispatches to `NPUMXFP8LinearMethod`. Weights are quantized once at load via `npu_dynamic_mx_quant(weight, dst_type=torch_npu.float8_e4m3fn)` and pre-transposed to `[in, out]`; activations are per-token quantized at inference and matmul runs via `npu_quant_matmul(..., group_sizes=[1, 1, 32])` (block_size = 32).
> - Offline path: `ModelSlimMXFP8Scheme` loads `float8_e4m3fn` weights + `float8_e8m0fnu` block scales pre-exported by msmodelslim. Transpose is kept as a non-contiguous view (`.data` assignment) — calling `.contiguous()` would physically reorder the pre-quantized layout and break the block-scale mapping.
-> - MoE MXFP8 (FusedMoE/TP) for LLMs is tracked separately and not part of this PR.
+> - MoE MXFP8 (FusedMoE) for LLMs is documented in **MXFP8 for LLM MoE models** below.
+
+**MXFP8 for LLM MoE models (e.g. Qwen3-30B-A3B / Qwen3.5 MoE):**
+
+LLM MoE W8A8 MXFP8 (FusedMoE) support builds on the dense MXFP8 path. Requires Ascend A5 series or newer — the fused MoE MX kernels (`npu_grouped_matmul_swiglu_quant_v2`, `npu_dynamic_mx_quant`) are A5-only.
+
+- Online MXFP8 quantization (BF16/FP16 expert weights → MXFP8 at load time):
+
+```bash Command
+python3 -m sglang.launch_server \
+ --model-path Qwen/Qwen3-30B-A3B \
+ --quantization mxfp8 \
+ --device npu --attention-backend ascend \
+ --host 0.0.0.0 --port 30000 \
+ --mem-fraction-static 0.8 --tp-size 1
+```
+
+- Offline MXFP8 quantization (msmodelslim pre-quantized weights, `W8A8_MXFP8` scheme). No `--quantization` flag is needed: the `quant_model_description.json` shipped with the checkpoint selects both the ModelSlim path and the scheme automatically.
+
+```bash Command
+python3 -m sglang.launch_server \
+ --model-path /path/to/Qwen3-30B-A3B-W8A8-MXFP8 \
+ --device npu --attention-backend ascend \
+ --host 0.0.0.0 --port 30000 \
+ --mem-fraction-static 0.8 --tp-size 1
+```
+
+> **Implementation Notes:**
+> - Both paths share the per-gmm kernel `NPUMXFP8MoEMethod` (`hardware_backend/npu/quantization/moe_methods.py`), which tells online from offline by weight dtype. Expert weights and their e8m0 block scales are kept as non-contiguous transpose views — calling `.contiguous()` would tank HBM bandwidth.
+> - Online path: `Fp8Config.get_quant_method()` dispatches FusedMoE layers to `NPUMXFP8OnlineMoEMethod`, which subclasses `UnquantizedFusedMoEMethod` and overrides only `create_moe_runner` to swap in the MXFP8 kernels — weight creation, weight post-processing and the forward pass are the unquantized Ascend ones. BF16 expert weights `w13`/`w2` are quantized once at load via `npu_dynamic_mx_quant(dst_type=torch.float8_e4m3fn)` (a 3D `[E, N, K]` input is accepted directly).
+> - Offline path: `ModelSlimMXFP8MoEScheme` (one instance per weight group) loads `float8_e4m3fn` expert weights + uint8 (e8m0, exponent + 127) block scales. The scale is reshaped `[E, N, K/32] → [E, N, K/64, 2]` (contiguous pairing, matching `npu_dynamic_mx_quant`) then transposed.
+> - Forward: `AscendTPDispatcher` runs `npu_moe_init_routing_v2(quant_mode=3)`, which fuses the per-token MX activation quant into routing (e4m3 payload + e8m0 block scale, reshaped to the pair-split layout). `AscendRunnerCore` then runs gmm1 `npu_grouped_matmul_swiglu_quant_v2` (cumulative `group_list`; fuses gate/up + swiglu + requant, so no separate activation step) → gmm2 `npu_grouped_matmul` (count `group_list`). The UE8M0 (`float8_e8m0fnu`) scale dtypes are passed explicitly; the e4m3 `x`/`weight` dtypes are left implicit.
+> - **Router gate**: msmodelslim may also quantize `mlp.gate` (`W8A8_MXFP8`). The gate is a `ReplicatedLinear`, so its quantization must be **description-driven**: for the offline `modelslim` path the gate is passed the quant config and dequantized correctly; the online path keeps it in BF16. Loading a quantized gate as BF16 without its block scale scrambles routing and produces garbage output.
+> - Where the activation quant happens depends on the dispatcher. On `ascend_tp` it is fused into routing as described above. DeepEP has no MXFP8 dispatch dtype, so it keeps dispatching BF16 and gmm1 quantizes the hidden states itself via `npu_dynamic_mx_quant` before the fused kernel — the two paths reach the same gmm1 input. Only the `ascend_tp` path has been validated end-to-end on Ascend A5.
**MXFP4 W4A8 for LLM dense models (e.g., Qwen3 / Qwen3.5):**
diff --git a/python/sglang/srt/arg_groups/overrides.py b/python/sglang/srt/arg_groups/overrides.py
index 842d7246d..7f430022d 100644
--- a/python/sglang/srt/arg_groups/overrides.py
+++ b/python/sglang/srt/arg_groups/overrides.py
@@ -2031,7 +2031,13 @@ def _moe_runner_backend_quant_constraints(view: Any) -> dict:
"--moe-runner-backend flashinfer_trtllm or "
"flashinfer_trtllm_routed."
)
- if view.quantization == "mxfp8":
+ # Ascend runs MXFP8 MoE on the Ascend runner; every backend selected below is
+ # CUDA/ROCm-only. Forcing one here would not merely pick the wrong runner:
+ # FusedMoE keys its w1/w3 shard swap ("flashinfer assumes w31") and its
+ # 128-alignment round-up off flashinfer_trtllm, so the experts would silently
+ # load with gate and up exchanged. Leave the backend at "auto" and let
+ # create_moe_runner resolve it to ASCEND.
+ if view.quantization == "mxfp8" and not is_npu():
from sglang.srt.server_args import MXFP8_MOE_RUNNER_BACKEND_CHOICES
is_gfx95_mxfp8 = is_hip() and is_gfx95_supported()
diff --git a/python/sglang/srt/hardware_backend/npu/moe/init_routing.py b/python/sglang/srt/hardware_backend/npu/moe/init_routing.py
index bce2acf19..e942a369e 100644
--- a/python/sglang/srt/hardware_backend/npu/moe/init_routing.py
+++ b/python/sglang/srt/hardware_backend/npu/moe/init_routing.py
@@ -11,6 +11,24 @@ from typing import Optional, Tuple
import torch
+# ``npu_moe_init_routing_v2`` quant_mode selecting MXFP8: the op emits an
+# float8_e4m3fn payload plus an e8m0 block scale, fusing the activation quant
+# that would otherwise need a separate ``npu_dynamic_mx_quant`` pass.
+MXFP8_QUANT_MODE = 3
+
+
+def _normalize_mxfp_scale(scale: torch.Tensor) -> torch.Tensor:
+ """Reshape a flat 2D e8m0 block scale ``[N, M]`` into pair-split ``[N, M//2, 2]``.
+
+ ``npu_moe_init_routing_v2(quant_mode=3)`` emits the scale flat, while the
+ grouped matmul wants the pair-split view. Already-3D scales (what
+ ``npu_dynamic_mx_quant`` returns) pass through untouched. Mirrors
+ vllm-ascend's ``maybe_normalize_mxfp_scale_layout``.
+ """
+ if scale is None or scale.ndim != 2:
+ return scale
+ return scale.reshape(scale.shape[0], scale.shape[1] // 2, 2)
+
class BaseInitRouting(ABC):
"""Abstract base for NPU MoE init routing."""
@@ -95,6 +113,8 @@ class NPUMoEInitRouting_v2(BaseInitRouting):
)
if self.quant_mode == -1:
pertoken_scale = None
+ elif self.quant_mode == MXFP8_QUANT_MODE:
+ pertoken_scale = _normalize_mxfp_scale(pertoken_scale)
expert_tokens = expert_tokens.to(torch.int64)
return hidden_states, expanded_row_idx, expert_tokens, pertoken_scale
diff --git a/python/sglang/srt/hardware_backend/npu/moe/matmul.py b/python/sglang/srt/hardware_backend/npu/moe/matmul.py
index 1ebcfc833..9075561a6 100644
--- a/python/sglang/srt/hardware_backend/npu/moe/matmul.py
+++ b/python/sglang/srt/hardware_backend/npu/moe/matmul.py
@@ -1,4 +1,5 @@
from abc import ABC, abstractmethod
+from typing import Tuple
import torch
@@ -47,3 +48,42 @@ class GroupedMatmul(BaseMatmul):
group_list=expert_tokens,
output_dtype=output_dtype,
)[0]
+
+
+class GroupedMatmulSwigluQuant(BaseMatmul):
+ """Grouped matmul with swiglu and requantisation fused into one kernel.
+
+ Used for the gate/up projection (gmm1) of block-scaled MoE: the kernel emits
+ activations already quantised for the following down projection, so the
+ caller has no separate activation step. Unlike ``GroupedMatmul`` it returns
+ ``(quantized_activations, block_scale)`` instead of a single tensor, and it
+ takes no ``output_dtype`` — the output dtype is set through ``quant_dtype``
+ in ``scale_args``.
+ """
+
+ def forward(
+ self,
+ layer: torch.nn.Module,
+ weight_prefix: str,
+ hidden_states: torch.Tensor,
+ expert_tokens: torch.Tensor,
+ output_dtype: torch.dtype = None,
+ group_list_type: int = 1,
+ transposed: bool = True,
+ **scale_args,
+ ) -> Tuple[torch.Tensor, torch.Tensor]:
+ weight = getattr(layer, f"{weight_prefix}_weight", None)
+ if weight is None:
+ raise AttributeError(
+ f"Weight attribute '{weight_prefix}_weight' not found in layer"
+ )
+ # This op wants a cumulative group_list while the plain grouped matmul
+ # keeps the COUNT form the dispatcher produces (group_list_type=1). The
+ # asymmetry is intentional.
+ group_list = expert_tokens.cumsum(0) if group_list_type == 1 else expert_tokens
+ return torch.ops.npu.npu_grouped_matmul_swiglu_quant_v2(
+ x=hidden_states,
+ weight=[weight] if transposed else [weight.transpose(1, 2)],
+ group_list=group_list,
+ **scale_args,
+ )
diff --git a/python/sglang/srt/hardware_backend/npu/moe/hidden_states_quant.py b/python/sglang/srt/hardware_backend/npu/moe/quant.py
similarity index 72%
rename from python/sglang/srt/hardware_backend/npu/moe/hidden_states_quant.py
rename to python/sglang/srt/hardware_backend/npu/moe/quant.py
index 99a5c5ed8..afb0bb6b9 100644
--- a/python/sglang/srt/hardware_backend/npu/moe/hidden_states_quant.py
+++ b/python/sglang/srt/hardware_backend/npu/moe/quant.py
@@ -1,5 +1,5 @@
"""
-Hidden state quantization utilities for NPU MoE.
+Quantization kernel wrappers for NPU MoE.
Each class quantises hidden states and returns a (quantized_tensor, scale) tuple.
For static quantization the scale is ``None``.
@@ -27,15 +27,26 @@ class HiddenStatesDynamicQuant(BaseHiddenStatesQuant):
"""
Dynamic per‑token quantisation of hidden states.
+ ``torch.float8_e4m3fn`` selects the MX (block-scaled) op, whose scale is a
+ ``float8_e8m0fnu`` block scale ``[N, K//64, 2]`` rather than one scalar per
+ token; the int8/int4 dtypes keep the plain per-token op.
+
Returns ``(quantized_hidden_states, per‑token_scale)``.
"""
+ def __init__(self, quant_dtype: torch.dtype) -> None:
+ super().__init__(quant_dtype)
+ if quant_dtype == torch.float8_e4m3fn:
+ self._op = torch.ops.npu.npu_dynamic_mx_quant
+ elif quant_dtype in (torch.int8, torch.quint4x2):
+ self._op = torch.ops.npu.npu_dynamic_quant
+ else:
+ raise ValueError(f"Unsupported dynamic quant dtype: {quant_dtype}")
+
def __call__(
self, hidden_states: torch.Tensor
) -> Tuple[torch.Tensor, torch.Tensor]:
- quantized, scale = torch.ops.npu.npu_dynamic_quant(
- hidden_states, dst_type=self.quant_dtype
- )
+ quantized, scale = self._op(hidden_states, dst_type=self.quant_dtype)
return quantized, scale
diff --git a/python/sglang/srt/hardware_backend/npu/quantization/moe_methods.py b/python/sglang/srt/hardware_backend/npu/quantization/moe_methods.py
index f49d53534..33fa8e54f 100644
--- a/python/sglang/srt/hardware_backend/npu/quantization/moe_methods.py
+++ b/python/sglang/srt/hardware_backend/npu/quantization/moe_methods.py
@@ -1,7 +1,8 @@
-from typing import TYPE_CHECKING, Any, Dict, Optional
+from typing import TYPE_CHECKING, Any, Dict, Optional, Tuple
import numpy as np
import torch
+from torch.nn.parameter import Parameter
from sglang.srt.environ import envs
from sglang.srt.hardware_backend.npu.utils import npu_format_cast
@@ -13,14 +14,56 @@ if TYPE_CHECKING:
import logging
-from sglang.srt.hardware_backend.npu.moe.hidden_states_quant import (
- HiddenStatesDynamicQuant,
+from sglang.srt.hardware_backend.npu.moe.matmul import (
+ GroupedMatmul,
+ GroupedMatmulSwigluQuant,
+)
+from sglang.srt.hardware_backend.npu.moe.quant import HiddenStatesDynamicQuant
+from sglang.srt.hardware_backend.npu.quantization.linear_method_npu import (
+ _get_float8_e8m0fnu_dtype,
)
-from sglang.srt.hardware_backend.npu.moe.matmul import GroupedMatmul
logger = logging.getLogger(__name__)
+_E8M0_DTYPE = None
+
+
+def _require_e8m0_dtype():
+ """Resolve the e8m0 block-scale dtype, failing loudly if it is unavailable.
+
+ The grouped matmuls validate their scale-dtype arguments against torch_npu's
+ own dtype enum (``torch_npu.float8_e8m0fnu``, 293 on A5) and reject the torch
+ dtype object with "weight_scale_dtype only supports float8_e8m0fnu or None,
+ but the actual value is Float8_e8m0fnu" — hence torch_npu first, torch only
+ as a fallback. Dense ``npu_quant_matmul`` accepts either, which is why
+ ``_get_float8_e8m0fnu_dtype`` reads it off torch.
+
+ The MXFP8 ops take the scale dtype explicitly; passing None silently changes
+ how they interpret the scales, so a missing dtype must raise rather than
+ propagate.
+
+ torch_npu is imported lazily (and cached) so this module stays importable on
+ CUDA/CPU/AMD/XPU CI.
+ """
+ global _E8M0_DTYPE
+ if _E8M0_DTYPE is None:
+ from sglang.srt.utils import is_npu
+
+ if is_npu():
+ import torch_npu
+
+ _E8M0_DTYPE = getattr(torch_npu, "float8_e8m0fnu", None)
+ if _E8M0_DTYPE is None:
+ _E8M0_DTYPE = _get_float8_e8m0fnu_dtype()
+ if _E8M0_DTYPE is None:
+ raise RuntimeError(
+ "float8_e8m0fnu dtype not found — MXFP8 MoE requires Ascend A5 "
+ "with a torch_npu build exposing float8_e8m0fnu (torch_npu >= 2.9)."
+ )
+ return _E8M0_DTYPE
+
+
# DEPRECATED METHOD
# TODO: Remove in future realeses
def fused_moe_npu(
@@ -714,3 +757,196 @@ class NPUUnquantMoEMethod(_NPUMoEMethodBase):
transposed=False,
**self._get_bias_args(quant_info, weight_prefix),
)
+
+
+# ---------------------------------------------------------------------------
+# NPUMXFP8MoEMethod
+# ---------------------------------------------------------------------------
+class NPUMXFP8MoEMethod(_NPUMoEMethodBase):
+ """MXFP8 MoE on Ascend A5 – float8_e4m3fn weights with e8m0 block scales.
+
+ Serves both the online config path (``--quantization mxfp8``, weights
+ quantised at load time) and the offline ModelSlim ``W8A8_MXFP8`` scheme
+ (weights already quantised); ``process_weights_after_loading`` tells the two
+ apart by weight dtype.
+
+ gmm1 re-quantises its own output, so it is a single fused kernel rather than
+ a matmul plus a separate activation: the runner calls
+ ``apply_fused_gmm1_swiglu`` for w13 and ``apply`` only for w2 — hence the
+ per-prefix matmul chosen here.
+
+ Where the *activation* quant happens depends on the dispatcher. On
+ ``ascend_tp`` it comes for free from ``npu_moe_init_routing_v2(quant_mode=3)``,
+ which emits the e4m3 payload and e8m0 scale as part of routing. DeepEP has no
+ mxfp8 dispatch dtype, so it hands over bf16 and w13 quantises the hidden
+ states itself before gmm1.
+ """
+
+ def __init__(self, weight_prefix: str):
+ super().__init__(quant_config=None)
+ if weight_prefix == "w13":
+ self.matmul = GroupedMatmulSwigluQuant()
+ self.hidden_states_quantizer = HiddenStatesDynamicQuant(
+ quant_dtype=torch.float8_e4m3fn
+ )
+ else:
+ self.matmul = GroupedMatmul()
+ self.hidden_states_quantizer = None
+
+ @staticmethod
+ def _quantize_weight_online(
+ weight: torch.Tensor, weight_prefix: str
+ ) -> Tuple[torch.Tensor, torch.Tensor]:
+ """Quantise BF16/FP16 expert weights ``[E, N, K]`` to MXFP8 at load time.
+
+ Returns the e4m3 payload ``[E, N, K]`` and its e8m0 block scale
+ ``[E, N, K//64, 2]`` (block_size=32, already pair-split by the op).
+ """
+ if weight.dtype not in (torch.float16, torch.bfloat16):
+ logger.warning(
+ "NPUMXFP8MoEMethod: %s_weight dtype %s is not float16/bfloat16; "
+ "casting to bfloat16 before MXFP8 quantisation.",
+ weight_prefix,
+ weight.dtype,
+ )
+ weight = weight.to(torch.bfloat16)
+ # cpu offload may have moved the weight back to host memory.
+ if not weight.is_npu:
+ weight = weight.to(f"npu:{torch.npu.current_device()}")
+ return torch.ops.npu.npu_dynamic_mx_quant(weight, dst_type=torch.float8_e4m3fn)
+
+ def process_weights_after_loading(
+ self, layer: torch.nn.Module, weight_prefix: str
+ ) -> None:
+ self._validate_weight_prefix(layer, weight_prefix)
+
+ weight: torch.Tensor = getattr(layer, f"{weight_prefix}_weight").data
+ if weight.dtype == torch.float8_e4m3fn:
+ # Offline (ModelSlim) path: the checkpoint already holds e4m3 weights
+ # and {prefix}_weight_scale holds uint8 block scales [E, N, K//32].
+ # Only re-layout: split the flat scale axis into pairs to match what
+ # npu_dynamic_mx_quant produces online.
+ scale: torch.Tensor = getattr(layer, f"{weight_prefix}_weight_scale").data
+ scale = scale.reshape(scale.shape[0], -1, scale.shape[-1] // 2, 2)
+ else:
+ weight, scale = self._quantize_weight_online(weight, weight_prefix)
+
+ # FRACTAL_NZ before the transpose, never after. gmm1 asserts that weight
+ # and weight_scale carry the SAME transpose flag (CheckMXTranspose: "the
+ # transposition of weightScale/weight should be equal"), and the cast
+ # yields a physically retiled — hence non-transposed — tensor. Casting
+ # the [E, K, N] view would therefore leave the weight at false against a
+ # true scale and fail outright, which is why this cannot copy the int8
+ # MoE methods above (they transpose first, but carry no MX scale to keep
+ # in sync). Same order as the dense W4A8 path in linear_method_npu.py.
+ #
+ # A5 measurement, Qwen3-30B-A3B shapes, 128 experts (see
+ # llm/probe_mxfp8_moe_nz.py): +1.4% decode, +3.8% prefill against a 0.2-
+ # 0.3% noise floor, bit-identical outputs. Set
+ # SGLANG_NPU_DISABLE_ACL_FORMAT_WEIGHT to fall back to plain ND.
+ weight = npu_format_cast(weight)
+
+ # Both paths hand the grouped matmul weight [E, K, N] and scale
+ # [E, K//64, N, 2] as strided transpose views — DO NOT call
+ # .contiguous(). Beyond breaking the transpose-flag match above, it
+ # measures slower on the same probe: making both sides contiguous costs
+ # 6.2% on decode. This matches NPUMXFP8LinearMethod, msmodelslim's
+ # offline layout and vllm-ascend's AscendW8A8MXFP8DynamicFusedMoEMethod.
+ setattr(
+ layer,
+ f"{weight_prefix}_weight",
+ Parameter(weight.transpose(1, 2), requires_grad=False),
+ )
+ setattr(
+ layer,
+ f"{weight_prefix}_weight_scale",
+ Parameter(scale.transpose(1, 2), requires_grad=False),
+ )
+
+ if weight_prefix == "w13":
+ from sglang.srt.layers.moe import get_moe_a2a_backend
+
+ # DeepEP has no mxfp8 entry in its dispatch dtype table, so let it
+ # keep sending bf16; apply_fused_gmm1_swiglu quantises instead.
+ dispatcher_dtype = "bf16" if get_moe_a2a_backend().is_deepep() else "mxfp8"
+ self._set_dispatcher_output_dtype(layer, dispatcher_dtype)
+
+ def apply_fused_gmm1_swiglu(
+ self,
+ quant_info: "AscendQuantInfo",
+ hidden_states: torch.Tensor,
+ expert_tokens: torch.Tensor,
+ pertoken_scale: Optional[torch.Tensor],
+ group_list_type,
+ ) -> Tuple[torch.Tensor, torch.Tensor]:
+ """Gate/up projection, swiglu and requantisation in one kernel (gmm1).
+
+ Returns the e4m3 activations and their e8m0 block scale, i.e. exactly
+ what the w2 gmm needs, which is why the runner skips its activation step
+ for MXFP8.
+
+ ``pertoken_scale`` is None when the dispatcher handed over unquantised
+ hidden states (the DeepEP path), in which case the activation quant that
+ ascend_tp fuses into routing is done here instead. Both dispatchers
+ therefore reach the kernel below with the same e4m3 + e8m0 input.
+ """
+ if pertoken_scale is None:
+ hidden_states, pertoken_scale = self.hidden_states_quantizer(hidden_states)
+
+ e8m0_dtype = _require_e8m0_dtype()
+ return self.matmul.forward(
+ quant_info,
+ "w13",
+ hidden_states,
+ expert_tokens,
+ group_list_type=group_list_type,
+ transposed=True,
+ weight_scale=[quant_info.w13_weight_scale],
+ x_scale=pertoken_scale,
+ dequant_mode=2,
+ quant_mode=2,
+ dequant_dtype=torch.float32,
+ quant_dtype=torch.float8_e4m3fn,
+ # e4m3 is implicit for these two — it is not in the op's QUANT_DTYPES.
+ x_dtype=None,
+ weight_dtype=None,
+ weight_scale_dtype=e8m0_dtype,
+ x_scale_dtype=e8m0_dtype,
+ )
+
+ def apply(
+ self,
+ quant_info: "AscendQuantInfo",
+ hidden_states: torch.Tensor,
+ expert_tokens: torch.Tensor,
+ pertoken_scale: torch.Tensor,
+ output_dtype: torch.dtype,
+ weight_prefix: str,
+ group_list_type,
+ ) -> torch.Tensor:
+ if weight_prefix != "w2":
+ raise ValueError(
+ f"NPUMXFP8MoEMethod.apply only serves the w2 gmm, got "
+ f"'{weight_prefix}'. gmm1 fuses swiglu into a single op and must "
+ f"go through apply_fused_gmm1_swiglu, which returns a scale too."
+ )
+
+ e8m0_dtype = _require_e8m0_dtype()
+ scale_args: Dict[str, Any] = {
+ "scale": [getattr(quant_info, f"{weight_prefix}_weight_scale", None)],
+ "per_token_scale": [pertoken_scale],
+ "scale_dtype": e8m0_dtype,
+ "per_token_scale_dtype": e8m0_dtype,
+ "x_dtype": None,
+ "weight_dtype": None,
+ }
+ return self.matmul.forward(
+ quant_info,
+ weight_prefix,
+ hidden_states,
+ expert_tokens,
+ output_dtype,
+ group_list_type=group_list_type,
+ transposed=True,
+ **scale_args,
+ )
diff --git a/python/sglang/srt/hardware_backend/npu/quantization/online_moe_methods.py b/python/sglang/srt/hardware_backend/npu/quantization/online_moe_methods.py
new file mode 100644
index 000000000..77da812ee
--- /dev/null
+++ b/python/sglang/srt/hardware_backend/npu/quantization/online_moe_methods.py
@@ -0,0 +1,65 @@
+"""Online (config-driven) quantized FusedMoE methods for Ascend NPU.
+
+These are the ``--quantization `` entry points: the checkpoint holds
+BF16/FP16 expert weights and the per-gmm kernels quantize them at load time.
+Offline (msmodelslim) checkpoints go through the ModelSlim schemes instead and
+reuse the same kernels.
+
+Kept out of ``moe_methods.py`` because ``unquant.py`` imports that module at
+module scope, so subclassing ``UnquantizedFusedMoEMethod`` there would be a
+circular import.
+"""
+
+from typing import TYPE_CHECKING, Optional
+
+import torch
+
+from sglang.srt.hardware_backend.npu.quantization.moe_methods import NPUMXFP8MoEMethod
+from sglang.srt.layers.moe.moe_runner import MoeRunner
+from sglang.srt.layers.moe.utils import MoeRunnerBackend, get_moe_runner_backend
+from sglang.srt.layers.quantization.unquant import UnquantizedFusedMoEMethod
+
+if TYPE_CHECKING:
+ from sglang.srt.layers.moe.moe_runner.base import MoeRunnerConfig
+ from sglang.srt.layers.quantization.base_config import QuantizationConfig
+
+
+class NPUMXFP8OnlineMoEMethod(UnquantizedFusedMoEMethod):
+ """Online MXFP8 FusedMoE entry point (``--quantization mxfp8`` on A5).
+
+ Weight creation, weight post-processing and the forward pass are identical
+ to the unquantized Ascend path — the only difference is which per-gmm kernel
+ the layer gets, so everything but ``create_moe_runner`` is inherited.
+ ``NPUMXFP8MoEMethod`` then quantizes the BF16 expert weights to MXFP8 in
+ ``process_weights_after_loading``.
+ """
+
+ def __init__(self, quant_config: Optional["QuantizationConfig"] = None):
+ super().__init__()
+ self.quant_config = quant_config
+
+ def create_moe_runner(
+ self, layer: torch.nn.Module, moe_runner_config: "MoeRunnerConfig"
+ ):
+ backend = get_moe_runner_backend()
+ if not (backend.is_auto() or backend.is_ascend()):
+ # Not merely a wrong-runner check. Because this method subclasses
+ # UnquantizedFusedMoEMethod it matches FusedMoE's shard-swap list, so
+ # a flashinfer backend would make the weight loader exchange the
+ # w1/w3 shards ("flashinfer assumes w31 format"). Every expert would
+ # then load with gate and up swapped, and gmm1's fused swiglu would
+ # compute silu(up) * gate — no error, just degenerate output.
+ raise ValueError(
+ "MXFP8 MoE on Ascend requires --moe-runner-backend 'auto' or "
+ f"'ascend', got {backend.value!r}."
+ )
+
+ # The kernels must be attached before the runner is built:
+ # AscendRunnerCore.__init__ reads layer.w2_kernel to pick its activation.
+ layer.w13_kernel = NPUMXFP8MoEMethod("w13")
+ layer.w2_kernel = NPUMXFP8MoEMethod("w2")
+ moe_runner_config.layer = layer
+ self.moe_runner_config = moe_runner_config
+ self.runner = MoeRunner(MoeRunnerBackend.ASCEND, moe_runner_config)
+ # Inherited apply() consults this; aiter is CUDA/ROCm-only.
+ self._aiter_runner = None
diff --git a/python/sglang/srt/hardware_backend/npu/utils.py b/python/sglang/srt/hardware_backend/npu/utils.py
index 54389f39a..664bb2881 100644
--- a/python/sglang/srt/hardware_backend/npu/utils.py
+++ b/python/sglang/srt/hardware_backend/npu/utils.py
@@ -113,18 +113,25 @@ def init_npu_backend():
def _is_nz_aligned(tensor: torch.Tensor) -> bool:
"""Check whether the last two dims satisfy FRACTAL_NZ alignment rules.
- Ascend FRACTAL_NZ requires:
- BF16 / FP16 : both dims divisible by 16
- INT8 : k % 16 == 0 and n % 32 == 0
+ A fractal tile is 16 rows by 32 bytes (the C0_32 in the op's error strings),
+ so the row rule is always k % 16 and the column rule is 32 // itemsize:
+
+ BF16 / FP16 : k % 16 == 0 and n % 16 == 0
+ INT8 / FP8 : k % 16 == 0 and n % 32 == 0
INT4 : k % 16 == 0 and n % 64 == 0
- FP4 : both dims divisible by 64
+
+ Unlisted dtypes fall through to True: this is a cheap pre-filter for known
+ bad combinations, not an authority — the op itself is.
"""
if tensor.dim() < 2:
return False
k, n = tensor.shape[-2], tensor.shape[-1]
if tensor.dtype in (torch.bfloat16, torch.float16):
return k % 16 == 0 and n % 16 == 0
- if tensor.dtype == torch.int8:
+ if tensor.dtype in (torch.int8, torch.float8_e4m3fn):
+ # e4m3 is single-byte like int8, so it shares the column rule. Reached
+ # only by the MXFP8 MoE weights; the packed-FP4 callers pass
+ # customize_dtype and return before this check.
return k % 16 == 0 and n % 32 == 0
if tensor.dtype in (torch.uint8, torch.int32):
# INT4 is typically packed into uint8/int32; be conservative
diff --git a/python/sglang/srt/layers/moe/moe_runner/ascend.py b/python/sglang/srt/layers/moe/moe_runner/ascend.py
index 7b15281c3..6dfee6448 100644
--- a/python/sglang/srt/layers/moe/moe_runner/ascend.py
+++ b/python/sglang/srt/layers/moe/moe_runner/ascend.py
@@ -17,6 +17,7 @@ from sglang.srt.hardware_backend.npu.moe.activation import (
NPUSwigluStepAndMul,
)
from sglang.srt.hardware_backend.npu.quantization.moe_methods import (
+ NPUMXFP8MoEMethod,
NPUW4A8Int8MoEMethod,
NPUW8A8Int8MoEMethod,
)
@@ -87,7 +88,15 @@ class AscendRunnerCore(MoeRunnerCore):
kernel = config.layer.w2_kernel
- if get_moe_a2a_backend().is_deepep():
+ if isinstance(kernel, NPUMXFP8MoEMethod):
+ # MXFP8 fuses gate/up + swiglu + requant into gmm1, so there is no
+ # separate activation step — run() skips it. Left None on purpose so
+ # that reaching for it fails loudly instead of silently applying an
+ # unfused swiglu to already-requantised activations. This holds for
+ # both dispatchers: ascend_tp gets its activation quant fused into
+ # routing, DeepEP dispatches bf16 and gmm1 quantises it itself.
+ self.activation = None
+ elif get_moe_a2a_backend().is_deepep():
# DeepEP path: use a unified kernel that decides quantisation
is_quant_kernel = isinstance(
kernel, (NPUW4A8Int8MoEMethod, NPUW8A8Int8MoEMethod)
@@ -134,30 +143,44 @@ class AscendRunnerCore(MoeRunnerCore):
expert_tokens = runner_input.expert_tokens
group_list_type = runner_input.group_list_type
- # --- w13 (gate & up) projection ---
- hidden_states = self.config.layer.w13_kernel.apply(
- quant_info,
- x,
- expert_tokens,
- pertoken_scale=runner_input.hidden_states_scale,
- output_dtype=original_dtype,
- weight_prefix="w13",
- group_list_type=group_list_type,
- )
+ w13_kernel = self.config.layer.w13_kernel
- # --- Activation ---
- # The DeepEP kernel expects extra dispatch metadata
- if isinstance(self.activation, NPUSwigluDeepEPKernel):
- hidden_states, pertoken_scale = self.activation._apply_activation(
- hidden_states,
- group_list=expert_tokens,
+ if isinstance(w13_kernel, NPUMXFP8MoEMethod):
+ # --- w13 projection + activation, fused into one kernel ---
+ # MXFP8 gmm1 returns activations already requantised for gmm2, so
+ # there is no separate activation step to run.
+ hidden_states, pertoken_scale = w13_kernel.apply_fused_gmm1_swiglu(
+ quant_info,
+ x,
+ expert_tokens,
+ pertoken_scale=runner_input.hidden_states_scale,
group_list_type=group_list_type,
)
else:
- hidden_states, pertoken_scale = self.activation._apply_activation(
- hidden_states
+ # --- w13 (gate & up) projection ---
+ hidden_states = w13_kernel.apply(
+ quant_info,
+ x,
+ expert_tokens,
+ pertoken_scale=runner_input.hidden_states_scale,
+ output_dtype=original_dtype,
+ weight_prefix="w13",
+ group_list_type=group_list_type,
)
+ # --- Activation ---
+ # The DeepEP kernel expects extra dispatch metadata
+ if isinstance(self.activation, NPUSwigluDeepEPKernel):
+ hidden_states, pertoken_scale = self.activation._apply_activation(
+ hidden_states,
+ group_list=expert_tokens,
+ group_list_type=group_list_type,
+ )
+ else:
+ hidden_states, pertoken_scale = self.activation._apply_activation(
+ hidden_states
+ )
+
# --- w2 (down) projection ---
hidden_states = self.config.layer.w2_kernel.apply(
quant_info,
diff --git a/python/sglang/srt/layers/moe/token_dispatcher/ascend_tp.py b/python/sglang/srt/layers/moe/token_dispatcher/ascend_tp.py
index 96921d995..b8303a4fa 100644
--- a/python/sglang/srt/layers/moe/token_dispatcher/ascend_tp.py
+++ b/python/sglang/srt/layers/moe/token_dispatcher/ascend_tp.py
@@ -9,6 +9,7 @@ from sglang.srt.hardware_backend.npu.moe.finalize_routing import (
NPUFinalizeRouting,
)
from sglang.srt.hardware_backend.npu.moe.init_routing import (
+ MXFP8_QUANT_MODE,
NPUMoEInitRouting_v2,
)
from sglang.srt.layers.moe.moe_runner.base import MoeRunnerConfig
@@ -84,6 +85,10 @@ class AscendTPDispatcher(BaseDispatcher):
self.init = NPUMoEInitRouting_v2(quant_mode=1)
self.finalize = NPUFinalizeRouting(drop_pad_mode=2)
self.group_list_type = 1
+ elif self.ascend_dispatcher_output_dtype == DispatcherOutputDtype.MXFP8:
+ self.init = NPUMoEInitRouting_v2(quant_mode=MXFP8_QUANT_MODE)
+ self.finalize = NPUFinalizeRouting(drop_pad_mode=2)
+ self.group_list_type = 1
else:
raise ValueError(
f"Unsupported ascend_dispatcher_output_dtype: {self.ascend_dispatcher_output_dtype}"
diff --git a/python/sglang/srt/layers/moe/utils.py b/python/sglang/srt/layers/moe/utils.py
index 0f3f48f1e..3778d18e2 100644
--- a/python/sglang/srt/layers/moe/utils.py
+++ b/python/sglang/srt/layers/moe/utils.py
@@ -209,12 +209,14 @@ class DispatcherOutputDtype(Enum):
- FP8: dispatch hidden states in fp8
- INT8: dispatch hidden states in int8
- NVFP4: dispatch hidden states in nvfp4
+ - MXFP8: dispatch hidden states in mxfp8 (fp8_e4m3 + e8m0 block scale)
"""
BF16 = "bf16"
FP8 = "fp8"
INT8 = "int8"
NVFP4 = "nvfp4"
+ MXFP8 = "mxfp8"
def get_deepep_output_dtype(self) -> DispatcherOutputDtype:
diff --git a/python/sglang/srt/layers/quantization/fp8.py b/python/sglang/srt/layers/quantization/fp8.py
index 9581d93dc..861cc9432 100644
--- a/python/sglang/srt/layers/quantization/fp8.py
+++ b/python/sglang/srt/layers/quantization/fp8.py
@@ -360,6 +360,13 @@ class Fp8Config(QuantizationConfig):
layer.use_triton_kernels, layer.use_flashinfer_trtllm_moe
)
+ if is_npu() and self.use_mxfp8:
+ from sglang.srt.hardware_backend.npu.quantization.online_moe_methods import (
+ NPUMXFP8OnlineMoEMethod,
+ )
+
+ return NPUMXFP8OnlineMoEMethod(self)
+
fp8_method = Fp8MoEMethod(self)
if self.is_fp4_experts and self.dequant_fp4_to_fp8:
diff --git a/python/sglang/srt/layers/quantization/modelslim/modelslim.py b/python/sglang/srt/layers/quantization/modelslim/modelslim.py
index e02cdb667..227e20010 100644
--- a/python/sglang/srt/layers/quantization/modelslim/modelslim.py
+++ b/python/sglang/srt/layers/quantization/modelslim/modelslim.py
@@ -18,6 +18,7 @@ from sglang.srt.layers.quantization.base_config import (
from sglang.srt.layers.quantization.modelslim.schemes import (
ModelSlimMXFP4Scheme,
ModelSlimMXFP4W4A8Scheme,
+ ModelSlimMXFP8MoEScheme,
ModelSlimMXFP8Scheme,
ModelSlimW4A4Int4,
ModelSlimW4A4Int4MoE,
@@ -238,6 +239,7 @@ class ModelSlimConfig(QuantizationConfig):
("W4A4_DYNAMIC", ModelSlimW4A4Int4MoE),
("W4A8_DYNAMIC", ModelSlimW4A8Int8MoE),
("W8A8_DYNAMIC", ModelSlimW8A8Int8MoE),
+ ("W8A8_MXFP8", ModelSlimMXFP8MoEScheme),
]
# Try multiple naming conventions:
diff --git a/python/sglang/srt/layers/quantization/modelslim/schemes/__init__.py b/python/sglang/srt/layers/quantization/modelslim/schemes/__init__.py
index 5d78a2d43..ca03b17cd 100644
--- a/python/sglang/srt/layers/quantization/modelslim/schemes/__init__.py
+++ b/python/sglang/srt/layers/quantization/modelslim/schemes/__init__.py
@@ -10,6 +10,7 @@ from .modelslim_mxfp4_w4a8 import ModelSlimMXFP4W4A8Scheme
from .modelslim_mxfp4 import ModelSlimMXFP4Scheme
# isort: on
+from .modelslim_mxfp8_moe import ModelSlimMXFP8MoEScheme
from .modelslim_w4a4_int4 import ModelSlimW4A4Int4
from .modelslim_w4a4_int4_moe import ModelSlimW4A4Int4MoE
from .modelslim_w4a8_int8_moe import ModelSlimW4A8Int8MoE
@@ -22,6 +23,7 @@ __all__ = [
"ModelSlimMXFP8Scheme",
"ModelSlimMXFP4W4A8Scheme",
"ModelSlimMXFP4Scheme",
+ "ModelSlimMXFP8MoEScheme",
"ModelSlimW8A8Int8",
"ModelSlimW4A4Int4",
"ModelSlimW4A4Int4MoE",
diff --git a/python/sglang/srt/layers/quantization/modelslim/schemes/modelslim_mxfp8_moe.py b/python/sglang/srt/layers/quantization/modelslim/schemes/modelslim_mxfp8_moe.py
new file mode 100644
index 000000000..d49423e36
--- /dev/null
+++ b/python/sglang/srt/layers/quantization/modelslim/schemes/modelslim_mxfp8_moe.py
@@ -0,0 +1,114 @@
+"""ModelSlim MXFP8 offline scheme for MoE layers on Ascend NPU (SRT).
+
+Loads weights pre-quantised by msmodelslim: float8_e4m3fn weights + uint8 block
+scales (block_size=32). The layout transform and the forward pass are delegated
+to ``NPUMXFP8MoEMethod`` -- the same kernel the online MXFP8 MoE path uses.
+"""
+
+from __future__ import annotations
+
+import logging
+from typing import Any, Dict
+
+import torch
+
+from sglang.srt.hardware_backend.npu.quantization.moe_methods import NPUMXFP8MoEMethod
+from sglang.srt.layers.quantization.modelslim.schemes import ModelSlimMoEScheme
+from sglang.srt.utils import set_weight_attrs
+
+logger = logging.getLogger(__name__)
+
+__all__ = [
+ "ModelSlimMXFP8MoEScheme",
+]
+
+# Block (group) size of the msmodelslim MXFP8 export format.
+MXFP8_BLOCK_SIZE = 32
+
+
+class ModelSlimMXFP8MoEScheme(ModelSlimMoEScheme):
+ """
+ Offline MXFP8 MoE scheme that creates weights for either the
+ w13 (gate+up) or w2 (down) projection group.
+
+ Two instances of this class are used per MoE layer:
+ - weight_prefix="w13" → handles the fused gate_proj + up_proj weights
+ - weight_prefix="w2" → handles the down_proj weights
+
+ The float8_e4m3fn weight dtype allocated here is what tells
+ ``NPUMXFP8MoEMethod.process_weights_after_loading`` to take its offline
+ (re-layout only) branch instead of quantising the weights itself.
+ """
+
+ def __init__(
+ self,
+ quant_config: Dict[str, Any],
+ weight_prefix: str, # "w13" or "w2"
+ ) -> None:
+ if weight_prefix not in ("w13", "w2"):
+ raise ValueError(
+ f"weight_prefix must be 'w13' or 'w2', got '{weight_prefix}'"
+ )
+ self.quant_config = quant_config
+ self.weight_prefix = weight_prefix
+ self.kernel = NPUMXFP8MoEMethod(weight_prefix)
+
+ def create_weights(
+ self,
+ layer: torch.nn.Module,
+ num_experts: int,
+ hidden_size: int,
+ intermediate_size_per_partition: int,
+ **extra_weight_attrs,
+ ) -> None:
+ from sglang.srt.layers.moe.fused_moe_triton import FusedMoeWeightScaleSupported
+
+ self.num_experts = num_experts
+ extra_weight_attrs.update(
+ {"quant_method": FusedMoeWeightScaleSupported.BLOCK.value}
+ )
+
+ # Determine shape based on weight group
+ if self.weight_prefix == "w13":
+ a_dim = 2 * intermediate_size_per_partition
+ b_dim = hidden_size
+ else: # w2
+ a_dim = hidden_size
+ b_dim = intermediate_size_per_partition
+
+ prefix = self.weight_prefix
+
+ # Create and register weight: [E, N, K] float8_e4m3fn
+ weight_name = f"{prefix}_weight"
+ weight = torch.nn.Parameter(
+ torch.empty(num_experts, a_dim, b_dim, dtype=torch.float8_e4m3fn),
+ requires_grad=False,
+ )
+ layer.register_parameter(weight_name, weight)
+ set_weight_attrs(weight, extra_weight_attrs)
+
+ # Create and register block scale: [E, N, K // 32] uint8 (e8m0)
+ scale_name = f"{prefix}_weight_scale"
+ scale = torch.nn.Parameter(
+ torch.empty(
+ num_experts, a_dim, b_dim // MXFP8_BLOCK_SIZE, dtype=torch.uint8
+ ),
+ requires_grad=False,
+ )
+ layer.register_parameter(scale_name, scale)
+ set_weight_attrs(scale, extra_weight_attrs)
+
+ # MXFP8 is a pure scale format: the e8m0 block exponent above carries
+ # everything, there is no zero point. The int8/int4 schemes do have one,
+ # so ModelSlimMoEMethod.apply reads layer.{w13,w2}_weight_offset
+ # unconditionally when it builds AscendQuantInfo (where the field is
+ # Optional). Register it as None so the attribute exists and resolves to
+ # "no offset" rather than raising AttributeError. A None parameter is
+ # skipped by named_parameters(), so no weight loader looks for it.
+ layer.register_parameter(f"{prefix}_weight_offset", None)
+
+ def process_weights_after_loading(self, layer: torch.nn.Module) -> None:
+ """
+ Delegate weight processing to the NPU kernel for the fixed weight group.
+ """
+ self.kernel.process_weights_after_loading(layer, self.weight_prefix)
diff --git a/python/sglang/srt/layers/quantization/unquant.py b/python/sglang/srt/layers/quantization/unquant.py
index 725f4d6bd..1834b038c 100644
--- a/python/sglang/srt/layers/quantization/unquant.py
+++ b/python/sglang/srt/layers/quantization/unquant.py
@@ -438,10 +438,12 @@ class UnquantizedFusedMoEMethod(FusedMoEMethodBase, MultiPlatformOp):
layer.num_local_experts, *new_shape_w2
)
if _is_npu:
+ # The kernels set the dispatcher output dtype themselves -- they are
+ # the ones that know what their gmms expect. NPUUnquantMoEMethod
+ # already sets bf16 here, and hardcoding it a second time would
+ # clobber a subclass that attached a quantized kernel instead.
layer.w13_kernel.process_weights_after_loading(layer, "w13")
layer.w2_kernel.process_weights_after_loading(layer, "w2")
- if hasattr(layer, "dispatcher"):
- layer.dispatcher.set_quant_config({"dispatcher_output_dtype": "bf16"})
return
diff --git a/python/sglang/srt/models/qwen3_moe.py b/python/sglang/srt/models/qwen3_moe.py
index bb8a4d660..9524a7ba7 100644
--- a/python/sglang/srt/models/qwen3_moe.py
+++ b/python/sglang/srt/models/qwen3_moe.py
@@ -271,11 +271,21 @@ class Qwen3MoeSparseMoeBlock(nn.Module):
routing_method_type=RoutingMethodType.Renormalize,
)
+ # Router gate: description-driven quant, mirroring vllm-ascend. Only the
+ # offline ModelSlim path (which carries a per-layer quant_model_description)
+ # may quantise the gate — if the checkpoint stored it as MXFP8 it is loaded
+ # and dequantised correctly instead of cast to bf16 without its block scale.
+ # The online Fp8/mxfp8 path keeps the gate in bf16 (unchanged, verified).
+ gate_quant_config = (
+ quant_config
+ if (quant_config is not None and quant_config.get_name() == "modelslim")
+ else None
+ )
self.gate = ReplicatedLinear(
config.hidden_size,
config.num_experts,
bias=False,
- quant_config=None,
+ quant_config=gate_quant_config,
prefix=add_prefix("gate", prefix),
)