[llm][npu][quant] Add W8A8 MXFP8 quantization support for Qwen3 Dense on Ascend NPU (#22352)

Co-authored-by: ronnie_zheng <zl19940307@163.com>
This commit is contained in:
Junlin Wu
2026-06-16 09:45:18 +03:00
committed by GitHub
co-authored by ronnie_zheng
parent 72d962be88
commit 2a8ea70059
8 changed files with 376 additions and 11 deletions
@@ -30,7 +30,7 @@ The following table summarizes quantization method support across NVIDIA and AMD
<th>Method</th>
<th>NVIDIA GPUs</th>
<th>AMD GPUs (MI300X/MI325X/MI350X)</th>
<th>Ascend NPUs (A2/A3)</th>
<th>Ascend NPUs (A2/A3/A5)</th>
<th>Notes</th>
</tr>
</thead>
@@ -49,6 +49,13 @@ The following table summarizes quantization method support across NVIDIA and AMD
<td>WIP</td>
<td>Requires CDNA3/CDNA4 with MXFP support; uses Aiter</td>
</tr>
<tr>
<td><code>mxfp8</code></td>
<td>No</td>
<td>No</td>
<td>Yes (A5 for Diffusion and LLM Dense Linear)</td>
<td>Ascend NPU only; online MXFP8 quantization for Diffusion models (e.g., Wan2.2) and LLM Dense Linear on A5 series; uses CANN <code>npu_dynamic_mx_quant</code> / <code>npu_quant_matmul</code> kernels</td>
</tr>
<tr>
<td><code>blockwise_int8</code></td>
<td>Yes</td>
@@ -182,13 +189,7 @@ The following table summarizes quantization method support across NVIDIA and AMD
<td>Yes</td>
<td>Ascend quantization; Uses CANN kernels</td>
</tr>
<tr>
<td><code>mxfp8</code> (diffusion)</td>
<td>No</td>
<td>No</td>
<td>Yes (A2/A3)</td>
<td>Ascend NPU only; online MXFP8 quantization for diffusion models (e.g., Wan2.2); requires CANN ≥ 8.0.RC3</td>
</tr>
</tbody>
</table>
@@ -45,7 +45,7 @@ SGLang support **mix-bits** quantization (independently defines and loads each l
<td><strong style={{color: 'green'}}>√</strong></td>
</tr>
<tr>
<td><a href="https://github.com/sgl-project/sglang/pull/20922">MXFP8</a></td>
<td>MXFP8 (<a href="https://github.com/sgl-project/sglang/pull/20922">Diffusion</a>, <a href="https://github.com/sgl-project/sglang/pull/22352">LLM dense</a>)</td>
<td>Linear</td>
<td><strong style={{color: 'red'}}>x</strong></td>
<td><strong style={{color: 'red'}}>x</strong></td>
@@ -316,6 +316,36 @@ python3 -m sglang.launch_server \
> - MoE layers use `npu_grouped_matmul` and `npu_moe_init_routing` / `npu_moe_finalize_routing` for high-performance expert computation.
> - TP (tensor parallelism) sharding is supported for both dense and MoE GGUF models.
**MXFP8 for LLM dense models (e.g. Qwen3 / Qwen3.5):**
LLM dense W8A8 MXFP8 Linear support on Ascend was added in [PR #22352](https://github.com/sgl-project/sglang/pull/22352). Requires Ascend A5 series or newer (`npu_dynamic_mx_quant` is not available on A2 / A3).
- Online MXFP8 quantization (BF16/FP16 weights → MXFP8 at load time):
```bash Command
python3 -m sglang.launch_server \
--model-path Qwen/Qwen3-8B \
--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 needed — auto-detected from `quant_model_description.json`):
```bash Command
python3 -m sglang.launch_server \
--model-path /path/to/Qwen3-8B-W8A8-MXFP8 \
--device npu --attention-backend ascend \
--host 0.0.0.0 --port 30000 \
--mem-fraction-static 0.8 --tp-size 1
```
> **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.
## 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.
@@ -1,13 +1,30 @@
import logging
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.layers.quantization.base_config import LinearMethodBase
from sglang.srt.platforms import current_platform
_is_npu = current_platform.is_npu()
if _is_npu:
import torch_npu
if TYPE_CHECKING:
from sglang.srt.layers.quantization.base_config import QuantizationConfig
logger = logging.getLogger(__name__)
MXFP8_BLOCK_SIZE = 32
_FLOAT8_E8M0FNU_DTYPE = (
getattr(torch_npu, "float8_e8m0fnu", getattr(torch, "float8_e8m0fnu", None))
if _is_npu
else getattr(torch, "float8_e8m0fnu", None)
)
class _NPULinearMethodBase(LinearMethodBase):
@@ -111,6 +128,135 @@ class NPUW8A8Int8DynamicLinearMethod(_NPULinearMethodBase):
)
class NPUMXFP8LinearMethod(_NPULinearMethodBase):
"""Ascend NPU MXFP8 linear method for LLM (SRT) models.
Online mode: loads FP16/BF16 weights → quantises to MXFP8 at load time.
Inference: dynamic MXFP8 activation quant + MXFP8 matmul (block_size=32).
"""
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,
):
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 later 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:
weight_fp = layer.weight.data
if weight_fp.dtype not in (torch.float16, torch.bfloat16):
logger.warning(
"NPUMXFP8LinearMethod: weight dtype %s is not float16/bfloat16; "
"casting to bfloat16 before MXFP8 quantisation.",
weight_fp.dtype,
)
weight_fp = weight_fp.to(torch.bfloat16)
# Move weight to NPU if needed (cpu offload may have moved it back to CPU)
if not weight_fp.is_npu:
weight_fp = weight_fp.to(f"npu:{torch.npu.current_device()}")
# Online MXFP8 quantisation of weights (block_size=32).
# qw: [out, in] float8_e4m3fn, w_scale: [out, in//64, 2] uint8.
qw, w_scale = torch_npu.npu_dynamic_mx_quant(
weight_fp, dst_type=torch_npu.float8_e4m3fn
)
# Transpose to [in, out] / [in//64, out, 2] as a strided view — DO NOT
# call .contiguous(). The matmul reduction loop scans the in-dim per
# output column; the [out, in] row-major layout gives stride-1 access
# for that scan via the transpose view (matches msmodelslim's offline
# layout and vllm-ascend's AscendW8A8MXFP8DynamicLinearMethod). Calling
# .contiguous() physically reorders to [in, out] row-major, which makes
# the inner-loop stride = out and tanks HBM bandwidth.
layer.weight = Parameter(qw.transpose(0, 1), requires_grad=False)
layer.weight_scale_inv = Parameter(w_scale.transpose(0, 1), 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:
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
qx, input_scale = torch_npu.npu_dynamic_mx_quant(
x_2d, dst_type=torch_npu.float8_e4m3fn
)
# MXFP8 matmul (weight & scale already transposed at load time)
# 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)
output = torch_npu.npu_quant_matmul(
qx,
layer.weight,
layer.weight_scale_inv,
scale_dtype=_FLOAT8_E8M0FNU_DTYPE,
pertoken_scale=input_scale,
pertoken_scale_dtype=_FLOAT8_E8M0FNU_DTYPE,
bias=quant_bias,
output_dtype=original_dtype,
group_sizes=[1, 1, MXFP8_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 NPU_W4A4DynamicLinearMethod(_NPULinearMethodBase):
def process_weights_after_loading(self, layer):
@@ -205,6 +205,8 @@ class Fp8Config(QuantizationConfig):
return [torch.bfloat16, torch.half]
def get_min_capability(self) -> int:
if is_npu():
return 0 # NPU bypasses CUDA capability checks
if _is_musa:
return 31
@@ -261,6 +263,12 @@ class Fp8Config(QuantizationConfig):
prefix, self.ignored_layers, fused_mapping=self.packed_modules_mapping
):
return UnquantizedLinearMethod()
if is_npu() and self.use_mxfp8:
from sglang.srt.hardware_backend.npu.quantization.linear_method_npu import (
NPUMXFP8LinearMethod,
)
return NPUMXFP8LinearMethod(self)
return Fp8LinearMethod(self)
elif isinstance(layer, FusedMoE):
if is_layer_skipped(
@@ -14,6 +14,7 @@ from sglang.srt.layers.quantization.base_config import (
QuantizationConfig,
)
from sglang.srt.layers.quantization.modelslim.schemes import (
ModelSlimMXFP8Scheme,
ModelSlimW4A4Int4,
ModelSlimW4A4Int4MoE,
ModelSlimW4A8Int8MoE,
@@ -183,6 +184,7 @@ class ModelSlimConfig(QuantizationConfig):
("W4A4_DYNAMIC", ModelSlimW4A4Int4),
("W8A8", ModelSlimW8A8Int8),
("W8A8_DYNAMIC", ModelSlimW8A8Int8),
("W8A8_MXFP8", ModelSlimMXFP8Scheme),
]
quant_schemes = [self.quant_description.get(prefix + ".weight", "")]
@@ -1,6 +1,13 @@
# SPDX-License-Identifier: Apache-2.0
# NOTE: Import order is critical to avoid circular dependency.
# modelslim_mxfp8 imports ModelSlimLinearScheme from this package,
# so the base class must be imported first.
# isort: off
from .modelslim_scheme import ModelSlimLinearScheme, ModelSlimMoEScheme
from .modelslim_mxfp8 import ModelSlimMXFP8Scheme
# isort: on
from .modelslim_w4a4_int4 import ModelSlimW4A4Int4
from .modelslim_w4a4_int4_moe import ModelSlimW4A4Int4MoE
from .modelslim_w4a8_int8_moe import ModelSlimW4A8Int8MoE
@@ -10,6 +17,7 @@ from .modelslim_w8a8_int8_moe import ModelSlimW8A8Int8MoE
__all__ = [
"ModelSlimLinearScheme",
"ModelSlimMoEScheme",
"ModelSlimMXFP8Scheme",
"ModelSlimW8A8Int8",
"ModelSlimW4A4Int4",
"ModelSlimW4A4Int4MoE",
@@ -0,0 +1,146 @@
"""ModelSlim MXFP8 scheme for pre-quantized weight inference on Ascend NPU (SRT).
Loads weights pre-quantized by msmodelslim (float8_e4m3fn weights,
uint8 scales) and runs MXFP8 matmul at inference.
"""
from typing import Dict, List, Optional
import torch
from sglang.srt.layers.parameter import GroupQuantScaleParameter, ModelWeightParameter
from sglang.srt.layers.quantization.modelslim.schemes import ModelSlimLinearScheme
from sglang.srt.platforms import current_platform
_is_npu = current_platform.is_npu()
if _is_npu:
import torch_npu
MXFP8_BLOCK_SIZE = 32
_FLOAT8_E8M0FNU_DTYPE = (
getattr(torch_npu, "float8_e8m0fnu", getattr(torch, "float8_e8m0fnu", None))
if _is_npu
else getattr(torch, "float8_e8m0fnu", None)
)
class ModelSlimMXFP8Scheme(ModelSlimLinearScheme):
def __init__(
self,
quant_config: Optional[Dict[str, any]] = None,
prefix: Optional[str] = None,
):
# quant_config / prefix are accepted to match the linear-scheme
# dispatch signature used by ModelSlimConfig.get_linear_scheme;
# MXFP8 needs no per-layer config beyond what create_weights derives.
del quant_config, prefix
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)
# msmodelslim exports weight as float8_e4m3fn, shape [out, in]
weight = ModelWeightParameter(
data=torch.empty(
(output_size_per_partition, input_size_per_partition),
dtype=torch.float8_e4m3fn,
),
input_dim=1,
output_dim=0,
weight_loader=weight_loader,
)
layer.register_parameter("weight", weight)
# msmodelslim exports weight_scale as uint8, shape [out, in/32].
# NOTE: Named "weight_scale" (not "weight_scale_inv") to match the
# checkpoint key exported by msmodelslim.
scale_dim = input_size_per_partition // MXFP8_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):
# Pre-transpose weight and scale to [in, out] for npu_quant_matmul.
# Use .data assignment without .contiguous() to preserve the transpose
# view strides — npu_quant_matmul reads strides correctly and calling
# .contiguous() would reorder data, breaking the block-scale mapping.
n_dim, k_dim = layer.weight_scale.data.shape
layer.weight_scale.data = layer.weight_scale.data.reshape(n_dim, k_dim // 2, 2)
layer.weight.data = layer.weight.data.transpose(0, 1)
layer.weight_scale.data = layer.weight_scale.data.transpose(0, 1)
# 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 = torch.nn.Parameter(
layer.bias.data.to(torch.float32), requires_grad=False
)
else:
layer.bias_fp32 = None
def apply_weights(
self,
layer: torch.nn.Module,
x: torch.Tensor,
bias: Optional[torch.Tensor] = None,
) -> torch.Tensor:
original_dtype = x.dtype
if original_dtype not in (torch.float16, torch.bfloat16):
x = x.to(torch.bfloat16)
original_dtype = torch.bfloat16
# npu_dynamic_mx_quant requires a 2D input [tokens, hidden_size]
input_shape = x.shape
x_2d = x.reshape(-1, x.shape[-1])
# Dynamic MXFP8 activation quantisation
qx, input_scale = torch_npu.npu_dynamic_mx_quant(
x_2d, dst_type=torch_npu.float8_e4m3fn
)
# MXFP8 matmul (weight & scale already transposed at load time).
# Use the cached FP32 bias from process_weights_after_loading.
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)
output = torch_npu.npu_quant_matmul(
qx,
layer.weight,
layer.weight_scale,
scale_dtype=_FLOAT8_E8M0FNU_DTYPE,
pertoken_scale=input_scale,
pertoken_scale_dtype=_FLOAT8_E8M0FNU_DTYPE,
bias=quant_bias,
output_dtype=original_dtype,
group_sizes=[1, 1, MXFP8_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)
@@ -2,6 +2,7 @@
from __future__ import annotations
import logging
from typing import TYPE_CHECKING, Dict, List, Optional, Tuple, Union
import torch
@@ -25,6 +26,8 @@ from sglang.srt.utils import (
if TYPE_CHECKING:
from sglang.jit_kernel.rope import FusedSetKVBufferArg # For type check-only
logger = logging.getLogger(__name__)
_is_cuda = is_cuda()
_is_hip = is_hip()
_use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip
@@ -40,7 +43,25 @@ if _is_cuda:
if _is_npu:
import torch_npu
from sgl_kernel_npu.norm.fused_rope_qk_mqa import fused_rope_qk_mqa
# `fused_rope_qk_mqa` is an optional fast-path kernel shipped with
# `sgl_kernel_npu`. Older NPU CANN / sgl_kernel_npu builds may not include
# it. If we let the ImportError propagate, importing this module fails,
# which in turn causes `ModelRegistry` to silently skip every model that
# depends on it (and fall back to HF Transformers without quantisation
# awareness — see PR #22352). We tolerate the missing kernel so model
# loading still works; call sites must check for `None` and use the
# generic rope path. A warning is emitted so the missing kernel is
# visible in logs instead of being silently swallowed.
try:
from sgl_kernel_npu.norm.fused_rope_qk_mqa import fused_rope_qk_mqa
except ImportError:
fused_rope_qk_mqa = None
logger.warning(
"sgl_kernel_npu.norm.fused_rope_qk_mqa is unavailable; "
"falling back to the generic rope implementation. Upgrade "
"sgl_kernel_npu to enable the fused kernel."
)
if _is_hip:
from sglang.srt.layers.attention.utils import (
@@ -272,7 +293,10 @@ class RotaryEmbedding(MultiPlatformOp):
else:
cos_sin = self.cos_sin_cache.index_select(0, positions)
if query.shape[0] * query.shape[1] < 65535:
if (
fused_rope_qk_mqa is not None
and query.shape[0] * query.shape[1] < 65535
):
return fused_rope_qk_mqa(
query,
key,