♻️ [llm][npu][quant] Delegate MXFP8 dense scheme to kernel and use torch.ops.npu (#28505)
This commit is contained in:
@@ -6,12 +6,6 @@ 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 npu_format_cast
|
||||||
from sglang.srt.layers.quantization.base_config import LinearMethodBase
|
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:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.layers.quantization.base_config import QuantizationConfig
|
from sglang.srt.layers.quantization.base_config import QuantizationConfig
|
||||||
@@ -19,11 +13,16 @@ if TYPE_CHECKING:
|
|||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
MXFP8_BLOCK_SIZE = 32
|
MXFP8_BLOCK_SIZE = 32
|
||||||
_FLOAT8_E8M0FNU_DTYPE = (
|
|
||||||
getattr(torch_npu, "float8_e8m0fnu", getattr(torch, "float8_e8m0fnu", None))
|
|
||||||
if _is_npu
|
# NPU ops are reached via torch.ops.npu.* (registered when torch_npu is imported
|
||||||
else getattr(torch, "float8_e8m0fnu", None)
|
# 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):
|
class _NPULinearMethodBase(LinearMethodBase):
|
||||||
@@ -131,8 +130,12 @@ class NPUW8A8Int8DynamicLinearMethod(_NPULinearMethodBase):
|
|||||||
class NPUMXFP8LinearMethod(_NPULinearMethodBase):
|
class NPUMXFP8LinearMethod(_NPULinearMethodBase):
|
||||||
"""Ascend NPU MXFP8 linear method for LLM (SRT) models.
|
"""Ascend NPU MXFP8 linear method for LLM (SRT) models.
|
||||||
|
|
||||||
Online mode: loads FP16/BF16 weights → quantises to MXFP8 at load time.
|
Shared kernel for both the online config path (``--quantization mxfp8``) and
|
||||||
Inference: dynamic MXFP8 activation quant + MXFP8 matmul (block_size=32).
|
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(
|
def create_weights(
|
||||||
@@ -169,33 +172,52 @@ class NPUMXFP8LinearMethod(_NPULinearMethodBase):
|
|||||||
layer.register_parameter("weight", weight)
|
layer.register_parameter("weight", weight)
|
||||||
|
|
||||||
def process_weights_after_loading(self, layer: torch.nn.Module) -> None:
|
def process_weights_after_loading(self, layer: torch.nn.Module) -> None:
|
||||||
weight_fp = layer.weight.data
|
weight = layer.weight.data
|
||||||
if weight_fp.dtype not in (torch.float16, torch.bfloat16):
|
if weight.dtype == torch.float8_e4m3fn:
|
||||||
logger.warning(
|
# Offline (ModelSlim) path: weight is already MXFP8-quantised and
|
||||||
"NPUMXFP8LinearMethod: weight dtype %s is not float16/bfloat16; "
|
# layer.weight_scale holds the uint8 block scales [out, in/32]. Only
|
||||||
"casting to bfloat16 before MXFP8 quantisation.",
|
# re-layout to [in, out] / [in//64, out, 2] strided views below.
|
||||||
weight_fp.dtype,
|
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)
|
# Both paths produce weight [in, out] and weight_scale_inv [in//64, out,
|
||||||
if not weight_fp.is_npu:
|
# 2] as strided transpose views — DO NOT call .contiguous(). The matmul
|
||||||
weight_fp = weight_fp.to(f"npu:{torch.npu.current_device()}")
|
# 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.
|
# Cache FP32 bias once to avoid a per-forward dtype conversion + alloc.
|
||||||
if (
|
if (
|
||||||
getattr(layer, "bias", None) is not None
|
getattr(layer, "bias", None) is not None
|
||||||
@@ -223,8 +245,8 @@ class NPUMXFP8LinearMethod(_NPULinearMethodBase):
|
|||||||
x_2d = x.reshape(-1, x.shape[-1])
|
x_2d = x.reshape(-1, x.shape[-1])
|
||||||
|
|
||||||
# Dynamic MXFP8 activation quantisation
|
# Dynamic MXFP8 activation quantisation
|
||||||
qx, input_scale = torch_npu.npu_dynamic_mx_quant(
|
qx, input_scale = torch.ops.npu.npu_dynamic_mx_quant(
|
||||||
x_2d, dst_type=torch_npu.float8_e4m3fn
|
x_2d, dst_type=torch.float8_e4m3fn
|
||||||
)
|
)
|
||||||
|
|
||||||
# MXFP8 matmul (weight & scale already transposed at load time)
|
# MXFP8 matmul (weight & scale already transposed at load time)
|
||||||
@@ -240,13 +262,14 @@ class NPUMXFP8LinearMethod(_NPULinearMethodBase):
|
|||||||
else:
|
else:
|
||||||
quant_bias = bias.to(torch.float32)
|
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,
|
qx,
|
||||||
layer.weight,
|
layer.weight,
|
||||||
layer.weight_scale_inv,
|
layer.weight_scale_inv,
|
||||||
scale_dtype=_FLOAT8_E8M0FNU_DTYPE,
|
scale_dtype=e8m0_dtype,
|
||||||
pertoken_scale=input_scale,
|
pertoken_scale=input_scale,
|
||||||
pertoken_scale_dtype=_FLOAT8_E8M0FNU_DTYPE,
|
pertoken_scale_dtype=e8m0_dtype,
|
||||||
bias=quant_bias,
|
bias=quant_bias,
|
||||||
output_dtype=original_dtype,
|
output_dtype=original_dtype,
|
||||||
group_sizes=[1, 1, MXFP8_BLOCK_SIZE],
|
group_sizes=[1, 1, MXFP8_BLOCK_SIZE],
|
||||||
|
|||||||
@@ -2,27 +2,25 @@
|
|||||||
|
|
||||||
Loads weights pre-quantized by msmodelslim (float8_e4m3fn weights,
|
Loads weights pre-quantized by msmodelslim (float8_e4m3fn weights,
|
||||||
uint8 scales) and runs MXFP8 matmul at inference.
|
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
|
from typing import Dict, List, Optional
|
||||||
|
|
||||||
import torch
|
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.parameter import GroupQuantScaleParameter, ModelWeightParameter
|
||||||
from sglang.srt.layers.quantization.modelslim.schemes import ModelSlimLinearScheme
|
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
|
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):
|
class ModelSlimMXFP8Scheme(ModelSlimLinearScheme):
|
||||||
@@ -36,6 +34,7 @@ class ModelSlimMXFP8Scheme(ModelSlimLinearScheme):
|
|||||||
# dispatch signature used by ModelSlimConfig.get_linear_scheme;
|
# dispatch signature used by ModelSlimConfig.get_linear_scheme;
|
||||||
# MXFP8 needs no per-layer config beyond what create_weights derives.
|
# MXFP8 needs no per-layer config beyond what create_weights derives.
|
||||||
del quant_config, prefix
|
del quant_config, prefix
|
||||||
|
self.kernel = NPUMXFP8LinearMethod()
|
||||||
|
|
||||||
def create_weights(
|
def create_weights(
|
||||||
self,
|
self,
|
||||||
@@ -64,7 +63,8 @@ class ModelSlimMXFP8Scheme(ModelSlimLinearScheme):
|
|||||||
|
|
||||||
# msmodelslim exports weight_scale as uint8, shape [out, in/32].
|
# msmodelslim exports weight_scale as uint8, shape [out, in/32].
|
||||||
# NOTE: Named "weight_scale" (not "weight_scale_inv") to match the
|
# 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
|
scale_dim = input_size_per_partition // MXFP8_BLOCK_SIZE
|
||||||
weight_scale = GroupQuantScaleParameter(
|
weight_scale = GroupQuantScaleParameter(
|
||||||
data=torch.empty(
|
data=torch.empty(
|
||||||
@@ -78,24 +78,7 @@ class ModelSlimMXFP8Scheme(ModelSlimLinearScheme):
|
|||||||
layer.register_parameter("weight_scale", weight_scale)
|
layer.register_parameter("weight_scale", weight_scale)
|
||||||
|
|
||||||
def process_weights_after_loading(self, layer: torch.nn.Module):
|
def process_weights_after_loading(self, layer: torch.nn.Module):
|
||||||
# Pre-transpose weight and scale to [in, out] for npu_quant_matmul.
|
self.kernel.process_weights_after_loading(layer)
|
||||||
# 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(
|
def apply_weights(
|
||||||
self,
|
self,
|
||||||
@@ -103,44 +86,4 @@ class ModelSlimMXFP8Scheme(ModelSlimLinearScheme):
|
|||||||
x: torch.Tensor,
|
x: torch.Tensor,
|
||||||
bias: Optional[torch.Tensor] = None,
|
bias: Optional[torch.Tensor] = None,
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
original_dtype = x.dtype
|
return self.kernel.apply(layer, x, bias)
|
||||||
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)
|
|
||||||
|
|||||||
Reference in New Issue
Block a user