♻️ [llm][npu][quant] Delegate MXFP8 dense scheme to kernel and use torch.ops.npu (#28505)

This commit is contained in:
Junlin Wu
2026-06-17 10:18:26 -07:00
committed by GitHub
parent 735a256f98
commit 873196f7fa
2 changed files with 79 additions and 113 deletions
@@ -6,12 +6,6 @@ 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
@@ -19,11 +13,16 @@ if TYPE_CHECKING:
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)
)
# NPU ops are reached via torch.ops.npu.* (registered when torch_npu is imported
# by the runtime), so this module needs no top-level `import torch_npu` and stays
# importable on CUDA/CPU/AMD/XPU CI.
def _get_float8_e8m0fnu_dtype():
# Resolve lazily rather than as a module-level constant: this module is
# imported early (during quant-scheme registration), so reading the dtype at
# call time keeps it correct regardless of import order / platform.
return getattr(torch, "float8_e8m0fnu", None)
class _NPULinearMethodBase(LinearMethodBase):
@@ -131,8 +130,12 @@ 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).
Shared kernel for both the online config path (``--quantization mxfp8``) and
the offline ModelSlimMXFP8Scheme (which delegates to this as ``self.kernel``).
process_weights_after_loading branches on weight dtype: FP16/BF16 weights are
quantised to MXFP8 at load time (online); pre-quantised float8_e4m3fn weights
are only re-laid-out (offline). Inference: dynamic MXFP8 activation quant +
MXFP8 matmul (block_size=32).
"""
def create_weights(
@@ -169,33 +172,52 @@ class NPUMXFP8LinearMethod(_NPULinearMethodBase):
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 = layer.weight.data
if weight.dtype == torch.float8_e4m3fn:
# Offline (ModelSlim) path: weight is already MXFP8-quantised and
# layer.weight_scale holds the uint8 block scales [out, in/32]. Only
# re-layout to [in, out] / [in//64, out, 2] strided views below.
n_dim, k_dim = layer.weight_scale.data.shape
scale = layer.weight_scale.data.reshape(n_dim, k_dim // 2, 2)
layer.weight = Parameter(weight.transpose(0, 1), requires_grad=False)
layer.weight_scale_inv = Parameter(
scale.transpose(0, 1), requires_grad=False
)
# weight_scale is now folded into weight_scale_inv (which keeps the
# underlying storage alive via its view); drop the stale parameter so
# it doesn't linger in named_parameters() / state_dict().
del layer.weight_scale
else:
# Online path: quantise FP16/BF16 weights to MXFP8 at load time.
if weight.dtype not in (torch.float16, torch.bfloat16):
logger.warning(
"NPUMXFP8LinearMethod: weight dtype %s is not float16/bfloat16; "
"casting to bfloat16 before MXFP8 quantisation.",
weight.dtype,
)
weight = weight.to(torch.bfloat16)
# Move weight to NPU if needed (cpu offload may move it back to CPU).
if not weight.is_npu:
weight = weight.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.ops.npu.npu_dynamic_mx_quant(
weight, dst_type=torch.float8_e4m3fn
)
layer.weight = Parameter(qw.transpose(0, 1), requires_grad=False)
layer.weight_scale_inv = Parameter(
w_scale.transpose(0, 1), requires_grad=False
)
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()}")
# Both paths produce weight [in, out] and weight_scale_inv [in//64, out,
# 2] as strided transpose views — DO NOT call .contiguous(). The matmul
# reduction loop scans the in-dim per output column; the [out, in]
# row-major source 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, making the inner-loop stride = out and
# tanking HBM bandwidth.
# 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
@@ -223,8 +245,8 @@ class NPUMXFP8LinearMethod(_NPULinearMethodBase):
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
qx, input_scale = torch.ops.npu.npu_dynamic_mx_quant(
x_2d, dst_type=torch.float8_e4m3fn
)
# MXFP8 matmul (weight & scale already transposed at load time)
@@ -240,13 +262,14 @@ class NPUMXFP8LinearMethod(_NPULinearMethodBase):
else:
quant_bias = bias.to(torch.float32)
output = torch_npu.npu_quant_matmul(
e8m0_dtype = _get_float8_e8m0fnu_dtype()
output = torch.ops.npu.npu_quant_matmul(
qx,
layer.weight,
layer.weight_scale_inv,
scale_dtype=_FLOAT8_E8M0FNU_DTYPE,
scale_dtype=e8m0_dtype,
pertoken_scale=input_scale,
pertoken_scale_dtype=_FLOAT8_E8M0FNU_DTYPE,
pertoken_scale_dtype=e8m0_dtype,
bias=quant_bias,
output_dtype=original_dtype,
group_sizes=[1, 1, MXFP8_BLOCK_SIZE],
@@ -2,27 +2,25 @@
Loads weights pre-quantized by msmodelslim (float8_e4m3fn weights,
uint8 scales) and runs MXFP8 matmul at inference.
Following the modelslim-scheme convention (see ModelSlimW8A8Int8), this scheme
owns only the hardware-agnostic weight creation; weight post-processing and the
forward pass are delegated to an NPUMXFP8LinearMethod kernel (self.kernel). Its
process_weights_after_loading detects the pre-quantized float8_e4m3fn weight and
takes the offline (transpose-only) branch.
"""
from typing import Dict, List, Optional
import torch
from sglang.srt.hardware_backend.npu.quantization.linear_method_npu import (
NPUMXFP8LinearMethod,
)
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):
@@ -36,6 +34,7 @@ class ModelSlimMXFP8Scheme(ModelSlimLinearScheme):
# dispatch signature used by ModelSlimConfig.get_linear_scheme;
# MXFP8 needs no per-layer config beyond what create_weights derives.
del quant_config, prefix
self.kernel = NPUMXFP8LinearMethod()
def create_weights(
self,
@@ -64,7 +63,8 @@ class ModelSlimMXFP8Scheme(ModelSlimLinearScheme):
# 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.
# checkpoint key exported by msmodelslim; the kernel re-layouts it into
# weight_scale_inv during process_weights_after_loading.
scale_dim = input_size_per_partition // MXFP8_BLOCK_SIZE
weight_scale = GroupQuantScaleParameter(
data=torch.empty(
@@ -78,24 +78,7 @@ class ModelSlimMXFP8Scheme(ModelSlimLinearScheme):
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
self.kernel.process_weights_after_loading(layer)
def apply_weights(
self,
@@ -103,44 +86,4 @@ class ModelSlimMXFP8Scheme(ModelSlimLinearScheme):
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)
return self.kernel.apply(layer, x, bias)