init v0.23.0

Signed-off-by: Sun Ruoxi <sunruoxi@4paradigm.com>
This commit is contained in:
2026-08-27 15:11:51 +08:00
parent b582a8e7d1
commit 7f8a1b1f7a
2849 changed files with 712887 additions and 22001 deletions

View File

@@ -0,0 +1,103 @@
#
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
# This file is a part of the vllm-ascend project.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
#
"""Ascend quantization scheme implementations.
This module provides all quantization scheme implementations for Ascend NPU.
Schemes are automatically registered via the @register_scheme decorator.
Usage:
from vllm_ascend.quantization.methods import get_scheme_class
# Get a scheme class by quant_type and layer_type
scheme_cls = get_scheme_class("W8A8_DYNAMIC", "linear")
scheme = scheme_cls()
"""
from typing import Any
# Import base classes
from .base import AscendAttentionScheme, AscendLinearScheme, AscendMoEScheme, QuantType
# Import all scheme classes for external access
from .fp8 import AscendW4A8MXFPDSDynamicFusedMoEMethod, AscendW8A8MXFP8DSDynamicLinearMethod
from .kv_c8 import AscendFAQuantAttentionMethod
# Import registry functions
from .registry import get_scheme_class, register_scheme
from .w4a4_flatquant import AscendW4A4FlatQuantDynamicLinearMethod
from .w4a4_laos_dynamic import AscendW4A4LaosDynamicLinearMethod
from .w4a4_mxfp4 import AscendW4A4MXFP4DynamicFusedMoEMethod, AscendW4A4MXFP4DynamicLinearMethod
from .w4a4_mxfp4_flatquant import AscendW4A4MXFP4FlatQuantDynamicLinearMethod
from .w4a8 import AscendW4A8DynamicFusedMoEMethod, AscendW4A8DynamicLinearMethod
from .w4a8_mxfp4 import AscendW4A8MXFPDynamicFusedMoEMethod, AscendW4A8MXFPDynamicLinearMethod
from .w4a16 import AscendW4A16FusedMoEMethod
from .w4a16_mxfp4 import AscendW4A16MXFP4FusedMoEMethod
from .w8a8_dynamic import AscendW8A8DynamicFusedMoEMethod, AscendW8A8DynamicLinearMethod
from .w8a8_mxfp8 import AscendW8A8MXFP8DynamicLinearMethod
from .w8a8_pdmix import AscendW8A8PDMixFusedMoeMethod, AscendW8A8PDMixLinearMethod
from .w8a8_static import AscendW8A8LinearMethod
from .w8a8fp8_dynamic import AscendW8A8FP8DynamicFusedMoEMethod, AscendW8A8FP8DynamicLinearMethod
from .w8a16 import AscendW8A16LinearMethod
def is_mx_quant_type(instance: Any) -> bool:
"""Checks if the quantization method is a microscaling (MX) type."""
MX_QUANT_TYPES = (
AscendW8A8MXFP8DynamicLinearMethod,
AscendW4A4MXFP4DynamicLinearMethod,
AscendW4A4MXFP4DynamicFusedMoEMethod,
AscendW4A4MXFP4FlatQuantDynamicLinearMethod,
AscendW4A8MXFPDynamicLinearMethod,
AscendW4A8MXFPDynamicFusedMoEMethod,
AscendW4A16MXFP4FusedMoEMethod,
)
return isinstance(instance, MX_QUANT_TYPES)
__all__ = [
# Base classes
"AscendAttentionScheme",
"AscendLinearScheme",
"AscendMoEScheme",
"QuantType",
# Registry functions
"register_scheme",
"get_scheme_class",
# Utility functions
"is_mx_quant_type",
# Scheme classes
"AscendW8A8LinearMethod",
"AscendW8A8DynamicLinearMethod",
"AscendW8A8DynamicFusedMoEMethod",
"AscendW8A8FP8DynamicLinearMethod",
"AscendW8A8FP8DynamicFusedMoEMethod",
"AscendW8A8MXFP8DynamicLinearMethod",
"AscendW8A8PDMixLinearMethod",
"AscendW8A8PDMixFusedMoeMethod",
"AscendW8A16LinearMethod",
"AscendW4A8DynamicLinearMethod",
"AscendW4A8DynamicFusedMoEMethod",
"AscendW4A16FusedMoEMethod",
"AscendW4A4FlatQuantDynamicLinearMethod",
"AscendW4A4LaosDynamicLinearMethod",
"AscendFAQuantAttentionMethod",
"AscendW4A4MXFP4DynamicLinearMethod",
"AscendW4A4MXFP4DynamicFusedMoEMethod",
"AscendW4A4MXFP4FlatQuantDynamicLinearMethod",
"AscendW8A8MXFP8DSDynamicLinearMethod",
"AscendW4A8MXFPDSDynamicFusedMoEMethod",
]

View File

@@ -0,0 +1,298 @@
#
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
# This file is a part of the vllm-ascend project.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
#
"""Abstract base classes for Ascend quantization schemes."""
from abc import ABC, abstractmethod
from collections.abc import Callable
from typing import Any
import torch
from vllm_ascend.quantization.quant_type import QuantType
def get_moe_num_logical_experts(
layer: torch.nn.Module,
num_experts: int,
global_redundant_expert_num: int = 0,
num_shared_experts: int = 0,
) -> int:
moe_config = getattr(layer, "moe_config", None)
num_logical_experts = getattr(moe_config, "num_logical_experts", None)
if num_logical_experts is not None:
return int(num_logical_experts)
return int(num_experts - global_redundant_expert_num - num_shared_experts)
class AscendLinearScheme(ABC):
"""Base class for all linear quantization schemes.
Subclasses must implement get_weight() and apply() methods.
Other methods have default implementations that return empty dicts
or do nothing.
"""
@abstractmethod
def get_weight(self, input_size: int, output_size: int, params_dtype: torch.dtype) -> dict[str, Any]:
"""Return weight tensor specifications.
Args:
input_size: Input dimension of the linear layer.
output_size: Output dimension of the linear layer.
params_dtype: Data type for parameters.
Returns:
Dictionary mapping parameter names to empty tensors with
the correct shape and dtype.
"""
...
def get_pertensor_param(self, params_dtype: torch.dtype, **kwargs: Any) -> dict[str, Any]:
"""Return per-tensor parameter specifications (e.g., input_scale).
Args:
params_dtype: Data type for parameters.
**kwargs: Additional keyword arguments for subclass extensions
Returns:
Dictionary mapping parameter names to empty tensors.
"""
return {}
def get_perchannel_param(self, output_size: int, params_dtype: torch.dtype) -> dict[str, Any]:
"""Return per-channel parameter specifications (e.g., weight_scale).
Args:
output_size: Output dimension of the linear layer.
params_dtype: Data type for parameters.
Returns:
Dictionary mapping parameter names to empty tensors.
"""
return {}
def get_pergroup_param(
self, input_size: int, output_size: int, params_dtype: torch.dtype, layer_type: str | None = None
) -> dict[str, Any]:
"""Return per-group parameter specifications.
Args:
input_size: Input dimension of the linear layer.
output_size: Output dimension of the linear layer.
params_dtype: Data type for parameters.
layer_type: Type of layer (e.g., "row" for RowParallelLinear).
Returns:
Dictionary mapping parameter names to empty tensors.
"""
return {}
@abstractmethod
def apply(
self, layer: torch.nn.Module, x: torch.Tensor, bias: torch.Tensor | None = None, tp_rank: int | None = 0
) -> torch.Tensor:
"""Forward computation.
Args:
layer: The linear layer module.
x: Input tensor.
bias: Optional bias tensor.
tp_rank: Tensor parallel rank.
Returns:
Output tensor after quantized linear operation.
"""
...
def process_weights_after_loading(self, layer: torch.nn.Module) -> None:
"""Post-loading weight processing (transpose, format conversion, etc.).
Args:
layer: The linear layer module.
"""
return
class AscendAttentionScheme(ABC):
"""Base class for all attention quantization schemes.
Subclasses must implement apply() method.
Other methods have default implementations.
"""
def create_weights(self, layer: torch.nn.Module) -> None:
"""Create weights for attention quantization.
Args:
layer: The attention layer module.
"""
return
def process_weights_after_loading(self, layer: torch.nn.Module) -> None:
"""Post-loading weight processing for attention layer.
Args:
layer: The attention layer module.
"""
return
@abstractmethod
def apply(
self,
layer: torch.nn.Module,
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
kv_cache,
attn_metadata,
attn_type,
scale,
output,
) -> torch.Tensor:
"""Forward computation for attention layer.
Args:
layer: The attention layer module.
query: Query tensor.
key: Key tensor.
value: Value tensor.
kv_cache: KV cache.
attn_metadata: Attention metadata.
attn_type: Attention type.
scale: Scale factor.
output: Output tensor.
Returns:
Output tensor after attention computation.
"""
...
class AscendMoEScheme(ABC):
"""Base class for all MoE quantization schemes.
Subclasses must implement get_weight(), get_dynamic_quant_param(),
and apply() methods.
Attributes:
quant_type: The quantization type for this scheme. Subclasses should
override this class attribute to declare their quant type.
"""
# Default quant type - subclasses should override this
quant_type: QuantType = QuantType.NONE
@abstractmethod
def get_weight(
self, num_experts: int, intermediate_size_per_partition: int, hidden_sizes: int, params_dtype: torch.dtype
) -> dict[str, Any]:
"""Return weight tensor specifications for MoE layer.
Args:
num_experts: Number of experts.
intermediate_size_per_partition: Intermediate size per partition.
hidden_sizes: Hidden dimension size.
params_dtype: Data type for parameters.
Returns:
Dictionary mapping parameter names to empty tensors.
"""
...
@abstractmethod
def get_dynamic_quant_param(
self, num_experts: int, intermediate_size_per_partition: int, hidden_sizes: int, params_dtype: torch.dtype
) -> dict[str, Any]:
"""Return dynamic quantization parameters for MoE layer.
Args:
num_experts: Number of experts.
intermediate_size_per_partition: Intermediate size per partition.
hidden_sizes: Hidden dimension size.
params_dtype: Data type for parameters.
Returns:
Dictionary mapping parameter names to empty tensors.
"""
...
@abstractmethod
def apply(
self,
layer: torch.nn.Module,
x: torch.Tensor,
router_logits: torch.Tensor,
top_k: int,
renormalize: bool,
use_grouped_topk: bool = False,
num_experts: int = -1,
expert_map: torch.Tensor | None = None,
topk_group: int | None = None,
num_expert_group: int | None = None,
custom_routing_function: Callable | None = None,
scoring_func: str = "softmax",
routed_scaling_factor: float = 1.0,
e_score_correction_bias: torch.Tensor | None = None,
is_prefill: bool = True,
enable_force_load_balance: bool = False,
log2phy: torch.Tensor | None = None,
global_redundant_expert_num: int = 0,
pertoken_scale: Any | None = None,
activation: str = "silu",
apply_router_weight_on_input: bool = False,
mc2_mask: torch.Tensor | None = None,
tid2eid: Any | None = None,
) -> torch.Tensor:
"""Forward computation for MoE layer.
Args:
layer: The MoE layer module.
x: Input hidden states.
router_logits: Router logits for expert selection.
top_k: Number of experts to select per token.
renormalize: Whether to renormalize expert weights.
use_grouped_topk: Whether to use grouped top-k selection.
num_experts: Number of experts.
expert_map: Mapping from local to global expert indices.
topk_group: Group size for grouped top-k.
num_expert_group: Number of expert groups.
custom_routing_function: Custom routing function.
scoring_func: Scoring function name.
routed_scaling_factor: Scaling factor for routed experts.
e_score_correction_bias: Expert score correction bias.
is_prefill: Whether in prefill phase.
enable_force_load_balance: Whether to force load balancing.
log2phy: Logical to physical expert mapping.
global_redundant_expert_num: Number of redundant experts.
pertoken_scale: Optional per-token activation scale from prepare stage.
activation: Expert MLP activation type.
apply_router_weight_on_input: Whether to pre-scale hidden states by router weights.
mc2_mask: Optional mask used by MC2 dispatch.
Returns:
Output tensor after MoE computation.
"""
...
def process_weights_after_loading(self, layer: torch.nn.Module) -> None:
"""Post-loading weight processing for MoE layer.
Args:
layer: The MoE layer module.
"""
return

View File

@@ -0,0 +1,130 @@
#
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
# This file is a part of the vllm-ascend project.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
#
from typing import Any
import torch
import torch_npu
from vllm.config import get_current_vllm_config
from .base import QuantType
from .registry import register_scheme
from .w4a8_mxfp4 import AscendW4A8MXFPDynamicFusedMoEMethod
from .w8a8_mxfp8 import AscendW8A8MXFP8DynamicLinearMethod
@register_scheme("FP8", "ds_linear")
class AscendW8A8MXFP8DSDynamicLinearMethod(AscendW8A8MXFP8DynamicLinearMethod):
"""Linear method for DS original W8A8 mxfp(blocksize: 128 * 128) quantization.
scales are reorganize as blocksize 32 * 1 in process_weights_after_loading
"""
model_dtype = None
def __init__(self, quant_config):
super().__init__()
self.block_size = quant_config.get("weight_block_size", [128, 128])[0]
vllm_config = get_current_vllm_config()
tp_size = vllm_config.parallel_config.tensor_parallel_size
hf_config = vllm_config.model_config.hf_config
self.n_groups = hf_config.o_groups
self.n_local_groups = self.n_groups // tp_size
self.o_lora_rank = hf_config.o_lora_rank
def get_pergroup_param(
self, input_size: int, output_size: int, params_dtype: torch.dtype, layer_type: str | None = None
) -> dict[str, Any]:
params_dict = {}
params_dict["weight_scale"] = torch.empty(
output_size // self.block_size, input_size // self.block_size, dtype=torch.float32
)
params_dict["_packed_dim"] = 0
params_dict["_packed_factor"] = self.block_size
return params_dict
def process_weights_after_loading(self, layer):
layer.weight_scale.data = layer.weight_scale.data.view(torch.int32) >> 23 & 0xFF
layer.weight_scale.data = layer.weight_scale.data.to(torch.uint8)
layer.weight_scale.data = layer.weight_scale.data.repeat_interleave(4, dim=1).repeat_interleave(128, dim=0)
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)
if layer.prefix.endswith("wo_a"):
layer.weight.data = (
layer.weight.data.T.reshape(self.n_local_groups, self.o_lora_rank, -1).transpose(1, 2).contiguous()
)
layer.weight_scale.data = (
layer.weight_scale.data.transpose(0, 1)
.reshape(self.n_local_groups, self.o_lora_rank, -1, 2)
.transpose(1, 2)
.contiguous()
)
@register_scheme("FP8", "w4a8_moe")
class AscendW4A8MXFPDSDynamicFusedMoEMethod(AscendW4A8MXFPDynamicFusedMoEMethod):
"""FusedMoe method for DS original w4a8 mxfp quantization."""
model_dtype = None
quant_type: QuantType = QuantType.W4A8MXFP
def __init__(self, quant_config, tid2eid=None):
super().__init__()
self.tid2eid = tid2eid
def get_dynamic_quant_param(
self, num_experts: int, intermediate_size_per_partition: int, hidden_sizes: int, params_dtype: torch.dtype
) -> dict[str, Any]:
param_dict = {}
param_dict["w13_weight_scale"] = torch.empty(
num_experts,
2 * intermediate_size_per_partition,
hidden_sizes // self.group_size,
dtype=torch.float8_e8m0fnu,
)
param_dict["w2_weight_scale"] = torch.empty(
num_experts, hidden_sizes, intermediate_size_per_partition // self.group_size, dtype=torch.float8_e8m0fnu
)
return param_dict
def process_weights_after_loading(self, layer):
layer.w13_weight.data = torch_npu.npu_format_cast(
layer.w13_weight.data.view(torch.uint8),
29,
customize_dtype=torch.float8_e4m3fn,
input_dtype=torch_npu.float4_e2m1fn_x2,
)
layer.w2_weight.data = torch_npu.npu_format_cast(
layer.w2_weight.data.view(torch.uint8),
29,
customize_dtype=torch.float8_e4m3fn,
input_dtype=torch_npu.float4_e2m1fn_x2,
)
layer.w13_weight.data = layer.w13_weight.data.transpose(1, 2)
layer.w2_weight.data = layer.w2_weight.data.transpose(1, 2)
g, n, k = layer.w13_weight_scale.shape
layer.w13_weight_scale.data = (
layer.w13_weight_scale.data.reshape(g, n, k // 2, 2).view(torch.uint8).transpose(-3, -2)
)
g, n, k = layer.w2_weight_scale.shape
layer.w2_weight_scale.data = (
layer.w2_weight_scale.data.reshape(g, n, k // 2, 2).view(torch.uint8).transpose(-3, -2)
)

View File

@@ -0,0 +1,164 @@
import torch
from vllm.config import get_current_vllm_config
from vllm.distributed import get_tensor_model_parallel_rank, get_tensor_model_parallel_world_size
from vllm.logger import logger
from vllm_ascend.utils import AscendDeviceType, get_ascend_device_type
from .base import AscendAttentionScheme
from .registry import register_scheme
def _fa_quant_weight_loader(param: torch.Tensor, loaded_weight: torch.Tensor):
"""Weight loader for MLA-based C8 (FAKQuant) models."""
if param.numel() == 1 and loaded_weight.numel() == 1:
param.data.fill_(loaded_weight.item())
else:
tp_rank = get_tensor_model_parallel_rank()
tp_size = get_tensor_model_parallel_world_size()
shard_size = loaded_weight.shape[0] // tp_size
loaded_weight = loaded_weight.narrow(0, shard_size * tp_rank, shard_size)
assert param.size() == loaded_weight.size(), (
"[vllm-ascend/FAKQuant] Attempted to load weight "
f"({loaded_weight.size()}) into parameter ({param.size()}) "
f"when TP size is {tp_size} and TP rank is {tp_rank}."
)
param.data.copy_(loaded_weight)
@register_scheme("FAKQuant", "attention")
class AscendFAQuantAttentionMethod:
def __init__(self):
vllm_config = get_current_vllm_config()
config = vllm_config.model_config.hf_config
self.kv_lora_rank = getattr(config, "kv_lora_rank", 0)
self.qk_rope_head_dim = getattr(config, "qk_rope_head_dim", 0)
def create_weights(self, layer: torch.nn.Module) -> None:
extra_module_names = ["fa_q", "fa_k", "fa_v"]
for name in extra_module_names:
setattr(layer, name, torch.nn.Module())
params_dict = {}
dtype = torch.get_default_dtype()
params_dict["fa_q.scale"] = torch.empty((layer.num_heads, 1), dtype=dtype)
params_dict["fa_k.scale"] = torch.empty((layer.num_kv_heads, 1), dtype=dtype)
params_dict["fa_v.scale"] = torch.empty((layer.num_kv_heads, 1), dtype=dtype)
params_dict["fa_q.offset"] = torch.empty((layer.num_heads, 1), dtype=torch.int8)
params_dict["fa_k.offset"] = torch.empty((layer.num_kv_heads, 1), dtype=torch.int8)
params_dict["fa_v.offset"] = torch.empty((layer.num_kv_heads, 1), dtype=torch.int8)
for name, weight in params_dict.items():
module_name, weight_name = name.rsplit(".", 1)
module = getattr(layer, module_name)
weight_param = torch.nn.Parameter(weight, requires_grad=False)
module.register_parameter(weight_name, weight_param)
# When loading weights, segment them according to TP
weight_param.weight_loader = _fa_quant_weight_loader
def process_weights_after_loading(self, layer: torch.nn.Module) -> None:
fa_k_scale = torch.squeeze(layer.fa_k.scale).unsqueeze(0)
layer.fak_descale_float = torch.nn.Parameter(fa_k_scale.to(torch.float), requires_grad=False)
layer.fak_descale = torch.nn.Parameter(fa_k_scale, requires_grad=False)
if get_ascend_device_type() == AscendDeviceType.A5:
layer.fak_descale_reciprocal = 1.0 / torch.nn.Parameter(fa_k_scale.to(torch.float), requires_grad=False)
else:
layer.fak_descale_reciprocal = 1.0 / torch.nn.Parameter(fa_k_scale, requires_grad=False)
fa_k_offset = torch.squeeze(layer.fa_k.offset).unsqueeze(0)
layer.fak_offset = torch.nn.Parameter(fa_k_offset.to(layer.fak_descale.dtype), requires_grad=False)
repeated_quant_kscale = fa_k_scale.repeat(self.kv_lora_rank)
layer.quant_kscale = repeated_quant_kscale.view(1, self.kv_lora_rank)
layer.quant_kscale = 1.0 / torch.nn.Parameter(layer.quant_kscale.to(torch.float), requires_grad=False)
@register_scheme("INT8_DYNAMIC", "attention")
class AscendSFAQuantAttentionMethod:
def __init__(self):
vllm_config = get_current_vllm_config()
config = vllm_config.model_config.hf_config
self.index_head_dim = config.index_head_dim
def create_weights(self, layer: torch.nn.Module) -> None:
extra_module_names = ["indexer"]
for name in extra_module_names:
setattr(layer, name, torch.nn.Module())
params_dict = {}
params_dict["indexer.q_rot"] = torch.empty((self.index_head_dim, self.index_head_dim), dtype=torch.float32)
params_dict["indexer.k_rot"] = torch.empty((self.index_head_dim, self.index_head_dim), dtype=torch.float32)
for name, weight in params_dict.items():
module_name, weight_name = name.split(".")
module = getattr(layer, module_name)
weight_param = torch.nn.Parameter(weight, requires_grad=False)
module.register_parameter(weight_name, weight_param)
def process_weights_after_loading(self, layer: torch.nn.Module) -> None:
pass
def _c8_kv_scale_weight_loader(param: torch.nn.Parameter, loaded_weight: torch.Tensor) -> None:
"""Weight loader for dense-attention C8 KV cache scales/offsets."""
loaded_weight = loaded_weight.squeeze()
if param.data.shape != loaded_weight.shape:
param.data = loaded_weight.to(param.dtype).clone()
else:
param.data.copy_(loaded_weight)
class AscendC8KVCacheAttentionMethod(AscendAttentionScheme):
"""C8 INT8 KV cache quantization for dense-attention models (e.g. Qwen3)."""
def __init__(self, quant_description: dict, prefix: str):
self.quant_description = quant_description
self.prefix = prefix
vllm_config = get_current_vllm_config()
self.is_kv_producer = False
if vllm_config.kv_transfer_config is not None:
self.is_kv_producer = vllm_config.kv_transfer_config.is_kv_producer
def create_weights(self, layer: torch.nn.Module) -> None:
# Returns int8 if the P node is not a PD detachment node.
if not self.is_kv_producer:
logger.info_once(
"[vllm-ascend/C8_KV] KV cache producer is disabled; setting kv_cache_torch_dtype to torch.int8."
)
layer.kv_cache_torch_dtype = torch.int8
# Upgrade impl to the C8-specific subclass so the C8 forward path is always used.
if hasattr(layer, "impl"):
from vllm_ascend.attention.attention_v1 import AscendC8AttentionBackendImpl
layer.impl.__class__ = AscendC8AttentionBackendImpl
dtype = torch.get_default_dtype()
layer.k_cache_scale = torch.nn.Parameter(torch.ones(1, dtype=dtype), requires_grad=False)
layer.k_cache_scale.weight_loader = _c8_kv_scale_weight_loader
layer.k_cache_offset = torch.nn.Parameter(torch.zeros(1, dtype=dtype), requires_grad=False)
layer.k_cache_offset.weight_loader = _c8_kv_scale_weight_loader
layer.v_cache_scale = torch.nn.Parameter(torch.ones(1, dtype=dtype), requires_grad=False)
layer.v_cache_scale.weight_loader = _c8_kv_scale_weight_loader
layer.v_cache_offset = torch.nn.Parameter(torch.zeros(1, dtype=dtype), requires_grad=False)
layer.v_cache_offset.weight_loader = _c8_kv_scale_weight_loader
def process_weights_after_loading(self, layer: torch.nn.Module) -> None:
layer.k_cache_scale.data = layer.k_cache_scale.data.flatten()
layer.k_cache_offset.data = layer.k_cache_offset.data.flatten()
layer.v_cache_scale.data = layer.v_cache_scale.data.flatten()
layer.v_cache_offset.data = layer.v_cache_offset.data.flatten()
def apply(
self,
layer: torch.nn.Module,
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
kv_cache,
attn_metadata,
attn_type,
scale,
output,
) -> torch.Tensor:
err_msg = (
"[vllm-ascend/C8_KV] AscendC8KVCacheAttentionMethod.apply should "
"not be called. C8 KV cache quantization is handled by the "
"attention backend."
)
raise RuntimeError(err_msg)

View File

@@ -0,0 +1,62 @@
#
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
# This file is a part of the vllm-ascend project.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
#
from typing import Any
# Registry: maps (quant_type, layer_type) -> SchemeClass
_SCHEME_REGISTRY: dict[tuple[str, str], type[Any]] = {}
def register_scheme(quant_type: str, layer_type: str):
"""Decorator to register a quantization scheme.
Args:
quant_type: Quantization type (e.g., "W8A8", "W8A8_DYNAMIC").
layer_type: Layer type (e.g., "linear", "moe").
Returns:
Decorator function that registers the class.
Example:
@register_scheme("W8A8_DYNAMIC", "linear")
class W8A8DynamicLinearScheme(AscendLinearScheme):
...
"""
def decorator(cls: type[Any]) -> type[Any]:
key = (quant_type, layer_type)
if key in _SCHEME_REGISTRY:
raise ValueError(
f"Scheme already registered for {quant_type}/{layer_type}: {_SCHEME_REGISTRY[key].__name__}"
)
_SCHEME_REGISTRY[key] = cls
return cls
return decorator
def get_scheme_class(quant_type: str, layer_type: str) -> type[Any] | None:
"""Get scheme class for given quant_type and layer_type.
Args:
quant_type: Quantization type (e.g., "W8A8", "W8A8_DYNAMIC").
layer_type: Layer type (e.g., "linear", "moe").
Returns:
The registered scheme class, or None if not found.
"""
return _SCHEME_REGISTRY.get((quant_type, layer_type))

View File

@@ -0,0 +1,364 @@
#
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
# This file is a part of the vllm-ascend project.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
#
"""Ascend W4A16 quantization helpers and fused MoE method."""
from collections.abc import Callable
from typing import Any
import torch
import torch_npu
from vllm.config import get_current_vllm_config
from vllm_ascend.ascend_config import get_ascend_config
from vllm_ascend.ascend_forward_context import _EXTRA_CTX
from vllm_ascend.ops.fused_moe.experts_selector import select_experts
from vllm_ascend.ops.fused_moe.moe_runtime_args import build_fused_experts_input
from .base import AscendMoEScheme, QuantType, get_moe_num_logical_experts
from .registry import register_scheme
def unpack_from_int32(
weight: torch.Tensor,
shape: torch.Size,
num_bits: int,
packed_dim: int = 1,
) -> torch.Tensor:
"""Unpacks quantized weights from int32 format back to original bits.
:param weight: The packed int32 tensor containing quantized weights
:param shape: Original shape to restore, defaults to None
:param num_bits: The number of bits used for quantization (<= 8)
:param packed_dim: Dimension along which weights are packed (0 or 1), defaults to 1
:return: Unpacked tensor with int8 dtype after applying offset correction
"""
assert weight.dtype == torch.int32, f"Expecting `weight.dtype` is torch.int32 but got {weight.dtype}."
assert num_bits > 0, f"Expecting `num_bits` should be positive but got {num_bits}."
assert num_bits <= 8, f"Expecting `num_bits` should not be larger than 8 but got {num_bits}."
assert 32 % num_bits == 0, f"Expecting `num_bits` {num_bits} to divide 32 exactly."
assert packed_dim in [0, 1], f"Expecting `packed_dim` is 0 or 1 but got {packed_dim}."
pack_factor = 32 // num_bits
mask = (1 << num_bits) - 1
if packed_dim == 1:
unpacked_weight = torch.zeros(
(weight.shape[0], weight.shape[1] * pack_factor),
device=weight.device,
dtype=torch.int32,
)
for i in range(pack_factor):
unpacked_weight[:, i::pack_factor] = (weight >> (num_bits * i)) & mask
original_row_size = int(shape[1])
unpacked_weight = unpacked_weight[:, :original_row_size]
else:
unpacked_weight = torch.zeros(
(weight.shape[0] * pack_factor, weight.shape[1]),
device=weight.device,
dtype=torch.int32,
)
for i in range(pack_factor):
unpacked_weight[i::pack_factor, :] = (weight >> (num_bits * i)) & mask
original_row_size = int(shape[0])
unpacked_weight = unpacked_weight[:original_row_size, :]
offset = pow(2, num_bits) // 2
unpacked_weight = (unpacked_weight - offset).to(torch.int8)
return unpacked_weight
def pack_to_int32(weight: torch.Tensor) -> torch.Tensor:
"""Packs quantized weights into int32 format for storage.
:param weight: The 3D tensor to pack, must be int8 or int32 dtype
:return: Packed tensor with int32 dtype optimized for storage
"""
assert weight.dim() == 3, (
"Expecting `weight.dim()` is 3 ([expert, output_channel, input_channel] or "
"[expert, input_channel, output_channel]) but got "
f"{weight.dim()}."
)
assert weight.dtype in [torch.int8, torch.int32], (
f"Expecting `weight.dtype` is torch.int8 or torch.int32 but got {weight.dtype}."
)
if weight.dtype == torch.int32:
assert weight.shape[-1] % 8 == 0, "the last dim of weight needs to be divided by 8."
packed_weight = torch_npu.npu_convert_weight_to_int4pack(weight.flatten(0, 1))
packed_weight = packed_weight.view(weight.shape[0], weight.shape[1], -1)
else:
assert weight.shape[-1] % 4 == 0, "the last dim of weight needs to be divided by 4."
packed_weight = weight.view(torch.int32).contiguous()
return packed_weight
@register_scheme("W4A16", "moe")
class AscendW4A16FusedMoEMethod(AscendMoEScheme):
"""FusedMoE method for Ascend W4A16.
This method supports only weights generated by LLM-Compressor, for
example ``moonshotai/Kimi-K2-Thinking``.
Each original routed MoE expert in the checkpoint stores separate
LLM-Compressor tensors. The names below use ``L`` for the layer index and
``E`` for the expert index. For these 4-bit weights, ``pack_factor`` is
8, so one int32 element stores eight 4-bit weight values.
- ``model.layers.L.mlp.experts.E.gate_proj.weight_packed``:
``torch.int32``,
``[moe_intermediate_size, hidden_sizes // pack_factor]``.
- ``model.layers.L.mlp.experts.E.gate_proj.weight_scale``:
``torch.bfloat16``,
``[moe_intermediate_size, hidden_sizes // group_size]``.
- ``model.layers.L.mlp.experts.E.gate_proj.weight_shape``:
``torch.int32``, ``[2]``.
- ``model.layers.L.mlp.experts.E.up_proj.weight_packed``:
``torch.int32``,
``[moe_intermediate_size, hidden_sizes // pack_factor]``.
- ``model.layers.L.mlp.experts.E.up_proj.weight_scale``:
``torch.bfloat16``,
``[moe_intermediate_size, hidden_sizes // group_size]``.
- ``model.layers.L.mlp.experts.E.up_proj.weight_shape``:
``torch.int32``, ``[2]``.
- ``model.layers.L.mlp.experts.E.down_proj.weight_packed``:
``torch.int32``,
``[hidden_sizes, moe_intermediate_size // pack_factor]``.
- ``model.layers.L.mlp.experts.E.down_proj.weight_scale``:
``torch.bfloat16``,
``[hidden_sizes, moe_intermediate_size // group_size]``.
- ``model.layers.L.mlp.experts.E.down_proj.weight_shape``:
``torch.int32``, ``[2]``.
During loading, the gate and up projections are fused into ``w13`` and the
down projection is loaded as ``w2``. In
:meth:`process_weights_after_loading`, weight tensors are unpacked,
transposed into the data layout required by the Ascend fused MoE operator,
and repacked into the int32 dtype. The offset tensors are not loaded from the
checkpoint; they are all-zero tensors constructed because the operator
requires offset inputs.
After :meth:`process_weights_after_loading`, ``apply`` consumes:
- ``w13_weight_packed``: ``torch.int32``,
``[num_experts, hidden_sizes,
2 * moe_intermediate_size // pack_factor]``.
- ``w2_weight_packed``: ``torch.int32``,
``[num_experts, moe_intermediate_size,
hidden_sizes // pack_factor]``.
- ``w13_weight_scale``: ``torch.bfloat16``,
``[num_experts, hidden_sizes // group_size,
2 * moe_intermediate_size]``.
- ``w2_weight_scale``: ``torch.bfloat16``,
``[num_experts, moe_intermediate_size // group_size,
hidden_sizes]``.
- ``w13_weight_offset``: ``torch.bfloat16``, all zeros,
``[num_experts, hidden_sizes // group_size,
2 * moe_intermediate_size]``.
- ``w2_weight_offset``: ``torch.bfloat16``, all zeros,
``[num_experts, moe_intermediate_size // group_size,
hidden_sizes]``.
"""
quant_type: QuantType = QuantType.W4A16
def __init__(self) -> None:
self.num_bits = 4 # dtype = torch.int4
self.pack_factor = 8 # pack 8 of torch.int4 tensors to torch.int32
vllm_config = get_current_vllm_config()
self.group_size = vllm_config.quant_config.quant_description.get("group_size", 32)
self.dynamic_eplb = get_ascend_config().eplb_config.dynamic_eplb
def get_weight(
self,
num_experts: int,
intermediate_size_per_partition: int,
hidden_sizes: int,
params_dtype: torch.dtype,
) -> dict[str, Any]:
assert intermediate_size_per_partition % self.pack_factor == 0, (
f"Expecting `intermediate_size_per_partition` {intermediate_size_per_partition} "
f"can be divided by `pack_factor` {self.pack_factor}"
)
assert hidden_sizes % self.pack_factor == 0, (
f"Expecting `hidden_sizes` {hidden_sizes} can be divided by `pack_factor` {self.pack_factor}"
)
param_dict = {}
param_dict["w13_weight_packed"] = torch.empty(
num_experts, 2 * intermediate_size_per_partition, hidden_sizes // self.pack_factor, dtype=torch.int32
)
param_dict["w2_weight_packed"] = torch.empty(
num_experts, hidden_sizes, intermediate_size_per_partition // self.pack_factor, dtype=torch.int32
)
return param_dict
def get_dynamic_quant_param(
self,
num_experts: int,
intermediate_size_per_partition: int,
hidden_sizes: int,
params_dtype: torch.dtype,
) -> dict[str, Any]:
assert intermediate_size_per_partition % self.group_size == 0, (
f"Expecting `intermediate_size_per_partition` {intermediate_size_per_partition} "
f"can be divided by `group_size` {self.group_size}"
)
assert hidden_sizes % self.group_size == 0, (
f"Expecting `hidden_sizes` {hidden_sizes} can be divided by `group_size` {self.group_size}"
)
param_dict = {}
param_dict["w13_weight_scale"] = torch.empty(
num_experts, 2 * intermediate_size_per_partition, hidden_sizes // self.group_size, dtype=params_dtype
)
param_dict["w2_weight_scale"] = torch.empty(
num_experts, hidden_sizes, intermediate_size_per_partition // self.group_size, dtype=params_dtype
)
param_dict["w13_weight_shape"] = torch.empty(num_experts, 2, dtype=torch.int32)
param_dict["w2_weight_shape"] = torch.empty(num_experts, 2, dtype=torch.int32)
param_dict["w13_weight_offset"] = torch.zeros(
num_experts, 2 * intermediate_size_per_partition, hidden_sizes // self.group_size, dtype=params_dtype
)
param_dict["w2_weight_offset"] = torch.zeros(
num_experts, hidden_sizes, intermediate_size_per_partition // self.group_size, dtype=params_dtype
)
return param_dict
def apply(
self,
layer: torch.nn.Module,
x: torch.Tensor,
router_logits: torch.Tensor,
top_k: int,
renormalize: bool,
use_grouped_topk: bool = False,
num_experts: int = -1,
expert_map: torch.Tensor | None = None,
topk_group: int | None = None,
num_expert_group: int | None = None,
custom_routing_function: Callable | None = None,
scoring_func: str = "softmax",
routed_scaling_factor: float = 1.0,
e_score_correction_bias: torch.Tensor | None = None,
is_prefill: bool = True,
enable_force_load_balance: bool = True,
log2phy: torch.Tensor | None = None,
global_redundant_expert_num: int = 0,
pertoken_scale: Any | None = None,
activation: str = "silu",
apply_router_weight_on_input: bool = False,
mc2_mask: torch.Tensor | None = None,
tid2eid: torch.Tensor | None = None,
) -> torch.Tensor:
num_shared_experts = getattr(layer, "n_shared_experts", 0)
if num_shared_experts is None:
num_shared_experts = 0
num_logical_experts = get_moe_num_logical_experts(
layer,
num_experts,
global_redundant_expert_num=global_redundant_expert_num,
num_shared_experts=num_shared_experts,
)
assert router_logits.shape[1] == num_logical_experts, (
"Number of global experts mismatch (excluding redundancy): "
f"router_logits.shape[1]={router_logits.shape[1]}, num_logical_experts={num_logical_experts}"
)
topk_weights, topk_ids = select_experts(
hidden_states=x,
router_logits=router_logits,
top_k=top_k,
use_grouped_topk=use_grouped_topk,
renormalize=renormalize,
topk_group=topk_group,
num_expert_group=num_expert_group,
custom_routing_function=custom_routing_function,
scoring_func=scoring_func,
routed_scaling_factor=routed_scaling_factor,
e_score_correction_bias=e_score_correction_bias,
num_experts=num_logical_experts,
tid2eid=tid2eid,
)
topk_ids = topk_ids.to(torch.int32)
topk_weights = topk_weights.to(x.dtype)
moe_comm_method = _EXTRA_CTX.moe_comm_method
return moe_comm_method.fused_experts(
fused_experts_input=build_fused_experts_input(
hidden_states=x,
topk_weights=topk_weights,
topk_ids=topk_ids,
w1=layer.w13_weight_packed,
w2=layer.w2_weight_packed,
quant_type=self.quant_type,
dynamic_eplb=self.dynamic_eplb,
expert_map=expert_map,
global_redundant_expert_num=global_redundant_expert_num,
mc2_mask=mc2_mask,
apply_router_weight_on_input=apply_router_weight_on_input,
log2phy=log2phy,
pertoken_scale=pertoken_scale,
activation=activation,
w1_scale=layer.w13_weight_scale,
w2_scale=layer.w2_weight_scale,
w1_offset=layer.w13_weight_offset,
w2_offset=layer.w2_weight_offset,
swiglu_limit=layer.swiglu_limit,
)
)
def process_weights_after_loading(self, layer: torch.nn.Module) -> None:
w13_shape = layer.w13_weight_packed.data.shape
w2_shape = layer.w2_weight_packed.data.shape
unpacked_w13_weight = (
unpack_from_int32(
layer.w13_weight_packed.data.flatten(0, 1),
torch.Size([w13_shape[0] * w13_shape[1], w13_shape[2] * self.pack_factor]),
self.num_bits,
)
.view(w13_shape[0], w13_shape[1], -1)
.transpose(1, 2)
.contiguous()
.int()
)
unpacked_w2_weight = (
unpack_from_int32(
layer.w2_weight_packed.data.flatten(0, 1),
torch.Size([w2_shape[0] * w2_shape[1], w2_shape[2] * self.pack_factor]),
self.num_bits,
)
.view(w2_shape[0], w2_shape[1], -1)
.transpose(1, 2)
.contiguous()
.int()
)
layer.w13_weight_packed.data = pack_to_int32(unpacked_w13_weight)
layer.w2_weight_packed.data = pack_to_int32(unpacked_w2_weight)
layer.w13_weight_scale.data = layer.w13_weight_scale.data.transpose(1, 2).contiguous()
layer.w2_weight_scale.data = layer.w2_weight_scale.data.transpose(1, 2).contiguous()
layer.w13_weight_offset.data = layer.w13_weight_offset.data.transpose(1, 2).contiguous()
layer.w2_weight_offset.data = layer.w2_weight_offset.data.transpose(1, 2).contiguous()

View File

@@ -0,0 +1,215 @@
#
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
# This file is a part of the vllm-ascend project.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
#
"""Ascend W4A16_MXFP4 quantization helpers and fused MoE method."""
from collections.abc import Callable
from typing import Any
import torch
import torch_npu
from vllm.config import CompilationMode, get_current_vllm_config
from vllm.distributed import get_ep_group
from vllm_ascend.ascend_config import get_ascend_config
from vllm_ascend.ascend_forward_context import _EXTRA_CTX
from vllm_ascend.device.mxfp_compat import (
FLOAT8_E8M0FNU_DTYPE,
ensure_mxfp4_moe_available,
)
from vllm_ascend.ops.fused_moe.experts_selector import select_experts
from vllm_ascend.ops.fused_moe.moe_runtime_args import build_fused_experts_input
from .base import AscendMoEScheme, QuantType, get_moe_num_logical_experts
from .registry import register_scheme
# Unpack the weights to FP4 and return them in float32 format
def unpack_uint8_to_fp4_return_float32(packed: torch.Tensor) -> torch.Tensor:
low = packed & 0x0F
high = packed // 16
# The high 4 bits and low 4 bits are arranged alternately, with the low 4 bits in front.
unpacked = torch.stack([low, high], dim=-1).reshape(*packed.shape[:-1], -1)
# A 4-digit integer is mapped to mxfp4 based on its value.
fp4_values = torch.tensor(
[0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0, -0.0, -0.5, -1.0, -1.5, -2.0, -3.0, -4.0, -6.0],
dtype=torch.float32,
device=packed.device,
)
return fp4_values[unpacked.to(torch.long)]
@register_scheme("W4A16_MXFP4", "moe")
class AscendW4A16MXFP4FusedMoEMethod(AscendMoEScheme):
"""FusedMoE method for Ascend W4A16_MXFP4."""
quant_type: QuantType = QuantType.W4A16MXFP4
def __init__(self) -> None:
ensure_mxfp4_moe_available("W4A16_MXFP4 MoE quantization")
self.ep_group = get_ep_group()
vllm_config = get_current_vllm_config()
self.group_size = vllm_config.quant_config.quant_description.get("group_size", 32)
ascend_config = get_ascend_config()
self.use_aclgraph = (
vllm_config.compilation_config.mode == CompilationMode.VLLM_COMPILE
and not vllm_config.model_config.enforce_eager
)
self.dynamic_eplb = ascend_config.eplb_config.dynamic_eplb
def get_weight(
self,
num_experts: int,
intermediate_size_per_partition: int,
hidden_sizes: int,
params_dtype: torch.dtype,
) -> dict[str, Any]:
param_dict = {}
param_dict["w13_weight"] = torch.empty(
num_experts,
2 * intermediate_size_per_partition,
hidden_sizes // 2,
dtype=torch.uint8,
)
param_dict["w2_weight"] = torch.empty(
num_experts,
hidden_sizes,
intermediate_size_per_partition // 2,
dtype=torch.uint8,
)
return param_dict
def get_dynamic_quant_param(
self,
num_experts: int,
intermediate_size_per_partition: int,
hidden_sizes: int,
params_dtype: torch.dtype,
) -> dict[str, Any]:
param_dict = {}
param_dict["w13_weight_scale"] = torch.empty(
num_experts, 2 * intermediate_size_per_partition, hidden_sizes // self.group_size, dtype=torch.uint8
)
param_dict["w2_weight_scale"] = torch.empty(
num_experts, hidden_sizes, intermediate_size_per_partition // self.group_size, dtype=torch.uint8
)
return param_dict
def apply(
self,
layer: torch.nn.Module,
x: torch.Tensor,
router_logits: torch.Tensor,
top_k: int,
renormalize: bool,
use_grouped_topk: bool = False,
num_experts: int = -1,
expert_map: torch.Tensor | None = None,
topk_group: int | None = None,
num_expert_group: int | None = None,
custom_routing_function: Callable | None = None,
scoring_func: str = "softmax",
routed_scaling_factor: float = 1.0,
e_score_correction_bias: torch.Tensor | None = None,
is_prefill: bool = True,
enable_force_load_balance: bool = True,
log2phy: torch.Tensor | None = None,
global_redundant_expert_num: int = 0,
pertoken_scale: Any | None = None,
activation: str = "silu",
apply_router_weight_on_input: bool = False,
mc2_mask: torch.Tensor | None = None,
tid2eid: torch.Tensor | None = None,
) -> torch.Tensor:
num_shared_experts = getattr(layer, "n_shared_experts", 0)
if num_shared_experts is None:
num_shared_experts = 0
num_logical_experts = get_moe_num_logical_experts(
layer,
num_experts,
global_redundant_expert_num=global_redundant_expert_num,
num_shared_experts=num_shared_experts,
)
assert router_logits.shape[1] == num_logical_experts, (
"Number of global experts mismatch (excluding redundancy): "
f"router_logits.shape[1]={router_logits.shape[1]}, num_logical_experts={num_logical_experts}"
)
topk_weights, topk_ids = select_experts(
hidden_states=x,
router_logits=router_logits,
top_k=top_k,
use_grouped_topk=use_grouped_topk,
renormalize=renormalize,
topk_group=topk_group,
num_expert_group=num_expert_group,
custom_routing_function=custom_routing_function,
scoring_func=scoring_func,
e_score_correction_bias=e_score_correction_bias,
routed_scaling_factor=routed_scaling_factor,
num_experts=num_logical_experts,
tid2eid=tid2eid,
)
if enable_force_load_balance:
random_matrix = torch.rand(topk_ids.size(0), num_logical_experts, device=topk_ids.device)
topk_ids = torch.argsort(random_matrix, dim=1)[:, : topk_ids.size(1)].to(topk_ids.dtype)
topk_weights = topk_weights.to(x.dtype)
moe_comm_method = _EXTRA_CTX.moe_comm_method
return moe_comm_method.fused_experts(
fused_experts_input=build_fused_experts_input(
hidden_states=x,
topk_weights=topk_weights,
topk_ids=topk_ids,
w1=layer.w13_weight,
w2=layer.w2_weight,
quant_type=self.quant_type,
dynamic_eplb=self.dynamic_eplb,
expert_map=expert_map,
global_redundant_expert_num=global_redundant_expert_num,
mc2_mask=mc2_mask,
apply_router_weight_on_input=apply_router_weight_on_input,
log2phy=log2phy,
pertoken_scale=pertoken_scale,
activation=activation,
mxfp_act_quant_type=None,
mxfp_weight_quant_type=torch_npu.float4_e2m1fn_x2,
mxfp_scale_dtype=FLOAT8_E8M0FNU_DTYPE,
mxfp_per_token_scale_dtype=None,
mxfp_use_bf16=(x.dtype == torch.bfloat16),
w1_scale=layer.w13_weight_scale,
w2_scale=layer.w2_weight_scale,
swiglu_limit=layer.swiglu_limit,
)
)
def process_weights_after_loading(self, layer):
layer.w13_weight.data = unpack_uint8_to_fp4_return_float32(layer.w13_weight.data)
layer.w13_weight.data = layer.w13_weight.data.transpose(1, 2)
layer.w13_weight.data = torch_npu.npu_format_cast(layer.w13_weight.data, 29, customize_dtype=torch.bfloat16)
layer.w13_weight.data = torch_npu.npu_convert_weight_to_int4pack(layer.w13_weight.data).contiguous()
layer.w2_weight.data = unpack_uint8_to_fp4_return_float32(layer.w2_weight.data)
layer.w2_weight.data = layer.w2_weight.data.transpose(1, 2)
layer.w2_weight.data = torch_npu.npu_format_cast(layer.w2_weight.data, 29, customize_dtype=torch.bfloat16)
layer.w2_weight.data = torch_npu.npu_convert_weight_to_int4pack(layer.w2_weight.data).contiguous()
layer.w13_weight_scale.data = layer.w13_weight_scale.data.transpose(1, 2).contiguous()
layer.w2_weight_scale.data = layer.w2_weight_scale.data.transpose(1, 2).contiguous()

View File

@@ -0,0 +1,171 @@
#
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
# This file is a part of the vllm-ascend project.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
#
import math
from typing import Any
import torch
import torch_npu
from vllm.logger import logger
from .base import AscendLinearScheme
from .registry import register_scheme
KRONECKER_QUANT_MAX_BATCH_SIZE = 32768
def pack_int4_weights(weight_tensor: torch.Tensor) -> torch.Tensor:
"""Pack int4 weights for NPU."""
original_device = weight_tensor.device
weight_tensor_npu = weight_tensor.npu()
weight_int4_packed = torch_npu.npu_convert_weight_to_int4pack(weight_tensor_npu.to(torch.int32), inner_k_tiles=1)
return weight_int4_packed.to(original_device)
def get_decompose_dim(n):
"""Get decomposed dimensions for Kronecker quantization."""
a = int(math.sqrt(n))
if a * a < n:
a += 1
while True:
tmp = a * a - n
b = int(math.sqrt(tmp))
if b * b == tmp:
break
a += 1
return a - b, a + b
# TODO: This function is a temporary workaround for the npu_kronecker_quant operator,
# which has a limitation on the maximum batch size (dim0). This wrapper should be
# removed once the operator supports larger inputs natively.
def batched_kronecker_quant(
x: torch.Tensor,
left_trans: torch.Tensor,
right_trans: torch.Tensor,
clip_ratio: float,
) -> tuple[torch.Tensor, torch.Tensor]:
"""Batched Kronecker quantization with batch size limit handling."""
batch_tokens = x.shape[0]
if batch_tokens <= KRONECKER_QUANT_MAX_BATCH_SIZE:
return torch_npu.npu_kronecker_quant(x, left_trans, right_trans, clip_ratio=clip_ratio, dst_dtype=torch.int32)
x_chunks = torch.split(x, KRONECKER_QUANT_MAX_BATCH_SIZE, dim=0)
processed_chunks = [
torch_npu.npu_kronecker_quant(chunk, left_trans, right_trans, clip_ratio=clip_ratio, dst_dtype=torch.int32)
for chunk in x_chunks
]
quantized_list, scale_list = zip(*processed_chunks)
x_quantized_int4 = torch.cat(quantized_list, dim=0)
activation_scale = torch.cat(scale_list, dim=0)
return x_quantized_int4, activation_scale
@register_scheme("W4A4_FLATQUANT_DYNAMIC", "linear")
class AscendW4A4FlatQuantDynamicLinearMethod(AscendLinearScheme):
"""Linear method for Ascend W4A4_FLATQUANT_DYNAMIC.
This class implements W4A4 quantization with FlatQuant approach and dynamic activation quantization.
- Weight: 4-bit quantization (per-channel) with scale and offset, stored as int8 and packed to int32 during loading
- Activation: 4-bit dynamic quantization with FlatQuant transform matrices (left_trans, right_trans) for
distribution smoothing
- Parameters: clip_ratio for controlling quantization clipping, weight_offset for asymmetric quantization, loaded
from external weights
"""
input_size = 0
def __init__(self):
self.sym = True
def get_weight(self, input_size: int, output_size: int, params_dtype: torch.dtype) -> dict[str, Any]:
if input_size % 8 != 0:
err_msg = f"input_size ({input_size}) must be divisible by 8 for int4 packing"
logger.error(err_msg)
raise ValueError(err_msg)
AscendW4A4FlatQuantDynamicLinearMethod.input_size = input_size
params_dict = {"weight": torch.empty(output_size, input_size, dtype=torch.int8)}
return params_dict
def get_pertensor_param(self, params_dtype: torch.dtype, **kwargs: Any) -> dict[str, Any]:
params_dict = {}
left_trans_dim, right_trans_dim = get_decompose_dim(AscendW4A4FlatQuantDynamicLinearMethod.input_size)
params_dict["left_trans"] = torch.empty(left_trans_dim, left_trans_dim, dtype=params_dtype)
params_dict["right_trans"] = torch.empty(right_trans_dim, right_trans_dim, dtype=params_dtype)
params_dict["clip_ratio"] = torch.empty(1, dtype=torch.float32)
return params_dict
def get_perchannel_param(
self,
output_size: int,
params_dtype: torch.dtype,
) -> dict[str, Any]:
params_dict = {}
params_dict["weight_scale"] = torch.empty(output_size, 1, dtype=torch.float32)
params_dict["weight_offset"] = torch.empty(output_size, 1, dtype=torch.float32)
return params_dict
def apply(
self,
layer: torch.nn.Module,
x: torch.Tensor,
bias: torch.Tensor | None = None,
tp_rank: int | None = 0,
) -> torch.Tensor:
original_dtype = x.dtype
input_shape = x.shape
in_features = input_shape[-1]
left_dim = layer.left_trans.shape[0]
right_dim = layer.right_trans.shape[0]
if left_dim * right_dim != in_features:
err_msg = (
f"FlatQuant transform matrices dimension mismatch: "
f"left_dim({left_dim}) * right_dim({right_dim}) != in_features({in_features})"
)
logger.error(err_msg)
raise ValueError(err_msg)
left_trans_matched = layer.left_trans.to(original_dtype)
right_trans_matched = layer.right_trans.to(original_dtype)
x_reshaped = x.view(-1, left_dim, right_dim)
x_quantized_int4, activation_scale = batched_kronecker_quant(
x_reshaped, left_trans_matched, right_trans_matched, layer.aclnn_clip_ratio
)
x_quantized_reshaped = x_quantized_int4.view(-1, left_dim * right_dim // 8)
pertoken_scale = activation_scale.view(-1).to(torch.float32)
output = torch_npu.npu_quant_matmul(
x_quantized_reshaped,
layer.weight_packed.t(),
layer.weight_scale.view(-1).to(torch.float32),
pertoken_scale=pertoken_scale,
bias=None,
output_dtype=original_dtype,
)
output = output.view(*input_shape[:-1], -1)
if bias is not None:
output = output + bias.to(original_dtype)
return output
def process_weights_after_loading(self, layer):
# NOTE: Currently, w4a4 can't support weight nz
weight_packed = pack_int4_weights(layer.weight.data)
layer.register_parameter("weight_packed", torch.nn.Parameter(weight_packed, requires_grad=False))
del layer.weight
layer.weight_scale.data = layer.weight_scale.data.to(torch.float32)
layer.weight_offset.data = layer.weight_offset.data.to(torch.float32)
layer.left_trans = torch.nn.Parameter(layer.left_trans.data.t().contiguous())
layer.right_trans = torch.nn.Parameter(layer.right_trans.data)
layer.clip_ratio = torch.nn.Parameter(layer.clip_ratio.data.to(torch.float32))
layer.aclnn_clip_ratio = layer.clip_ratio.item()

View File

@@ -0,0 +1,76 @@
#
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
# This file is a part of the vllm-ascend project.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
#
from typing import Any
import torch
import torch_npu
from .base import AscendLinearScheme
from .registry import register_scheme
@register_scheme("W4A4_DYNAMIC", "linear")
class AscendW4A4LaosDynamicLinearMethod(AscendLinearScheme):
"""Linear method for Ascend W4A4_DYNAMIC.
This class implements W4A4 quantization with LAOS approach and dynamic activation quantization.
- Weight: 4-bit quantization (per-channel) with scale and offset, stored as int8.
- Activation: 4-bit dynamic quantization.
"""
def __init__(self):
pass
def get_weight(self, 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
def get_perchannel_param(self, output_size: int, params_dtype: torch.dtype) -> dict[str, Any]:
params_dict = {}
params_dict["weight_scale"] = torch.empty(output_size, 1, dtype=torch.float32)
params_dict["weight_offset"] = torch.empty(output_size, 1, dtype=torch.float32)
return params_dict
def apply(
self,
layer: torch.nn.Module,
x: torch.Tensor,
bias: torch.Tensor | None = None,
tp_rank: int | None = 0,
) -> torch.Tensor:
dtype = x.dtype
x, pertoken_scale = torch_npu.npu_dynamic_quant(x, dst_type=torch.quint4x2)
pertoken_scale = pertoken_scale.reshape(-1, 1)
pertoken_scale = pertoken_scale.squeeze(-1)
output = torch_npu.npu_quant_matmul(
x,
layer.weight.data,
scale=layer.weight_scale.data.view(-1),
pertoken_scale=pertoken_scale,
bias=None,
output_dtype=torch.float16,
)
output = output.to(dtype)
if bias is not None:
output = output + bias.to(dtype)
return output
def process_weights_after_loading(self, layer: torch.nn.Module) -> None:
layer.weight_scale.data = layer.weight_scale.data.to(torch.float32)
layer.weight.data = torch_npu.npu_convert_weight_to_int4pack(layer.weight.data.to(torch.int32))
layer.weight.data = layer.weight.data.transpose(-1, -2)

View File

@@ -0,0 +1,261 @@
#
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
# This file is a part of the vllm-ascend project.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
#
from collections.abc import Callable
from typing import Any
import torch
import torch_npu
from vllm.config import CompilationMode, get_current_vllm_config
from vllm_ascend.ascend_config import get_ascend_config
from vllm_ascend.ascend_forward_context import _EXTRA_CTX
from vllm_ascend.device.mxfp_compat import (
FLOAT8_E8M0FNU_DTYPE,
ensure_mxfp4_linear_available,
ensure_mxfp4_moe_available,
)
from vllm_ascend.ops.fused_moe.experts_selector import select_experts
from vllm_ascend.ops.fused_moe.moe_runtime_args import build_fused_experts_input
from .base import AscendLinearScheme, AscendMoEScheme, QuantType, get_moe_num_logical_experts
from .registry import register_scheme
@register_scheme("W4A4_MXFP4", "linear")
class AscendW4A4MXFP4DynamicLinearMethod(AscendLinearScheme):
"""Linear method for Ascend W4A4_MXFP4 (Microscaling FP4) quantization.
This scheme uses microscaling FP4 quantization with per-group scales.
The activation is dynamically quantized to FP4 with microscaling, and
weights are stored in packed FP4-compatible format with per-group scales.
"""
model_dtype = None
def __init__(self):
ensure_mxfp4_linear_available("W4A4_MXFP4 linear quantization")
vllm_config = get_current_vllm_config()
self.group_size = vllm_config.quant_config.quant_description.get("group_size", 32)
def get_weight(self, input_size: int, output_size: int, params_dtype: torch.dtype) -> dict[str, Any]:
params_dict = {"weight": torch.empty(output_size, input_size // 2, dtype=torch.uint8)}
return params_dict
def get_pergroup_param(
self, input_size: int, output_size: int, params_dtype: torch.dtype, layer_type: str | None = None
) -> dict[str, Any]:
params_dict = {}
params_dict["weight_scale"] = torch.empty(output_size, input_size // self.group_size, dtype=torch.uint8)
return params_dict
def apply(
self,
layer: torch.nn.Module,
x: torch.Tensor,
bias: torch.Tensor | None = None,
tp_rank: int | None = 0,
) -> torch.Tensor:
# reshape x for Qwen VL models
original_shape = x.shape
if x.dim() > 2:
x = x.view(-1, x.shape[-1])
quantized_x, dynamic_scale = torch_npu.npu_dynamic_mx_quant(
x, dst_type=torch_npu.float4_e2m1fn_x2, round_mode="round"
)
pertoken_scale = dynamic_scale
output_dtype = x.dtype
if bias is not None and bias.dtype != torch.float32:
bias = bias.to(torch.float32)
output = torch_npu.npu_quant_matmul(
quantized_x,
layer.weight,
layer.weight_scale,
scale_dtype=FLOAT8_E8M0FNU_DTYPE,
pertoken_scale=pertoken_scale,
pertoken_scale_dtype=FLOAT8_E8M0FNU_DTYPE,
bias=bias,
output_dtype=output_dtype,
x1_dtype=torch_npu.float4_e2m1fn_x2,
x2_dtype=torch_npu.float4_e2m1fn_x2,
group_sizes=[1, 1, self.group_size],
)
# reshape output for Qwen VL models
if len(original_shape) > 2:
output = output.view(*original_shape[:-1], -1)
return output
def process_weights_after_loading(self, layer):
"""Process weights after loading for MXFP4 inference.
This method transforms weights for NPU MXFP4 computation:
- weight: (output_size, input_size) -> (input_size, output_size)
- weight_scale: (n_dim, k_dim) -> (k_dim//2, n_dim, 2)
"""
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)
@register_scheme("W4A4_MXFP4", "moe")
class AscendW4A4MXFP4DynamicFusedMoEMethod(AscendMoEScheme):
"""FusedMoe method for Ascend W4A4_MXFP4."""
model_dtype = None
quant_type: QuantType = QuantType.MXFP4
def __init__(self):
ensure_mxfp4_moe_available("W4A4_MXFP4 MoE quantization")
vllm_config = get_current_vllm_config()
self.group_size = vllm_config.quant_config.quant_description.get("group_size", 32)
ascend_config = get_ascend_config()
self.use_aclgraph = (
vllm_config.compilation_config.mode == CompilationMode.VLLM_COMPILE
and not vllm_config.model_config.enforce_eager
)
self.dynamic_eplb = ascend_config.eplb_config.dynamic_eplb
@staticmethod
def get_weight(
num_experts: int, intermediate_size_per_partition: int, hidden_sizes: int, params_dtype: torch.dtype
) -> dict[str, Any]:
param_dict = {}
param_dict["w13_weight"] = torch.empty(
num_experts, 2 * intermediate_size_per_partition, hidden_sizes // 2, dtype=torch.uint8
)
param_dict["w2_weight"] = torch.empty(
num_experts, hidden_sizes, intermediate_size_per_partition // 2, dtype=torch.uint8
)
return param_dict
def get_dynamic_quant_param(
self, num_experts: int, intermediate_size_per_partition: int, hidden_sizes: int, params_dtype: torch.dtype
) -> dict[str, Any]:
param_dict = {}
param_dict["w13_weight_scale"] = torch.empty(
num_experts, 2 * intermediate_size_per_partition, hidden_sizes // self.group_size, dtype=torch.uint8
)
param_dict["w2_weight_scale"] = torch.empty(
num_experts, hidden_sizes, intermediate_size_per_partition // self.group_size, dtype=torch.uint8
)
return param_dict
def apply(
self,
layer: torch.nn.Module,
x: torch.Tensor,
router_logits: torch.Tensor,
top_k: int,
renormalize: bool,
use_grouped_topk: bool = False,
num_experts: int = -1,
expert_map: torch.Tensor | None = None,
topk_group: int | None = None,
num_expert_group: int | None = None,
custom_routing_function: Callable | None = None,
scoring_func: str = "softmax",
routed_scaling_factor: float = 1.0,
e_score_correction_bias: torch.Tensor | None = None,
is_prefill: bool = True,
enable_force_load_balance: bool = True,
log2phy: torch.Tensor = None,
global_redundant_expert_num: int = 0,
pertoken_scale: Any | None = None,
activation: str = "silu",
apply_router_weight_on_input: bool = False,
mc2_mask: torch.Tensor | None = None,
tid2eid: Any | None = None,
) -> torch.Tensor:
num_shared_experts = getattr(layer, "n_shared_experts", 0)
if num_shared_experts is None:
num_shared_experts = 0
num_logical_experts = get_moe_num_logical_experts(
layer,
num_experts,
global_redundant_expert_num=global_redundant_expert_num,
num_shared_experts=num_shared_experts,
)
assert router_logits.shape[1] == num_logical_experts, "Number of global experts mismatch (excluding redundancy)"
topk_weights, topk_ids = select_experts(
hidden_states=x,
router_logits=router_logits,
top_k=top_k,
use_grouped_topk=use_grouped_topk,
renormalize=renormalize,
topk_group=topk_group,
num_expert_group=num_expert_group,
custom_routing_function=custom_routing_function,
scoring_func=scoring_func,
routed_scaling_factor=routed_scaling_factor,
e_score_correction_bias=e_score_correction_bias,
num_experts=num_logical_experts,
)
# this is a naive implementation for experts load balance so as
# to avoid accumulating too much tokens on a single rank.
# currently it is only activated when doing profile runs.
if enable_force_load_balance:
random_matrix = torch.rand(topk_ids.size(0), num_logical_experts, device=topk_ids.device)
topk_ids = torch.argsort(random_matrix, dim=1)[:, : topk_ids.size(1)].to(topk_ids.dtype)
if x.dtype not in [torch.uint8]:
topk_weights = topk_weights.to(x.dtype)
moe_comm_method = _EXTRA_CTX.moe_comm_method
return moe_comm_method.fused_experts(
fused_experts_input=build_fused_experts_input(
hidden_states=x,
topk_weights=topk_weights,
topk_ids=topk_ids,
w1=layer.w13_weight,
w2=layer.w2_weight,
quant_type=self.quant_type,
dynamic_eplb=self.dynamic_eplb,
expert_map=expert_map,
global_redundant_expert_num=global_redundant_expert_num,
mc2_mask=mc2_mask,
apply_router_weight_on_input=apply_router_weight_on_input,
log2phy=log2phy,
pertoken_scale=pertoken_scale,
activation=activation,
mxfp_act_quant_type=torch_npu.float4_e2m1fn_x2,
mxfp_weight_quant_type=torch_npu.float4_e2m1fn_x2,
mxfp_scale_dtype=FLOAT8_E8M0FNU_DTYPE,
mxfp_per_token_scale_dtype=FLOAT8_E8M0FNU_DTYPE,
mxfp_use_bf16=(x.dtype in [torch.bfloat16, torch.uint8]),
w1_scale=layer.w13_weight_scale,
w2_scale=layer.w2_weight_scale,
)
)
def process_weights_after_loading(self, layer):
g_num, n_size, k_size = layer.w13_weight_scale.shape
layer.w13_weight_scale.data = layer.w13_weight_scale.data.reshape(g_num, n_size, k_size // 2, 2)
g_num, n_size, k_size = layer.w2_weight_scale.shape
layer.w2_weight_scale.data = layer.w2_weight_scale.data.reshape(g_num, n_size, k_size // 2, 2)
# The A5 MXFP4 fused grouped-matmul-swiglu op relies on the
# transpose stride to interpret packed FP4 weights as logical K.
layer.w13_weight.data = layer.w13_weight.data.transpose(1, 2)
layer.w2_weight.data = layer.w2_weight.data.transpose(1, 2)
layer.w13_weight_scale.data = layer.w13_weight_scale.data.transpose(1, 2)
layer.w2_weight_scale.data = layer.w2_weight_scale.data.transpose(1, 2)

View File

@@ -0,0 +1,180 @@
#
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
# This file is a part of the vllm-ascend project.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
#
import math
from typing import Any
import torch
import torch_npu
from vllm.config import get_current_vllm_config
from vllm.distributed import get_tensor_model_parallel_world_size
from vllm.model_executor.layers.linear import RowParallelLinear
from vllm_ascend.device.mxfp_compat import ensure_mxfp4_flatquant_linear_available
from .base import AscendLinearScheme
from .registry import register_scheme
# Maximum supported dimension for Kronecker quantization left_trans_dim and right_trans_dim
MAX_SUPPORT_DIM = 256
def get_decompose_dim(n: int, m: int) -> tuple[int, int]:
"""Get decomposed dimensions for Kronecker quantization.
Args:
n: Dimension to decompose
m: Tensor parallelism size
Returns:
tuple[int, int]: Left decomposed dim, right decomposed dim
Raises:
ValueError: If decomposed dimension exceeds MAX_SUPPORT_DIM
"""
a = int(math.sqrt(n))
if a * a < n:
a += 1
while True:
tmp = a * a - n
b = int(math.sqrt(tmp))
if b * b == tmp:
break
a += 1
if (a + b) > MAX_SUPPORT_DIM:
raise ValueError(
f"Kronecker quantization left_trans_dim and right_trans_dim should be less than {MAX_SUPPORT_DIM}"
)
if (a - b) * m > MAX_SUPPORT_DIM:
return MAX_SUPPORT_DIM, m * n // MAX_SUPPORT_DIM
return a - b, a + b
@register_scheme("W4A4_MXFP4_FLATQUANT", "linear")
class AscendW4A4MXFP4FlatQuantDynamicLinearMethod(AscendLinearScheme):
"""Linear method for Ascend W4A4_MXFP4_FLATQUANT_DYNAMIC."""
def __init__(self):
ensure_mxfp4_flatquant_linear_available("W4A4_MXFP4_FLATQUANT linear quantization")
vllm_config = get_current_vllm_config()
self.group_size = vllm_config.quant_config.quant_description.get("group_size", 32)
self.max_supported_tp = vllm_config.quant_config.quant_description.get("max_supported_tp", 4)
self.tp_size = get_tensor_model_parallel_world_size()
if self.tp_size > self.max_supported_tp:
raise ValueError(
f"For W4A4_MXFP4_FLATQUANT, TP size ({self.tp_size}) is not supported. "
f"Max supported TP size is {self.max_supported_tp}, "
f"according to the max_supported_tp parameter in quant_description."
)
def get_weight(self, input_size: int, output_size: int, params_dtype: torch.dtype) -> dict[str, Any]:
if input_size % 2 != 0:
raise ValueError(f"input_size ({input_size}) must be divisible by 2 for fp4 packing")
self.input_size = input_size
params_dict = {"weight": torch.empty(output_size, input_size // 2, dtype=torch.uint8)}
return params_dict
def get_pertensor_param(self, params_dtype: torch.dtype, **kwargs: Any) -> dict[str, Any]:
params_dict = {}
layer_type = kwargs.get("layer_type")
if layer_type == "row":
origin_size = self.input_size * self.tp_size
_, right_trans_dim = get_decompose_dim(origin_size // self.max_supported_tp, self.max_supported_tp)
left_trans_dim = origin_size // right_trans_dim
else:
left_trans_dim, right_trans_dim = get_decompose_dim(self.input_size, 1)
params_dict["left_trans"] = torch.empty(left_trans_dim, left_trans_dim, dtype=params_dtype)
params_dict["right_trans"] = torch.empty(right_trans_dim, right_trans_dim, dtype=params_dtype)
params_dict["clip_ratio"] = torch.empty(1, dtype=torch.float32)
return params_dict
def get_pergroup_param(
self, input_size: int, output_size: int, params_dtype: torch.dtype, layer_type: str | None = None
) -> dict[str, Any]:
params_dict = {}
params_dict["weight_scale"] = torch.empty(output_size, input_size // self.group_size, dtype=torch.uint8)
return params_dict
def apply(
self,
layer: torch.nn.Module,
x: torch.Tensor,
bias: torch.Tensor | None = None,
tp_rank: int | None = 0,
) -> torch.Tensor:
original_dtype = x.dtype
input_shape = x.shape
in_features = input_shape[-1]
left_dim = layer.left_trans.shape[0]
right_dim = layer.right_trans.shape[0]
if left_dim * right_dim != in_features:
raise ValueError(
f"FlatQuant transform matrices dimension mismatch: "
f"left_dim({left_dim}) * right_dim({right_dim}) != in_features({in_features})"
)
x_reshaped = x.view(-1, left_dim, right_dim)
x_quantized_fp4, pertoken_scale = torch_npu.npu_kronecker_quant(
x_reshaped,
layer.left_trans,
layer.right_trans,
layer.aclnn_clip_ratio,
dst_dtype=torch_npu.float4_e2m1fn_x2,
)
output = torch_npu.npu_quant_matmul(
x_quantized_fp4,
layer.weight,
layer.weight_scale,
scale_dtype=torch_npu.float8_e8m0fnu,
pertoken_scale=pertoken_scale,
pertoken_scale_dtype=torch_npu.float8_e8m0fnu,
bias=bias,
output_dtype=original_dtype,
x1_dtype=torch_npu.float4_e2m1fn_x2,
x2_dtype=torch_npu.float4_e2m1fn_x2,
group_sizes=[1, 1, self.group_size],
)
output = output.view(*input_shape[:-1], -1)
return output
def process_weights_after_loading(self, layer):
if isinstance(layer, RowParallelLinear):
"""
Process weights after loading with TP diagonal block extraction.
This is the special weight loading logic for FlatQuant row parallelism.
"""
left_dim = layer.left_trans.data.shape[0]
# Calculate block sizes
left_block_size = left_dim // layer.tp_size
# Extract diagonal block for current rank
layer.left_trans.data = layer.left_trans.data[
layer.tp_rank * left_block_size : (layer.tp_rank + 1) * left_block_size,
layer.tp_rank * left_block_size : (layer.tp_rank + 1) * left_block_size,
]
layer.weight_scale.data = layer.weight_scale.data.view(-1, layer.weight_scale.shape[-1] // 2, 2)
layer.weight.data = layer.weight.data.transpose(0, 1)
layer.weight_scale.data = layer.weight_scale.data.transpose(0, 1)
layer.left_trans = torch.nn.Parameter(layer.left_trans.data.t().contiguous())
layer.right_trans = torch.nn.Parameter(layer.right_trans.data)
layer.clip_ratio = torch.nn.Parameter(layer.clip_ratio.data.to(torch.float32))
layer.aclnn_clip_ratio = layer.clip_ratio.item()

View File

@@ -0,0 +1,778 @@
#
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
# This file is a part of the vllm-ascend project.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
#
from collections.abc import Callable
from typing import Any
import numpy as np
import torch
import torch_npu
from vllm.config import get_current_vllm_config
from vllm.distributed import get_tensor_model_parallel_world_size
from vllm_ascend.ascend_config import get_ascend_config
from vllm_ascend.ascend_forward_context import _EXTRA_CTX
from vllm_ascend.distributed.parallel_state import get_mc2_group
from vllm_ascend.ops.fused_moe.experts_selector import select_experts
from vllm_ascend.ops.fused_moe.moe_runtime_args import build_fused_experts_input
from vllm_ascend.utils import COMPRESSED_TENSORS_METHOD, maybe_trans_nz
from .base import AscendLinearScheme, AscendMoEScheme, QuantType, get_moe_num_logical_experts
from .registry import register_scheme
@register_scheme("W4A8_DYNAMIC", "linear")
class AscendW4A8DynamicLinearMethod(AscendLinearScheme):
"""Linear method for Ascend W4A8_DYNAMIC.
This method supports only weights quantized by msModelSlim. It supports two
weight layouts, distinguished by ``quant_version`` which comes from
``quant_description["version"]`` in the vLLM quantization config. Version
``"1.0.0"`` is the newer layout: it reduces the checkpoint weight size and
precomputes the ``scale_bias`` offline, reducing weight loading time.
The names below use ``linear`` as the checkpoint prefix of a linear layer,
``input_size`` as the logical input dimension, ``output_size`` as the
logical output dimension, and ``group_size`` as the number of input
channels per weight quantization group.
For ``quant_version != "1.0.0"``, the original linear weights are:
- ``linear.weight``: ``torch.int8``, ``[output_size, input_size]``.
Each int8 element stores one 4-bit weight value.
- ``linear.weight_scale``: ``params_dtype``, ``[output_size, 1]``.
- ``linear.weight_offset``: ``params_dtype``, ``[output_size, 1]``.
- ``linear.weight_scale_second``: ``params_dtype``,
``[output_size, input_size // group_size]``.
- ``linear.weight_offset_second``: ``torch.int64``,
``[output_size, input_size // group_size]``.
For ``quant_version == "1.0.0"``, the original linear weights are:
- ``linear.weight``: ``torch.int8``, ``[output_size // 2, input_size]``.
Each int8 element stores two packed 4-bit weight values along the output
dimension.
- ``linear.weight_scale``: ``params_dtype``, ``[output_size, 1]``.
- ``linear.weight_offset``: ``params_dtype``, ``[output_size, 1]``.
- ``linear.weight_scale_second``: ``params_dtype``,
``[output_size, input_size // group_size]``.
- ``linear.weight_offset_second``: ``torch.int64``,
``[output_size, input_size // group_size]``.
- ``linear.scale_bias``: ``torch.float32``, ``[output_size, 1]`` for
column-parallel linear layers and ``[output_size, 16]`` for
row-parallel linear layers.
In :meth:`process_weights_after_loading`, ``linear.weight`` is transposed
from ``[output, input]`` to the operator-oriented ``[input, output]``
layout. Old-version weights are converted with
``torch_npu.npu_convert_weight_to_int4pack``; new-version weights are
already packed as int4 pairs in int8 storage and are reinterpreted as int32
by grouping four int8 values.
After processing, ``torch_npu.npu_weight_quant_batchmatmul`` is called with
``weight`` as ``torch.int32`` in the operator-required packed layout
with shape ``[input_size, output_size // 8]`` and
``antiquant_scale`` as ``weight_scale * weight_scale_second`` converted to
``x.dtype`` with shape ``[input_size // group_size, output_size]``.
"""
def __init__(self):
vllm_config = get_current_vllm_config()
self.group_size = vllm_config.quant_config.quant_description.get("group_size", 256)
quant_version = vllm_config.quant_config.quant_description.get("version", "0")
self.new_quant_version = quant_version == "1.0.0"
self.tp_size = get_tensor_model_parallel_world_size()
def get_weight(self, input_size: int, output_size: int, params_dtype: torch.dtype) -> dict[str, Any]:
"""Create weight parameters.
For new quantization version (double int4 pack into int8), the output dimension
is compressed by factor 2 (e.g., [2048, 3072] -> [1024, 3072]). The returned
dict includes "_packed_dim" and "_packed_factor" for vLLM's weight loader.
"""
params_dict = {}
if self.new_quant_version:
# double int4 pack into int8: output dimension is compressed
pack_factor = 2
actual_output_size = output_size // pack_factor
params_dict["weight"] = torch.empty(actual_output_size, input_size, dtype=torch.int8)
# Add packing information for vLLM's weight_loader
params_dict["_packed_dim"] = 0
params_dict["_packed_factor"] = pack_factor
else:
params_dict["weight"] = torch.empty(output_size, input_size, dtype=torch.int8)
return params_dict
def get_pergroup_param(
self, input_size: int, output_size: int, params_dtype: torch.dtype, layer_type: str | None = None
) -> dict[str, Any]:
"""Create per-group quantization parameters."""
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)
params_dict["weight_scale_second"] = torch.empty(output_size, input_size // self.group_size, dtype=params_dtype)
params_dict["weight_offset_second"] = torch.empty(
output_size, input_size // self.group_size, dtype=params_dtype
)
# NOTE: In w4a8 quantization implementation,
# for down_proj and o_proj(layer_type == "row") scale_bias shape is [output_size, 16],
# others are [output_size, 1]
if self.new_quant_version:
scale_bias_dim = 16 if layer_type == "row" else 1
params_dict["scale_bias"] = torch.empty(output_size, scale_bias_dim, dtype=torch.float32)
return params_dict
@staticmethod
def process_scale_second(
weight: torch.Tensor, scale: torch.Tensor, per_group_scale: torch.Tensor, is_new_quant: bool = False
):
"""Process the scale for second-level quantization.
Args:
weight: weight tensor [k, n] (in new version, n is already compressed to n/2)
scale: first-level quantization scale [output_size]
per_group_scale: second-level per-group quantization scale [group_num, n_scale]
is_new_quant: whether it's the new quantization version (weight already compressed)
Returns:
(antiquant_scale, bias): dequantization scale and bias (bias=None for new version)
"""
k, n = weight.shape
group_num, n_scale = per_group_scale.shape
if is_new_quant:
# Restore logical dimension for compressed weight
n = n * 2
bias = None
if not is_new_quant:
weight_high = weight.to(torch.float32).reshape(group_num, -1, n) * per_group_scale.reshape(group_num, 1, n)
weight_high = weight_high.reshape(k, n)
bias = 8 * (weight_high.to(torch.float32) * scale).sum(dim=0)
# NOTE: scale_bias is not used currently
# because in msmodelslim w4a8 uses symmetric quantization
# TODO: support potential future asymmetric quantization
antiquant_scale = (scale * per_group_scale).reshape(group_num, n)
return antiquant_scale.npu(), bias
def apply(
self,
layer: torch.nn.Module,
x: torch.Tensor,
bias: torch.Tensor | None = None,
tp_rank: int | None = None,
) -> torch.Tensor:
# NOTE: activation `x` is not quantized
return torch_npu.npu_weight_quant_batchmatmul(
x,
layer.weight,
antiquant_scale=layer.weight_scale_second.to(x.dtype),
antiquant_group_size=self.group_size,
)
def process_weights_after_loading(self, layer: torch.nn.Module):
layer.weight.data = layer.weight.data.transpose(0, 1).contiguous()
layer.weight.data = maybe_trans_nz(layer.weight.data)
layer.weight_scale.data = layer.weight_scale.data.flatten().to(torch.float32)
layer.weight_offset.data = layer.weight_offset.data.flatten()
layer.weight_scale_second.data, scale_bias = self.process_scale_second(
layer.weight.data,
layer.weight_scale.data,
layer.weight_scale_second.data.transpose(0, 1).contiguous(),
is_new_quant=self.new_quant_version,
)
if self.new_quant_version:
# Process the loaded data based on layer type
if hasattr(layer, "scale_bias"):
if layer.scale_bias.data.shape[1] == 1:
layer.scale_bias.data = layer.scale_bias.data.flatten()
else:
layer.scale_bias.data = layer.scale_bias.data.contiguous()
else:
if scale_bias is not None:
param = torch.nn.Parameter(scale_bias, requires_grad=False)
layer.register_parameter("weight_scale_bias", param)
# Convert to NPU-specific int4pack format
if self.new_quant_version:
# weights on disk are already in packed int4 format
# pack 4 int8(int4*2) to int32
assert layer.weight.data.shape[-1] % 4 == 0, (
f"the last dim of weight needs to be divided by 4 but got shape {layer.weight.data.shape}"
)
layer.weight.data = layer.weight.data.view(torch.int32).contiguous()
else:
# weights are not compressed
# need to be packed via npu_convert_weight_to_int4pack
layer.weight.data = torch_npu.npu_convert_weight_to_int4pack(layer.weight.data.to(torch.int32))
@register_scheme("W4A8_DYNAMIC", "moe")
class AscendW4A8DynamicFusedMoEMethod(AscendMoEScheme):
"""FusedMoE method for Ascend W4A8_DYNAMIC.
This method supports four MoE weight formats: three generated by
msModelSlim and one generated by LLM-Compressor. The LLM-Compressor path
is selected when ``ascend_quant_method`` in ``quant_description`` is
``COMPRESSED_TENSORS_METHOD``. Otherwise, the msModelSlim path is used.
msModelSlim layouts are first distinguished by
``quant_description["version"] == "1.0.0"``; for version ``"1.0.0"``,
``group_size == 0`` selects per-channel weight quantization and
``group_size > 0`` selects per-group weight quantization.
The names below use ``L`` for the layer index, ``E`` for the expert index,
``num_experts`` for the routed expert count, ``hidden_sizes`` for the
hidden dimension, ``moe_intermediate_size`` for the expert intermediate
dimension, ``group_size`` for per-group weight quantization, and
``tp_size`` for tensor parallel size.
Original MoE layer weights generated by msModelSlim with
``quant_version != "1.0.0"``:
- ``model.layers.L.mlp.experts.E.gate_proj.weight``:
``torch.int8``, ``[moe_intermediate_size, hidden_sizes]``.
- ``model.layers.L.mlp.experts.E.up_proj.weight``:
``torch.int8``, ``[moe_intermediate_size, hidden_sizes]``.
- ``model.layers.L.mlp.experts.E.down_proj.weight``:
``torch.int8``, ``[hidden_sizes, moe_intermediate_size]``.
- Each linear also has ``weight_scale`` and ``weight_offset``:
``torch.float32``, ``[out_features, 1]``.
- Each linear also has ``weight_scale_second`` and
``weight_offset_second``. The ``weight_scale_second`` dtype is
``torch.float32`` and the ``weight_offset_second`` dtype is
``torch.int64``; both use shape
``[out_features, in_features // group_size]``.
Original MoE layer weights generated by msModelSlim with
``quant_version == "1.0.0"`` and per-group quantization:
- Compared with the previous msModelSlim layout, ``weight`` stores two
packed 4-bit values in each int8 element along the output dimension.
Therefore ``gate_proj.weight`` and ``up_proj.weight`` are
``torch.int8`` with shape
``[moe_intermediate_size // 2, hidden_sizes]``, and
``down_proj.weight`` is ``torch.int8`` with shape
``[hidden_sizes // 2, moe_intermediate_size]``.
- Each linear additionally has ``scale_bias``: ``torch.float32``,
``[moe_intermediate_size, 1]`` for ``gate_proj`` and ``up_proj``, and
``[hidden_sizes, 16 // tp_size]`` for ``down_proj``.
Original MoE layer weights generated by msModelSlim with
``quant_version == "1.0.0"`` and per-channel quantization:
- ``weight`` has the same packed shape as the previous msModelSlim
``1.0.0`` per-group layout.
- ``weight_scale`` and ``weight_offset`` are per-channel tensors:
``torch.float32``, ``[out_features, 1]``. There are no
``weight_scale_second`` or ``weight_offset_second`` tensors.
- Each linear also has ``scale_bias``: ``torch.float32``,
``[moe_intermediate_size, 1]`` for ``gate_proj`` and ``up_proj``, and
``[hidden_sizes, 16 // tp_size]`` for ``down_proj``.
Original MoE layer weights generated by LLM-Compressor:
- ``model.layers.L.mlp.experts.E.gate_proj.weight``:
``torch.int8``, ``[moe_intermediate_size, hidden_sizes]``.
- ``model.layers.L.mlp.experts.E.up_proj.weight``:
``torch.int8``, ``[moe_intermediate_size, hidden_sizes]``.
- ``model.layers.L.mlp.experts.E.down_proj.weight``:
``torch.int8``, ``[hidden_sizes, moe_intermediate_size]``.
- Each linear also has ``weight_scale``: ``torch.bfloat16``,
``[out_features, in_features // group_size]`` for group quantization, or
``[out_features, 1]`` for channel quantization.
During loading, ``gate_proj`` and ``up_proj`` are fused into ``w13`` and
``down_proj`` is loaded as ``w2``. Before
:meth:`process_weights_after_loading`, their logical shapes are:
- msModelSlim old: ``w13_weight`` ``torch.int8``,
``[num_experts, 2 * moe_intermediate_size, hidden_sizes]``; and
``w2_weight`` ``torch.int8``,
``[num_experts, hidden_sizes, moe_intermediate_size]``.
- msModelSlim ``1.0.0`` per-group and per-channel: ``w13_weight`` ``torch.int8``,
``[num_experts, moe_intermediate_size, hidden_sizes]``; and
``w2_weight`` ``torch.int8``,
``[num_experts, hidden_sizes // 2, moe_intermediate_size]``.
- LLM-Compressor: ``w13_weight`` ``torch.int8``,
``[num_experts, 2 * moe_intermediate_size, hidden_sizes]``; and
``w2_weight`` ``torch.int8``,
``[num_experts, hidden_sizes, moe_intermediate_size]``.
After processing, ``apply`` passes these tensors to the fused MoE operator:
- Shared by all formats:
``w13_weight``: ``torch.int32``,
``[num_experts, hidden_sizes, moe_intermediate_size // 4]``.
``w2_weight``: ``torch.int32``,
``[num_experts, moe_intermediate_size, hidden_sizes // 8]``.
``w13_scale_bias``: ``torch.float32``, ``[num_experts, 2 * moe_intermediate_size]``.
``w2_scale_bias``: ``torch.float32``, ``[num_experts, hidden_sizes]``.
- per-group:
``w13_weight_scale``: ``torch.int64``,
``[num_experts, hidden_sizes // group_size,
2 * moe_intermediate_size]``.
``w2_weight_scale``: ``torch.int64``,
``[num_experts, moe_intermediate_size // group_size, hidden_sizes]``.
- per-channel:
``w13_weight_scale``: ``torch.int64``,
``[num_experts, 2 * moe_intermediate_size]``.
``w2_weight_scale``: ``torch.int64``,
``[num_experts, 1, hidden_sizes]``.
"""
# Declare the quantization type for this scheme
quant_type: QuantType = QuantType.W4A8
def __init__(self):
self.supports_eplb = True
vllm_config = get_current_vllm_config()
self.group_size = vllm_config.quant_config.quant_description.get("group_size", 256)
# NOTE: the weights are quantized from bf16 to int4 through a per-channel quantization process
self.is_per_channel_weight = self.group_size == 0
quant_version = vllm_config.quant_config.quant_description.get("version", "0")
# NOTE: new quantize weights: 2 int4 pack into int8
self.new_quant_version = quant_version == "1.0.0"
self.quant_method = vllm_config.quant_config.quant_description.get("ascend_quant_method", "")
if self.quant_method == COMPRESSED_TENSORS_METHOD:
self.weight_strategy = vllm_config.quant_config.quant_description.get("weight_strategy", "group")
self.tp_size = (
1 if vllm_config.parallel_config.enable_expert_parallel else get_tensor_model_parallel_world_size()
)
self.dynamic_eplb = get_ascend_config().eplb_config.dynamic_eplb
if self.new_quant_version and self.tp_size > 16:
raise ValueError("The current weight does not support moe part tp>16.")
try:
device_group = get_mc2_group().device_group
# TODO: Try local_rank = ep_group.rank_in_group
local_rank = torch.distributed.get_rank(group=device_group)
backend = device_group._get_backend(torch.device("npu"))
self.moe_all_to_all_group_name = backend.get_hccl_comm_name(local_rank)
except AttributeError:
self.moe_all_to_all_group_name = ""
def get_weight(
self, num_experts: int, intermediate_size_per_partition: int, hidden_sizes: int, params_dtype: torch.dtype
) -> dict[str, Any]:
if self.quant_method == COMPRESSED_TENSORS_METHOD:
return self.get_weight_compressed_tensors(
num_experts, intermediate_size_per_partition, hidden_sizes, params_dtype
)
else:
return self.get_weight_modelslim(num_experts, intermediate_size_per_partition, hidden_sizes, params_dtype)
def get_weight_compressed_tensors(
self, num_experts: int, intermediate_size_per_partition: int, hidden_sizes: int, params_dtype: torch.dtype
) -> dict[str, Any]:
param_dict = {}
E = num_experts
H = hidden_sizes
IN = intermediate_size_per_partition
param_dict["w13_weight"] = torch.empty(E, 2 * IN, H, dtype=torch.int8)
param_dict["w2_weight"] = torch.empty(E, H, IN, dtype=torch.int8)
return param_dict
def get_weight_modelslim(
self, num_experts: int, intermediate_size_per_partition: int, hidden_sizes: int, params_dtype: torch.dtype
) -> dict[str, Any]:
param_dict = {}
if self.new_quant_version:
w13_output_size = intermediate_size_per_partition
w2_output_size = hidden_sizes // 2
else:
w13_output_size = 2 * intermediate_size_per_partition
w2_output_size = hidden_sizes
param_dict["w13_weight"] = torch.empty(num_experts, w13_output_size, hidden_sizes, dtype=torch.int8)
param_dict["w2_weight"] = torch.empty(
num_experts, w2_output_size, intermediate_size_per_partition, dtype=torch.int8
)
return param_dict
def get_dynamic_quant_param(
self, num_experts: int, intermediate_size_per_partition: int, hidden_sizes: int, params_dtype: torch.dtype
) -> dict[str, Any]:
if self.quant_method == COMPRESSED_TENSORS_METHOD:
return self.get_dynamic_quant_param_compressed_tensors(
num_experts, intermediate_size_per_partition, hidden_sizes, params_dtype
)
else:
return self.get_dynamic_quant_param_modelslim(
num_experts, intermediate_size_per_partition, hidden_sizes, params_dtype
)
def get_dynamic_quant_param_compressed_tensors(
self, num_experts: int, intermediate_size_per_partition: int, hidden_sizes: int, params_dtype: torch.dtype
) -> dict[str, Any]:
param_dict = {}
E = num_experts
H = hidden_sizes
IN = intermediate_size_per_partition
g = self.group_size
# Per-row scale columns
def _n_scale_cols(in_features: int) -> int:
return 1 if g <= 0 else (in_features // g)
param_dict["w13_weight_scale"] = torch.empty(E, 2 * IN, _n_scale_cols(H), dtype=torch.bfloat16)
param_dict["w2_weight_scale"] = torch.empty(E, H, _n_scale_cols(IN), dtype=torch.bfloat16)
return param_dict
def get_dynamic_quant_param_modelslim(
self, num_experts: int, intermediate_size_per_partition: int, hidden_sizes: int, params_dtype: torch.dtype
) -> dict[str, Any]:
param_dict = {}
param_dict["w13_weight_scale"] = torch.empty(
num_experts, 2 * intermediate_size_per_partition, 1, dtype=torch.float32
)
param_dict["w13_weight_offset"] = torch.empty(
num_experts, 2 * intermediate_size_per_partition, 1, dtype=torch.float32
)
param_dict["w2_weight_scale"] = torch.empty(num_experts, hidden_sizes, 1, dtype=torch.float32)
param_dict["w2_weight_offset"] = torch.empty(num_experts, hidden_sizes, 1, dtype=torch.float32)
if not self.is_per_channel_weight:
param_dict["w13_weight_scale_second"] = torch.empty(
num_experts, 2 * intermediate_size_per_partition, hidden_sizes // self.group_size, dtype=torch.float32
)
param_dict["w13_weight_offset_second"] = torch.empty(
num_experts, 2 * intermediate_size_per_partition, hidden_sizes // self.group_size, dtype=torch.float32
)
param_dict["w2_weight_scale_second"] = torch.empty(
num_experts, hidden_sizes, intermediate_size_per_partition // self.group_size, dtype=torch.float32
)
param_dict["w2_weight_offset_second"] = torch.empty(
num_experts, hidden_sizes, intermediate_size_per_partition // self.group_size, dtype=torch.float32
)
if self.new_quant_version:
param_dict["w13_scale_bias"] = torch.empty(
num_experts, 2 * intermediate_size_per_partition, 1, dtype=torch.float32
)
param_dict["w2_scale_bias"] = torch.empty(
num_experts, hidden_sizes, 16 // self.tp_size, dtype=torch.float32
)
return param_dict
def apply(
self,
layer: torch.nn.Module,
x: torch.Tensor,
router_logits: torch.Tensor,
top_k: int,
renormalize: bool,
use_grouped_topk: bool = False,
num_experts: int = -1,
expert_map: torch.Tensor | None = None,
topk_group: int | None = None,
num_expert_group: int | None = None,
custom_routing_function: Callable | None = None,
scoring_func: str = "softmax",
routed_scaling_factor: float = 1.0,
e_score_correction_bias: torch.Tensor | None = None,
is_prefill: bool = True,
enable_force_load_balance: bool = False,
log2phy: torch.Tensor | None = None,
global_redundant_expert_num: int = 0,
pertoken_scale: torch.Tensor | None = None,
activation: str = "silu",
apply_router_weight_on_input: bool = False,
mc2_mask: torch.Tensor | None = None,
tid2eid: torch.Tensor | None = None,
) -> torch.Tensor:
num_shared_experts = getattr(layer, "n_shared_experts", 0)
if num_shared_experts is None:
num_shared_experts = 0
num_logical_experts = get_moe_num_logical_experts(
layer,
num_experts,
global_redundant_expert_num=global_redundant_expert_num,
num_shared_experts=num_shared_experts,
)
assert router_logits.shape[1] == num_logical_experts, (
"Number of global experts mismatch (excluding redundancy): "
f"router_logits.shape[1]={router_logits.shape[1]}, num_logical_experts={num_logical_experts}"
)
# NOTE: now npu_moe_gating_top_k can only support `group_count=256` pattern
topk_weights, topk_ids = select_experts(
hidden_states=x,
router_logits=router_logits,
top_k=top_k,
use_grouped_topk=use_grouped_topk,
renormalize=renormalize,
topk_group=topk_group,
num_expert_group=num_expert_group,
custom_routing_function=custom_routing_function,
scoring_func=scoring_func,
routed_scaling_factor=routed_scaling_factor,
e_score_correction_bias=e_score_correction_bias,
num_experts=num_logical_experts,
tid2eid=tid2eid,
)
# this is a naive implementation for experts load balance so as
# to avoid accumulating too much tokens on a single rank.
# currently it is only activated when doing profile runs.
if enable_force_load_balance:
random_matrix = torch.rand(topk_ids.size(0), num_logical_experts, device=topk_ids.device)
topk_ids = torch.argsort(random_matrix, dim=1)[:, : topk_ids.size(1)].to(topk_ids.dtype)
topk_weights = topk_weights.to(x.dtype)
if self.dynamic_eplb:
w1 = [i.view(torch.int32) for i in layer.w13_weight_list]
w1_scale = layer.w13_weight_scale_list
w2 = [i.view(torch.int32) for i in layer.w2_weight_list]
w2_scale = layer.w2_weight_scale_list
w1_scale_bias = layer.w13_scale_bias_list
w2_scale_bias = layer.w2_scale_bias_list
else:
w1 = [layer.w13_weight]
w1_scale = [layer.w13_weight_scale]
w2 = [layer.w2_weight]
w2_scale = [layer.w2_weight_scale]
w1_scale_bias = [layer.w13_scale_bias.detach()] if hasattr(layer, "w13_scale_bias") else None
w2_scale_bias = [layer.w2_scale_bias.detach()] if hasattr(layer, "w2_scale_bias") else None
moe_comm_method = _EXTRA_CTX.moe_comm_method
return moe_comm_method.fused_experts(
fused_experts_input=build_fused_experts_input(
hidden_states=x,
topk_weights=topk_weights,
topk_ids=topk_ids,
w1=w1,
w2=w2,
quant_type=self.quant_type,
dynamic_eplb=self.dynamic_eplb,
expert_map=expert_map,
global_redundant_expert_num=global_redundant_expert_num,
mc2_mask=mc2_mask,
apply_router_weight_on_input=apply_router_weight_on_input,
log2phy=log2phy,
pertoken_scale=pertoken_scale,
activation=activation,
w1_scale=w1_scale,
w2_scale=w2_scale,
w1_scale_bias=w1_scale_bias,
w2_scale_bias=w2_scale_bias,
is_per_channel_weight=self.is_per_channel_weight,
swiglu_limit=layer.swiglu_limit,
)
)
def process_scale(self, weight: torch.Tensor, scale, per_group_scale):
scale = scale.transpose(1, 2).contiguous()
if self.is_per_channel_weight:
scale_np = scale.cpu().numpy()
scale_np.dtype = np.uint32
scale_uint64_tensor = torch.from_numpy(scale_np.astype(np.int64)).npu()
return scale_uint64_tensor, None
per_group_scale = per_group_scale.transpose(1, 2).contiguous()
group_num, k, n = weight.shape
# the weight of the new version is reduced by half by pack n, so it needs to be restored
if self.new_quant_version:
n = n * 2
per_group_scale = per_group_scale.reshape(group_num, -1, n)
group_num, quantgroup_num, n = per_group_scale.shape
bias = None
if not self.new_quant_version:
weight_high = weight.to(torch.float32).reshape(
[group_num, quantgroup_num, -1, n]
) * per_group_scale.reshape([group_num, quantgroup_num, 1, n])
weight_high = weight_high.reshape([group_num, k, n])
bias = 8 * (weight_high.to(torch.float32) * scale).sum(axis=1)
scale_fp32 = (scale * per_group_scale).to(torch.float16).to(torch.float32)
scale_fp32_np = scale_fp32.cpu().numpy()
scale_fp32_np.dtype = np.uint32
sscale_uint64 = np.zeros((group_num, quantgroup_num, n * 2), dtype=np.uint32)
sscale_uint64[..., ::2] = scale_fp32_np
sscale_uint64_buffer = np.frombuffer(sscale_uint64.tobytes(), dtype=np.int64).copy()
sscale_uint64_tensor = torch.from_numpy(sscale_uint64_buffer).reshape(group_num, quantgroup_num, n)
sscale_uint64_tensor = sscale_uint64_tensor.npu()
return sscale_uint64_tensor, bias
def update_bias(self, layer, w13_bias, w2_bias):
if self.new_quant_version:
layer.w13_scale_bias.data = layer.w13_scale_bias.data.transpose(1, 2).contiguous().sum(axis=1)
layer.w2_scale_bias.data = layer.w2_scale_bias.data.transpose(1, 2).contiguous().sum(axis=1)
else:
w13_scale_bias = torch.nn.Parameter(w13_bias, requires_grad=False)
layer.register_parameter("w13_scale_bias", w13_scale_bias)
w2_scale_bias = torch.nn.Parameter(w2_bias, requires_grad=False)
layer.register_parameter("w2_scale_bias", w2_scale_bias)
def pack_to_int32(self, weight: torch.Tensor):
if self.new_quant_version or self.quant_method == COMPRESSED_TENSORS_METHOD:
# pack 4 int8(int4*2) to int32, because in pytorch, we need to use int32 to represent int4
assert weight.shape[-1] % 4 == 0, (
f"the last dim of weight needs to be divided by 4 but got shape {weight.shape}"
)
return weight.view(torch.int32).contiguous()
else:
return torch_npu.npu_quantize(
weight.to(torch.float32), torch.tensor([1.0]).npu(), None, torch.quint4x2, -1, False
)
def pack_int4_to_int8(self, weight: torch.Tensor) -> torch.Tensor:
shape = weight.shape
weight = weight.reshape(-1, 2)
weight0 = weight[:, :1]
weight1 = weight[:, 1:]
weight1_4 = torch.bitwise_left_shift(weight1, 4)
weight2_4 = weight0 & 0b00001111
weight_add = torch.bitwise_or(weight1_4, weight2_4)
# The clone() call is used to break the view chain
return weight_add.reshape(shape[:-1] + (shape[-1] // 2,)).clone()
@staticmethod
def maybe_squeeze_per_channel_weight_scale(scale: torch.Tensor) -> torch.Tensor:
if scale.dim() > 1 and scale.shape[1] == 1:
return scale.squeeze(1)
return scale
def process_weights_after_loading(self, layer):
if self.quant_method == COMPRESSED_TENSORS_METHOD:
self.process_weights_after_loading_compressed_tensors(layer)
else:
self.process_weights_after_loading_modelslim(layer)
def process_weights_after_loading_compressed_tensors(self, layer):
layer.w13_weight.data = layer.w13_weight.data.transpose(1, 2).contiguous()
layer.w2_weight.data = layer.w2_weight.data.transpose(1, 2).contiguous()
def process_scale_compressed_tensors(scale: torch.Tensor, squeeze: bool = True):
scale = scale.transpose(1, 2).to(torch.float32).contiguous()
scale_np = scale.cpu().numpy()
scale_np.dtype = np.uint32
scale_uint64_tensor = torch.from_numpy(scale_np.astype(np.int64)).npu()
if self.is_per_channel_weight and squeeze:
return self.maybe_squeeze_per_channel_weight_scale(scale_uint64_tensor)
return scale_uint64_tensor
def update_bias_compressed_tensors(weight: torch.Tensor, scale: torch.Tensor, strategy: str):
group_num, k, n = weight.shape
scale = scale.transpose(1, 2).contiguous()
scale = scale.reshape(group_num, -1, n)
group_num, quantgroup_num, n = scale.shape
bias = None
if strategy == "group":
tmp = weight.to(torch.float32).reshape([group_num, quantgroup_num, -1, n]) * scale.reshape(
[group_num, quantgroup_num, 1, n]
)
tmp = tmp.reshape([group_num, k, n])
bias = 8 * tmp.sum(axis=1)
elif strategy == "channel":
bias = 8 * (weight.to(torch.float32) * scale).sum(axis=1)
else:
raise ValueError(f"Unsupported weight strategy: {strategy}")
return bias
w13_bias = update_bias_compressed_tensors(
layer.w13_weight.data, layer.w13_weight_scale.data, self.weight_strategy
)
w2_bias = update_bias_compressed_tensors(layer.w2_weight.data, layer.w2_weight_scale.data, self.weight_strategy)
layer.w13_weight_scale.data = process_scale_compressed_tensors(layer.w13_weight_scale.data)
# To use torch_npu.npu_grouped_matmul, keep w2_weigh_scale unsqueezed
layer.w2_weight_scale.data = process_scale_compressed_tensors(layer.w2_weight_scale.data, False)
w13_scale_bias = torch.nn.Parameter(w13_bias, requires_grad=False)
layer.register_parameter("w13_scale_bias", w13_scale_bias)
w2_scale_bias = torch.nn.Parameter(w2_bias, requires_grad=False)
layer.register_parameter("w2_scale_bias", w2_scale_bias)
# Packs 2 int4 into 1 int8 on-the-fly to mirror the modelslim new_quant_version path
layer.w13_weight.data = self.pack_int4_to_int8(layer.w13_weight.data)
layer.w2_weight.data = self.pack_int4_to_int8(layer.w2_weight.data)
layer.w13_weight.data = maybe_trans_nz(layer.w13_weight.data)
layer.w2_weight.data = maybe_trans_nz(layer.w2_weight.data)
layer.w13_weight.data = self.pack_to_int32(layer.w13_weight.data)
layer.w2_weight.data = self.pack_to_int32(layer.w2_weight.data)
def process_weights_after_loading_modelslim(self, layer):
layer.w13_weight.data = layer.w13_weight.data.transpose(1, 2).contiguous()
layer.w2_weight.data = layer.w2_weight.data.transpose(1, 2).contiguous()
w13_weight_scale_second = (
layer.w13_weight_scale_second.data if hasattr(layer, "w13_weight_scale_second") else None
)
w2_weight_scale_second = layer.w2_weight_scale_second.data if hasattr(layer, "w2_weight_scale_second") else None
layer.w13_weight_scale.data, w13_bias = self.process_scale(
layer.w13_weight, layer.w13_weight_scale.data, w13_weight_scale_second
)
layer.w2_weight_scale.data, w2_bias = self.process_scale(
layer.w2_weight, layer.w2_weight_scale.data, w2_weight_scale_second
)
if hasattr(layer, "w13_weight_scale_second"):
# scale_second is no longer used, release this part of the memory
del layer.w13_weight_scale_second
del layer.w2_weight_scale_second
del layer.w13_weight_offset_second
del layer.w2_weight_offset_second
self.update_bias(layer, w13_bias, w2_bias)
if self.is_per_channel_weight:
layer.w13_weight_scale.data = self.maybe_squeeze_per_channel_weight_scale(layer.w13_weight_scale.data)
layer.w13_weight.data = maybe_trans_nz(layer.w13_weight.data)
layer.w2_weight.data = maybe_trans_nz(layer.w2_weight.data)
if self.dynamic_eplb:
layer.w13_weight_list = [weight.clone() for weight in layer.w13_weight.data.unbind(dim=0)]
layer.w2_weight_list = [weight.clone() for weight in layer.w2_weight.data.unbind(dim=0)]
layer.w13_weight_scale_list = [weight.clone() for weight in layer.w13_weight_scale.data.unbind(dim=0)]
layer.w2_weight_scale_list = [weight.clone() for weight in layer.w2_weight_scale.data.unbind(dim=0)]
layer.w13_scale_bias_list = (
[weight.clone() for weight in layer.w13_scale_bias.data.unbind(dim=0)]
if hasattr(layer, "w13_scale_bias")
else None
)
layer.w2_scale_bias_list = (
[weight.clone() for weight in layer.w2_scale_bias.data.unbind(dim=0)]
if hasattr(layer, "w2_scale_bias")
else None
)
del layer.w13_weight
del layer.w2_weight
del layer.w13_weight_scale
del layer.w2_weight_scale
del layer.w13_scale_bias
del layer.w2_scale_bias
else:
layer.w13_weight.data = self.pack_to_int32(layer.w13_weight.data)
layer.w2_weight.data = self.pack_to_int32(layer.w2_weight.data)

View File

@@ -0,0 +1,245 @@
#
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
# This file is a part of the vllm-ascend project.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
#
from collections.abc import Callable
from typing import Any
import torch
import torch_npu
from vllm.config import CompilationMode, get_current_vllm_config
from vllm.distributed import get_ep_group
from vllm.forward_context import get_forward_context
from vllm_ascend.ascend_config import get_ascend_config
from vllm_ascend.device.mxfp_compat import (
FLOAT8_E8M0FNU_DTYPE,
ensure_mxfp4_linear_available,
)
from vllm_ascend.ops.fused_moe.experts_selector import select_experts
from vllm_ascend.ops.fused_moe.moe_runtime_args import build_fused_experts_input
from .base import AscendLinearScheme, AscendMoEScheme, QuantType, get_moe_num_logical_experts
from .registry import register_scheme
@register_scheme("W4A8_MXFP", "linear")
class AscendW4A8MXFPDynamicLinearMethod(AscendLinearScheme):
"""Linear method for Ascend W4A8_MXFP (Microscaling) quantization."""
def __init__(self):
ensure_mxfp4_linear_available("W8A8_MXFP8 linear quantization")
vllm_config = get_current_vllm_config()
self.group_size = vllm_config.quant_config.quant_description.get("group_size", 32)
@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 // 2, dtype=torch.uint8)}
return params_dict
def get_pergroup_param(
self, input_size: int, output_size: int, params_dtype: torch.dtype, layer_type: str | None = None
) -> dict[str, Any]:
params_dict = {}
params_dict["weight_scale"] = torch.empty(output_size, input_size // self.group_size, dtype=torch.uint8)
return params_dict
def apply(
self,
layer: torch.nn.Module,
x: torch.Tensor | tuple[torch.Tensor, torch.Tensor],
bias: torch.Tensor | None = None,
tp_rank: int | None = 0,
) -> torch.Tensor:
if isinstance(x, tuple):
quantized_x, dynamic_scale = x
output_dtype = torch.bfloat16
else:
quantized_x, dynamic_scale = torch_npu.npu_dynamic_mx_quant(x, dst_type=torch.float8_e4m3fn)
output_dtype = x.dtype
output = torch_npu.npu_quant_matmul(
quantized_x,
layer.weight,
layer.weight_scale,
scale_dtype=torch_npu.float8_e8m0fnu,
pertoken_scale=dynamic_scale,
pertoken_scale_dtype=torch_npu.float8_e8m0fnu,
bias=bias,
output_dtype=output_dtype,
x2_dtype=torch_npu.float4_e2m1fn_x2,
group_sizes=[0, 0, self.group_size],
)
return output
def process_weights_after_loading(self, layer):
layer.weight.data = torch_npu.npu_format_cast(
layer.weight.data, 29, customize_dtype=torch.float8_e4m3fn, input_dtype=torch_npu.float4_e2m1fn_x2
)
layer.weight.data = layer.weight.data.transpose(-1, -2)
n, k = layer.weight_scale.shape
layer.weight_scale.data = layer.weight_scale.data.reshape(n, k // 2, 2).transpose(-3, -2)
@register_scheme("W4A8_MXFP", "moe")
class AscendW4A8MXFPDynamicFusedMoEMethod(AscendMoEScheme):
"""FusedMoe method for Ascend W4A8_DYNAMIC."""
quant_type: QuantType = QuantType.W4A8MXFP
def __init__(self):
self.ep_group = get_ep_group()
vllm_config = get_current_vllm_config()
self.group_size = vllm_config.quant_config.quant_description.get("group_size", 32)
ascend_config = get_ascend_config()
self.use_aclgraph = (
vllm_config.compilation_config.mode == CompilationMode.VLLM_COMPILE
and not vllm_config.model_config.enforce_eager
)
self.dynamic_eplb = ascend_config.eplb_config.dynamic_eplb
@staticmethod
def get_weight(
num_experts: int, intermediate_size_per_partition: int, hidden_sizes: int, params_dtype: torch.dtype
) -> dict[str, Any]:
param_dict = {}
param_dict["w13_weight"] = torch.empty(
num_experts, 2 * intermediate_size_per_partition, hidden_sizes // 2, dtype=torch.uint8
)
param_dict["w2_weight"] = torch.empty(
num_experts, hidden_sizes, intermediate_size_per_partition // 2, dtype=torch.uint8
)
return param_dict
def get_dynamic_quant_param(
self, num_experts: int, intermediate_size_per_partition: int, hidden_sizes: int, params_dtype: torch.dtype
) -> dict[str, Any]:
param_dict = {}
param_dict["w13_weight_scale"] = torch.empty(
num_experts, 2 * intermediate_size_per_partition, hidden_sizes // self.group_size, dtype=torch.uint8
)
param_dict["w2_weight_scale"] = torch.empty(
num_experts, hidden_sizes, intermediate_size_per_partition // self.group_size, dtype=torch.uint8
)
return param_dict
def apply(
self,
layer: torch.nn.Module,
x: torch.Tensor,
router_logits: torch.Tensor,
top_k: int,
renormalize: bool,
use_grouped_topk: bool = False,
num_experts: int = -1,
expert_map: torch.Tensor | None = None,
topk_group: int | None = None,
num_expert_group: int | None = None,
custom_routing_function: Callable | None = None,
scoring_func: str = "softmax",
routed_scaling_factor: float = 1.0,
e_score_correction_bias: torch.Tensor | None = None,
is_prefill: bool = True,
enable_force_load_balance: bool = True,
log2phy: torch.Tensor = None,
global_redundant_expert_num: int = 0,
pertoken_scale: Any | None = None,
activation: str = "silu",
apply_router_weight_on_input: bool = False,
mc2_mask: torch.Tensor | None = None,
tid2eid: torch.Tensor | None = None,
) -> torch.Tensor:
num_shared_experts = getattr(layer, "n_shared_experts", 0)
if num_shared_experts is None:
num_shared_experts = 0
num_logical_experts = get_moe_num_logical_experts(
layer,
num_experts,
global_redundant_expert_num=global_redundant_expert_num,
num_shared_experts=num_shared_experts,
)
assert router_logits.shape[1] == num_logical_experts, "Number of global experts mismatch (excluding redundancy)"
topk_weights, topk_ids = select_experts(
hidden_states=x,
router_logits=router_logits,
top_k=top_k,
use_grouped_topk=use_grouped_topk,
renormalize=renormalize,
topk_group=topk_group,
num_expert_group=num_expert_group,
custom_routing_function=custom_routing_function,
scoring_func=scoring_func,
e_score_correction_bias=e_score_correction_bias,
routed_scaling_factor=routed_scaling_factor,
num_experts=num_logical_experts,
tid2eid=tid2eid,
)
# this is a naive implementation for experts load balance so as
# to avoid accumulating too much tokens on a single rank.
# currently it is only activated when doing profile runs.
if enable_force_load_balance:
random_matrix = torch.rand(topk_ids.size(0), num_logical_experts, device=topk_ids.device)
topk_ids = torch.argsort(random_matrix, dim=1)[:, : topk_ids.size(1)].to(topk_ids.dtype)
if x.dtype not in [torch.float8_e4m3fn]:
topk_weights = topk_weights.to(x.dtype)
moe_comm_method = get_forward_context().moe_comm_method
return moe_comm_method.fused_experts(
fused_experts_input=build_fused_experts_input(
hidden_states=x,
topk_weights=topk_weights,
topk_ids=topk_ids,
w1=layer.w13_weight,
w2=layer.w2_weight,
quant_type=self.quant_type,
dynamic_eplb=self.dynamic_eplb,
expert_map=expert_map,
global_redundant_expert_num=global_redundant_expert_num,
mc2_mask=mc2_mask,
apply_router_weight_on_input=apply_router_weight_on_input,
log2phy=log2phy,
pertoken_scale=pertoken_scale,
activation=activation,
mxfp_act_quant_type=torch.float8_e4m3fn,
mxfp_weight_quant_type=torch_npu.float4_e2m1fn_x2,
mxfp_scale_dtype=FLOAT8_E8M0FNU_DTYPE,
mxfp_per_token_scale_dtype=FLOAT8_E8M0FNU_DTYPE,
mxfp_use_bf16=(x.dtype in [torch.bfloat16, torch.float8_e4m3fn]),
w1_scale=layer.w13_weight_scale,
w2_scale=layer.w2_weight_scale,
swiglu_limit=layer.swiglu_limit,
)
)
def process_weights_after_loading(self, layer):
layer.w13_weight.data = torch_npu.npu_format_cast(
layer.w13_weight.data, 29, customize_dtype=torch.float8_e4m3fn, input_dtype=torch_npu.float4_e2m1fn_x2
)
layer.w2_weight.data = torch_npu.npu_format_cast(
layer.w2_weight.data, 29, customize_dtype=torch.float8_e4m3fn, input_dtype=torch_npu.float4_e2m1fn_x2
)
layer.w13_weight.data = layer.w13_weight.data.transpose(1, 2)
layer.w2_weight.data = layer.w2_weight.data.transpose(1, 2)
g, n, k = layer.w13_weight_scale.shape
layer.w13_weight_scale.data = layer.w13_weight_scale.data.reshape(g, n, k // 2, 2).transpose(-3, -2)
g, n, k = layer.w2_weight_scale.shape
layer.w2_weight_scale.data = layer.w2_weight_scale.data.reshape(g, n, k // 2, 2).transpose(-3, -2)

View File

@@ -0,0 +1,78 @@
#
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
# This file is a part of the vllm-ascend project.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
#
from typing import Any
import torch
import torch_npu
from vllm_ascend.utils import maybe_trans_nz
from .base import AscendLinearScheme
from .registry import register_scheme
@register_scheme("W8A16", "linear")
class AscendW8A16LinearMethod(AscendLinearScheme):
"""Linear method for Ascend W8A16.
This scheme uses 8-bit quantized weights with 16-bit activations.
"""
def __init__(self) -> None:
pass
def get_weight(
self,
input_size: int,
output_size: int,
params_dtype: torch.dtype = torch.bfloat16,
) -> dict[str, Any]:
params_dict = {"weight": torch.empty(output_size, input_size, dtype=torch.int8)}
return params_dict
def get_perchannel_param(
self,
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 apply(
self,
layer: torch.nn.Module,
x: torch.Tensor,
bias: torch.Tensor | None = None,
tp_rank: int | None = 0,
) -> torch.Tensor:
output = torch_npu.npu_weight_quant_batchmatmul(
x=x,
weight=layer.weight,
antiquant_scale=layer.weight_scale,
antiquant_offset=layer.weight_offset,
bias=bias,
)
return output
def process_weights_after_loading(self, layer):
layer.weight.data = layer.weight.data.transpose(0, 1).contiguous()
layer.weight.data = maybe_trans_nz(layer.weight.data)
layer.weight_scale.data = torch.flatten(layer.weight_scale.data)
layer.weight_offset.data = torch.flatten(layer.weight_offset.data)

View File

@@ -0,0 +1,395 @@
#
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
# This file is a part of the vllm-ascend project.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
#
from collections.abc import Callable
from typing import Any
import torch
import torch_npu
from vllm.config import CompilationMode, get_current_vllm_config
from vllm.logger import logger
from vllm_ascend.ascend_config import get_ascend_config
from vllm_ascend.ascend_forward_context import _EXTRA_CTX, MoECommType
from vllm_ascend.distributed.parallel_state import get_mc2_group
from vllm_ascend.flash_common3_context import get_flash_common3_context
from vllm_ascend.ops.fused_moe.experts_selector import select_experts, zero_experts_compute
from vllm_ascend.ops.fused_moe.moe_runtime_args import build_fused_experts_input
from vllm_ascend.utils import ACL_FORMAT_FRACTAL_NZ, enable_dsa_cp, maybe_trans_nz
from .base import AscendLinearScheme, AscendMoEScheme, QuantType, get_moe_num_logical_experts
from .registry import register_scheme
def scale_from_float_to_int64(scale):
"""Convert float32 scale to int64 representation."""
import numpy as np
scale = torch.from_numpy(
np.frombuffer(scale.cpu().to(torch.float32).numpy().tobytes(), dtype=np.int32).astype(np.int64)
).to(scale.device)
return scale
@register_scheme("W8A8_DYNAMIC", "linear")
class AscendW8A8DynamicLinearMethod(AscendLinearScheme):
"""Linear method for Ascend W8A8_DYNAMIC.
This scheme uses dynamic per-token quantization for activations
and per-channel quantization for weights.
"""
act_quant_type: torch.dtype = torch.int8
def __init__(self):
pass
def get_weight(self, 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
def get_perchannel_param(
self,
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 apply(
self,
layer: torch.nn.Module,
x: torch.Tensor,
bias: torch.Tensor | None = None,
tp_rank: int | None = 0,
) -> torch.Tensor:
quantized_x, pertoken_scale = torch_npu.npu_dynamic_quant(x, dst_type=self.act_quant_type)
need_unsqz = False
if pertoken_scale.dim() == 2:
need_unsqz = True
quantized_x = quantized_x.squeeze(dim=1)
pertoken_scale = pertoken_scale.squeeze(dim=1)
chunk_size = getattr(layer, "_chunk_size", 0)
if isinstance(chunk_size, int) and chunk_size > 0:
bias_1 = bias[:chunk_size] if bias is not None else None
bias_2 = bias[chunk_size:] if bias is not None else None
output = torch.cat(
[
torch_npu.npu_quant_matmul(
quantized_x,
layer.weight_1,
layer.weight_1_scale,
pertoken_scale=pertoken_scale,
bias=bias_1,
output_dtype=x.dtype,
),
torch_npu.npu_quant_matmul(
quantized_x,
layer.weight_2,
layer.weight_2_scale,
pertoken_scale=pertoken_scale,
bias=bias_2,
output_dtype=x.dtype,
),
],
dim=-1,
)
else:
output = torch_npu.npu_quant_matmul(
quantized_x,
layer.weight,
layer.weight_scale,
pertoken_scale=pertoken_scale,
bias=bias if self.act_quant_type == torch.int8 else None,
output_dtype=x.dtype,
)
if need_unsqz:
output = output.unsqueeze(dim=1)
return output
def process_weights_after_loading(self, layer):
layer.weight.data = layer.weight.data.transpose(0, 1).contiguous()
if "wq_b" in getattr(layer, "prefix", "") and layer.weight.shape[1] >= 65536 and enable_dsa_cp():
# TODO(jianzs): Remove this workaround after
# `torch_npu.npu_quant_matmul` supports large weight dimensions.
chunk_size = layer.weight.shape[1] // 2
assert chunk_size < 65536, "Even after chunking, the weight dimension is still larger than 65536."
layer._chunk_size = chunk_size
layer.weight_1 = maybe_trans_nz(layer.weight.data[:, :chunk_size].contiguous())
layer.weight_2 = maybe_trans_nz(layer.weight.data[:, chunk_size:].contiguous())
layer.weight_1_scale = layer.weight_scale.data[:chunk_size].flatten().contiguous()
layer.weight_2_scale = layer.weight_scale.data[chunk_size:].flatten().contiguous()
layer.weight_1_scale_fp32 = layer.weight_1_scale.to(torch.float32)
layer.weight_2_scale_fp32 = layer.weight_2_scale.to(torch.float32)
layer.weight_1_offset = layer.weight_offset.data[:chunk_size].flatten().contiguous()
layer.weight_2_offset = layer.weight_offset.data[chunk_size:].flatten().contiguous()
del layer.weight
del layer.weight_scale
del layer.weight_offset
else:
# cast quantized weight tensors in NZ format for higher inference speed
if self.act_quant_type == torch.int8:
layer.weight.data = maybe_trans_nz(layer.weight.data)
layer.weight_scale.data = layer.weight_scale.data.flatten()
layer.weight_scale_fp32 = layer.weight_scale.data.to(torch.float32)
layer.weight_offset.data = layer.weight_offset.data.flatten()
@register_scheme("W8A8_DYNAMIC", "moe")
class AscendW8A8DynamicFusedMoEMethod(AscendMoEScheme):
"""FusedMoE method for Ascend W8A8_DYNAMIC."""
# Declare the quantization type for this scheme
quant_type: QuantType = QuantType.W8A8
def __init__(self):
vllm_config = get_current_vllm_config()
ascend_config = get_ascend_config()
self.use_aclgraph = (
vllm_config.compilation_config.mode == CompilationMode.VLLM_COMPILE
and not vllm_config.model_config.enforce_eager
)
self.multistream_overlap_gate = ascend_config.multistream_overlap_gate
self.dynamic_eplb = ascend_config.eplb_config.dynamic_eplb
self.in_dtype = vllm_config.model_config.dtype
self.supports_eplb = True
try:
device_group = get_mc2_group().device_group
# TODO: Try local_rank = ep_group.rank_in_group
local_rank = torch.distributed.get_rank(group=device_group)
backend = device_group._get_backend(torch.device("npu"))
self.moe_all_to_all_group_name = backend.get_hccl_comm_name(local_rank)
except AttributeError:
logger.warning_once(
"[vllm-ascend/W8A8_DYNAMIC] MC2 group metadata unavailable, "
"falling back to empty moe_all_to_all_group_name."
)
self.moe_all_to_all_group_name = ""
def get_weight(
self, num_experts: int, intermediate_size_per_partition: int, hidden_sizes: int, params_dtype: torch.dtype
) -> dict[str, Any]:
param_dict = {}
param_dict["w13_weight"] = torch.empty(
num_experts, 2 * intermediate_size_per_partition, hidden_sizes, dtype=torch.int8
)
param_dict["w2_weight"] = torch.empty(
num_experts, hidden_sizes, intermediate_size_per_partition, dtype=torch.int8
)
return param_dict
def get_dynamic_quant_param(
self, num_experts: int, intermediate_size_per_partition: int, hidden_sizes: int, params_dtype: torch.dtype
) -> dict[str, Any]:
param_dict = {}
param_dict["w13_weight_scale"] = torch.empty(
num_experts, 2 * intermediate_size_per_partition, 1, dtype=params_dtype
)
param_dict["w13_weight_offset"] = torch.empty(
num_experts, 2 * intermediate_size_per_partition, 1, dtype=params_dtype
)
param_dict["w2_weight_scale"] = torch.empty(num_experts, hidden_sizes, 1, dtype=params_dtype)
param_dict["w2_weight_offset"] = torch.empty(num_experts, hidden_sizes, 1, dtype=params_dtype)
return param_dict
def apply(
self,
layer: torch.nn.Module,
x: torch.Tensor,
router_logits: torch.Tensor,
top_k: int,
renormalize: bool,
use_grouped_topk: bool = False,
num_experts: int = -1,
expert_map: torch.Tensor | None = None,
topk_group: int | None = None,
num_expert_group: int | None = None,
custom_routing_function: Callable | None = None,
scoring_func: str = "softmax",
routed_scaling_factor: float = 1.0,
e_score_correction_bias: torch.Tensor | None = None,
is_prefill: bool = True,
enable_force_load_balance: bool = False,
log2phy: torch.Tensor | None = None,
global_redundant_expert_num: int = 0,
pertoken_scale: Any | None = None,
activation: str = "silu",
apply_router_weight_on_input: bool = False,
mc2_mask: torch.Tensor | None = None,
tid2eid: torch.Tensor | None = None,
) -> torch.Tensor:
zero_expert_num = getattr(layer, "zero_expert_num", 0)
zero_expert_type = getattr(layer, "zero_expert_type", None)
n_shared_experts = getattr(layer, "n_shared_experts", 0)
mix_placement = getattr(layer, "mix_placement", False)
if n_shared_experts is None:
n_shared_experts = 0
num_logical_experts = get_moe_num_logical_experts(
layer,
num_experts,
global_redundant_expert_num=global_redundant_expert_num,
num_shared_experts=n_shared_experts,
)
if zero_expert_num == 0 or zero_expert_type is None:
assert router_logits.shape[1] == num_logical_experts, (
"[vllm-ascend/W8A8_DYNAMIC] Number of global experts mismatch "
"(excluding redundancy). "
f"router_experts={router_logits.shape[1]}, "
f"expected_experts={num_logical_experts}, "
f"zero_expert_num={zero_expert_num}, "
f"zero_expert_type={zero_expert_type}"
)
if self.multistream_overlap_gate:
fc3_context = get_flash_common3_context()
assert fc3_context is not None, (
"[vllm-ascend/W8A8_DYNAMIC] flash_common3 context is required when multistream_overlap_gate is enabled."
)
topk_weights = fc3_context.topk_weights
topk_ids = fc3_context.topk_ids
else:
topk_weights, topk_ids = select_experts(
hidden_states=x,
router_logits=router_logits,
top_k=top_k,
use_grouped_topk=use_grouped_topk,
renormalize=renormalize,
topk_group=topk_group,
num_expert_group=num_expert_group,
custom_routing_function=custom_routing_function,
scoring_func=scoring_func,
routed_scaling_factor=routed_scaling_factor,
e_score_correction_bias=e_score_correction_bias,
mix_placement=mix_placement,
num_logical_experts=router_logits.shape[1],
num_shared_experts=n_shared_experts,
num_experts=num_logical_experts,
tid2eid=tid2eid,
)
assert topk_ids is not None
assert topk_weights is not None
if zero_expert_num > 0 and zero_expert_type is not None:
topk_ids, topk_weights, zero_expert_result = zero_experts_compute(
expert_indices=topk_ids,
expert_scales=topk_weights,
num_experts=num_logical_experts,
zero_expert_type=zero_expert_type,
hidden_states=x,
)
# this is a naive implementation for experts load balance so as
# to avoid accumulating too much tokens on a single rank.
# currently it is only activated when doing profile runs.
if enable_force_load_balance:
random_matrix = torch.rand(topk_ids.size(0), num_logical_experts, device=topk_ids.device)
topk_ids = torch.argsort(random_matrix, dim=1)[:, : topk_ids.size(1)].to(topk_ids.dtype)
assert topk_weights is not None
topk_weights = topk_weights.to(self.in_dtype)
moe_comm_method = _EXTRA_CTX.moe_comm_method
fused_scale_flag = (
_EXTRA_CTX.moe_comm_type == MoECommType.FUSED_MC2 and get_ascend_config().enable_fused_mc2 == 1
)
if self.dynamic_eplb:
w1 = layer.w13_weight_list
w1_scale = layer.fused_w1_scale_list if fused_scale_flag else layer.w13_weight_scale_fp32_list
w2 = layer.w2_weight_list
w2_scale = layer.fused_w2_scale_list if fused_scale_flag else layer.w2_weight_scale_list
else:
w1 = [layer.w13_weight]
w1_scale = [layer.fused_w1_scale] if fused_scale_flag else [layer.w13_weight_scale_fp32]
w2 = [layer.w2_weight]
w2_scale = [layer.fused_w2_scale] if fused_scale_flag else [layer.w2_weight_scale]
w1_scale_bias = [torch.tensor([], dtype=torch.float32)] if fused_scale_flag else None
w2_scale_bias = [torch.tensor([], dtype=torch.float32)] if fused_scale_flag else None
final_hidden_states = moe_comm_method.fused_experts(
fused_experts_input=build_fused_experts_input(
hidden_states=x,
topk_weights=topk_weights,
topk_ids=topk_ids,
w1=w1,
w2=w2,
quant_type=self.quant_type,
dynamic_eplb=self.dynamic_eplb,
expert_map=expert_map,
global_redundant_expert_num=global_redundant_expert_num,
mc2_mask=mc2_mask,
apply_router_weight_on_input=apply_router_weight_on_input,
log2phy=log2phy,
pertoken_scale=pertoken_scale,
activation=activation,
w1_scale=w1_scale,
w2_scale=w2_scale,
w1_scale_bias=w1_scale_bias,
w2_scale_bias=w2_scale_bias,
swiglu_limit=layer.swiglu_limit,
)
)
if zero_expert_num > 0 and zero_expert_type is not None:
final_hidden_states += zero_expert_result
return final_hidden_states
def process_weights_after_loading(self, layer):
layer.w13_weight.data = layer.w13_weight.data.transpose(1, 2).contiguous()
layer.w2_weight.data = layer.w2_weight.data.transpose(1, 2).contiguous()
# TODO(zzzzwwjj): Currently, `torch_npu.npu_grouped_matmul_swiglu_quant`
# can only support weight nz.
if self.quant_type == QuantType.W8A8:
layer.w13_weight.data = torch_npu.npu_format_cast(layer.w13_weight.data, ACL_FORMAT_FRACTAL_NZ)
layer.w2_weight.data = torch_npu.npu_format_cast(layer.w2_weight.data, ACL_FORMAT_FRACTAL_NZ)
layer.w13_weight_scale.data = layer.w13_weight_scale.data.view(layer.w13_weight_scale.data.shape[0], -1)
layer.w13_weight_scale_fp32 = layer.w13_weight_scale.data.to(torch.float32)
layer.w13_weight_offset.data = layer.w13_weight_offset.data.view(layer.w13_weight_offset.data.shape[0], -1)
layer.w2_weight_scale.data = layer.w2_weight_scale.data.view(layer.w2_weight_scale.data.shape[0], -1)
layer.w2_weight_offset.data = layer.w2_weight_offset.data.view(layer.w2_weight_offset.data.shape[0], -1)
if get_ascend_config().enable_fused_mc2 == 1:
layer.fused_w1_scale = scale_from_float_to_int64(layer.w13_weight_scale.data)
layer.fused_w2_scale = scale_from_float_to_int64(layer.w2_weight_scale.data)
if self.dynamic_eplb:
layer.w13_weight_list = [weight.clone() for weight in layer.w13_weight.data.unbind(dim=0)]
layer.w2_weight_list = [weight.clone() for weight in layer.w2_weight.data.unbind(dim=0)]
layer.w13_weight_scale_fp32_list = [
weight.clone() for weight in layer.w13_weight_scale_fp32.data.unbind(dim=0)
]
layer.w2_weight_scale_list = [weight.clone() for weight in layer.w2_weight_scale.data.unbind(dim=0)]
if get_ascend_config().enable_fused_mc2 == 1:
layer.fused_w1_scale_list = [
weight.clone()
for weight in layer.fused_w1_scale.view(len(layer.w13_weight_list), -1).data.unbind(dim=0)
]
layer.fused_w2_scale_list = [
weight.clone()
for weight in layer.fused_w2_scale.view(len(layer.w2_weight_list), -1).data.unbind(dim=0)
]
del layer.w13_weight
del layer.w2_weight
del layer.w13_weight_scale
del layer.w13_weight_scale_fp32
del layer.w2_weight_scale
if get_ascend_config().enable_fused_mc2 == 1:
del layer.fused_w1_scale
del layer.fused_w2_scale
torch.npu.empty_cache()

View File

@@ -0,0 +1,432 @@
#
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
# This file is a part of the vllm-ascend project.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
#
from collections.abc import Callable
from typing import Any
import torch
import torch.nn.functional as F
import torch_npu
from vllm.config import CompilationMode, get_current_vllm_config
from vllm.logger import logger
from vllm.utils.math_utils import cdiv
from vllm_ascend.ascend_config import get_ascend_config
from vllm_ascend.ascend_forward_context import _EXTRA_CTX
from vllm_ascend.device.mxfp_compat import (
FLOAT8_E8M0FNU_DTYPE,
ensure_mxfp8_linear_available,
ensure_mxfp8_moe_available,
)
from vllm_ascend.flash_common3_context import get_flash_common3_context
from vllm_ascend.ops.fused_moe.experts_selector import select_experts
from vllm_ascend.ops.fused_moe.moe_runtime_args import build_fused_experts_input
from .base import AscendLinearScheme, AscendMoEScheme, QuantType, get_moe_num_logical_experts
from .registry import register_scheme
@register_scheme("W8A8_MXFP8", "linear")
class AscendW8A8MXFP8DynamicLinearMethod(AscendLinearScheme):
"""Linear method for Ascend W8A8_MXFP8 (Microscaling FP8) quantization.
This scheme uses microscaling FP8 quantization with per-group scales.
The activation is dynamically quantized to FP8 (E4M3FN format) with
microscaling, and weights are stored in FP8 format with per-group scales.
"""
model_dtype = None
def __init__(self):
ensure_mxfp8_linear_available("W8A8_MXFP8 linear quantization")
vllm_config = get_current_vllm_config()
self.group_size = vllm_config.quant_config.quant_description.get("group_size", 32)
def get_weight(self, input_size: int, output_size: int, params_dtype: torch.dtype) -> dict[str, Any]:
params_dict = {"weight": torch.empty(output_size, input_size, dtype=torch.float8_e4m3fn)}
return params_dict
def get_pergroup_param(
self, input_size: int, output_size: int, params_dtype: torch.dtype, layer_type: str | None = None
) -> dict[str, Any]:
params_dict = {}
params_dict["weight_scale"] = torch.empty(output_size, cdiv(input_size, self.group_size), dtype=torch.uint8)
return params_dict
def apply(
self,
layer: torch.nn.Module,
x: torch.Tensor | tuple[torch.Tensor, torch.Tensor],
bias: torch.Tensor | None = None,
tp_rank: int | None = 0,
) -> torch.Tensor:
if isinstance(x, tuple):
quantized_x, pertoken_scale = x
original_shape = quantized_x.shape
output_dtype = torch.bfloat16
else:
# reshape x for Qwen VL models
original_shape = x.shape
if x.dim() > 2:
x = x.view(-1, x.shape[-1])
quantized_x, pertoken_scale = torch_npu.npu_dynamic_mx_quant(x, dst_type=torch.float8_e4m3fn)
output_dtype = x.dtype
if bias is not None and bias.dtype != torch.float32:
bias = bias.to(torch.float32)
output = torch_npu.npu_quant_matmul(
quantized_x,
layer.weight,
layer.weight_scale,
scale_dtype=FLOAT8_E8M0FNU_DTYPE,
pertoken_scale=pertoken_scale,
pertoken_scale_dtype=FLOAT8_E8M0FNU_DTYPE,
bias=bias,
output_dtype=output_dtype,
group_sizes=[1, 1, self.group_size],
)
# reshape output for Qwen VL models
if len(original_shape) > 2:
output = output.view(*original_shape[:-1], -1)
return output
def process_weights_after_loading(self, layer):
"""Process weights after loading for MXFP8 inference.
This method transforms weights for NPU MXFP8 computation:
- weight: (output_size, input_size) -> (input_size, output_size)
- weight_scale: (n_dim, k_dim) -> (k_dim//2, n_dim, 2)
For RL training scenarios where weights need to be reloaded multiple times,
this method stores original shapes and can be called multiple times safely.
Use restore_weights_for_rl_loading() before weight reload, then call this
method again after loading.
"""
# Check if already transformed to avoid double transformation
if getattr(layer, "_mxfp8_transformed", False):
return
# Store original shapes for RL weight reloading
# Only store on first call (when shapes are in original format)
if not hasattr(layer, "_mxfp8_original_shapes"):
layer._mxfp8_original_shapes = {
"weight": tuple(layer.weight.data.shape),
"weight_scale": tuple(layer.weight_scale.data.shape),
}
n_dim, k_dim = layer.weight_scale.data.shape
# Shape should be padded if it cannot be divided by 2
if layer.weight_scale.data.shape[-1] % 2 != 0:
layer.weight_scale.data = F.pad(layer.weight_scale.data, (0, 1), mode="constant", value=0)
layer.weight_scale.data = layer.weight_scale.data.reshape(n_dim, k_dim // 2 + 1, 2)
else:
layer.weight_scale.data = layer.weight_scale.data.reshape(n_dim, k_dim // 2, 2)
layer.weight.data = layer.weight.data.transpose(0, 1).contiguous()
layer.weight_scale.data = layer.weight_scale.data.transpose(0, 1).contiguous()
# Mark as transformed
layer._mxfp8_transformed = True
def restore_weights_for_rl_loading(self, layer):
"""Restore weights to original shapes for RL weight reloading.
This method must be called BEFORE model.load_weights() in RL training
loops to restore the tensors to their original shapes that the weight
loader expects.
After weight loading, call process_weights_after_loading() again to
re-apply the MXFP8 transformations.
Shape transformations reversed:
- weight: (input_size, output_size) -> (output_size, input_size)
- weight_scale: (k_dim//2, n_dim, 2) -> (n_dim, k_dim)
"""
if not getattr(layer, "_mxfp8_transformed", False):
# Not transformed, nothing to restore
return
if not hasattr(layer, "_mxfp8_original_shapes"):
err_msg = (
"[vllm-ascend/W8A8_MXFP8] Cannot restore weights: original "
"shapes not recorded. "
"This should not happen if process_weights_after_loading was called first."
)
logger.error(err_msg)
raise RuntimeError(err_msg)
orig_shapes = layer._mxfp8_original_shapes
orig_scale_shape = orig_shapes["weight_scale"]
# Restore weight: (input_size, output_size) -> (output_size, input_size)
target_weight = layer.weight.data.transpose(0, 1).contiguous()
layer.weight.data = layer.weight.data.transpose(0, 1)
layer.weight.data.copy_(target_weight)
# Restore weight_scale: (k_dim//2, n_dim, 2) -> (n_dim, k_dim)
# Current shape: (k_dim//2, n_dim, 2)
# Target shape: (n_dim, k_dim)
target_scale = layer.weight_scale.data.transpose(0, 1).reshape(orig_scale_shape).contiguous()
layer.weight_scale.data = layer.weight_scale.data.transpose(0, 1).reshape(orig_scale_shape)
layer.weight_scale.data.copy_(target_scale)
# Mark as not transformed (ready for weight loading)
layer._mxfp8_transformed = False
@register_scheme("W8A8_MXFP8", "moe")
class AscendW8A8MXFP8DynamicFusedMoEMethod(AscendMoEScheme):
"""FusedMoe method for Ascend W8A8_DYNAMIC."""
model_dtype = None
quant_type: QuantType = QuantType.MXFP8
def __init__(self):
ensure_mxfp8_moe_available("W8A8_MXFP8 MoE quantization")
vllm_config = get_current_vllm_config()
self.group_size = vllm_config.quant_config.quant_description.get("group_size", 32)
ascend_config = get_ascend_config()
self.use_aclgraph = (
vllm_config.compilation_config.mode == CompilationMode.VLLM_COMPILE
and not vllm_config.model_config.enforce_eager
)
self.dynamic_eplb = ascend_config.eplb_config.dynamic_eplb
self.multistream_overlap_gate = ascend_config.multistream_overlap_gate
@staticmethod
def get_weight(
num_experts: int, intermediate_size_per_partition: int, hidden_sizes: int, params_dtype: torch.dtype
) -> dict[str, Any]:
param_dict = {}
param_dict["w13_weight"] = torch.empty(
num_experts, 2 * intermediate_size_per_partition, hidden_sizes, dtype=torch.float8_e4m3fn
)
param_dict["w2_weight"] = torch.empty(
num_experts, hidden_sizes, intermediate_size_per_partition, dtype=torch.float8_e4m3fn
)
return param_dict
def get_dynamic_quant_param(
self, num_experts: int, intermediate_size_per_partition: int, hidden_sizes: int, params_dtype: torch.dtype
) -> dict[str, Any]:
param_dict = {}
param_dict["w13_weight_scale"] = torch.empty(
num_experts, 2 * intermediate_size_per_partition, hidden_sizes // self.group_size, dtype=torch.uint8
)
param_dict["w2_weight_scale"] = torch.empty(
num_experts, hidden_sizes, intermediate_size_per_partition // self.group_size, dtype=torch.uint8
)
return param_dict
def apply(
self,
layer: torch.nn.Module,
x: torch.Tensor,
router_logits: torch.Tensor,
top_k: int,
renormalize: bool,
use_grouped_topk: bool = False,
num_experts: int = -1,
expert_map: torch.Tensor | None = None,
topk_group: int | None = None,
num_expert_group: int | None = None,
custom_routing_function: Callable | None = None,
scoring_func: str = "softmax",
routed_scaling_factor: float = 1.0,
e_score_correction_bias: torch.Tensor | None = None,
is_prefill: bool = True,
enable_force_load_balance: bool = True,
log2phy: torch.Tensor = None,
global_redundant_expert_num: int = 0,
pertoken_scale: Any | None = None,
activation: str = "silu",
apply_router_weight_on_input: bool = False,
mc2_mask: torch.Tensor | None = None,
tid2eid: torch.Tensor | None = None,
) -> torch.Tensor:
num_shared_experts = getattr(layer, "n_shared_experts", 0)
if num_shared_experts is None:
num_shared_experts = 0
num_logical_experts = get_moe_num_logical_experts(
layer,
num_experts,
global_redundant_expert_num=global_redundant_expert_num,
num_shared_experts=num_shared_experts,
)
assert router_logits.shape[1] == num_logical_experts, "Number of global experts mismatch (excluding redundancy)"
if self.multistream_overlap_gate:
fc3_context = get_flash_common3_context()
assert fc3_context is not None
topk_weights = fc3_context.topk_weights
topk_ids = fc3_context.topk_ids
else:
topk_weights, topk_ids = select_experts(
hidden_states=x,
router_logits=router_logits,
top_k=top_k,
use_grouped_topk=use_grouped_topk,
renormalize=renormalize,
topk_group=topk_group,
num_expert_group=num_expert_group,
custom_routing_function=custom_routing_function,
scoring_func=scoring_func,
routed_scaling_factor=routed_scaling_factor,
e_score_correction_bias=e_score_correction_bias,
num_experts=num_logical_experts,
tid2eid=tid2eid,
)
if topk_weights is None or topk_ids is None:
raise RuntimeError("topk_weights and topk_ids must be set before fused MoE execution.")
# this is a naive implementation for experts load balance so as
# to avoid accumulating too much tokens on a single rank.
# currently it is only activated when doing profile runs.
if enable_force_load_balance:
random_matrix = torch.rand(topk_ids.size(0), num_logical_experts, device=topk_ids.device)
topk_ids = torch.argsort(random_matrix, dim=1)[:, : topk_ids.size(1)].to(topk_ids.dtype)
if x.dtype not in [torch.float8_e4m3fn]:
topk_weights = topk_weights.to(x.dtype)
moe_comm_method = _EXTRA_CTX.moe_comm_method
return moe_comm_method.fused_experts(
fused_experts_input=build_fused_experts_input(
hidden_states=x,
topk_weights=topk_weights,
topk_ids=topk_ids,
w1=layer.w13_weight,
w2=layer.w2_weight,
quant_type=self.quant_type,
dynamic_eplb=self.dynamic_eplb,
expert_map=expert_map,
global_redundant_expert_num=global_redundant_expert_num,
mc2_mask=mc2_mask,
apply_router_weight_on_input=apply_router_weight_on_input,
log2phy=log2phy,
pertoken_scale=pertoken_scale,
activation=activation,
mxfp_act_quant_type=torch.float8_e4m3fn,
mxfp_weight_quant_type=torch.float8_e4m3fn,
mxfp_scale_dtype=FLOAT8_E8M0FNU_DTYPE,
mxfp_per_token_scale_dtype=FLOAT8_E8M0FNU_DTYPE,
mxfp_use_bf16=(x.dtype in [torch.bfloat16, torch.float8_e4m3fn]),
w1_scale=layer.w13_weight_scale,
w2_scale=layer.w2_weight_scale,
swiglu_limit=layer.swiglu_limit,
)
)
def process_weights_after_loading(self, layer):
"""Process weights after loading for MXFP8 inference.
This method transforms weights for NPU MXFP8 computation:
- w13_weight: (g_num, n_size, k_size) -> (g_num, k_size, n_size)
- w2_weight: (g_num, n_size, k_size) -> (g_num, k_size, n_size)
- w13_weight_scale: (g_num, n_size, k_size) -> (g_num, k_size//2, n_size, 2)
- w2_weight_scale: (g_num, n_size, k_size) -> (g_num, k_size//2, n_size, 2)
For RL training scenarios where weights need to be reloaded multiple times,
this method stores original shapes and can be called multiple times safely.
Use restore_weights_for_rl_loading() before weight reload, then call this
method again after loading.
"""
# Check if already transformed to avoid double transformation
if getattr(layer, "_mxfp8_transformed", False):
return
# Store original shapes for RL weight reloading
# Only store on first call (when shapes are in original format)
if not hasattr(layer, "_mxfp8_original_shapes"):
layer._mxfp8_original_shapes = {
"w13_weight": tuple(layer.w13_weight.data.shape),
"w13_weight_scale": tuple(layer.w13_weight_scale.data.shape),
"w2_weight": tuple(layer.w2_weight.data.shape),
"w2_weight_scale": tuple(layer.w2_weight_scale.data.shape),
}
g_num, n_size, k_size = layer.w13_weight_scale.shape
layer.w13_weight_scale.data = layer.w13_weight_scale.data.reshape(g_num, n_size, k_size // 2, 2)
g_num, n_size, k_size = layer.w2_weight_scale.shape
layer.w2_weight_scale.data = layer.w2_weight_scale.data.reshape(g_num, n_size, k_size // 2, 2)
layer.w13_weight.data = layer.w13_weight.data.transpose(1, 2)
layer.w2_weight.data = layer.w2_weight.data.transpose(1, 2)
layer.w13_weight_scale.data = layer.w13_weight_scale.data.transpose(1, 2)
layer.w2_weight_scale.data = layer.w2_weight_scale.data.transpose(1, 2)
# Mark as transformed
layer._mxfp8_transformed = True
def restore_weights_for_rl_loading(self, layer):
"""Restore weights to original shapes for RL weight reloading.
This method must be called BEFORE model.load_weights() in RL training
loops to restore the tensors to their original shapes that the weight
loader expects.
After weight loading, call process_weights_after_loading() again to
re-apply the MXFP8 transformations.
Shape transformations reversed:
- w13_weight: (g_num, k_size, n_size) -> (g_num, n_size, k_size)
- w2_weight: (g_num, k_size, n_size) -> (g_num, n_size, k_size)
- w13_weight_scale: (g_num, k_size//2, n_size, 2) -> (g_num, n_size, k_size)
- w2_weight_scale: (g_num, k_size//2, n_size, 2) -> (g_num, n_size, k_size)
"""
if not getattr(layer, "_mxfp8_transformed", False):
# Not transformed, nothing to restore
return
if not hasattr(layer, "_mxfp8_original_shapes"):
err_msg = (
"[vllm-ascend/W8A8_MXFP8] Cannot restore weights: original "
"shapes not recorded. "
"This should not happen if process_weights_after_loading was called first."
)
logger.error(err_msg)
raise RuntimeError(err_msg)
orig_shapes = layer._mxfp8_original_shapes
def _restore(weight_key: str, scale_key: str):
"""Helper to restore a single MoE weight and its scale using safe memory copies."""
# --- 1. Restore Weight ---
weight_tensor = getattr(layer, weight_key)
target_weight = weight_tensor.data.transpose(1, 2).contiguous()
weight_tensor.data = weight_tensor.data.transpose(1, 2)
weight_tensor.data.copy_(target_weight)
# --- 2. Restore Weight Scale ---
scale_tensor = getattr(layer, scale_key)
orig_scale_shape = orig_shapes[scale_key]
target_scale = scale_tensor.data.transpose(1, 2).reshape(orig_scale_shape).contiguous()
scale_tensor.data = scale_tensor.data.transpose(1, 2).view(orig_scale_shape)
scale_tensor.data.copy_(target_scale)
_restore("w13_weight", "w13_weight_scale")
_restore("w2_weight", "w2_weight_scale")
# Mark as not transformed (ready for weight loading)
layer._mxfp8_transformed = False

View File

@@ -0,0 +1,101 @@
#
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
# This file is a part of the vllm-ascend project.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
#
"""W8A8 Prefill-Decode Mix quantization methods.
This module provides quantization methods that use different strategies
for prefill and decode phases:
- Prefill: Uses dynamic W8A8 quantization
- Decode (KV consumer): Uses static W8A8 quantization
"""
from typing import Any
import torch
from vllm.config import get_current_vllm_config
from .base import AscendLinearScheme
from .registry import register_scheme
from .w8a8_dynamic import AscendW8A8DynamicFusedMoEMethod, AscendW8A8DynamicLinearMethod
from .w8a8_static import AscendW8A8LinearMethod
@register_scheme("W8A8_MIX", "linear")
class AscendW8A8PDMixLinearMethod(AscendLinearScheme):
"""Linear method for W8A8 prefill-decode mix quantization.
This scheme uses composition to delegate to the appropriate quantization
method based on the execution phase:
- Static W8A8 for KV consumer (decode phase)
- Dynamic W8A8 for prefill phase
The static method is used for weight/parameter specifications since
it requires more parameters (input_scale, deq_scale, etc.) that are
needed for static quantization during decode.
"""
def __init__(self):
self._static_method = AscendW8A8LinearMethod()
self._dynamic_method = AscendW8A8DynamicLinearMethod()
kv_transfer_config = get_current_vllm_config().kv_transfer_config
self._is_kv_consumer = kv_transfer_config is not None and kv_transfer_config.is_kv_consumer
def get_weight(self, input_size: int, output_size: int, params_dtype: torch.dtype) -> dict[str, Any]:
return self._static_method.get_weight(input_size, output_size, params_dtype)
def get_pertensor_param(self, params_dtype: torch.dtype, **kwargs: Any) -> dict[str, Any]:
return self._static_method.get_pertensor_param(params_dtype)
def get_perchannel_param(
self,
output_size: int,
params_dtype: torch.dtype,
) -> dict[str, Any]:
return self._static_method.get_perchannel_param(output_size, params_dtype)
def apply(
self,
layer: torch.nn.Module,
x: torch.Tensor,
bias: torch.Tensor | None = None,
tp_rank: int | None = 0,
) -> torch.Tensor:
if layer.is_kv_consumer:
return self._static_method.apply(layer, x, bias, tp_rank)
else:
return self._dynamic_method.apply(layer, x, bias, tp_rank)
def process_weights_after_loading(self, layer: torch.nn.Module) -> None:
self._static_method.process_weights_after_loading(layer)
layer.weight_scale_fp32 = layer.weight_scale.data.to(torch.float32)
layer.is_kv_consumer = self._is_kv_consumer
@register_scheme("W8A8_MIX", "moe")
class AscendW8A8PDMixFusedMoeMethod(AscendW8A8DynamicFusedMoEMethod):
def get_dynamic_quant_param(
self, num_experts: int, intermediate_size_per_partition: int, hidden_sizes: int, params_dtype: torch.dtype
) -> dict[str, Any]:
param_dict = super().get_dynamic_quant_param(
num_experts, intermediate_size_per_partition, hidden_sizes, params_dtype
)
param_dict["w2_deq_scale"] = torch.empty(num_experts, hidden_sizes, dtype=torch.float32)
param_dict["w13_deq_scale"] = torch.empty(num_experts, 2 * intermediate_size_per_partition, dtype=torch.float32)
param_dict["w2_input_offset"] = torch.empty(num_experts, 1, dtype=torch.int8)
param_dict["w13_input_offset"] = torch.empty(num_experts, 1, dtype=torch.int8)
return param_dict

View File

@@ -0,0 +1,161 @@
#
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
# This file is a part of the vllm-ascend project.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
#
from typing import Any
import torch
import torch_npu
from vllm_ascend.utils import (
COMPRESSED_TENSORS_METHOD,
get_weight_prefetch_method,
maybe_trans_nz,
)
from .base import AscendLinearScheme
from .registry import register_scheme
@register_scheme("W8A8", "linear")
class AscendW8A8LinearMethod(AscendLinearScheme):
"""Linear method for Ascend W8A8 static quantization.
This scheme uses static per-tensor quantization for activations
and per-channel quantization for weights.
"""
def __init__(self) -> None:
pass
def get_weight(
self,
input_size: int,
output_size: int,
params_dtype: torch.dtype = torch.bfloat16,
) -> dict[str, Any]:
params_dict = {"weight": torch.empty(output_size, input_size, dtype=torch.int8)}
return params_dict
def get_pertensor_param(self, params_dtype: torch.dtype, **kwargs: Any) -> dict[str, Any]:
params_dict = {}
params_dict["input_scale"] = torch.empty(1, dtype=params_dtype)
params_dict["input_offset"] = torch.empty(1, dtype=torch.int8)
return params_dict
def get_perchannel_param(
self,
output_size: int,
params_dtype: torch.dtype,
) -> dict[str, Any]:
params_dict = {}
params_dict["quant_bias"] = torch.empty(output_size, dtype=torch.int32)
if params_dtype == torch.bfloat16:
params_dict["deq_scale"] = torch.empty(output_size, dtype=torch.float32)
elif params_dtype == torch.float16:
params_dict["deq_scale"] = torch.empty(output_size, dtype=torch.int64)
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 apply(
self,
layer: torch.nn.Module,
x: torch.Tensor,
bias: torch.Tensor | None = None,
tp_rank: int | None = 0,
) -> torch.Tensor:
if x.dtype != torch.int8:
layer_cls_name = layer.__class__.__name__
weight_prefetch_method = get_weight_prefetch_method()
# prefetch qkvo_proj.weight preprocess
weight_prefetch_method.maybe_prefetch_attn_weight_preprocess(
layer_cls_name=layer_cls_name,
weight=layer.weight,
start_flag=x,
)
try:
quant_comm_config = layer._quant_comm_config
except AttributeError:
quant_comm_config = {}
comm_fn = quant_comm_config.get("communication_fn")
enable_flashcomm2_quant_comm = comm_fn is not None and (
"o_proj" in layer.prefix or "out_proj" in layer.prefix
)
if enable_flashcomm2_quant_comm:
quant_input_x = x.contiguous().view(-1, layer.aclnn_input_scale_reciprocal.size(0))
quant_x = torch.ops.vllm.quantize(
quant_input_x,
layer.aclnn_input_scale,
layer.aclnn_input_scale_reciprocal,
layer.aclnn_input_offset,
)
comm_input = quant_x.view(x.size(0), -1)
assert comm_fn is not None
x = comm_fn(comm_input)
else:
# quant
x = torch.ops.vllm.quantize(
x,
layer.aclnn_input_scale,
layer.aclnn_input_scale_reciprocal,
layer.aclnn_input_offset,
)
# prefetch qkvo_proj.weight postprocess
weight_prefetch_method.maybe_prefetch_attn_weight_postprocess(
layer_cls_name=layer_cls_name,
stop_flag=x,
)
quant_bias = layer.quant_bias if tp_rank == 0 else None
try:
ascend_quant_method = layer.ascend_quant_method
except AttributeError:
ascend_quant_method = ""
if ascend_quant_method == COMPRESSED_TENSORS_METHOD:
quant_bias = bias
output = torch_npu.npu_quant_matmul(
x,
layer.weight,
layer.deq_scale,
bias=quant_bias,
output_dtype=layer.params_dtype,
)
return output
def process_weights_after_loading(self, layer):
expanding_factor = layer.weight.data.shape[1]
layer.aclnn_input_scale = torch.nn.Parameter(
layer.input_scale.data.repeat(expanding_factor), requires_grad=False
)
layer.aclnn_input_scale_reciprocal = 1 / torch.nn.Parameter(
layer.input_scale.data.repeat(expanding_factor), requires_grad=False
)
layer.aclnn_input_offset = torch.nn.Parameter(
layer.input_offset.data.repeat(expanding_factor), requires_grad=False
).to(layer.aclnn_input_scale.dtype)
layer.weight.data = layer.weight.data.transpose(0, 1).contiguous()
layer.weight.data = maybe_trans_nz(layer.weight.data)
layer.weight_scale.data = torch.flatten(layer.weight_scale.data)
layer.weight_offset.data = torch.flatten(layer.weight_offset.data)
ascend_quant_method = getattr(layer, "ascend_quant_method", "")
if ascend_quant_method == COMPRESSED_TENSORS_METHOD:
deq_scale = layer.input_scale.data * layer.weight_scale.data
layer.deq_scale = torch.nn.Parameter(deq_scale, requires_grad=False)

View File

@@ -0,0 +1,102 @@
#
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
# This file is a part of the vllm-ascend project.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
#
from typing import Any
import torch
from .base import QuantType
from .registry import register_scheme
from .w8a8_dynamic import AscendW8A8DynamicFusedMoEMethod, AscendW8A8DynamicLinearMethod
@register_scheme("W8A8FP8_DYNAMIC", "linear")
class AscendW8A8FP8DynamicLinearMethod(AscendW8A8DynamicLinearMethod):
"""Linear method for Ascend W8A8FP8_DYNAMIC.
This scheme uses FP8 dynamic per-token quantization for activations
and FP8 per-channel quantization for weights.
"""
act_quant_type: torch.dtype = torch.float8_e4m3fn
def __init__(self):
pass
def get_weight(self, input_size: int, output_size: int, params_dtype: torch.dtype) -> dict[str, Any]:
params_dict = {"weight": torch.empty(output_size, input_size, dtype=torch.float8_e4m3fn)}
return params_dict
def get_perchannel_param(
self,
output_size: int,
params_dtype: torch.dtype,
) -> dict[str, Any]:
params_dict = {}
params_dict["weight_scale"] = torch.empty(output_size, 1, dtype=torch.float32)
params_dict["weight_offset"] = torch.empty(output_size, 1, dtype=params_dtype)
return params_dict
def apply(
self,
layer: torch.nn.Module,
x: torch.Tensor,
bias: torch.Tensor | None = None,
tp_rank: int | None = 0,
) -> torch.Tensor:
output = super().apply(layer, x, bias, tp_rank)
# TODO: there is a bug in npu_quant_matmul for fp8 with bias
# after the bug is fixed, the whole apply method can be removed.
if bias is not None:
output = (output + bias).to(x.dtype)
return output
@register_scheme("W8A8FP8_DYNAMIC", "moe")
class AscendW8A8FP8DynamicFusedMoEMethod(AscendW8A8DynamicFusedMoEMethod):
"""FusedMoE method for Ascend W8A8FP8_DYNAMIC."""
quant_type: QuantType = QuantType.W8A8FP8
def __init__(self):
super().__init__()
def get_weight(
self, num_experts: int, intermediate_size_per_partition: int, hidden_sizes: int, params_dtype: torch.dtype
) -> dict[str, Any]:
param_dict = {}
param_dict["w13_weight"] = torch.empty(
num_experts, 2 * intermediate_size_per_partition, hidden_sizes, dtype=torch.float8_e4m3fn
)
param_dict["w2_weight"] = torch.empty(
num_experts, hidden_sizes, intermediate_size_per_partition, dtype=torch.float8_e4m3fn
)
return param_dict
def get_dynamic_quant_param(
self, num_experts: int, intermediate_size_per_partition: int, hidden_sizes: int, params_dtype: torch.dtype
) -> dict[str, Any]:
param_dict = {}
param_dict["w13_weight_scale"] = torch.empty(
num_experts, 2 * intermediate_size_per_partition, 1, dtype=torch.float32
)
param_dict["w13_weight_offset"] = torch.empty(
num_experts, 2 * intermediate_size_per_partition, 1, dtype=params_dtype
)
param_dict["w2_weight_scale"] = torch.empty(num_experts, hidden_sizes, 1, dtype=torch.float32)
param_dict["w2_weight_offset"] = torch.empty(num_experts, hidden_sizes, 1, dtype=params_dtype)
return param_dict