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

View File

@@ -0,0 +1,104 @@
# Copyright (c) 2024; NVIDIA CORPORATION. All rights reserved.
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
# Copyright 2023 The vLLM team.
# 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 torch
import torch.distributed
import torch.distributed as dist
import torch_npu
COMM_STREAM = None
def async_all_to_all(input_, output_split_sizes, input_split_sizes, group, event=None):
if output_split_sizes is None:
# Equal split (all2all)
a2a_out = torch.empty_like(input_)
else:
# Unequal split (all2all-v)
a2a_out = input_.new_empty(
size=[sum(output_split_sizes)] + list(input_.size()[1:]),
dtype=input_.dtype,
device=torch.npu.current_device(),
)
if event:
# multi stream wait event
global COMM_STREAM
if COMM_STREAM is None:
COMM_STREAM = torch_npu.npu.Stream(device=torch.npu.current_device())
with torch_npu.npu.stream(COMM_STREAM):
event.wait()
handle = dist.all_to_all_single(
a2a_out,
input_.contiguous(),
output_split_sizes=output_split_sizes,
input_split_sizes=input_split_sizes,
group=group,
async_op=True,
)
else:
handle = dist.all_to_all_single(
a2a_out,
input_.contiguous(),
output_split_sizes=output_split_sizes,
input_split_sizes=input_split_sizes,
group=group,
async_op=True,
)
return input_, a2a_out, handle
def _gather_along_first_dim(input_, group, output_split_sizes=None):
"""Gather tensors and concatenate along the first dimension.
Args:
input_tensor (torch.Tensor):
A tensor to be gathered.
output_split_sizes (List[int], optional):
A list specifying the sizes of the output splits along the first dimension.
If None, equal splitting is assumed. Default: None.
Returns:
torch.Tensor: Gathered tensor.
"""
world_size = torch.distributed.get_world_size(group)
# Bypass the function if we are using only 1 GPU.
if world_size == 1:
return input_
dim_size = list(input_.size())
if output_split_sizes is None:
dim_size[0] = dim_size[0] * world_size
output = torch.empty(dim_size, dtype=input_.dtype, device=torch.npu.current_device())
torch.distributed.all_gather_into_tensor(output, input_.contiguous(), group=group)
else:
dim_size[0] = sum(output_split_sizes)
output = torch.empty(dim_size, dtype=input_.dtype, device=torch.npu.current_device())
output_tensor_list = list(torch.split(output, output_split_sizes, dim=0))
torch.distributed.all_gather(output_tensor_list, input_, group=group)
return output
def gather_from_sequence_parallel_region(
input_,
group,
output_split_sizes=None,
):
"""Wrapper for autograd function: forward: AG, backward: RS <first dim>"""
return _gather_along_first_dim(input_, group, output_split_sizes)

View File

@@ -0,0 +1,417 @@
#
# 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
import torch
import torch.nn.functional as F
from vllm.distributed import get_tp_group
from vllm.forward_context import get_forward_context
from vllm_ascend.ascend_forward_context import MoECommType
from vllm_ascend.device.device_op import DeviceOperator
from vllm_ascend.distributed.utils import split_tensor_along_first_dim
from vllm_ascend.utils import get_weight_prefetch_method
def select_experts(
hidden_states: torch.Tensor,
router_logits: torch.Tensor,
top_k: int,
use_grouped_topk: bool,
renormalize: bool,
topk_group: int | None = None,
num_expert_group: int | None = None,
custom_routing_function: Callable | None = None,
scoring_func: str = "softmax",
routed_scaling_factor=1.0,
e_score_correction_bias: torch.Tensor | None = None,
indices_type: torch.dtype | None = None,
mix_placement: bool = False,
num_logical_experts: int = -1,
num_shared_experts: int = 0,
num_experts: int = -1,
input_ids: torch.Tensor | None = None,
tid2eid: torch.Tensor | None = None,
):
"""
Fused experts with select experts.
Args:
router_logits: router logits of shape (num_tokens, hidden_size).
hidden_states: Hidden states of shape (num_tokens, hidden_size).
top_k: number of top k experts.
use_grouped_topk: Whether to group experts before selecting top-k.
renormalize: Whether to renormalize the routing weights.
topk_group: Number of expert groups to select from.
num_expert_group: Number of experts in each group.
custom_routing_function: Custom routing function.
scoring_func: Scoring function to use.
e_score_correction_bias: Correction bias to apply to expert scores.
indices_type: dtype of indices
num_experts: Number of experts.
Returns:
topk_weights: router weights of shape (num_tokens, top_k).
topk_ids: selected expert IDs of shape (num_tokens, top_k).
"""
# prefetch w1_w3_proj.weight preprocess
weight_prefetch_method = get_weight_prefetch_method()
if weight_prefetch_method:
weight_prefetch_method.maybe_prefetch_moe_weight_preprocess(hidden_states, "gate_up")
is_support_npu_moe_gating_top_k = check_npu_moe_gating_top_k(
hidden_states=hidden_states,
top_k=top_k,
renormalize=renormalize,
topk_group=topk_group,
num_expert_group=num_expert_group,
scoring_func=scoring_func,
custom_routing_function=custom_routing_function,
)
if is_support_npu_moe_gating_top_k:
topk_weights, topk_ids = _select_experts_with_fusion_ops(
hidden_states=hidden_states,
router_logits=router_logits,
top_k=top_k,
use_grouped_topk=use_grouped_topk,
topk_group=topk_group,
renormalize=renormalize,
e_score_correction_bias=e_score_correction_bias,
num_expert_group=num_expert_group,
scoring_func=scoring_func,
routed_scaling_factor=routed_scaling_factor,
tid2eid=tid2eid,
input_ids=input_ids,
)
else:
topk_weights, topk_ids = _native_select_experts(
hidden_states=hidden_states,
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,
tid2eid=None,
input_ids=None,
)
# Apply routed scaling factor to weights
if routed_scaling_factor != 1.0:
topk_weights = topk_weights * routed_scaling_factor
if mix_placement:
shared_expert_routing_factor = 1.0 if is_support_npu_moe_gating_top_k else (1 / routed_scaling_factor)
batch_size = topk_ids.shape[0]
pad_shared_expert_ids = torch.arange(
num_logical_experts, num_logical_experts + num_shared_experts, dtype=topk_ids.dtype, device=topk_ids.device
).repeat(batch_size, 1)
pad_shared_expert_weights = torch.full(
(topk_weights.shape[0], num_shared_experts),
shared_expert_routing_factor,
dtype=topk_weights.dtype,
device=topk_weights.device,
)
topk_ids = torch.cat([topk_ids, pad_shared_expert_ids], dim=1)
topk_weights = torch.cat([topk_weights, pad_shared_expert_weights], dim=1)
return topk_weights, topk_ids
def check_npu_moe_gating_top_k(
hidden_states: torch.Tensor,
top_k: int,
renormalize: bool,
topk_group: int | None = None,
num_expert_group: int | None = None,
scoring_func: str = "softmax",
custom_routing_function: Callable | None = None,
):
if scoring_func == "sigmoid" and not renormalize: # sigmoid + renorm=0 is not supported in current branch
return False
if custom_routing_function is not None:
return False
if scoring_func != "softmax" and scoring_func != "sigmoid" and scoring_func != "sqrtsoftplus":
return False
topk_group = topk_group if topk_group is not None else 1
num_expert_group = num_expert_group if num_expert_group is not None else 1
if not (
num_expert_group > 0
and hidden_states.shape[-1] % num_expert_group == 0
and hidden_states.shape[-1] // num_expert_group > 2
):
return False
if topk_group < 1 or topk_group > num_expert_group:
return False
if top_k < 1 or top_k > (hidden_states.shape[-1] / (num_expert_group * topk_group)):
return False
if topk_group * hidden_states.shape[-1] / num_expert_group < top_k: # noqa: SIM103
return False
return True
def _native_grouped_topk(
topk_weights: torch.Tensor,
num_expert_group: int | None,
topk_group: int | None,
):
topk_group = 0 if topk_group is None else topk_group
num_expert_group = 0 if num_expert_group is None else num_expert_group
num_token = topk_weights.shape[0]
grouped_weights = topk_weights.view(num_token, num_expert_group, -1).max(dim=-1).values
topk_group_indices = torch.topk(grouped_weights.to(torch.float32), k=topk_group, dim=-1, sorted=False)[1]
topk_group_mask = torch.zeros_like(grouped_weights)
topk_group_mask.scatter_(1, topk_group_indices, 1)
topk_weight_mask = (
topk_group_mask.unsqueeze(-1)
.expand(num_token, num_expert_group, topk_weights.shape[-1] // num_expert_group)
.reshape(num_token, -1)
)
topk_weights = topk_weights.masked_fill(~topk_weight_mask.bool(), 0.0)
return topk_weights
def _renormalize_topk_weights(
topk_weights: torch.Tensor,
renormalize: bool,
):
if renormalize:
topk_weights = topk_weights / topk_weights.sum(dim=-1, keepdim=True)
return topk_weights
def _select_expert_use_group_topk(
topk_weights: torch.Tensor,
topk_group: int | None,
renormalize: bool,
top_k: int,
num_expert_group: int | None,
e_score_correction_bias: torch.Tensor | None,
):
assert topk_group is not None
assert num_expert_group is not None
if e_score_correction_bias is not None:
# Store original scores before applying correction bias. We use biased
# scores for expert selection but original scores for routing weights
original_weights = topk_weights
topk_weights = topk_weights + e_score_correction_bias.unsqueeze(0)
# TODO: Change to npu_group_topk when the latest CANN and NNAL is available
# >>> torch_npu._npu_group_topk(topk_weights, group_num=num_expert_group, k=topk_group)
topk_weights = _native_grouped_topk(topk_weights, num_expert_group, topk_group)
# TODO bfloat16 is not supported in torch.topk with ge graph.
if e_score_correction_bias is not None:
topk_ids = torch.topk(topk_weights.to(torch.float32), k=top_k, dim=-1, sorted=False)[1]
# Use original unbiased scores for the routing weights
topk_weights = original_weights.gather(1, topk_ids)
else:
topk_weights, topk_ids = torch.topk(topk_weights.to(torch.float32), k=top_k, dim=-1, sorted=False)
topk_ids = topk_ids.to(torch.int32)
topk_weights = _renormalize_topk_weights(topk_weights, renormalize)
return topk_weights, topk_ids
def _select_experts_with_fusion_ops(
hidden_states: torch.Tensor,
router_logits: torch.Tensor,
top_k: int,
use_grouped_topk: bool,
renormalize: bool,
e_score_correction_bias: torch.Tensor | None,
topk_group: int | None,
num_expert_group: int | None,
scoring_func: str = "softmax",
routed_scaling_factor=1.0,
tid2eid=None,
input_ids=None,
):
topk_group = topk_group if topk_group is not None else 1
num_expert_group = num_expert_group if num_expert_group is not None else 1
renorm = int(renormalize)
if scoring_func == "sqrtsoftplus":
if tid2eid is not None:
forward_context = get_forward_context()
input_ids = forward_context.input_ids.to(torch.int64)
# tid2eid_ones = torch.ones(tid2eid.shape[0],tid2eid.shape[1],device=router_logits.device,dtype=torch.int32)
tid2eid_ones = tid2eid.to(torch.int32)
if forward_context.moe_comm_type == MoECommType.ALLGATHER:
prepare_finalize = forward_context.moe_comm_method.prepare_finalize
input_ids = prepare_finalize.all_gather_input_id_with_dp_group(input_ids)
else:
input_ids = forward_context.moe_comm_method.pad_and_split_input_ids(input_ids)
if forward_context.flash_comm_v1_enabled and forward_context.moe_comm_type != MoECommType.ALLGATHER:
# Process for Flash Comm V1
tp_size = get_tp_group().world_size
tp_rank = get_tp_group().rank_in_group
splitted_input = split_tensor_along_first_dim(input_ids, num_partitions=tp_size)
input_ids = splitted_input[tp_rank].contiguous()
input_ids = torch.where(input_ids == -1, 0, input_ids)
else:
input_ids = None
tid2eid_ones = None
topk_weights, topk_ids, _ = torch.ops._C_ascend.moe_gating_top_k_hash(
x=router_logits,
k=top_k,
bias=e_score_correction_bias,
input_ids=input_ids,
tid2eid=tid2eid_ones,
k_group=topk_group,
group_count=num_expert_group,
routed_scaling_factor=routed_scaling_factor,
eps=1e-20,
group_select_mode=1,
# The hash custom op currently rejects renorm != 0. Apply
# norm_topk_prob in Python below before returning to MoE compute.
renorm=0,
norm_type=2,
out_flag=False,
)
return topk_weights, topk_ids
norm_type = 0 if scoring_func == "softmax" else 1
if e_score_correction_bias is not None and e_score_correction_bias.dtype != router_logits.dtype:
e_score_correction_bias = e_score_correction_bias.to(router_logits.dtype)
topk_weights, topk_ids, _ = DeviceOperator.moe_gating_top_k(
router_logits,
k=top_k,
k_group=topk_group,
group_count=num_expert_group,
group_select_mode=1,
renorm=renorm,
norm_type=norm_type, # 0: softmax; 1: sigmoid
out_flag=False,
routed_scaling_factor=routed_scaling_factor,
eps=1e-20,
bias_opt=e_score_correction_bias,
)
return topk_weights, topk_ids
def _native_select_experts(
hidden_states: torch.Tensor,
router_logits: torch.Tensor,
top_k: int,
use_grouped_topk: bool,
renormalize: bool,
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,
use_hash: bool = False,
tid2eid: dict[int, int] | None = None,
input_ids: torch.Tensor | None = None,
) -> tuple[torch.Tensor, torch.Tensor]:
"""
Select top-k experts based on router logits.
Args:
hidden_states: Hidden states of shape (num_tokens, hidden_size).
router_logits: Router logits of shape (num_tokens, num_experts).
top_k: Number of experts to select.
use_grouped_topk: Whether to group experts before selecting top-k.
renormalize: Whether to renormalize the routing weights.
topk_group: Number of expert groups to select from.
num_expert_group: Number of experts in each group.
custom_routing_function: Custom routing function.
scoring_func: Scoring function to use.
e_score_correction_bias: Correction bias to apply to expert scores.
Returns:
topk_weights: Routing weights of shape (num_tokens, top_k).
topk_ids: Selected expert IDs of shape (num_tokens, top_k).
Raises:
ValueError: If an unsupported scoring function is provided.
"""
if scoring_func == "softmax":
topk_weights = router_logits.softmax(dim=-1)
elif scoring_func == "sigmoid":
topk_weights = router_logits.sigmoid()
elif scoring_func == "sqrtsoftplus":
topk_weights = F.softplus(router_logits).sqrt()
else:
raise ValueError(f"Unsupported scoring function: {scoring_func}")
if use_grouped_topk:
topk_weights, topk_ids = _select_expert_use_group_topk(
topk_weights=topk_weights,
top_k=top_k,
renormalize=renormalize,
topk_group=topk_group,
num_expert_group=num_expert_group,
e_score_correction_bias=e_score_correction_bias,
)
return topk_weights * routed_scaling_factor, topk_ids
if e_score_correction_bias is not None:
topk_weights = topk_weights + e_score_correction_bias
if custom_routing_function is not None:
topk_weights, topk_ids = custom_routing_function(
hidden_states=hidden_states,
gating_output=router_logits,
topk=top_k,
renormalize=renormalize,
)
# Required by npu_moe_init_routing
topk_ids = topk_ids.to(torch.int32)
return topk_weights, topk_ids
topk_weights, topk_ids = topk_weights.topk(top_k, dim=-1)
topk_weights = topk_weights.to(hidden_states.dtype)
# Required by npu_moe_init_routing
topk_ids = topk_ids.to(torch.int32)
topk_weights = _renormalize_topk_weights(topk_weights, renormalize)
topk_weights = topk_weights * routed_scaling_factor
return topk_weights, topk_ids
def zero_experts_compute(
expert_indices: torch.Tensor,
expert_scales: torch.Tensor,
num_experts: int,
zero_expert_type: str,
hidden_states: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
if zero_expert_type == "identity":
zero_expert_mask = expert_indices < num_experts
zero_expert_scales = expert_scales.clone()
zero_expert_scales = torch.where(zero_expert_mask, 0.0, zero_expert_scales)
hidden_states = hidden_states.unsqueeze(1)
zero_expert_scales = zero_expert_scales.unsqueeze(2)
result = hidden_states * zero_expert_scales
result = result.sum(dim=1)
normal_expert_mask = expert_indices >= num_experts
expert_indices = torch.where(normal_expert_mask, 0, expert_indices)
expert_scales = torch.where(normal_expert_mask, 0.0, expert_scales)
return expert_indices, expert_scales, result

View File

@@ -0,0 +1,894 @@
#
# 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.
#
# ruff: noqa: E501
from collections.abc import Callable
from dataclasses import dataclass, field
from functools import wraps
import torch
import torch.nn.functional as F
import torch_npu
from vllm.config import get_current_vllm_config
from vllm.distributed import get_dp_group, get_ep_group, get_tp_group, tensor_model_parallel_all_reduce
from vllm.forward_context import get_forward_context
from vllm.logger import logger
from vllm.model_executor.layers.fused_moe.config import FusedMoEConfig
from vllm.model_executor.layers.fused_moe.layer import (
FusedMoE, # noqa: F401
MoERunner,
)
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.eplb.adaptor.vllm_adaptor import VllmEplbAdaptor
from vllm_ascend.eplb.core.eplb_utils import init_eplb_config
from vllm_ascend.flash_common3_context import get_flash_common3_context, set_flash_common3_context
from vllm_ascend.ops.fused_moe.experts_selector import select_experts, zero_experts_compute
from vllm_ascend.ops.fused_moe.moe_comm_method import AllGatherCommImpl, FusedExpertsResult, setup_moe_comm_method
from vllm_ascend.ops.fused_moe.moe_runtime_args import build_fused_experts_input
from vllm_ascend.quantization.methods.base import get_moe_num_logical_experts
from vllm_ascend.quantization.quant_type import QuantType
from vllm_ascend.utils import (
ACL_FORMAT_FRACTAL_NZ,
maybe_trans_nz,
npu_stream_switch,
shared_expert_dp_enabled,
shared_experts_calculation_stream,
vllm_version_is,
)
if vllm_version_is("0.23.0"):
from vllm.model_executor.layers.fused_moe.layer import UnquantizedFusedMoEMethod
else:
from vllm.model_executor.layers.fused_moe.unquantized_fused_moe_method import UnquantizedFusedMoEMethod
def get_compressed_expert_map(expert_map: torch.Tensor) -> str:
global_indices = torch.where(expert_map != -1)[0]
local_indices = expert_map[global_indices]
return ", ".join(
f"{local_index.item()}->{global_index.item()}"
for local_index, global_index in zip(local_indices, global_indices)
)
@dataclass
class FusedMoEResult:
routed_out: torch.Tensor
before_dispatch_evt: torch.npu.Event | None = None
before_gmm2_evt: torch.npu.Event | None = None
before_combine_evt: torch.npu.Event | None = None
swiglu_limit: float = 0.0
@dataclass
class FusedMoEEvents:
before_routed_experts: torch.npu.Event
after_routed_experts: torch.npu.Event | None = field(default=None)
before_dispatch: torch.npu.Event | None = field(default=None)
before_gmm2: torch.npu.Event | None = field(default=None)
before_combine: torch.npu.Event | None = field(default=None)
swiglu_limit: float = 0.0
def mock_false():
return False
def mock_true():
return True
class AscendUnquantizedFusedMoEMethod(UnquantizedFusedMoEMethod):
def __init__(self, moe: FusedMoEConfig = None, tid2eid=None):
super().__init__(moe=moe)
self.dynamic_eplb = get_ascend_config().eplb_config.dynamic_eplb
self.tid2eid = tid2eid
@property
def is_monolithic(self) -> bool:
return False
def maybe_make_prepare_finalize(self, routing_tables=None):
# Ascend uses its own MoE communication and forward_impl path.
# Do not let upstream modular-kernel initialization replace it.
return None
def process_weights_after_loading(self, layer):
super(UnquantizedFusedMoEMethod, self).process_weights_after_loading(layer)
w13_data = self._maybe_pad_weight(layer.w13_weight.data).transpose(1, 2).contiguous()
layer.w13_weight = torch.nn.Parameter(w13_data, requires_grad=False)
w2_data = self._maybe_pad_weight(layer.w2_weight.data).transpose(1, 2).contiguous()
layer.w2_weight = torch.nn.Parameter(w2_data, requires_grad=False)
# TODO: Current dispatch_ffn_combine fusion operator ONLY supports NZ format.
# Therefore, we must cast weights to NZ when fusion is enabled.
# Once the underlying dispatch_ffn_combine operator is updated to support
# ND format (or other formats), remove this specific 'if' check and the forced
# npu_format_cast. At that point, the operator should be able to handle weights
# in their native format without explicit casting here.
enable_fused_mc2 = get_ascend_config().enable_fused_mc2
if enable_fused_mc2:
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)
if enable_fused_mc2 == 1 and 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)]
del layer.w13_weight
del layer.w2_weight
torch.npu.empty_cache()
else:
layer.w13_weight.data = maybe_trans_nz(layer.w13_weight.data)
layer.w2_weight.data = maybe_trans_nz(layer.w2_weight.data)
def apply(
self,
layer: torch.nn.Module,
x: torch.Tensor,
use_grouped_topk: bool,
top_k: int,
router_logits: torch.Tensor,
renormalize: bool,
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,
num_experts: int = -1,
expert_map: torch.Tensor | None = None,
apply_router_weight_on_input: bool = False,
activation: str = "silu",
enable_force_load_balance: bool = False,
log2phy: torch.Tensor = None,
global_redundant_expert_num: int = 0,
pertoken_scale: torch.Tensor | None = None,
mc2_mask: torch.Tensor | None = None,
) -> torch.Tensor:
zero_expert_num = getattr(layer, "zero_expert_num", 0)
zero_expert_type = getattr(layer, "zero_expert_type", None)
input_ids = getattr(get_forward_context(), "input_ids", None)
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,
)
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=self.tid2eid,
input_ids=input_ids,
)
if vllm_version_is("0.23.0"):
model_config = layer.vllm_config.model_config
else:
try:
_vllm_config = get_current_vllm_config()
except AssertionError:
_vllm_config = None
model_config = None if _vllm_config is None else _vllm_config.model_config
if model_config is not None and model_config.enable_return_routed_experts:
capturer = getattr(layer, "_ascend_routed_experts_capturer", None)
if capturer is not None:
capturer.capture(layer_id=layer.layer_id, topk_ids=topk_ids)
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,
)
topk_weights = topk_weights.to(x.dtype)
# 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)
moe_comm_method = _EXTRA_CTX.moe_comm_method
# NOTE: In the MoECommType.FUSED_MC2 branch, we wrap weights (w1, w2) into lists
# and provide dummy scales (w1_scale, w2_scale). This is required because:
# The underlying Ascend fused operator (e.g., dispatch_ffn_combine) expects
# inputs in a list format.
# TODO: Passing an empty tensor as scale for float (BF16) cases is semantically
# incorrect. The ideal solution is to pass None. However, if the underlying
# dispatch_ffn_combine C++ operator does not support None for the scale argument
# (due to signature constraints), we are forced to use a placeholder empty tensor.
# This TODO tracks the requirement to update the C++ operator to accept Optional[Tensor]
# or None for scales in non-quantized scenarios.
w13_weight_list = getattr(layer, "w13_weight_list", None)
w2_weight_list = getattr(layer, "w2_weight_list", None)
has_split_weight_lists = isinstance(w13_weight_list, list) and isinstance(w2_weight_list, list)
if _EXTRA_CTX.moe_comm_type == MoECommType.FUSED_MC2:
if self.dynamic_eplb and not has_split_weight_lists:
logger.warning_once(
"FUSED_MC2 is enabled with dynamic EPLB, but unquantized MoE weights are not split into "
"tensor lists. This may cause accuracy issues or communication hangs."
)
w1 = w13_weight_list if isinstance(w13_weight_list, list) else [layer.w13_weight]
w2 = w2_weight_list if isinstance(w2_weight_list, list) else [layer.w2_weight]
w1_scale = [torch.tensor([], dtype=torch.int64)]
w2_scale = [torch.tensor([], dtype=torch.int64)]
w1_scale_bias = [torch.tensor([], dtype=torch.float32)]
w2_scale_bias = [torch.tensor([], dtype=torch.float32)]
else:
w1 = w13_weight_list if isinstance(w13_weight_list, list) else layer.w13_weight
w1_scale = None
w2 = w2_weight_list if isinstance(w2_weight_list, list) else layer.w2_weight
w2_scale = None
w1_scale_bias = None
w2_scale_bias = 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,
w1_bias=layer.w13_bias if self.moe.has_bias else None,
w2_bias=layer.w2_bias if self.moe.has_bias else None,
quant_type=QuantType.NONE,
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,
# Per-layer MoE LoRA state, set once by AscendFusedMoEWithLoRA
# when an adapter wraps this layer; None for non-LoRA layers.
lora_context=getattr(layer, "_ascend_moe_lora_context", None),
)
)
if zero_expert_num > 0 and zero_expert_type is not None:
final_hidden_states += zero_expert_result
return final_hidden_states
if vllm_version_is("0.23.0"):
from vllm_ascend.ops.fused_moe.fused_moe_0_23_0 import AscendFusedMoE, AscendMoERunner
AscendFusedMoE.__module__ = __name__
AscendMoERunner.__module__ = __name__
else:
class AscendMoERunner(MoERunner): # type: ignore[no-redef]
moe_counter = -1
gate_stream: torch.npu.Stream | None = None
def __init__(
self,
layer_name,
moe_config,
router,
routed_experts,
enable_dbo=False,
gate=None,
shared_experts=None,
shared_expert_gate=None,
routed_input_transform=None,
routed_output_transform=None,
routed_scaling_factor=1,
tid2eid=None,
n_shared_experts: int = 0,
):
super().__init__(
layer_name,
moe_config,
router,
routed_experts,
enable_dbo,
gate,
shared_experts,
shared_expert_gate,
routed_input_transform,
routed_output_transform,
routed_scaling_factor,
)
self.top_k = moe_config.experts_per_token
self._gate = gate
self.hidden_size = moe_config.hidden_dim
# Routing params — all stored on routed_experts by the factory function.
self.use_grouped_topk = routed_experts.use_grouped_topk
self.renormalize = routed_experts.renormalize
self.topk_group = routed_experts.topk_group
self.num_expert_group = routed_experts.num_expert_group
self.custom_routing_function = routed_experts.custom_routing_function
self.scoring_func = routed_experts.scoring_func
self._original_routed_scaling_factor = routed_experts.routed_scaling_factor
self.e_score_correction_bias = routed_experts.e_score_correction_bias
self.apply_router_weight_on_input = routed_experts.apply_router_weight_on_input
# Ascend-specific: not stored on RoutedExperts, passed via runner_args.
self.tid2eid = tid2eid
# Replace quant_method on routed_experts with the Ascend version.
# Must NOT set self.quant_method (instance attr) — vllm loader scans all
# nn.Module children for quant_method and would call process_weights_after_loading
# on the runner with wrong expectations. Set on routed_experts so that
# self._quant_method (the MoERunner property) returns the Ascend version.
if routed_experts.quant_config is None:
routed_experts.quant_method = AscendUnquantizedFusedMoEMethod(self.moe_config, tid2eid=self.tid2eid)
else:
routed_experts.quant_method = routed_experts.quant_config.get_quant_method(
routed_experts, self.layer_name, tid2eid=self.tid2eid
)
self.quant_type = self._get_quant_type()
# Can be removed after vllm fixes the issue.
if self._needs_routed_expert_parameter_aliases():
self._register_routed_expert_parameter_aliases()
self.moe_config.tp_group = get_tp_group()
self.moe_config.dp_group = get_dp_group()
if self.moe_config.ep_size > 1:
self.moe_config.ep_group = get_ep_group()
self.moe_config.mc2_group = get_mc2_group()
ascend_config = get_ascend_config()
self._shared_experts = shared_experts
has_shared_experts = self._shared_experts is not None
self.shared_multistream_overlap_gate = ascend_config.multistream_overlap_gate and has_shared_experts
self.enable_npugraph_ex_static_kernel = ascend_config.ascend_compilation_config.enable_static_kernel
# flashcommon3 gate stream
self.multistream_overlap_gate = ascend_config.multistream_overlap_gate
if self.multistream_overlap_gate and AscendMoERunner.gate_stream is None:
AscendMoERunner.gate_stream = torch.npu.Stream()
vllm_config = get_current_vllm_config()
if (
self.custom_routing_function is None
and self.e_score_correction_bias is not None
and not vllm_config.model_config.is_deepseek_mla
):
self.e_score_correction_bias.data = self.e_score_correction_bias.data.to(
dtype=vllm_config.model_config.dtype
)
self.enable_shared_expert_dp = ascend_config.enable_shared_expert_dp
self.multistream_overlap_shared_expert = (
ascend_config.multistream_overlap_shared_expert and shared_experts is not None
)
mix_placement = getattr(ascend_config, "mix_placement", False)
# EPLB initialization (Ascend-specific; mirrors old AscendFusedMoE logic).
AscendMoERunner.moe_counter += 1
self.moe_instance_id = AscendMoERunner.moe_counter
eplb_config = ascend_config.eplb_config
if mix_placement:
moe_config.num_experts += n_shared_experts
(
self.global_expert_map,
self._expert_map,
self.log2phy,
self.global_redundant_expert_num,
) = init_eplb_config(
eplb_config,
AscendMoERunner.moe_counter,
moe_config,
mix_placement,
n_shared_experts,
tp_size=vllm_config.parallel_config.tensor_parallel_size,
)
moe_config.global_redundant_expert_num = self.global_redundant_expert_num
local_num_experts = (moe_config.num_experts + self.global_redundant_expert_num) // moe_config.ep_size
moe_config.num_local_experts = local_num_experts
routed_experts.expert_map_manager._local_num_experts = local_num_experts
routed_experts.expert_map_manager._expert_map = self._expert_map
self.dynamic_eplb = eplb_config.dynamic_eplb and (self.log2phy is not None)
self.multi_stage = False
self.moe_load = torch.zeros(local_num_experts, dtype=torch.int64).npu()
if self.dynamic_eplb and eplb_config.expert_heat_collection_interval > 1:
self.multi_stage = True
self.load_counter = torch.tensor(0, dtype=torch.int32, device="npu")
self.num_iter = eplb_config.expert_heat_collection_interval
self.moe_load = torch.zeros((self.num_iter, local_num_experts), dtype=torch.int32, device="npu")
setup_moe_comm_method(self.moe_config)
if self.multistream_overlap_shared_expert:
# Wrap the quant_method's process_weights_after_loading to validate that
# splitting shared expert computation (gate_up projection + activation,
# then down projection) yields identical results to integrated
# computation after weight loading.
original_process_weights = self._quant_method.process_weights_after_loading
@wraps(original_process_weights)
def wrapped_process_weights(*args, **kwargs):
result = original_process_weights(*args, **kwargs)
self._validate_shared_expert_consistency()
return result
self._quant_method.process_weights_after_loading = wrapped_process_weights # type: ignore
# Register this MoE layer with EPLB for PP compatibility.
# PPMissingLayer (nn.Identity) never calls AscendFusedMoE.__init__,
# so only real MoE layers on this rank are registered.
VllmEplbAdaptor.register_layer(self)
def _validate_shared_expert_consistency(self):
"""Validate that split shared expert computation matches integrated computation."""
test_input = (
torch.rand(10, self.hidden_size, device="npu", dtype=self.moe_config.in_dtype) * 2 - 1
) # Random input for testing, scoped to [-1, 1]
assert self._shared_experts is not None
integrated_out = self._shared_experts(test_input)
part1_out = self._shared_experts_part1(test_input)
split_out = self._shared_experts_part2(test_input, part1_out)
if not torch.allclose(integrated_out, split_out):
diff = (integrated_out - split_out).abs()
logger.error(
"[fused_moe/layer] Shared expert split computation validation failed."
" The split-path computation does not match the integrated-path result."
" max_abs_diff=%s, integrated_sum=%s, integrated_norm=%s,"
" split_sum=%s, split_norm=%s, hidden_size=%s, dtype=%s.",
diff.max().item(),
integrated_out.sum().item(),
integrated_out.norm().item(),
split_out.sum().item(),
split_out.norm().item(),
self.hidden_size,
self.moe_config.in_dtype,
)
raise ValueError("FusedMoE shared experts split computation does not match the integrated computation.")
logger.info_once(
"[fused_moe/layer] Shared expert split computation validation passed."
" Integrated and split-path results are consistent."
)
def _shared_experts_part1(self, hidden_states: torch.Tensor):
shared_gate_up, _ = self._shared_experts.gate_up_proj(hidden_states) # type: ignore
return shared_gate_up
def _shared_experts_part2(self, hidden_states: torch.Tensor, shared_gate_up: torch.Tensor):
shared_act = self._shared_experts.act_fn(shared_gate_up) # type: ignore
shared_out, _ = self._shared_experts.down_proj(shared_act) # type: ignore
# Qwen3-Next specific gating mechanism
assert self._shared_experts is not None
if hasattr(self._shared_experts, "expert_gate") and self._shared_experts.expert_gate is not None:
gate_out, _ = self._shared_experts.expert_gate(hidden_states) # type: ignore
shared_out = F.sigmoid(gate_out) * shared_out
return shared_out
def _get_quant_type(self) -> QuantType:
quant_type = QuantType.NONE
method = getattr(self._quant_method, "quant_method", None)
if method is not None:
quant_type = getattr(method, "quant_type", QuantType.NONE)
return quant_type
def _register_routed_expert_parameter_aliases(self) -> None:
alias_names = []
for name, param in self.routed_experts.named_parameters(recurse=False):
alias_param = torch.nn.Parameter(param.data, requires_grad=param.requires_grad)
alias_param.__dict__.update(param.__dict__)
self.register_parameter(name, alias_param)
alias_names.append(name)
original_process_weights = self._quant_method.process_weights_after_loading
@wraps(original_process_weights)
def wrapped_process_weights(layer, *args, **kwargs):
for name in alias_names:
self._parameters.pop(name, None)
return original_process_weights(layer, *args, **kwargs)
self._quant_method.process_weights_after_loading = wrapped_process_weights # type: ignore[method-assign]
def _needs_routed_expert_parameter_aliases(self) -> bool:
vllm_config = get_current_vllm_config()
hf_config = getattr(vllm_config.model_config, "hf_config", None)
return getattr(hf_config, "model_type", None) == "gpt_oss"
@property
def is_internal_router(self) -> bool:
gate = self.gate
return gate is not None and hasattr(gate, "weight_fp32")
@property
def use_dp_chunking(self) -> bool:
"""Ascend uses its own forward_impl path, not the FlashInfer Cutlass
chunked path. Always return False to stay on forward_impl."""
return False
@property
def _fused_output_is_reduced(self) -> bool:
# For MC2/ALLTOALL/FUSED_MC2 comm types, finalize() already includes
# TP all-reduce for the routed output, and _forward_shared_experts
# handles it for the shared output. Signal this to the upstream
# MoERunner.forward() so _maybe_reduce_final_output does not apply a
# second TP all-reduce (which would double-count the contributions).
moe_comm_type = _EXTRA_CTX.moe_comm_type
return moe_comm_type in {
MoECommType.ALLTOALL,
MoECommType.MC2,
MoECommType.FUSED_MC2,
} or (moe_comm_type == MoECommType.ALLGATHER and _EXTRA_CTX.flash_comm_v1_enabled)
def _maybe_reduce_shared_expert_output(
self,
shared_output: torch.Tensor | None,
) -> torch.Tensor | None:
# _forward_shared_experts already handles shared expert TP all-reduce
# for MC2/ALLTOALL/FUSED_MC2. For AllGather the reduction is done
# via _maybe_reduce_final_output on the combined (shared + routed)
# output. Skip any additional reduction here.
return shared_output
def _maybe_reduce_final_output(
self,
states: torch.Tensor,
trunc_size: int,
) -> torch.Tensor:
states = torch.ops.vllm.maybe_all_reduce_tensor_model_parallel(states)
return states[..., :trunc_size]
def set_lora_context(self, lora_context):
self.routed_experts._ascend_moe_lora_context = lora_context
def no_shared_forward_impl( # type: ignore[override]
self, hidden_states: torch.Tensor, router_logits: torch.Tensor, return_with_event: bool = False
) -> torch.Tensor | FusedMoEResult:
forward_context = get_forward_context()
# When static kernels are enabled, the forward pass runs twice (compilation + capture),
# causing moe_layer_index to overflow. Wrap the index to prevent out-of-bounds errors.
if self.enable_npugraph_ex_static_kernel and forward_context.all_moe_layers:
moe_layer_index = forward_context.moe_layer_index % (len(forward_context.all_moe_layers))
forward_context.moe_layer_index = moe_layer_index
# Load balancing for token distribution among experts in dummy_run
# TODO: The community only considers load balancing when DP > 1.
# This approach may overlook some extreme scenarios.
enable_force_load_balance = _EXTRA_CTX.in_profile_run
forward_context = get_forward_context()
if self.multistream_overlap_gate:
fc3_context = get_flash_common3_context()
assert fc3_context is not None
assert AscendMoERunner.gate_stream is not None
AscendMoERunner.gate_stream.wait_stream(torch.npu.current_stream())
with npu_stream_switch(AscendMoERunner.gate_stream, enabled=self.multistream_overlap_gate):
# share_expert
assert fc3_context.shared_experts is not None
shared_out = fc3_context.shared_experts(hidden_states)
# NOTE: This is exactly the opposite of `maybe_all_reduce_tensor_model_parallel`
moe_comm_type = _EXTRA_CTX.moe_comm_type
if (
moe_comm_type in {MoECommType.ALLTOALL, MoECommType.MC2, MoECommType.FUSED_MC2}
and not shared_expert_dp_enabled()
):
shared_out = tensor_model_parallel_all_reduce(shared_out)
set_flash_common3_context(shared_out=shared_out)
input_ids = getattr(get_forward_context(), "input_ids", None)
topk_weights, topk_ids = select_experts(
hidden_states=hidden_states,
router_logits=router_logits,
top_k=self.top_k,
use_grouped_topk=self.use_grouped_topk,
renormalize=self.renormalize,
topk_group=self.topk_group,
num_expert_group=self.num_expert_group,
custom_routing_function=self.custom_routing_function,
scoring_func=self.scoring_func,
routed_scaling_factor=self._original_routed_scaling_factor,
e_score_correction_bias=self.e_score_correction_bias,
num_experts=self.moe_config.num_experts,
input_ids=input_ids,
tid2eid=self.tid2eid,
)
if isinstance(_EXTRA_CTX.moe_comm_method, AllGatherCommImpl):
topk_weights = torch.ops.vllm.maybe_all_gather_and_maybe_unpad(topk_weights, True, True)
topk_ids = torch.ops.vllm.maybe_all_gather_and_maybe_unpad(topk_ids, True, True)
set_flash_common3_context(topk_weights=topk_weights, topk_ids=topk_ids)
prepare_output = _EXTRA_CTX.moe_comm_method.prepare(
hidden_states=hidden_states,
router_logits=router_logits,
replace_allreduce=_EXTRA_CTX.flash_comm_v1_enabled,
enable_shared_expert_dp=self.enable_shared_expert_dp,
quant_type=self.quant_type,
)
hidden_states = prepare_output.hidden_states
router_logits = prepare_output.router_logits
mc2_mask = prepare_output.mc2_mask
padded_hidden_states_shape = prepare_output.padded_hidden_states_shape
pertoken_scale = prepare_output.pertoken_scale
# Make sure the default stream waits for the gate stream to finish.
if self.multistream_overlap_gate:
assert AscendMoERunner.gate_stream is not None
torch.npu.current_stream().wait_stream(AscendMoERunner.gate_stream)
# Matrix multiply.
# apply() expects a RoutedExperts-like layer for weight access
# (w13_weight, w2_weight, swiglu_limit, etc.). Pass routed_experts,
# not self; the routing params come through the other kwargs.
fused_experts_results: FusedExpertsResult = self._quant_method.apply(
layer=self.routed_experts,
x=hidden_states,
router_logits=router_logits,
pertoken_scale=pertoken_scale,
top_k=self.top_k,
renormalize=self.renormalize,
use_grouped_topk=self.use_grouped_topk,
num_experts=self.moe_config.num_experts,
expert_map=self._expert_map,
topk_group=self.topk_group,
num_expert_group=self.num_expert_group,
custom_routing_function=self.custom_routing_function,
scoring_func=self.scoring_func,
routed_scaling_factor=self._original_routed_scaling_factor,
e_score_correction_bias=self.e_score_correction_bias,
activation=self.activation,
apply_router_weight_on_input=self.apply_router_weight_on_input,
enable_force_load_balance=enable_force_load_balance,
log2phy=self.log2phy,
global_redundant_expert_num=self.global_redundant_expert_num,
mc2_mask=mc2_mask,
)
if self.dynamic_eplb and _EXTRA_CTX.eplb_heat_collection_status:
expert_tokens = fused_experts_results.expert_tokens
group_list_type = fused_experts_results.group_list_type
assert expert_tokens is not None and group_list_type is not None, (
"expert_tokens and group_list_type should not be None when dynamic_eplb is enabled."
)
local_load = (
expert_tokens
if group_list_type == 1
else torch.cat([expert_tokens[:1], expert_tokens[1:] - expert_tokens[:-1]])
)
if self.multi_stage:
cur_iter = torch.remainder(self.load_counter, self.num_iter)
self.moe_load.index_add_(
dim=0, index=cur_iter, source=local_load.to(torch.int32, non_blocking=True).view(1, -1)
)
self.load_counter.add_(1)
else:
self.moe_load.add_(local_load)
routed_out = _EXTRA_CTX.moe_comm_method.finalize(
hidden_states=fused_experts_results.routed_out,
reduce_results=isinstance(_EXTRA_CTX.moe_comm_method, AllGatherCommImpl),
padded_hidden_states_shape=padded_hidden_states_shape,
)
if return_with_event:
return FusedMoEResult(
routed_out=routed_out,
before_dispatch_evt=fused_experts_results.before_dispatch_evt,
before_gmm2_evt=fused_experts_results.before_gmm2_evt,
before_combine_evt=fused_experts_results.before_combine_evt,
swiglu_limit=fused_experts_results.swiglu_limit,
)
else:
# The vLLM FusedMoE forward_impl does not return events.
return routed_out
def _forward_shared_experts(self, hidden_states: torch.Tensor, fused_moe_evts: FusedMoEEvents):
if self._shared_experts is None:
return None
def maybe_wait_event(evt: torch.npu.Event | None):
if evt is not None:
torch.npu.current_stream().wait_event(evt)
with npu_stream_switch(shared_experts_calculation_stream(), enabled=self.multistream_overlap_shared_expert):
# Only used for int quantization
has_quantized_shared = hasattr(self._shared_experts.gate_up_proj, "weight_scale") and hasattr(
self._shared_experts.down_proj, "weight_scale"
)
if has_quantized_shared and self.quant_type in (QuantType.W8A8, QuantType.W4A8):
original_dtype = hidden_states.dtype
# Execute dynamic quant concurrently with MoE gate.
torch.npu.current_stream().wait_event(fused_moe_evts.before_routed_experts)
quantized_x, pertoken_scale = torch_npu.npu_dynamic_quant(hidden_states)
# Execute the gate projection and activation concurrently with the
# dispatch communication.
maybe_wait_event(fused_moe_evts.after_routed_experts)
hidden_states = torch_npu.npu_quant_matmul(
quantized_x,
self._shared_experts.gate_up_proj.weight,
self._shared_experts.gate_up_proj.weight_scale,
pertoken_scale=None,
bias=None,
output_dtype=torch.int32,
)
# Execute activation concurrently with gmm2.
maybe_wait_event(fused_moe_evts.before_gmm2)
quantized_x, swiglu_out_scale = torch.ops._C_ascend.npu_dequant_swiglu_quant(
x=hidden_states,
weight_scale=self._shared_experts.gate_up_proj.weight_scale_fp32,
activation_scale=pertoken_scale,
bias=None,
quant_scale=None,
quant_offset=None,
group_index=None,
activate_left=True,
quant_mode=1,
swiglu_mode=1,
clamp_limit=fused_moe_evts.swiglu_limit,
)
# Execute the down projection concurrently with the combine
# communication.
maybe_wait_event(fused_moe_evts.before_combine)
shared_out = torch_npu.npu_quant_matmul(
quantized_x,
self._shared_experts.down_proj.weight,
self._shared_experts.down_proj.weight_scale,
pertoken_scale=swiglu_out_scale,
bias=None,
output_dtype=original_dtype,
)
elif has_quantized_shared and self.quant_type == QuantType.W4A8MXFP:
original_dtype = hidden_states.dtype
# Execute dynamic quant concurrently with MoE gate.
torch.npu.current_stream().wait_event(fused_moe_evts.before_routed_experts)
quantized_x, pertoken_scale = torch_npu.npu_dynamic_mx_quant(
hidden_states, dst_type=torch.float8_e4m3fn
)
# Execute the gate projection and activation concurrently with the
# dispatch communication.
maybe_wait_event(fused_moe_evts.before_dispatch)
hidden_states = self._shared_experts.gate_up_proj((quantized_x, pertoken_scale))[0]
# Execute activation concurrently with gmm2.
maybe_wait_event(fused_moe_evts.before_gmm2)
quantized_x, swiglu_out_scale, _ = torch.ops._C_ascend.npu_swiglu_group_quant(
hidden_states,
topk_weight=None,
group_index=None,
dst_type=torch.float8_e4m3fn,
quant_mode=2,
clamp_value=fused_moe_evts.swiglu_limit,
)
# Execute the down projection concurrently with the combine
# communication.
maybe_wait_event(fused_moe_evts.before_combine)
shared_out = self._shared_experts.down_proj((quantized_x, swiglu_out_scale))[0]
else:
# Ensure the shared experts wait for hidden_states to be ready.
torch.npu.current_stream().wait_event(fused_moe_evts.before_routed_experts)
# Execute the gate projection and activation concurrently with the
# dispatch communication.
maybe_wait_event(fused_moe_evts.before_dispatch)
part1_out = self._shared_experts_part1(hidden_states)
# Execute the down projection concurrently with the combine
# communication.
maybe_wait_event(fused_moe_evts.before_combine)
shared_out = self._shared_experts_part2(hidden_states, part1_out)
# Make sure the default stream waits for the shared experts stream to
# finish.
if self.multistream_overlap_shared_expert:
torch.npu.current_stream().wait_stream(shared_experts_calculation_stream())
# NOTE: This is exactly the opposite of
# `maybe_all_reduce_tensor_model_parallel`
moe_comm_type = _EXTRA_CTX.moe_comm_type
if (
moe_comm_type in {MoECommType.ALLTOALL, MoECommType.MC2, MoECommType.FUSED_MC2}
and not shared_expert_dp_enabled()
):
shared_out = tensor_model_parallel_all_reduce(shared_out)
return shared_out
def shared_forward_impl( # type: ignore[override]
self, hidden_states: torch.Tensor, router_logits: torch.Tensor
):
if self.shared_multistream_overlap_gate:
set_flash_common3_context(shared_experts=self._shared_experts)
if self.is_internal_router:
gate = self.gate
assert gate is not None
# NOTE(Angazenn): To make this cast explicitly, the hbm usage might
# increase with extra hidden states. We also assume that all gate
# linear is unquantized so that we the weight is pre-casted in
# process_weights_after_loading of AscendUnquantizedLinearMethod.
hidden_states_fp32 = hidden_states.float()
before_routed_experts = torch.npu.current_stream().record_event()
router_logits = F.linear(hidden_states_fp32, gate.weight_fp32)
after_routed_experts = torch.npu.current_stream().record_event()
else:
before_routed_experts = torch.npu.current_stream().record_event()
after_routed_experts = None
fused_moe_results = self.no_shared_forward_impl(
hidden_states,
router_logits,
return_with_event=True,
)
routed_out = fused_moe_results.routed_out
if self._shared_experts is None:
return routed_out
if self.shared_multistream_overlap_gate:
fc3_context = get_flash_common3_context()
assert fc3_context is not None
shared_out = fc3_context.shared_out
else:
shared_out = self._forward_shared_experts(
hidden_states,
FusedMoEEvents(
after_routed_experts=after_routed_experts,
before_routed_experts=before_routed_experts,
before_dispatch=fused_moe_results.before_dispatch_evt,
before_gmm2=fused_moe_results.before_gmm2_evt,
before_combine=fused_moe_results.before_combine_evt,
swiglu_limit=fused_moe_results.swiglu_limit,
),
)
return shared_out, routed_out
def _forward_impl(
self,
hidden_states: torch.Tensor,
router_logits: torch.Tensor,
shared_experts_input: torch.Tensor | None,
input_ids: torch.Tensor | None = None,
) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]:
with self._sequence_parallel_context():
if self.shared_experts is None:
return self.no_shared_forward_impl(hidden_states, router_logits)
else:
return self.shared_forward_impl(hidden_states, router_logits)

View File

@@ -0,0 +1,722 @@
#
# 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.
#
# ruff: noqa: E501
"""Legacy vLLM 0.23.0 FusedMoE implementation.
Private module. Import AscendFusedMoE and AscendMoERunner through
vllm_ascend.ops.fused_moe.fused_moe only.
"""
from __future__ import annotations
from vllm_ascend.ops.fused_moe.fused_moe import (
_EXTRA_CTX,
AllGatherCommImpl,
AscendUnquantizedFusedMoEMethod,
F,
FusedExpertsResult,
FusedMoE,
FusedMoEEvents,
FusedMoEResult,
MoECommType,
MoERunner,
QuantType,
VllmEplbAdaptor,
get_ascend_config,
get_compressed_expert_map,
get_current_vllm_config,
get_dp_group,
get_ep_group,
get_flash_common3_context,
get_forward_context,
get_mc2_group,
get_tp_group,
init_eplb_config,
logger,
npu_stream_switch,
select_experts,
set_flash_common3_context,
setup_moe_comm_method,
shared_expert_dp_enabled,
shared_experts_calculation_stream,
tensor_model_parallel_all_reduce,
torch,
torch_npu,
wraps,
)
from vllm_ascend.utils import enable_sp
class AscendMoERunner(MoERunner):
@property
def use_dp_chunking(self) -> bool:
"""Ascend uses its own forward_impl path, not the FlashInfer Cutlass
chunked path. Always return False to stay on forward_impl."""
return False
@property
def _fused_output_is_reduced(self) -> bool:
# For MC2/ALLTOALL/FUSED_MC2 comm types, finalize() already includes
# TP all-reduce for the routed output, and _forward_shared_experts
# handles it for the shared output. Signal this to the upstream
# MoERunner.forward() so _maybe_reduce_final_output does not apply a
# second TP all-reduce (which would double-count the contributions).
moe_comm_type = _EXTRA_CTX.moe_comm_type
return moe_comm_type in {
MoECommType.ALLTOALL,
MoECommType.MC2,
MoECommType.FUSED_MC2,
} or (moe_comm_type == MoECommType.ALLGATHER and _EXTRA_CTX.flash_comm_v1_enabled)
def _maybe_reduce_shared_expert_output(
self,
shared_output: torch.Tensor | None,
) -> torch.Tensor | None:
# _forward_shared_experts already handles shared expert TP all-reduce
# for MC2/ALLTOALL/FUSED_MC2. For AllGather the reduction is done
# via _maybe_reduce_final_output on the combined (shared + routed)
# output. Skip any additional reduction here.
return shared_output
def _maybe_reduce_final_output(
self,
states: torch.Tensor,
trunc_size: int,
) -> torch.Tensor:
states = torch.ops.vllm.maybe_all_reduce_tensor_model_parallel(states)
return states[..., :trunc_size]
# TODO: Remove this after drop v0.19.1 support
def forward_impl(
self,
layer: torch.nn.Module,
hidden_states: torch.Tensor,
router_logits: torch.Tensor,
shared_input: torch.Tensor | None,
) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]:
"""
Override the default forward_impl to use Ascend-specific implementation.
This delegates to the layer's forward_impl method which contains the
Ascend-specific MoE computation logic.
"""
if self.shared_experts is None:
result = layer.forward_impl(hidden_states, router_logits)
# If the layer has shared experts, forward_impl returns a tuple (shared_out, routed_out)
# Otherwise, it returns just routed_out
# The torch op expects the same return type based on whether it's moe_forward or moe_forward_shared
else:
result = layer.shared_forward_impl(hidden_states, router_logits)
return result
def _forward_impl(
self,
layer: torch.nn.Module,
hidden_states: torch.Tensor,
router_logits: torch.Tensor,
shared_experts_input: torch.Tensor | None,
input_ids: torch.Tensor | None,
) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]:
with self._sequence_parallel_context():
return self.forward_impl(
layer,
hidden_states,
router_logits,
shared_experts_input,
)
class AscendFusedMoE(FusedMoE):
moe_counter = -1
gate_stream: torch.npu.Stream | None = None
def __init__(self, *args, **kwargs):
# Save original routed_scaling_factor before super().__init__ modifies it.
# When apply_routed_scale_to_output=True, vLLM sets self.routed_scaling_factor
# to 1.0 and expects the runner to apply scaling to output. But vllm-ascend
# uses its own forward path, so we need the original value.
_ = kwargs.pop("hash") if "hash" in kwargs else None
tid2eid = kwargs.pop("tid2eid") if "tid2eid" in kwargs else None
self._original_routed_scaling_factor = kwargs.get("routed_scaling_factor", 1.0)
super().__init__(*args, **kwargs)
self.use_overlapped = True
self._routed_input_transform = kwargs.get("routed_input_transform")
self._shared_experts = kwargs.get("shared_experts")
self.shared_expert_stream = None
has_shared_experts = self._shared_experts is not None
num_experts = kwargs["num_experts"]
intermediate_size = kwargs["intermediate_size"]
num_shared_experts = kwargs.get("n_shared_experts", 0)
AscendFusedMoE.moe_counter += 1
self.moe_instance_id = AscendFusedMoE.moe_counter
self._expert_map = None
self.log2phy = None
self.tid2eid = tid2eid
if self.quant_config is None:
self.quant_method = AscendUnquantizedFusedMoEMethod(self.moe_config, tid2eid=self.tid2eid)
else:
self.quant_method = self.quant_config.get_quant_method(self, self.layer_name, tid2eid=self.tid2eid)
assert self.quant_method is not None
# Keep base_quant_method in sync with the swapped-in Ascend method,
# otherwise FusedMoE.maybe_init_modular_kernel (called via the V2
# model runner's prepare_communication_buffer_for_model) would dispatch
# to the upstream UnquantizedFusedMoEMethod.maybe_make_prepare_finalize,
# which raises by design.
self.base_quant_method = self.quant_method
self.moe_config.tp_group = get_tp_group()
self.moe_config.dp_group = get_dp_group()
if self.moe_config.ep_size > 1:
self.moe_config.ep_group = get_ep_group()
self.moe_config.mc2_group = get_mc2_group()
self.moe_config.supports_eplb = self.quant_method.supports_eplb
ascend_config = get_ascend_config()
self.multistream_overlap_shared_expert = ascend_config.multistream_overlap_shared_expert and has_shared_experts
self.shared_multistream_overlap_gate = ascend_config.multistream_overlap_gate and has_shared_experts
if self.multistream_overlap_shared_expert:
logger.info_once("[fused_moe/layer] Multistream overlap shared expert is enabled.")
if enable_sp() and has_shared_experts:
logger.info_once(
"[fused_moe/layer] Sequence parallelism is enabled, shared experts are replicated for best performance."
)
# flashcommon3 gate stream
self.multistream_overlap_gate = ascend_config.multistream_overlap_gate
if self.multistream_overlap_gate and AscendFusedMoE.gate_stream is None:
AscendFusedMoE.gate_stream = torch.npu.Stream()
if self.multistream_overlap_gate:
logger.info_once("[fused_moe/layer] Multistream overlap gate is enabled.")
vllm_config = get_current_vllm_config()
if (
self.custom_routing_function is None
and self.e_score_correction_bias is not None
and not vllm_config.model_config.is_deepseek_mla
):
self.e_score_correction_bias.data = self.e_score_correction_bias.data.to(
dtype=vllm_config.model_config.dtype
)
self._gate = kwargs.get("gate")
# init moe
eplb_config = ascend_config.eplb_config
self.mix_placement = getattr(ascend_config, "mix_placement", False)
self.n_shared_experts = num_shared_experts
num_experts += num_shared_experts if self.mix_placement else 0
self.moe_config.num_experts = num_experts
self.global_expert_map, self._expert_map, self.log2phy, self.global_redundant_expert_num = init_eplb_config(
eplb_config,
self.moe_instance_id,
self.moe_config,
self.mix_placement,
num_shared_experts,
tp_size=self.vllm_config.parallel_config.tensor_parallel_size,
)
self.global_num_experts = num_experts + self.global_redundant_expert_num
self.dynamic_eplb = eplb_config.dynamic_eplb and (self.log2phy is not None)
self.local_num_experts = self.global_num_experts // self.ep_size
self.expert_map_manager._local_num_experts = self.local_num_experts
self.expert_map_manager._expert_map = self._expert_map
if self._expert_map is not None:
logger.info_once(
"[fused_moe/layer] Expert parallelism is enabled."
" ep_rank=%s/%s, local_num_experts=%s, global_num_experts=%s,"
" expert_map=%s",
self.ep_rank,
self.ep_size,
self.local_num_experts,
self.global_num_experts,
get_compressed_expert_map(self._expert_map),
)
if self.dynamic_eplb:
self.multi_stage = False
self.moe_load = torch.zeros(self.local_num_experts, dtype=torch.int64).npu()
if eplb_config.eplb_policy_type == 3:
self.multi_stage = True
self.load_counter = torch.tensor(0, dtype=torch.int32, device="npu")
self.num_iter = eplb_config.expert_heat_collection_interval
self.moe_load = torch.zeros((self.num_iter, self.local_num_experts), dtype=torch.int32, device="npu")
self.moe_config.num_experts = self.global_num_experts
self.moe_config.num_local_experts = self.local_num_experts
self.moe_config.global_redundant_expert_num = self.global_redundant_expert_num
self.swiglu_limit = getattr(self.vllm_config.model_config.hf_config, "swiglu_limit", 0)
moe_quant_params = {
"num_experts": self.local_num_experts,
"hidden_size": self.hidden_size,
"intermediate_size_per_partition": self.intermediate_size_per_partition,
"params_dtype": self.params_dtype,
"weight_loader": self.weight_loader,
}
# need full intermediate size pre-sharding for WNA16 act order
if self.quant_method.__class__.__name__ in ("GPTQMarlinMoEMethod", "CompressedTensorsWNA16MoEMethod"):
moe_quant_params["intermediate_size_full"] = intermediate_size
self.quant_method.create_weights(layer=self, **moe_quant_params)
self.enable_shared_expert_dp = ascend_config.enable_shared_expert_dp
self.enable_npugraph_ex_static_kernel = ascend_config.ascend_compilation_config.enable_static_kernel
setup_moe_comm_method(self.moe_config)
self.quant_type = self._get_quant_type()
self.runner = AscendMoERunner(
self.layer_name,
self.moe_config,
self.router,
self._routed_input_transform,
kwargs.pop("gate", None),
kwargs.pop("shared_experts", None),
self.quant_method,
self.vllm_config.parallel_config.enable_dbo,
)
if self.multistream_overlap_shared_expert:
# Wrap the quant_method's process_weights_after_loading to validate that
# splitting shared expert computation (gate_up projection + activation,
# then down projection) yields identical results to integrated
# computation after weight loading.
original_process_weights = self.quant_method.process_weights_after_loading
@wraps(original_process_weights)
def wrapped_process_weights(*args, **kwargs):
result = original_process_weights(*args, **kwargs)
self._validate_shared_expert_consistency()
return result
self.quant_method.process_weights_after_loading = wrapped_process_weights # type: ignore
# Register this MoE layer with EPLB for PP compatibility.
# PPMissingLayer (nn.Identity) never calls AscendFusedMoE.__init__,
# so only real MoE layers on this rank are registered.
VllmEplbAdaptor.register_layer(self)
def _validate_shared_expert_consistency(self):
"""Validate that split shared expert computation matches integrated
computation."""
test_input = (
torch.rand(10, self.hidden_size, device="npu", dtype=self.moe_config.in_dtype) * 2 - 1
) # Random input for testing, scoped to [-1, 1]
assert self._shared_experts is not None
integrated_out = self._shared_experts(test_input)
part1_out = self._shared_experts_part1(test_input)
split_out = self._shared_experts_part2(test_input, part1_out)
if not torch.allclose(integrated_out, split_out):
diff = (integrated_out - split_out).abs()
logger.error(
"[fused_moe/layer] Shared expert split computation validation failed."
" The split-path computation does not match the integrated-path result."
" max_abs_diff=%s, integrated_sum=%s, integrated_norm=%s,"
" split_sum=%s, split_norm=%s, hidden_size=%s, dtype=%s.",
diff.max().item(),
integrated_out.sum().item(),
integrated_out.norm().item(),
split_out.sum().item(),
split_out.norm().item(),
self.hidden_size,
self.moe_config.in_dtype,
)
raise ValueError("FusedMoE shared experts split computation does not match the integrated computation.")
logger.info_once(
"[fused_moe/layer] Shared expert split computation validation passed."
" Integrated and split-path results are consistent."
)
def _shared_experts_part1(self, hidden_states: torch.Tensor):
shared_gate_up, _ = self._shared_experts.gate_up_proj(hidden_states) # type: ignore
return shared_gate_up
def _shared_experts_part2(self, hidden_states: torch.Tensor, shared_gate_up: torch.Tensor):
shared_act = self._shared_experts.act_fn(shared_gate_up) # type: ignore
shared_out, _ = self._shared_experts.down_proj(shared_act) # type: ignore
# Qwen3-Next specific gating mechanism
assert self._shared_experts is not None
if hasattr(self._shared_experts, "expert_gate") and self._shared_experts.expert_gate is not None:
gate_out, _ = self._shared_experts.expert_gate(hidden_states) # type: ignore
shared_out = F.sigmoid(gate_out) * shared_out
return shared_out
def _get_quant_type(self) -> QuantType:
quant_type = QuantType.NONE
method = getattr(self.quant_method, "quant_method", None)
if method is not None:
quant_type = getattr(method, "quant_type", QuantType.NONE)
return quant_type
def set_lora_context(self, lora_context):
self._ascend_moe_lora_context = lora_context
def update_expert_map(self, new_expert_map):
self._expert_map = new_expert_map
def get_log2phy_map(self):
return self.log2phy
def clear_moe_load(self):
if self.moe_load is not None:
self.moe_load.zero_()
if self.multi_stage:
self.load_counter.zero_()
def maybe_all_reduce_tensor_model_parallel(self, final_hidden_states: torch.Tensor):
"""NOTE(Yizhou): This is to override the parent class method. In `mc2commimpl`,
and `alltoallcommimpl`, we do not need to all-reduce the final outputs since
the outputs are already aggregated across tensor parallel ranks in the
`finalize` function. In `allgathercommimpl`, we still need to all-reduce the
outputs since each rank only has partial outputs.
"""
return torch.ops.vllm.maybe_all_reduce_tensor_model_parallel(final_hidden_states)
@property
def gate(self) -> torch.nn.Module | None:
return self._gate if self.use_overlapped else None
@property
def is_internal_router(self) -> bool:
gate = self.gate
return gate is not None and hasattr(gate, "weight_fp32")
@property
def use_dp_chunking(self) -> bool:
"""This func routes to the chunked forward path using the FlashInfer Cutlass kernel
only when data parallelism (DP) is enabled. Thus just returning False in vllm-ascend
"""
return False
def forward(
self,
hidden_states: torch.Tensor,
router_logits: torch.Tensor,
) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]:
self.ensure_moe_quant_config_init()
return self.runner.forward(
hidden_states,
router_logits,
)
def forward_impl( # type: ignore[override]
self, hidden_states: torch.Tensor, router_logits: torch.Tensor, return_with_event: bool = False
) -> torch.Tensor | FusedMoEResult:
assert self.quant_method is not None
forward_context = get_forward_context()
# When static kernels are enabled, the forward pass runs twice (compilation + capture),
# causing moe_layer_index to overflow. Wrap the index to prevent out-of-bounds errors.
if self.enable_npugraph_ex_static_kernel and forward_context.all_moe_layers:
moe_layer_index = forward_context.moe_layer_index % (len(forward_context.all_moe_layers))
forward_context.moe_layer_index = moe_layer_index
# Load balancing for token distribution among experts in dummy_run
# TODO: The community only considers load balancing when DP > 1.
# This approach may overlook some extreme scenarios.
enable_force_load_balance = _EXTRA_CTX.in_profile_run
forward_context = get_forward_context()
if self.multistream_overlap_gate:
assert AscendFusedMoE.gate_stream is not None
fc3_context = get_flash_common3_context()
assert fc3_context is not None
AscendFusedMoE.gate_stream.wait_stream(torch.npu.current_stream())
with npu_stream_switch(AscendFusedMoE.gate_stream, enabled=self.multistream_overlap_gate):
# share_expert
assert fc3_context.shared_experts is not None
shared_out = fc3_context.shared_experts(hidden_states)
# NOTE: This is exactly the opposite of `maybe_all_reduce_tensor_model_parallel`
moe_comm_type = _EXTRA_CTX.moe_comm_type
if (
moe_comm_type in {MoECommType.ALLTOALL, MoECommType.MC2, MoECommType.FUSED_MC2}
and not shared_expert_dp_enabled()
):
shared_out = tensor_model_parallel_all_reduce(shared_out)
set_flash_common3_context(shared_out=shared_out)
input_ids = getattr(get_forward_context(), "input_ids", None)
topk_weights, topk_ids = select_experts(
hidden_states=hidden_states,
router_logits=router_logits,
top_k=self.top_k,
use_grouped_topk=self.use_grouped_topk,
renormalize=self.renormalize,
topk_group=self.topk_group,
num_expert_group=self.num_expert_group,
custom_routing_function=self.custom_routing_function,
scoring_func=self.scoring_func,
routed_scaling_factor=self._original_routed_scaling_factor,
e_score_correction_bias=self.e_score_correction_bias,
num_experts=self.moe_config.num_experts,
input_ids=input_ids,
tid2eid=self.tid2eid,
)
if isinstance(_EXTRA_CTX.moe_comm_method, AllGatherCommImpl):
topk_weights = torch.ops.vllm.maybe_all_gather_and_maybe_unpad(topk_weights, True, True)
topk_ids = torch.ops.vllm.maybe_all_gather_and_maybe_unpad(topk_ids, True, True)
set_flash_common3_context(topk_weights=topk_weights, topk_ids=topk_ids)
prepare_output = _EXTRA_CTX.moe_comm_method.prepare(
hidden_states=hidden_states,
router_logits=router_logits,
replace_allreduce=_EXTRA_CTX.flash_comm_v1_enabled,
enable_shared_expert_dp=self.enable_shared_expert_dp,
quant_type=self.quant_type,
)
hidden_states = prepare_output.hidden_states
router_logits = prepare_output.router_logits
mc2_mask = prepare_output.mc2_mask
padded_hidden_states_shape = prepare_output.padded_hidden_states_shape
pertoken_scale = prepare_output.pertoken_scale
# Make sure the default stream waits for the gate stream to finish.
if self.multistream_overlap_gate:
torch.npu.current_stream().wait_stream(AscendFusedMoE.gate_stream)
# Matrix multiply.
fused_experts_results: FusedExpertsResult = self.quant_method.apply(
layer=self,
x=hidden_states,
router_logits=router_logits,
pertoken_scale=pertoken_scale,
top_k=self.top_k,
renormalize=self.renormalize,
use_grouped_topk=self.use_grouped_topk,
num_experts=self.moe_config.num_experts,
expert_map=self._expert_map,
topk_group=self.topk_group,
num_expert_group=self.num_expert_group,
custom_routing_function=self.custom_routing_function,
scoring_func=self.scoring_func,
routed_scaling_factor=self._original_routed_scaling_factor,
e_score_correction_bias=self.e_score_correction_bias,
activation=self.activation,
apply_router_weight_on_input=self.apply_router_weight_on_input,
enable_force_load_balance=enable_force_load_balance,
log2phy=self.log2phy,
global_redundant_expert_num=self.global_redundant_expert_num,
mc2_mask=mc2_mask,
)
if self.dynamic_eplb and _EXTRA_CTX.eplb_heat_collection_status:
expert_tokens = fused_experts_results.expert_tokens
group_list_type = fused_experts_results.group_list_type
assert expert_tokens is not None and group_list_type is not None, (
"expert_tokens and group_list_type should not be None when dynamic_eplb is enabled."
)
local_load = (
expert_tokens
if group_list_type == 1
else torch.cat([expert_tokens[:1], expert_tokens[1:] - expert_tokens[:-1]])
)
if self.multi_stage:
cur_iter = torch.remainder(self.load_counter, self.num_iter)
self.moe_load.index_add_(
dim=0, index=cur_iter, source=local_load.to(torch.int32, non_blocking=True).view(1, -1)
)
self.load_counter.add_(1)
else:
self.moe_load.add_(local_load)
routed_out = _EXTRA_CTX.moe_comm_method.finalize(
hidden_states=fused_experts_results.routed_out,
reduce_results=isinstance(_EXTRA_CTX.moe_comm_method, AllGatherCommImpl),
padded_hidden_states_shape=padded_hidden_states_shape,
)
if return_with_event:
return FusedMoEResult(
routed_out=routed_out,
before_dispatch_evt=fused_experts_results.before_dispatch_evt,
before_gmm2_evt=fused_experts_results.before_gmm2_evt,
before_combine_evt=fused_experts_results.before_combine_evt,
swiglu_limit=fused_experts_results.swiglu_limit,
)
else:
# The vLLM FusedMoE forward_impl does not return events.
return routed_out
def _forward_shared_experts(self, hidden_states: torch.Tensor, fused_moe_evts: FusedMoEEvents):
if self._shared_experts is None:
return None
def maybe_wait_event(evt: torch.npu.Event | None):
if evt is not None:
torch.npu.current_stream().wait_event(evt)
with npu_stream_switch(shared_experts_calculation_stream(), enabled=self.multistream_overlap_shared_expert):
# Only used for int quantization
has_quantized_shared = hasattr(self._shared_experts.gate_up_proj, "weight_scale") and hasattr(
self._shared_experts.down_proj, "weight_scale"
)
if has_quantized_shared and self.quant_type in (QuantType.W8A8, QuantType.W4A8):
original_dtype = hidden_states.dtype
# Execute dynamic quant concurrently with MoE gate.
torch.npu.current_stream().wait_event(fused_moe_evts.before_routed_experts)
quantized_x, pertoken_scale = torch_npu.npu_dynamic_quant(hidden_states)
# Execute the gate projection and activation concurrently with the
# dispatch communication.
maybe_wait_event(fused_moe_evts.after_routed_experts)
hidden_states = torch_npu.npu_quant_matmul(
quantized_x,
self._shared_experts.gate_up_proj.weight,
self._shared_experts.gate_up_proj.weight_scale,
pertoken_scale=None,
bias=None,
output_dtype=torch.int32,
)
# Execute activation concurrently with gmm2.
maybe_wait_event(fused_moe_evts.before_gmm2)
clamp_limit = fused_moe_evts.swiglu_limit or 0.0
group_index = None
if clamp_limit <= 0.0:
group_index = torch.empty((1,), dtype=torch.int64, device=hidden_states.device)
group_index.fill_(hidden_states.shape[0])
quantized_x, swiglu_out_scale = torch.ops._C_ascend.npu_dequant_swiglu_quant(
x=hidden_states,
weight_scale=self._shared_experts.gate_up_proj.weight_scale_fp32,
activation_scale=pertoken_scale,
bias=None,
quant_scale=None,
quant_offset=None,
group_index=group_index,
activate_left=True,
quant_mode=1,
swiglu_mode=1,
clamp_limit=clamp_limit,
)
# Execute the down projection concurrently with the combine
# communication.
maybe_wait_event(fused_moe_evts.before_combine)
shared_out = torch_npu.npu_quant_matmul(
quantized_x,
self._shared_experts.down_proj.weight,
self._shared_experts.down_proj.weight_scale,
pertoken_scale=swiglu_out_scale,
bias=None,
output_dtype=original_dtype,
)
elif has_quantized_shared and self.quant_type == QuantType.W4A8MXFP:
original_dtype = hidden_states.dtype
# Execute dynamic quant concurrently with MoE gate.
torch.npu.current_stream().wait_event(fused_moe_evts.before_routed_experts)
quantized_x, pertoken_scale = torch_npu.npu_dynamic_mx_quant(
hidden_states, dst_type=torch.float8_e4m3fn
)
# Execute the gate projection and activation concurrently with the
# dispatch communication.
maybe_wait_event(fused_moe_evts.before_dispatch)
hidden_states = self._shared_experts.gate_up_proj((quantized_x, pertoken_scale))[0]
# Execute activation concurrently with gmm2.
maybe_wait_event(fused_moe_evts.before_gmm2)
quantized_x, swiglu_out_scale, _ = torch.ops._C_ascend.npu_swiglu_group_quant(
hidden_states,
topk_weight=None,
group_index=None,
dst_type=torch.float8_e4m3fn,
quant_mode=2,
clamp_value=fused_moe_evts.swiglu_limit,
)
# Execute the down projection concurrently with the combine
# communication.
maybe_wait_event(fused_moe_evts.before_combine)
shared_out = self._shared_experts.down_proj((quantized_x, swiglu_out_scale))[0]
else:
# Ensure the shared experts wait for hidden_states to be ready.
torch.npu.current_stream().wait_event(fused_moe_evts.before_routed_experts)
# Execute the gate projection and activation concurrently with the
# dispatch communication.
maybe_wait_event(fused_moe_evts.before_dispatch)
part1_out = self._shared_experts_part1(hidden_states)
# Execute the down projection concurrently with the combine
# communication.
maybe_wait_event(fused_moe_evts.before_combine)
shared_out = self._shared_experts_part2(hidden_states, part1_out)
# Make sure the default stream waits for the shared experts stream to
# finish.
if self.multistream_overlap_shared_expert:
torch.npu.current_stream().wait_stream(shared_experts_calculation_stream())
# NOTE: This is exactly the opposite of
# `maybe_all_reduce_tensor_model_parallel`
moe_comm_type = _EXTRA_CTX.moe_comm_type
if (
moe_comm_type in {MoECommType.ALLTOALL, MoECommType.MC2, MoECommType.FUSED_MC2}
and not shared_expert_dp_enabled()
):
shared_out = tensor_model_parallel_all_reduce(shared_out)
return shared_out
def shared_forward_impl( # type: ignore[override]
self, hidden_states: torch.Tensor, router_logits: torch.Tensor
):
if self.shared_multistream_overlap_gate:
set_flash_common3_context(shared_experts=self._shared_experts)
if self.is_internal_router:
gate = self.gate
assert gate is not None
# NOTE(Angazenn): To make this cast explicitly, the hbm usage might
# increase with extra hidden states. We also assume that all gate
# linear is unquantized so that we the weight is pre-casted in
# process_weights_after_loading of AscendUnquantizedLinearMethod.
hidden_states_fp32 = hidden_states.float()
before_routed_experts = torch.npu.current_stream().record_event()
router_logits = F.linear(hidden_states_fp32, gate.weight_fp32)
after_routed_experts = torch.npu.current_stream().record_event()
else:
before_routed_experts = torch.npu.current_stream().record_event()
after_routed_experts = None
fused_moe_results = self.forward_impl(
hidden_states=hidden_states,
router_logits=router_logits,
return_with_event=True,
)
routed_out = fused_moe_results.routed_out
if self._shared_experts is None:
return routed_out
if self.shared_multistream_overlap_gate:
fc3_context = get_flash_common3_context()
assert fc3_context is not None
shared_out = fc3_context.shared_out
else:
shared_out = self._forward_shared_experts(
hidden_states,
FusedMoEEvents(
after_routed_experts=after_routed_experts,
before_routed_experts=before_routed_experts,
before_dispatch=fused_moe_results.before_dispatch_evt,
before_gmm2=fused_moe_results.before_gmm2_evt,
before_combine=fused_moe_results.before_combine_evt,
swiglu_limit=fused_moe_results.swiglu_limit,
),
)
return shared_out, routed_out
__all__ = ["AscendFusedMoE", "AscendMoERunner"]

View File

@@ -0,0 +1,63 @@
#
# 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 __future__ import annotations
import torch
from vllm.model_executor.layers.fused_moe.router.gate_linear import GateLinear
from vllm.model_executor.layers.linear import ReplicatedLinear
class AscendGateLinear(GateLinear):
"""Ascend replacement for vLLM GateLinear.
Router logits are sensitive to numerical precision because they directly
affect expert selection in MoE models. On NPU, computing the router gate in
lower precision may lead to accuracy issues in some agent workloads.
Therefore, this layer forces the gate input and weights to fp32 for the
router linear computation, and keeps the router logits in fp32.
"""
def __init__(
self,
input_size: int,
output_size: int,
bias: bool = False,
out_dtype: torch.dtype | None = None,
params_dtype: torch.dtype | None = None,
force_fp32_compute: bool = False,
prefix: str = "",
):
super().__init__(
input_size=input_size,
output_size=output_size,
bias=bias,
params_dtype=torch.float32,
out_dtype=out_dtype,
force_fp32_compute=True,
prefix=prefix,
)
def forward(self, x: torch.Tensor):
# TODO: Remove this workaround after upgrading to a vLLM version that
# no longer forces router logits to bf16 via
# self.gate.set_out_dtype(torch.bfloat16).
if x.dtype != torch.float32:
x = x.to(torch.float32)
output, output_bias = ReplicatedLinear.forward(self, x)
return output, output_bias

View File

@@ -0,0 +1,349 @@
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
# Copyright 2023 The vLLM team.
#
# 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.
# This file is a part of the vllm-ascend project.
from __future__ import annotations
from abc import ABC, abstractmethod
from dataclasses import dataclass
import torch
from vllm.model_executor.layers.fused_moe import FusedMoEConfig
from vllm_ascend.ascend_config import get_ascend_config
from vllm_ascend.ascend_forward_context import _EXTRA_CTX, MoECommType
from vllm_ascend.ops.fused_moe.moe_mlp import unified_apply_mlp
from vllm_ascend.ops.fused_moe.moe_runtime_args import (
MoEFusedExpertsInput,
MoEMlpComputeInput,
MoEPrepareOutput,
build_mlp_compute_input,
build_token_dispatch_input,
)
from vllm_ascend.ops.fused_moe.prepare_finalize import (
PrepareAndFinalize,
PrepareAndFinalizeWithAll2All,
PrepareAndFinalizeWithAllGather,
PrepareAndFinalizeWithMC2,
)
from vllm_ascend.ops.fused_moe.token_dispatcher import (
MoETokenDispatcher,
TokenDispatcherWithAll2AllV,
TokenDispatcherWithAllGather,
TokenDispatcherWithMC2,
)
from vllm_ascend.quantization.quant_type import QuantType
_MoECommMethods: dict[MoECommType | None, MoECommMethod] = {}
def get_moe_comm_method(moe_comm_type: MoECommType | None) -> MoECommMethod | None:
return _MoECommMethods.get(moe_comm_type)
def setup_moe_comm_method(moe_config):
if moe_config.ep_size > 1:
_MoECommMethods[MoECommType.ALLTOALL] = AlltoAllCommImpl(moe_config)
_MoECommMethods[MoECommType.ALLGATHER] = AllGatherCommImpl(moe_config)
_MoECommMethods[MoECommType.MC2] = MC2CommImpl(moe_config)
_MoECommMethods[MoECommType.FUSED_MC2] = FusedMC2CommImpl(moe_config)
else:
_MoECommMethods[MoECommType.ALLGATHER] = AllGatherCommImpl(moe_config)
def set_gmmswigluquant_method():
from vllm_ascend.ascend_config import get_ascend_config
ascend_config = get_ascend_config()
return ascend_config.ascend_fusion_config.fusion_ops_gmmswigluquant
@dataclass
class FusedExpertsResult:
routed_out: torch.Tensor
# This field is for shared experts and should be set by the MoE
# communication method that supports shared experts in parallel with routed
# experts.
before_dispatch_evt: torch.npu.Event | None = None
before_gmm2_evt: torch.npu.Event | None = None
before_combine_evt: torch.npu.Event | None = None
# For dynamic_eplb
group_list_type: int = 1
expert_tokens: torch.Tensor | None = None
swiglu_limit: float = 0.0
class MoECommMethod(ABC):
"""Base class for MoE communication methods."""
def __init__(self, moe_config: FusedMoEConfig):
self.moe_config = moe_config
self.token_dispatcher = self._get_token_dispatcher()
self.prepare_finalize = self._get_prepare_finalize()
self.use_fusion_ops = set_gmmswigluquant_method()
def prepare(
self,
hidden_states: torch.Tensor,
router_logits: torch.Tensor,
enable_shared_expert_dp: bool = False,
replace_allreduce: bool = False,
quant_type: QuantType = QuantType.NONE,
) -> MoEPrepareOutput:
return self.prepare_finalize.prepare(
hidden_states,
router_logits,
enable_shared_expert_dp,
replace_allreduce,
quant_type,
)
def finalize(
self,
hidden_states: torch.Tensor,
reduce_results: bool,
padded_hidden_states_shape: torch.Size | None = None,
) -> torch.Tensor:
hidden_states = self.prepare_finalize.finalize(hidden_states, reduce_results, padded_hidden_states_shape)
return hidden_states
def fused_experts(
self,
fused_experts_input: MoEFusedExpertsInput,
):
# Check constraints
assert fused_experts_input.hidden_states.dtype in [
torch.float32,
torch.float16,
torch.bfloat16,
torch.int8,
torch.float8_e4m3fn,
torch.uint8,
], f"Unsupported hidden_states dtype: {fused_experts_input.hidden_states.dtype}"
moe_comm_method = _EXTRA_CTX.moe_comm_method
assert moe_comm_method is not None, "Missing communication context"
before_dispatch_evt = torch.npu.current_stream().record_event()
routed_topk_ids = fused_experts_input.topk_ids
if fused_experts_input.routing.log2phy is not None:
routed_topk_ids = fused_experts_input.routing.log2phy[routed_topk_ids]
token_dispatch_input = build_token_dispatch_input(
fused_experts_input=fused_experts_input,
topk_ids=routed_topk_ids,
)
token_dispatch_output = self.token_dispatcher.token_dispatch(token_dispatch_input=token_dispatch_input)
mlp_compute_input = build_mlp_compute_input(
fused_experts_input=fused_experts_input,
token_dispatch_output=token_dispatch_output,
use_fusion_ops=self.use_fusion_ops,
)
mlp_output, before_gmm2_evt = self._apply_mlp(mlp_compute_input)
before_combine_evt = torch.npu.current_stream().record_event()
routed_out = self.token_dispatcher.token_combine(
hidden_states=mlp_output,
combine_metadata=token_dispatch_output.combine_metadata,
)
return FusedExpertsResult(
routed_out=routed_out,
before_dispatch_evt=before_dispatch_evt,
before_gmm2_evt=before_gmm2_evt,
before_combine_evt=before_combine_evt,
group_list_type=token_dispatch_output.group_list_type,
expert_tokens=token_dispatch_output.group_list,
swiglu_limit=fused_experts_input.swiglu_limit,
)
def _apply_mlp(self, mlp_compute_input: MoEMlpComputeInput) -> torch.Tensor:
return unified_apply_mlp(mlp_compute_input=mlp_compute_input)
@abstractmethod
def _get_token_dispatcher(self) -> MoETokenDispatcher:
raise NotImplementedError("_get_token_dispatcher function not implemented.")
@abstractmethod
def _get_prepare_finalize(self) -> PrepareAndFinalize:
raise NotImplementedError("_get_prepare_finalize function not implemented.")
class AllGatherCommImpl(MoECommMethod):
"""This implementation is the same as NativeAllGatherCommImpl,
but uses NPU-specific ops for better performance.
This implementation should be compatible with all scenarios, and
thus it is the default implementation for MoE communication methods.
It uses `torch_npu.npu_moe_init_routing_v2` for pre-processing
and `torch_npu.npu_moe_token_unpermute` for post-processing
to handle the token-to-expert mapping and communication efficiently.
NOTE(Yizhou): TBH, it is really weird that we were supposed to use
`torch_npu.npu_moe_init_routing_v2` and `torch_npu.npu_moe_finalize_routing`
or `torch_npu.npu_moe_token_permute` and `torch_npu.npu_moe_token_unpermute`
for pre-processing and post-processing, respectively.
But `npu_moe_finalize_routing` will lead to accuracy issues so we have to
use `torch_npu.npu_moe_token_unpermute` instead.
This is a workaround and should be removed after the issue is fixed.
"""
def _get_token_dispatcher(self):
return TokenDispatcherWithAllGather(
top_k=self.moe_config.experts_per_token,
num_experts=self.moe_config.num_experts,
num_local_experts=self.moe_config.num_local_experts,
)
def _get_prepare_finalize(self):
return PrepareAndFinalizeWithAllGather(self.moe_config)
class MC2CommImpl(MoECommMethod):
"""This implementation is for the scenarios listed below:
1. `enable_expert_parallel=True`.
2. `npu_moe_distribute_dispatch` and `npu_moe_distribute_combine` are available.
3. `enable_expert_parallel=False` is not supported.
This implementation uses the MC2 communication method, which is optimized for
Communication and Computation parallelism on Ascend devices.
"""
def pad_and_split_input_ids(self, input_ids):
return self.prepare_finalize.pad_and_split_input_ids(input_ids) # type: ignore[attr-defined]
def _get_token_dispatcher(self):
return TokenDispatcherWithMC2()
def _get_prepare_finalize(self):
return PrepareAndFinalizeWithMC2(self.moe_config)
class AlltoAllCommImpl(MoECommMethod):
"""This implementation is for the scenarios listed below:
1. `enable_expert_parallel=True`.
2. `npu_grouped_matmul` is available.
This implementation uses all-to-all communication to exchange tokens
between data parallel ranks before and after the MLP computation. It should
have better performance than AllGatherCommImpl when DP size > 1.
"""
def pad_and_split_input_ids(self, input_ids):
return self.prepare_finalize.pad_and_split_input_ids(input_ids) # type: ignore[attr-defined]
def _get_token_dispatcher(self):
return TokenDispatcherWithAll2AllV(
top_k=self.moe_config.experts_per_token,
num_experts=self.moe_config.num_experts,
num_local_experts=self.moe_config.num_local_experts,
)
def _get_prepare_finalize(self):
return PrepareAndFinalizeWithAll2All(self.moe_config)
class FusedMC2CommImpl(MoECommMethod):
"""This implementation is for the scenarios listed below:
1. `enable_expert_parallel=True`.
2. `npu_moe_distribute_dispatch` and `npu_moe_distribute_combine` are available.
3. `enable_expert_parallel=False` is not supported.
This implementation uses the MC2 communication method, which is optimized for
Communication and Computation parallelism on Ascend devices.
"""
def __init__(self, moe_config):
super().__init__(moe_config)
if get_ascend_config().enable_fused_mc2 == 1:
self.expert_token_nums = torch.zeros([self.moe_config.num_local_experts], dtype=torch.int32, device="npu")
else:
self.expert_token_nums = None
def pad_and_split_input_ids(self, input_ids):
return self.prepare_finalize.pad_and_split_input_ids(input_ids) # type: ignore[attr-defined]
def _get_token_dispatcher(self):
return TokenDispatcherWithMC2()
def _get_prepare_finalize(self):
return PrepareAndFinalizeWithMC2(self.moe_config)
def fused_experts(
self,
fused_experts_input: MoEFusedExpertsInput,
):
assert not (fused_experts_input.weights.w1_scale is None or fused_experts_input.weights.w2_scale is None), (
"w1_scale and w2_scale cannot be None for FusedMC2CommImpl."
)
assert isinstance(self.token_dispatcher, TokenDispatcherWithMC2), (
"token_dispatcher must be an instance of TokenDispatcherWithMC2."
)
# Apply log2phy if needed
topk_ids = fused_experts_input.topk_ids
if fused_experts_input.routing.log2phy is not None:
topk_ids = fused_experts_input.routing.log2phy[topk_ids]
expert_tokens = None
if get_ascend_config().enable_fused_mc2 == 1:
assert not (
fused_experts_input.weights.w1_scale_bias is None or fused_experts_input.weights.w2_scale_bias is None
), "w1_scale_bias and w2_scale_bias cannot be None when enable_fused_mc2=1."
out = torch.empty_like(fused_experts_input.hidden_states)
torch.ops._C_ascend.dispatch_ffn_combine( # type: ignore
x=fused_experts_input.hidden_states,
weight1=fused_experts_input.weights.w1,
weight2=fused_experts_input.weights.w2,
expert_idx=topk_ids,
scale1=fused_experts_input.weights.w1_scale,
scale2=fused_experts_input.weights.w2_scale,
bias1=fused_experts_input.weights.w1_scale_bias,
bias2=fused_experts_input.weights.w2_scale_bias,
probs=fused_experts_input.topk_weights.to(torch.float32),
group=self.token_dispatcher.moe_all_to_all_group_name,
max_output_size=get_ascend_config().mega_moe_max_tokens,
swiglu_limit=fused_experts_input.swiglu_limit,
x_active_mask=fused_experts_input.routing.mc2_mask,
out=out,
expert_token_nums=self.expert_token_nums,
)
expert_tokens = self.expert_token_nums
elif get_ascend_config().enable_fused_mc2 == 2:
assert fused_experts_input.routing.expert_map is not None, "expert_map cannot be None."
out, expert_tokens = torch.ops._C_ascend.dispatch_gmm_combine_decode( # type: ignore
x=fused_experts_input.hidden_states,
expert_ids=topk_ids,
gmm1_permuted_weight=fused_experts_input.weights.w1,
gmm1_permuted_weight_scale=fused_experts_input.weights.w1_scale,
gmm2_weight=fused_experts_input.weights.w2,
gmm2_weight_scale=fused_experts_input.weights.w2_scale,
expert_smooth_scales=None,
expert_scales=fused_experts_input.topk_weights.to(torch.float32),
group_ep=self.token_dispatcher.moe_all_to_all_group_name,
ep_rank_size=self.token_dispatcher.ep_world_size,
ep_rank_id=self.token_dispatcher.ep_rank_id,
moe_expert_num=self.moe_config.num_experts,
global_bs=self.token_dispatcher.global_bs,
)
else:
raise ValueError(f"Wrong value of {get_ascend_config().enable_fused_mc2=}")
return FusedExpertsResult(
routed_out=out, expert_tokens=expert_tokens, swiglu_limit=fused_experts_input.swiglu_limit
)

View File

@@ -0,0 +1,552 @@
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
# Copyright 2023 The vLLM team.
#
# 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.
# This file is a part of the vllm-ascend project.
import torch
import torch_npu
from torch.nn.functional import pad
from vllm.model_executor.layers.fused_moe.activation import MoEActivation
from vllm.triton_utils import HAS_TRITON
from vllm_ascend.ascend_forward_context import _EXTRA_CTX, MoECommType
from vllm_ascend.device.device_op import DeviceOperator
from vllm_ascend.device.mxfp_compat import (
ensure_mxfp8_moe_available,
)
from vllm_ascend.ops.activation import AscendSwigluOAIAndMul, AscendSwigluStepAndMul
from vllm_ascend.ops.fused_moe.moe_runtime_args import MoEMlpComputeInput
from vllm_ascend.quantization.quant_type import QuantType
from vllm_ascend.utils import (
dispose_tensor,
enable_custom_op,
get_ascend_device_type,
get_weight_prefetch_method,
)
ASCEND_DEVICE_TYPE = get_ascend_device_type()
def _custom_gmm_swiglu_enabled(fusion, dynamic_eplb):
return fusion and dynamic_eplb and enable_custom_op()
def cumsum_group_list(
group_list: torch.Tensor, src_list_type: int, dst_list_type: int, active_num: int = 0, expert_num: int = 0
) -> torch.Tensor:
if src_list_type not in [0, 1, 2]:
raise ValueError(f"group_list_type should be in [0, 1, 2], but received {src_list_type}")
if src_list_type == dst_list_type:
return group_list
if src_list_type == 1 and dst_list_type == 0:
return group_list.cumsum(dim=0)
if src_list_type == 0 and dst_list_type == 1:
group_diff = torch.diff(group_list)
new_group = torch.cat([group_list[0].unsqueeze(0), group_diff], dim=0)
return new_group
if src_list_type == 2 and dst_list_type == 0:
experts = pad(group_list[:, 0], (1, 0))
tokens = pad(group_list[:, 1].cumsum(dim=0), (1, 0))
cumsum_group_list = torch.full(
size=(expert_num,), fill_value=active_num, dtype=group_list.dtype, device=group_list.device
)
for i, (start, end) in enumerate(zip(experts[:-1], experts[1:])):
if end > start:
cumsum_group_list[start:end] = tokens[i]
return cumsum_group_list
raise NotImplementedError(
f"Conversion from src_list_type={src_list_type} to dst_list_type={dst_list_type} is not implemented yet. "
"This feature is under development."
)
def _require_single_tensor_for_swiglu_quant(
tensor_or_list: list[torch.Tensor] | torch.Tensor, *, name: str
) -> torch.Tensor:
if isinstance(tensor_or_list, list):
if len(tensor_or_list) != 1:
raise ValueError(f"{name} must be a tensor or a single-element list, but got {len(tensor_or_list)}.")
return tensor_or_list[0]
return tensor_or_list
def quant_apply_mlp(
hidden_states: torch.Tensor,
w1: list[torch.Tensor] | torch.Tensor,
w1_scale: list[torch.Tensor] | torch.Tensor,
w2: list[torch.Tensor] | torch.Tensor,
w2_scale: list[torch.Tensor] | torch.Tensor,
group_list: torch.Tensor,
group_list_type: int = 1,
dynamic_scale: torch.Tensor = None,
w1_scale_bias: torch.Tensor = None,
w2_scale_bias: torch.Tensor = None,
w1_offset: torch.Tensor | None = None,
w2_offset: torch.Tensor | None = None,
fusion: bool = False,
dynamic_eplb: bool = False,
use_mxfp_quant: bool = False,
mxfp_quant_dtype: QuantType | None = None,
act_quant_type: torch.dtype = torch.float8_e4m3fn,
weight_quant_type: torch.dtype | None = None,
scale_type: torch.dtype | None = None,
per_token_scale_type: torch.dtype | None = None,
use_bf16: bool = True,
activation: str | None = None,
swiglu_limit: float = 0.0,
use_w4a8_per_channel_gmm_swiglu: bool = False,
) -> torch.Tensor:
input_hidden_dtype = hidden_states.dtype
use_gmm_swiglu_quant_fusion = use_mxfp_quant or (fusion and not dynamic_eplb)
if use_mxfp_quant:
ensure_mxfp8_moe_available("MXFP MoE MLP path")
if w1_scale_bias is not None or w2_scale_bias is not None:
raise NotImplementedError("MXFP path does not support scale_bias yet.")
if w1_offset is not None or w2_offset is not None:
raise NotImplementedError("MXFP path does not support antiquant offset yet.")
if w1_offset is not None:
unquantized_hidden_states = hidden_states
quantized_hidden_states = None
elif mxfp_quant_dtype == QuantType.W4A16MXFP4:
quantized_hidden_states = None
pertoken_scale = None
elif dynamic_scale is None:
unquantized_hidden_states = hidden_states
hidden_states, pertoken_scale = DeviceOperator.npu_dynamic_quant(
hidden_states=hidden_states,
dynamic_scale=None,
act_quant_type=act_quant_type,
use_mxfp_quant=use_mxfp_quant,
)
dispose_tensor(unquantized_hidden_states)
quantized_hidden_states = None
else:
unquantized_hidden_states = None
pertoken_scale = (
DeviceOperator.maybe_normalize_mxfp_scale_layout(dynamic_scale) if use_mxfp_quant else dynamic_scale
)
quantized_hidden_states = hidden_states
bias1, bias2 = None, None
_output_dtype = w2_scale[0].dtype if isinstance(w2_scale, list) else w2_scale.dtype
weight_prefetch_method = get_weight_prefetch_method()
if weight_prefetch_method:
weight_prefetch_method.maybe_prefetch_moe_weight_postprocess(hidden_states)
is_mc2 = _EXTRA_CTX.moe_comm_type == MoECommType.MC2
if w1_scale_bias is None and w1_offset is None and is_mc2:
if _custom_gmm_swiglu_enabled(fusion, dynamic_eplb) and not use_mxfp_quant:
# gmm1: gate_up_proj & act_fn: swiglu
hidden_states, swiglu_out_scale, _ = torch.ops._C_ascend.grouped_matmul_swiglu_quant_weight_nz_tensor_list(
x=hidden_states,
weight=w1,
weight_scale=w1_scale,
x_scale=pertoken_scale,
group_list=cumsum_group_list(group_list, group_list_type, 0),
swiglu_limit=swiglu_limit,
)
elif use_gmm_swiglu_quant_fusion:
# gmm1: gate_up_proj & act_fn: swiglu
hidden_states, swiglu_out_scale, _ = DeviceOperator.npu_grouped_matmul_swiglu_quant(
x=hidden_states,
weight=_require_single_tensor_for_swiglu_quant(w1, name="w1"),
group_list=cumsum_group_list(group_list, group_list_type, 0),
weight_scale=_require_single_tensor_for_swiglu_quant(w1_scale, name="w1_scale"),
x_scale=pertoken_scale,
bias=None,
use_mxfp_quant=use_mxfp_quant,
act_quant_type=act_quant_type,
weight_quant_type=weight_quant_type,
swiglu_limit=swiglu_limit,
mxfp_quant_dtype=mxfp_quant_dtype,
)
if quantized_hidden_states is not None:
dispose_tensor(quantized_hidden_states)
else:
if w1_scale[0].dtype != torch.float32:
w1_scale[0] = w1_scale[0].to(torch.float32)
# gmm1: gate_up_proj
hidden_states = torch_npu.npu_grouped_matmul(
x=[hidden_states],
weight=w1,
split_item=3,
group_list_type=group_list_type,
group_type=0,
group_list=group_list,
output_dtype=torch.int32,
)[0]
if quantized_hidden_states is not None:
dispose_tensor(quantized_hidden_states)
# act_fn: swiglu
hidden_states, swiglu_out_scale = torch.ops._C_ascend.npu_dequant_swiglu_quant(
x=hidden_states,
weight_scale=w1_scale[0],
activation_scale=pertoken_scale,
bias=None,
quant_scale=None,
quant_offset=None,
group_index=cumsum_group_list(group_list, group_list_type, 1),
activate_left=True,
quant_mode=1,
)
before_gmm2_evt = torch.npu.current_stream().record_event()
# gmm2: down_proj
hidden_states = DeviceOperator.npu_grouped_matmul_gmm2(
hidden_states=hidden_states,
weight=w2,
weight_scale=w2_scale,
per_token_scale=swiglu_out_scale,
group_list=group_list,
group_list_type=group_list_type,
input_dtype=input_hidden_dtype,
act_quant_type=act_quant_type,
weight_quant_type=weight_quant_type,
scale_type=scale_type,
per_token_scale_type=per_token_scale_type,
use_bf16=use_bf16,
use_mxfp_quant=use_mxfp_quant,
bias=None,
fallback_output_dtype=w2_scale[0].dtype if isinstance(w2_scale, list) else w2_scale.dtype,
mxfp_quant_dtype=mxfp_quant_dtype,
)
elif w1_offset is not None:
# gmm1: gate_up_proj
hidden_states = torch_npu.npu_grouped_matmul(
x=[unquantized_hidden_states],
weight=[w1],
antiquant_scale=[w1_scale],
antiquant_offset=[w1_offset],
split_item=2,
group_list_type=group_list_type,
group_type=0,
group_list=group_list,
output_dtype=_output_dtype,
)[0]
dispose_tensor(unquantized_hidden_states)
# act_fn: swiglu
if activation == MoEActivation.SWIGLUSTEP:
hidden_states = AscendSwigluStepAndMul.swiglustep_forward(hidden_states, limit=swiglu_limit or 7.0)
else:
hidden_states = torch_npu.npu_swiglu(hidden_states)
before_gmm2_evt = torch.npu.current_stream().record_event()
# gmm2: down_proj
hidden_states = torch_npu.npu_grouped_matmul(
x=[hidden_states],
weight=[w2],
antiquant_scale=[w2_scale],
antiquant_offset=[w2_offset],
split_item=2,
group_list_type=group_list_type,
group_type=0,
group_list=group_list,
output_dtype=_output_dtype,
)[0]
else:
if w1_scale_bias is not None:
if group_list_type == 0:
group_list = torch.cat([group_list[:1], torch.diff(group_list, dim=0)])
group_list_type = 1
bias1 = w1_scale_bias
bias2 = w2_scale_bias
# TODO w4a8 scene: dynamic acquisition of dtype in the future
_output_dtype = torch.bfloat16
if use_w4a8_per_channel_gmm_swiglu and enable_custom_op() and activation != MoEActivation.SWIGLUSTEP:
hidden_states, swiglu_out_scale = torch.ops._C_ascend.grouped_matmul_swiglu_quant_v2(
x=hidden_states,
weight=w1,
weight_scale=w1_scale if isinstance(w1_scale, list) else [w1_scale],
x_scale=pertoken_scale,
group_list=group_list,
weight_assist_matrix=bias1,
dequant_mode=0,
group_list_type=group_list_type,
swiglu_limit=swiglu_limit,
)
elif _custom_gmm_swiglu_enabled(fusion, dynamic_eplb) and not use_mxfp_quant:
# gmm1: gate_up_proj & act_fn: swiglu
hidden_states, swiglu_out_scale, _ = torch.ops._C_ascend.grouped_matmul_swiglu_quant_weight_nz_tensor_list(
x=hidden_states,
weight=w1,
weight_scale=w1_scale,
x_scale=pertoken_scale,
group_list=cumsum_group_list(group_list, group_list_type, 0),
bias=bias1,
swiglu_limit=swiglu_limit,
)
elif use_gmm_swiglu_quant_fusion and activation != MoEActivation.SWIGLUSTEP:
hidden_states, swiglu_out_scale, _ = DeviceOperator.npu_grouped_matmul_swiglu_quant(
x=hidden_states,
weight=_require_single_tensor_for_swiglu_quant(w1, name="w1"),
group_list=cumsum_group_list(group_list, group_list_type, 0),
weight_scale=_require_single_tensor_for_swiglu_quant(w1_scale, name="w1_scale"),
x_scale=pertoken_scale,
bias=bias1,
use_mxfp_quant=use_mxfp_quant,
act_quant_type=act_quant_type,
weight_quant_type=weight_quant_type,
swiglu_limit=swiglu_limit,
mxfp_quant_dtype=mxfp_quant_dtype,
)
if quantized_hidden_states is not None:
dispose_tensor(quantized_hidden_states)
else:
w1_scale[0] = w1_scale[0].to(w2_scale[0].dtype)
# gmm1: gate_up_proj
hidden_states = torch_npu.npu_grouped_matmul(
x=[hidden_states],
weight=w1,
scale=w1_scale,
bias=bias1,
per_token_scale=[pertoken_scale],
split_item=2,
group_list_type=group_list_type,
group_type=0,
group_list=group_list,
output_dtype=_output_dtype,
)[0]
if quantized_hidden_states is not None:
dispose_tensor(quantized_hidden_states)
# act_fn: swiglu
if activation == MoEActivation.SWIGLUSTEP:
hidden_states = AscendSwigluStepAndMul.swiglustep_forward(hidden_states, limit=swiglu_limit or 7.0)
hidden_states, swiglu_out_scale = torch_npu.npu_dynamic_quant(hidden_states)
elif HAS_TRITON:
from vllm_ascend.ops.triton.activation.swiglu_quant import swiglu_quant
hidden_states, swiglu_out_scale = swiglu_quant(
hidden_states, group_list=group_list, group_list_type=group_list_type
)
else:
hidden_states = torch_npu.npu_swiglu(hidden_states)
hidden_states, swiglu_out_scale = torch_npu.npu_dynamic_quant(hidden_states)
before_gmm2_evt = torch.npu.current_stream().record_event()
# gmm2: down_proj
hidden_states = DeviceOperator.npu_grouped_matmul_gmm2(
hidden_states=hidden_states,
weight=w2,
weight_scale=w2_scale,
per_token_scale=swiglu_out_scale,
group_list=group_list,
group_list_type=group_list_type,
input_dtype=input_hidden_dtype,
act_quant_type=act_quant_type,
weight_quant_type=weight_quant_type,
scale_type=scale_type,
per_token_scale_type=per_token_scale_type,
use_bf16=use_bf16,
use_mxfp_quant=use_mxfp_quant,
bias=bias2,
fallback_output_dtype=_output_dtype,
mxfp_quant_dtype=mxfp_quant_dtype,
)
return hidden_states, before_gmm2_evt
def unquant_apply_mlp(
hidden_states: torch.Tensor,
w1: torch.Tensor,
w2: torch.Tensor,
group_list: torch.Tensor,
w1_bias: torch.Tensor = None,
w2_bias: torch.Tensor = None,
activation: str | None = None,
group_list_type: int = 1,
topk_scales: torch.Tensor | None = None,
need_trans: bool = True,
swiglu_limit: float = 0.0,
lora_context=None,
expanded_row_idx: torch.Tensor | None = None,
topk_ids: torch.Tensor | None = None,
) -> torch.Tensor:
if need_trans:
w1 = w1.transpose(1, 2)
w2 = w2.transpose(1, 2)
gate_up_out = torch_npu.npu_grouped_matmul(
x=[hidden_states],
weight=[w1],
bias=[w1_bias.to(dtype=torch.float32)] if w1_bias is not None else None,
split_item=2,
group_list_type=group_list_type,
group_type=0,
group_list=group_list,
)[0]
# MoE LoRA: only attempt injection when an adapter wraps this layer and the
# comm method provided AllGather routing metadata (expanded_row_idx). Lazy
# import keeps the core MLP free of any LoRA dependency on the common path.
lora_routing = None
if lora_context is not None: # LoRA applied
if expanded_row_idx is None or topk_ids is None:
raise AssertionError(
"MoE LoRA requires expanded_row_idx and topk_ids metadata, "
"which are only available in AllGather communication mode. "
"Please ensure you are running in a supported configuration."
)
from vllm_ascend.lora.fused_moe import moe_lora_apply_w2, moe_lora_apply_w13
# LoRA w13 delta: applied to gate_up_out before activation, with the MLP
# input as the lora_a input (mirrors the base gate_up GMM above).
lora_routing = moe_lora_apply_w13(
lora_context,
gate_up_out=gate_up_out,
hidden_states=hidden_states,
expanded_row_idx=expanded_row_idx,
topk_ids=topk_ids,
)
if activation == MoEActivation.SWIGLUOAI:
num_experts, _, hidden_size = w1.shape
gate_up_out = AscendSwigluOAIAndMul.swiglu_oai_forward(gate_up_out.view(-1, hidden_size))
elif activation == MoEActivation.SWIGLUSTEP:
gate_up_out = AscendSwigluStepAndMul.swiglustep_forward(gate_up_out, limit=swiglu_limit or 7.0)
elif activation == MoEActivation.GELU:
gate, up = gate_up_out.chunk(2, dim=-1)
gate_up_out = torch.nn.functional.gelu(gate) * up
elif activation == MoEActivation.GELU_TANH:
gate, up = gate_up_out.chunk(2, dim=-1)
gate_up_out = torch.nn.functional.gelu(gate, approximate="tanh") * up
else:
if swiglu_limit > 0:
gate, up = gate_up_out.chunk(2, dim=-1)
gate.clamp_(max=swiglu_limit)
up.clamp_(min=-swiglu_limit, max=swiglu_limit)
gate_up_out = torch_npu.npu_swiglu(gate_up_out)
if topk_scales is not None:
gate_up_out *= topk_scales
hidden_states = torch_npu.npu_grouped_matmul(
x=[gate_up_out],
weight=[w2],
bias=[w2_bias.to(dtype=torch.float32)] if w2_bias is not None else None,
split_item=2,
group_list_type=group_list_type,
group_type=0,
group_list=group_list,
)[0]
# LoRA w2 delta: applied to the down-proj output, with the activation output
# as the lora_a input. Reuses the per-row routing computed for w13.
if lora_routing is not None:
moe_lora_apply_w2(
lora_context,
down_out=hidden_states,
silu_out=gate_up_out,
lora_routing=lora_routing,
)
return hidden_states, None
def unified_apply_mlp(*, mlp_compute_input: MoEMlpComputeInput) -> torch.Tensor:
"""
Unified MoE MLP entry.
Quant path is dispatched by DeviceOperator with explicit typed kernel flags.
"""
hidden_states = mlp_compute_input.hidden_states
group_list = mlp_compute_input.group_list
group_list_type = mlp_compute_input.group_list_type
dynamic_scale = mlp_compute_input.dynamic_scale
topk_scales = mlp_compute_input.topk_scales
w1 = mlp_compute_input.weights.w1
w2 = mlp_compute_input.weights.w2
w1_bias = mlp_compute_input.weights.w1_bias
w2_bias = mlp_compute_input.weights.w2_bias
w1_scale = mlp_compute_input.weights.w1_scale
w2_scale = mlp_compute_input.weights.w2_scale
w1_scale_bias = mlp_compute_input.weights.w1_scale_bias
w2_scale_bias = mlp_compute_input.weights.w2_scale_bias
w1_offset = mlp_compute_input.weights.w1_offset
w2_offset = mlp_compute_input.weights.w2_offset
activation = mlp_compute_input.activation
need_trans = mlp_compute_input.need_trans
dynamic_eplb = mlp_compute_input.dynamic_eplb
fusion = mlp_compute_input.fusion
swiglu_limit = mlp_compute_input.swiglu_limit
if not mlp_compute_input.quant.is_quant:
return unquant_apply_mlp(
hidden_states=hidden_states,
w1=w1,
w2=w2,
w1_bias=w1_bias,
w2_bias=w2_bias,
activation=activation,
group_list=group_list,
group_list_type=group_list_type,
topk_scales=topk_scales,
need_trans=need_trans,
swiglu_limit=swiglu_limit,
lora_context=mlp_compute_input.lora_context,
expanded_row_idx=mlp_compute_input.expanded_row_idx,
topk_ids=mlp_compute_input.topk_ids,
)
assert w1_scale is not None and w2_scale is not None
act_quant_type = torch.int8 if mlp_compute_input.quant.is_int_quant else torch.float8_e4m3fn
weight_quant_type = torch.float8_e4m3fn
scale_type = None
per_token_scale_type = None
use_bf16 = hidden_states.dtype == torch.bfloat16
use_mxfp_quant = mlp_compute_input.quant.is_mxfp
mxfp_quant_dtype = mlp_compute_input.quant.quant_type
if use_mxfp_quant:
mxfp = mlp_compute_input.quant.mxfp
assert mxfp is not None, "mlp_compute_input.quant.mxfp is required when quant_type is MXFP8."
act_quant_type = mxfp.act_quant_type or act_quant_type
if mxfp_quant_dtype == QuantType.W4A16MXFP4:
act_quant_type = mxfp.act_quant_type
weight_quant_type = mxfp.weight_quant_type or weight_quant_type
if mxfp_quant_dtype in [QuantType.W4A8MXFP, QuantType.W4A16MXFP4]:
weight_quant_type = mxfp.weight_quant_type
scale_type = mxfp.scale_dtype
per_token_scale_type = mxfp.per_token_scale_dtype
use_bf16 = mxfp.use_bf16
return quant_apply_mlp(
hidden_states=hidden_states,
w1=w1,
w1_scale=w1_scale,
w2=w2,
w2_scale=w2_scale,
group_list=group_list,
dynamic_scale=dynamic_scale,
group_list_type=group_list_type,
w1_scale_bias=w1_scale_bias,
w2_scale_bias=w2_scale_bias,
w1_offset=w1_offset,
w2_offset=w2_offset,
fusion=fusion,
dynamic_eplb=dynamic_eplb,
use_mxfp_quant=use_mxfp_quant,
mxfp_quant_dtype=mxfp_quant_dtype,
act_quant_type=act_quant_type,
weight_quant_type=weight_quant_type,
scale_type=scale_type,
per_token_scale_type=per_token_scale_type,
use_bf16=use_bf16,
activation=activation,
swiglu_limit=swiglu_limit,
use_w4a8_per_channel_gmm_swiglu=mlp_compute_input.quant.use_w4a8_per_channel_gmm_swiglu,
)

View File

@@ -0,0 +1,270 @@
#
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
# Copyright 2023 The vLLM team.
#
# 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.
#
"""Typed runtime contracts and builders for fused MoE execution.
This module is the single entry point for the runtime payloads used across the
fused MoE pipeline.
Relationship overview:
stage params: reusable sub-payloads
- MoERoutingParams
- MoEQuantParams
- internal MXFP leaf: MoEMxfpParams
stage contracts: stage input/output payloads
prepare
-> MoEPrepareOutput
fused_experts input
-> MoEFusedExpertsInput
|- weights: MoEWeights
|- routing: MoERoutingParams
|- quant: MoEQuantParams
dispatch
input -> MoETokenDispatchInput
output -> MoETokenDispatchOutput[TMoECombineMetadata]
TMoECombineMetadata is one of:
- MoEAllGatherCombineMetadata
- MoEAllToAllCombineMetadata
- MoEMC2CombineMetadata
mlp
input -> MoEMlpComputeInput
combine
output -> torch.Tensor
The helper builders below adapt legacy call sites into these typed contracts.
Only the fused_moe package should need to know about the internal MXFP leaf
dataclass directly.
"""
from __future__ import annotations
import torch
import vllm_ascend.ops.fused_moe.moe_stage_params as _stage_params
from vllm_ascend.ops.fused_moe.moe_stage_contracts import (
MoEAllGatherCombineMetadata,
MoEAllToAllCombineMetadata,
MoEFusedExpertsInput,
MoEMC2CombineMetadata,
MoEMlpComputeInput,
MoEPrepareOutput,
MoETokenDispatchInput,
MoETokenDispatchOutput,
MoEWeights,
TMoECombineMetadata,
)
from vllm_ascend.ops.fused_moe.moe_stage_params import (
MoEQuantParams,
MoERoutingParams,
)
from vllm_ascend.quantization.quant_type import QuantType
from vllm_ascend.utils import vllm_version_is
def _build_mxfp_params(
*,
quant_type: QuantType,
mxfp_act_quant_type: torch.dtype | None = None,
mxfp_weight_quant_type: torch.dtype | None = None,
mxfp_scale_dtype: torch.dtype | None = None,
mxfp_per_token_scale_dtype: torch.dtype | None = None,
mxfp_use_bf16: bool | None = None,
) -> _stage_params.MoEMxfpParams | None:
if quant_type not in [QuantType.MXFP8, QuantType.MXFP4, QuantType.W4A8MXFP, QuantType.W4A16MXFP4]:
return None
has_explicit_mxfp_args = any(
value is not None
for value in (
mxfp_act_quant_type,
mxfp_weight_quant_type,
mxfp_scale_dtype,
mxfp_per_token_scale_dtype,
mxfp_use_bf16,
)
)
if not has_explicit_mxfp_args:
raise ValueError("primitive MXFP params are required when quant_type is an MXFP quant type.")
return _stage_params.MoEMxfpParams(
act_quant_type=mxfp_act_quant_type,
weight_quant_type=mxfp_weight_quant_type,
scale_dtype=mxfp_scale_dtype,
per_token_scale_dtype=mxfp_per_token_scale_dtype,
use_bf16=True if mxfp_use_bf16 is None else mxfp_use_bf16,
)
def build_fused_experts_input(
*,
hidden_states: torch.Tensor,
topk_weights: torch.Tensor,
topk_ids: torch.Tensor,
w1: torch.Tensor | list[torch.Tensor],
w2: torch.Tensor | list[torch.Tensor],
quant_type: QuantType,
dynamic_eplb: bool,
expert_map: torch.Tensor | None = None,
global_redundant_expert_num: int = 0,
mc2_mask: torch.Tensor | None = None,
apply_router_weight_on_input: bool = False,
log2phy: torch.Tensor | None = None,
pertoken_scale: torch.Tensor | None = None,
activation: str = "silu",
need_trans: bool = False,
w1_bias: torch.Tensor | None = None,
w2_bias: torch.Tensor | None = None,
comm_quant_mode: int | None = None,
mxfp_act_quant_type: torch.dtype | None = None,
mxfp_weight_quant_type: torch.dtype | None = None,
mxfp_scale_dtype: torch.dtype | None = None,
mxfp_per_token_scale_dtype: torch.dtype | None = None,
mxfp_use_bf16: bool | None = None,
is_per_channel_weight: bool = False,
w1_scale: list[torch.Tensor] | torch.Tensor | None = None,
w2_scale: list[torch.Tensor] | torch.Tensor | None = None,
w1_scale_bias: list[torch.Tensor] | torch.Tensor | None = None,
w2_scale_bias: list[torch.Tensor] | torch.Tensor | None = None,
w1_offset: torch.Tensor | None = None,
w2_offset: torch.Tensor | None = None,
swiglu_limit: float | None = 0.0,
lora_context=None,
) -> MoEFusedExpertsInput:
if not vllm_version_is("0.23.0") and swiglu_limit is None:
swiglu_limit = 0.0
assert swiglu_limit is not None
return MoEFusedExpertsInput(
hidden_states=hidden_states,
topk_weights=topk_weights,
topk_ids=topk_ids,
weights=MoEWeights(
w1=w1,
w2=w2,
w1_bias=w1_bias,
w2_bias=w2_bias,
w1_scale=w1_scale,
w2_scale=w2_scale,
w1_scale_bias=w1_scale_bias,
w2_scale_bias=w2_scale_bias,
w1_offset=w1_offset,
w2_offset=w2_offset,
),
routing=MoERoutingParams(
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,
need_trans=need_trans,
dynamic_eplb=dynamic_eplb,
quant=MoEQuantParams(
quant_type=quant_type,
comm_quant_mode=comm_quant_mode,
mxfp=_build_mxfp_params(
quant_type=quant_type,
mxfp_act_quant_type=mxfp_act_quant_type,
mxfp_weight_quant_type=mxfp_weight_quant_type,
mxfp_scale_dtype=mxfp_scale_dtype,
mxfp_per_token_scale_dtype=mxfp_per_token_scale_dtype,
mxfp_use_bf16=mxfp_use_bf16,
),
is_per_channel_weight=is_per_channel_weight,
),
swiglu_limit=swiglu_limit,
lora_context=lora_context,
)
def build_token_dispatch_input(
*,
fused_experts_input: MoEFusedExpertsInput,
topk_ids: torch.Tensor | None = None,
) -> MoETokenDispatchInput:
return MoETokenDispatchInput(
hidden_states=fused_experts_input.hidden_states,
topk_weights=fused_experts_input.topk_weights,
topk_ids=fused_experts_input.topk_ids if topk_ids is None else topk_ids,
routing=fused_experts_input.routing,
quant=fused_experts_input.quant,
)
def build_mlp_compute_input(
*,
fused_experts_input: MoEFusedExpertsInput,
token_dispatch_output: MoETokenDispatchOutput[TMoECombineMetadata],
use_fusion_ops: bool,
) -> MoEMlpComputeInput:
if fused_experts_input.quant.is_mxfp and fused_experts_input.quant.mxfp is None:
raise ValueError("fused_experts_input.quant.mxfp is required for MXFP quant types.")
expanded_row_idx = getattr(token_dispatch_output.combine_metadata, "expanded_row_idx", None)
return MoEMlpComputeInput(
hidden_states=token_dispatch_output.hidden_states,
group_list=token_dispatch_output.group_list,
group_list_type=token_dispatch_output.group_list_type,
dynamic_scale=token_dispatch_output.dynamic_scale,
topk_scales=token_dispatch_output.topk_scales,
weights=fused_experts_input.weights,
quant=fused_experts_input.quant,
fusion=fused_experts_input.quant.quant_type
in (
QuantType.W8A8,
QuantType.MXFP8,
QuantType.MXFP4,
QuantType.W4A8MXFP,
QuantType.W8A8FP8,
QuantType.W4A16MXFP4,
)
and use_fusion_ops,
activation=fused_experts_input.activation,
need_trans=fused_experts_input.need_trans,
dynamic_eplb=fused_experts_input.dynamic_eplb,
swiglu_limit=fused_experts_input.swiglu_limit,
expanded_row_idx=expanded_row_idx,
topk_ids=fused_experts_input.topk_ids,
lora_context=fused_experts_input.lora_context,
)
__all__ = [
"MoEAllGatherCombineMetadata",
"MoEAllToAllCombineMetadata",
"MoEFusedExpertsInput",
"MoEMC2CombineMetadata",
"MoEMlpComputeInput",
"MoEPrepareOutput",
"MoEQuantParams",
"MoERoutingParams",
"MoETokenDispatchInput",
"MoETokenDispatchOutput",
"MoEWeights",
"TMoECombineMetadata",
"build_fused_experts_input",
"build_token_dispatch_input",
"build_mlp_compute_input",
]

View File

@@ -0,0 +1,165 @@
#
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
# Copyright 2023 The vLLM team.
#
# 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 __future__ import annotations
from dataclasses import dataclass
from typing import Any, Generic, TypeVar
import numpy as np
import torch
from vllm_ascend.ops.fused_moe.moe_stage_params import MoEQuantParams, MoERoutingParams
TMoECombineMetadata = TypeVar("TMoECombineMetadata")
# prepare -> fused_experts
@dataclass(frozen=True, slots=True)
class MoEPrepareOutput:
"""Typed output from prepare stage."""
hidden_states: torch.Tensor
router_logits: torch.Tensor
mc2_mask: torch.Tensor | None
padded_hidden_states_shape: torch.Size | None
pertoken_scale: torch.Tensor | None = None
@dataclass(frozen=True, slots=True)
class MoEWeights:
"""Dense and quantized weight payloads consumed by MoE execution."""
w1: torch.Tensor | list[torch.Tensor]
w2: torch.Tensor | list[torch.Tensor]
w1_bias: torch.Tensor | None = None
w2_bias: torch.Tensor | None = None
w1_scale: torch.Tensor | list[torch.Tensor] | None = None
w2_scale: torch.Tensor | list[torch.Tensor] | None = None
w1_scale_bias: torch.Tensor | list[torch.Tensor] | None = None
w2_scale_bias: torch.Tensor | list[torch.Tensor] | None = None
w1_offset: torch.Tensor | None = None
w2_offset: torch.Tensor | None = None
@dataclass(frozen=True, slots=True)
class MoEFusedExpertsInput:
"""Top-level input for the routed experts pipeline."""
hidden_states: torch.Tensor
topk_weights: torch.Tensor
topk_ids: torch.Tensor
weights: MoEWeights
routing: MoERoutingParams
quant: MoEQuantParams
activation: str = "silu"
need_trans: bool = False
dynamic_eplb: bool = False
swiglu_limit: float = 0.0
# Optional per-layer MoE LoRA state (vllm_ascend.lora MoELoRAContext).
# ``Any`` avoids coupling the core contracts to the LoRA module; only the
# unquant MLP path reads it, and only when a LoRA adapter is active.
lora_context: Any = None
@dataclass(frozen=True, slots=True)
class MoETokenDispatchInput:
"""Input to token dispatch."""
hidden_states: torch.Tensor
topk_weights: torch.Tensor
topk_ids: torch.Tensor
routing: MoERoutingParams
quant: MoEQuantParams
# dispatch carry-over state consumed by combine
@dataclass(frozen=True, slots=True)
class MoEMC2CombineMetadata:
topk_ids: torch.Tensor
topk_weights: torch.Tensor
expert_map: torch.Tensor | None
ep_recv_counts: torch.Tensor
tp_recv_counts: torch.Tensor
assist_info_for_combine: torch.Tensor
expand_scales: torch.Tensor | None
quant: MoEQuantParams
mc2_mask: torch.Tensor | None = None
@dataclass(frozen=True, slots=True)
class MoEAllGatherCombineMetadata:
topk_weights: torch.Tensor
expanded_row_idx: torch.Tensor
restore_shape: torch.Size
@dataclass(frozen=True, slots=True)
class MoEAllToAllCombineMetadata:
input_splits: np.ndarray
output_splits: np.ndarray
topk_weights: torch.Tensor
reversed_local_input_permutation_mapping: torch.Tensor
reversed_global_input_permutation_mapping: torch.Tensor | None
hidden_shape: torch.Size
hidden_shape_before_permute: torch.Size
@dataclass(frozen=True, slots=True)
class MoETokenDispatchOutput(Generic[TMoECombineMetadata]):
hidden_states: torch.Tensor
group_list: torch.Tensor
group_list_type: int
combine_metadata: TMoECombineMetadata
dynamic_scale: torch.Tensor | None = None
topk_scales: torch.Tensor | None = None
# dispatch -> mlp -> combine
@dataclass(frozen=True, slots=True)
class MoEMlpComputeInput:
"""Input to MLP compute."""
hidden_states: torch.Tensor
group_list: torch.Tensor
group_list_type: int
dynamic_scale: torch.Tensor | None
topk_scales: torch.Tensor | None
weights: MoEWeights
quant: MoEQuantParams
fusion: bool
activation: str = "silu"
need_trans: bool = False
dynamic_eplb: bool = False
swiglu_limit: float = 0.0
expanded_row_idx: torch.Tensor | None = None
topk_ids: torch.Tensor | None = None
# Optional per-layer MoE LoRA state, propagated from MoEFusedExpertsInput.
lora_context: Any = None
__all__ = [
"MoEPrepareOutput",
"MoEWeights",
"MoEFusedExpertsInput",
"MoETokenDispatchInput",
"MoEMC2CombineMetadata",
"MoEAllGatherCombineMetadata",
"MoEAllToAllCombineMetadata",
"MoETokenDispatchOutput",
"MoEMlpComputeInput",
"TMoECombineMetadata",
]

View File

@@ -0,0 +1,127 @@
#
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
# Copyright 2023 The vLLM team.
#
# 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 __future__ import annotations
from dataclasses import dataclass
import torch
import torch_npu
from vllm_ascend.quantization.quant_type import QuantType
@dataclass(frozen=True, slots=True)
class MoERoutingParams:
"""Routing and dispatch side inputs for one MoE invocation.
`pertoken_scale` is intentionally kept here even though it is not a pure
routing concept. It is used by pre-quantized activation flows, currently
the AllGather + EP W8A8 prepare path, where prepare emits per-token
activation scales and dispatch needs to carry them forward so the MLP
quant path can reuse those scales instead of requantizing activations.
"""
expert_map: torch.Tensor | None
global_redundant_expert_num: int
mc2_mask: torch.Tensor | None
apply_router_weight_on_input: bool
log2phy: torch.Tensor | None = None
# Precomputed activation scales from prepare stage for quantized dispatch.
pertoken_scale: torch.Tensor | None = None
@dataclass(frozen=True, slots=True)
class MoEMxfpParams:
"""Internal MXFP-only precision settings used by fused_moe runtime."""
act_quant_type: torch.dtype | None = None
weight_quant_type: torch.dtype | None = None
scale_dtype: torch.dtype | None = None
per_token_scale_dtype: torch.dtype | None = None
use_bf16: bool = True
@dataclass(frozen=True, slots=True)
class MoEQuantParams:
"""Quant mode, backend override, and optional internal MXFP leaf config."""
quant_type: QuantType = QuantType.NONE
comm_quant_mode: int | None = None
mxfp: MoEMxfpParams | None = None
is_per_channel_weight: bool = False
@property
def is_quant(self) -> bool:
return self.quant_type != QuantType.NONE
@property
def is_mxfp(self) -> bool:
return self.quant_type in (QuantType.MXFP8, QuantType.MXFP4, QuantType.W4A8MXFP, QuantType.W4A16MXFP4)
@property
def is_w4a4_mxfp(self) -> bool:
return self.quant_type == QuantType.MXFP4
@property
def is_int_quant(self) -> bool:
return self.quant_type in (QuantType.W8A8, QuantType.W4A8)
@property
def is_fp8(self) -> bool:
return self.quant_type == QuantType.W8A8FP8
@property
def use_w4a8_per_channel_gmm_swiglu(self) -> bool:
return self.quant_type == QuantType.W4A8 and self.is_per_channel_weight
@property
def dispatch_with_quant(self) -> bool:
return self.quant_type in (
QuantType.W8A8,
QuantType.W4A8,
QuantType.MXFP8,
QuantType.MXFP4,
QuantType.W4A8MXFP,
QuantType.W8A8FP8,
)
@property
def get_dst_type(self):
if self.is_w4a4_mxfp:
return torch_npu.float4_e2m1fn_x2
elif self.is_mxfp or self.is_fp8:
return torch.float8_e4m3fn
elif self.dispatch_with_quant:
return torch.int8
else:
return None
@property
def get_scale_type(self):
if self.is_mxfp:
return torch.float8_e8m0fnu
elif self.dispatch_with_quant:
return torch.float32
else:
return None
__all__ = [
"MoERoutingParams",
"MoEMxfpParams",
"MoEQuantParams",
]

View File

@@ -0,0 +1,548 @@
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
# Copyright 2023 The vLLM team.
#
# 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.
# This file is a part of the vllm-ascend project.
from abc import ABC, abstractmethod
import torch
import torch.distributed as dist
import torch.nn as nn
import torch_npu
from vllm.distributed.parallel_state import (
get_dp_group,
get_pcp_group,
get_tensor_model_parallel_rank,
get_tensor_model_parallel_world_size,
)
from vllm.forward_context import get_forward_context
from vllm.model_executor.layers.fused_moe import FusedMoEConfig
from vllm_ascend.ascend_config import get_ascend_config
from vllm_ascend.ascend_forward_context import _EXTRA_CTX
from vllm_ascend.distributed.utils import fc3_all_gather_and_maybe_unpad_impl
from vllm_ascend.ops.fused_moe.moe_runtime_args import MoEPrepareOutput
from vllm_ascend.quantization.quant_type import QuantType
from vllm_ascend.utils import enable_sp, enable_sp_by_pass, npu_stream_switch
class PrepareAndFinalize(ABC):
"""
Abstract base class for MoE (Mixture-of-Experts) tensor preparation and finalization
in distributed environments. Subclasses implement specific communication strategies
(e.g., AllGather, All2All, MC2) to handle tensor padding, slicing,
broadcasting, and reduction across TP/DP/EP groups.
Attributes:
moe_config (FusedMoEConfig): Configuration object containing TP/DP/EP group info,
sizes, ranks, and communication settings.
"""
quant_stream: torch.npu.Stream | None = None
def __init__(self, moe_config: FusedMoEConfig):
self.moe_config = moe_config
ascend_config = get_ascend_config()
self.multistream_overlap_gate = ascend_config.multistream_overlap_gate
if self.multistream_overlap_gate and PrepareAndFinalize.quant_stream is None:
PrepareAndFinalize.quant_stream = torch.npu.Stream()
@abstractmethod
def prepare(
self,
hidden_states: torch.Tensor,
router_logits: torch.Tensor,
enable_shared_expert_dp: bool = False,
replace_allreduce: bool = False,
quant_type: QuantType = QuantType.NONE,
) -> MoEPrepareOutput:
"""
Prepare tensors before MoE computation. May involve:
- Padding to align communication boundaries
- Slicing across tensor-parallel ranks
- Broadcasting across data-parallel ranks
Args:
hidden_states (torch.Tensor): Input features, shape [num_tokens, hidden_size]
router_logits (torch.Tensor): Router outputs, shape [num_tokens, num_experts]
enable_shared_expert_dp (bool): Skip DP communication for shared experts
replace_allreduce (bool): Bypass default all-reduce behavior
quant_type: none, w8a8, w4a8, mxfp8, or mxfp4
Returns:
MoEPrepareOutput:
- processed hidden_states (may be padded/sliced/broadcasted)
- processed router_logits (may be recomputed or broadcasted)
- optional communication mask (e.g., mc2_mask for sparse ops)
- optional padded hidden state shape for finalization
- optional per-token scale for quantized path
"""
raise NotImplementedError("Prepare not implemented.")
def finalize(
self,
hidden_states: torch.Tensor,
reduce_results: bool,
padded_hidden_states_shape: torch.Size | None = None,
) -> torch.Tensor:
"""
Finalize MoE output. May involve:
- Gathering sliced tensors across TP ranks
- Reducing or scattering across DP ranks
- Unpadding to original token count
- Applying all-reduce across TP/EP if requested
Args:
hidden_states (torch.Tensor): MoE layer output, possibly padded or sliced
reduce_results (bool): Whether to apply all-reduce across TP/EP groups
Returns:
torch.Tensor: Final output with shape [original_num_tokens, hidden_size]
"""
raise NotImplementedError("Finalize function not implemented.")
class PrepareAndFinalizeWithAll2All(PrepareAndFinalize):
"""
MoE communication strategy using All-to-All style slicing.
Similar to MC2 but does not use mc2_mask; instead pads to TP size for uniform slicing.
Will be used when num_tokens exceed mc2's limitation (512 tokens/rank).
"""
def __init__(self, moe_config: FusedMoEConfig):
super().__init__(moe_config)
self._restore_tp_across_dp()
def _restore_tp_across_dp(self):
"""Restore original TP configuration (same as MC2)."""
self.tp_size = get_tensor_model_parallel_world_size()
self.tp_rank = get_tensor_model_parallel_rank()
def prepare(
self,
hidden_states: torch.Tensor,
router_logits: torch.Tensor,
enable_shared_expert_dp: bool = False,
replace_allreduce: bool = False,
quant_type=QuantType.NONE,
) -> MoEPrepareOutput:
"""
Preparation steps:
1. Pad hidden_states and router_logits to next multiple of TP size.
2. If TP > 1, split along token dim and select current TP rank's slice.
3. Save splits for later all-gather in finalize.
Skips if `enable_shared_expert_dp` or `replace_allreduce` is True.
Returns:
MoEPrepareOutput where `mc2_mask` is None for All2All path.
"""
self.replace_allreduce = replace_allreduce
self.enable_shared_expert_dp = enable_shared_expert_dp
padded_hidden_states_shape = hidden_states.shape
if not (self.replace_allreduce or self.enable_shared_expert_dp):
self.num_tokens, _ = hidden_states.shape
pad_size = self.tp_size - self.num_tokens # Pad to TP size (cyclic)
if pad_size > 0:
hidden_states = nn.functional.pad(hidden_states, (0, 0, 0, pad_size))
router_logits = nn.functional.pad(router_logits, (0, 0, 0, pad_size))
padded_hidden_states_shape = hidden_states.shape
if self.tp_size > 1:
split_hidden_states = torch.tensor_split(hidden_states, self.tp_size, dim=0)
split_router_logits = torch.tensor_split(router_logits, self.tp_size, dim=0)
hidden_states = split_hidden_states[self.tp_rank]
router_logits = split_router_logits[self.tp_rank]
return MoEPrepareOutput(
hidden_states=hidden_states,
router_logits=router_logits,
mc2_mask=None,
padded_hidden_states_shape=padded_hidden_states_shape,
pertoken_scale=None,
)
def pad_and_split_input_ids(
self,
input_ids,
):
if not (self.replace_allreduce or self.enable_shared_expert_dp):
pad_size = self.tp_size - self.num_tokens
if pad_size > 0:
input_ids = nn.functional.pad(input_ids, (0, pad_size))
if self.tp_size > 1:
input_ids = torch.tensor_split(input_ids, self.tp_size, dim=0)
input_ids = input_ids[self.tp_rank]
return input_ids
def finalize(
self,
hidden_states: torch.Tensor,
reduce_results: bool,
padded_hidden_states_shape: torch.Size | None = None,
) -> torch.Tensor:
"""
Finalization steps:
1. If TP > 1, all-gather slices to reconstruct full tensor.
2. Unpad to original token count.
3. Return [original_num_tokens, hidden_size] tensor.
Skips if `enable_shared_expert_dp` or `replace_allreduce` is True.
"""
if not (self.enable_shared_expert_dp or self.replace_allreduce):
if self.tp_size > 1:
assert padded_hidden_states_shape is not None
# Cannot reuse `split_hidden_states` from prepare phase as it
# may share memory with original hidden_states. Since shared
# experts may use the original tensor, reusing it would cause
# in-place modification during all_gather, corrupting the data.
gathered_hidden_states = torch.empty(
padded_hidden_states_shape, device=hidden_states.device, dtype=hidden_states.dtype
)
split_hidden_states = torch.tensor_split(gathered_hidden_states, self.tp_size, dim=0)
dist.all_gather(list(split_hidden_states), hidden_states, self.moe_config.tp_group.device_group)
hidden_states = gathered_hidden_states
if self.num_tokens < hidden_states.shape[0]:
hidden_states = hidden_states[: self.num_tokens]
return hidden_states
class PrepareAndFinalizeWithMC2(PrepareAndFinalizeWithAll2All):
"""
MoE communication strategy using MC2, which is based on All2All. Hence, it inherits
All2All and share the same finalize method.
Designed for Ascend or environments requiring explicit padding and slicing control.
Relies on `mc2_mask` and `padded_num_tokens` from forward_context for alignment.
"""
def __init__(self, moe_config: FusedMoEConfig):
super().__init__(moe_config)
self._restore_tp_across_dp()
def _restore_tp_across_dp(self):
"""
Restore original TP configuration.
vLLM flattens TP and DP into a single dimension; this method recovers
the true TP world size and rank for correct tensor slicing.
"""
self.tp_size = get_tensor_model_parallel_world_size()
self.tp_rank = get_tensor_model_parallel_rank()
def prepare(
self,
hidden_states: torch.Tensor,
router_logits: torch.Tensor,
enable_shared_expert_dp: bool = False,
replace_allreduce: bool = False,
quant_type=QuantType.NONE,
) -> MoEPrepareOutput:
"""
Preparation steps:
1. Fetch `mc2_mask` and target padding length from forward context.
2. Pad `hidden_states` and `router_logits` to target length if needed.
3. If TP > 1, split tensors along token dimension and select current TP rank's slice.
4. Split and return corresponding `mc2_mask`.
Skips padding/slicing if `enable_shared_expert_dp` or `replace_allreduce` is True.
Returns:
MoEPrepareOutput, possibly sliced/padded.
"""
self.replace_allreduce = replace_allreduce
self.enable_shared_expert_dp = enable_shared_expert_dp
mc2_mask = _EXTRA_CTX.mc2_mask
if self.tp_size > 1:
# Also slice mc2_mask
split_mc2_mask = torch.tensor_split(mc2_mask, self.tp_size, dim=0)
mc2_mask = split_mc2_mask[self.tp_rank]
padded_hidden_states_shape = hidden_states.shape
if not self.replace_allreduce:
self.num_tokens, _ = hidden_states.shape
target_pad_length = _EXTRA_CTX.padded_num_tokens
pad_size = target_pad_length - self.num_tokens
# Pad if necessary (unless shared expert DP is enabled)
if pad_size > 0 and not self.enable_shared_expert_dp:
hidden_states = nn.functional.pad(hidden_states, (0, 0, 0, pad_size))
router_logits = nn.functional.pad(router_logits, (0, 0, 0, pad_size))
padded_hidden_states_shape = hidden_states.shape
# Slice across TP ranks
if self.tp_size > 1 and not self.enable_shared_expert_dp:
split_hidden_states = torch.tensor_split(hidden_states, self.tp_size, dim=0)
split_router_logits = torch.tensor_split(router_logits, self.tp_size, dim=0)
hidden_states = split_hidden_states[self.tp_rank]
router_logits = split_router_logits[self.tp_rank]
return MoEPrepareOutput(
hidden_states=hidden_states,
router_logits=router_logits,
mc2_mask=mc2_mask,
padded_hidden_states_shape=padded_hidden_states_shape,
pertoken_scale=None,
)
def pad_and_split_input_ids(
self,
input_ids,
):
if not self.replace_allreduce:
forward_context = get_forward_context()
target_pad_length = forward_context.padded_num_tokens
pad_size = target_pad_length - self.num_tokens
if pad_size > 0 and not self.enable_shared_expert_dp:
input_ids = nn.functional.pad(input_ids, (0, pad_size))
if self.tp_size > 1 and not self.enable_shared_expert_dp:
input_ids = torch.tensor_split(input_ids, self.tp_size, dim=0)
input_ids = input_ids[self.tp_rank]
return input_ids
class PrepareAndFinalizeWithAllGather(PrepareAndFinalize):
"""
MoE communication strategy using All-Gather + Reduce-Scatter on EP group.
There are two sets of prepare and finalize:
1. _prepare_with_dp_group/_finalize_with_dp_group: When sequence parallelism is not enabled,
we gather inputs across DP ranks before MoE, scatter outputs after.
The communication and calculation process is as follows (AG, AR and RS
are abbreviations for All-Gather, All-Reduce and Reduce-Scatter, respectively):
Attn → TP AR → DP AG → MoE → DP RS → TP AR
2. _prepare_with_ep_group/_finalize_with_ep_group: When sequence parallelism is enabled,
the above process becomes:
TP AG → Attn → TP RS → TP AG → DP AG → MoE → DP RS → TP RS
This strategy further combines TP AG + DP AG into EP All-Gather and TP RS + DP RS
into EP Reduce-Scatter to improve communication performance. The optimized process is as follows:
TP AG → Attn → TP RS → EP AG → MoE → EP RS
"""
def prepare(
self,
hidden_states: torch.Tensor,
router_logits: torch.Tensor,
enable_shared_expert_dp: bool = False,
replace_allreduce: bool = False,
quant_type=QuantType.NONE,
) -> MoEPrepareOutput:
"""
Preparation steps:
AllGather hidden_states and router_logits to form global tensors.
Returns:
MoEPrepareOutput with global tensors.
"""
if enable_sp() or enable_sp_by_pass():
return self._prepare_with_ep_group(hidden_states, router_logits, quant_type)
return self._prepare_with_dp_group(hidden_states, router_logits, enable_shared_expert_dp, replace_allreduce)
def _prepare_with_ep_group(
self, hidden_states: torch.Tensor, router_logits: torch.Tensor, quant_type=QuantType.NONE
) -> MoEPrepareOutput:
pertoken_scale = None
if quant_type == QuantType.W8A8:
hidden_states, pertoken_scale = torch_npu.npu_dynamic_quant(hidden_states)
elif quant_type in (QuantType.MXFP8, QuantType.W4A8MXFP):
hidden_states, pertoken_scale = torch_npu.npu_dynamic_mx_quant(
hidden_states,
dst_type=torch.float8_e4m3fn,
)
elif quant_type == QuantType.MXFP4:
hidden_states, pertoken_scale = torch_npu.npu_dynamic_mx_quant(
hidden_states,
dst_type=torch_npu.float4_e2m1fn_x2,
round_mode="round",
)
if self.multistream_overlap_gate:
assert PrepareAndFinalize.quant_stream is not None
PrepareAndFinalize.quant_stream.wait_stream(torch.npu.current_stream())
with npu_stream_switch(PrepareAndFinalize.quant_stream, enabled=self.multistream_overlap_gate):
hidden_states = fc3_all_gather_and_maybe_unpad_impl(hidden_states)
else:
hidden_states = torch.ops.vllm.maybe_all_gather_and_maybe_unpad(hidden_states, True, True)
router_logits = torch.ops.vllm.maybe_all_gather_and_maybe_unpad(router_logits, True, True)
# TODO(fuzhihong): To adapt to self.num_token in the all_gather_input_id_with_dp_group method,
# when flashcomm1 is used and dp = N(N >=2).
self.num_tokens = hidden_states.shape[0]
if pertoken_scale is not None:
pertoken_scale = torch.ops.vllm.maybe_all_gather_and_maybe_unpad(pertoken_scale, True, True)
if self.multistream_overlap_gate:
torch.npu.current_stream().wait_stream(PrepareAndFinalize.quant_stream)
if self.moe_config.pcp_size > 1:
max_tokens_across_pcp = _EXTRA_CTX.max_tokens_across_pcp
self.num_tokens_pcp = hidden_states.shape[0]
pad_size = max_tokens_across_pcp - self.num_tokens_pcp
if pad_size > 0:
hidden_states = nn.functional.pad(hidden_states, (0, 0, 0, pad_size))
router_logits = nn.functional.pad(router_logits, (0, 0, 0, pad_size))
if pertoken_scale is not None:
pertoken_scale = (
nn.functional.pad(pertoken_scale, (0, pad_size))
if pertoken_scale.dim() == 1
else nn.functional.pad(pertoken_scale, (0, 0, 0, pad_size))
)
hidden_states = get_pcp_group().all_gather(hidden_states, dim=0)
router_logits = get_pcp_group().all_gather(router_logits, dim=0)
if pertoken_scale is not None:
pertoken_scale = get_pcp_group().all_gather(pertoken_scale, dim=0)
return MoEPrepareOutput(
hidden_states=hidden_states,
router_logits=router_logits,
mc2_mask=None,
padded_hidden_states_shape=None,
pertoken_scale=pertoken_scale,
)
def _prepare_with_dp_group(
self,
hidden_states: torch.Tensor,
router_logits: torch.Tensor,
enable_shared_expert_dp: bool = False,
replace_allreduce: bool = False,
quant_type=QuantType.NONE,
) -> MoEPrepareOutput:
"""
Preparation steps:
1. Fetch max token count across DP group from forward context.
2. Pad local tensors to that size.
3. All-gather across DP group to form global input tensor.
Returns:
MoEPrepareOutput with global tensors.
"""
self.enable_shared_expert_dp = enable_shared_expert_dp
if self.moe_config.dp_size > 1:
max_tokens_across_dp = _EXTRA_CTX.max_tokens_across_dp
self.num_tokens = hidden_states.shape[0]
pad_size = max_tokens_across_dp - self.num_tokens
if pad_size > 0:
hidden_states = nn.functional.pad(hidden_states, (0, 0, 0, pad_size))
router_logits = nn.functional.pad(router_logits, (0, 0, 0, pad_size))
# All-gather across DP group
hidden_states = self.moe_config.dp_group.all_gather(hidden_states, 0)
router_logits = self.moe_config.dp_group.all_gather(router_logits, 0)
if self.moe_config.pcp_size > 1:
max_tokens_across_pcp = _EXTRA_CTX.max_tokens_across_pcp
self.num_tokens_pcp = hidden_states.shape[0]
pad_size = max_tokens_across_pcp - self.num_tokens_pcp
if pad_size > 0:
hidden_states = nn.functional.pad(hidden_states, (0, 0, 0, pad_size))
router_logits = nn.functional.pad(router_logits, (0, 0, 0, pad_size))
hidden_states = get_pcp_group().all_gather(
hidden_states,
dim=0,
)
router_logits = get_pcp_group().all_gather(
router_logits,
dim=0,
)
return MoEPrepareOutput(
hidden_states=hidden_states,
router_logits=router_logits,
mc2_mask=None,
padded_hidden_states_shape=None,
pertoken_scale=None,
)
def all_gather_input_id_with_dp_group(self, input_ids: torch.Tensor) -> torch.Tensor:
if self.moe_config.dp_size > 1:
max_tokens_across_dp = _EXTRA_CTX.max_tokens_across_dp
pad_size = max_tokens_across_dp - self.num_tokens
if pad_size > 0:
input_ids = nn.functional.pad(input_ids, (0, pad_size))
input_ids = self.moe_config.dp_group.all_gather(input_ids, 0)
return input_ids
def finalize(
self,
hidden_states: torch.Tensor,
reduce_results: bool,
padded_hidden_states_shape: torch.Size | None = None,
) -> torch.Tensor:
"""
Finalization steps:
Reduce Scatter hidden states.
Returns:
Tensor with shape [local_num_tokens, hidden_size]
"""
if enable_sp() or enable_sp_by_pass():
return self._finalize_with_ep_group(hidden_states)
return self._finalize_with_dp_group(hidden_states, reduce_results)
def _finalize_with_ep_group(self, hidden_states: torch.Tensor) -> torch.Tensor:
"""
Argument `reduce_results` is not needed in this func. Given sequence parallelism is enabled:
1. Reduce_results is False usually happens when models have shared experts and need to
allreduce hidden states after results of shared experts and routed experts are added in FusedMoe.
We do reduce scatter for hidden states here, then skip allreudce in FusedMoe and add it to the
result of shared experts.
2 Reduce_results is True usually happens when model has no shared experts. We still do reduce scatter
here, then skip allreudce in FusedMoe.
"""
if self.moe_config.pcp_size > 1:
hidden_states = get_pcp_group().reduce_scatter(hidden_states, dim=0)
hidden_states = hidden_states[: self.num_tokens_pcp]
hidden_states = torch.ops.vllm.maybe_pad_and_reduce(hidden_states, True)
return hidden_states
def _finalize_with_dp_group(self, hidden_states: torch.Tensor, reduce_results: bool) -> torch.Tensor:
"""
Finalization steps:
1. If DP > 1 and not shared expert, reduce-scatter output across DP group.
2. Slice to original local token count.
3. If `reduce_results=True` and TP/EP > 1, apply tensor_model_parallel_all_reduce.
Returns:
Tensor with shape [original_local_num_tokens, hidden_size]
"""
if self.moe_config.dp_size > 1 and not self.enable_shared_expert_dp:
hidden_states = get_dp_group().reduce_scatter(hidden_states, 0)
hidden_states = hidden_states[: self.num_tokens]
if self.moe_config.pcp_size > 1:
hidden_states = get_pcp_group().reduce_scatter(hidden_states, dim=0)
hidden_states = hidden_states[: self.num_tokens_pcp]
return hidden_states

View File

@@ -0,0 +1,709 @@
# SPDX-License-Identifier: Apache-2.0
# Copyright (c) 2024; NVIDIA CORPORATION. All rights reserved.
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
# Copyright 2023 The vLLM team.
# Copyright 2023 DeepSeek-AI and the HuggingFace Inc. team. All rights reserved.
#
# This code is based on EleutherAI's GPT-NeoX library and the GPT-NeoX
# and OPT implementations in this library. It has been modified from its
# original forms to accommodate minor architectural differences compared
# to GPT-NeoX and OPT used by the Meta AI team that trained the model.
#
# 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 abc import ABC, abstractmethod
from typing import Generic
import torch
import torch_npu
from vllm.config import get_current_vllm_config
from vllm.distributed.parallel_state import get_ep_group
from vllm_ascend.ascend_config import get_ascend_config
from vllm_ascend.ascend_forward_context import get_mc2_tokens_capacity
from vllm_ascend.device.device_op import DeviceOperator
from vllm_ascend.distributed.parallel_state import get_mc2_group
from vllm_ascend.ops.fused_moe.comm_utils import async_all_to_all, gather_from_sequence_parallel_region
from vllm_ascend.ops.fused_moe.moe_runtime_args import (
MoEAllGatherCombineMetadata,
MoEAllToAllCombineMetadata,
MoEMC2CombineMetadata,
MoETokenDispatchInput,
MoETokenDispatchOutput,
TMoECombineMetadata,
)
from vllm_ascend.quantization.quant_type import QuantType
from vllm_ascend.utils import (
AscendDeviceType,
get_ascend_device_type,
is_hierarchical_communication_enabled,
should_skip_allreduce_across_dp_group,
)
EXPERT_TOKEN_NUMS_TYPE_CUMSUM = 0
EXPERT_TOKEN_NUMS_TYPE_COUNT = 1
def _get_expert_token_nums_type(token_dispatch_input: MoETokenDispatchInput) -> int:
# grouped_matmul_swiglu_quant_v2 consumes per-expert counts; existing
# MC2 grouped-matmul paths consume prefix sums.
if token_dispatch_input.quant.use_w4a8_per_channel_gmm_swiglu:
return EXPERT_TOKEN_NUMS_TYPE_COUNT
return EXPERT_TOKEN_NUMS_TYPE_CUMSUM
class MoETokenDispatcher(ABC, Generic[TMoECombineMetadata]):
def __init__(self, **kwargs) -> None:
"""
Initialize the MoE Token Dispatcher.
"""
self.top_k = kwargs.get("top_k", 0)
self.num_experts = kwargs.get("num_experts", 0)
@property
def ep_group(self):
"""Get expert model parallel group."""
return get_ep_group().device_group
@property
def ep_rank(self):
return get_ep_group().rank_in_group
@property
def ep_size(self):
return get_ep_group().world_size
@abstractmethod
def token_dispatch(
self,
token_dispatch_input: MoETokenDispatchInput,
) -> MoETokenDispatchOutput[TMoECombineMetadata]:
raise NotImplementedError("Dispatch function not implemented.")
@abstractmethod
def token_combine(
self,
hidden_states: torch.Tensor,
combine_metadata: TMoECombineMetadata,
bias: torch.Tensor | None = None,
) -> torch.Tensor:
raise NotImplementedError("Combine function not implemented.")
class TokenDispatcherWithMC2(MoETokenDispatcher[MoEMC2CombineMetadata]):
def __init__(self, **kwargs):
super().__init__(**kwargs)
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)
self.ep_rank_id = get_mc2_group().rank_in_group
self.ep_world_size = get_mc2_group().world_size
self.enable_dispatch_v2 = hasattr(torch_npu, "npu_moe_distribute_dispatch_v2")
self.need_extra_args = get_ascend_device_type() in [AscendDeviceType.A3, AscendDeviceType.A5]
self.a5_need_extra_args = get_ascend_device_type() == AscendDeviceType.A5
# NOTE: When in A2, setting the environment variables HCCL_INTRA_PCIE_ENABLE=1 and
# HCCL_INTRA_ROCE_ENABLE=0 can reduce cross-machine communication traffic and significantly
# improve communication performance.
# When enable hierarchical communication, param `expert_scales` need to be passed in.
self.need_expert_scale = is_hierarchical_communication_enabled()
# Here we need to calculate the global_bs = max_bs_per_rank * ep_world_size to execute
# dispatch & combine operators with different input num_tokens per rank.
vllm_config = get_current_vllm_config()
tp_size = vllm_config.parallel_config.tensor_parallel_size
mc2_tokens_capacity = get_mc2_tokens_capacity()
num_tokens_per_tp_rank = mc2_tokens_capacity // tp_size
_max_global_bs = num_tokens_per_tp_rank * self.ep_world_size
# When allreduce across DP is not skipped, tokens are uniform across ranks:
# use global_bs=0 (uniform mode) and pass mc2_mask.
# When allreduce is skipped, tokens may differ per rank:
# use the real global_bs and do NOT pass mc2_mask.
self.global_bs = _max_global_bs if should_skip_allreduce_across_dp_group(vllm_config) else 0
# NOTE: When enable_mc2_hierarchy_comm is true, we need pass in `comm_alg` to mc2 op.
self.need_comm_alg = get_ascend_config().enable_mc2_hierarchy_comm
if not self.enable_dispatch_v2 and self.need_comm_alg:
raise RuntimeError(
"PTA and CANN version is too old to support mc2 hierarchy comm, please upgrade your version."
)
def refresh_hccl_group(self) -> None:
"""Refresh MC2 communicator metadata after HCCL groups are recreated."""
device_group = get_mc2_group().device_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)
def get_dispatch_mc2_kwargs(
self,
token_dispatch_input: MoETokenDispatchInput,
):
hidden_states = token_dispatch_input.hidden_states
topk_weights = token_dispatch_input.topk_weights
topk_ids = token_dispatch_input.topk_ids
expert_map = token_dispatch_input.routing.expert_map
global_redundant_expert_num = token_dispatch_input.routing.global_redundant_expert_num
comm_quant_mode = token_dispatch_input.quant.comm_quant_mode
assert expert_map is not None, "expert_map is required for MC2 token dispatch."
# NOTE: quant_mode differs by quant feature:
# - Legacy int communication quantization uses quant_mode=2.
# - A5 MXFP communication uses quant_mode=4.
if comm_quant_mode is not None:
quant_mode = comm_quant_mode
elif token_dispatch_input.quant.dispatch_with_quant:
quant_mode = 4 if self.a5_need_extra_args and token_dispatch_input.quant.is_mxfp else 2
else:
quant_mode = 0
self.moe_expert_num = len(expert_map) + global_redundant_expert_num
expert_token_nums_type = _get_expert_token_nums_type(token_dispatch_input)
kwargs_mc2 = {
"x": hidden_states,
"expert_ids": topk_ids,
"expert_shard_type": 0,
"shared_expert_rank_num": 0,
"moe_expert_num": self.moe_expert_num,
"global_bs": self.global_bs,
"expert_token_nums_type": expert_token_nums_type,
}
if self.global_bs == 0:
kwargs_mc2["x_active_mask"] = token_dispatch_input.routing.mc2_mask
stage1_kwargs = {
"scales": None,
"quant_mode": quant_mode,
"group_ep": self.moe_all_to_all_group_name,
"ep_world_size": self.ep_world_size,
"ep_rank_id": self.ep_rank_id,
}
if self.need_extra_args:
stage1_kwargs.update(
{
"group_tp": self.moe_all_to_all_group_name,
"tp_world_size": 1,
"tp_rank_id": 0,
}
)
# Only dispatch-enabled MXFP paths pass y_dtype through MC2.
if (
self.a5_need_extra_args
and (token_dispatch_input.quant.is_mxfp or token_dispatch_input.quant.is_fp8)
and token_dispatch_input.quant.dispatch_with_quant
):
y_dtype = torch.float8_e4m3fn
if (
token_dispatch_input.quant.mxfp is not None
and token_dispatch_input.quant.mxfp.act_quant_type is not None
):
y_dtype = token_dispatch_input.quant.mxfp.act_quant_type
stage1_kwargs.update({"tp_world_size": 1, "tp_rank_id": 0, "y_dtype": y_dtype})
if self.need_expert_scale or self.a5_need_extra_args:
stage1_kwargs.update(
{
"expert_scales": topk_weights.to(torch.float32),
}
)
if self.need_comm_alg:
stage1_kwargs.update({"comm_alg": "hierarchy"})
kwargs_mc2.update(stage1_kwargs)
return kwargs_mc2
def token_dispatch(
self,
token_dispatch_input: MoETokenDispatchInput,
):
kwargs_mc2 = self.get_dispatch_mc2_kwargs(token_dispatch_input)
output = (
torch_npu.npu_moe_distribute_dispatch_v2(**kwargs_mc2)
if self.enable_dispatch_v2
else torch_npu.npu_moe_distribute_dispatch(**kwargs_mc2)
)
# comm_stream.wait_stream(torch.npu.current_stream())
(
expand_x,
dynamic_scale,
assist_info_for_combine,
expert_token_nums,
ep_recv_counts,
tp_recv_counts,
expand_scales,
) = output[0:7]
group_list_type = kwargs_mc2["expert_token_nums_type"]
return MoETokenDispatchOutput(
hidden_states=expand_x,
dynamic_scale=dynamic_scale,
group_list=expert_token_nums,
group_list_type=group_list_type,
combine_metadata=MoEMC2CombineMetadata(
topk_ids=token_dispatch_input.topk_ids,
topk_weights=token_dispatch_input.topk_weights,
expert_map=token_dispatch_input.routing.expert_map,
ep_recv_counts=ep_recv_counts,
tp_recv_counts=tp_recv_counts,
assist_info_for_combine=assist_info_for_combine,
expand_scales=expand_scales,
quant=token_dispatch_input.quant,
mc2_mask=token_dispatch_input.routing.mc2_mask if self.global_bs == 0 else None,
),
)
def get_combine_mc_kwargs(self, hidden_states: torch.Tensor, combine_metadata: MoEMC2CombineMetadata):
expert_map = combine_metadata.expert_map
topk_ids = combine_metadata.topk_ids
topk_weights = combine_metadata.topk_weights
ep_recv_counts = combine_metadata.ep_recv_counts
tp_recv_counts = combine_metadata.tp_recv_counts
assist_info_for_combine = combine_metadata.assist_info_for_combine
expand_scales = combine_metadata.expand_scales
quant_type = combine_metadata.quant.quant_type
comm_quant_mode = combine_metadata.quant.comm_quant_mode
assert expert_map is not None
# NOTE: quant_mode differs by quant features:
# - A5 MXFP communication uses quant_mode=4 only for MXFP8 currently.
if comm_quant_mode is not None:
quant_mode = comm_quant_mode
elif quant_type == QuantType.MXFP8:
quant_mode = 4
else:
quant_mode = 0
kwargs_mc2 = {
"expand_x": hidden_states,
"expert_ids": topk_ids,
"expert_scales": topk_weights.to(torch.float32),
"expert_shard_type": 0,
"shared_expert_rank_num": 0,
"moe_expert_num": self.moe_expert_num,
"global_bs": self.global_bs,
}
if self.global_bs == 0:
kwargs_mc2["x_active_mask"] = combine_metadata.mc2_mask
if combine_metadata.quant.dispatch_with_quant:
tp_recv_counts = torch.empty(1, dtype=torch.int32, device=hidden_states.device)
stage3_kwargs = {
"ep_send_counts": ep_recv_counts,
"group_ep": self.moe_all_to_all_group_name,
"ep_world_size": self.ep_world_size,
"ep_rank_id": self.ep_rank_id,
"expand_scales": expand_scales,
"comm_quant_mode": quant_mode,
}
if self.enable_dispatch_v2:
stage3_kwargs["assist_info_for_combine"] = assist_info_for_combine
else:
stage3_kwargs["expand_idx"] = assist_info_for_combine
if self.need_extra_args:
stage3_kwargs.update(
{
"tp_send_counts": tp_recv_counts,
"group_tp": self.moe_all_to_all_group_name,
"tp_world_size": 1,
"tp_rank_id": 0,
}
)
if self.need_comm_alg:
stage3_kwargs.update({"comm_alg": "hierarchy"})
kwargs_mc2.update(stage3_kwargs)
return kwargs_mc2
def token_combine(self, hidden_states, combine_metadata, bias=None):
assert bias is None, "Bias is not supported in MoEAlltoAllvTokenDispatcher."
kwargs_mc2 = self.get_combine_mc_kwargs(hidden_states, combine_metadata)
combined_output = (
torch_npu.npu_moe_distribute_combine_v2(**kwargs_mc2)
if self.enable_dispatch_v2
else torch_npu.npu_moe_distribute_combine(**kwargs_mc2)
)
return combined_output
class TokenDispatcherWithAllGather(MoETokenDispatcher[MoEAllGatherCombineMetadata]):
def __init__(self, **kwargs):
super().__init__(**kwargs)
self.max_num_tokens = kwargs.get("max_num_tokens")
num_experts_local = kwargs.get("num_local_experts", 0)
self.num_experts_local = (
num_experts_local.item() if torch.is_tensor(num_experts_local) else int(num_experts_local)
)
def token_dispatch(
self,
token_dispatch_input: MoETokenDispatchInput,
):
quant_type = token_dispatch_input.quant.quant_type
dynamic_scale = token_dispatch_input.routing.pertoken_scale
unquantized_mxfp4_dispatch = quant_type == QuantType.MXFP4 and dynamic_scale is None
# Without prepare-stage scales, MXFP4 stays unquantized in dispatch and
# is quantized again inside the MLP path.
with_quant = token_dispatch_input.quant.dispatch_with_quant and quant_type != QuantType.W8A8FP8
with_quant = with_quant and not unquantized_mxfp4_dispatch
is_mxfp = token_dispatch_input.quant.is_mxfp
hidden_states = token_dispatch_input.hidden_states
topk_weights = token_dispatch_input.topk_weights
topk_ids = token_dispatch_input.topk_ids
expert_map = token_dispatch_input.routing.expert_map
act_quant_type = (
token_dispatch_input.quant.mxfp.act_quant_type
if token_dispatch_input.quant.mxfp is not None and not unquantized_mxfp4_dispatch
else None
)
global_redundant_expert_num = token_dispatch_input.routing.global_redundant_expert_num
restore_shape = hidden_states.shape
# Fuse the first dynamic quant of moe_mlp into initrouting when
# dispatch_with_quant is on but got a None dynamic_scale.
if with_quant and dynamic_scale is None:
if quant_type == QuantType.MXFP4:
quant_mode = 9
else:
quant_mode = 3 if is_mxfp else 1
else:
quant_mode = -1
num_tokens = hidden_states.shape[:-1].numel()
apply_router_weight_on_input = token_dispatch_input.routing.apply_router_weight_on_input
if apply_router_weight_on_input:
assert topk_weights.dim() == 2, "`topk_weights` should be in shape (num_tokens, topk)"
_, topk = topk_weights.shape
assert topk == 1, "Only support topk=1 when `apply_router_weight_on_input` is True"
hidden_states = hidden_states * topk_weights.to(hidden_states.dtype)
if expert_map is not None:
global_num_experts = len(expert_map) + global_redundant_expert_num
mask = expert_map[topk_ids] != -1
topk_weights = topk_weights * mask
first_expert_idx = get_ep_group().rank_in_group * self.num_experts_local
last_expert_idx = first_expert_idx + self.num_experts_local
else:
first_expert_idx = 0
last_expert_idx = self.num_experts_local
global_num_experts = self.num_experts_local
sorted_hidden_states, expanded_row_idx, expert_tokens, dynamic_scale = DeviceOperator.npu_moe_init_routing(
hidden_states,
topk_ids,
scale=dynamic_scale,
active_num=num_tokens * self.top_k,
expert_num=global_num_experts,
expert_tokens_num_type=1,
expert_tokens_num_flag=True,
active_expert_range=[first_expert_idx, last_expert_idx],
quant_mode=quant_mode,
act_quant_type=act_quant_type,
)
expert_tokens = expert_tokens.to(torch.int64)
group_list_type = 1 # `count` mode
return MoETokenDispatchOutput(
hidden_states=sorted_hidden_states,
dynamic_scale=dynamic_scale if with_quant else None,
group_list=expert_tokens,
group_list_type=group_list_type,
combine_metadata=MoEAllGatherCombineMetadata(
topk_weights=topk_weights,
expanded_row_idx=expanded_row_idx,
restore_shape=restore_shape,
),
)
def token_combine(self, hidden_states, combine_metadata, bias=None):
final_hidden_states = DeviceOperator.npu_moe_token_unpermute(
permuted_tokens=hidden_states,
sorted_indices=combine_metadata.expanded_row_idx,
probs=combine_metadata.topk_weights,
)
if len(combine_metadata.restore_shape) == 3:
final_hidden_states = final_hidden_states.view(combine_metadata.restore_shape)
# these values are no longer used, so they need to be set to None for memory release.
return final_hidden_states
class TokenDispatcherWithAll2AllV(MoETokenDispatcher[MoEAllToAllCombineMetadata]):
"""
The implementation of the AlltoAll-based token dispatcher, which handles token
dispatching on the sequence level instead of token level. The core of this implementation
lies in each device dispatching on the entire sequence, with the hidden state being partitioned.
"""
def __init__(self, **kwargs):
super().__init__(**kwargs)
self.num_local_experts = kwargs.get("num_local_experts", 0)
assert self.num_local_experts > 0, "Expected at least one expert"
if self.num_local_experts > 1:
self.expert_ids_per_ep_rank = torch.tensor(
[i % self.num_local_experts for i in range(self.num_experts)],
dtype=torch.int32,
device=torch.npu.current_device(),
)
local_expert_indices_offset = self.ep_rank * self.num_local_experts
self.local_expert_indices = [local_expert_indices_offset + i for i in range(self.num_local_experts)]
assert len(self.local_expert_indices) == self.num_local_experts, "Invalid local expert indices"
for i in range(len(self.local_expert_indices) - 1):
assert self.local_expert_indices[i] == self.local_expert_indices[i + 1] - 1, (
"local_expert_indices must be continuous"
)
# TODO: Try local_rank = ep_group.rank_in_group
local_rank = torch.distributed.get_rank(group=self.ep_group)
backend = self.ep_group._get_backend(torch.device("npu"))
self.moe_all_to_all_group_name = backend.get_hccl_comm_name(local_rank)
def token_dispatch(
self,
token_dispatch_input: MoETokenDispatchInput,
):
use_mxfp_quant = token_dispatch_input.quant.is_mxfp
with_quant = token_dispatch_input.quant.dispatch_with_quant
dst_type = token_dispatch_input.quant.get_dst_type
scale_type = token_dispatch_input.quant.get_scale_type
hidden_states = token_dispatch_input.hidden_states
topk_weights = token_dispatch_input.topk_weights
topk_ids = token_dispatch_input.topk_ids
(
permutated_local_input_tokens,
reversed_local_input_permutation_mapping,
tokens_per_expert,
input_splits,
output_splits,
global_input_tokens_local_experts_indices,
hidden_shape,
hidden_shape_before_permute,
) = self._dispatch_preprocess(hidden_states, topk_ids)
dynamic_scale_after_all2all = None
if with_quant:
permutated_local_input_tokens, dynamic_scale = DeviceOperator.npu_dynamic_quant(
permutated_local_input_tokens, act_quant_type=dst_type, use_mxfp_quant=use_mxfp_quant
)
_, dynamic_scale_after_all2all, permute2_ep_all_to_all_handle = async_all_to_all(
dynamic_scale, output_splits, input_splits, self.ep_group
)
permute2_ep_all_to_all_handle.wait()
dynamic_scale.untyped_storage().resize_(0)
_, global_input_tokens, permute1_ep_all_to_all_handle = async_all_to_all(
permutated_local_input_tokens, output_splits, input_splits, self.ep_group
)
permute1_ep_all_to_all_handle.wait()
permutated_local_input_tokens.untyped_storage().resize_(0)
# Postprocess
global_input_tokens, dynamic_scale_final, reversed_global_input_permutation_mapping = (
self._dispatch_postprocess(
global_input_tokens,
dynamic_scale_after_all2all,
global_input_tokens_local_experts_indices,
with_quant,
dst_type,
scale_type,
)
)
return MoETokenDispatchOutput(
hidden_states=global_input_tokens,
dynamic_scale=dynamic_scale_final,
group_list=tokens_per_expert,
group_list_type=1,
combine_metadata=MoEAllToAllCombineMetadata(
input_splits=input_splits,
output_splits=output_splits,
topk_weights=topk_weights,
reversed_local_input_permutation_mapping=reversed_local_input_permutation_mapping,
reversed_global_input_permutation_mapping=reversed_global_input_permutation_mapping,
hidden_shape=hidden_shape,
hidden_shape_before_permute=hidden_shape_before_permute,
),
)
def token_combine(self, hidden_states, combine_metadata, bias=None):
assert bias is None, "Bias is not supported in MoEAlltoAllvTokenDispatcher."
# 1. Preprocess using metadata
hidden_states = self._combine_preprocess(hidden_states, combine_metadata)
# 2. AllToAll
_, permutated_local_input_tokens, handle = async_all_to_all(
hidden_states,
combine_metadata.input_splits,
combine_metadata.output_splits,
self.ep_group,
)
handle.wait()
hidden_states.untyped_storage().resize_(0)
# 3. Postprocess using metadata
output = self._combine_postprocess(permutated_local_input_tokens, combine_metadata)
return output
def _dispatch_preprocess(self, hidden_states, topk_ids):
hidden_shape = hidden_states.shape
hidden_states = hidden_states.view(-1, hidden_states.size(-1))
(
tokens_per_expert,
input_splits,
output_splits,
global_input_tokens_local_experts_indices,
num_out_tokens,
) = self._preprocess(topk_ids)
hidden_shape_before_permute = hidden_states.shape
permutated_local_input_tokens, reversed_local_input_permutation_mapping = torch_npu.npu_moe_token_permute(
tokens=hidden_states,
indices=topk_ids,
num_out_tokens=num_out_tokens,
)
return (
permutated_local_input_tokens,
reversed_local_input_permutation_mapping,
tokens_per_expert,
input_splits,
output_splits,
global_input_tokens_local_experts_indices,
hidden_shape,
hidden_shape_before_permute,
)
def _preprocess(self, topk_ids: torch.Tensor):
num_local_tokens_per_expert = torch.histc(topk_ids, bins=self.num_experts, min=0, max=self.num_experts)
ep_size = self.ep_size
num_out_tokens = topk_ids.numel()
input_splits = (
num_local_tokens_per_expert.reshape(ep_size, self.num_local_experts)
.sum(axis=1)
.to(torch.device("cpu"), non_blocking=True)
.numpy()
)
num_global_tokens_per_expert = gather_from_sequence_parallel_region(
num_local_tokens_per_expert, group=self.ep_group
).reshape(ep_size, self.num_experts)
num_global_tokens_per_local_expert = num_global_tokens_per_expert[
:, self.local_expert_indices[0] : self.local_expert_indices[-1] + 1
]
if num_global_tokens_per_local_expert is None:
raise ValueError("num_global_tokens_per_local_expert must be set before sum.")
output_splits = (
num_global_tokens_per_local_expert.sum(axis=-1).to(torch.device("cpu"), non_blocking=True).numpy()
)
num_tokens_per_local_expert = num_global_tokens_per_local_expert.sum(axis=0)
global_input_tokens_local_experts_indices = None
if self.num_local_experts > 1:
if num_global_tokens_per_local_expert is None:
raise ValueError("num_global_tokens_per_local_expert must be set before operations.")
global_input_tokens_local_experts_indices = torch.repeat_interleave(
self.expert_ids_per_ep_rank, num_global_tokens_per_local_expert.ravel()
)
else:
torch.npu.synchronize()
return (
num_tokens_per_local_expert,
input_splits,
output_splits,
global_input_tokens_local_experts_indices,
num_out_tokens,
)
def _dispatch_postprocess(
self,
global_input_tokens,
dynamic_scale_after_all2all,
global_input_tokens_local_experts_indices,
with_quant,
dst_type,
scale_type,
):
# Early return if no local experts or no tokens
if self.num_local_experts <= 1:
return global_input_tokens, dynamic_scale_after_all2all, None
assert global_input_tokens_local_experts_indices is not None, (
"global_input_tokens_local_experts_indices must be provided"
)
if with_quant:
if scale_type == torch.float8_e8m0fnu:
experts_indices_2d_copy = global_input_tokens_local_experts_indices.reshape(
global_input_tokens_local_experts_indices.shape[0], 1
)
dynamic_scale_for_routing = dynamic_scale_after_all2all.view(torch.float8_e8m0fnu)
global_input_tokens, reversed_global_input_permutation_mapping, _, routed_scale = (
torch_npu.npu_moe_init_routing_v2(
global_input_tokens,
experts_indices_2d_copy,
scale=dynamic_scale_for_routing,
active_num=experts_indices_2d_copy.shape[0],
expert_num=self.num_local_experts,
expert_tokens_num_type=1,
expert_tokens_num_flag=True,
active_expert_range=[0, self.num_local_experts],
x_dtype=dst_type,
)
)
dynamic_scale_after_all2all = routed_scale.view(torch.uint8)
experts_indices_2d_copy.untyped_storage().resize_(0)
return global_input_tokens, dynamic_scale_after_all2all, reversed_global_input_permutation_mapping
dynamic_scale_after_all2all, _ = torch_npu.npu_moe_token_permute(
dynamic_scale_after_all2all.unsqueeze(-1), global_input_tokens_local_experts_indices
)
dynamic_scale_after_all2all = dynamic_scale_after_all2all.squeeze(-1)
# Non-quantized case
global_input_tokens, reversed_global_input_permutation_mapping = torch_npu.npu_moe_token_permute(
global_input_tokens, global_input_tokens_local_experts_indices
)
return global_input_tokens, dynamic_scale_after_all2all, reversed_global_input_permutation_mapping
def _combine_preprocess(
self, hidden_states: torch.Tensor, combine_metadata: MoEAllToAllCombineMetadata
) -> torch.Tensor:
# Unpermutation 2: expert output to AlltoAll input
rev_global = combine_metadata.reversed_global_input_permutation_mapping
if hidden_states.shape[0] > 0 and self.num_local_experts > 1 and rev_global is not None:
hidden_states = torch_npu.npu_moe_token_unpermute(hidden_states, rev_global)
return hidden_states
def _combine_postprocess(
self,
permutated_local_input_tokens: torch.Tensor,
combine_metadata: MoEAllToAllCombineMetadata,
) -> torch.Tensor:
# Unpermutation 1: AlltoAll output to output
output = torch_npu.npu_moe_token_unpermute(
permuted_tokens=permutated_local_input_tokens,
sorted_indices=combine_metadata.reversed_local_input_permutation_mapping.to(torch.int32),
probs=combine_metadata.topk_weights,
restore_shape=combine_metadata.hidden_shape_before_permute,
)
output = output.view(combine_metadata.hidden_shape)
return output