0
vllm_ascend/ops/fused_moe/__init__.py
Normal file
0
vllm_ascend/ops/fused_moe/__init__.py
Normal file
104
vllm_ascend/ops/fused_moe/comm_utils.py
Normal file
104
vllm_ascend/ops/fused_moe/comm_utils.py
Normal 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)
|
||||
417
vllm_ascend/ops/fused_moe/experts_selector.py
Normal file
417
vllm_ascend/ops/fused_moe/experts_selector.py
Normal 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
|
||||
894
vllm_ascend/ops/fused_moe/fused_moe.py
Normal file
894
vllm_ascend/ops/fused_moe/fused_moe.py
Normal 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)
|
||||
722
vllm_ascend/ops/fused_moe/fused_moe_0_23_0.py
Normal file
722
vllm_ascend/ops/fused_moe/fused_moe_0_23_0.py
Normal 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"]
|
||||
63
vllm_ascend/ops/fused_moe/gate_linear.py
Normal file
63
vllm_ascend/ops/fused_moe/gate_linear.py
Normal 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
|
||||
349
vllm_ascend/ops/fused_moe/moe_comm_method.py
Normal file
349
vllm_ascend/ops/fused_moe/moe_comm_method.py
Normal 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
|
||||
)
|
||||
552
vllm_ascend/ops/fused_moe/moe_mlp.py
Normal file
552
vllm_ascend/ops/fused_moe/moe_mlp.py
Normal 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,
|
||||
)
|
||||
270
vllm_ascend/ops/fused_moe/moe_runtime_args.py
Normal file
270
vllm_ascend/ops/fused_moe/moe_runtime_args.py
Normal 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",
|
||||
]
|
||||
165
vllm_ascend/ops/fused_moe/moe_stage_contracts.py
Normal file
165
vllm_ascend/ops/fused_moe/moe_stage_contracts.py
Normal 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",
|
||||
]
|
||||
127
vllm_ascend/ops/fused_moe/moe_stage_params.py
Normal file
127
vllm_ascend/ops/fused_moe/moe_stage_params.py
Normal 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",
|
||||
]
|
||||
548
vllm_ascend/ops/fused_moe/prepare_finalize.py
Normal file
548
vllm_ascend/ops/fused_moe/prepare_finalize.py
Normal 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
|
||||
709
vllm_ascend/ops/fused_moe/token_dispatcher.py
Normal file
709
vllm_ascend/ops/fused_moe/token_dispatcher.py
Normal 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
|
||||
Reference in New Issue
Block a user