[NPU] NPU quantization refactoring & more quantization formats support (#14504)
Co-authored-by: TamirBaydasov <mr.jeijy@gmail.com> Co-authored-by: Tamir Baydasov <41994229+TamirBaydasov@users.noreply.github.com> Co-authored-by: Савкин Артем <savkinartem@MacBook-Air-Viktoria.local> Co-authored-by: Edward Shogulin <edward.shogulin@gmail.com>
This commit is contained in:
@@ -7,9 +7,6 @@ import torch
|
||||
|
||||
from sglang.srt.compilation.piecewise_context_manager import is_in_piecewise_cuda_graph
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.hardware_backend.npu.quantization.fused_moe_method_npu import (
|
||||
NPUW4A16Int4DynamicMoEMethod,
|
||||
)
|
||||
from sglang.srt.layers import deep_gemm_wrapper
|
||||
from sglang.srt.layers.moe import (
|
||||
get_deepep_mode,
|
||||
@@ -27,6 +24,9 @@ from sglang.srt.layers.moe.token_dispatcher.deepep import (
|
||||
)
|
||||
from sglang.srt.layers.moe.topk import TopKOutput, TopKOutputChecker
|
||||
from sglang.srt.layers.quantization.base_config import QuantizationConfig
|
||||
from sglang.srt.layers.quantization.compressed_tensors.compressed_tensors_moe import (
|
||||
NPUCompressedTensorsW4A16Int4DynamicMoEMethod,
|
||||
)
|
||||
from sglang.srt.layers.quantization.fp8 import Fp8Config
|
||||
from sglang.srt.layers.quantization.fp8_kernel import is_fp8_fnuz
|
||||
from sglang.srt.layers.quantization.w4afp8 import W4AFp8Config, W4AFp8MoEMethod
|
||||
@@ -374,7 +374,7 @@ class DeepEPMoE(FusedMoE):
|
||||
else:
|
||||
input_quant = get_bool_env_var("DEEP_NORMAL_MODE_USE_INT8_QUANT")
|
||||
if not input_quant and not isinstance(
|
||||
self.quant_method, NPUW4A16Int4DynamicMoEMethod
|
||||
self.quant_method, NPUCompressedTensorsW4A16Int4DynamicMoEMethod
|
||||
):
|
||||
hidden_states, hidden_states_scale = torch_npu.npu_dynamic_quant(
|
||||
hidden_states
|
||||
|
||||
@@ -31,6 +31,7 @@ from sglang.srt.layers.quantization.modelopt_quant import (
|
||||
ModelOptFp4Config,
|
||||
ModelOptFp8Config,
|
||||
)
|
||||
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.petit import PetitNvFp4Config
|
||||
@@ -69,6 +70,7 @@ BASE_QUANTIZATION_METHODS: Dict[str, Type[QuantizationConfig]] = {
|
||||
"fbgemm_fp8": FBGEMMFp8Config,
|
||||
"quark": QuarkConfig,
|
||||
"auto-round": AutoRoundConfig,
|
||||
"modelslim": ModelSlimConfig,
|
||||
"quark_int4fp8_moe": QuarkInt4Fp8Config,
|
||||
}
|
||||
|
||||
@@ -80,15 +82,6 @@ if is_cuda() or (_is_mxfp_supported and is_hip()):
|
||||
}
|
||||
)
|
||||
|
||||
if is_npu():
|
||||
from sglang.srt.hardware_backend.npu.quantization.modelslim import ModelSlimConfig
|
||||
|
||||
BASE_QUANTIZATION_METHODS.update(
|
||||
{
|
||||
"modelslim": ModelSlimConfig,
|
||||
}
|
||||
)
|
||||
|
||||
QUANTIZATION_METHODS = {**BASE_QUANTIZATION_METHODS}
|
||||
|
||||
|
||||
|
||||
@@ -628,8 +628,8 @@ class AWQLinearAscendMethod(AWQLinearMethod):
|
||||
qzeros_tmp = -(qzeros_tmp - 8)
|
||||
qzeros_tmp = qzeros_tmp.to(layer.scales.data.dtype)
|
||||
|
||||
layer.qzeros = torch.nn.Parameter(qzeros_tmp, requires_grad=False)
|
||||
layer.qweight = torch.nn.Parameter(qweight_tmp, requires_grad=False)
|
||||
layer.zeros = torch.nn.Parameter(qzeros_tmp, requires_grad=False)
|
||||
layer.weight = torch.nn.Parameter(qweight_tmp, requires_grad=False)
|
||||
|
||||
def apply(
|
||||
self,
|
||||
@@ -637,9 +637,9 @@ class AWQLinearAscendMethod(AWQLinearMethod):
|
||||
x: torch.Tensor,
|
||||
bias: Optional[torch.Tensor] = None,
|
||||
) -> torch.Tensor:
|
||||
qweight = layer.qweight
|
||||
qweight = layer.weight
|
||||
scales = layer.scales
|
||||
qzeros = layer.qzeros
|
||||
qzeros = layer.zeros
|
||||
pack_factor = self.quant_config.pack_factor
|
||||
out_shape = x.shape[:-1] + (qweight.shape[-1] * pack_factor,)
|
||||
reshaped_x = x.reshape(-1, x.shape[-1])
|
||||
|
||||
@@ -17,7 +17,6 @@ if TYPE_CHECKING:
|
||||
class QuantizeMethodBase(ABC):
|
||||
"""Base class for different quantized methods."""
|
||||
|
||||
@abstractmethod
|
||||
def create_weights(
|
||||
self, layer: torch.nn.Module, *weight_args, **extra_weight_attrs
|
||||
):
|
||||
@@ -44,7 +43,6 @@ class QuantizeMethodBase(ABC):
|
||||
class LinearMethodBase(QuantizeMethodBase):
|
||||
"""Base class for different (maybe quantized) linear methods."""
|
||||
|
||||
@abstractmethod
|
||||
def create_weights(
|
||||
self,
|
||||
layer: torch.nn.Module,
|
||||
@@ -84,7 +82,6 @@ class LinearMethodBase(QuantizeMethodBase):
|
||||
|
||||
class FusedMoEMethodBase(QuantizeMethodBase):
|
||||
|
||||
@abstractmethod
|
||||
def create_weights(
|
||||
self,
|
||||
layer: torch.nn.Module,
|
||||
@@ -96,7 +93,6 @@ class FusedMoEMethodBase(QuantizeMethodBase):
|
||||
):
|
||||
raise NotImplementedError
|
||||
|
||||
@abstractmethod
|
||||
def create_moe_runner(
|
||||
self, layer: torch.nn.Module, moe_runner_config: MoeRunnerConfig
|
||||
):
|
||||
|
||||
@@ -45,6 +45,7 @@ from sglang.srt.layers.quantization.compressed_tensors.schemes import (
|
||||
CompressedTensorsW8A8Int8,
|
||||
CompressedTensorsW8A16Fp8,
|
||||
CompressedTensorsWNA16,
|
||||
NPUCompressedTensorsW8A8Int8,
|
||||
)
|
||||
from sglang.srt.layers.quantization.compressed_tensors.utils import (
|
||||
find_matched_target,
|
||||
@@ -53,6 +54,10 @@ from sglang.srt.layers.quantization.compressed_tensors.utils import (
|
||||
)
|
||||
from sglang.srt.layers.quantization.fp8 import Fp8LinearMethod
|
||||
from sglang.srt.layers.quantization.unquant import UnquantizedLinearMethod
|
||||
from sglang.srt.utils import is_cuda, is_npu
|
||||
|
||||
_is_cuda = is_cuda()
|
||||
_is_npu = is_npu()
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.models.utils import WeightsMapper
|
||||
@@ -305,6 +310,28 @@ class CompressedTensorsConfig(QuantizationConfig):
|
||||
else:
|
||||
return False
|
||||
|
||||
def _is_dynamic_token_w4a8(
|
||||
self, weight_quant: BaseModel, input_quant: BaseModel
|
||||
) -> bool:
|
||||
is_weight_4_bits = weight_quant.num_bits == 4
|
||||
is_activation_8_bits = input_quant.num_bits == 8
|
||||
weight_strategy = (
|
||||
weight_quant.strategy == QuantizationStrategy.GROUP.value
|
||||
or weight_quant.strategy == QuantizationStrategy.CHANNEL.value
|
||||
)
|
||||
is_token = (
|
||||
weight_strategy and input_quant.strategy == QuantizationStrategy.TOKEN.value
|
||||
)
|
||||
is_dynamic = not weight_quant.dynamic and input_quant.dynamic
|
||||
|
||||
return (
|
||||
is_weight_4_bits
|
||||
and is_activation_8_bits
|
||||
and is_token
|
||||
and weight_quant.symmetric
|
||||
and is_dynamic
|
||||
)
|
||||
|
||||
def _is_static_tensor_w8a8(
|
||||
self, weight_quant: BaseModel, input_quant: BaseModel
|
||||
) -> bool:
|
||||
@@ -444,6 +471,29 @@ class CompressedTensorsConfig(QuantizationConfig):
|
||||
|
||||
return is_channel_group and input_quant_none and is_symmetric and is_static
|
||||
|
||||
def _is_dynamic_token_w4(
|
||||
self, weight_quant: BaseModel, input_quant: BaseModel
|
||||
) -> bool:
|
||||
is_w4 = weight_quant.num_bits == 4
|
||||
weight_strategy = (
|
||||
weight_quant.strategy == QuantizationStrategy.TENSOR.value
|
||||
or weight_quant.strategy == QuantizationStrategy.CHANNEL.value
|
||||
or weight_quant.strategy == QuantizationStrategy.GROUP.value
|
||||
)
|
||||
if input_quant is not None:
|
||||
is_token = (
|
||||
weight_strategy
|
||||
and input_quant.strategy == QuantizationStrategy.TOKEN.value
|
||||
)
|
||||
is_dynamic = not weight_quant.dynamic and input_quant.dynamic
|
||||
else:
|
||||
is_token = weight_strategy
|
||||
is_dynamic = not weight_quant.dynamic
|
||||
|
||||
# Both symmetric and asymmetric input quantization supported.
|
||||
# Only symmetric weight quantization supported.
|
||||
return is_w4 and weight_quant.symmetric and is_token and is_dynamic
|
||||
|
||||
def _get_scheme_from_parts(
|
||||
self, weight_quant: BaseModel, input_quant: BaseModel
|
||||
) -> CompressedTensorsScheme:
|
||||
@@ -505,18 +555,32 @@ class CompressedTensorsConfig(QuantizationConfig):
|
||||
)
|
||||
|
||||
if self._is_static_tensor_w8a8(weight_quant, input_quant):
|
||||
return CompressedTensorsW8A8Int8(
|
||||
strategy=weight_quant.strategy,
|
||||
is_static_input_scheme=True,
|
||||
input_symmetric=input_quant.symmetric,
|
||||
)
|
||||
if not _is_npu:
|
||||
return CompressedTensorsW8A8Int8(
|
||||
strategy=weight_quant.strategy,
|
||||
is_static_input_scheme=True,
|
||||
input_symmetric=input_quant.symmetric,
|
||||
)
|
||||
else:
|
||||
return NPUCompressedTensorsW8A8Int8(
|
||||
strategy=weight_quant.strategy,
|
||||
is_static_input_scheme=True,
|
||||
input_symmetric=input_quant.symmetric,
|
||||
)
|
||||
|
||||
if self._is_dynamic_token_w8a8(weight_quant, input_quant):
|
||||
return CompressedTensorsW8A8Int8(
|
||||
strategy=weight_quant.strategy,
|
||||
is_static_input_scheme=False,
|
||||
input_symmetric=input_quant.symmetric,
|
||||
)
|
||||
if not _is_npu:
|
||||
return CompressedTensorsW8A8Int8(
|
||||
strategy=weight_quant.strategy,
|
||||
is_static_input_scheme=False,
|
||||
input_symmetric=input_quant.symmetric,
|
||||
)
|
||||
else:
|
||||
return NPUCompressedTensorsW8A8Int8(
|
||||
strategy=weight_quant.strategy,
|
||||
is_static_input_scheme=False,
|
||||
input_symmetric=input_quant.symmetric,
|
||||
)
|
||||
|
||||
raise NotImplementedError("No compressed-tensors compatible scheme was found.")
|
||||
|
||||
@@ -594,7 +658,9 @@ class CompressedTensorsConfig(QuantizationConfig):
|
||||
|
||||
# Raise error if device does not support the scheme
|
||||
# (e.g. fp8 needs ada lovelace)
|
||||
self._check_scheme_supported(scheme.get_min_capability())
|
||||
# Note: NPU devices do not support min_capability function
|
||||
if not _is_npu:
|
||||
self._check_scheme_supported(scheme.get_min_capability())
|
||||
logger.debug("Using scheme: %s for %s", scheme.__class__.__name__, layer_name)
|
||||
return scheme
|
||||
|
||||
|
||||
@@ -15,6 +15,11 @@ from sglang.srt.distributed import get_tensor_model_parallel_world_size, get_tp_
|
||||
from sglang.srt.distributed.device_communicators.pynccl_allocator import (
|
||||
use_symmetric_memory,
|
||||
)
|
||||
from sglang.srt.hardware_backend.npu.quantization.fused_moe_method_npu import (
|
||||
NPUW4A8Int8DynamicMoEMethod,
|
||||
NPUW4A16Int4DynamicMoEMethod,
|
||||
NPUW8A8Int8DynamicMoEMethod,
|
||||
)
|
||||
from sglang.srt.layers.dp_attention import is_allocation_symmetric
|
||||
from sglang.srt.layers.moe import MoeRunner, MoeRunnerBackend, MoeRunnerConfig
|
||||
from sglang.srt.layers.moe.cutlass_moe_params import CutlassMoEParams, CutlassMoEType
|
||||
@@ -43,6 +48,7 @@ from sglang.srt.utils import (
|
||||
get_bool_env_var,
|
||||
is_cuda,
|
||||
is_hip,
|
||||
is_npu,
|
||||
next_power_of_2,
|
||||
set_weight_attrs,
|
||||
)
|
||||
@@ -58,6 +64,7 @@ if TYPE_CHECKING:
|
||||
)
|
||||
|
||||
_is_hip = is_hip()
|
||||
_is_npu = is_npu()
|
||||
_is_cuda = is_cuda()
|
||||
|
||||
_use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip
|
||||
@@ -79,8 +86,11 @@ class GPTQMarlinState(Enum):
|
||||
__all__ = [
|
||||
"CompressedTensorsMoEMethod",
|
||||
"CompressedTensorsW4A4Nvfp4MoEMethod",
|
||||
"NPUCompressedTensorsW4A8Int8DynamicMoEMethod",
|
||||
"CompressedTensorsW8A8Fp8MoEMethod",
|
||||
"NPUCompressedTensorsW8A8Int8MoEMethod",
|
||||
"CompressedTensorsWNA16MoEMethod",
|
||||
"NPUCompressedTensorsW4A16Int4DynamicMoEMethod",
|
||||
]
|
||||
|
||||
|
||||
@@ -103,14 +113,40 @@ class CompressedTensorsMoEMethod(FusedMoEMethodBase):
|
||||
input_quant = quant_config.target_scheme_map["Linear"].get("input_activations")
|
||||
|
||||
if quant_config._is_wNa16_group_channel(weight_quant, input_quant):
|
||||
logger.info_once("Using CompressedTensorsWNA16MarlinMoEMethod")
|
||||
return CompressedTensorsWNA16MoEMethod(quant_config)
|
||||
if not _is_npu:
|
||||
logger.info_once("Using CompressedTensorsWNA16MarlinMoEMethod")
|
||||
return CompressedTensorsWNA16MoEMethod(quant_config)
|
||||
else:
|
||||
if (
|
||||
quant_config._is_dynamic_token_w4(weight_quant, input_quant)
|
||||
and input_quant is None
|
||||
):
|
||||
logger.info_once(
|
||||
"Using NPUCompressedTensorsW4A16Int4DynamicMoEMethod"
|
||||
)
|
||||
return NPUCompressedTensorsW4A16Int4DynamicMoEMethod(quant_config)
|
||||
elif quant_config._is_fp4a4_nvfp4(weight_quant, input_quant):
|
||||
logger.info_once("Using CompressedTensorsW4A4Nvfp4MoEMethod")
|
||||
return CompressedTensorsW4A4Nvfp4MoEMethod(quant_config)
|
||||
elif quant_config._is_fp8_w8a8(weight_quant, input_quant):
|
||||
logger.info_once("Using CompressedTensorsW8A8Fp8MoEMethod")
|
||||
return CompressedTensorsW8A8Fp8MoEMethod(quant_config)
|
||||
elif quant_config._is_dynamic_token_w8a8(weight_quant, input_quant):
|
||||
if _is_npu:
|
||||
logger.info_once("Using NPUCompressedTensorsW8A8Int8DynamicMoEMethod")
|
||||
return NPUCompressedTensorsW8A8Int8DynamicMoEMethod(quant_config)
|
||||
else:
|
||||
raise NotImplementedError(
|
||||
f"The W8A8Int8 Fused MoE scheme is implemented only for NPU for now."
|
||||
)
|
||||
elif quant_config._is_dynamic_token_w4a8(weight_quant, input_quant):
|
||||
if _is_npu:
|
||||
logger.info_once("Using NPUCompressedTensorsW4A8Int8DynamicMoEMethod")
|
||||
return NPUCompressedTensorsW4A8Int8DynamicMoEMethod(quant_config)
|
||||
else:
|
||||
raise NotImplementedError(
|
||||
f"The W4A8Int8 Fused MoE scheme is implemented only for NPU for now."
|
||||
)
|
||||
else:
|
||||
raise RuntimeError(
|
||||
f"Unsupported FusedMoe scheme: {weight_quant}, {input_quant}"
|
||||
@@ -853,6 +889,117 @@ class CompressedTensorsW8A8Fp8MoEMethod(CompressedTensorsMoEMethod):
|
||||
return self.runner.run(dispatch_output, quant_info)
|
||||
|
||||
|
||||
class NPUCompressedTensorsW8A8Int8DynamicMoEMethod(CompressedTensorsMoEMethod):
|
||||
|
||||
def __init__(self, quant_config: CompressedTensorsConfig):
|
||||
self.quant_config = quant_config
|
||||
self.weight_quant = self.quant_config.target_scheme_map["Linear"].get("weights")
|
||||
self.input_quant = self.quant_config.target_scheme_map["Linear"].get(
|
||||
"input_activations"
|
||||
)
|
||||
self.kernel = NPUW8A8Int8DynamicMoEMethod()
|
||||
|
||||
self.static_input_scales = not self.input_quant.dynamic
|
||||
per_channel = (
|
||||
self.weight_quant.strategy == QuantizationStrategy.CHANNEL
|
||||
and self.input_quant.strategy == QuantizationStrategy.TOKEN
|
||||
)
|
||||
if not per_channel:
|
||||
raise ValueError(
|
||||
"For INT8 Fused MoE layers, we require channelwise, "
|
||||
"dynamic per token quantization. Found "
|
||||
f"{self.weight_quant}, {self.input_quant}"
|
||||
)
|
||||
|
||||
self.static_input_scales = not self.input_quant.dynamic
|
||||
if self.static_input_scales:
|
||||
raise ValueError(
|
||||
"For INT8 Fused MoE layers, we require channelwise, "
|
||||
"dynamic per token quantization. Found static input scales."
|
||||
)
|
||||
|
||||
def create_weights(
|
||||
self,
|
||||
layer: torch.nn.Module,
|
||||
num_experts: int,
|
||||
hidden_size: int,
|
||||
intermediate_size_per_partition: int,
|
||||
params_dtype: torch.dtype,
|
||||
**extra_weight_attrs,
|
||||
):
|
||||
|
||||
from sglang.srt.layers.moe.fused_moe_triton import FusedMoeWeightScaleSupported
|
||||
|
||||
params_dtype = torch.int8
|
||||
|
||||
# WEIGHTS
|
||||
w13_weight = torch.nn.Parameter(
|
||||
torch.empty(
|
||||
num_experts,
|
||||
2 * intermediate_size_per_partition,
|
||||
hidden_size,
|
||||
dtype=params_dtype,
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w13_weight", w13_weight)
|
||||
set_weight_attrs(w13_weight, extra_weight_attrs)
|
||||
|
||||
w2_weight = torch.nn.Parameter(
|
||||
torch.empty(
|
||||
num_experts,
|
||||
hidden_size,
|
||||
intermediate_size_per_partition,
|
||||
dtype=params_dtype,
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w2_weight", w2_weight)
|
||||
set_weight_attrs(w2_weight, extra_weight_attrs)
|
||||
|
||||
# WEIGHT_SCALES
|
||||
assert self.weight_quant.strategy == QuantizationStrategy.CHANNEL
|
||||
w13_weight_scale = torch.nn.Parameter(
|
||||
torch.ones(
|
||||
num_experts, 2 * intermediate_size_per_partition, 1, dtype=torch.float32
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w13_weight_scale", w13_weight_scale)
|
||||
w2_weight_scale = torch.nn.Parameter(
|
||||
torch.ones(num_experts, hidden_size, 1, dtype=torch.float32),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w2_weight_scale", w2_weight_scale)
|
||||
# Add PER-CHANNEL quantization for FusedMoE.weight_loader.
|
||||
extra_weight_attrs.update(
|
||||
{"quant_method": FusedMoeWeightScaleSupported.CHANNEL.value}
|
||||
)
|
||||
set_weight_attrs(w13_weight_scale, extra_weight_attrs)
|
||||
set_weight_attrs(w2_weight_scale, extra_weight_attrs)
|
||||
|
||||
# INPUT_SCALES
|
||||
assert not self.static_input_scales
|
||||
layer.w13_input_scale = None
|
||||
layer.w2_input_scale = None
|
||||
|
||||
def process_weights_after_loading(self, layer: torch.nn.Module) -> None:
|
||||
self.kernel.process_weights_after_loading(layer)
|
||||
|
||||
def create_moe_runner(
|
||||
self, layer: torch.nn.Module, moe_runner_config: MoeRunnerConfig
|
||||
):
|
||||
self.moe_runner_config = moe_runner_config
|
||||
|
||||
def apply(
|
||||
self,
|
||||
layer: torch.nn.Module,
|
||||
dispatch_output: StandardDispatchOutput,
|
||||
) -> CombineInput:
|
||||
|
||||
return self.kernel.apply(layer, dispatch_output)
|
||||
|
||||
|
||||
class CompressedTensorsWNA16MoEMethod(CompressedTensorsMoEMethod):
|
||||
|
||||
def __init__(self, quant_config: CompressedTensorsConfig, num_gpu_experts=-1):
|
||||
@@ -1200,3 +1347,421 @@ class CompressedTensorsWNA16MoEMethod(CompressedTensorsMoEMethod):
|
||||
routed_scaling_factor=self.moe_runner_config.routed_scaling_factor,
|
||||
)
|
||||
return StandardCombineInput(hidden_states=output)
|
||||
|
||||
|
||||
class NPUCompressedTensorsW4A8Int8DynamicMoEMethod(CompressedTensorsMoEMethod):
|
||||
|
||||
### TODO: Get rid of code duplication with python/sglang/srt/modelslim/modelslim_moe.py @OrangeRedeng @TamirBaydasov
|
||||
def __init__(self, quantization_config) -> None:
|
||||
self.group_size = 0
|
||||
self.is_per_channel_weight = self.group_size == 0
|
||||
self.tp_size = 1
|
||||
self.activation_use_clip = (
|
||||
self.quantization_config.get("config_groups", {})
|
||||
.get("group_1", {})
|
||||
.get("activation_use_clip", False)
|
||||
)
|
||||
self.kernel = NPUW4A8Int8DynamicMoEMethod()
|
||||
|
||||
def create_weights(
|
||||
self,
|
||||
layer: torch.nn.Module,
|
||||
num_experts: int,
|
||||
hidden_size: int,
|
||||
intermediate_size_per_partition: int,
|
||||
params_dtype: torch.dtype,
|
||||
**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.CHANNEL.value}
|
||||
)
|
||||
|
||||
# >> weight
|
||||
w13_output_size = intermediate_size_per_partition
|
||||
w2_output_size = hidden_size // 2
|
||||
w13_weight = torch.nn.Parameter(
|
||||
torch.empty(num_experts, w13_output_size, hidden_size, dtype=torch.int8),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w13_weight", w13_weight)
|
||||
set_weight_attrs(w13_weight, extra_weight_attrs)
|
||||
w2_weight = torch.nn.Parameter(
|
||||
torch.empty(
|
||||
num_experts,
|
||||
w2_output_size,
|
||||
intermediate_size_per_partition,
|
||||
dtype=torch.int8,
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w2_weight", w2_weight)
|
||||
set_weight_attrs(w2_weight, extra_weight_attrs)
|
||||
|
||||
# >> scale
|
||||
weight_scale_dtype = torch.int64 if self.activation_use_clip else torch.float32
|
||||
w13_weight_scale = torch.nn.Parameter(
|
||||
torch.empty(
|
||||
num_experts,
|
||||
2 * intermediate_size_per_partition,
|
||||
1,
|
||||
dtype=weight_scale_dtype,
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w13_weight_scale", w13_weight_scale)
|
||||
set_weight_attrs(w13_weight_scale, extra_weight_attrs)
|
||||
|
||||
w2_weight_scale = torch.nn.Parameter(
|
||||
torch.empty(num_experts, hidden_size, 1, dtype=weight_scale_dtype),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w2_weight_scale", w2_weight_scale)
|
||||
set_weight_attrs(w2_weight_scale, extra_weight_attrs)
|
||||
|
||||
# >> offset
|
||||
w13_weight_offset = torch.nn.Parameter(
|
||||
torch.empty(
|
||||
num_experts, 2 * intermediate_size_per_partition, 1, dtype=torch.float32
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w13_weight_offset", w13_weight_offset)
|
||||
set_weight_attrs(w13_weight_offset, extra_weight_attrs)
|
||||
|
||||
w2_weight_offset = torch.nn.Parameter(
|
||||
torch.empty(num_experts, hidden_size, 1, dtype=torch.float32),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w2_weight_offset", w2_weight_offset)
|
||||
set_weight_attrs(w2_weight_offset, extra_weight_attrs)
|
||||
|
||||
# >>> special param for w4a8
|
||||
if self.activation_use_clip:
|
||||
self._init_activation_clip_params(
|
||||
layer,
|
||||
num_experts,
|
||||
hidden_size,
|
||||
intermediate_size_per_partition,
|
||||
extra_weight_attrs,
|
||||
)
|
||||
else:
|
||||
self._init_extra_scale_params(
|
||||
layer,
|
||||
num_experts,
|
||||
hidden_size,
|
||||
intermediate_size_per_partition,
|
||||
extra_weight_attrs,
|
||||
)
|
||||
|
||||
def _init_activation_clip_params(
|
||||
self,
|
||||
layer: torch.nn.Module,
|
||||
num_experts: int,
|
||||
hidden_size: int,
|
||||
intermediate_size_per_partition: int,
|
||||
extra_weight_attrs: dict,
|
||||
) -> None:
|
||||
"""
|
||||
Initializes bias and alpha parameters for quantization schemes that use activation clipping.
|
||||
|
||||
This helper registers `w13_bias`, `w2_bias`, and `w2_alpha`, which are required to
|
||||
shift and scale the activations or outputs to compensate for the precision loss
|
||||
introduced by clamping activations.
|
||||
"""
|
||||
w13_bias = torch.nn.Parameter(
|
||||
torch.ones(
|
||||
num_experts, 2 * intermediate_size_per_partition, dtype=torch.float
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w13_bias", w13_bias)
|
||||
set_weight_attrs(w13_bias, extra_weight_attrs)
|
||||
|
||||
w2_bias = torch.nn.Parameter(
|
||||
torch.ones(num_experts, hidden_size, dtype=torch.float),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w2_bias", w2_bias)
|
||||
set_weight_attrs(w2_bias, extra_weight_attrs)
|
||||
|
||||
w2_alpha = torch.nn.Parameter(
|
||||
torch.ones(num_experts, dtype=torch.float), requires_grad=False
|
||||
)
|
||||
layer.register_parameter("w2_alpha", w2_alpha)
|
||||
set_weight_attrs(w2_alpha, extra_weight_attrs)
|
||||
|
||||
def _init_extra_scale_params(
|
||||
self,
|
||||
layer: torch.nn.Module,
|
||||
num_experts: int,
|
||||
hidden_size: int,
|
||||
intermediate_size_per_partition: int,
|
||||
extra_weight_attrs: dict,
|
||||
) -> None:
|
||||
"""
|
||||
Initializes additional scaling, offset, and bias parameters for quantization schemes without activation clipping.
|
||||
|
||||
This method registers the following parameters:
|
||||
1. Scale Biases: `w13_scale_bias` and `w2_scale_bias`.
|
||||
2. Secondary Quantization Params (initialized only for grouped quantization):
|
||||
`w13_weight_scale_second`, `w13_weight_offset_second`,
|
||||
`w2_weight_scale_second`, and `w2_weight_offset_second`.
|
||||
"""
|
||||
if not self.is_per_channel_weight:
|
||||
w13_weight_scale_second = torch.nn.Parameter(
|
||||
torch.empty(
|
||||
num_experts,
|
||||
2 * intermediate_size_per_partition,
|
||||
hidden_size // self.group_size,
|
||||
dtype=torch.float32,
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w13_weight_scale_second", w13_weight_scale_second)
|
||||
set_weight_attrs(w13_weight_scale_second, extra_weight_attrs)
|
||||
|
||||
w13_weight_offset_second = torch.nn.Parameter(
|
||||
torch.empty(
|
||||
num_experts,
|
||||
2 * intermediate_size_per_partition,
|
||||
hidden_size // self.group_size,
|
||||
dtype=torch.float32,
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter(
|
||||
"w13_weight_offset_second", w13_weight_offset_second
|
||||
)
|
||||
set_weight_attrs(w13_weight_offset_second, extra_weight_attrs)
|
||||
|
||||
w2_weight_scale_second = torch.nn.Parameter(
|
||||
torch.empty(
|
||||
num_experts,
|
||||
hidden_size,
|
||||
intermediate_size_per_partition // self.group_size,
|
||||
dtype=torch.float32,
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w2_weight_scale_second", w2_weight_scale_second)
|
||||
set_weight_attrs(w2_weight_scale_second, extra_weight_attrs)
|
||||
|
||||
w2_weight_offset_second = torch.nn.Parameter(
|
||||
torch.empty(
|
||||
num_experts,
|
||||
hidden_size,
|
||||
intermediate_size_per_partition // self.group_size,
|
||||
dtype=torch.float32,
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w2_weight_offset_second", w2_weight_offset_second)
|
||||
set_weight_attrs(w2_weight_offset_second, extra_weight_attrs)
|
||||
|
||||
w13_scale_bias = torch.nn.Parameter(
|
||||
torch.empty(
|
||||
num_experts, 2 * intermediate_size_per_partition, 1, dtype=torch.float32
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w13_scale_bias", w13_scale_bias)
|
||||
set_weight_attrs(w13_scale_bias, extra_weight_attrs)
|
||||
|
||||
w2_scale_bias = torch.nn.Parameter(
|
||||
torch.empty(
|
||||
num_experts, hidden_size, 16 // self.tp_size, dtype=torch.float32
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w2_scale_bias", w2_scale_bias)
|
||||
set_weight_attrs(w2_scale_bias, extra_weight_attrs)
|
||||
|
||||
def process_weights_after_loading(self, layer: torch.nn.Module) -> None:
|
||||
self.kernel.process_weights_after_loading(
|
||||
layer, self.is_per_channel_weight, self.activation_use_clip
|
||||
)
|
||||
|
||||
def create_moe_runner(
|
||||
self, layer: torch.nn.Module, moe_runner_config: MoeRunnerConfig
|
||||
):
|
||||
self.moe_runner_config = moe_runner_config
|
||||
|
||||
def apply(
|
||||
self,
|
||||
layer: torch.nn.Module,
|
||||
dispatch_output: StandardDispatchOutput,
|
||||
) -> CombineInput:
|
||||
|
||||
return self.kernel.apply(layer, dispatch_output)
|
||||
|
||||
def apply_without_routing_weights(
|
||||
self,
|
||||
layer,
|
||||
hidden_states,
|
||||
hidden_states_scale,
|
||||
group_list_type,
|
||||
group_list,
|
||||
output_dtype,
|
||||
):
|
||||
return self.kernel.apply_without_routing_weights(
|
||||
layer,
|
||||
hidden_states,
|
||||
hidden_states_scale,
|
||||
group_list_type,
|
||||
group_list,
|
||||
output_dtype,
|
||||
)
|
||||
|
||||
|
||||
class NPUCompressedTensorsW4A16Int4DynamicMoEMethod(CompressedTensorsMoEMethod):
|
||||
|
||||
def __init__(self, quantization_config) -> None:
|
||||
self.pack_factor = 8 # weight dtype is int4, but use int32 to create
|
||||
target = (
|
||||
"MoEGMM" if "MoEGMM" in quantization_config.target_scheme_map else "Linear"
|
||||
)
|
||||
if target in quantization_config.target_scheme_map:
|
||||
self.group_size = quantization_config.target_scheme_map[target][
|
||||
"weights"
|
||||
].group_size
|
||||
else:
|
||||
self.group_size = 128
|
||||
|
||||
self.kernel = NPUW4A16Int4DynamicMoEMethod()
|
||||
|
||||
# TODO: See if we can merge this method's logic
|
||||
# with CompressedTensorsWNA16MoEMethod. Need more models and tests.
|
||||
# @OrangeRedeng @TamirBaydasov
|
||||
def create_weights(
|
||||
self,
|
||||
layer: torch.nn.Module,
|
||||
num_experts: int,
|
||||
hidden_size: int,
|
||||
intermediate_size_per_partition: int,
|
||||
params_dtype: torch.dtype,
|
||||
**extra_weight_attrs,
|
||||
) -> None:
|
||||
from sglang.srt.layers.moe.fused_moe_triton import FusedMoeWeightScaleSupported
|
||||
|
||||
self.num_experts = num_experts
|
||||
if (
|
||||
extra_weight_attrs.get(
|
||||
"intermediate_size_full", intermediate_size_per_partition
|
||||
)
|
||||
// intermediate_size_per_partition
|
||||
> 1
|
||||
):
|
||||
quant_method = FusedMoeWeightScaleSupported.GROUP.value
|
||||
else:
|
||||
quant_method = FusedMoeWeightScaleSupported.CHANNEL.value
|
||||
extra_weight_attrs.update({"quant_method": quant_method})
|
||||
# weight
|
||||
w13_weight = torch.nn.Parameter(
|
||||
torch.empty(
|
||||
num_experts,
|
||||
2 * intermediate_size_per_partition,
|
||||
hidden_size // self.pack_factor,
|
||||
dtype=torch.int32,
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w13_weight", w13_weight)
|
||||
set_weight_attrs(w13_weight, extra_weight_attrs)
|
||||
w2_weight = torch.nn.Parameter(
|
||||
torch.empty(
|
||||
num_experts,
|
||||
hidden_size,
|
||||
intermediate_size_per_partition // self.pack_factor,
|
||||
dtype=torch.int32,
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w2_weight", w2_weight)
|
||||
set_weight_attrs(w2_weight, extra_weight_attrs)
|
||||
|
||||
# scale
|
||||
weight_scale_dtype = torch.bfloat16
|
||||
w13_weight_scale = torch.nn.Parameter(
|
||||
torch.empty(
|
||||
num_experts,
|
||||
2 * intermediate_size_per_partition,
|
||||
hidden_size // self.group_size,
|
||||
dtype=weight_scale_dtype,
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w13_weight_scale", w13_weight_scale)
|
||||
set_weight_attrs(w13_weight_scale, extra_weight_attrs)
|
||||
w2_weight_scale = torch.nn.Parameter(
|
||||
torch.empty(
|
||||
num_experts,
|
||||
hidden_size,
|
||||
intermediate_size_per_partition // self.group_size,
|
||||
dtype=weight_scale_dtype,
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w2_weight_scale", w2_weight_scale)
|
||||
set_weight_attrs(w2_weight_scale, extra_weight_attrs)
|
||||
|
||||
# offset
|
||||
w13_weight_offset = torch.nn.Parameter(
|
||||
torch.zeros(
|
||||
num_experts,
|
||||
2 * intermediate_size_per_partition,
|
||||
hidden_size // self.group_size,
|
||||
dtype=weight_scale_dtype,
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w13_weight_offset", w13_weight_offset)
|
||||
set_weight_attrs(w13_weight_offset, extra_weight_attrs)
|
||||
|
||||
w2_weight_offset = torch.nn.Parameter(
|
||||
torch.zeros(
|
||||
num_experts,
|
||||
hidden_size,
|
||||
intermediate_size_per_partition // self.group_size,
|
||||
dtype=weight_scale_dtype,
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w2_weight_offset", w2_weight_offset)
|
||||
set_weight_attrs(w2_weight_offset, extra_weight_attrs)
|
||||
|
||||
def process_weights_after_loading(self, layer: torch.nn.Module) -> None:
|
||||
self.kernel.process_weights_after_loading(layer)
|
||||
|
||||
def create_moe_runner(
|
||||
self, layer: torch.nn.Module, moe_runner_config: MoeRunnerConfig
|
||||
):
|
||||
self.moe_runner_config = moe_runner_config
|
||||
|
||||
def apply(
|
||||
self,
|
||||
layer: torch.nn.Module,
|
||||
dispatch_output: StandardDispatchOutput,
|
||||
) -> CombineInput:
|
||||
|
||||
return self.kernel.apply(layer, dispatch_output)
|
||||
|
||||
def apply_without_routing_weights(
|
||||
self,
|
||||
layer,
|
||||
hidden_states,
|
||||
hidden_states_scale,
|
||||
group_list_type,
|
||||
group_list,
|
||||
output_dtype,
|
||||
):
|
||||
return self.kernel.apply_without_routing_weights(
|
||||
layer,
|
||||
hidden_states,
|
||||
hidden_states_scale,
|
||||
group_list_type,
|
||||
group_list,
|
||||
output_dtype,
|
||||
)
|
||||
|
||||
@@ -3,7 +3,10 @@
|
||||
from .compressed_tensors_scheme import CompressedTensorsScheme
|
||||
from .compressed_tensors_w4a4_nvfp4 import CompressedTensorsW4A4Fp4
|
||||
from .compressed_tensors_w8a8_fp8 import CompressedTensorsW8A8Fp8
|
||||
from .compressed_tensors_w8a8_int8 import CompressedTensorsW8A8Int8
|
||||
from .compressed_tensors_w8a8_int8 import (
|
||||
CompressedTensorsW8A8Int8,
|
||||
NPUCompressedTensorsW8A8Int8,
|
||||
)
|
||||
from .compressed_tensors_w8a16_fp8 import CompressedTensorsW8A16Fp8
|
||||
from .compressed_tensors_wNa16 import WNA16_SUPPORTED_BITS, CompressedTensorsWNA16
|
||||
|
||||
@@ -12,6 +15,7 @@ __all__ = [
|
||||
"CompressedTensorsW8A8Fp8",
|
||||
"CompressedTensorsW8A16Fp8",
|
||||
"CompressedTensorsW8A8Int8",
|
||||
"NPUCompressedTensorsW8A8Int8",
|
||||
"CompressedTensorsWNA16",
|
||||
"WNA16_SUPPORTED_BITS",
|
||||
"CompressedTensorsW4A4Fp4",
|
||||
|
||||
@@ -7,6 +7,9 @@ import torch
|
||||
from compressed_tensors.quantization import QuantizationStrategy
|
||||
from torch.nn import Parameter
|
||||
|
||||
from sglang.srt.hardware_backend.npu.quantization.linear_method_npu import (
|
||||
NPUW8A8Int8DynamicLinearMethod,
|
||||
)
|
||||
from sglang.srt.layers.parameter import (
|
||||
ChannelQuantScaleParameter,
|
||||
ModelWeightParameter,
|
||||
@@ -33,6 +36,61 @@ class CompressedTensorsW8A8Int8(CompressedTensorsScheme):
|
||||
self.is_static_input_scheme = is_static_input_scheme
|
||||
self.input_symmetric = input_symmetric
|
||||
|
||||
def create_weights(
|
||||
self,
|
||||
layer: torch.nn.Module,
|
||||
output_partition_sizes: list[int],
|
||||
input_size_per_partition: int,
|
||||
params_dtype: torch.dtype,
|
||||
weight_loader: Callable,
|
||||
**kwargs,
|
||||
):
|
||||
output_size_per_partition = sum(output_partition_sizes)
|
||||
layer.logical_widths = output_partition_sizes
|
||||
|
||||
# WEIGHT
|
||||
weight = ModelWeightParameter(
|
||||
data=torch.empty(
|
||||
output_size_per_partition, input_size_per_partition, dtype=torch.int8
|
||||
),
|
||||
input_dim=1,
|
||||
output_dim=0,
|
||||
weight_loader=weight_loader,
|
||||
)
|
||||
|
||||
layer.register_parameter("weight", weight)
|
||||
|
||||
# WEIGHT SCALE
|
||||
if self.strategy == QuantizationStrategy.CHANNEL:
|
||||
weight_scale = ChannelQuantScaleParameter(
|
||||
data=torch.empty((sum(output_partition_sizes), 1), dtype=torch.float32),
|
||||
output_dim=0,
|
||||
weight_loader=weight_loader,
|
||||
)
|
||||
else:
|
||||
assert self.strategy == QuantizationStrategy.TENSOR
|
||||
weight_scale = PerTensorScaleParameter(
|
||||
data=torch.empty(len(output_partition_sizes), dtype=torch.float32),
|
||||
weight_loader=weight_loader,
|
||||
)
|
||||
layer.register_parameter("weight_scale", weight_scale)
|
||||
|
||||
# INPUT SCALE
|
||||
if self.is_static_input_scheme:
|
||||
input_scale = PerTensorScaleParameter(
|
||||
data=torch.empty(1, dtype=torch.float32), weight_loader=weight_loader
|
||||
)
|
||||
layer.register_parameter("input_scale", input_scale)
|
||||
|
||||
if not self.input_symmetric:
|
||||
# Note: compressed-tensors stores the zp using the same dtype
|
||||
# as the weights
|
||||
# AZP loaded as int8 but used as int32
|
||||
input_zero_point = PerTensorScaleParameter(
|
||||
data=torch.empty(1, dtype=torch.int8), weight_loader=weight_loader
|
||||
)
|
||||
layer.register_parameter("input_zero_point", input_zero_point)
|
||||
|
||||
@classmethod
|
||||
def get_min_capability(cls) -> int:
|
||||
# ampere and up
|
||||
@@ -107,61 +165,6 @@ class CompressedTensorsW8A8Int8(CompressedTensorsScheme):
|
||||
else:
|
||||
layer.azp_adj = None
|
||||
|
||||
def create_weights(
|
||||
self,
|
||||
layer: torch.nn.Module,
|
||||
output_partition_sizes: list[int],
|
||||
input_size_per_partition: int,
|
||||
params_dtype: torch.dtype,
|
||||
weight_loader: Callable,
|
||||
**kwargs,
|
||||
):
|
||||
output_size_per_partition = sum(output_partition_sizes)
|
||||
layer.logical_widths = output_partition_sizes
|
||||
|
||||
# WEIGHT
|
||||
weight = ModelWeightParameter(
|
||||
data=torch.empty(
|
||||
output_size_per_partition, input_size_per_partition, dtype=torch.int8
|
||||
),
|
||||
input_dim=1,
|
||||
output_dim=0,
|
||||
weight_loader=weight_loader,
|
||||
)
|
||||
|
||||
layer.register_parameter("weight", weight)
|
||||
|
||||
# WEIGHT SCALE
|
||||
if self.strategy == QuantizationStrategy.CHANNEL:
|
||||
weight_scale = ChannelQuantScaleParameter(
|
||||
data=torch.empty((sum(output_partition_sizes), 1), dtype=torch.float32),
|
||||
output_dim=0,
|
||||
weight_loader=weight_loader,
|
||||
)
|
||||
else:
|
||||
assert self.strategy == QuantizationStrategy.TENSOR
|
||||
weight_scale = PerTensorScaleParameter(
|
||||
data=torch.empty(len(output_partition_sizes), dtype=torch.float32),
|
||||
weight_loader=weight_loader,
|
||||
)
|
||||
layer.register_parameter("weight_scale", weight_scale)
|
||||
|
||||
# INPUT SCALE
|
||||
if self.is_static_input_scheme:
|
||||
input_scale = PerTensorScaleParameter(
|
||||
data=torch.empty(1, dtype=torch.float32), weight_loader=weight_loader
|
||||
)
|
||||
layer.register_parameter("input_scale", input_scale)
|
||||
|
||||
if not self.input_symmetric:
|
||||
# Note: compressed-tensors stores the zp using the same dtype
|
||||
# as the weights
|
||||
# AZP loaded as int8 but used as int32
|
||||
input_zero_point = PerTensorScaleParameter(
|
||||
data=torch.empty(1, dtype=torch.int8), weight_loader=weight_loader
|
||||
)
|
||||
layer.register_parameter("input_zero_point", input_zero_point)
|
||||
|
||||
def apply_weights(
|
||||
self, layer: torch.nn.Module, x: torch.Tensor, bias: Optional[torch.Tensor]
|
||||
) -> torch.Tensor:
|
||||
@@ -171,3 +174,28 @@ class CompressedTensorsW8A8Int8(CompressedTensorsScheme):
|
||||
return int8_scaled_mm(
|
||||
x_q, layer.weight, x_scale, layer.weight_scale, out_dtype=x.dtype, bias=bias
|
||||
)
|
||||
|
||||
|
||||
class NPUCompressedTensorsW8A8Int8(CompressedTensorsW8A8Int8):
|
||||
|
||||
def __init__(
|
||||
self, strategy: str, is_static_input_scheme: bool, input_symmetric: bool
|
||||
):
|
||||
super().__init__(strategy, is_static_input_scheme, input_symmetric)
|
||||
# TODO: Currently, NPU kernel for static quant requires quant_bias field,
|
||||
# which can't be replicated in compressed-tensors.
|
||||
if self.is_static_input_scheme:
|
||||
raise NotImplementedError(
|
||||
"Static compressed-tensors scheme is not yet supported on NPU."
|
||||
)
|
||||
self.kernel = NPUW8A8Int8DynamicLinearMethod()
|
||||
|
||||
@classmethod
|
||||
def get_min_capability(cls) -> int:
|
||||
return NotImplementedError
|
||||
|
||||
def process_weights_after_loading(self, layer):
|
||||
return self.kernel.process_weights_after_loading(layer)
|
||||
|
||||
def apply_weights(self, layer, x, bias):
|
||||
return self.kernel.apply(layer, x, bias)
|
||||
|
||||
14
python/sglang/srt/layers/quantization/modelslim/README.md
Normal file
14
python/sglang/srt/layers/quantization/modelslim/README.md
Normal file
@@ -0,0 +1,14 @@
|
||||
Quantization [ModelSlim](https://gitcode.com/Ascend/msit) module.
|
||||
|
||||
`--quantization modelslim` flag introduced. To load already quantized models, simply load the model weights. For models quantized with ModelSlim, there's no need to add `--quantization modelslim` argument when starting the engine. The quantization method will be automatically parsed from the downloaded `quant_model_description.json` config.
|
||||
|
||||
ModelSlim was developed in the format of compressed_tensors and includes support for various quantization schemes, such as:
|
||||
- [x] W4A4 dynamic linear
|
||||
- [x] W8A8 static linear
|
||||
- [x] W8A8 dynamic linear
|
||||
- [x] W4A8 dynamic MOE
|
||||
- [x] W8A8 dynamic MOE
|
||||
|
||||
Also ModelSlim module include:
|
||||
- [x] Automated config detection for modelslim format (without the need to specify --quantization modelslim flag)
|
||||
- [x] Unit-tests for w4a4 modelslim, w8a8 modelslim
|
||||
273
python/sglang/srt/layers/quantization/modelslim/modelslim.py
Normal file
273
python/sglang/srt/layers/quantization/modelslim/modelslim.py
Normal file
@@ -0,0 +1,273 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from types import MappingProxyType
|
||||
from typing import Any, Dict, List, Mapping, Optional, Tuple, Union, cast
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.hardware_backend.npu.quantization.linear_method_npu import (
|
||||
_NPULinearMethodBase,
|
||||
)
|
||||
from sglang.srt.layers.quantization.base_config import (
|
||||
QuantizationConfig,
|
||||
QuantizeMethodBase,
|
||||
)
|
||||
from sglang.srt.layers.quantization.compressed_tensors.utils import should_ignore_layer
|
||||
from sglang.srt.layers.quantization.modelslim.modelslim_moe import ModelSlimMoEMethod
|
||||
from sglang.srt.layers.quantization.modelslim.schemes import (
|
||||
ModelSlimScheme,
|
||||
ModelSlimW4A4Int4,
|
||||
ModelSlimW8A8Int8,
|
||||
)
|
||||
from sglang.srt.layers.quantization.unquant import UnquantizedLinearMethod
|
||||
from sglang.srt.utils import apply_module_patch
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
# func refers to RMSNorm.__init__
|
||||
def npu_wrapper_rmsnorm_init(func):
|
||||
def init(self, hidden_size: int, **extra_args) -> None:
|
||||
func(self, hidden_size, **extra_args)
|
||||
self.ignore_anti = True
|
||||
# The Ascend w8a8_int8 quantization requires adding a bias in rmsnorm
|
||||
self.bias = torch.nn.Parameter(torch.zeros(hidden_size), requires_grad=False)
|
||||
|
||||
return init
|
||||
|
||||
|
||||
# func refers to RMSNorm.forward_oot
|
||||
def npu_wrapper_rmsnorm_forward(func):
|
||||
def _rmsnorm_forward_oot(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
residual: Optional[torch.Tensor] = None,
|
||||
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
|
||||
from sgl_kernel_npu.norm.add_rmsnorm_bias import add_rmsnorm_bias
|
||||
|
||||
if not x.is_contiguous():
|
||||
x = x.contiguous()
|
||||
if residual is not None:
|
||||
out, residual_out = add_rmsnorm_bias(
|
||||
x,
|
||||
residual,
|
||||
self.weight.data,
|
||||
self.bias,
|
||||
self.variance_epsilon,
|
||||
)
|
||||
return out.to(x.dtype), residual_out
|
||||
|
||||
out = torch.ops.npu.npu_rms_norm(x, self.weight.data, self.variance_epsilon)[0]
|
||||
out = out + self.bias
|
||||
return out.to(x.dtype)
|
||||
|
||||
return _rmsnorm_forward_oot
|
||||
|
||||
|
||||
class ModelSlimConfig(QuantizationConfig):
|
||||
"""
|
||||
Config class for ModelSlim Quantization, a NPU-specific quantization type.
|
||||
"""
|
||||
|
||||
def __init__(self, quant_config: Dict[str, Any] = {}):
|
||||
super().__init__()
|
||||
self.quant_description = quant_config
|
||||
ignore = cast(List[str], quant_config.get("ignore", []))
|
||||
self.ignore = ignore if ignore is not None else []
|
||||
packed_modules_mapping = quant_config.get("packed_modules_mapping", {})
|
||||
self.packed_modules_mapping = (
|
||||
packed_modules_mapping if packed_modules_mapping is not None else {}
|
||||
)
|
||||
|
||||
for name in self.quant_description.keys():
|
||||
if "norm.bias" in name:
|
||||
apply_module_patch(
|
||||
"sglang.srt.layers.layernorm.RMSNorm",
|
||||
"__init__",
|
||||
[npu_wrapper_rmsnorm_init],
|
||||
)
|
||||
apply_module_patch(
|
||||
"sglang.srt.layers.layernorm.RMSNorm",
|
||||
"forward_npu",
|
||||
[npu_wrapper_rmsnorm_forward],
|
||||
)
|
||||
|
||||
def get_linear_method(self) -> ModelSlimLinearMethod:
|
||||
return ModelSlimLinearMethod(self)
|
||||
|
||||
@classmethod
|
||||
def get_supported_act_dtypes(cls) -> List[torch.dtype]:
|
||||
return [torch.int8, torch.float16, torch.bfloat16]
|
||||
|
||||
@classmethod
|
||||
def get_min_capability(cls) -> int:
|
||||
return 0
|
||||
|
||||
@classmethod
|
||||
def get_name(cls) -> str:
|
||||
return "modelslim"
|
||||
|
||||
@classmethod
|
||||
def get_config_filenames(cls) -> List[str]:
|
||||
filenames = ["quant_model_description.json"]
|
||||
return filenames
|
||||
|
||||
@classmethod
|
||||
def from_config(cls, config: Dict[str, Any]) -> ModelSlimConfig:
|
||||
return cls(config)
|
||||
|
||||
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 should_ignore_layer(
|
||||
prefix,
|
||||
ignore=self.ignore,
|
||||
fused_mapping=self.packed_modules_mapping,
|
||||
):
|
||||
return UnquantizedLinearMethod()
|
||||
key = "model"
|
||||
if "vision_model" in prefix:
|
||||
key = "vision_model"
|
||||
elif "visual" in prefix:
|
||||
key = "visual"
|
||||
packed_modules_mapping_subset = self.packed_modules_mapping.get(key, {})
|
||||
prefix_in_quant_config = prefix
|
||||
proj_name = prefix.split(".")[-1]
|
||||
if proj_name in packed_modules_mapping_subset:
|
||||
prefix_in_quant_config = prefix.replace(
|
||||
proj_name, packed_modules_mapping_subset[proj_name][0]
|
||||
)
|
||||
|
||||
if self.is_layer_skipped(prefix, packed_modules_mapping_subset):
|
||||
return UnquantizedLinearMethod()
|
||||
scheme = self.get_scheme(layer=layer, layer_name=prefix_in_quant_config)
|
||||
layer.scheme = scheme
|
||||
return ModelSlimLinearMethod(self)
|
||||
elif isinstance(layer, FusedMoE):
|
||||
return ModelSlimMoEMethod.get_moe_method(self, layer, prefix)
|
||||
return None
|
||||
|
||||
def _get_scheme_from_parts(
|
||||
self,
|
||||
layer_name: str,
|
||||
) -> ModelSlimScheme:
|
||||
|
||||
quant_type = self.quant_description.get(layer_name + ".weight", "")
|
||||
if quant_type == "W8A8_DYNAMIC" or quant_type == "W8A8":
|
||||
return ModelSlimW8A8Int8(
|
||||
quant_config=self.quant_description, prefix=layer_name
|
||||
)
|
||||
elif quant_type == "W4A4_DYNAMIC":
|
||||
return ModelSlimW4A4Int4(
|
||||
quant_config=self.quant_description, prefix=layer_name
|
||||
)
|
||||
raise NotImplementedError("No modelslim compatible scheme was found.")
|
||||
|
||||
def get_scheme(
|
||||
self, layer: torch.nn.Module, layer_name: Optional[str] = None
|
||||
) -> Optional[ModelSlimScheme]:
|
||||
"""
|
||||
get_scheme method adjusted for modelslim, taken from
|
||||
python/sglang/srt/layers/quantization/compressed_tensors/compressed_tensors.py
|
||||
"""
|
||||
scheme = self._get_scheme_from_parts(
|
||||
layer_name=layer_name,
|
||||
)
|
||||
|
||||
# Ascend doesn't support device capability
|
||||
logger.debug("Using scheme: %s for %s", scheme.__class__.__name__, layer_name)
|
||||
return scheme
|
||||
|
||||
def is_layer_skipped(
|
||||
self, prefix: str, fused_mapping: Mapping[str, List[str]] = MappingProxyType({})
|
||||
):
|
||||
# adapted from vllm.model_executor.layers.quantization.utils.quant_utils.is_layer_skipped
|
||||
proj_name = prefix.split(".")[-1]
|
||||
if proj_name in fused_mapping:
|
||||
shard_prefixes = [
|
||||
prefix.replace(proj_name, shard_proj_name)
|
||||
for shard_proj_name in fused_mapping[proj_name]
|
||||
]
|
||||
|
||||
is_skipped = None
|
||||
for shard_prefix in shard_prefixes:
|
||||
is_shard_skipped = (
|
||||
self.quant_description.get(shard_prefix + ".weight", "") == "FLOAT"
|
||||
)
|
||||
|
||||
if is_skipped is None:
|
||||
is_skipped = is_shard_skipped
|
||||
elif is_shard_skipped != is_skipped:
|
||||
raise ValueError(
|
||||
f"Detected some but not all shards of {prefix} "
|
||||
"are quantized. All shards of fused layers "
|
||||
"to have the same precision."
|
||||
)
|
||||
else:
|
||||
is_skipped = self.quant_description.get(prefix + ".weight", "") == "FLOAT"
|
||||
|
||||
assert is_skipped is not None
|
||||
return is_skipped
|
||||
|
||||
def get_scaled_act_names(self) -> List[str]:
|
||||
return []
|
||||
|
||||
|
||||
class ModelSlimLinearMethod(_NPULinearMethodBase):
|
||||
|
||||
def __init__(self, quantization_config: ModelSlimConfig):
|
||||
self.quantization_config = quantization_config
|
||||
|
||||
def process_weights_after_loading(self, layer: torch.nn.Module) -> None:
|
||||
layer.scheme.process_weights_after_loading(layer)
|
||||
|
||||
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,
|
||||
):
|
||||
"""
|
||||
Use the ModelSlimScheme associated with each layer to create
|
||||
the necessary parameters for the layer. See LinearMethodBase for param
|
||||
details
|
||||
"""
|
||||
weight_loader = extra_weight_attrs.get("weight_loader")
|
||||
layer.scheme.create_weights(
|
||||
layer=layer,
|
||||
input_size=input_size,
|
||||
input_size_per_partition=input_size_per_partition,
|
||||
output_partition_sizes=output_partition_sizes,
|
||||
output_size=output_size,
|
||||
params_dtype=params_dtype,
|
||||
weight_loader=weight_loader,
|
||||
)
|
||||
|
||||
def apply(
|
||||
self,
|
||||
layer: torch.nn.Module,
|
||||
x: torch.Tensor,
|
||||
bias: Optional[torch.Tensor] = None,
|
||||
):
|
||||
"""
|
||||
Use the output of create_weights and the CompressedTensorsScheme
|
||||
associated with the layer to apply the forward pass with the
|
||||
layer input. See LinearMethodBase for param details
|
||||
|
||||
"""
|
||||
|
||||
scheme = layer.scheme
|
||||
if scheme is None:
|
||||
raise ValueError("A scheme must be defined for each layer")
|
||||
return scheme.apply_weights(layer, x, bias=bias)
|
||||
377
python/sglang/srt/layers/quantization/modelslim/modelslim_moe.py
Normal file
377
python/sglang/srt/layers/quantization/modelslim/modelslim_moe.py
Normal file
@@ -0,0 +1,377 @@
|
||||
# Adapted from https://github.com/vllm-project/vllm/tree/v0.8.2/vllm/model_executor/layers/quantization/compressed_tensors
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from typing import TYPE_CHECKING, Any, Dict
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.hardware_backend.npu.quantization.fused_moe_method_npu import (
|
||||
NPUW4A8Int8DynamicMoEMethod,
|
||||
NPUW8A8Int8DynamicMoEMethod,
|
||||
)
|
||||
from sglang.srt.layers.quantization.base_config import FusedMoEMethodBase
|
||||
from sglang.srt.utils import set_weight_attrs
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.layers.moe import MoeRunnerConfig
|
||||
from sglang.srt.layers.moe.token_dispatcher import (
|
||||
CombineInput,
|
||||
StandardDispatchOutput,
|
||||
)
|
||||
from sglang.srt.layers.quantization.modelslim.modelslim import ModelSlimConfig
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"ModelSlimMoEMethod",
|
||||
"ModelSlimW4A8Int8MoE",
|
||||
"ModelSlimW8A8Int8MoE",
|
||||
]
|
||||
|
||||
|
||||
class ModelSlimMoEMethod(FusedMoEMethodBase):
|
||||
def __new__(cls, *args, **kwargs):
|
||||
if cls is ModelSlimMoEMethod:
|
||||
return super().__new__(cls)
|
||||
return super().__new__(cls)
|
||||
|
||||
@staticmethod
|
||||
def get_moe_method(
|
||||
quant_config: ModelSlimConfig,
|
||||
layer: torch.nn.Module,
|
||||
prefix: str,
|
||||
) -> "ModelSlimMoEMethod":
|
||||
# TODO: @dsikka: refactor this to use schemes as other kernels
|
||||
# are supported + check if the layer is being ignored.
|
||||
|
||||
prefix_in_quant_config = prefix + ".0.down_proj.weight"
|
||||
is_moe_w4a8_dynamic = (
|
||||
quant_config.quant_description.get(prefix_in_quant_config, "STATIC")
|
||||
== "W4A8_DYNAMIC"
|
||||
)
|
||||
is_moe_w8a8_dynamic = (
|
||||
quant_config.quant_description.get(prefix_in_quant_config, "STATIC")
|
||||
== "W8A8_DYNAMIC"
|
||||
)
|
||||
if is_moe_w4a8_dynamic:
|
||||
logger.info_once("Using ModelSlimW4A8Int8MoE")
|
||||
return ModelSlimW4A8Int8MoE(quant_config)
|
||||
elif is_moe_w8a8_dynamic:
|
||||
logger.info_once("Using ModelSlimW8A8Int8MoE")
|
||||
return ModelSlimW8A8Int8MoE(quant_config)
|
||||
else:
|
||||
logger.warning(
|
||||
f"Unsupported FusedMoe modelslim scheme: \
|
||||
{quant_config.quant_description.get(prefix_in_quant_config.strip())} \
|
||||
in layer: {prefix}"
|
||||
)
|
||||
return None
|
||||
|
||||
|
||||
class ModelSlimW4A8Int8MoE(ModelSlimMoEMethod):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
quant_config: Dict[str, Any],
|
||||
prefix: str = None,
|
||||
):
|
||||
self.quant_config = quant_config
|
||||
self.group_size = 0
|
||||
self.is_per_channel_weight = self.group_size == 0
|
||||
self.tp_size = 1
|
||||
self.activation_use_clip = False
|
||||
self.kernel = NPUW4A8Int8DynamicMoEMethod()
|
||||
|
||||
def create_weights(
|
||||
self,
|
||||
layer: torch.nn.Module,
|
||||
num_experts: int,
|
||||
hidden_size: int,
|
||||
intermediate_size_per_partition: int,
|
||||
params_dtype: torch.dtype,
|
||||
**extra_weight_attrs,
|
||||
) -> None:
|
||||
from sglang.srt.layers.moe.fused_moe_triton import FusedMoeWeightScaleSupported
|
||||
|
||||
self.is_per_channel_weight = self.group_size == 0
|
||||
self.num_experts = num_experts
|
||||
extra_weight_attrs.update(
|
||||
{"quant_method": FusedMoeWeightScaleSupported.CHANNEL.value}
|
||||
)
|
||||
|
||||
# >> weight
|
||||
w13_output_size = intermediate_size_per_partition
|
||||
w2_output_size = hidden_size // 2
|
||||
w13_weight = torch.nn.Parameter(
|
||||
torch.empty(num_experts, w13_output_size, hidden_size, dtype=torch.int8),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w13_weight", w13_weight)
|
||||
set_weight_attrs(w13_weight, extra_weight_attrs)
|
||||
w2_weight = torch.nn.Parameter(
|
||||
torch.empty(
|
||||
num_experts,
|
||||
w2_output_size,
|
||||
intermediate_size_per_partition,
|
||||
dtype=torch.int8,
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w2_weight", w2_weight)
|
||||
set_weight_attrs(w2_weight, extra_weight_attrs)
|
||||
|
||||
# >> scale
|
||||
w13_weight_scale = torch.nn.Parameter(
|
||||
torch.empty(
|
||||
num_experts, 2 * intermediate_size_per_partition, 1, dtype=torch.float32
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w13_weight_scale", w13_weight_scale)
|
||||
set_weight_attrs(w13_weight_scale, extra_weight_attrs)
|
||||
|
||||
w2_weight_scale = torch.nn.Parameter(
|
||||
torch.empty(num_experts, hidden_size, 1, dtype=torch.float32),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w2_weight_scale", w2_weight_scale)
|
||||
set_weight_attrs(w2_weight_scale, extra_weight_attrs)
|
||||
|
||||
# >> offset
|
||||
w13_weight_offset = torch.nn.Parameter(
|
||||
torch.empty(
|
||||
num_experts, 2 * intermediate_size_per_partition, 1, dtype=torch.float32
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w13_weight_offset", w13_weight_offset)
|
||||
set_weight_attrs(w13_weight_offset, extra_weight_attrs)
|
||||
|
||||
w2_weight_offset = torch.nn.Parameter(
|
||||
torch.empty(num_experts, hidden_size, 1, dtype=torch.float32),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w2_weight_offset", w2_weight_offset)
|
||||
set_weight_attrs(w2_weight_offset, extra_weight_attrs)
|
||||
|
||||
# >>> special param for w4a8
|
||||
if not self.is_per_channel_weight:
|
||||
w13_weight_scale_second = torch.nn.Parameter(
|
||||
torch.empty(
|
||||
num_experts,
|
||||
2 * intermediate_size_per_partition,
|
||||
hidden_size // self.group_size,
|
||||
dtype=torch.float32,
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w13_weight_scale_second", w13_weight_scale_second)
|
||||
set_weight_attrs(w13_weight_scale_second, extra_weight_attrs)
|
||||
w13_weight_offset_second = torch.nn.Parameter(
|
||||
torch.empty(
|
||||
num_experts,
|
||||
2 * intermediate_size_per_partition,
|
||||
hidden_size // self.group_size,
|
||||
dtype=torch.float32,
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter(
|
||||
"w13_weight_offset_second", w13_weight_offset_second
|
||||
)
|
||||
set_weight_attrs(w13_weight_offset_second, extra_weight_attrs)
|
||||
|
||||
w2_weight_scale_second = torch.nn.Parameter(
|
||||
torch.empty(
|
||||
num_experts,
|
||||
hidden_size,
|
||||
intermediate_size_per_partition // self.group_size,
|
||||
dtype=torch.float32,
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w2_weight_scale_second", w2_weight_scale_second)
|
||||
set_weight_attrs(w2_weight_scale_second, extra_weight_attrs)
|
||||
|
||||
w2_weight_offset_second = torch.nn.Parameter(
|
||||
torch.empty(
|
||||
num_experts,
|
||||
hidden_size,
|
||||
intermediate_size_per_partition // self.group_size,
|
||||
dtype=torch.float32,
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w2_weight_offset_second", w2_weight_offset_second)
|
||||
set_weight_attrs(w2_weight_offset_second, extra_weight_attrs)
|
||||
|
||||
w13_scale_bias = torch.nn.Parameter(
|
||||
torch.empty(
|
||||
num_experts, 2 * intermediate_size_per_partition, 1, dtype=torch.float32
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w13_scale_bias", w13_scale_bias)
|
||||
set_weight_attrs(w13_scale_bias, extra_weight_attrs)
|
||||
|
||||
w2_scale_bias = torch.nn.Parameter(
|
||||
torch.empty(
|
||||
num_experts, hidden_size, 16 // self.tp_size, dtype=torch.float32
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w2_scale_bias", w2_scale_bias)
|
||||
set_weight_attrs(w2_scale_bias, extra_weight_attrs)
|
||||
|
||||
def process_weights_after_loading(self, layer: torch.nn.Module) -> None:
|
||||
self.kernel.process_weights_after_loading(
|
||||
layer, self.is_per_channel_weight, self.activation_use_clip
|
||||
)
|
||||
|
||||
def create_moe_runner(
|
||||
self, layer: torch.nn.Module, moe_runner_config: "MoeRunnerConfig"
|
||||
):
|
||||
self.moe_runner_config = moe_runner_config
|
||||
|
||||
def apply(
|
||||
self,
|
||||
layer,
|
||||
dispatch_output: "StandardDispatchOutput",
|
||||
) -> "CombineInput":
|
||||
# FIXME W4A8 without EP can give 0 accuracy
|
||||
return self.kernel.apply(layer, dispatch_output)
|
||||
|
||||
def apply_without_routing_weights(
|
||||
self,
|
||||
layer,
|
||||
hidden_states,
|
||||
hidden_states_scale,
|
||||
group_list_type,
|
||||
group_list,
|
||||
output_dtype,
|
||||
):
|
||||
return self.kernel.apply_without_routing_weights(
|
||||
layer,
|
||||
hidden_states,
|
||||
hidden_states_scale,
|
||||
group_list_type,
|
||||
group_list,
|
||||
output_dtype,
|
||||
)
|
||||
|
||||
|
||||
class ModelSlimW8A8Int8MoE(ModelSlimMoEMethod):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
quant_config: Dict[str, Any],
|
||||
prefix: str = None,
|
||||
):
|
||||
self.quant_config = quant_config
|
||||
self.kernel = NPUW8A8Int8DynamicMoEMethod()
|
||||
|
||||
def create_weights(
|
||||
self,
|
||||
layer: torch.nn.Module,
|
||||
num_experts: int,
|
||||
hidden_size: int,
|
||||
intermediate_size_per_partition: int,
|
||||
params_dtype: torch.dtype,
|
||||
**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.CHANNEL.value}
|
||||
)
|
||||
|
||||
# weight
|
||||
w13_weight = torch.nn.Parameter(
|
||||
torch.empty(
|
||||
num_experts,
|
||||
2 * intermediate_size_per_partition,
|
||||
hidden_size,
|
||||
dtype=torch.int8,
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w13_weight", w13_weight)
|
||||
set_weight_attrs(w13_weight, extra_weight_attrs)
|
||||
w2_weight = torch.nn.Parameter(
|
||||
torch.empty(
|
||||
num_experts,
|
||||
hidden_size,
|
||||
intermediate_size_per_partition,
|
||||
dtype=torch.int8,
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w2_weight", w2_weight)
|
||||
set_weight_attrs(w2_weight, extra_weight_attrs)
|
||||
# scale
|
||||
w13_weight_scale = torch.nn.Parameter(
|
||||
torch.empty(
|
||||
num_experts, 2 * intermediate_size_per_partition, 1, dtype=torch.float32
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w13_weight_scale", w13_weight_scale)
|
||||
set_weight_attrs(w13_weight_scale, extra_weight_attrs)
|
||||
w2_weight_scale = torch.nn.Parameter(
|
||||
torch.empty(num_experts, hidden_size, 1, dtype=torch.float32),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w2_weight_scale", w2_weight_scale)
|
||||
set_weight_attrs(w2_weight_scale, extra_weight_attrs)
|
||||
# offset
|
||||
w13_weight_offset = torch.nn.Parameter(
|
||||
torch.empty(
|
||||
num_experts, 2 * intermediate_size_per_partition, 1, dtype=torch.float32
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w13_weight_offset", w13_weight_offset)
|
||||
set_weight_attrs(w13_weight_offset, extra_weight_attrs)
|
||||
w2_weight_offset = torch.nn.Parameter(
|
||||
torch.empty(num_experts, hidden_size, 1, dtype=torch.float32),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w2_weight_offset", w2_weight_offset)
|
||||
set_weight_attrs(w2_weight_offset, extra_weight_attrs)
|
||||
|
||||
def process_weights_after_loading(self, layer: torch.nn.Module) -> None:
|
||||
self.kernel.process_weights_after_loading(layer)
|
||||
|
||||
def create_moe_runner(
|
||||
self, layer: torch.nn.Module, moe_runner_config: "MoeRunnerConfig"
|
||||
):
|
||||
self.moe_runner_config = moe_runner_config
|
||||
|
||||
def apply(
|
||||
self,
|
||||
layer,
|
||||
dispatch_output: "StandardDispatchOutput",
|
||||
) -> "CombineInput":
|
||||
return self.kernel.apply(layer, dispatch_output)
|
||||
|
||||
def apply_without_routing_weights(
|
||||
self,
|
||||
layer,
|
||||
hidden_states,
|
||||
hidden_states_scale,
|
||||
group_list_type,
|
||||
group_list,
|
||||
output_dtype,
|
||||
):
|
||||
return self.kernel.apply_without_routing_weights(
|
||||
layer,
|
||||
hidden_states,
|
||||
hidden_states_scale,
|
||||
group_list_type,
|
||||
group_list,
|
||||
output_dtype,
|
||||
)
|
||||
@@ -0,0 +1,11 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from .modelslim_scheme import ModelSlimScheme
|
||||
from .modelslim_w4a4_int4 import ModelSlimW4A4Int4
|
||||
from .modelslim_w8a8_int8 import ModelSlimW8A8Int8
|
||||
|
||||
__all__ = [
|
||||
"ModelSlimScheme",
|
||||
"ModelSlimW8A8Int8",
|
||||
"ModelSlimW4A4Int4",
|
||||
]
|
||||
@@ -0,0 +1,48 @@
|
||||
# Adapted from https://github.com/vllm-project/vllm/tree/main/vllm/model_executor/layers/quantization/compressed_tensors
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Optional
|
||||
|
||||
import torch
|
||||
|
||||
__all__ = ["ModelSlimScheme"]
|
||||
|
||||
|
||||
class ModelSlimScheme(ABC):
|
||||
"""
|
||||
Abstract class used to describe the weight creation and forward pass
|
||||
of different quantization schemes supported by CompressedTensors.
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
def create_weights(self, *args, **kwargs):
|
||||
"""
|
||||
Weight creation for the particular scheme. Inputs to this function
|
||||
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
@abstractmethod
|
||||
def apply_weights(
|
||||
self, layer: torch.nn.Module, x: torch.Tensor, bias: Optional[torch.Tensor]
|
||||
):
|
||||
"""
|
||||
Run the forward pass for the particular scheme. This is where
|
||||
scheme-specific dequant/quant steps/kernels should be applied.
|
||||
|
||||
:param layer: torch.nn.Module with the registered weights and
|
||||
other parameters relevant to the particular scheme.
|
||||
:param x: input to the layer
|
||||
:param bias: bias parameter
|
||||
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
@abstractmethod
|
||||
def process_weights_after_loading(self, layer: torch.nn.Module):
|
||||
"""
|
||||
Called after weight loading is complete for any cleanup that
|
||||
needs to occur.
|
||||
"""
|
||||
raise NotImplementedError
|
||||
@@ -0,0 +1,99 @@
|
||||
# Adapted from https://github.com/vllm-project/vllm/tree/main/vllm/model_executor/layers/quantization/compressed_tensors
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.hardware_backend.npu.quantization.linear_method_npu import (
|
||||
NPU_W4A4DynamicLinearMethod,
|
||||
)
|
||||
from sglang.srt.layers.parameter import PerTensorScaleParameter
|
||||
from sglang.srt.layers.quantization.modelslim.schemes import ModelSlimScheme
|
||||
from sglang.srt.utils import set_weight_attrs
|
||||
|
||||
|
||||
class ModelSlimW4A4Int4(ModelSlimScheme):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
quant_config: Dict[str, any],
|
||||
prefix: str,
|
||||
):
|
||||
self.quant_config = quant_config
|
||||
self.is_dynamic = self.quant_config[prefix + ".weight"] == "W4A4_DYNAMIC"
|
||||
self.kernel = NPU_W4A4DynamicLinearMethod()
|
||||
|
||||
@staticmethod
|
||||
def get_weight(
|
||||
input_size: int, output_size: int, params_dtype: torch.dtype
|
||||
) -> Dict[str, Any]:
|
||||
params_dict = {"weight": torch.empty(output_size, input_size, dtype=torch.int8)}
|
||||
return params_dict
|
||||
|
||||
@staticmethod
|
||||
def get_perchannel_param(
|
||||
output_size: int,
|
||||
params_dtype: torch.dtype,
|
||||
) -> Dict[str, Any]:
|
||||
params_dict = {}
|
||||
params_dict["weight_scale"] = torch.empty(output_size, 1, dtype=params_dtype)
|
||||
params_dict["weight_offset"] = torch.empty(output_size, 1, dtype=params_dtype)
|
||||
return params_dict
|
||||
|
||||
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,
|
||||
) -> None:
|
||||
output_size_per_partition = sum(output_partition_sizes)
|
||||
weight_loader = extra_weight_attrs.get("weight_loader")
|
||||
|
||||
weight_dict = {
|
||||
"weight": torch.empty(
|
||||
output_size_per_partition, input_size_per_partition, dtype=torch.int8
|
||||
)
|
||||
}
|
||||
for weight_name, weight_param in weight_dict.items():
|
||||
param = torch.nn.Parameter(weight_param, requires_grad=False)
|
||||
set_weight_attrs(param, {"input_dim": 1, "output_dim": 0})
|
||||
layer.register_parameter(weight_name, param)
|
||||
set_weight_attrs(param, extra_weight_attrs)
|
||||
|
||||
pertensor_dict = {}
|
||||
for pertensor_name, pertensor_param in pertensor_dict.items():
|
||||
param = PerTensorScaleParameter(
|
||||
data=pertensor_param, weight_loader=weight_loader
|
||||
)
|
||||
# disable warning
|
||||
param.ignore_warning = True
|
||||
layer.register_parameter(pertensor_name, param)
|
||||
|
||||
perchannel_dict = {}
|
||||
perchannel_dict["weight_scale"] = torch.empty(
|
||||
output_size_per_partition, 1, dtype=params_dtype
|
||||
)
|
||||
perchannel_dict["weight_offset"] = torch.empty(
|
||||
output_size_per_partition, 1, dtype=params_dtype
|
||||
)
|
||||
for perchannel_name, perchannel_param in perchannel_dict.items():
|
||||
param = torch.nn.Parameter(perchannel_param, requires_grad=False)
|
||||
set_weight_attrs(param, {"output_dim": 0})
|
||||
layer.register_parameter(perchannel_name, param)
|
||||
set_weight_attrs(param, extra_weight_attrs)
|
||||
|
||||
def process_weights_after_loading(self, layer):
|
||||
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)
|
||||
@@ -0,0 +1,117 @@
|
||||
# Adapted from https://github.com/vllm-project/vllm/tree/main/vllm/model_executor/layers/quantization/compressed_tensors
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from typing import Dict, List, Optional
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.hardware_backend.npu.quantization.linear_method_npu import (
|
||||
NPUW8A8Int8DynamicLinearMethod,
|
||||
NPUW8A8Int8LinearMethod,
|
||||
)
|
||||
from sglang.srt.layers.parameter import (
|
||||
ChannelQuantScaleParameter,
|
||||
ModelWeightParameter,
|
||||
PerTensorScaleParameter,
|
||||
)
|
||||
from sglang.srt.layers.quantization.modelslim.schemes import ModelSlimScheme
|
||||
|
||||
|
||||
class ModelSlimW8A8Int8(ModelSlimScheme):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
quant_config: Dict[str, any],
|
||||
prefix: str,
|
||||
):
|
||||
self.quant_config = quant_config
|
||||
self.is_dynamic = (
|
||||
self.quant_config.get(prefix + ".weight", "") == "W8A8_DYNAMIC"
|
||||
)
|
||||
if self.is_dynamic:
|
||||
self.kernel = NPUW8A8Int8DynamicLinearMethod()
|
||||
else:
|
||||
self.kernel = NPUW8A8Int8LinearMethod()
|
||||
|
||||
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)
|
||||
|
||||
weight = ModelWeightParameter(
|
||||
data=torch.empty(
|
||||
(output_size_per_partition, input_size_per_partition), dtype=torch.int8
|
||||
),
|
||||
input_dim=1,
|
||||
output_dim=0,
|
||||
weight_loader=weight_loader,
|
||||
)
|
||||
layer.register_parameter("weight", weight)
|
||||
|
||||
weight_scale = ChannelQuantScaleParameter(
|
||||
data=torch.empty((output_size_per_partition, 1), dtype=params_dtype),
|
||||
output_dim=0,
|
||||
weight_loader=weight_loader,
|
||||
)
|
||||
layer.register_parameter("weight_scale", weight_scale)
|
||||
|
||||
weight_offset = ChannelQuantScaleParameter(
|
||||
data=torch.empty((output_size_per_partition, 1), dtype=params_dtype),
|
||||
output_dim=0,
|
||||
weight_loader=weight_loader,
|
||||
)
|
||||
layer.register_parameter("weight_offset", weight_offset)
|
||||
|
||||
if not self.is_dynamic:
|
||||
input_scale = PerTensorScaleParameter(
|
||||
data=torch.empty(1, dtype=params_dtype),
|
||||
weight_loader=weight_loader,
|
||||
)
|
||||
input_scale.ignore_warning = True
|
||||
layer.register_parameter("input_scale", input_scale)
|
||||
|
||||
input_offset = PerTensorScaleParameter(
|
||||
data=torch.empty(1, dtype=params_dtype),
|
||||
weight_loader=weight_loader,
|
||||
)
|
||||
input_offset.ignore_warning = True
|
||||
layer.register_parameter("input_offset", input_offset)
|
||||
|
||||
quant_bias = ChannelQuantScaleParameter(
|
||||
data=torch.empty(output_size_per_partition, dtype=torch.int32),
|
||||
output_dim=0,
|
||||
weight_loader=weight_loader,
|
||||
)
|
||||
layer.register_parameter("quant_bias", quant_bias)
|
||||
|
||||
if params_dtype == torch.bfloat16:
|
||||
deq_scale_dtype = torch.float32
|
||||
elif params_dtype == torch.float16:
|
||||
deq_scale_dtype = torch.int64
|
||||
else:
|
||||
raise ValueError(f"Unsupported params_dtype: {params_dtype}")
|
||||
deq_scale = ChannelQuantScaleParameter(
|
||||
data=torch.empty(output_size_per_partition, dtype=deq_scale_dtype),
|
||||
output_dim=0,
|
||||
weight_loader=weight_loader,
|
||||
)
|
||||
layer.register_parameter("deq_scale", deq_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)
|
||||
Reference in New Issue
Block a user