From 3abdbab9bbb4ef960be30809a62affd6c6d7a38f Mon Sep 17 00:00:00 2001 From: Junlin Wu Date: Tue, 7 Jul 2026 00:23:26 +0800 Subject: [PATCH] :sparkles: [llm][npu][quant] Add W4A8 MXFP quantization support for Qwen3 Dense on Ascend NPU (#23650) Co-authored-by: ronnie_zheng --- .../docs/advanced_features/quantization.mdx | 7 + .../ascend-npus/ascend_npu_quantization.mdx | 30 ++ .../npu/quantization/linear_method_npu.py | 312 +++++++++++++++++- .../sglang/srt/hardware_backend/npu/utils.py | 28 +- .../srt/layers/quantization/__init__.py | 2 + .../quantization/modelslim/modelslim.py | 2 + .../modelslim/schemes/__init__.py | 2 + .../modelslim/schemes/modelslim_mxfp4_w4a8.py | 107 ++++++ .../srt/layers/quantization/npu_mxfp4.py | 127 +++++++ python/sglang/srt/server_args.py | 1 + 10 files changed, 613 insertions(+), 5 deletions(-) create mode 100644 python/sglang/srt/layers/quantization/modelslim/schemes/modelslim_mxfp4_w4a8.py create mode 100644 python/sglang/srt/layers/quantization/npu_mxfp4.py diff --git a/docs_new/docs/advanced_features/quantization.mdx b/docs_new/docs/advanced_features/quantization.mdx index 9a4dc1af5..ee0d4b7e9 100644 --- a/docs_new/docs/advanced_features/quantization.mdx +++ b/docs_new/docs/advanced_features/quantization.mdx @@ -56,6 +56,13 @@ The following table summarizes quantization method support across NVIDIA and AMD 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 + + mxfp_w4a8 + No + No + Yes (A5) + Ascend NPU only; online W4A8 for Qwen3 dense LLM (MXFP4 weights + MXFP8 activations) on A5 series; offline W4A8_MXFP checkpoints are auto-detected via modelslim + blockwise_int8 Yes 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 18c6b8edf..f4ea12d64 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 @@ -60,6 +60,14 @@ SGLang supports **mix-bits** quantization (independently defines and loads each + + MXFP4 W4A8 + Linear + x + x + + x + W4A4 dynamic MoE @@ -347,6 +355,28 @@ python3 -m sglang.launch_server \ > - 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. +**MXFP4 W4A8 for LLM dense models (e.g. Qwen3 / Qwen3.5):** + +LLM dense W4A8 (MXFP4 4-bit weights + MXFP8 8-bit activations) Linear support was added in [PR #23650](https://github.com/sgl-project/sglang/pull/23650). Requires Ascend A5 series or newer. + +- Online W4A8 quantization (BF16/FP16 weights → MXFP4 at load time): + +```bash Command +python3 -m sglang.launch_server \ + --model-path Qwen/Qwen3-8B \ + --quantization mxfp_w4a8 \ + --device npu --attention-backend ascend \ + --host 0.0.0.0 --port 30000 \ + --mem-fraction-static 0.8 --tp-size 1 +``` + +- Offline W4A8 quantization (msmodelslim pre-quantized weights, `W4A8_MXFP` scheme; no `--quantization` flag needed — auto-detected from `quant_model_description.json`). + +> **Implementation Notes:** +> - Weights are packed FP4 (`float4_e2m1fn_x2`, two nibbles per byte) with a UE8M0 per-block shared exponent (block_size = 32); activations are per-token MXFP8. Matmul runs via `npu_quant_matmul(..., x2_dtype=torch_npu.float4_e2m1fn_x2, group_sizes=[0, 0, 32])`. +> - The packed-FP4 dtype passed to the NPU ops (`dst_type` / `x2_dtype` / `input_dtype`) must be resolved from `torch_npu.float4_e2m1fn_x2` (an int enum), not the `torch.float4_e2m1fn_x2` dtype object, which recent op-plugin builds reject. +> - Online and offline share the same kernel path and layout; they differ only in the weight source (RTN at load vs msmodelslim calibration). + ## Diffusion Model Quantization on Ascend NPU SGLang-Diffusion supports MXFP8 online and offline quantization for diffusion models (such as Wan2.2) on Ascend NPUs. MXFP8 requires A5; the ModelSlim W8A8/W4A4 schemes work on A2/A3. diff --git a/python/sglang/srt/hardware_backend/npu/quantization/linear_method_npu.py b/python/sglang/srt/hardware_backend/npu/quantization/linear_method_npu.py index 60066f5cd..2bd602f1d 100644 --- a/python/sglang/srt/hardware_backend/npu/quantization/linear_method_npu.py +++ b/python/sglang/srt/hardware_backend/npu/quantization/linear_method_npu.py @@ -4,7 +4,7 @@ from typing import TYPE_CHECKING, Optional import torch from torch.nn.parameter import Parameter -from sglang.srt.hardware_backend.npu.utils import npu_format_cast +from sglang.srt.hardware_backend.npu.utils import NPUACLFormat, npu_format_cast from sglang.srt.layers.quantization.base_config import LinearMethodBase if TYPE_CHECKING: @@ -13,6 +13,8 @@ if TYPE_CHECKING: logger = logging.getLogger(__name__) MXFP8_BLOCK_SIZE = 32 +# W4A8_MXFP block (group) size — fixed at 32 by the msmodelslim export format. +MXFP4_BLOCK_SIZE = 32 # NPU ops are reached via torch.ops.npu.* (registered when torch_npu is imported @@ -25,6 +27,30 @@ def _get_float8_e8m0fnu_dtype(): return getattr(torch, "float8_e8m0fnu", None) +def _get_float4_e2m1fn_x2_dtype(): + # The packed-FP4 dtype MUST come from torch_npu (an int enum, e.g. 296), not + # from torch. The NPU ops that consume it -- npu_dynamic_mx_quant(dst_type=), + # npu_quant_matmul(x2_dtype=), npu_format_cast(input_dtype=) -- REJECT the + # torch dtype object torch.float4_e2m1fn_x2 in op-plugin on recent torch_npu + # builds (it raises, or with None gives "output y must be same shape as input + # x"), even though torch.float4_e2m1fn_x2 exists. This is fp4-specific: fp8 / + # float8_e8m0fnu is accepted from torch either way. Verified on A5 / + # torch_npu 2.10.0.post2.dev20260704 (see llm/probe_fp4_w4a8_chain.py: dst=296 + # passes the full quant->format_cast->matmul chain, dst=torch dtype fails). + # + # Lazy import so this NPU-only path keeps the module importable on + # CUDA/CPU/AMD/XPU CI (no top-level torch_npu; see AGENTS.md known pitfalls). + from sglang.srt.utils import is_npu + + if is_npu(): + import torch_npu + + npu_dtype = getattr(torch_npu, "float4_e2m1fn_x2", None) + if npu_dtype is not None: + return npu_dtype + return getattr(torch, "float4_e2m1fn_x2", None) + + class _NPULinearMethodBase(LinearMethodBase): def __init__( @@ -310,3 +336,287 @@ class NPU_W4A4DynamicLinearMethod(_NPULinearMethodBase): bias=bias, output_dtype=original_dtype, ) + + +class NPUMXFP4W4A8LinearMethod(_NPULinearMethodBase): + """Ascend NPU W4A8 online quantization: MXFP4 weights + MXFP8 activations. + + This is a *true* W4(weight) A8(activation) path: it mirrors the offline + ``W4A8_MXFP`` kernel (``NPUMXFP4W4A8OfflineLinearMethod``) exactly — the only + difference is that the FP4 weights are produced online from BF16/FP16 + (round-to-nearest, no calibration) instead of being loaded from a msmodelslim + checkpoint. An earlier version of this method ran a *dual-level* scheme that + also compressed the activation to FP4 (W4A4 compute via + ``npu_dual_level_quant_matmul``); that was a large accuracy regression — 4-bit + activations — so it was replaced with the single-level FP8-activation path + below, aligned with the offline W4A8 implementation. + + Weight quantization (process_weights_after_loading): + BF16/FP16 weight → npu_dynamic_mx_quant(dst=float4_e2m1fn_x2) → packed FP4 + + UE8M0 block scale → npu_format_cast to FRACTAL_NZ → transpose [in//2, out] + + Inference (apply): + BF16/FP16 activation → npu_dynamic_mx_quant(dst=float8_e4m3fn) (A8, FP8) + → npu_quant_matmul(x2_dtype=float4_e2m1fn_x2, group_sizes=[0, 0, block]) + + Hardware: Ascend 950 (A5) + a recent torch_npu with the FP4 npu_quant_matmul + (same requirement as the offline W4A8 path — see that class's docstring). + """ + + def create_weights( + self, + layer: torch.nn.Module, + input_size_per_partition: int, + output_partition_sizes, + input_size: int, + output_size: int, + params_dtype: torch.dtype, + **extra_weight_attrs, + ): + """Register an unquantized (``params_dtype``) weight placeholder. + + Online quantization needs its own ``create_weights`` because the + checkpoint still holds full-precision BF16/FP16 weights: the loader + fills this buffer, then ``process_weights_after_loading`` quantizes it to + MXFP4 in place. This differs from the offline/int8 methods, whose weights + are created by the scheme's own ``create_weights`` to match the + already-quantized (FP8 / uint8-packed) layout the checkpoint provides. + """ + from sglang.srt.layers.parameter import ModelWeightParameter + + output_size_per_partition = sum(output_partition_sizes) + weight_loader = extra_weight_attrs.get("weight_loader") + + layer.logical_widths = output_partition_sizes + layer.input_size_per_partition = input_size_per_partition + layer.output_size_per_partition = output_size_per_partition + layer.orig_dtype = params_dtype + + # Load weights in original dtype; quantise to MXFP4 in + # process_weights_after_loading. + weight = ModelWeightParameter( + data=torch.empty( + output_size_per_partition, + input_size_per_partition, + dtype=params_dtype, + ), + input_dim=1, + output_dim=0, + weight_loader=weight_loader, + ) + layer.register_parameter("weight", weight) + + def process_weights_after_loading(self, layer: torch.nn.Module) -> None: + # Online single-level MXFP4 weight quant, then lay the weight out exactly + # like the offline W4A8 path so the same npu_quant_matmul(x2_dtype=fp4) + # kernel accepts it. All NPU ops go through torch.ops.npu.* (no torch_npu). + fp4_dtype = _get_float4_e2m1fn_x2_dtype() + + weight_fp = layer.weight.data + if weight_fp.dtype not in (torch.float16, torch.bfloat16): + weight_fp = weight_fp.to(torch.bfloat16) + # Move to NPU if needed (cpu offload may have put it on CPU). + if not weight_fp.is_npu: + weight_fp = weight_fp.to(f"npu:{torch.npu.current_device()}") + + # BF16 -> packed FP4 (float4_e2m1fn_x2, [out, in//2]) + UE8M0 block scale. + # npu_dynamic_mx_quant returns the scale as [out, in//64, 2] (3D); older + # builds may return [out, in//32] (2D) — handle both before the transpose. + qw, w_scale = torch.ops.npu.npu_dynamic_mx_quant( + weight_fp, dst_type=fp4_dtype, round_mode="round" + ) + + # weight: packed FP4 -> FRACTAL_NZ (float8_e4m3fn view) -> transpose + # [in//2, out]. Mirror the offline path (no .contiguous() on the NZ view); + # view as uint8 first because npu_format_cast only accepts int-dtype tensors. + qw_nz = npu_format_cast( + qw.view(torch.uint8), + NPUACLFormat.ACL_FORMAT_FRACTAL_NZ, + customize_dtype=torch.float8_e4m3fn, + input_dtype=fp4_dtype, + ) + layer.weight = Parameter(qw_nz.transpose(-1, -2), requires_grad=False) + + # weight_scale -> [in//64, out, 2] to match npu_quant_matmul. + if w_scale.dim() == 2: + n, k = w_scale.shape + w_scale = w_scale.reshape(n, k // 2, 2) + layer.weight_scale = Parameter(w_scale.transpose(-3, -2), requires_grad=False) + + # Cache FP32 bias once to avoid a per-forward dtype conversion + alloc. + if ( + getattr(layer, "bias", None) is not None + and layer.bias.dtype != torch.float32 + ): + layer.bias_fp32 = Parameter( + layer.bias.data.to(torch.float32), requires_grad=False + ) + else: + layer.bias_fp32 = None + + def apply( + self, + layer: torch.nn.Module, + x: torch.Tensor, + bias: Optional[torch.Tensor] = None, + ) -> torch.Tensor: + e8m0_dtype = _get_float8_e8m0fnu_dtype() + fp4_dtype = _get_float4_e2m1fn_x2_dtype() + + original_dtype = x.dtype + if original_dtype not in (torch.float16, torch.bfloat16): + x = x.to(torch.bfloat16) + original_dtype = torch.bfloat16 + + # Flatten to 2D [tokens, hidden] for npu_dynamic_mx_quant. + input_shape = x.shape + x_2d = x.reshape(-1, x.shape[-1]) + + # Dynamic MXFP8 activation quantisation (A8 — FP8, not FP4). + quantized_x, dynamic_scale = torch.ops.npu.npu_dynamic_mx_quant( + x_2d, dst_type=torch.float8_e4m3fn + ) + + # Use the cached FP32 bias from process_weights_after_loading; fall back + # to per-call conversion if the cache was bypassed (e.g. dynamic bias). + if bias is None: + quant_bias = None + elif ( + bias is getattr(layer, "bias", None) + and getattr(layer, "bias_fp32", None) is not None + ): + quant_bias = layer.bias_fp32 + else: + quant_bias = bias.to(torch.float32) + + # True W4(weight)A8(activation) matmul, identical to the offline path. + output = torch.ops.npu.npu_quant_matmul( + quantized_x, + layer.weight, + layer.weight_scale, + scale_dtype=e8m0_dtype, + pertoken_scale=dynamic_scale, + pertoken_scale_dtype=e8m0_dtype, + bias=quant_bias, + output_dtype=original_dtype, + x2_dtype=fp4_dtype, + group_sizes=[0, 0, MXFP4_BLOCK_SIZE], + ) + + # Restore original shape (replace last dim with output features). + output_shape = list(input_shape[:-1]) + [output.shape[-1]] + return output.reshape(output_shape) + + +class NPUMXFP4W4A8OfflineLinearMethod(_NPULinearMethodBase): + """Ascend NPU offline W4A8 (ModelSlim ``W4A8_MXFP``): packed-FP4 weights + MXFP8 activations. + + Kernel for the offline ModelSlimMXFP4W4A8Scheme (delegated as ``self.kernel``). + The msmodelslim ``W4A8_MXFP`` checkpoint stores weights as *packed FP4* + (``pack_fp4_to_uint8`` → ``uint8`` shape ``[out, in//2]``) plus UE8M0 block + scales (``uint8`` shape ``[out, in//group_size]``): + + process_weights_after_loading: + weight (uint8 packed FP4 [out, in//2]) → npu_format_cast(29, + customize_dtype=float8_e4m3fn, input_dtype=float4_e2m1fn_x2) → FRACTAL_NZ + → transpose [in//2, out] + weight_scale [out, in/32] → reshape [out, in/64, 2] → transpose → [in/64, out, 2] + + apply: + BF16/FP16 activation → npu_dynamic_mx_quant(dst=float8_e4m3fn) (A8, MXFP8) + → npu_quant_matmul(x2_dtype=float4_e2m1fn_x2, group_sizes=[0, 0, block]) + + Mirrors vllm-ascend ``AscendW4A8MXFPDynamicLinearMethod`` exactly (Ascend 950/A5). + The weight is cast to FRACTAL_NZ then transposed; ``npu_dynamic_mx_quant`` already + returns a 3D ``[tokens, in//64, 2]`` block scale so the matmul needs no extra + scale-layout normalization. + + ⚠️ REQUIRES a recent torch_npu build for the FP4 ``npu_quant_matmul``. On the + A5 this device forces ``allow_internal_format=False`` (the NZ cast still produces + a ``FRACTAL_NZ_C0_16`` tensor, which is fine). Older torch_npu (e.g. + ``2.10.0.dev20260320``) had a broken FP4 matmul that rejected the NZ weight in + *prefill* with ``x2 should be in ... nz format, but it is 2``; + ``2.10.0.post1.dev20260624`` (and later) runs the vllm-aligned NZ path + correctly. If you hit ``it is 2``, update torch_npu — do NOT "fix" it by + switching the weight to ND. + + ⚠️ A ``atb::OperationSetup`` *segfault during decode* (not prefill) is a + DIFFERENT, unrelated issue: it is the eager-decode ``ascend`` attention + backend, NOT this matmul (verified by stage-sync bisection — qkv's matmul + syncs clean, the fault surfaces at the entry-sync of the next layer, i.e. the + decode attention between qkv and o_proj). Run with the NPU decode graph (do + NOT pass ``--disable-cuda-graph``); graph mode is the NPU default and what + vllm uses. This attention issue is model-agnostic and out of scope for W4A8. + + This is a true W4(weight) A8(activation) single-level matmul. The *online* + ``NPUMXFP4W4A8LinearMethod`` now uses this exact apply path — the only + difference is that it quantizes BF16/FP16 weights to FP4 at load time instead + of loading them from a msmodelslim checkpoint. ``group_size`` is fixed at 32 + by the ``W4A8_MXFP`` export format. + """ + + def process_weights_after_loading(self, layer: torch.nn.Module) -> None: + # Mirror vllm-ascend AscendW4A8MXFPDynamicLinearMethod: cast the packed-FP4 + # weight to FRACTAL_NZ then transpose. All NPU ops go through + # torch.ops.npu.* (no torch_npu). Requires a recent torch_npu build (see + # class docstring): older builds reject the NZ weight ("x2 ... it is 2"). + fp4_dtype = _get_float4_e2m1fn_x2_dtype() + + # weight: packed-FP4 uint8 [out, in//2] -> FRACTAL_NZ (float8_e4m3fn view) + # -> transpose to [in//2, out]. + layer.weight.data = npu_format_cast( + layer.weight.data, + NPUACLFormat.ACL_FORMAT_FRACTAL_NZ, + customize_dtype=torch.float8_e4m3fn, + input_dtype=fp4_dtype, + ) + layer.weight.data = layer.weight.data.transpose(-1, -2) + # weight_scale: [out, in/32] uint8 -> [in/64, out, 2]. + n, k = layer.weight_scale.data.shape + layer.weight_scale.data = layer.weight_scale.data.reshape( + n, k // 2, 2 + ).transpose(-3, -2) + + def apply( + self, + layer: torch.nn.Module, + x: torch.Tensor, + bias: Optional[torch.Tensor] = None, + ) -> torch.Tensor: + e8m0_dtype = _get_float8_e8m0fnu_dtype() + fp4_dtype = _get_float4_e2m1fn_x2_dtype() + + original_dtype = x.dtype + if original_dtype not in (torch.float16, torch.bfloat16): + x = x.to(torch.bfloat16) + original_dtype = torch.bfloat16 + + # Flatten to 2D [tokens, hidden] for npu_dynamic_mx_quant. + input_shape = x.shape + x_2d = x.reshape(-1, x.shape[-1]) + + # Dynamic MXFP8 activation quantisation (A8). + quantized_x, dynamic_scale = torch.ops.npu.npu_dynamic_mx_quant( + x_2d, dst_type=torch.float8_e4m3fn + ) + + if bias is not None and bias.dtype != torch.float32: + bias = bias.to(torch.float32) + + # W4(weight)A8(activation) matmul, mirroring vllm-ascend exactly. + output = torch.ops.npu.npu_quant_matmul( + quantized_x, + layer.weight, + layer.weight_scale, + scale_dtype=e8m0_dtype, + pertoken_scale=dynamic_scale, + pertoken_scale_dtype=e8m0_dtype, + bias=bias, + output_dtype=original_dtype, + x2_dtype=fp4_dtype, + group_sizes=[0, 0, MXFP4_BLOCK_SIZE], + ) + + # Restore original shape (replace last dim with output features). + output_shape = list(input_shape[:-1]) + [output.shape[-1]] + return output.reshape(output_shape) diff --git a/python/sglang/srt/hardware_backend/npu/utils.py b/python/sglang/srt/hardware_backend/npu/utils.py index 4679c4970..15d117878 100644 --- a/python/sglang/srt/hardware_backend/npu/utils.py +++ b/python/sglang/srt/hardware_backend/npu/utils.py @@ -140,6 +140,9 @@ def _is_nz_aligned(tensor: torch.Tensor) -> bool: def npu_format_cast( tensor: torch.Tensor, acl_format: NPUACLFormat = NPUACLFormat.ACL_FORMAT_FRACTAL_NZ, + *, + customize_dtype=None, + input_dtype=None, ) -> torch.Tensor: """ Cast a tensor to a specific NPU ACL format. @@ -147,6 +150,12 @@ def npu_format_cast( Args: tensor (torch.Tensor): The input tensor. acl_format (NPUACLFormat): The target NPU ACL format. + customize_dtype / input_dtype: packed-FP4 unpack kwargs (e.g. + ``customize_dtype=torch.float8_e4m3fn``, + ``input_dtype=torch.float4_e2m1fn_x2``). When either is set the unpack + kwargs are forwarded to the op and the ``_is_nz_aligned`` ND fallback + is skipped: the FP4 matmul strictly requires FRACTAL_NZ, so a silent + ND fallback would corrupt results. Returns: torch.Tensor: The tensor cast to the specified NPU ACL format. @@ -166,6 +175,21 @@ def npu_format_cast( ) return tensor + # Skip format cast for meta tensors (used in offloader) + if tensor.device.type == "meta": + return tensor + + # Packed-FP4 → FRACTAL_NZ: forward the unpack kwargs to the op, and skip the + # _is_nz_aligned ND fallback — the FP4 matmul strictly requires NZ, so a + # silent ND fallback would corrupt results. + if customize_dtype is not None or input_dtype is not None: + return torch.ops.npu.npu_format_cast( + tensor, + int(acl_format), + customize_dtype=customize_dtype, + input_dtype=input_dtype, + ) + if acl_format == NPUACLFormat.ACL_FORMAT_FRACTAL_NZ and not _is_nz_aligned(tensor): k, n = tensor.shape[-2], tensor.shape[-1] logger.warning_once( @@ -178,10 +202,6 @@ def npu_format_cast( ) return tensor - # Skip format cast for meta tensors (used in offloader) - if tensor.device.type == "meta": - return tensor - return torch.ops.npu.npu_format_cast(tensor, acl_format.value) diff --git a/python/sglang/srt/layers/quantization/__init__.py b/python/sglang/srt/layers/quantization/__init__.py index b44042d9b..1762bcbb4 100644 --- a/python/sglang/srt/layers/quantization/__init__.py +++ b/python/sglang/srt/layers/quantization/__init__.py @@ -44,6 +44,7 @@ from sglang.srt.layers.quantization.modelopt_quant import ( from sglang.srt.layers.quantization.modelslim.modelslim import ModelSlimConfig from sglang.srt.layers.quantization.moe_wna16 import MoeWNA16Config from sglang.srt.layers.quantization.mxfp4 import Mxfp4Config +from sglang.srt.layers.quantization.npu_mxfp4 import Mxfp4W4A8Config from sglang.srt.layers.quantization.nvfp4_online import NvFp4OnlineConfig from sglang.srt.layers.quantization.petit import PetitNvFp4Config from sglang.srt.layers.quantization.qoq import QoQConfig @@ -98,6 +99,7 @@ BASE_QUANTIZATION_METHODS: Dict[str, Type[QuantizationConfig]] = { "auto-round-int8": W8A8Int8Config, "modelslim": ModelSlimConfig, "quark_int4fp8_moe": QuarkInt4Fp8Config, + "mxfp_w4a8": Mxfp4W4A8Config, } diff --git a/python/sglang/srt/layers/quantization/modelslim/modelslim.py b/python/sglang/srt/layers/quantization/modelslim/modelslim.py index a4074e5b8..b69936985 100644 --- a/python/sglang/srt/layers/quantization/modelslim/modelslim.py +++ b/python/sglang/srt/layers/quantization/modelslim/modelslim.py @@ -14,6 +14,7 @@ from sglang.srt.layers.quantization.base_config import ( QuantizationConfig, ) from sglang.srt.layers.quantization.modelslim.schemes import ( + ModelSlimMXFP4W4A8Scheme, ModelSlimMXFP8Scheme, ModelSlimW4A4Int4, ModelSlimW4A4Int4MoE, @@ -198,6 +199,7 @@ class ModelSlimConfig(QuantizationConfig): ("W8A8", ModelSlimW8A8Int8), ("W8A8_DYNAMIC", ModelSlimW8A8Int8), ("W8A8_MXFP8", ModelSlimMXFP8Scheme), + ("W4A8_MXFP", ModelSlimMXFP4W4A8Scheme), ] quant_schemes = [self.quant_description.get(prefix + ".weight", "")] diff --git a/python/sglang/srt/layers/quantization/modelslim/schemes/__init__.py b/python/sglang/srt/layers/quantization/modelslim/schemes/__init__.py index bfc2a350c..1b0233967 100644 --- a/python/sglang/srt/layers/quantization/modelslim/schemes/__init__.py +++ b/python/sglang/srt/layers/quantization/modelslim/schemes/__init__.py @@ -6,6 +6,7 @@ # isort: off from .modelslim_scheme import ModelSlimLinearScheme, ModelSlimMoEScheme from .modelslim_mxfp8 import ModelSlimMXFP8Scheme +from .modelslim_mxfp4_w4a8 import ModelSlimMXFP4W4A8Scheme # isort: on from .modelslim_w4a4_int4 import ModelSlimW4A4Int4 @@ -18,6 +19,7 @@ __all__ = [ "ModelSlimLinearScheme", "ModelSlimMoEScheme", "ModelSlimMXFP8Scheme", + "ModelSlimMXFP4W4A8Scheme", "ModelSlimW8A8Int8", "ModelSlimW4A4Int4", "ModelSlimW4A4Int4MoE", diff --git a/python/sglang/srt/layers/quantization/modelslim/schemes/modelslim_mxfp4_w4a8.py b/python/sglang/srt/layers/quantization/modelslim/schemes/modelslim_mxfp4_w4a8.py new file mode 100644 index 000000000..ba841d7a0 --- /dev/null +++ b/python/sglang/srt/layers/quantization/modelslim/schemes/modelslim_mxfp4_w4a8.py @@ -0,0 +1,107 @@ +"""ModelSlim W4A8_MXFP scheme for pre-quantized weight inference on Ascend NPU (SRT). + +The msmodelslim ``W4A8_MXFP`` checkpoint stores weights as **packed FP4**: + + weight: uint8 (pack_fp4_to_uint8), shape [out, in//2], group_size=32 + weight_scale: uint8 (UE8M0, +127 biased), shape [out, in//32] + +(verified on ``Qwen3-8B-mxw4a8-pack-full`` and matching the msmodelslim exporter +``ascendv1.py:on_w4a8_mx_dynamic_per_block``). This is a true W4(weight) A8(activation) +scheme: weights are 4-bit FP4, activations are dynamically quantised to MXFP8. + +This is NOT the same layout as ``W8A8_MXFP8`` (which stores float8_e4m3fn weights +of shape [out, in]) — so weight creation and the forward pass differ from MXFP8. +Weight post-processing and the matmul are delegated to ``NPUMXFP4W4A8OfflineLinearMethod`` +(``self.kernel``), mirroring vllm-ascend's ``AscendW4A8MXFPDynamicLinearMethod``: +``npu_format_cast`` the packed FP4 to FRACTAL_NZ + transpose, then ``x2_dtype= +float4_e2m1fn_x2`` matmul with ``group_sizes=[0, 0, 32]``. Requires a recent +torch_npu for the FP4 matmul on Ascend 950/A5 (older builds reject the NZ weight) — +see ``NPUMXFP4W4A8OfflineLinearMethod`` for the version caveat. +""" + +from typing import Dict, List, Optional + +import torch + +from sglang.srt.hardware_backend.npu.quantization.linear_method_npu import ( + NPUMXFP4W4A8OfflineLinearMethod, +) +from sglang.srt.layers.parameter import GroupQuantScaleParameter, ModelWeightParameter +from sglang.srt.layers.quantization.modelslim.schemes import ModelSlimLinearScheme + +# Fixed by the msmodelslim W4A8_MXFP export format (ascendv1.py sets group_size=32). +MXFP4_W4A8_BLOCK_SIZE = 32 +# FP4 weights are bit-packed two-per-byte along the input (reduction) dim. +MXFP4_W4A8_PACK_FACTOR = 2 + + +class ModelSlimMXFP4W4A8Scheme(ModelSlimLinearScheme): + """W4A8_MXFP offline scheme — packed-FP4 weights, MXFP8 activations.""" + + def __init__( + self, + quant_config: Optional[Dict[str, any]] = None, + prefix: Optional[str] = None, + ): + # quant_config / prefix accepted to match ModelSlimConfig.get_linear_scheme's + # dispatch signature; W4A8_MXFP needs no per-layer config beyond create_weights. + del quant_config, prefix + self.kernel = NPUMXFP4W4A8OfflineLinearMethod() + + def create_weights( + self, + layer: torch.nn.Module, + input_size_per_partition: int, + output_partition_sizes: List[int], + input_size: int, + output_size: int, + params_dtype: torch.dtype, + **extra_weight_attrs, + ): + weight_loader = extra_weight_attrs.get("weight_loader") + output_size_per_partition = sum(output_partition_sizes) + + # Packed-FP4 weight: uint8, shape [out, in//2] (two FP4 nibbles per byte + # along the input dim). input_dim=1 is the packed dim; TP row-parallel + # sharding narrows by self.data.shape[input_dim] (already halved), so a + # plain ModelWeightParameter shards correctly without packing metadata + # (FP4 packs the reduction dim only; the output dim stays unpacked). + weight = ModelWeightParameter( + data=torch.empty( + ( + output_size_per_partition, + input_size_per_partition // MXFP4_W4A8_PACK_FACTOR, + ), + dtype=torch.uint8, + ), + input_dim=1, + output_dim=0, + weight_loader=weight_loader, + ) + layer.register_parameter("weight", weight) + + # UE8M0 block scales: uint8, shape [out, in//32]. Named "weight_scale" to + # match the checkpoint key; the kernel re-layouts it into weight_scale_inv + # during process_weights_after_loading. + scale_dim = input_size_per_partition // MXFP4_W4A8_BLOCK_SIZE + weight_scale = GroupQuantScaleParameter( + data=torch.empty( + (output_size_per_partition, scale_dim), + dtype=torch.uint8, + ), + input_dim=1, + output_dim=0, + weight_loader=weight_loader, + ) + layer.register_parameter("weight_scale", weight_scale) + + def process_weights_after_loading(self, layer: torch.nn.Module): + self.kernel.process_weights_after_loading(layer) + + def apply_weights( + self, + layer: torch.nn.Module, + x: torch.Tensor, + bias: Optional[torch.Tensor] = None, + ) -> torch.Tensor: + return self.kernel.apply(layer, x, bias) diff --git a/python/sglang/srt/layers/quantization/npu_mxfp4.py b/python/sglang/srt/layers/quantization/npu_mxfp4.py new file mode 100644 index 000000000..1f0bfb5e2 --- /dev/null +++ b/python/sglang/srt/layers/quantization/npu_mxfp4.py @@ -0,0 +1,127 @@ +"""MXFP4 W4A8 online quantization config (MXFP4 weights + MXFP8 activations). + +Triggered by ``--quantization mxfp_w4a8``. + +Online mode: FP16/BF16 weights are quantised to MXFP4 in +``process_weights_after_loading``; activations are dynamically quantised to +MXFP8 (``float8_e4m3fn`` + UE8M0 block scale) at inference time and the matmul +runs via ``npu_quant_matmul`` with FP4 weights. + +The config is device-agnostic and dispatches per device in +``get_quant_method``; only the Ascend NPU backend (Ascend 950 / A5) is +implemented today. +""" + +from __future__ import annotations + +import logging +from typing import Dict, List, Optional + +import torch + +from sglang.srt.layers.quantization.base_config import ( + QuantizationConfig, + QuantizeMethodBase, +) +from sglang.srt.layers.quantization.unquant import ( + UnquantizedFusedMoEMethod, + UnquantizedLinearMethod, +) +from sglang.srt.layers.quantization.utils import is_layer_skipped +from sglang.srt.utils import is_npu + +logger = logging.getLogger(__name__) + + +class Mxfp4W4A8Config(QuantizationConfig): + """MXFP4 W4A8 online quantization config; dispatches per device. + + True W4(weight) A8(activation): weights are quantised online to MXFP4 and + activations to MXFP8 at inference time. The device-specific linear method + is selected in ``get_quant_method``; only Ascend NPU is wired up today. + """ + + def __init__( + self, + ignored_layers: Optional[List[str]] = None, + packed_modules_mapping: Optional[Dict[str, str]] = None, + ): + super().__init__() + self.ignored_layers = ignored_layers or [] + self.packed_modules_mapping = packed_modules_mapping or {} + + @classmethod + def get_name(cls) -> str: + return "mxfp_w4a8" + + @classmethod + def get_supported_act_dtypes(cls) -> List[torch.dtype]: + return [torch.bfloat16, torch.half] + + @classmethod + def get_min_capability(cls) -> int: + return 0 # NPU bypasses CUDA capability checks + + @classmethod + def get_config_filenames(cls) -> List[str]: + return [] + + @classmethod + def from_config(cls, config: Dict) -> Mxfp4W4A8Config: + ignored_layers = cls.get_from_keys_or( + config, ["ignored_layers", "modules_to_not_convert"], None + ) + if ignored_layers: + normalized: List[str] = [] + for layer in ignored_layers: + base = layer.removeprefix("model.") + normalized.append(base) + normalized.append(f"model.{base}") + ignored_layers = normalized + packed_modules_mapping = ( + cls.get_from_keys_or(config, ["packed_modules_mapping"], {}) or {} + ) + return cls( + ignored_layers=ignored_layers, + packed_modules_mapping=packed_modules_mapping, + ) + + def get_quant_method( + self, layer: torch.nn.Module, prefix: str + ) -> Optional[QuantizeMethodBase]: + from sglang.srt.layers.linear import LinearBase + from sglang.srt.layers.moe.fused_moe_triton import FusedMoE + + if isinstance(layer, LinearBase): + if is_layer_skipped( + prefix, + self.ignored_layers, + fused_mapping=self.packed_modules_mapping, + ): + return UnquantizedLinearMethod() + if is_npu(): + from sglang.srt.hardware_backend.npu.quantization.linear_method_npu import ( + NPUMXFP4W4A8LinearMethod, + ) + + return NPUMXFP4W4A8LinearMethod(self) + raise NotImplementedError( + "mxfp_w4a8 (MXFP4 weights + MXFP8 activations, W4A8) is currently " + "only implemented for the Ascend NPU backend; no CUDA/other-device " + "kernel exists yet. Add a device branch here when one lands." + ) + elif isinstance(layer, FusedMoE): + # MoE MXFP4 not yet implemented; fall back to unquantised + logger.warning( + "MXFP4 W4A8 quantization is not yet supported for FusedMoE layers " + "(prefix=%s). Falling back to unquantized MoE — MoE weights will " + "run in full precision (BF16/FP16).", + prefix, + ) + return UnquantizedFusedMoEMethod( + layer.use_triton_kernels, layer.use_flashinfer_trtllm_moe + ) + return None + + def get_scaled_act_names(self) -> List[str]: + return [] diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index 1f403309b..884ba1877 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -153,6 +153,7 @@ QUANTIZATION_CHOICES = [ "auto-round-int8", "compressed-tensors", # for Ktransformers "modelslim", # for NPU + "mxfp_w4a8", # for NPU W4A8 (MXFP4 weights + MXFP8 activations) "quark", # AMD Quark quantizer (FP8 / MXFP4 / Int4FP8 etc.) "quark_int4fp8_moe", "quark_mxfp4", # Online MOE + linear quantization.