init v0.23.0

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

0
tests/ut/ops/__init__.py Normal file
View File

View File

View File

@@ -0,0 +1,439 @@
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
# This file is a part of the vllm-ascend project.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
import pytest
import torch
from vllm_ascend.ops.triton.fla import chunk, chunk_o, chunk_o_update
from vllm_ascend.utils import enable_custom_op
enable_custom_op()
class _FakeKernel:
def __init__(self):
self.grid = None
self.grid_result = None
self.launch_kwargs: dict[str, object] | None = None
def __getitem__(self, grid):
self.grid = grid
self.grid_result = grid({"BV": 128})
def launch(**kwargs):
self.launch_kwargs = kwargs
return launch
class _DummyTensor:
def __init__(self, name: str):
self.name = name
self.shape = (1,)
self.dtype = torch.float32
def unsqueeze(self, dim: int):
return self
def new_empty(self, *shape):
return _DummyTensor(f"{self.name}.new_empty")
def __getitem__(self, item):
return self
def __setitem__(self, item, value):
return None
def __add__(self, other):
return self
def __sub__(self, other):
return self
def transpose(self, dim0, dim1):
return self
def contiguous(self):
return self
def to(self, *args, **kwargs):
return self
class _GatherResult:
def __init__(self, items):
self.items = items
def __getitem__(self, item):
if isinstance(item, tuple):
item = item[0]
return self.items[item]
def _patch_missing_cdiv(monkeypatch: pytest.MonkeyPatch, module) -> None:
if hasattr(module.triton, "cdiv"):
return
monkeypatch.setattr(
module.triton,
"cdiv",
lambda x, y: (x + y - 1) // y,
raising=False,
)
@pytest.mark.parametrize("target", ["chunk_o", "chunk_o_update"])
def test_chunk_leaf_wrappers_use_prebuilt_chunk_offsets(
monkeypatch: pytest.MonkeyPatch,
target: str,
):
fake_kernel = _FakeKernel()
sentinel = torch.tensor([0, 2, 5], dtype=torch.int32)
cu_seqlens = torch.tensor([0, 4, 7], dtype=torch.int32)
if target == "chunk_o":
_patch_missing_cdiv(monkeypatch, chunk_o)
monkeypatch.setattr(chunk_o, "chunk_fwd_kernel_o", fake_kernel)
monkeypatch.setattr(
chunk_o,
"prepare_chunk_offsets",
lambda *args, **kwargs: pytest.fail("prepare_chunk_offsets should not be called"),
)
chunk_o.chunk_fwd_o(
q=torch.zeros((2, 4, 1, 8), dtype=torch.float32),
k=torch.zeros((2, 4, 1, 8), dtype=torch.float32),
v=torch.zeros((2, 4, 1, 16), dtype=torch.float32),
h=torch.zeros((4, 1, 8, 16), dtype=torch.float32),
g=torch.zeros((2, 4, 1), dtype=torch.float32),
cu_seqlens=cu_seqlens,
chunk_offsets=sentinel,
)
else:
_patch_missing_cdiv(monkeypatch, chunk_o_update)
monkeypatch.setattr(chunk_o_update, "chunk_fwd_kernel_o_update", fake_kernel)
monkeypatch.setattr(
chunk_o_update,
"prepare_chunk_offsets",
lambda *args, **kwargs: pytest.fail("prepare_chunk_offsets should not be called"),
)
chunk_o_update.chunk_fwd_o_update(
q=torch.zeros((2, 4, 1, 8), dtype=torch.float32),
v=torch.zeros((2, 4, 1, 16), dtype=torch.float32),
h=torch.zeros((4, 1, 8, 16), dtype=torch.float32),
h_update=torch.zeros((5, 1, 8, 8), dtype=torch.float32),
updated_h_state=torch.zeros((1, 8, 16), dtype=torch.float32),
cu_seqlens=cu_seqlens,
chunk_offsets=sentinel,
)
assert fake_kernel.launch_kwargs is not None
assert fake_kernel.launch_kwargs["chunk_offsets"] is sentinel
def test_chunk_gated_delta_rule_fwd_threads_prebuilt_chunk_offsets(
monkeypatch: pytest.MonkeyPatch,
):
chunk_offsets = torch.tensor([0, 2, 5], dtype=torch.int32)
update_chunk_offsets = torch.tensor([0, 3, 7], dtype=torch.int32)
final_chunk_indices = torch.tensor([1, 3], dtype=torch.int32)
prebuilt_meta = type(
"PrebuiltMeta",
(),
{
"block_indices_cumsum": None,
"cu_seqlens_host": (0, 4, 7),
"chunk_indices_chunk64_host": (0, 0, 1, 0),
"chunk_indices_chunk64": None,
"chunk_offsets_chunk64": chunk_offsets,
"update_chunk_offsets_chunk64": update_chunk_offsets,
"final_chunk_indices_chunk64": final_chunk_indices,
"chunk_indices_large_block": None,
"keep_meta": None,
"cu_seqlens_kern": None,
},
)()
q = _DummyTensor("q")
k = _DummyTensor("k")
v = _DummyTensor("v")
g = _DummyTensor("g")
beta = _DummyTensor("beta")
initial_state = _DummyTensor("initial_state")
non_pcp_calls: list[tuple[str, object]] = []
pcp_calls: list[tuple[str, object]] = []
def run_case(world_size: int, calls: list[tuple[str, object]]):
group = type(
"Group",
(),
{
"world_size": world_size,
"rank_in_group": 0,
"all_gather": lambda self, value, dim: _GatherResult([_DummyTensor("g0"), _DummyTensor("g1")]),
},
)()
monkeypatch.setattr(chunk, "get_forward_context", lambda: type("Ctx", (), {"attn_metadata": None})())
monkeypatch.setattr(chunk, "get_pcp_group", lambda: group)
monkeypatch.setattr(chunk, "chunk_local_cumsum", lambda *args, **kwargs: _DummyTensor("g_cumsum"))
monkeypatch.setattr(chunk, "chunk_scaled_dot_kkt_fwd", lambda *args, **kwargs: _DummyTensor("A"))
monkeypatch.setattr(chunk, "solve_tril", lambda *args, **kwargs: _DummyTensor("A_solved"))
monkeypatch.setattr(chunk, "recompute_w_u_fwd", lambda *args, **kwargs: (_DummyTensor("w"), _DummyTensor("u")))
monkeypatch.setattr(
chunk,
"chunk_gated_delta_rule_fwd_h",
lambda *args, **kwargs: (_DummyTensor("h"), _DummyTensor("v_new"), _DummyTensor("final_state")),
)
monkeypatch.setattr(
chunk,
"chunk_gated_delta_rule_fwd_hupdate",
lambda *args, **kwargs: _DummyTensor("h_update"),
)
monkeypatch.setattr(
chunk.torch,
"matmul",
lambda *args, **kwargs: _DummyTensor("matmul"),
raising=False,
)
monkeypatch.setattr(
chunk.torch,
"zeros_like",
lambda *args, **kwargs: _DummyTensor("zeros_like"),
raising=False,
)
monkeypatch.setattr(
torch.ops._C_ascend,
"chunk_gated_delta_rule_fwd_h",
lambda *args, **kwargs: (_DummyTensor("h"), _DummyTensor("v_new"), _DummyTensor("final_state")),
raising=False,
)
monkeypatch.setattr(
torch.ops._C_ascend,
"chunk_fwd_o",
lambda *args, **kwargs: _DummyTensor("o_ascend"),
raising=False,
)
def fake_chunk_fwd_o(*args, **kwargs):
calls.append(("o", kwargs["chunk_offsets"]))
return _DummyTensor("o")
def fake_chunk_fwd_o_update(*args, **kwargs):
calls.append(("o_update", kwargs["chunk_offsets"]))
return _DummyTensor("h_updated")
monkeypatch.setattr(chunk, "chunk_fwd_o", fake_chunk_fwd_o)
if world_size > 1:
monkeypatch.setattr(chunk, "chunk_gated_delta_rule_fwd_hupdate", fake_chunk_fwd_o_update)
chunk.chunk_gated_delta_rule_fwd(
q=q,
k=k,
v=v,
g=g,
beta=beta,
scale=1.0,
initial_state=initial_state,
output_final_state=False,
cu_seqlens=torch.tensor([0, 4, 7], dtype=torch.int32),
prebuilt_meta=prebuilt_meta,
)
run_case(1, non_pcp_calls)
assert non_pcp_calls == []
run_case(2, pcp_calls)
assert pcp_calls == [("o_update", chunk_offsets)]
def test_chunk_gated_delta_rule_fwd_uses_prebuilt_metadata_without_runtime_tolist(
monkeypatch: pytest.MonkeyPatch,
):
prebuilt_meta = type(
"PrebuiltMeta",
(),
{
"block_indices_cumsum": None,
"cu_seqlens_host": (0, 4, 7),
"chunk_indices_chunk64_host": (0, 0, 1, 0),
"chunk_indices_chunk64": torch.tensor([[0, 0], [1, 0]], dtype=torch.int32),
"chunk_offsets_chunk64": torch.tensor([0, 1, 2], dtype=torch.int32),
"update_chunk_offsets_chunk64": torch.tensor([0, 2, 4], dtype=torch.int32),
"final_chunk_indices_chunk64": torch.tensor([1, 3], dtype=torch.int32),
"chunk_indices_large_block": None,
"keep_meta": None,
"cu_seqlens_kern": None,
},
)()
q = _DummyTensor("q")
k = _DummyTensor("k")
v = _DummyTensor("v")
g = _DummyTensor("g")
beta = _DummyTensor("beta")
initial_state = _DummyTensor("initial_state")
captured: dict[str, tuple[int, ...] | None] = {}
monkeypatch.setattr(chunk, "get_forward_context", lambda: type("Ctx", (), {"attn_metadata": None})())
monkeypatch.setattr(
chunk,
"get_pcp_group",
lambda: type("Group", (), {"world_size": 1, "rank_in_group": 0})(),
)
monkeypatch.setattr(chunk, "chunk_local_cumsum", lambda *args, **kwargs: _DummyTensor("g_cumsum"))
monkeypatch.setattr(chunk, "chunk_scaled_dot_kkt_fwd", lambda *args, **kwargs: _DummyTensor("A"))
monkeypatch.setattr(chunk, "solve_tril", lambda *args, **kwargs: _DummyTensor("A_solved"))
monkeypatch.setattr(chunk, "recompute_w_u_fwd", lambda *args, **kwargs: (_DummyTensor("w"), _DummyTensor("u")))
monkeypatch.setattr(
torch.ops._C_ascend,
"chunk_gated_delta_rule_fwd_h",
lambda *args, **kwargs: (
captured.update(
{
"cu_seqlens": kwargs["cu_seqlens"],
"chunk_indices": kwargs["chunk_indices"],
}
)
or (_DummyTensor("h"), _DummyTensor("v_new"), _DummyTensor("final_state"))
),
raising=False,
)
monkeypatch.setattr(
torch.ops._C_ascend,
"chunk_fwd_o",
lambda *args, **kwargs: _DummyTensor("o_ascend"),
raising=False,
)
monkeypatch.setattr(
torch.Tensor,
"tolist",
lambda self: pytest.fail("runtime should not convert device tensors to host tuples"),
)
chunk.chunk_gated_delta_rule_fwd(
q=q,
k=k,
v=v,
g=g,
beta=beta,
scale=1.0,
initial_state=initial_state,
output_final_state=False,
cu_seqlens=torch.tensor([0, 4, 7], dtype=torch.int32),
prebuilt_meta=prebuilt_meta,
)
assert captured["cu_seqlens"] == prebuilt_meta.cu_seqlens_host
assert captured["chunk_indices"] == prebuilt_meta.chunk_indices_chunk64_host
def test_chunk_gated_delta_rule_fwd_pcp_chaining_subtracts_initial_state(
monkeypatch: pytest.MonkeyPatch,
):
"""PCP chaining uses (updated_state[i-1] - initial_state), not updated_state[i-1].
With s0 != 0 (subsequent prefill chunk), the fix subtracts s0 to avoid
double-counting Φ_i·s0. Verified by checking the returned final_state
matches the sequential result Φ_1·(Φ_0·s0+p_0)+p_1.
"""
torch.manual_seed(42)
N, H, K, V = 1, 2, 4, 4
s0 = torch.randn(N, H, K, V)
phi_0 = torch.randn(N, H, K, K)
phi_1 = torch.randn(N, H, K, K)
p_0 = torch.randn(N, H, K, V)
p_1 = torch.randn(N, H, K, V)
# Each rank computes final_state = Φ_i · s0 + p_i (from shared s0)
rank0_fs = torch.matmul(phi_0, s0) + p_0
rank1_fs = torch.matmul(phi_1, s0) + p_1
# h_update shape [1, N, H, K, K]; after [:, [0], :, :, :] → [1, N, H, K, K]
h_update_tensor = phi_0.unsqueeze(0)
prebuilt_meta = type(
"PrebuiltMeta",
(),
{
"block_indices_cumsum": None,
"cu_seqlens_host": (0, N),
"chunk_indices_chunk64_host": (0, 0),
"chunk_indices_chunk64": None,
"chunk_offsets_chunk64": torch.tensor([0, 1], dtype=torch.int32),
"update_chunk_offsets_chunk64": torch.tensor([0, 2], dtype=torch.int32),
"final_chunk_indices_chunk64": torch.tensor([0], dtype=torch.int32),
"chunk_indices_large_block": None,
"num_decodes": 0,
"keep_meta": None,
"cu_seqlens_kern": None,
},
)()
all_gather_returns = [
torch.stack([rank0_fs, rank1_fs]), # all_final_state: [2, N, H, K, V]
torch.stack([phi_0, phi_1]), # all_final_h_update: [2, N, H, K, K]
]
group = type(
"Group",
(),
{
"world_size": 2,
"rank_in_group": 0,
"all_gather": lambda self, value, dim: all_gather_returns.pop(0),
},
)()
monkeypatch.setattr(chunk, "get_forward_context", lambda: type("Ctx", (), {"attn_metadata": None})())
monkeypatch.setattr(chunk, "get_pcp_group", lambda: group)
monkeypatch.setattr(chunk, "chunk_local_cumsum", lambda *a, **kw: _DummyTensor("g_cumsum"))
monkeypatch.setattr(chunk, "chunk_scaled_dot_kkt_fwd", lambda *a, **kw: _DummyTensor("A"))
monkeypatch.setattr(chunk, "solve_tril", lambda *a, **kw: _DummyTensor("A_solved"))
monkeypatch.setattr(chunk, "recompute_w_u_fwd", lambda *a, **kw: (_DummyTensor("w"), _DummyTensor("u")))
monkeypatch.setattr(
torch.ops._C_ascend,
"chunk_gated_delta_rule_fwd_h",
lambda *a, **kw: (_DummyTensor("h"), _DummyTensor("v_new"), rank0_fs),
raising=False,
)
monkeypatch.setattr(
chunk,
"chunk_gated_delta_rule_fwd_hupdate",
lambda *a, **kw: h_update_tensor,
)
monkeypatch.setattr(
torch.ops._C_ascend,
"chunk_fwd_o",
lambda *a, **kw: _DummyTensor("o_ascendc"),
raising=False,
)
result = chunk.chunk_gated_delta_rule_fwd(
q=_DummyTensor("q"),
k=_DummyTensor("k"),
v=_DummyTensor("v"),
g=_DummyTensor("g"),
beta=_DummyTensor("beta"),
scale=1.0,
initial_state=s0,
output_final_state=False,
cu_seqlens=torch.tensor([0, N], dtype=torch.int32),
prebuilt_meta=prebuilt_meta,
)
final_state = result[3]
# Sequential: Φ_1·(Φ_0·s0 + p_0) + p_1
expected = torch.matmul(phi_1, torch.matmul(phi_0, s0) + p_0) + p_1
torch.testing.assert_close(final_state, expected, rtol=1e-4, atol=1e-4)

View File

@@ -0,0 +1,827 @@
#
# 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 unittest.mock import MagicMock, PropertyMock, patch
import numpy as np
import pytest
import torch
from tests.ut.base import TestBase
from vllm_ascend.ops.fused_moe.moe_runtime_args import (
MoEAllGatherCombineMetadata,
MoEAllToAllCombineMetadata,
MoEMC2CombineMetadata,
MoEQuantParams,
MoERoutingParams,
MoETokenDispatchInput,
)
from vllm_ascend.ops.fused_moe.token_dispatcher import ( # isort: skip
AscendDeviceType,
EXPERT_TOKEN_NUMS_TYPE_COUNT,
EXPERT_TOKEN_NUMS_TYPE_CUMSUM,
TokenDispatcherWithAll2AllV,
TokenDispatcherWithAllGather,
TokenDispatcherWithMC2,
)
from vllm_ascend.ops.fused_moe.moe_stage_params import MoEMxfpParams
from vllm_ascend.quantization.quant_type import QuantType
MXFP4_TEST_DTYPE = getattr(torch, "float4_e2m1fn_x2", torch.float16)
def build_token_dispatch_input_fixture(
*,
hidden_states: torch.Tensor,
topk_weights: torch.Tensor,
topk_ids: torch.Tensor,
expert_map: torch.Tensor | None = None,
global_redundant_expert_num: int = 0,
apply_router_weight_on_input: bool = False,
pertoken_scale: torch.Tensor | None = None,
quant_type: QuantType = QuantType.NONE,
comm_quant_mode: int | None = None,
act_quant_type: torch.dtype | None = None,
is_per_channel_weight: bool = False,
mc2_mask: torch.Tensor | None = None,
) -> MoETokenDispatchInput:
mxfp_spec = None
if quant_type in (QuantType.MXFP8, QuantType.MXFP4):
mxfp_spec = MoEMxfpParams(act_quant_type=act_quant_type)
return MoETokenDispatchInput(
hidden_states=hidden_states,
topk_weights=topk_weights,
topk_ids=topk_ids,
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,
pertoken_scale=pertoken_scale,
),
quant=MoEQuantParams(
quant_type=quant_type,
comm_quant_mode=comm_quant_mode,
mxfp=mxfp_spec,
is_per_channel_weight=is_per_channel_weight,
),
)
class TestTokenDispatcherWithMC2(TestBase):
def setUp(self):
self.config_patcher = patch("vllm_ascend.ops.fused_moe.token_dispatcher.get_current_vllm_config")
self.mock_get_config = self.config_patcher.start()
mock_config = MagicMock()
mock_config.scheduler_config.max_num_seqs = 256
mock_config.compilation_config.custom_ops = ["all"]
mock_config.speculative_config = None
mock_config.parallel_config.tensor_parallel_size = 1
self.mock_get_config.return_value = mock_config
self.mc2_tokens_capacity = 128
self.mc2_capacity_patch = patch(
"vllm_ascend.ops.fused_moe.token_dispatcher.get_mc2_tokens_capacity",
return_value=self.mc2_tokens_capacity,
)
self.mock_get_mc2_tokens_capacity = self.mc2_capacity_patch.start()
self.mc2_group = MagicMock()
self.mc2_group.device_group.return_value._get_backend.return_value.get_hccl_comm_name.return_value = "hccl_123"
self.mc2_group.rank_in_group = 0
self.mc2_group.world_size = 8
self.mc2_group_patch = patch(
"vllm_ascend.ops.fused_moe.token_dispatcher.get_mc2_group", return_value=self.mc2_group
)
self.mc2_group_patch.start()
self.rank_group_patch = patch("torch.distributed.get_rank", return_value=0)
self.rank_group_patch.start()
# Mock get_forward_context().mc2_mask
self.forward_context = MagicMock()
self.forward_context.mc2_mask = torch.tensor([1, 0, 1])
self.forward_context_patch = patch(
"vllm.forward_context.get_forward_context", return_value=self.forward_context
)
self.forward_context_patch.start()
# Mock get_ascend_device_type()
self.ascend_soc_version_patch = patch(
"vllm_ascend.ops.fused_moe.token_dispatcher.get_ascend_device_type", return_value=AscendDeviceType.A3
)
self.ascend_soc_version_patch.start()
# Mock get_ascend_config() and is_hierarchical_communication_enabled()
mock_ascend_config = MagicMock()
mock_ascend_config.enable_mc2_hierarchy_comm = False
mock_ascend_config.eplb_config = MagicMock()
mock_ascend_config.eplb_config.dynamic_eplb = False
self.ascend_config_patch = patch(
"vllm_ascend.ops.fused_moe.token_dispatcher.get_ascend_config", return_value=mock_ascend_config
)
self.ascend_config_patch.start()
self.ascend_config_utils_patch = patch("vllm_ascend.utils.get_ascend_config", return_value=mock_ascend_config)
self.ascend_config_utils_patch.start()
self.hier_comm_patch = patch(
"vllm_ascend.ops.fused_moe.token_dispatcher.is_hierarchical_communication_enabled", return_value=False
)
self.hier_comm_patch.start()
self.skip_allreduce_patch = patch(
"vllm_ascend.ops.fused_moe.token_dispatcher.should_skip_allreduce_across_dp_group", return_value=False
)
self.mock_skip_allreduce = self.skip_allreduce_patch.start()
kwargs = {"with_quant": False, "top_k": 8, "num_experts": 128}
self.dispatcher = TokenDispatcherWithMC2(**kwargs)
def tearDown(self):
self.config_patcher.stop()
self.mc2_capacity_patch.stop()
self.mc2_group_patch.stop()
self.rank_group_patch.stop()
self.forward_context_patch.stop()
self.ascend_soc_version_patch.stop()
self.ascend_config_patch.stop()
self.ascend_config_utils_patch.stop()
self.hier_comm_patch.stop()
self.skip_allreduce_patch.stop()
def test_init(self):
self.assertEqual(self.dispatcher.ep_rank_id, 0)
self.assertEqual(self.dispatcher.ep_world_size, 8)
self.assertTrue(self.dispatcher.enable_dispatch_v2)
self.assertTrue(self.dispatcher.need_extra_args)
self.assertEqual(self.dispatcher.global_bs, 0)
def test_init_uses_mc2_capacity_for_non_uniform_global_bs(self):
self.mock_get_config.return_value.parallel_config.tensor_parallel_size = 4
self.mock_skip_allreduce.return_value = True
dispatcher = TokenDispatcherWithMC2(with_quant=False, top_k=8, num_experts=128)
self.assertEqual(dispatcher.global_bs, 256)
def test_get_dispatch_mc2_kwargs_with_skip_allreduce_omits_mc2_mask(self):
self.mock_get_config.return_value.parallel_config.tensor_parallel_size = 4
self.mock_skip_allreduce.return_value = True
dispatcher = TokenDispatcherWithMC2(with_quant=False, top_k=8, num_experts=128)
hidden_states = torch.randn(10, 128)
topk_ids = torch.randint(0, 8, (10, 1))
topk_weights = torch.randn(10, 1)
expert_map = torch.tensor([0, 1, 2, 3, 4, 5, 6, 7])
mc2_mask = torch.tensor([True, False, True, False])
token_dispatch_input = build_token_dispatch_input_fixture(
hidden_states=hidden_states,
topk_weights=topk_weights,
topk_ids=topk_ids,
expert_map=expert_map,
mc2_mask=mc2_mask,
)
kwargs = dispatcher.get_dispatch_mc2_kwargs(token_dispatch_input)
self.assertEqual(kwargs["global_bs"], 256)
self.assertNotIn("x_active_mask", kwargs)
def test_get_dispatch_mc2_kwargs_without_skip_allreduce_keeps_mc2_mask(self):
hidden_states = torch.randn(10, 128)
topk_ids = torch.randint(0, 8, (10, 1))
topk_weights = torch.randn(10, 1)
expert_map = torch.tensor([0, 1, 2, 3, 4, 5, 6, 7])
mc2_mask = torch.tensor([True, False, True, False])
token_dispatch_input = build_token_dispatch_input_fixture(
hidden_states=hidden_states,
topk_weights=topk_weights,
topk_ids=topk_ids,
expert_map=expert_map,
mc2_mask=mc2_mask,
)
kwargs = self.dispatcher.get_dispatch_mc2_kwargs(token_dispatch_input)
self.assertEqual(kwargs["global_bs"], 0)
self.assertIs(kwargs["x_active_mask"], mc2_mask)
def test_get_dispatch_mc2_kwargs_without_quant(self):
hidden_states = torch.randn(10, 128)
topk_ids = torch.randint(0, 8, (10, 1))
topk_weights = torch.randn(10, 1)
expert_map = torch.tensor([0, 1, 2, 3, 4, 5, 6, 7])
token_dispatch_input = build_token_dispatch_input_fixture(
hidden_states=hidden_states,
topk_weights=topk_weights,
topk_ids=topk_ids,
expert_map=expert_map,
global_redundant_expert_num=0,
apply_router_weight_on_input=False,
pertoken_scale=None,
)
kwargs = self.dispatcher.get_dispatch_mc2_kwargs(token_dispatch_input)
self.assertIn("x", kwargs)
self.assertIn("expert_ids", kwargs)
self.assertEqual(kwargs["moe_expert_num"], 8)
def test_token_permutation_dispatch(self):
hidden_states = torch.randn(10, 128)
topk_weights = torch.randn(10, 1)
topk_ids = torch.randint(0, 8, (10, 1))
expert_map = torch.tensor([0, 1, 2, 3, 4, 5, 6, 7])
with patch(
"torch_npu.npu_moe_distribute_dispatch_v2", return_value=(torch.randn(10, 128),) * 5 + (None, None)
) as mock_dispatch:
token_dispatch_input = build_token_dispatch_input_fixture(
hidden_states=hidden_states,
topk_weights=topk_weights,
topk_ids=topk_ids,
expert_map=expert_map,
)
output = self.dispatcher.token_dispatch(token_dispatch_input=token_dispatch_input)
mock_dispatch.assert_called_once()
self.assertEqual(output.group_list_type, 0) # group_list_type == 0
self.assertIsInstance(output.combine_metadata, MoEMC2CombineMetadata)
def test_w4a8_per_channel_dispatch_uses_count_group_list(self):
hidden_states = torch.randn(10, 128)
topk_weights = torch.randn(10, 1)
topk_ids = torch.randint(0, 8, (10, 1))
expert_map = torch.tensor([0, 1, 2, 3, 4, 5, 6, 7])
self.dispatcher.enable_dispatch_v2 = True
token_dispatch_input = build_token_dispatch_input_fixture(
hidden_states=hidden_states,
topk_weights=topk_weights,
topk_ids=topk_ids,
expert_map=expert_map,
quant_type=QuantType.W4A8,
is_per_channel_weight=True,
)
with patch(
"torch_npu.npu_moe_distribute_dispatch_v2", return_value=(torch.randn(10, 128),) * 5 + (None, None)
) as mock_dispatch:
output = self.dispatcher.token_dispatch(token_dispatch_input=token_dispatch_input)
mock_dispatch.assert_called_once()
self.assertEqual(mock_dispatch.call_args.kwargs["expert_token_nums_type"], EXPERT_TOKEN_NUMS_TYPE_COUNT)
self.assertEqual(output.group_list_type, EXPERT_TOKEN_NUMS_TYPE_COUNT)
def test_w4a8_group_dispatch_keeps_prefix_sum_group_list(self):
hidden_states = torch.randn(10, 128)
topk_weights = torch.randn(10, 1)
topk_ids = torch.randint(0, 8, (10, 1))
expert_map = torch.tensor([0, 1, 2, 3, 4, 5, 6, 7])
token_dispatch_input = build_token_dispatch_input_fixture(
hidden_states=hidden_states,
topk_weights=topk_weights,
topk_ids=topk_ids,
expert_map=expert_map,
quant_type=QuantType.W4A8,
is_per_channel_weight=False,
)
kwargs = self.dispatcher.get_dispatch_mc2_kwargs(token_dispatch_input)
self.assertEqual(kwargs["expert_token_nums_type"], EXPERT_TOKEN_NUMS_TYPE_CUMSUM)
def test_get_combine_mc_kwargs_with_quant(self):
hidden_states = torch.randn(10, 128)
topk_ids = torch.randint(0, 8, (10, 1))
topk_weights = torch.randn(10, 1)
expert_map = torch.tensor([0, 1, 2, 3, 4, 5, 6, 7])
ep_recv_counts = torch.tensor([0, 1, 2, 3, 4, 5, 6, 7])
tp_recv_counts = torch.tensor([0, 1, 2, 3, 4, 5, 6, 7])
assist_info_for_combine = torch.arange(10)
token_dispatch_input = build_token_dispatch_input_fixture(
hidden_states=hidden_states,
topk_weights=topk_weights,
topk_ids=topk_ids,
expert_map=expert_map,
)
combine_metadata = MoEMC2CombineMetadata(
topk_ids=topk_ids,
topk_weights=topk_weights,
expert_map=expert_map,
ep_recv_counts=ep_recv_counts,
tp_recv_counts=tp_recv_counts,
assist_info_for_combine=assist_info_for_combine,
expand_scales=None,
quant=token_dispatch_input.quant,
)
self.dispatcher.need_extra_args = True
self.dispatcher.enable_dispatch_v2 = True
self.dispatcher.moe_expert_num = len(expert_map)
kwargs = self.dispatcher.get_combine_mc_kwargs(hidden_states, combine_metadata)
self.assertIn("tp_send_counts", kwargs)
def test_get_dispatch_mc2_kwargs_with_mxfp8_quant(self):
hidden_states = torch.randn(10, 128)
topk_ids = torch.randint(0, 8, (10, 1))
topk_weights = torch.randn(10, 1)
expert_map = torch.tensor([0, 1, 2, 3, 4, 5, 6, 7])
self.dispatcher.a5_need_extra_args = True
token_dispatch_input = build_token_dispatch_input_fixture(
hidden_states=hidden_states,
topk_weights=topk_weights,
topk_ids=topk_ids,
expert_map=expert_map,
quant_type=QuantType.MXFP8,
act_quant_type=torch.float8_e4m3fn,
)
kwargs = self.dispatcher.get_dispatch_mc2_kwargs(token_dispatch_input)
self.assertTrue(token_dispatch_input.quant.dispatch_with_quant)
self.assertEqual(kwargs["quant_mode"], 4)
self.assertEqual(kwargs["y_dtype"], torch.float8_e4m3fn)
def test_get_dispatch_mc2_kwargs_with_mxfp4_quant(self):
hidden_states = torch.randn(10, 128)
topk_weights = torch.randn(10, 1)
topk_ids = torch.randint(0, 8, (10, 1))
expert_map = torch.tensor([0, 1, 2, 3, 4, 5, 6, 7])
self.dispatcher.a5_need_extra_args = True
token_dispatch_input = build_token_dispatch_input_fixture(
hidden_states=hidden_states,
topk_weights=topk_weights,
topk_ids=topk_ids,
expert_map=expert_map,
quant_type=QuantType.MXFP4,
act_quant_type=MXFP4_TEST_DTYPE,
)
kwargs = self.dispatcher.get_dispatch_mc2_kwargs(token_dispatch_input)
self.assertTrue(token_dispatch_input.quant.dispatch_with_quant)
self.assertEqual(kwargs["quant_mode"], 4)
self.assertIn("y_dtype", kwargs)
self.assertNotEqual(kwargs["y_dtype"], torch.float8_e4m3fn)
with patch(
"torch_npu.npu_moe_distribute_dispatch_v2",
return_value=(
torch.randn(10, 128),
torch.randn(10, 1),
torch.arange(10, dtype=torch.int32),
torch.tensor([10], dtype=torch.int64),
torch.tensor([10], dtype=torch.int64),
torch.tensor([10], dtype=torch.int64),
torch.randn(10, 1),
),
) as mock_dispatch:
output = self.dispatcher.token_dispatch(token_dispatch_input=token_dispatch_input)
mock_dispatch.assert_called_once()
self.assertIsNotNone(output.dynamic_scale)
self.assertTrue(output.combine_metadata.quant.dispatch_with_quant)
def test_allgather_token_dispatch_quant_mode_without_dynamic_scale():
dispatcher = TokenDispatcherWithAllGather(top_k=2, num_experts=128)
hidden_states = torch.randn(3, 128)
topk_weights = torch.tensor([[0.7, 0.3], [0.6, 0.4], [0.5, 0.5]])
topk_ids = torch.tensor([[0, 1], [1, 2], [2, 3]], dtype=torch.int32)
init_routing_output = (
torch.randn(6, 128),
torch.tensor([0, 1, 2, 3, 4, 5], dtype=torch.int32),
torch.tensor([2, 2, 2], dtype=torch.int32),
torch.randn(6, 4),
)
cases = [
{
"quant_type": QuantType.MXFP8,
"act_quant_type": torch.float8_e4m3fn,
"expected_quant_mode": 3,
"expected_act_quant_type": torch.float8_e4m3fn,
"expect_dynamic_scale": True,
},
{
"quant_type": QuantType.MXFP4,
"act_quant_type": MXFP4_TEST_DTYPE,
"expected_quant_mode": -1,
"expected_act_quant_type": None,
"expect_dynamic_scale": False,
},
]
for case in cases:
token_dispatch_input = build_token_dispatch_input_fixture(
hidden_states=hidden_states,
topk_weights=topk_weights,
topk_ids=topk_ids,
quant_type=case["quant_type"],
act_quant_type=case["act_quant_type"],
)
with patch(
"vllm_ascend.ops.fused_moe.token_dispatcher.DeviceOperator.npu_moe_init_routing",
return_value=init_routing_output,
) as mock_init_routing:
output = dispatcher.token_dispatch(token_dispatch_input=token_dispatch_input)
init_kwargs = mock_init_routing.call_args.kwargs
assert init_kwargs["quant_mode"] == case["expected_quant_mode"]
assert init_kwargs["act_quant_type"] == case["expected_act_quant_type"]
assert (output.dynamic_scale is not None) == case["expect_dynamic_scale"]
def test_allgather_token_dispatch_mxfp4_keeps_prequantized_scale():
dispatcher = TokenDispatcherWithAllGather(top_k=2, num_experts=128)
hidden_states = torch.randn(3, 128)
topk_weights = torch.tensor([[0.7, 0.3], [0.6, 0.4], [0.5, 0.5]])
topk_ids = torch.tensor([[0, 1], [1, 2], [2, 3]], dtype=torch.int32)
pertoken_scale = torch.randn(3, 4)
returned_scale = torch.randn(6, 4)
init_routing_output = (
torch.randn(6, 128),
torch.tensor([0, 1, 2, 3, 4, 5], dtype=torch.int32),
torch.tensor([2, 2, 2], dtype=torch.int32),
returned_scale,
)
token_dispatch_input = build_token_dispatch_input_fixture(
hidden_states=hidden_states,
topk_weights=topk_weights,
topk_ids=topk_ids,
pertoken_scale=pertoken_scale,
quant_type=QuantType.MXFP4,
act_quant_type=MXFP4_TEST_DTYPE,
)
with patch(
"vllm_ascend.ops.fused_moe.token_dispatcher.DeviceOperator.npu_moe_init_routing",
return_value=init_routing_output,
) as mock_init_routing:
output = dispatcher.token_dispatch(token_dispatch_input=token_dispatch_input)
init_kwargs = mock_init_routing.call_args.kwargs
assert init_kwargs["scale"] is pertoken_scale
assert init_kwargs["quant_mode"] == -1
assert init_kwargs["act_quant_type"] == MXFP4_TEST_DTYPE
assert output.dynamic_scale is returned_scale
class TestTokenDispatcherWithAllGather(TestBase):
def setUp(self):
# Mock dependencies
kwargs = {
"apply_router_weight_on_input": False,
"top_k": 2,
"max_num_tokens": 100,
"ep_size": 2,
"num_experts": 128,
"with_quant": False,
}
self.dispatcher = TokenDispatcherWithAllGather(**kwargs)
# Mock NPU functions
self.patcher_npu_moe_init_routing_custom = patch("torch.ops._C_ascend.npu_moe_init_routing_custom")
self.mock_npu_moe_init_routing_custom = self.patcher_npu_moe_init_routing_custom.start()
self.mock_npu_moe_init_routing_custom.return_value = (
torch.randn(6, 128), # sorted_hidden_states
torch.tensor([0, 1, 2, 3, 4, 5]), # expanded_row_idx
torch.tensor([0, 1, 0, 1, 0, 1]), # expanded_expert_idx
torch.tensor([0, 1, 0, 1, 0, 1]),
)
self.patcher_npu_moe_token_unpermute = patch("torch_npu.npu_moe_token_unpermute")
self.mock_npu_moe_token_unpermute = self.patcher_npu_moe_token_unpermute.start()
self.mock_npu_moe_token_unpermute.return_value = torch.randn(6, 128)
def tearDown(self):
self.patcher_npu_moe_init_routing_custom.stop()
self.patcher_npu_moe_token_unpermute.stop()
@pytest.mark.skip("Skip as register_kernels has NPU SocName checking in CANN 8.5.0.")
def test_token_dispatch_without_expert_map(self):
hidden_states = torch.randn(3, 128)
topk_weights = torch.tensor([[0.7, 0.3], [0.6, 0.4], [0.5, 0.5]])
topk_ids = torch.tensor([[0, 1], [1, 2], [2, 3]])
token_dispatch_input = build_token_dispatch_input_fixture(
hidden_states=hidden_states,
topk_weights=topk_weights,
topk_ids=topk_ids,
)
results = self.dispatcher.token_dispatch(token_dispatch_input=token_dispatch_input)
# Verify npu_moe_init_routing is called
self.mock_npu_moe_init_routing_custom.assert_called_once()
args, kwargs = self.mock_npu_moe_init_routing_custom.call_args
self.assertEqual(results.group_list_type, 1)
self.assertIsInstance(results.combine_metadata, MoEAllGatherCombineMetadata)
@pytest.mark.skip("Skip as register_kernels has NPU SocName checking in CANN 8.5.0.")
def test_token_dispatch_with_expert_map(self):
self.dispatcher.expert_map = torch.tensor([0, 1, 2, 3])
hidden_states = torch.randn(3, 128)
topk_weights = torch.tensor([[0.7, 0.3], [0.6, 0.4], [0.5, 0.5]])
topk_ids = torch.tensor([[0, 1], [1, 2], [2, 3]])
token_dispatch_input = build_token_dispatch_input_fixture(
hidden_states=hidden_states,
topk_weights=topk_weights,
topk_ids=topk_ids,
)
results = self.dispatcher.token_dispatch(token_dispatch_input=token_dispatch_input)
# Verify npu_moe_init_routing is called
self.mock_npu_moe_init_routing_custom.assert_called_once()
args, kwargs = self.mock_npu_moe_init_routing_custom.call_args
self.assertEqual(results.group_list_type, 1)
self.assertIsInstance(results.combine_metadata, MoEAllGatherCombineMetadata)
@pytest.mark.skip("Skip as register_kernels has NPU SocName checking in CANN 8.5.0.")
def test_token_dispatch_without_quant(self):
kwargs = {
"apply_router_weight_on_input": False,
"top_k": 2,
"max_num_tokens": 100,
"ep_size": 2,
"num_experts": 128,
}
self.dispatcher_quant = TokenDispatcherWithAllGather(**kwargs)
hidden_states = torch.randn(3, 128)
topk_weights = torch.tensor([[0.7, 0.3], [0.6, 0.4], [0.5, 0.5]])
topk_ids = torch.tensor([[0, 1], [1, 2], [2, 3]])
token_dispatch_input = build_token_dispatch_input_fixture(
hidden_states=hidden_states,
topk_weights=topk_weights,
topk_ids=topk_ids,
)
results = self.dispatcher_quant.token_dispatch(token_dispatch_input=token_dispatch_input)
self.assertEqual(results.group_list_type, 1)
@pytest.mark.skip("Skip as register_kernels has NPU SocName checking in CANN 8.5.0.")
def test_token_dispatch_with_quant(self):
kwargs = {
"apply_router_weight_on_input": False,
"top_k": 2,
"max_num_tokens": 100,
"ep_size": 2,
"num_experts": 128,
}
self.dispatcher_quant = TokenDispatcherWithAllGather(**kwargs)
hidden_states = torch.randn(3, 128)
topk_weights = torch.tensor([[0.7, 0.3], [0.6, 0.4], [0.5, 0.5]])
topk_ids = torch.tensor([[0, 1], [1, 2], [2, 3]])
token_dispatch_input = build_token_dispatch_input_fixture(
hidden_states=hidden_states,
topk_weights=topk_weights,
topk_ids=topk_ids,
quant_type=QuantType.W8A8,
)
results = self.dispatcher_quant.token_dispatch(token_dispatch_input=token_dispatch_input)
self.assertIsNotNone(results.hidden_states)
self.assertIsNotNone(results.group_list)
self.assertIsNotNone(results.dynamic_scale)
self.assertEqual(results.group_list_type, 1)
@pytest.mark.skip("Skip as register_kernels has NPU SocName checking in CANN 8.5.0.")
def test_token_combine_with_expert_map(self):
hidden_states = torch.randn(6, 128)
combine_metadata = MoEAllGatherCombineMetadata(
expanded_row_idx=torch.tensor([0, 1, 1, 1, 1, 1]),
topk_weights=torch.tensor([0.5, 0.5, 0.5, 0.5, 0.5, 0.5]),
restore_shape=torch.Size([6, 128]),
)
final_hidden_states = self.dispatcher.token_combine(hidden_states, combine_metadata)
self.assertEqual(final_hidden_states.shape, (6, 128))
@pytest.mark.skip("Skip as register_kernels has NPU SocName checking in CANN 8.5.0.")
def test_token_combine_without_expert_map(self):
hidden_states = torch.randn(6, 128)
combine_metadata = MoEAllGatherCombineMetadata(
expanded_row_idx=torch.tensor([0, 1, 1, 1, 1, 1]),
topk_weights=torch.tensor([0.5, 0.5, 0.5, 0.5, 0.5, 0.5]),
restore_shape=torch.Size([6, 128]),
)
final_hidden_states = self.dispatcher.token_combine(hidden_states, combine_metadata)
self.mock_npu_moe_token_unpermute.assert_called_once()
self.assertEqual(final_hidden_states.shape, (6, 128))
@pytest.mark.skip("Skip as register_kernels has NPU SocName checking in CANN 8.5.0.")
def test_token_dispatch_with_router_weight(self):
hidden_states = torch.randn(3, 128)
topk_weights = torch.tensor([[0.7], [0.6], [0.5]]) # topk=1
topk_ids = torch.tensor([[0], [1], [2]])
token_dispatch_input = build_token_dispatch_input_fixture(
hidden_states=hidden_states,
topk_weights=topk_weights,
topk_ids=topk_ids,
apply_router_weight_on_input=True,
)
results = self.dispatcher.token_dispatch(token_dispatch_input=token_dispatch_input)
self.assertEqual(results.hidden_states.shape, (6, 128))
self.assertIsInstance(results.combine_metadata, MoEAllGatherCombineMetadata)
class TestTokenDispatcherWithAll2AllV(TestBase):
def setUp(self):
# Patch properties
patcher1 = patch.object(
TokenDispatcherWithAll2AllV, "ep_group", new_callable=PropertyMock, return_value=MagicMock()
)
patcher2 = patch.object(TokenDispatcherWithAll2AllV, "ep_rank", new_callable=PropertyMock, return_value=0)
patcher3 = patch.object(TokenDispatcherWithAll2AllV, "ep_size", new_callable=PropertyMock, return_value=2)
self.addCleanup(patcher1.stop)
self.addCleanup(patcher2.stop)
self.addCleanup(patcher3.stop)
self.mock_ep_group_prop = patcher1.start()
self.mock_ep_rank_prop = patcher2.start()
self.mock_ep_size_prop = patcher3.start()
# Mock torch_npu.npu_moe_token_permute
patcher4 = patch("torch_npu.npu_moe_token_permute")
self.mock_npu_moe_token_permute = patcher4.start()
self.addCleanup(patcher4.stop)
self.mock_npu_moe_token_permute.return_value = (torch.randn(16, 16), torch.arange(16))
# Mock torch_npu.npu_moe_token_unpermute
patcher5 = patch("torch_npu.npu_moe_token_unpermute")
self.mock_npu_moe_token_unpermute = patcher5.start()
self.addCleanup(patcher5.stop)
self.mock_npu_moe_token_unpermute.return_value = torch.randn(8, 16)
# Mock async_all_to_all
patcher6 = patch("vllm_ascend.ops.fused_moe.comm_utils.async_all_to_all")
self.mock_async_all_to_all = patcher6.start()
self.addCleanup(patcher6.stop)
self.mock_async_all_to_all.return_value = (None, torch.randn(16, 16), MagicMock())
# Mock gather_from_sequence_parallel_region
patcher7 = patch("vllm_ascend.ops.fused_moe.token_dispatcher.gather_from_sequence_parallel_region")
self.mock_gather_from_sequence_parallel_region = patcher7.start()
self.addCleanup(patcher7.stop)
self.mock_gather_from_sequence_parallel_region.return_value = torch.tensor(
[[2, 2, 2, 2], [2, 2, 2, 2]], dtype=torch.int64
)
# Mock torch.histc
patcher8 = patch("torch.histc")
self.mock_histc = patcher8.start()
self.addCleanup(patcher8.stop)
self.mock_histc.return_value = torch.tensor([2, 2, 2, 2], dtype=torch.int64)
# Mock torch.npu.current_device
patcher9 = patch("torch.npu.current_device")
self.mock_current_device = patcher9.start()
self.addCleanup(patcher9.stop)
self.mock_current_device.return_value = "cpu"
# Mock torch_npu.npu_dynamic_quant
patcher10 = patch("torch_npu.npu_dynamic_quant")
self.mock_npu_dynamic_quant = patcher10.start()
self.addCleanup(patcher10.stop)
self.mock_npu_dynamic_quant.return_value = (torch.randn(16, 16), torch.randn(16))
# Mock torch.ops._C_ascend.npu_moe_init_routing_custom
patcher11 = patch("torch.ops._C_ascend.npu_moe_init_routing_custom")
self.mock_npu_moe_init_routing_custom = patcher11.start()
self.addCleanup(patcher11.stop)
self.mock_npu_moe_init_routing_custom.return_value = (
torch.randn(16, 16),
torch.arange(16),
None,
torch.randn(16),
)
# Mock torch.repeat_interleave
patcher12 = patch("torch.repeat_interleave")
self.mock_repeat_interleave = patcher12.start()
self.addCleanup(patcher12.stop)
self.mock_repeat_interleave.return_value = torch.arange(16)
self.dispatcher = TokenDispatcherWithAll2AllV(top_k=2, num_experts=4, num_local_experts=2, with_quant=False)
@pytest.mark.skip("Skip as register_kernels has NPU SocName checking in CANN 8.5.0.")
def test_token_dispatch(self):
hidden_states = torch.randn(8, 16)
topk_weights = torch.rand(8, 4)
topk_ids = torch.randint(0, 4, (8, 2)).long()
expert_map = torch.tensor([0, 1, 2, 3])
self.dispatcher.expert_ids_per_ep_rank = torch.tensor([0, 1], dtype=torch.int32)
self.dispatcher.local_expert_indices = [0, 1]
token_dispatch_input = build_token_dispatch_input_fixture(
hidden_states=hidden_states,
topk_weights=topk_weights,
topk_ids=topk_ids,
expert_map=expert_map,
)
result = self.dispatcher.token_dispatch(token_dispatch_input=token_dispatch_input)
self.assertIsNotNone(result.hidden_states)
self.assertIsNotNone(result.group_list)
self.assertEqual(result.group_list_type, 1)
self.assertIsInstance(result.combine_metadata, MoEAllToAllCombineMetadata)
@pytest.mark.skip("Skip as register_kernels has NPU SocName checking in CANN 8.5.0.")
def test_token_combine(self):
hidden_states = torch.randn(16, 16)
combine_metadata = MoEAllToAllCombineMetadata(
input_splits=np.array([4, 4]),
output_splits=np.array([4, 4]),
topk_weights=torch.rand(8, 4),
reversed_local_input_permutation_mapping=torch.arange(8),
reversed_global_input_permutation_mapping=torch.arange(16),
hidden_shape=torch.Size([8, 16]),
hidden_shape_before_permute=torch.Size([8, 16]),
)
self.dispatcher.expert_ids_per_ep_rank = torch.tensor([0, 1], dtype=torch.int32)
self.dispatcher.local_expert_indices = [0, 1]
output = self.dispatcher.token_combine(hidden_states, combine_metadata)
self.assertIsNotNone(output)
self.assertEqual(output.shape, (8, 16))
@pytest.mark.skip("Skip as register_kernels has NPU SocName checking in CANN 8.5.0.")
def test_token_dispatch_with_quant(self):
self.dispatcher = TokenDispatcherWithAll2AllV(top_k=2, num_experts=4, num_local_experts=2)
hidden_states = torch.randn(8, 16)
topk_weights = torch.rand(8, 4)
topk_ids = torch.randint(0, 4, (8, 2)).long()
expert_map = torch.tensor([0, 1, 2, 3])
self.dispatcher.expert_ids_per_ep_rank = torch.tensor([0, 1], dtype=torch.int32)
self.dispatcher.local_expert_indices = [0, 1]
token_dispatch_input = build_token_dispatch_input_fixture(
hidden_states=hidden_states,
topk_weights=topk_weights,
topk_ids=topk_ids,
expert_map=expert_map,
quant_type=QuantType.W8A8,
)
result = self.dispatcher.token_dispatch(token_dispatch_input=token_dispatch_input)
self.assertIsNotNone(result.hidden_states)
self.assertIsNotNone(result.group_list)
self.assertIsNotNone(result.dynamic_scale)
self.assertEqual(result.group_list_type, 1)
self.assertIsInstance(result.combine_metadata, MoEAllToAllCombineMetadata)
@pytest.mark.skip("Skip as register_kernels has NPU SocName checking in CANN 8.5.0.")
def test_token_dispatch_with_quant_no_active_tokens(self):
self.dispatcher = TokenDispatcherWithAll2AllV(top_k=2, num_experts=4, num_local_experts=2)
self.mock_repeat_interleave.return_value = torch.tensor([], dtype=torch.long)
hidden_states = torch.randn(8, 16)
topk_weights = torch.rand(8, 4)
topk_ids = torch.randint(0, 4, (8, 2)).long()
expert_map = torch.tensor([0, 1, 2, 3])
self.dispatcher.expert_ids_per_ep_rank = torch.tensor([0, 1], dtype=torch.int32)
self.dispatcher.local_expert_indices = [0, 1]
token_dispatch_input = build_token_dispatch_input_fixture(
hidden_states=hidden_states,
topk_weights=topk_weights,
topk_ids=topk_ids,
expert_map=expert_map,
quant_type=QuantType.W8A8,
)
result = self.dispatcher.token_dispatch(token_dispatch_input=token_dispatch_input)
self.assertIsNotNone(result.hidden_states)
self.assertIsNotNone(result.group_list)
self.assertIsNotNone(result.dynamic_scale)
self.assertEqual(result.group_list_type, 1)

View File

@@ -0,0 +1,458 @@
#
# 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 unittest.mock import MagicMock, patch
import pytest
import torch
from vllm_ascend.ops.weight_prefetch import (
MAX_PREFETCH_WEIGHT_SIZE,
MOE_PREFETCH_TOKEN_THRESHOLD,
SUPPORTED_MODULES,
ModuleWeightPrefetchConfig,
WeightPrefetchMethod,
maybe_npu_prefetch,
)
class TestModuleWeightPrefetchConfig:
def test_init_with_valid_module_name(self):
for module_name in SUPPORTED_MODULES:
config = ModuleWeightPrefetchConfig(module_name=module_name)
assert config.module_name == module_name
assert config.enable is False
assert config.is_active_this_forward is False
assert config.prefetch_ratio == {}
assert config.linear_prefix_map == {}
def test_init_with_invalid_module_name(self):
with pytest.raises(AssertionError, match="Invalid module name"):
ModuleWeightPrefetchConfig(module_name="invalid_module")
def test_prefetch_ratio_filtering(self):
config = ModuleWeightPrefetchConfig(
module_name="attn",
prefetch_ratio={"qkv": 0.8, "o": 1.2, "invalid": -0.5},
)
assert "qkv" in config.prefetch_ratio
assert "o" not in config.prefetch_ratio
assert "invalid" not in config.prefetch_ratio
def test_enable_logic_with_prefetch_ratio(self):
config = ModuleWeightPrefetchConfig(
module_name="attn",
enable=True,
prefetch_ratio={"qkv": 0.8},
)
assert config.enable is True
def test_enable_logic_without_prefetch_ratio(self):
config = ModuleWeightPrefetchConfig(
module_name="attn",
enable=True,
prefetch_ratio={},
)
assert config.enable is False
class TestWeightPrefetchMethod:
@pytest.fixture
def mock_weight_prefetch_config(self):
config = MagicMock()
config.enabled = True
config.prefetch_ratio = {
"attn": {"qkv": 0.8, "o": 0.8},
"moe": {"gate_up": 0.8},
"mlp": {"gate_up": 1.0, "down": 1.0},
}
return config
@pytest.fixture
def mock_vllm_config(self):
config = MagicMock()
config.model_config = MagicMock()
config.model_config.hf_config = MagicMock()
config.model_config.hf_config.model_type = "llama"
return config
@patch("vllm_ascend.ops.weight_prefetch.is_moe_model", return_value=False)
@patch("vllm_ascend.ops.weight_prefetch.get_current_vllm_config")
def test_init_non_moe_model(
self,
mock_get_config,
mock_is_moe,
mock_weight_prefetch_config,
mock_vllm_config,
):
mock_get_config.return_value = mock_vllm_config
method = WeightPrefetchMethod(mock_weight_prefetch_config)
assert method.is_moe is False
assert method.attn.enable is True
assert method.moe.enable is False
assert method.mlp.enable is True
@patch("vllm_ascend.ops.weight_prefetch.is_moe_model", return_value=True)
@patch("vllm_ascend.ops.weight_prefetch.get_current_vllm_config")
def test_init_moe_model(
self,
mock_get_config,
mock_is_moe,
mock_weight_prefetch_config,
mock_vllm_config,
):
mock_get_config.return_value = mock_vllm_config
method = WeightPrefetchMethod(mock_weight_prefetch_config)
assert method.is_moe is True
assert method.attn.enable is True
assert method.moe.enable is True
assert method.mlp.enable is False
@patch("vllm_ascend.ops.weight_prefetch.is_moe_model", return_value=False)
@patch("vllm_ascend.ops.weight_prefetch.get_current_vllm_config")
def test_init_disabled_config(
self,
mock_get_config,
mock_is_moe,
mock_vllm_config,
):
mock_get_config.return_value = mock_vllm_config
disabled_config = MagicMock()
disabled_config.enabled = False
disabled_config.prefetch_ratio = {}
method = WeightPrefetchMethod(disabled_config)
assert method.attn.enable is False
assert method.moe.enable is False
assert method.mlp.enable is False
@patch("vllm_ascend.ops.weight_prefetch.is_moe_model", return_value=False)
@patch("vllm_ascend.ops.weight_prefetch.get_current_vllm_config")
@patch("torch.ops.vllm.prefetch_preprocess")
def test_maybe_prefetch_attn_weight_preprocess_enabled(
self,
mock_prefetch,
mock_get_config,
mock_is_moe,
mock_weight_prefetch_config,
mock_vllm_config,
):
mock_get_config.return_value = mock_vllm_config
method = WeightPrefetchMethod(mock_weight_prefetch_config)
weight = torch.randn(1024, 1024)
start_flag = torch.tensor([1])
method.maybe_prefetch_attn_weight_preprocess(
layer_cls_name="AscendQKVParallelLinear",
weight=weight,
start_flag=start_flag,
)
mock_prefetch.assert_called_once()
call_kwargs = mock_prefetch.call_args[1]
assert call_kwargs["weight"] is weight
assert call_kwargs["start_flag"] is start_flag
@patch("vllm_ascend.ops.weight_prefetch.is_moe_model", return_value=False)
@patch("vllm_ascend.ops.weight_prefetch.get_current_vllm_config")
@patch("torch.ops.vllm.prefetch_preprocess")
def test_maybe_prefetch_attn_weight_preprocess_disabled(
self,
mock_prefetch,
mock_get_config,
mock_is_moe,
mock_weight_prefetch_config,
mock_vllm_config,
):
mock_get_config.return_value = mock_vllm_config
disabled_config = MagicMock()
disabled_config.enabled = False
disabled_config.prefetch_ratio = {}
method = WeightPrefetchMethod(disabled_config)
weight = torch.randn(1024, 1024)
start_flag = torch.tensor([1])
method.maybe_prefetch_attn_weight_preprocess(
layer_cls_name="AscendQKVParallelLinear",
weight=weight,
start_flag=start_flag,
)
mock_prefetch.assert_not_called()
@patch("vllm_ascend.ops.weight_prefetch.is_moe_model", return_value=False)
@patch("vllm_ascend.ops.weight_prefetch.get_current_vllm_config")
@patch("torch.ops.vllm.prefetch_postprocess")
def test_maybe_prefetch_attn_weight_postprocess_enabled(
self,
mock_postprocess,
mock_get_config,
mock_is_moe,
mock_weight_prefetch_config,
mock_vllm_config,
):
mock_get_config.return_value = mock_vllm_config
method = WeightPrefetchMethod(mock_weight_prefetch_config)
stop_flag = torch.tensor([1])
method.maybe_prefetch_attn_weight_postprocess(
layer_cls_name="AscendQKVParallelLinear",
stop_flag=stop_flag,
)
mock_postprocess.assert_called_once_with(stop_flag)
@patch("vllm_ascend.ops.weight_prefetch.is_moe_model", return_value=False)
@patch("vllm_ascend.ops.weight_prefetch.get_current_vllm_config")
@patch("torch.ops.vllm.prefetch_postprocess")
def test_maybe_prefetch_attn_weight_postprocess_disabled(
self,
mock_postprocess,
mock_get_config,
mock_is_moe,
mock_weight_prefetch_config,
mock_vllm_config,
):
mock_get_config.return_value = mock_vllm_config
disabled_config = MagicMock()
disabled_config.enabled = False
disabled_config.prefetch_ratio = {}
method = WeightPrefetchMethod(disabled_config)
stop_flag = torch.tensor([1])
method.maybe_prefetch_attn_weight_postprocess(
layer_cls_name="AscendQKVParallelLinear",
stop_flag=stop_flag,
)
mock_postprocess.assert_not_called()
@patch("vllm_ascend.ops.weight_prefetch.is_moe_model", return_value=True)
@patch("vllm_ascend.ops.weight_prefetch.get_current_vllm_config")
@patch("vllm_ascend.ops.weight_prefetch.get_forward_context")
@patch("vllm_ascend.ops.weight_prefetch._EXTRA_CTX")
@patch("torch.ops.vllm.prefetch_preprocess")
def test_maybe_prefetch_moe_weight_preprocess_enabled(
self,
mock_prefetch,
mock_extra_ctx,
mock_get_forward_context,
mock_get_config,
mock_is_moe,
mock_weight_prefetch_config,
mock_vllm_config,
):
mock_get_config.return_value = mock_vllm_config
mock_get_forward_context.return_value = MagicMock()
mock_model_instance = MagicMock()
mock_layer = MagicMock()
mock_layer.mlp.experts.w13_weight = torch.randn(1024, 1024)
mock_model_instance.model.layers = [mock_layer]
mock_extra_ctx.model_instance = mock_model_instance
mock_extra_ctx.layer_idx = 1
method = WeightPrefetchMethod(mock_weight_prefetch_config)
hidden_states = torch.randn(MOE_PREFETCH_TOKEN_THRESHOLD + 10, 1024)
method.maybe_prefetch_moe_weight_preprocess(hidden_states, "gate_up")
assert method.moe.is_active_this_forward is True
mock_prefetch.assert_called_once()
@patch("vllm_ascend.ops.weight_prefetch.is_moe_model", return_value=True)
@patch("vllm_ascend.ops.weight_prefetch.get_current_vllm_config")
def test_maybe_prefetch_moe_weight_preprocess_below_threshold(
self,
mock_get_config,
mock_is_moe,
mock_weight_prefetch_config,
mock_vllm_config,
):
mock_get_config.return_value = mock_vllm_config
method = WeightPrefetchMethod(mock_weight_prefetch_config)
hidden_states = torch.randn(MOE_PREFETCH_TOKEN_THRESHOLD - 10, 1024)
method.maybe_prefetch_moe_weight_preprocess(hidden_states, "gate_up")
assert method.moe.is_active_this_forward is False
@patch("vllm_ascend.ops.weight_prefetch.is_moe_model", return_value=True)
@patch("vllm_ascend.ops.weight_prefetch.get_current_vllm_config")
@patch("torch.ops.vllm.prefetch_postprocess")
def test_maybe_prefetch_moe_weight_postprocess_enabled(
self,
mock_postprocess,
mock_get_config,
mock_is_moe,
mock_weight_prefetch_config,
mock_vllm_config,
):
mock_get_config.return_value = mock_vllm_config
method = WeightPrefetchMethod(mock_weight_prefetch_config)
method.moe.is_active_this_forward = True
stop_flag = torch.tensor([1])
method.maybe_prefetch_moe_weight_postprocess(stop_flag)
mock_postprocess.assert_called_once_with(stop_flag)
@patch("vllm_ascend.ops.weight_prefetch.is_moe_model", return_value=False)
@patch("vllm_ascend.ops.weight_prefetch.get_current_vllm_config")
@patch("torch.ops.vllm.prefetch_postprocess")
def test_maybe_prefetch_moe_weight_postprocess_disabled(
self,
mock_postprocess,
mock_get_config,
mock_is_moe,
mock_weight_prefetch_config,
mock_vllm_config,
):
mock_get_config.return_value = mock_vllm_config
method = WeightPrefetchMethod(mock_weight_prefetch_config)
method.moe.is_active_this_forward = False
stop_flag = torch.tensor([1])
method.maybe_prefetch_moe_weight_postprocess(stop_flag)
mock_postprocess.assert_not_called()
@patch("vllm_ascend.ops.weight_prefetch.is_moe_model", return_value=False)
@patch("vllm_ascend.ops.weight_prefetch.get_current_vllm_config")
@patch("torch.ops.vllm.prefetch_preprocess")
def test_maybe_prefetch_mla_or_sla_weight_enabled(
self,
mock_prefetch,
mock_get_config,
mock_is_moe,
mock_weight_prefetch_config,
mock_vllm_config,
):
mock_get_config.return_value = mock_vllm_config
method = WeightPrefetchMethod(mock_weight_prefetch_config)
inputs = torch.randn(1024, 1024)
dependency = torch.tensor([1])
method.maybe_prefetch_mla_or_sla_weight_in_current_stream(
inputs=inputs,
dependency=dependency,
max_size=1024,
)
mock_prefetch.assert_called_once()
@patch("vllm_ascend.ops.weight_prefetch.is_moe_model", return_value=False)
@patch("vllm_ascend.ops.weight_prefetch.get_current_vllm_config")
@patch("torch.ops.vllm.prefetch_preprocess")
def test_maybe_prefetch_mla_or_sla_weight_disabled(
self,
mock_prefetch,
mock_get_config,
mock_is_moe,
mock_vllm_config,
):
mock_get_config.return_value = mock_vllm_config
disabled_config = MagicMock()
disabled_config.enabled = False
disabled_config.prefetch_ratio = {}
method = WeightPrefetchMethod(disabled_config)
inputs = torch.randn(1024, 1024)
dependency = torch.tensor([1])
method.maybe_prefetch_mla_or_sla_weight_in_current_stream(
inputs=inputs,
dependency=dependency,
)
mock_prefetch.assert_not_called()
class TestMaybeNpuPrefetch:
@patch("torch_npu.npu_prefetch")
def test_maybe_npu_prefetch_enabled(self, mock_prefetch):
inputs = torch.randn(1024, 1024)
dependency = torch.tensor([1])
maybe_npu_prefetch(inputs, dependency, enabled=True)
mock_prefetch.assert_called_once()
call_args = mock_prefetch.call_args[0]
assert call_args[0] is inputs
assert call_args[1] is dependency
@patch("torch_npu.npu_prefetch")
def test_maybe_npu_prefetch_disabled(self, mock_prefetch):
inputs = torch.randn(1024, 1024)
dependency = torch.tensor([1])
maybe_npu_prefetch(inputs, dependency, enabled=False)
mock_prefetch.assert_not_called()
@patch("torch_npu.npu_prefetch")
def test_maybe_npu_prefetch_max_size_calculation(self, mock_prefetch):
inputs = torch.randn(100, 100, dtype=torch.float32)
dependency = torch.tensor([1])
maybe_npu_prefetch(inputs, dependency, max_size=0, enabled=True)
expected_size = inputs.element_size() * inputs.numel()
call_args = mock_prefetch.call_args[0]
assert call_args[2] == expected_size
@patch("torch_npu.npu_prefetch")
def test_maybe_npu_prefetch_with_custom_max_size(self, mock_prefetch):
inputs = torch.randn(100, 100, dtype=torch.float32)
dependency = torch.tensor([1])
custom_max_size = 1000
maybe_npu_prefetch(inputs, dependency, max_size=custom_max_size, enabled=True)
call_args = mock_prefetch.call_args[0]
assert call_args[2] == custom_max_size
@patch("torch_npu.npu_prefetch")
def test_maybe_npu_prefetch_with_offset(self, mock_prefetch):
inputs = torch.randn(1024, 1024)
dependency = torch.tensor([1])
offset = 100
maybe_npu_prefetch(inputs, dependency, offset=offset, enabled=True)
call_args = mock_prefetch.call_args[0]
assert call_args[3] == offset
class TestConstants:
def test_supported_modules(self):
assert "attn" in SUPPORTED_MODULES
assert "mlp" in SUPPORTED_MODULES
assert "moe" in SUPPORTED_MODULES
def test_moe_prefetch_token_threshold(self):
assert MOE_PREFETCH_TOKEN_THRESHOLD == 96
def test_max_prefetch_weight_size(self):
assert MAX_PREFETCH_WEIGHT_SIZE == 18 * 1024 * 1024

View File

View File

@@ -0,0 +1,343 @@
#
# 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 unittest.mock import MagicMock, patch
import pytest
import torch
from vllm.config import set_current_vllm_config
from vllm.model_executor.layers.activation import QuickGELU, SiluAndMul
from vllm_ascend.ops.activation import (
AscendQuickGELU,
AscendSiluAndMul,
AscendSwigluOAIAndMul,
AscendSwigluStepAndMul,
)
from vllm_ascend.utils import is_310p as is_310p_hw
@pytest.fixture
def dummy_tensor():
return torch.randn(4, 8, dtype=torch.float16)
@pytest.fixture
def default_vllm_config():
mock_config = MagicMock()
mock_config.compilation_config.dispatch_forward_backend = "eager"
mock_config.compilation_config.custom_ops = ["all"]
with set_current_vllm_config(mock_config):
yield mock_config
@patch("torch_npu.npu_fast_gelu", side_effect=lambda x: x + 1)
def test_QuickGELU_forward(mock_gelu, dummy_tensor, default_vllm_config):
layer = QuickGELU()
out = layer.forward(dummy_tensor)
expected_out = dummy_tensor + 1
assert torch.allclose(out, expected_out)
mock_gelu.assert_called_once()
@patch("torch_npu.npu_fast_gelu", side_effect=lambda x: x + 1)
def test_AscendQuickGELU_forward_oot(mock_gelu, dummy_tensor, default_vllm_config):
layer = AscendQuickGELU()
out = layer.forward_oot(dummy_tensor)
assert torch.allclose(out, dummy_tensor + 1)
mock_gelu.assert_called_once_with(dummy_tensor)
@patch("vllm_ascend.ops.activation.get_weight_prefetch_method", return_value=MagicMock())
@patch("torch_npu.npu_swiglu", side_effect=lambda x: x + 1)
def test_SiluAndMul_forward(
mock_swiglu,
mock_get_weight_prefetch_method,
dummy_tensor,
default_vllm_config,
):
layer = SiluAndMul()
out = layer.forward(dummy_tensor)
expected_arg = dummy_tensor
# assert mock_swiglu.call_count == 1
mock_swiglu.assert_called_once()
actual_arg = mock_swiglu.call_args[0][0]
assert torch.allclose(actual_arg, expected_arg), "npu_swiglu called with unexpected input"
expected_out = dummy_tensor + 1
assert torch.allclose(out, expected_out)
@patch("vllm_ascend.ops.activation.get_weight_prefetch_method")
@patch("torch_npu.npu_swiglu", side_effect=lambda x: x + 1)
def test_AscendSiluAndMul_forward_oot_prefetch(
mock_swiglu,
mock_get_weight_prefetch_method,
dummy_tensor,
default_vllm_config,
):
weight_prefetch_method = MagicMock()
weight_prefetch_method.MLP_DOWN = "mlp_down"
mock_get_weight_prefetch_method.return_value = weight_prefetch_method
layer = AscendSiluAndMul()
out = layer.forward_oot(dummy_tensor)
weight_prefetch_method.maybe_prefetch_mlp_weight_preprocess.assert_called_once_with(
weight_prefetch_method.MLP_DOWN, dummy_tensor
)
weight_prefetch_method.maybe_prefetch_mlp_weight_postprocess.assert_called_once_with(out)
mock_swiglu.assert_called_once_with(dummy_tensor)
assert torch.allclose(out, dummy_tensor + 1)
@pytest.mark.skipif(not is_310p_hw(), reason="310P device unittest case.")
@patch("torch.nn.functional.silu", side_effect=lambda x: x + 1)
def test_SiluAndMul_forward_310p(
mock_silu,
dummy_tensor,
default_vllm_config,
):
layer = SiluAndMul()
out = layer.forward(dummy_tensor)
h = dummy_tensor.shape[-1] // 2
expected_arg = dummy_tensor[..., :h]
# assert mock_silu.call_count == 1
mock_silu.assert_called_once()
actual_arg = mock_silu.call_args[0][0]
assert torch.allclose(actual_arg, expected_arg), "swiglu called with unexpected input"
expected_out = (dummy_tensor[..., :h] + 1) * dummy_tensor[..., h:]
assert torch.allclose(out, expected_out)
def _swiglu_oai_reference(x: torch.Tensor, alpha: float = 1.702, limit: float = 7.0) -> torch.Tensor:
gate = x[..., ::2].clamp(max=limit)
up = x[..., 1::2].clamp(min=-limit, max=limit)
return (up + 1) * gate * torch.sigmoid(gate * alpha)
def _quick_gelu_reference(x: torch.Tensor) -> torch.Tensor:
return x * torch.sigmoid(1.702 * x)
def _silu_and_mul_reference(x: torch.Tensor) -> torch.Tensor:
d = x.shape[-1] // 2
return torch.nn.functional.silu(x[..., :d]) * x[..., d:]
def _swiglustep_and_mul_reference(x: torch.Tensor, limit: float = 7.0) -> torch.Tensor:
# Independent of the chunk-based implementation under test: slice the gate
# (first half) and up (second half) halves and apply silu + symmetric clamp.
d = x.shape[-1] // 2
gate = torch.nn.functional.silu(x[..., :d]).clamp(max=limit)
up = x[..., d:].clamp(min=-limit, max=limit)
return gate * up
class TestAscendSwigluOAIAndMul:
def test_swiglu_oai_forward_matches_reference_formula(self):
x = torch.tensor([[8.0, 9.0, -2.0, -8.0, 3.0, 0.5, -9.0, 10.0]], dtype=torch.float32)
result = AscendSwigluOAIAndMul.swiglu_oai_forward(x)
expected = _swiglu_oai_reference(x)
assert result.shape == (1, x.shape[-1] // 2)
assert result.dtype == x.dtype
assert torch.allclose(result, expected)
def test_swiglu_oai_forward_uses_interleaved_gate_and_up_layout(self):
x = torch.tensor([[1.0, 10.0, 2.0, 20.0, 3.0, 30.0, 4.0, 40.0]], dtype=torch.float32)
result = AscendSwigluOAIAndMul.swiglu_oai_forward(x, alpha=1.5, limit=100.0)
expected = _swiglu_oai_reference(x, alpha=1.5, limit=100.0)
chunk_based = (x[..., 4:] + 1) * x[..., :4] * torch.sigmoid(x[..., :4] * 1.5)
assert torch.allclose(result, expected)
assert not torch.allclose(result, chunk_based)
def test_swiglu_oai_forward_with_custom_alpha_and_limit_matches_reference(self):
x = torch.tensor([[9.0, 8.0, -5.0, -9.0]], dtype=torch.float32)
alpha = 2.0
limit = 5.0
result = AscendSwigluOAIAndMul.swiglu_oai_forward(x, alpha=alpha, limit=limit)
expected = _swiglu_oai_reference(x, alpha=alpha, limit=limit)
assert torch.allclose(result, expected)
def test_swiglu_oai_forward_clamps_gate_and_up_values(self):
x = torch.tensor([[100.0, 100.0, -100.0, -100.0]], dtype=torch.float32)
result = AscendSwigluOAIAndMul.swiglu_oai_forward(x)
expected = _swiglu_oai_reference(x)
assert torch.allclose(result, expected)
assert not torch.isnan(result).any()
assert not torch.isinf(result).any()
def test_swiglu_oai_forward_large_input(self):
x = torch.randn(64, 128, dtype=torch.float32)
result = AscendSwigluOAIAndMul.swiglu_oai_forward(x)
expected = _swiglu_oai_reference(x)
assert result.shape == (64, 64)
assert torch.allclose(result, expected)
assert not torch.isnan(result).any()
class TestSwiglustepAndMul:
def test_swiglustep_and_mul_matches_reference_formula(self):
x = torch.tensor([[1.0, 2.0, -3.0, 4.0, 5.0, -6.0, 7.0, -8.0]], dtype=torch.float32)
result = AscendSwigluStepAndMul.swiglustep_forward(x)
expected = _swiglustep_and_mul_reference(x)
assert result.shape == (1, x.shape[-1] // 2)
assert result.dtype == x.dtype
assert torch.allclose(result, expected, atol=1e-6)
def test_swiglustep_and_mul_uses_contiguous_gate_up_layout(self):
# gate = first half, up = second half (contiguous split via chunk),
# NOT the interleaved layout used by SwigluOAI.
x = torch.tensor([[1.0, 2.0, 3.0, 4.0, 10.0, 20.0, 30.0, 40.0]], dtype=torch.float32)
result = AscendSwigluStepAndMul.swiglustep_forward(x, limit=100.0)
expected = _swiglustep_and_mul_reference(x, limit=100.0)
interleaved = torch.nn.functional.silu(x[..., ::2]) * x[..., 1::2]
assert torch.allclose(result, expected, atol=1e-6)
assert not torch.allclose(result, interleaved, atol=1e-6)
def test_swiglustep_and_mul_with_custom_limit_matches_reference(self):
x = torch.tensor([[9.0, 8.0, -5.0, -9.0]], dtype=torch.float32)
limit = 3.0
result = AscendSwigluStepAndMul.swiglustep_forward(x, limit=limit)
expected = _swiglustep_and_mul_reference(x, limit=limit)
assert torch.allclose(result, expected, atol=1e-6)
def test_swiglustep_and_mul_clamps_gate_and_up_values(self):
# gate = [100, 100] -> silu(~100) clamped to 7.0;
# up = [-100, -100] -> clamped to -7.0 => 7.0 * -7.0 = -49.0
x = torch.tensor([[100.0, 100.0, -100.0, -100.0]], dtype=torch.float32)
result = AscendSwigluStepAndMul.swiglustep_forward(x)
expected = _swiglustep_and_mul_reference(x)
assert torch.allclose(result, expected, atol=1e-5)
assert torch.allclose(result, torch.tensor([[-49.0, -49.0]]), atol=1e-4)
assert not torch.isnan(result).any()
assert not torch.isinf(result).any()
def test_swiglustep_and_mul_large_input(self):
x = torch.randn(64, 128, dtype=torch.float32)
result = AscendSwigluStepAndMul.swiglustep_forward(x)
expected = _swiglustep_and_mul_reference(x)
assert result.shape == (64, 64)
assert torch.allclose(result, expected, atol=1e-5)
assert not torch.isnan(result).any()
def test_swiglustep_and_mul_validates_limit(self):
x = torch.tensor([[8.0, 9.0, -2.0, -8.0]], dtype=torch.float32)
with pytest.raises(ValueError, match="requires limit"):
AscendSwigluStepAndMul.swiglustep_forward(x, limit=None)
class TestActivationNPUPrecision:
@pytest.mark.parametrize(
"dtype,atol,rtol",
[
(torch.float32, 1e-4, 1e-4),
(torch.float16, 5e-3, 5e-3),
(torch.bfloat16, 2e-2, 2e-2),
],
)
def test_ascend_quick_gelu_matches_cpu_reference_on_npu(self, dtype, atol, rtol, default_vllm_config):
x_cpu = torch.linspace(-6, 6, steps=128, dtype=torch.float32).reshape(16, 8)
x_npu = x_cpu.to(dtype=dtype, device="npu")
result = AscendQuickGELU().forward_oot(x_npu).cpu()
expected = _quick_gelu_reference(x_cpu.to(dtype=dtype)).float()
assert torch.allclose(result.float(), expected, atol=atol, rtol=rtol)
@pytest.mark.parametrize(
"dtype,atol,rtol",
[
(torch.float32, 1e-4, 1e-4),
(torch.float16, 5e-3, 5e-3),
(torch.bfloat16, 2e-2, 2e-2),
],
)
@patch("vllm_ascend.ops.activation.get_weight_prefetch_method")
def test_ascend_silu_and_mul_matches_cpu_reference_on_npu(
self,
mock_get_weight_prefetch_method,
dtype,
atol,
rtol,
default_vllm_config,
):
weight_prefetch_method = MagicMock()
weight_prefetch_method.MLP_DOWN = "mlp_down"
mock_get_weight_prefetch_method.return_value = weight_prefetch_method
x_cpu = torch.randn(16, 16, dtype=torch.float32)
x_npu = x_cpu.to(dtype=dtype, device="npu")
result = AscendSiluAndMul().forward_oot(x_npu).cpu()
expected = _silu_and_mul_reference(x_cpu.to(dtype=dtype)).float()
assert torch.allclose(result.float(), expected, atol=atol, rtol=rtol)
@pytest.mark.parametrize(
"dtype,atol,rtol",
[
(torch.float32, 1e-5, 1e-5),
(torch.float16, 5e-3, 5e-3),
(torch.bfloat16, 2e-2, 2e-2),
],
)
def test_ascend_swiglu_oai_matches_cpu_reference_on_npu(self, dtype, atol, rtol):
x_cpu = torch.randn(16, 16, dtype=torch.float32) * 4
x_npu = x_cpu.to(dtype=dtype, device="npu")
result = AscendSwigluOAIAndMul.swiglu_oai_forward(x_npu).cpu()
expected = _swiglu_oai_reference(x_cpu.to(dtype=dtype)).float()
assert result.shape == (16, 8)
assert torch.allclose(result.float(), expected, atol=atol, rtol=rtol)
@pytest.mark.parametrize(
"dtype,atol,rtol",
[
(torch.float32, 1e-5, 1e-5),
(torch.float16, 5e-3, 5e-3),
(torch.bfloat16, 2e-2, 2e-2),
],
)
def test_swiglustep_and_mul_matches_cpu_reference_on_npu(self, dtype, atol, rtol):
x_cpu = torch.randn(16, 16, dtype=torch.float32) * 4
x_npu = x_cpu.to(dtype=dtype, device="npu")
result = AscendSwigluStepAndMul.swiglustep_forward(x_npu).cpu()
expected = _swiglustep_and_mul_reference(x_cpu.to(dtype=dtype)).float()
assert result.shape == (16, 8)
assert torch.allclose(result.float(), expected, atol=atol, rtol=rtol)

View File

@@ -0,0 +1,710 @@
#
# 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 unittest.mock import MagicMock, patch
import pytest
import torch
import torch.nn.functional as F
import torch_npu
from pytest_mock import MockerFixture
from tests.ut.base import TestBase
from vllm_ascend.ascend_forward_context import MoECommType
from vllm_ascend.ops.fused_moe.experts_selector import select_experts, zero_experts_compute
from vllm_ascend.ops.fused_moe.moe_runtime_args import MoEPrepareOutput
from vllm_ascend.utils import AscendDeviceType, adapt_patch, enable_custom_op
adapt_patch(True)
def mock_ep_and_mc2_group(mocker):
mock_group = mocker.MagicMock()
mock_group.rank_in_group = 0
mock_group.rank = 0
mock_group.world_size = 4
mock_group.device_group = "mock_group_ep"
mock_group.all_to_all = MagicMock(return_value=torch.randn(8, 8))
return mock_group
def mock_dp_and_tp_group(mocker):
mock_group = mocker.MagicMock()
mock_group.rank_in_group = 0
mock_group.world_size = 2
mock_group.device_group = "mock_group"
mock_group.all_gather = MagicMock(return_value=torch.randn(10, 32))
return mock_group
def _require_ascend_custom_op(op_name: str):
try:
custom_op_enabled = enable_custom_op()
except Exception as exc:
pytest.skip(f"requires vllm_ascend custom ops: {exc}")
if not custom_op_enabled:
pytest.skip("requires vllm_ascend custom ops")
try:
getattr(torch.ops._C_ascend, op_name)
except AttributeError:
pytest.skip(f"requires torch.ops._C_ascend.{op_name}")
def _sort_topk_by_ids(topk_weights: torch.Tensor, topk_ids: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
sorted_ids, order = torch.sort(topk_ids.to(torch.int64), dim=-1)
sorted_weights = topk_weights.gather(1, order)
return sorted_weights, sorted_ids.to(torch.int32)
@pytest.fixture(autouse=True)
def setup_vllm_config_mock(mocker: MockerFixture):
mock_hf_config = MagicMock()
mock_hf_config.model_type = "llama"
mock_model_config = MagicMock()
mock_model_config.hf_config = mock_hf_config
mock_vllm_config = MagicMock()
mock_vllm_config.model_config = mock_model_config
mock_vllm_config.parallel_config = MagicMock(tensor_parallel_size=2)
mock_vllm_config.scheduler_config = MagicMock(max_num_seqs=4)
mock_vllm_config.model_config.max_model_len = 2048
mocker.patch("vllm_ascend.ops.fused_moe.fused_moe.get_current_vllm_config", return_value=mock_vllm_config)
@pytest.fixture
def mock_dist_env(mocker: MockerFixture):
mock_moe_comm_method = MagicMock()
def mock_prepare(hidden_states, router_logits, **kwargs):
return MoEPrepareOutput(
hidden_states=hidden_states,
router_logits=router_logits,
mc2_mask=kwargs.get("mc2_mask"),
padded_hidden_states_shape=None,
pertoken_scale=None,
)
mock_moe_comm_method.prepare.side_effect = mock_prepare
mock_fused_experts_result = torch.randn(16, 2)
mock_moe_comm_method.fused_experts.return_value = mock_fused_experts_result
def mock_finalize(hidden_states, **kwargs):
return hidden_states
mock_moe_comm_method.finalize.side_effect = mock_finalize
dp_metadata = MagicMock(num_tokens_across_dp_cpu=[5, 5])
mock_weight_prefetch_method = MagicMock()
mock_forward_context_obj = MagicMock(
moe_comm_method=mock_moe_comm_method,
moe_comm_type=MoECommType.MC2,
max_tokens_across_dp=10,
dp_metadata=dp_metadata,
mc2_mask=torch.zeros(16, dtype=torch.bool),
padded_num_tokens=16,
with_quant=False,
)
with (
patch("torch.distributed.get_rank", return_value=0),
patch("torch.distributed.get_world_size", return_value=4),
patch("vllm_ascend.ops.fused_moe.fused_moe.get_ep_group", return_value=mock_ep_and_mc2_group(mocker)),
patch("vllm_ascend.ops.fused_moe.token_dispatcher.get_ep_group", return_value=mock_ep_and_mc2_group(mocker)),
patch("vllm_ascend.ops.fused_moe.fused_moe.get_mc2_group", return_value=mock_ep_and_mc2_group(mocker)),
patch("vllm_ascend.ops.fused_moe.fused_moe.get_tp_group", return_value=mock_dp_and_tp_group(mocker)),
patch("vllm.distributed.parallel_state.get_tp_group", return_value=mock_dp_and_tp_group(mocker)),
patch("vllm_ascend.ops.fused_moe.fused_moe.get_dp_group", return_value=mock_dp_and_tp_group(mocker)),
patch("vllm.model_executor.layers.fused_moe.layer.get_dp_group", return_value=mock_dp_and_tp_group(mocker)),
patch("vllm.model_executor.layers.fused_moe.config.get_dp_group", return_value=mock_dp_and_tp_group(mocker)),
patch(
"vllm_ascend.ops.fused_moe.fused_moe.get_ascend_config",
return_value=MagicMock(enable_multistream_moe=False, expert_map_path=None),
),
patch(
"vllm_ascend.ops.fused_moe.fused_moe.init_eplb_config",
return_value=(torch.tensor([0, 1, 2, -1, -1, -1, -1, -1]), None, 0),
),
patch("vllm_ascend.ops.fused_moe.fused_moe.get_forward_context", return_value=mock_forward_context_obj),
patch("vllm_ascend.ascend_forward_context.get_forward_context", return_value=mock_forward_context_obj),
patch("vllm_ascend.utils.get_ascend_device_type", return_value=AscendDeviceType.A3),
patch("vllm_ascend.ops.fused_moe.moe_comm_method.MC2CommImpl._get_token_dispatcher", return_value=None),
patch("vllm_ascend.ops.fused_moe.moe_comm_method.AlltoAllCommImpl._get_token_dispatcher", return_value=None),
patch("vllm_ascend.ops.fused_moe.moe_comm_method.AllGatherCommImpl._get_token_dispatcher", return_value=None),
patch(
"vllm_ascend.ops.fused_moe.experts_selector.get_weight_prefetch_method",
return_value=mock_weight_prefetch_method,
),
):
yield {
"mock_forward_context_obj": mock_forward_context_obj,
"mock_moe_comm_method": mock_moe_comm_method,
}
@pytest.fixture
def mock_moe_env(mocker: MockerFixture):
with (
patch(
"torch_npu.npu_moe_init_routing",
return_value=(torch.randn(8, 2), torch.randint(0, 8, (8, 2)), torch.tensor([0, 1, 2, 4, 6, 2, 7, 1])),
),
patch("torch_npu.npu_moe_compute_expert_tokens", return_value=(torch.randn(8, 2))),
patch("torch_npu.npu_moe_distribute_dispatch", return_value=(torch.randn(16, 2))),
patch("torch_npu.npu_moe_distribute_combine", return_value=(torch.randn(16, 2))),
patch("torch_npu.npu_grouped_matmul", return_value=([torch.randn(16, 2)])),
patch("torch_npu.npu_swiglu", return_value=(torch.randn(16, 2))),
patch("torch_npu.npu_moe_finalize_routing", return_value=(torch.randn(16, 2))),
):
if hasattr(torch_npu, "npu_moe_distribute_dispatch_v2"):
with (
patch("torch_npu.npu_moe_distribute_dispatch_v2", return_value=(torch.randn(16, 2))),
patch("torch_npu.npu_moe_distribute_combine_v2", return_value=(torch.randn(16, 2))),
):
yield
else:
yield
class TestExpertsSelector:
@pytest.mark.parametrize("num_experts", [256, 128])
def test_select_experts(self, mock_dist_env, mock_moe_env, num_experts):
x = torch.randn(8, 2)
router_logits = torch.randn(8, 2)
topk_weights, topk_ids = select_experts(
hidden_states=x,
router_logits=router_logits,
top_k=2,
use_grouped_topk=False,
renormalize=True,
topk_group=None,
num_expert_group=None,
custom_routing_function=None,
scoring_func="softmax",
e_score_correction_bias=None,
num_experts=num_experts,
)
assert topk_weights.shape == (8, 2)
assert topk_ids.shape == (8, 2)
@pytest.mark.parametrize("scoring_func", ["softmax", "sigmoid"])
@pytest.mark.parametrize("renormalize", [True, False])
def test_select_experts_with_different_scoring_func(self, mock_dist_env, mock_moe_env, scoring_func, renormalize):
num_tokens = 16
num_experts = 8
hidden_size = 32
hidden_states = torch.randn(num_tokens, hidden_size)
router_logits = torch.randn(num_tokens, num_experts)
def simple_custom_routing(hidden_states, gating_output, topk, renormalize):
if scoring_func == "softmax":
weights = gating_output.softmax(dim=-1)
else:
weights = gating_output.sigmoid()
topk_weights, topk_ids = weights.topk(topk, dim=-1)
if renormalize:
topk_weights = topk_weights / topk_weights.sum(dim=-1, keepdim=True)
return topk_weights, topk_ids.to(torch.int32)
topk_weights, topk_ids = select_experts(
hidden_states=hidden_states,
router_logits=router_logits,
top_k=2,
use_grouped_topk=False,
renormalize=renormalize,
topk_group=None,
num_expert_group=None,
custom_routing_function=simple_custom_routing,
scoring_func=scoring_func,
e_score_correction_bias=None,
num_experts=num_experts,
)
assert topk_weights.shape == (num_tokens, 2)
assert topk_ids.shape == (num_tokens, 2)
assert topk_weights.dtype == hidden_states.dtype
assert topk_ids.dtype == torch.int32
if renormalize:
weight_sum = topk_weights.sum(dim=-1)
torch.testing.assert_close(weight_sum, torch.ones_like(weight_sum), rtol=1e-4, atol=1e-4)
def test_select_experts_with_grouped_topk(self, mock_dist_env, mock_moe_env):
_require_ascend_custom_op("moe_gating_top_k")
num_tokens = 16
num_experts = 8
hidden_size = 32
num_expert_group = 4
topk_group = 2
hidden_states = torch.randn(num_tokens, hidden_size, device="npu")
router_logits = torch.randn(num_tokens, num_experts, device="npu")
topk_weights, topk_ids = select_experts(
hidden_states=hidden_states,
router_logits=router_logits,
top_k=2,
use_grouped_topk=True,
renormalize=True,
topk_group=topk_group,
num_expert_group=num_expert_group,
custom_routing_function=None,
scoring_func="softmax",
e_score_correction_bias=None,
num_experts=num_experts,
)
assert topk_weights.shape == (num_tokens, 2)
assert topk_ids.shape == (num_tokens, 2)
def test_select_experts_with_e_score_correction_bias(self, mock_dist_env, mock_moe_env):
_require_ascend_custom_op("moe_gating_top_k")
num_tokens = 16
num_experts = 8
hidden_size = 32
hidden_states = torch.randn(num_tokens, hidden_size, device="npu")
router_logits = torch.randn(num_tokens, num_experts, device="npu")
e_score_correction_bias = torch.randn(num_experts, device="npu")
topk_weights, topk_ids = select_experts(
hidden_states=hidden_states,
router_logits=router_logits,
top_k=2,
use_grouped_topk=False,
renormalize=True,
topk_group=None,
num_expert_group=None,
custom_routing_function=None,
scoring_func="softmax",
e_score_correction_bias=e_score_correction_bias,
num_experts=num_experts,
)
assert topk_weights.shape == (num_tokens, 2)
assert topk_ids.shape == (num_tokens, 2)
def test_select_experts_with_custom_routing_function(self, mock_dist_env, mock_moe_env):
num_tokens = 16
num_experts = 8
hidden_size = 32
hidden_states = torch.randn(num_tokens, hidden_size)
router_logits = torch.randn(num_tokens, num_experts)
def custom_routing(hidden_states, gating_output, topk, renormalize):
weights = gating_output.softmax(dim=-1)
topk_weights, topk_ids = weights.topk(topk, dim=-1)
if renormalize:
topk_weights = topk_weights / topk_weights.sum(dim=-1, keepdim=True)
return topk_weights, topk_ids.to(torch.int32)
topk_weights, topk_ids = select_experts(
hidden_states=hidden_states,
router_logits=router_logits,
top_k=2,
use_grouped_topk=False,
renormalize=True,
topk_group=None,
num_expert_group=None,
custom_routing_function=custom_routing,
scoring_func="softmax",
e_score_correction_bias=None,
num_experts=num_experts,
)
assert topk_weights.shape == (num_tokens, 2)
assert topk_ids.shape == (num_tokens, 2)
def test_select_experts_weight_sum_range(self, mock_dist_env, mock_moe_env):
num_tokens = 16
num_experts = 8
hidden_size = 32
hidden_states = torch.randn(num_tokens, hidden_size)
router_logits = torch.randn(num_tokens, num_experts)
def simple_routing(hidden_states, gating_output, topk, renormalize):
weights = gating_output.softmax(dim=-1)
topk_weights, topk_ids = weights.topk(topk, dim=-1)
if renormalize:
topk_weights = topk_weights / topk_weights.sum(dim=-1, keepdim=True)
return topk_weights, topk_ids.to(torch.int32)
topk_weights, topk_ids = select_experts(
hidden_states=hidden_states,
router_logits=router_logits,
top_k=2,
use_grouped_topk=False,
renormalize=True,
topk_group=None,
num_expert_group=None,
custom_routing_function=simple_routing,
scoring_func="softmax",
e_score_correction_bias=None,
num_experts=num_experts,
)
assert (topk_weights >= 0).all()
assert (topk_weights <= 1).all()
weight_sum = topk_weights.sum(dim=-1)
torch.testing.assert_close(weight_sum, torch.ones_like(weight_sum), rtol=1e-4, atol=1e-4)
def test_select_experts_expert_id_range(self, mock_dist_env, mock_moe_env):
num_tokens = 16
num_experts = 8
hidden_size = 32
hidden_states = torch.randn(num_tokens, hidden_size)
router_logits = torch.randn(num_tokens, num_experts)
def simple_routing(hidden_states, gating_output, topk, renormalize):
weights = gating_output.softmax(dim=-1)
topk_weights, topk_ids = weights.topk(topk, dim=-1)
if renormalize:
topk_weights = topk_weights / topk_weights.sum(dim=-1, keepdim=True)
return topk_weights, topk_ids.to(torch.int32)
topk_weights, topk_ids = select_experts(
hidden_states=hidden_states,
router_logits=router_logits,
top_k=2,
use_grouped_topk=False,
renormalize=True,
topk_group=None,
num_expert_group=None,
custom_routing_function=simple_routing,
scoring_func="softmax",
e_score_correction_bias=None,
num_experts=num_experts,
)
assert (topk_ids >= 0).all()
assert (topk_ids < num_experts).all()
@patch("vllm_ascend.ops.fused_moe.experts_selector.get_weight_prefetch_method")
@patch("vllm_ascend.ops.fused_moe.experts_selector.check_npu_moe_gating_top_k", return_value=False)
def test_select_experts_native_softmax_matches_expected(self, _, mock_get_weight_prefetch_method):
hidden_states = torch.tensor([[1.0, 0.0, -1.0, 2.0], [0.5, 1.5, -0.5, 1.0]], dtype=torch.float32)
router_logits = torch.tensor([[3.0, 1.0, 0.0, 2.0], [0.0, 4.0, 1.0, 2.0]], dtype=torch.float32)
prefetch_method = MagicMock()
mock_get_weight_prefetch_method.return_value = prefetch_method
topk_weights, topk_ids = select_experts(
hidden_states=hidden_states,
router_logits=router_logits,
top_k=2,
use_grouped_topk=False,
renormalize=True,
topk_group=None,
num_expert_group=None,
custom_routing_function=None,
scoring_func="softmax",
e_score_correction_bias=None,
num_experts=4,
)
expected_probs = router_logits.softmax(dim=-1)
expected_weights, expected_ids = expected_probs.topk(2, dim=-1)
expected_weights = expected_weights / expected_weights.sum(dim=-1, keepdim=True)
prefetch_method.maybe_prefetch_moe_weight_preprocess.assert_called_once_with(hidden_states, "gate_up")
torch.testing.assert_close(topk_weights, expected_weights.to(hidden_states.dtype))
assert torch.equal(topk_ids, expected_ids.to(torch.int32))
def test_select_experts_grouped_topk_bias_uses_original_weights(self):
_require_ascend_custom_op("moe_gating_top_k")
hidden_states = torch.tensor([[1.0, 0.0, -1.0, 2.0]], dtype=torch.float32, device="npu")
router_logits = torch.tensor([[4.0, 3.0, 1.0, 0.0]], dtype=torch.float32, device="npu")
e_score_correction_bias = torch.tensor([-10.0, -10.0, 5.0, 5.0], dtype=torch.float32, device="npu")
topk_weights, topk_ids = select_experts(
hidden_states=hidden_states,
router_logits=router_logits,
top_k=2,
use_grouped_topk=True,
renormalize=False,
topk_group=1,
num_expert_group=2,
custom_routing_function=None,
scoring_func="softmax",
e_score_correction_bias=e_score_correction_bias,
num_experts=4,
)
original_weights = router_logits.softmax(dim=-1)
actual_weights, actual_ids = _sort_topk_by_ids(topk_weights.cpu(), topk_ids.cpu())
expected_ids = torch.tensor([[2, 3]], dtype=torch.int32)
expected_weights = original_weights[:, 2:4].to("cpu")
assert torch.equal(actual_ids, expected_ids)
torch.testing.assert_close(actual_weights, expected_weights)
@patch("vllm_ascend.ops.fused_moe.experts_selector.get_weight_prefetch_method", return_value=MagicMock())
@patch("vllm_ascend.ops.fused_moe.experts_selector.check_npu_moe_gating_top_k", return_value=False)
def test_select_experts_invalid_scoring_func_raises(self, _, __):
hidden_states = torch.randn(2, 4)
router_logits = torch.randn(2, 4)
with pytest.raises(ValueError, match="Unsupported scoring function"):
select_experts(
hidden_states=hidden_states,
router_logits=router_logits,
top_k=2,
use_grouped_topk=False,
renormalize=True,
topk_group=None,
num_expert_group=None,
custom_routing_function=None,
scoring_func="unsupported",
e_score_correction_bias=None,
num_experts=4,
)
class TestFusedMoENPUOpsAccuracy:
def test_moe_gating_top_k_matches_native_softmax(self, mock_dist_env):
_require_ascend_custom_op("moe_gating_top_k")
hidden_states = torch.randn(4, 16, dtype=torch.float32, device="npu")
router_logits_cpu = torch.tensor(
[[3.0, 1.0, 0.0, 2.0], [0.0, 4.0, 1.0, 2.0], [1.0, 0.5, 5.0, -1.0], [2.5, -0.5, 1.5, 0.25]],
dtype=torch.float32,
)
router_logits = router_logits_cpu.to("npu")
topk_weights, topk_ids = select_experts(
hidden_states=hidden_states,
router_logits=router_logits,
top_k=2,
use_grouped_topk=False,
renormalize=True,
topk_group=None,
num_expert_group=None,
custom_routing_function=None,
scoring_func="softmax",
e_score_correction_bias=None,
num_experts=4,
)
expected_weights, expected_ids = router_logits_cpu.softmax(dim=-1).topk(2, dim=-1)
expected_weights = expected_weights / expected_weights.sum(dim=-1, keepdim=True)
actual_weights, actual_ids = _sort_topk_by_ids(topk_weights.cpu(), topk_ids.cpu())
expected_weights, expected_ids = _sort_topk_by_ids(expected_weights, expected_ids.to(torch.int32))
assert torch.equal(actual_ids, expected_ids)
torch.testing.assert_close(actual_weights, expected_weights, rtol=1e-3, atol=1e-3)
def test_npu_swiglu_matches_torch_reference(self):
gate_up_cpu = torch.tensor(
[[1.0, -2.0, 0.5, 3.0], [-0.5, 2.0, 1.5, -1.0], [4.0, -3.0, -2.0, 0.25]],
dtype=torch.float16,
)
gate_up = gate_up_cpu.to("npu")
actual = torch_npu.npu_swiglu(gate_up)
left, right = gate_up_cpu.chunk(2, dim=-1)
expected = F.silu(left.float()) * right.float()
torch.testing.assert_close(actual.cpu().float(), expected, rtol=2e-3, atol=2e-3)
def test_npu_grouped_matmul_matches_torch_reference(self):
hidden_states_cpu = torch.tensor(
[[1.0, 2.0, -1.0], [0.5, -0.5, 1.0], [2.0, 0.0, 1.0], [-1.0, 3.0, 0.5]],
dtype=torch.float16,
)
weight_cpu = torch.tensor(
[[[1.0, 0.0], [0.5, -1.0], [-0.5, 2.0]], [[-1.0, 1.5], [2.0, 0.25], [0.0, -0.5]]],
dtype=torch.float16,
)
group_list = torch.tensor([2, 2], dtype=torch.int64, device="npu")
actual = torch_npu.npu_grouped_matmul(
x=[hidden_states_cpu.to("npu")],
weight=[weight_cpu.to("npu")],
split_item=3,
group_list_type=1,
group_type=0,
group_list=group_list,
output_dtype=torch.float16,
)[0]
expected = torch.cat(
[
hidden_states_cpu[:2].float() @ weight_cpu[0].float(),
hidden_states_cpu[2:].float() @ weight_cpu[1].float(),
],
dim=0,
)
torch.testing.assert_close(actual.cpu().float(), expected, rtol=2e-3, atol=2e-3)
class TestZeroExpertsCompute(TestBase):
def test_zero_experts_compute_identity_type(self):
num_experts = 8
num_tokens = 4
top_k = 2
hidden_size = 16
expert_indices = torch.tensor([[0, 1], [2, 3], [4, 5], [6, 7]], dtype=torch.int32)
expert_scales = torch.ones(num_tokens, top_k, dtype=torch.float32)
hidden_states = torch.randn(num_tokens, hidden_size, dtype=torch.float32)
result_indices, result_scales, result_hidden = zero_experts_compute(
expert_indices=expert_indices,
expert_scales=expert_scales,
num_experts=num_experts,
zero_expert_type="identity",
hidden_states=hidden_states,
)
assert result_indices.shape == expert_indices.shape
assert result_scales.shape == expert_scales.shape
assert result_hidden.shape == hidden_states.shape
def test_zero_experts_compute_with_zero_experts(self):
num_experts = 4
num_tokens = 3
hidden_size = 8
expert_indices = torch.tensor([[0, 1], [1, 2], [2, 3]], dtype=torch.int32)
expert_scales = torch.tensor([[0.5, 0.5], [0.3, 0.7], [0.4, 0.6]], dtype=torch.float32)
hidden_states = torch.randn(num_tokens, hidden_size, dtype=torch.float32)
result_indices, result_scales, result_hidden = zero_experts_compute(
expert_indices=expert_indices,
expert_scales=expert_scales,
num_experts=num_experts,
zero_expert_type="identity",
hidden_states=hidden_states,
)
assert result_indices.shape == expert_indices.shape
assert result_scales.shape == expert_scales.shape
assert result_hidden.shape == hidden_states.shape
def test_zero_experts_compute_normal_experts_masked(self):
num_experts = 4
num_tokens = 2
top_k = 2
hidden_size = 8
expert_indices = torch.tensor([[0, 5], [6, 7]], dtype=torch.int32)
expert_scales = torch.ones(num_tokens, top_k, dtype=torch.float32)
hidden_states = torch.randn(num_tokens, hidden_size, dtype=torch.float32)
result_indices, result_scales, _ = zero_experts_compute(
expert_indices=expert_indices,
expert_scales=expert_scales,
num_experts=num_experts,
zero_expert_type="identity",
hidden_states=hidden_states,
)
normal_expert_mask = expert_indices >= num_experts
for i in range(num_tokens):
for j in range(top_k):
if normal_expert_mask[i, j]:
assert result_scales[i, j] == 0.0
assert result_indices[i, j] == 0
def test_zero_experts_compute_output_sum(self):
num_experts = 2
num_tokens = 2
hidden_size = 4
expert_indices = torch.tensor([[0, 1], [0, 1]], dtype=torch.int32)
expert_scales = torch.tensor([[0.5, 0.5], [0.3, 0.7]], dtype=torch.float32)
hidden_states = torch.randn(num_tokens, hidden_size, dtype=torch.float32)
_, _, result_hidden = zero_experts_compute(
expert_indices=expert_indices,
expert_scales=expert_scales,
num_experts=num_experts,
zero_expert_type="identity",
hidden_states=hidden_states,
)
assert result_hidden.shape == (num_tokens, hidden_size)
def test_zero_experts_compute_all_zero_experts(self):
num_experts = 4
num_tokens = 2
top_k = 2
hidden_size = 8
expert_indices = torch.tensor([[0, 1], [2, 3]], dtype=torch.int32)
expert_scales = torch.ones(num_tokens, top_k, dtype=torch.float32)
hidden_states = torch.randn(num_tokens, hidden_size, dtype=torch.float32)
result_indices, result_scales, result_hidden = zero_experts_compute(
expert_indices=expert_indices,
expert_scales=expert_scales,
num_experts=num_experts,
zero_expert_type="identity",
hidden_states=hidden_states,
)
assert torch.equal(result_indices, expert_indices)
assert torch.equal(result_scales, expert_scales)
assert result_hidden.shape == hidden_states.shape
def test_zero_experts_compute_mixed_experts(self):
num_experts = 3
num_tokens = 2
hidden_size = 8
expert_indices = torch.tensor([[0, 1, 4], [2, 5, 6]], dtype=torch.int32)
expert_scales = torch.tensor([[0.3, 0.3, 0.4], [0.2, 0.5, 0.3]], dtype=torch.float32)
hidden_states = torch.randn(num_tokens, hidden_size, dtype=torch.float32)
result_indices, result_scales, _ = zero_experts_compute(
expert_indices=expert_indices,
expert_scales=expert_scales,
num_experts=num_experts,
zero_expert_type="identity",
hidden_states=hidden_states,
)
assert result_indices.shape == expert_indices.shape
assert result_scales.shape == expert_scales.shape
def test_zero_experts_compute_identity_values_match_expected(self):
expert_indices = torch.tensor([[0, 2], [3, 1]], dtype=torch.int32)
expert_scales = torch.tensor([[0.25, 0.75], [0.60, 0.40]], dtype=torch.float32)
hidden_states = torch.tensor([[1.0, 2.0], [3.0, 4.0]], dtype=torch.float32)
result_indices, result_scales, result_hidden = zero_experts_compute(
expert_indices=expert_indices,
expert_scales=expert_scales,
num_experts=2,
zero_expert_type="identity",
hidden_states=hidden_states,
)
expected_indices = torch.tensor([[0, 0], [0, 1]], dtype=torch.int32)
expected_scales = torch.tensor([[0.25, 0.0], [0.0, 0.40]], dtype=torch.float32)
expected_hidden = torch.tensor([[0.75, 1.50], [1.80, 2.40]], dtype=torch.float32)
assert torch.equal(result_indices, expected_indices)
torch.testing.assert_close(result_scales, expected_scales)
torch.testing.assert_close(result_hidden, expected_hidden)

View File

@@ -20,13 +20,14 @@ import torch
from pytest_mock import MockerFixture
from tests.ut.base import PytestBase
from vllm_ascend.ops.moe.comm_utils import (
_gather_along_first_dim, async_all_to_all,
gather_from_sequence_parallel_region)
from vllm_ascend.ops.fused_moe.comm_utils import (
_gather_along_first_dim,
async_all_to_all,
gather_from_sequence_parallel_region,
)
class TestDistributedCommunication(PytestBase):
@pytest.fixture(autouse=True)
def context(self, mocker: MockerFixture):
mocker.patch("torch.npu.current_device", return_value="cpu")
@@ -36,63 +37,52 @@ class TestDistributedCommunication(PytestBase):
@pytest.mark.parametrize(
"input_tensor, output_split_sizes, input_split_sizes",
[(torch.randn(8, 16), [2, 2, 2, 2], [2, 2, 2, 2]),
(torch.randn(16, 32), None, None)])
def test_async_all_to_all(self, input_tensor, output_split_sizes,
input_split_sizes, mocker: MockerFixture):
[(torch.randn(8, 16), [2, 2, 2, 2], [2, 2, 2, 2]), (torch.randn(16, 32), None, None)],
)
def test_async_all_to_all(self, input_tensor, output_split_sizes, input_split_sizes, mocker: MockerFixture):
"""Test async_all_to_all"""
mock_group = mocker.MagicMock()
mocker.patch("torch.distributed.all_to_all_single",
return_value=mocker.MagicMock())
mocker.patch("torch.distributed.all_to_all_single", return_value=mocker.MagicMock())
_, a2a_out, handle = async_all_to_all(input_tensor, output_split_sizes,
input_split_sizes, mock_group)
_, a2a_out, handle = async_all_to_all(input_tensor, output_split_sizes, input_split_sizes, mock_group)
# Check if the output tensor is created properly
if output_split_sizes is None:
assert a2a_out.shape == input_tensor.shape
else:
total_output_size = sum(output_split_sizes)
expected_shape = [total_output_size] + list(
input_tensor.size())[1:]
expected_shape = [total_output_size] + list(input_tensor.size())[1:]
assert a2a_out.shape == torch.Size(expected_shape)
# Ensure handle is returned from async operation
assert handle is not None
assert isinstance(handle, mocker.MagicMock)
@pytest.mark.parametrize("world_size, test_tensor, expected",
[(1, torch.randn(8, 16), (8, 16)),
(4, torch.randn(8, 16), (32, 16))])
def test_gather_along_first_dim(self, test_tensor, expected, world_size,
mocker: MockerFixture):
@pytest.mark.parametrize(
"world_size, test_tensor, expected", [(1, torch.randn(8, 16), (8, 16)), (4, torch.randn(8, 16), (32, 16))]
)
def test_gather_along_first_dim(self, test_tensor, expected, world_size, mocker: MockerFixture):
"""Test _gather_along_first_dim"""
mocker.patch("torch.distributed.get_world_size",
return_value=world_size)
mocker.patch("torch.distributed.get_world_size", return_value=world_size)
result = _gather_along_first_dim(test_tensor, mocker.MagicMock())
assert result.shape == expected
@pytest.mark.parametrize("input_tensor, output_split_sizes",
[(torch.randn(8, 16), None),
(torch.randn(8, 16), [2, 2, 2, 2])])
def test_gather_from_sequence_parallel_region(self, input_tensor,
output_split_sizes,
mocker: MockerFixture):
@pytest.mark.parametrize(
"input_tensor, output_split_sizes", [(torch.randn(8, 16), None), (torch.randn(8, 16), [2, 2, 2, 2])]
)
def test_gather_from_sequence_parallel_region(self, input_tensor, output_split_sizes, mocker: MockerFixture):
"""Test gather_from_sequence_parallel_region"""
mock_group = mocker.MagicMock()
result = gather_from_sequence_parallel_region(input_tensor, mock_group,
output_split_sizes)
result = gather_from_sequence_parallel_region(input_tensor, mock_group, output_split_sizes)
# If output_split_sizes is not provided, result should have expanded first dimension by world size
if output_split_sizes is None:
expected_shape = [input_tensor.shape[0] * 4] + list(
input_tensor.shape[1:])
expected_shape = [input_tensor.shape[0] * 4] + list(input_tensor.shape[1:])
assert result.shape == torch.Size(expected_shape)
else:
# If output_split_sizes is provided, result shape is dictated by sum of output_split_sizes
expected_shape = [sum(output_split_sizes)] + list(
input_tensor.shape[1:])
expected_shape = [sum(output_split_sizes)] + list(input_tensor.shape[1:])
assert result.shape == torch.Size(expected_shape)

View File

@@ -0,0 +1,487 @@
#
# 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 unittest.mock import MagicMock, patch
import pytest
from vllm_ascend.ops.flashcomm2_oshard_manager import (
Flashcomm2OShardManager,
flashcomm2_oshard_manager,
)
class TestFlashcomm2OShardManager:
@pytest.fixture
def manager(self):
return Flashcomm2OShardManager()
def test_init(self, manager):
assert manager._shard_layers == {}
assert isinstance(manager._shard_layers, dict)
@patch("vllm_ascend.ops.flashcomm2_oshard_manager.flashcomm2_enable")
@patch("vllm_ascend.ops.flashcomm2_oshard_manager.o_shard_enable")
def test_flashcomm2_oshard_enable_both_enabled(
self,
mock_o_shard_enable,
mock_flashcomm2_enable,
manager,
):
mock_flashcomm2_enable.return_value = True
mock_o_shard_enable.return_value = True
result = manager.flashcomm2_oshard_enable()
assert result is True
@patch("vllm_ascend.ops.flashcomm2_oshard_manager.flashcomm2_enable")
@patch("vllm_ascend.ops.flashcomm2_oshard_manager.o_shard_enable")
def test_flashcomm2_oshard_enable_flashcomm2_disabled(
self,
mock_o_shard_enable,
mock_flashcomm2_enable,
manager,
):
mock_flashcomm2_enable.return_value = False
mock_o_shard_enable.return_value = True
result = manager.flashcomm2_oshard_enable()
assert result is False
@patch("vllm_ascend.ops.flashcomm2_oshard_manager.flashcomm2_enable")
@patch("vllm_ascend.ops.flashcomm2_oshard_manager.o_shard_enable")
def test_flashcomm2_oshard_enable_o_shard_disabled(
self,
mock_o_shard_enable,
mock_flashcomm2_enable,
manager,
):
mock_flashcomm2_enable.return_value = True
mock_o_shard_enable.return_value = False
result = manager.flashcomm2_oshard_enable()
assert result is False
@patch("vllm_ascend.ops.flashcomm2_oshard_manager.flashcomm2_enable")
@patch("vllm_ascend.ops.flashcomm2_oshard_manager.o_shard_enable")
def test_flashcomm2_oshard_enable_both_disabled(
self,
mock_o_shard_enable,
mock_flashcomm2_enable,
manager,
):
mock_flashcomm2_enable.return_value = False
mock_o_shard_enable.return_value = False
result = manager.flashcomm2_oshard_enable()
assert result is False
@patch("vllm_ascend.ops.flashcomm2_oshard_manager.register_layer_to_shard_weight_series")
@patch("vllm_ascend.ops.flashcomm2_oshard_manager.get_shard_weight_group")
@patch("vllm_ascend.ops.flashcomm2_oshard_manager.is_hidden_layer")
@patch("vllm_ascend.ops.flashcomm2_oshard_manager.extract_layer_index")
def test_register_layer_hidden_layer(
self,
mock_extract_index,
mock_is_hidden,
mock_get_group,
mock_register,
manager,
):
mock_is_hidden.return_value = True
mock_extract_index.return_value = 5
mock_group = MagicMock()
mock_get_group.return_value = mock_group
layer = MagicMock()
layer.prefix = "model.layers.5.self_attn.o_proj"
manager.register_layer(layer, prefetch_step=2)
assert 5 in manager._shard_layers
assert manager._shard_layers[5] is layer
mock_register.assert_called_once_with(
series_name="o_proj",
group=mock_group,
layer=layer,
prefetch_step=2,
)
@patch("vllm_ascend.ops.flashcomm2_oshard_manager.register_layer_to_shard_weight_series")
@patch("vllm_ascend.ops.flashcomm2_oshard_manager.get_shard_weight_group")
@patch("vllm_ascend.ops.flashcomm2_oshard_manager.is_hidden_layer")
def test_register_layer_non_hidden_layer(
self,
mock_is_hidden,
mock_get_group,
mock_register,
manager,
):
mock_is_hidden.return_value = False
layer = MagicMock()
layer.prefix = "model.layers.100.self_attn.o_proj"
manager.register_layer(layer)
assert len(manager._shard_layers) == 0
mock_register.assert_not_called()
@patch("vllm_ascend.ops.flashcomm2_oshard_manager.register_layer_to_shard_weight_series")
@patch("vllm_ascend.ops.flashcomm2_oshard_manager.get_shard_weight_group")
@patch("vllm_ascend.ops.flashcomm2_oshard_manager.is_hidden_layer")
@patch("vllm_ascend.ops.flashcomm2_oshard_manager.extract_layer_index")
def test_register_layer_default_prefetch_step(
self,
mock_extract_index,
mock_is_hidden,
mock_get_group,
mock_register,
manager,
):
mock_is_hidden.return_value = True
mock_extract_index.return_value = 0
mock_group = MagicMock()
mock_get_group.return_value = mock_group
layer = MagicMock()
layer.prefix = "model.layers.0.self_attn.o_proj"
manager.register_layer(layer)
mock_register.assert_called_once_with(
series_name="o_proj",
group=mock_group,
layer=layer,
prefetch_step=1,
)
@patch("vllm_ascend.ops.flashcomm2_oshard_manager.register_layer_to_shard_weight_series")
@patch("vllm_ascend.ops.flashcomm2_oshard_manager.get_shard_weight_group")
@patch("vllm_ascend.ops.flashcomm2_oshard_manager.is_hidden_layer")
@patch("vllm_ascend.ops.flashcomm2_oshard_manager.extract_layer_index")
def test_register_layer_overwrites_existing_layer_with_same_index(
self,
mock_extract_index,
mock_is_hidden,
mock_get_group,
mock_register,
manager,
):
mock_is_hidden.return_value = True
mock_extract_index.return_value = 5
mock_group = MagicMock()
mock_get_group.return_value = mock_group
first_layer = MagicMock()
first_layer.prefix = "model.layers.5.self_attn.o_proj"
second_layer = MagicMock()
second_layer.prefix = "model.layers.5.self_attn.o_proj"
manager.register_layer(first_layer)
manager.register_layer(second_layer)
assert len(manager._shard_layers) == 1
assert manager.get_layer(5) is second_layer
assert mock_register.call_count == 2
@patch("vllm_ascend.ops.flashcomm2_oshard_manager.register_layer_to_shard_weight_series")
@patch("vllm_ascend.ops.flashcomm2_oshard_manager.get_shard_weight_group")
@patch("vllm_ascend.ops.flashcomm2_oshard_manager.is_hidden_layer")
def test_register_layer_missing_prefix_raises(
self,
mock_is_hidden,
mock_get_group,
mock_register,
manager,
):
mock_is_hidden.return_value = True
layer = MagicMock(spec=[])
with pytest.raises(AttributeError):
manager.register_layer(layer)
assert manager._shard_layers == {}
mock_get_group.assert_not_called()
mock_register.assert_not_called()
@patch("vllm_ascend.ops.flashcomm2_oshard_manager.register_layer_to_shard_weight_series")
@patch("vllm_ascend.ops.flashcomm2_oshard_manager.get_shard_weight_group")
@patch("vllm_ascend.ops.flashcomm2_oshard_manager.is_hidden_layer")
@patch("vllm_ascend.ops.flashcomm2_oshard_manager.extract_layer_index")
def test_register_layer_extract_layer_index_failure_propagates(
self,
mock_extract_index,
mock_is_hidden,
mock_get_group,
mock_register,
manager,
):
mock_is_hidden.return_value = True
mock_extract_index.side_effect = ValueError("invalid layer prefix")
layer = MagicMock()
layer.prefix = "invalid-prefix"
with pytest.raises(ValueError, match="invalid layer prefix"):
manager.register_layer(layer)
assert manager._shard_layers == {}
mock_get_group.assert_not_called()
mock_register.assert_not_called()
def test_get_layer_existing(self, manager):
layer = MagicMock()
manager._shard_layers[5] = layer
result = manager.get_layer(5)
assert result is layer
def test_get_layer_non_existing(self, manager):
result = manager.get_layer(999)
assert result is None
def test_get_layer_empty_dict(self, manager):
result = manager.get_layer(0)
assert result is None
@patch("vllm_ascend.ops.flashcomm2_oshard_manager.reach_layer_for_shard_weight_series")
@patch("vllm_ascend.ops.flashcomm2_oshard_manager.is_hidden_layer")
@patch("vllm_ascend.ops.flashcomm2_oshard_manager.extract_layer_index")
def test_trigger_broadcast_for_layer_success(
self,
mock_extract_index,
mock_is_hidden,
mock_reach_layer,
manager,
):
mock_extract_index.return_value = 3
mock_is_hidden.return_value = True
layer = MagicMock()
manager._shard_layers[3] = layer
manager.trigger_broadcast_for_layer("model.layers.3.self_attn.o_proj")
mock_reach_layer.assert_called_once_with(layer)
@patch("vllm_ascend.ops.flashcomm2_oshard_manager.reach_layer_for_shard_weight_series")
@patch("vllm_ascend.ops.flashcomm2_oshard_manager.is_hidden_layer")
@patch("vllm_ascend.ops.flashcomm2_oshard_manager.extract_layer_index")
def test_trigger_broadcast_for_layer_not_hidden(
self,
mock_extract_index,
mock_is_hidden,
mock_reach_layer,
manager,
):
mock_extract_index.return_value = 3
mock_is_hidden.return_value = False
layer = MagicMock()
manager._shard_layers[3] = layer
manager.trigger_broadcast_for_layer("model.layers.3.self_attn.o_proj")
mock_reach_layer.assert_not_called()
@patch("vllm_ascend.ops.flashcomm2_oshard_manager.reach_layer_for_shard_weight_series")
@patch("vllm_ascend.ops.flashcomm2_oshard_manager.extract_layer_index")
def test_trigger_broadcast_for_layer_not_registered(
self,
mock_extract_index,
mock_reach_layer,
manager,
):
mock_extract_index.return_value = 999
manager.trigger_broadcast_for_layer("model.layers.999.self_attn.o_proj")
mock_reach_layer.assert_not_called()
@patch("vllm_ascend.ops.flashcomm2_oshard_manager.reach_layer_for_shard_weight_series")
@patch("vllm_ascend.ops.flashcomm2_oshard_manager.is_hidden_layer")
@patch("vllm_ascend.ops.flashcomm2_oshard_manager.extract_layer_index")
def test_trigger_broadcast_for_layer_not_registered_short_circuits_hidden_check(
self,
mock_extract_index,
mock_is_hidden,
mock_reach_layer,
manager,
):
mock_extract_index.return_value = 999
manager.trigger_broadcast_for_layer("model.layers.999.self_attn.o_proj")
mock_is_hidden.assert_not_called()
mock_reach_layer.assert_not_called()
@patch("vllm_ascend.ops.flashcomm2_oshard_manager.reach_layer_for_shard_weight_series")
@patch("vllm_ascend.ops.flashcomm2_oshard_manager.is_hidden_layer")
@patch("vllm_ascend.ops.flashcomm2_oshard_manager.extract_layer_index")
def test_trigger_broadcast_for_layer_empty_manager(
self,
mock_extract_index,
mock_is_hidden,
mock_reach_layer,
manager,
):
mock_extract_index.return_value = 0
manager.trigger_broadcast_for_layer("model.layers.0.self_attn.o_proj")
mock_is_hidden.assert_not_called()
mock_reach_layer.assert_not_called()
@patch("vllm_ascend.ops.flashcomm2_oshard_manager.reach_layer_for_shard_weight_series")
@patch("vllm_ascend.ops.flashcomm2_oshard_manager.extract_layer_index")
def test_trigger_broadcast_for_layer_extract_layer_index_failure_propagates(
self,
mock_extract_index,
mock_reach_layer,
manager,
):
mock_extract_index.side_effect = ValueError("invalid layer prefix")
with pytest.raises(ValueError, match="invalid layer prefix"):
manager.trigger_broadcast_for_layer("invalid-prefix")
mock_reach_layer.assert_not_called()
@patch("vllm_ascend.ops.flashcomm2_oshard_manager.post_process_after_loading_for_shard_weight_series")
def test_post_process_after_loading_with_layers(self, mock_post_process, manager):
layer1 = MagicMock()
layer2 = MagicMock()
manager._shard_layers[0] = layer1
manager._shard_layers[1] = layer2
manager.post_process_after_loading()
mock_post_process.assert_called_once()
called_layer = mock_post_process.call_args[0][0]
assert called_layer in [layer1, layer2]
@patch("vllm_ascend.ops.flashcomm2_oshard_manager.post_process_after_loading_for_shard_weight_series")
def test_post_process_after_loading_uses_first_registered_layer(self, mock_post_process, manager):
first_layer = MagicMock()
second_layer = MagicMock()
manager._shard_layers[1] = first_layer
manager._shard_layers[2] = second_layer
manager.post_process_after_loading()
mock_post_process.assert_called_once_with(first_layer)
@patch("vllm_ascend.ops.flashcomm2_oshard_manager.post_process_after_loading_for_shard_weight_series")
def test_post_process_after_loading_empty(self, mock_post_process, manager):
manager.post_process_after_loading()
mock_post_process.assert_not_called()
@patch("vllm_ascend.ops.flashcomm2_oshard_manager.post_process_after_loading_for_shard_weight_series")
def test_post_process_after_loading_single_layer(self, mock_post_process, manager):
layer = MagicMock()
manager._shard_layers[5] = layer
manager.post_process_after_loading()
mock_post_process.assert_called_once_with(layer)
class TestGlobalInstance:
def test_global_instance_exists(self):
assert flashcomm2_oshard_manager is not None
assert isinstance(flashcomm2_oshard_manager, Flashcomm2OShardManager)
def test_global_instance_has_shard_layers(self):
assert hasattr(flashcomm2_oshard_manager, "_shard_layers")
assert isinstance(flashcomm2_oshard_manager._shard_layers, dict)
class TestIntegration:
@patch("vllm_ascend.ops.flashcomm2_oshard_manager.register_layer_to_shard_weight_series")
@patch("vllm_ascend.ops.flashcomm2_oshard_manager.get_shard_weight_group")
@patch("vllm_ascend.ops.flashcomm2_oshard_manager.is_hidden_layer")
@patch("vllm_ascend.ops.flashcomm2_oshard_manager.extract_layer_index")
@patch("vllm_ascend.ops.flashcomm2_oshard_manager.reach_layer_for_shard_weight_series")
def test_full_workflow(
self,
mock_reach_layer,
mock_extract_index,
mock_is_hidden,
mock_get_group,
mock_register,
):
manager = Flashcomm2OShardManager()
mock_is_hidden.return_value = True
mock_extract_index.side_effect = lambda x: int(x.split(".")[2])
mock_group = MagicMock()
mock_get_group.return_value = mock_group
layer0 = MagicMock()
layer0.prefix = "model.layers.0.self_attn.o_proj"
layer1 = MagicMock()
layer1.prefix = "model.layers.1.self_attn.o_proj"
manager.register_layer(layer0)
manager.register_layer(layer1)
assert len(manager._shard_layers) == 2
assert manager.get_layer(0) is layer0
assert manager.get_layer(1) is layer1
manager.trigger_broadcast_for_layer("model.layers.0.self_attn.o_proj")
mock_reach_layer.assert_called_with(layer0)
manager.trigger_broadcast_for_layer("model.layers.1.self_attn.o_proj")
mock_reach_layer.assert_called_with(layer1)
@patch("vllm_ascend.ops.flashcomm2_oshard_manager.register_layer_to_shard_weight_series")
@patch("vllm_ascend.ops.flashcomm2_oshard_manager.get_shard_weight_group")
@patch("vllm_ascend.ops.flashcomm2_oshard_manager.is_hidden_layer")
@patch("vllm_ascend.ops.flashcomm2_oshard_manager.extract_layer_index")
@patch("vllm_ascend.ops.flashcomm2_oshard_manager.post_process_after_loading_for_shard_weight_series")
def test_register_and_post_process_workflow(
self,
mock_post_process,
mock_extract_index,
mock_is_hidden,
mock_get_group,
mock_register,
):
manager = Flashcomm2OShardManager()
mock_is_hidden.return_value = True
mock_extract_index.return_value = 0
mock_group = MagicMock()
mock_get_group.return_value = mock_group
layer = MagicMock()
layer.prefix = "model.layers.0.self_attn.o_proj"
manager.register_layer(layer)
manager.post_process_after_loading()
mock_post_process.assert_called_once()

View File

@@ -0,0 +1,950 @@
#
# 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 ast
import inspect
import textwrap
from types import SimpleNamespace
from typing import TypedDict
from unittest.mock import MagicMock, patch
import pytest
import torch
import torch.nn as nn
import torch.nn.functional as F
from pytest_mock import MockerFixture
from vllm_ascend.ascend_forward_context import MoECommType
from vllm_ascend.ops.fused_moe import fused_moe as fused_moe_module
from vllm_ascend.ops.fused_moe.moe_comm_method import FusedExpertsResult
from vllm_ascend.ops.fused_moe.moe_runtime_args import (
MoEMlpComputeInput,
MoEPrepareOutput,
MoEQuantParams,
MoEWeights,
)
from vllm_ascend.quantization.quant_type import QuantType
from vllm_ascend.utils import AscendDeviceType, adapt_patch, vllm_version_is
if vllm_version_is("0.23.0"):
from vllm_ascend.ops.fused_moe import fused_moe_0_23_0 as fused_moe_legacy_module
from vllm_ascend.ops.fused_moe.fused_moe import (
AscendFusedMoE,
AscendMoERunner,
AscendUnquantizedFusedMoEMethod,
)
adapt_patch(True)
else:
pytest.skip(
"Legacy AscendFusedMoE UTs are only for vLLM 0.23.0.",
allow_module_level=True,
)
def mock_ep_and_mc2_group(mocker):
mock_group = mocker.MagicMock()
mock_group.rank_in_group = 0
mock_group.rank = 0
mock_group.world_size = 4
mock_group.device_group = "mock_group_ep"
mock_group.all_to_all = MagicMock(return_value=torch.randn(8, 8))
return mock_group
def mock_dp_and_tp_group(mocker):
mock_group = mocker.MagicMock()
mock_group.rank_in_group = 0
mock_group.world_size = 2
mock_group.device_group = "mock_group"
mock_group.all_gather = MagicMock(return_value=torch.randn(10, 32))
return mock_group
def mock_npu_format_cast(weight_data, format):
return weight_data
def build_mlp_compute_input_fixture(
*,
hidden_states: torch.Tensor,
w1: torch.Tensor | list[torch.Tensor],
w2: torch.Tensor | list[torch.Tensor],
group_list: torch.Tensor,
with_quant: bool,
group_list_type: int = 1,
dynamic_scale: torch.Tensor | None = None,
topk_scales: 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 | None = None,
w2_scale_bias: torch.Tensor | None = None,
w1_offset: torch.Tensor | None = None,
w2_offset: torch.Tensor | None = None,
fusion: bool = False,
activation: str = "silu",
need_trans: bool = True,
dynamic_eplb: bool = False,
) -> MoEMlpComputeInput:
return MoEMlpComputeInput(
hidden_states=hidden_states,
group_list=group_list,
group_list_type=group_list_type,
dynamic_scale=dynamic_scale,
topk_scales=topk_scales,
weights=MoEWeights(
w1=w1,
w2=w2,
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,
),
quant=MoEQuantParams(quant_type=QuantType.W8A8 if with_quant else QuantType.NONE),
fusion=fusion,
activation=activation,
need_trans=need_trans,
dynamic_eplb=dynamic_eplb,
)
@pytest.fixture(autouse=True)
def setup_vllm_config_mock(mocker: MockerFixture):
mock_hf_config = MagicMock()
mock_hf_config.model_type = "llama"
mock_model_config = MagicMock()
mock_model_config.hf_config = mock_hf_config
mock_vllm_config = MagicMock()
mock_vllm_config.model_config = mock_model_config
mock_vllm_config.parallel_config = MagicMock(tensor_parallel_size=2)
mock_vllm_config.scheduler_config = MagicMock(max_num_seqs=4)
mock_vllm_config.model_config.max_model_len = 2048
mocker.patch("vllm_ascend.ops.fused_moe.fused_moe.get_current_vllm_config", return_value=mock_vllm_config)
@pytest.fixture
def mock_dist_env(mocker: MockerFixture):
mock_moe_comm_method = MagicMock()
def mock_prepare(hidden_states, router_logits, **kwargs):
return MoEPrepareOutput(
hidden_states=hidden_states,
router_logits=router_logits,
mc2_mask=kwargs.get("mc2_mask"),
padded_hidden_states_shape=None,
pertoken_scale=None,
)
mock_moe_comm_method.prepare.side_effect = mock_prepare
mock_fused_experts_result = torch.randn(16, 2)
mock_moe_comm_method.fused_experts.return_value = mock_fused_experts_result
def mock_finalize(hidden_states, **kwargs):
return hidden_states
mock_moe_comm_method.finalize.side_effect = mock_finalize
dp_metadata = MagicMock(num_tokens_across_dp_cpu=[5, 5])
mock_weight_prefetch_method = MagicMock()
mock_forward_context_obj = MagicMock(
moe_comm_method=mock_moe_comm_method,
moe_comm_type=MoECommType.MC2,
max_tokens_across_dp=10,
dp_metadata=dp_metadata,
mc2_mask=torch.zeros(16, dtype=torch.bool),
padded_num_tokens=16,
with_quant=False,
)
with (
patch("torch.distributed.get_rank", return_value=0),
patch("torch.distributed.get_world_size", return_value=4),
patch("vllm_ascend.ops.fused_moe.fused_moe.get_ep_group", return_value=mock_ep_and_mc2_group(mocker)),
patch("vllm_ascend.ops.fused_moe.token_dispatcher.get_ep_group", return_value=mock_ep_and_mc2_group(mocker)),
patch("vllm_ascend.ops.fused_moe.fused_moe.get_mc2_group", return_value=mock_ep_and_mc2_group(mocker)),
patch("vllm_ascend.ops.fused_moe.fused_moe.get_tp_group", return_value=mock_dp_and_tp_group(mocker)),
patch("vllm.distributed.parallel_state.get_tp_group", return_value=mock_dp_and_tp_group(mocker)),
patch("vllm_ascend.ops.fused_moe.fused_moe.get_dp_group", return_value=mock_dp_and_tp_group(mocker)),
patch("vllm.model_executor.layers.fused_moe.layer.get_dp_group", return_value=mock_dp_and_tp_group(mocker)),
patch("vllm.model_executor.layers.fused_moe.config.get_dp_group", return_value=mock_dp_and_tp_group(mocker)),
patch(
"vllm_ascend.ops.fused_moe.fused_moe.get_ascend_config",
return_value=MagicMock(enable_multistream_moe=False, expert_map_path=None),
),
patch(
"vllm_ascend.ops.fused_moe.fused_moe.init_eplb_config",
return_value=(torch.tensor([0, 1, 2, -1, -1, -1, -1, -1]), None, 0),
),
patch("vllm_ascend.ops.fused_moe.fused_moe.get_forward_context", return_value=mock_forward_context_obj),
patch("vllm_ascend.ascend_forward_context.get_forward_context", return_value=mock_forward_context_obj),
patch("vllm_ascend.utils.get_ascend_device_type", return_value=AscendDeviceType.A3),
patch("vllm_ascend.ops.fused_moe.moe_comm_method.MC2CommImpl._get_token_dispatcher", return_value=None),
patch("vllm_ascend.ops.fused_moe.moe_comm_method.AlltoAllCommImpl._get_token_dispatcher", return_value=None),
patch("vllm_ascend.ops.fused_moe.moe_comm_method.AllGatherCommImpl._get_token_dispatcher", return_value=None),
patch(
"vllm_ascend.ops.fused_moe.experts_selector.get_weight_prefetch_method",
return_value=mock_weight_prefetch_method,
),
):
yield {
"mock_forward_context_obj": mock_forward_context_obj,
"mock_moe_comm_method": mock_moe_comm_method,
}
@pytest.fixture
def default_moe_config():
return {"num_experts": 8, "top_k": 2, "hidden_size": 512, "intermediate_size": 1024}
@pytest.fixture
def moe_method(mock_dist_env):
moe = MagicMock()
moe.moe_parallel_config.return_value = MagicMock(ep_size=4)
moe.moe_parallel_config.use_ep = False
moe.moe_parallel_config.dp_size = 1
return AscendUnquantizedFusedMoEMethod(moe)
def test_ascend_unquantized_skips_upstream_modular_kernel_init():
method = AscendUnquantizedFusedMoEMethod.maybe_make_prepare_finalize
assert method(object()) is None
class Device(TypedDict):
device_id: int
device_expert: list[int]
class Layer(TypedDict):
layer_id: int
device_count: int
device_list: list[Device]
class MockData(TypedDict):
moe_layer_count: int
layer_list: list[Layer]
class MockQuantMethod(nn.Module):
def __init__(self, shared_experts, num_tokens):
super().__init__()
if shared_experts:
self.apply = MagicMock(return_value=(torch.randn(num_tokens, 32), torch.randn(num_tokens, 10)))
else:
self.apply = MagicMock(return_value=(torch.randn(num_tokens, 32)))
def _drop_self(signature: inspect.Signature) -> list[inspect.Parameter]:
params = list(signature.parameters.values())
if params and params[0].name == "self":
return params[1:]
return params
def _format_signature_mismatch(method_name: str, issues: list[str]) -> str:
return f"{method_name} signature is not aligned with vLLM parent: " + "; ".join(issues)
def _assert_child_signature_accepts_parent_interface(child_method, parent_method):
child_params = _drop_self(inspect.signature(child_method))
parent_params = _drop_self(inspect.signature(parent_method))
child_by_name = {
param.name: param
for param in child_params
if param.kind not in (inspect.Parameter.VAR_POSITIONAL, inspect.Parameter.VAR_KEYWORD)
}
child_has_var_positional = any(param.kind == inspect.Parameter.VAR_POSITIONAL for param in child_params)
child_has_var_keyword = any(param.kind == inspect.Parameter.VAR_KEYWORD for param in child_params)
issues: list[str] = []
for parent_param in parent_params:
if parent_param.kind == inspect.Parameter.VAR_POSITIONAL:
if not child_has_var_positional:
issues.append("child is missing *args from parent")
continue
if parent_param.kind == inspect.Parameter.VAR_KEYWORD:
if not child_has_var_keyword:
issues.append("child is missing **kwargs from parent")
continue
child_param = child_by_name.get(parent_param.name)
if child_param is None:
if parent_param.kind == inspect.Parameter.KEYWORD_ONLY:
if not child_has_var_keyword:
issues.append(f"missing keyword-only parameter {parent_param.name!r}")
elif not child_has_var_positional and not child_has_var_keyword:
issues.append(f"missing parameter {parent_param.name!r}")
continue
if parent_param.kind != child_param.kind:
issues.append(
f"parameter {parent_param.name!r} has kind {child_param.kind!s}, expected {parent_param.kind!s}"
)
parent_param_names = {param.name for param in parent_params}
for child_param in child_params:
if child_param.kind in (inspect.Parameter.VAR_POSITIONAL, inspect.Parameter.VAR_KEYWORD):
continue
if child_param.name in parent_param_names:
continue
if child_param.default is inspect.Parameter.empty:
issues.append(f"extra parameter {child_param.name!r} must be optional")
assert not issues, _format_signature_mismatch(parent_method.__qualname__, issues)
def _method_uses_super(method) -> bool:
try:
source = inspect.getsource(method)
except (OSError, TypeError):
return False
tree = ast.parse(textwrap.dedent(source))
return any(
isinstance(node, ast.Call) and isinstance(node.func, ast.Name) and node.func.id == "super"
for node in ast.walk(tree)
)
class TestVllmParentInterfaceCompatibility:
@pytest.mark.parametrize(
"child_cls,parent_cls,method_name",
[
(AscendUnquantizedFusedMoEMethod, fused_moe_module.UnquantizedFusedMoEMethod, "__init__"),
(
AscendUnquantizedFusedMoEMethod,
fused_moe_module.UnquantizedFusedMoEMethod,
"process_weights_after_loading",
),
(AscendUnquantizedFusedMoEMethod, fused_moe_module.UnquantizedFusedMoEMethod, "apply"),
(AscendMoERunner, fused_moe_module.MoERunner, "__init__"),
(AscendMoERunner, fused_moe_module.MoERunner, "forward_impl"),
(AscendMoERunner, fused_moe_module.MoERunner, "_forward_impl"),
(AscendFusedMoE, fused_moe_module.FusedMoE, "__init__"),
(AscendFusedMoE, fused_moe_module.FusedMoE, "forward"),
(AscendFusedMoE, fused_moe_module.FusedMoE, "forward_impl"),
(AscendFusedMoE, fused_moe_module.FusedMoE, "maybe_all_reduce_tensor_model_parallel"),
],
)
def test_overridden_method_signature_accepts_parent_interface(self, child_cls, parent_cls, method_name):
child_method = getattr(child_cls, method_name)
if not _method_uses_super(child_method):
pytest.skip(
f"{child_cls.__name__}.{method_name} does not call "
"super(), so parent interface alignment is not "
"required"
)
if not hasattr(parent_cls, method_name):
pytest.fail(
f"{child_cls.__name__}.{method_name} calls super(), but {parent_cls.__name__} has no {method_name}"
)
_assert_child_signature_accepts_parent_interface(
child_method,
getattr(parent_cls, method_name),
)
class TestAscendUnquantizedFusedMoEMethod:
def _build_layer(self, *, has_bias=True, zero_expert_num=0):
layer = MagicMock()
layer.w13_weight = nn.Parameter(torch.randn(2, 3, 4))
layer.w2_weight = nn.Parameter(torch.randn(2, 4, 3))
layer.w13_bias = torch.randn(2, 4) if has_bias else None
layer.w2_bias = torch.randn(2, 3) if has_bias else None
layer.zero_expert_num = zero_expert_num
layer.zero_expert_type = "identity" if zero_expert_num > 0 else None
layer.n_shared_experts = 0
layer.moe_config = SimpleNamespace(num_logical_experts=None)
layer.layer_id = 3
layer.vllm_config = SimpleNamespace(model_config=SimpleNamespace(enable_return_routed_experts=False))
return layer
@pytest.mark.parametrize("enable_fused_mc2", [True, False])
def test_process_weights_after_loading_transposes_and_formats(self, monkeypatch, enable_fused_mc2):
method = AscendUnquantizedFusedMoEMethod.__new__(AscendUnquantizedFusedMoEMethod)
method.dynamic_eplb = False
method._maybe_pad_weight = MagicMock(side_effect=lambda weight: weight)
layer = self._build_layer()
original_w13 = layer.w13_weight.detach().clone()
original_w2 = layer.w2_weight.detach().clone()
format_cast = MagicMock(side_effect=lambda weight, _: weight)
maybe_trans_nz = MagicMock(side_effect=lambda weight: weight)
mock_ascend_config = MagicMock()
mock_ascend_config.enable_fused_mc2 = enable_fused_mc2
monkeypatch.setattr(fused_moe_module, "get_ascend_config", lambda: mock_ascend_config)
monkeypatch.setattr(fused_moe_module.torch_npu, "npu_format_cast", format_cast)
monkeypatch.setattr(fused_moe_module, "maybe_trans_nz", maybe_trans_nz)
method.process_weights_after_loading(layer)
torch.testing.assert_close(layer.w13_weight, original_w13.transpose(1, 2).contiguous())
torch.testing.assert_close(layer.w2_weight, original_w2.transpose(1, 2).contiguous())
if enable_fused_mc2:
assert format_cast.call_count == 2
maybe_trans_nz.assert_not_called()
else:
assert maybe_trans_nz.call_count == 2
format_cast.assert_not_called()
def test_process_weights_after_loading_splits_dynamic_eplb_fused_mc2_weights(self, monkeypatch):
method = AscendUnquantizedFusedMoEMethod.__new__(AscendUnquantizedFusedMoEMethod)
method.dynamic_eplb = True
method._maybe_pad_weight = MagicMock(side_effect=lambda weight: weight)
layer = nn.Module()
layer.w13_weight = nn.Parameter(torch.randn(2, 3, 4))
layer.w2_weight = nn.Parameter(torch.randn(2, 4, 3))
expected_w13 = layer.w13_weight.detach().clone().transpose(1, 2).contiguous()
expected_w2 = layer.w2_weight.detach().clone().transpose(1, 2).contiguous()
format_cast = MagicMock(side_effect=lambda weight, _: weight)
empty_cache = MagicMock()
mock_ascend_config = MagicMock()
mock_ascend_config.enable_fused_mc2 = True
monkeypatch.setattr(fused_moe_module, "get_ascend_config", lambda: mock_ascend_config)
monkeypatch.setattr(fused_moe_module.torch_npu, "npu_format_cast", format_cast)
monkeypatch.setattr(fused_moe_module.torch, "npu", SimpleNamespace(empty_cache=empty_cache), raising=False)
method.process_weights_after_loading(layer)
assert "w13_weight" not in layer._parameters
assert "w2_weight" not in layer._parameters
assert len(layer.w13_weight_list) == 2
assert len(layer.w2_weight_list) == 2
torch.testing.assert_close(layer.w13_weight_list[0], expected_w13[0])
torch.testing.assert_close(layer.w2_weight_list[1], expected_w2[1])
assert layer.w13_weight_list[0].untyped_storage().data_ptr() != expected_w13[0].untyped_storage().data_ptr()
assert format_cast.call_count == 2
empty_cache.assert_called_once()
@pytest.mark.parametrize("moe_comm_type", [MoECommType.MC2, MoECommType.FUSED_MC2])
def test_apply_builds_fused_experts_input(self, monkeypatch, moe_comm_type):
method = AscendUnquantizedFusedMoEMethod.__new__(AscendUnquantizedFusedMoEMethod)
method.moe = SimpleNamespace(has_bias=True)
method.dynamic_eplb = False
method.tid2eid = None
layer = self._build_layer(has_bias=True)
hidden_states = torch.randn(2, 4, dtype=torch.float16)
router_logits = torch.randn(2, 4)
topk_weights = torch.tensor([[0.25, 0.75], [0.6, 0.4]], dtype=torch.float32)
topk_ids = torch.tensor([[0, 1], [1, 0]], dtype=torch.int64)
moe_comm_method = MagicMock()
moe_comm_method.fused_experts.return_value = torch.ones_like(hidden_states)
monkeypatch.setattr(
fused_moe_module,
"_EXTRA_CTX",
SimpleNamespace(moe_comm_type=moe_comm_type, moe_comm_method=moe_comm_method),
)
select_experts_mock = MagicMock(return_value=(topk_weights, topk_ids))
monkeypatch.setattr(fused_moe_module, "select_experts", select_experts_mock)
monkeypatch.setattr(fused_moe_module, "get_forward_context", MagicMock(return_value=MagicMock(input_ids=None)))
result = method.apply(
layer=layer,
x=hidden_states,
use_grouped_topk=False,
top_k=2,
router_logits=router_logits,
renormalize=True,
num_experts=4,
apply_router_weight_on_input=True,
activation="gelu",
pertoken_scale=torch.ones(2),
mc2_mask=torch.tensor([True, False]),
)
torch.testing.assert_close(result, torch.ones_like(hidden_states))
select_experts_mock.assert_called_once()
fused_input = moe_comm_method.fused_experts.call_args.kwargs["fused_experts_input"]
assert fused_input.hidden_states is hidden_states
torch.testing.assert_close(fused_input.topk_weights, topk_weights.to(hidden_states.dtype))
assert torch.equal(fused_input.topk_ids, topk_ids)
assert fused_input.weights.w1_bias is layer.w13_bias
assert fused_input.weights.w2_bias is layer.w2_bias
assert fused_input.routing.apply_router_weight_on_input
assert fused_input.activation == "gelu"
if moe_comm_type == MoECommType.FUSED_MC2:
assert fused_input.weights.w1[0] is layer.w13_weight
assert fused_input.weights.w2[0] is layer.w2_weight
assert isinstance(fused_input.weights.w1_scale, list)
assert isinstance(fused_input.weights.w2_scale, list)
assert fused_input.weights.w1_scale[0].dtype == torch.int64
assert fused_input.weights.w2_scale[0].dtype == torch.int64
assert fused_input.weights.w1_scale_bias[0].dtype == torch.float32
assert fused_input.weights.w2_scale_bias[0].dtype == torch.float32
else:
assert fused_input.weights.w1 is layer.w13_weight
assert fused_input.weights.w2 is layer.w2_weight
assert fused_input.weights.w1_scale is None
assert fused_input.weights.w2_scale is None
@pytest.mark.parametrize("moe_comm_type", [MoECommType.MC2, MoECommType.FUSED_MC2])
def test_apply_uses_weight_lists_when_dynamic_eplb_splits_weights(self, monkeypatch, moe_comm_type):
method = AscendUnquantizedFusedMoEMethod.__new__(AscendUnquantizedFusedMoEMethod)
method.moe = SimpleNamespace(has_bias=False)
method.dynamic_eplb = True
method.tid2eid = None
layer = self._build_layer(has_bias=False)
layer.w13_weight_list = [torch.randn(4, 6), torch.randn(4, 6)]
layer.w2_weight_list = [torch.randn(3, 4), torch.randn(3, 4)]
hidden_states = torch.randn(2, 4, dtype=torch.float16)
topk_weights = torch.ones(2, 2, dtype=torch.float32)
topk_ids = torch.tensor([[0, 1], [1, 0]], dtype=torch.int64)
moe_comm_method = MagicMock()
moe_comm_method.fused_experts.return_value = torch.ones_like(hidden_states)
monkeypatch.setattr(
fused_moe_module,
"_EXTRA_CTX",
SimpleNamespace(moe_comm_type=moe_comm_type, moe_comm_method=moe_comm_method),
)
monkeypatch.setattr(fused_moe_module, "select_experts", MagicMock(return_value=(topk_weights, topk_ids)))
monkeypatch.setattr(fused_moe_module, "get_forward_context", MagicMock(return_value=MagicMock(input_ids=None)))
method.apply(
layer=layer,
x=hidden_states,
use_grouped_topk=False,
top_k=2,
router_logits=torch.randn(2, 4),
renormalize=True,
num_experts=4,
)
fused_input = moe_comm_method.fused_experts.call_args.kwargs["fused_experts_input"]
assert fused_input.weights.w1 is layer.w13_weight_list
assert fused_input.weights.w2 is layer.w2_weight_list
if moe_comm_type == MoECommType.FUSED_MC2:
assert len(fused_input.weights.w1_scale) == 1
assert len(fused_input.weights.w2_scale) == 1
assert fused_input.weights.w1_scale[0].dtype == torch.int64
assert fused_input.weights.w2_scale[0].dtype == torch.int64
assert fused_input.weights.w1_scale[0].numel() == 0
assert fused_input.weights.w2_scale[0].numel() == 0
assert fused_input.weights.w1_scale_bias[0].dtype == torch.float32
assert fused_input.weights.w2_scale_bias[0].dtype == torch.float32
assert fused_input.weights.w1_scale_bias[0].numel() == 0
assert fused_input.weights.w2_scale_bias[0].numel() == 0
else:
assert fused_input.weights.w1_scale is None
assert fused_input.weights.w2_scale is None
def test_apply_warns_when_dynamic_eplb_fused_mc2_weights_are_not_split(self, monkeypatch):
method = AscendUnquantizedFusedMoEMethod.__new__(AscendUnquantizedFusedMoEMethod)
method.moe = SimpleNamespace(has_bias=False)
method.dynamic_eplb = True
method.tid2eid = None
layer = self._build_layer(has_bias=False)
hidden_states = torch.randn(2, 4, dtype=torch.float16)
topk_weights = torch.ones(2, 2, dtype=torch.float32)
topk_ids = torch.tensor([[0, 1], [1, 0]], dtype=torch.int64)
moe_comm_method = MagicMock()
moe_comm_method.fused_experts.return_value = torch.ones_like(hidden_states)
warning_once = MagicMock()
monkeypatch.setattr(
fused_moe_module,
"_EXTRA_CTX",
SimpleNamespace(moe_comm_type=MoECommType.FUSED_MC2, moe_comm_method=moe_comm_method),
)
monkeypatch.setattr(fused_moe_module, "select_experts", MagicMock(return_value=(topk_weights, topk_ids)))
monkeypatch.setattr(fused_moe_module, "get_forward_context", MagicMock(return_value=MagicMock(input_ids=None)))
monkeypatch.setattr(fused_moe_module.logger, "warning_once", warning_once)
method.apply(
layer=layer,
x=hidden_states,
use_grouped_topk=False,
top_k=2,
router_logits=torch.randn(2, 4),
renormalize=True,
num_experts=4,
)
warning_once.assert_called_once()
warning_msg = warning_once.call_args.args[0]
assert "dynamic EPLB" in warning_msg
assert "not split into tensor lists" in warning_msg
fused_input = moe_comm_method.fused_experts.call_args.kwargs["fused_experts_input"]
assert fused_input.weights.w1[0] is layer.w13_weight
assert fused_input.weights.w2[0] is layer.w2_weight
def test_apply_adds_zero_expert_result_and_force_balances(self, monkeypatch):
method = AscendUnquantizedFusedMoEMethod.__new__(AscendUnquantizedFusedMoEMethod)
method.moe = SimpleNamespace(has_bias=False)
method.dynamic_eplb = True
method.tid2eid = None
layer = self._build_layer(has_bias=False, zero_expert_num=1)
hidden_states = torch.randn(2, 4)
topk_weights = torch.ones(2, 2)
topk_ids = torch.tensor([[0, 1], [1, 0]], dtype=torch.int32)
zero_hidden = torch.full_like(hidden_states, 3.0)
routed_hidden = torch.full_like(hidden_states, 5.0)
expected = routed_hidden + zero_hidden
moe_comm_method = MagicMock()
moe_comm_method.fused_experts.return_value = routed_hidden
monkeypatch.setattr(
fused_moe_module,
"_EXTRA_CTX",
SimpleNamespace(moe_comm_type=MoECommType.MC2, moe_comm_method=moe_comm_method),
)
monkeypatch.setattr(fused_moe_module, "select_experts", MagicMock(return_value=(topk_weights, topk_ids)))
zero_experts_mock = MagicMock(return_value=(topk_ids, topk_weights, zero_hidden))
monkeypatch.setattr(fused_moe_module, "zero_experts_compute", zero_experts_mock)
monkeypatch.setattr(torch, "rand", MagicMock(return_value=torch.tensor([[0.2, 0.1], [0.4, 0.3]])))
monkeypatch.setattr(fused_moe_module, "get_forward_context", MagicMock(return_value=MagicMock(input_ids=None)))
result = method.apply(
layer=layer,
x=hidden_states,
use_grouped_topk=False,
top_k=2,
router_logits=torch.randn(2, 2),
renormalize=False,
num_experts=2,
enable_force_load_balance=True,
)
torch.testing.assert_close(result, expected)
zero_experts_mock.assert_called_once()
fused_input = moe_comm_method.fused_experts.call_args.kwargs["fused_experts_input"]
assert fused_input.dynamic_eplb
assert fused_input.weights.w1_bias is None
assert fused_input.weights.w2_bias is None
class TestAscendMoERunner:
@pytest.mark.parametrize(
"moe_comm_type, flash_comm_v1_enabled, expected",
[
(MoECommType.ALLTOALL, False, True),
(MoECommType.MC2, False, True),
(MoECommType.FUSED_MC2, False, True),
(MoECommType.ALLGATHER, False, False),
(MoECommType.ALLGATHER, True, True),
],
)
def test_runner_reduction_properties(self, monkeypatch, moe_comm_type, flash_comm_v1_enabled, expected):
runner = AscendMoERunner.__new__(AscendMoERunner)
monkeypatch.setattr(fused_moe_legacy_module, "_EXTRA_CTX", SimpleNamespace(moe_comm_type=moe_comm_type))
monkeypatch.setattr(
fused_moe_legacy_module,
"_EXTRA_CTX",
SimpleNamespace(moe_comm_type=moe_comm_type, flash_comm_v1_enabled=flash_comm_v1_enabled),
)
assert runner.use_dp_chunking is False
if hasattr(type(runner), "_fused_output_is_reduced"):
assert runner._fused_output_is_reduced is expected
if hasattr(runner, "_maybe_reduce_shared_expert_output"):
assert runner._maybe_reduce_shared_expert_output("shared") == "shared"
@pytest.mark.parametrize("has_shared_experts", [False, True])
def test_forward_impl_delegates_to_layer(self, monkeypatch, has_shared_experts):
runner = AscendMoERunner.__new__(AscendMoERunner)
shared_experts = MagicMock() if has_shared_experts else None
shared_experts_owner = next(
(cls for cls in type(runner).__mro__ if "shared_experts" in cls.__dict__),
AscendMoERunner,
)
monkeypatch.setattr(shared_experts_owner, "shared_experts", property(lambda _: shared_experts), raising=False)
layer = MagicMock()
hidden_states = torch.randn(2, 4)
router_logits = torch.randn(2, 3)
layer.forward_impl.return_value = "routed"
layer.shared_forward_impl.return_value = ("shared", "routed")
result = runner.forward_impl(layer, hidden_states, router_logits, None)
if has_shared_experts:
assert result == ("shared", "routed")
layer.shared_forward_impl.assert_called_once_with(hidden_states, router_logits)
layer.forward_impl.assert_not_called()
else:
assert result == "routed"
layer.forward_impl.assert_called_once_with(hidden_states, router_logits)
layer.shared_forward_impl.assert_not_called()
class TestAscendFusedMoE:
def _build_layer(self):
layer = AscendFusedMoE.__new__(AscendFusedMoE)
layer.quant_method = MagicMock()
layer.ensure_moe_quant_config_init = MagicMock()
layer.runner = MagicMock()
layer.moe_load = torch.zeros(2, dtype=torch.int64)
layer.multi_stage = False
layer.log2phy = torch.tensor([1, 0])
return layer
def test_simple_helpers(self, monkeypatch):
layer = self._build_layer()
layer.quant_method.quant_method = SimpleNamespace(quant_type=QuantType.W8A8)
layer.update_expert_map(torch.tensor([0, -1]))
assert torch.equal(layer._expert_map, torch.tensor([0, -1]))
assert torch.equal(layer.get_log2phy_map(), torch.tensor([1, 0]))
assert layer._get_quant_type() == QuantType.W8A8
layer.clear_moe_load()
assert torch.equal(layer.moe_load, torch.zeros_like(layer.moe_load))
layer.multi_stage = True
layer.load_counter = torch.tensor(4)
layer.clear_moe_load()
assert layer.load_counter.item() == 0
maybe_all_reduce = MagicMock(return_value="reduced")
monkeypatch.setattr(
fused_moe_module.torch.ops,
"vllm",
SimpleNamespace(maybe_all_reduce_tensor_model_parallel=maybe_all_reduce),
raising=False,
)
assert layer.maybe_all_reduce_tensor_model_parallel(torch.ones(1)) == "reduced"
def test_forward_delegates_to_runner(self):
layer = self._build_layer()
hidden_states = torch.randn(2, 4)
router_logits = torch.randn(2, 3)
layer.runner.forward.return_value = "forwarded"
assert layer.forward(hidden_states, router_logits) == "forwarded"
layer.ensure_moe_quant_config_init.assert_called_once()
layer.runner.forward.assert_called_once_with(hidden_states, router_logits)
@pytest.mark.parametrize("return_with_event", [True, False])
def test_forward_impl_prepare_apply_finalize(self, monkeypatch, return_with_event):
layer = self._build_layer()
layer.enable_npugraph_ex_static_kernel = True
layer.multistream_overlap_gate = False
layer.enable_shared_expert_dp = False
layer.quant_type = QuantType.NONE
layer.top_k = 2
layer.renormalize = True
layer.use_grouped_topk = False
layer.moe_config = SimpleNamespace(num_experts=4)
layer._expert_map = None
layer.topk_group = None
layer.num_expert_group = None
layer.custom_routing_function = None
layer.scoring_func = "softmax"
layer._original_routed_scaling_factor = 1.0
layer.routed_scaling_factor = 1.0
layer.e_score_correction_bias = None
layer.activation = "silu"
layer.apply_router_weight_on_input = False
layer.global_redundant_expert_num = 0
layer.dynamic_eplb = True
layer.reduce_results = True
forward_context = SimpleNamespace(moe_layer_index=5, all_moe_layers=[0, 1])
hidden_states = torch.randn(2, 4)
router_logits = torch.randn(2, 4)
prepared_hidden = hidden_states + 1
prepared_logits = router_logits + 1
prepare_output = MoEPrepareOutput(
hidden_states=prepared_hidden,
router_logits=prepared_logits,
mc2_mask=torch.tensor([True, False]),
padded_hidden_states_shape=torch.Size([4, 4]),
pertoken_scale=torch.ones(2),
)
moe_comm_method = MagicMock()
moe_comm_method.prepare.return_value = prepare_output
moe_comm_method.finalize.side_effect = lambda hidden_states, **_: hidden_states + 2
before_dispatch_evt = MagicMock()
before_combine_evt = MagicMock()
layer.quant_method.apply.return_value = FusedExpertsResult(
routed_out=torch.ones_like(hidden_states),
before_dispatch_evt=before_dispatch_evt,
before_combine_evt=before_combine_evt,
expert_tokens=torch.tensor([2, 5]),
group_list_type=0,
)
monkeypatch.setattr(fused_moe_legacy_module, "get_forward_context", MagicMock(return_value=forward_context))
monkeypatch.setattr(
fused_moe_legacy_module,
"_EXTRA_CTX",
SimpleNamespace(
in_profile_run=True,
moe_comm_method=moe_comm_method,
flash_comm_v1_enabled=True,
eplb_heat_collection_status=True,
),
)
result = layer.forward_impl(hidden_states, router_logits, return_with_event=return_with_event)
assert forward_context.moe_layer_index == 1
moe_comm_method.prepare.assert_called_once_with(
hidden_states=hidden_states,
router_logits=router_logits,
replace_allreduce=True,
enable_shared_expert_dp=False,
quant_type=QuantType.NONE,
)
apply_kwargs = layer.quant_method.apply.call_args.kwargs
assert apply_kwargs["x"] is prepared_hidden
assert apply_kwargs["router_logits"] is prepared_logits
assert apply_kwargs["num_experts"] == 4
assert apply_kwargs["enable_force_load_balance"] is True
assert torch.equal(apply_kwargs["mc2_mask"], prepare_output.mc2_mask)
torch.testing.assert_close(layer.moe_load, torch.tensor([2, 3]))
if return_with_event:
assert result.routed_out.shape == hidden_states.shape
assert result.before_dispatch_evt is before_dispatch_evt
assert result.before_combine_evt is before_combine_evt
else:
torch.testing.assert_close(result, torch.ones_like(hidden_states) + 2)
def test_forward_impl_dynamic_eplb_multi_stage(self, monkeypatch):
layer = self._build_layer()
layer.enable_npugraph_ex_static_kernel = False
layer.multistream_overlap_gate = False
layer.enable_shared_expert_dp = False
layer.quant_type = QuantType.NONE
layer.top_k = 1
layer.renormalize = False
layer.use_grouped_topk = False
layer.moe_config = SimpleNamespace(num_experts=2)
layer._expert_map = None
layer.topk_group = None
layer.num_expert_group = None
layer.custom_routing_function = None
layer.scoring_func = "softmax"
layer._original_routed_scaling_factor = 1.0
layer.routed_scaling_factor = 1.0
layer.e_score_correction_bias = None
layer.activation = "silu"
layer.apply_router_weight_on_input = False
layer.global_redundant_expert_num = 0
layer.dynamic_eplb = True
layer.multi_stage = True
layer.moe_load = torch.zeros((2, 2), dtype=torch.int32)
layer.load_counter = torch.tensor([1], dtype=torch.int64)
layer.num_iter = 2
layer.reduce_results = False
moe_comm_method = MagicMock()
moe_comm_method.prepare.return_value = MoEPrepareOutput(
hidden_states=torch.ones(2, 4),
router_logits=torch.ones(2, 2),
mc2_mask=None,
padded_hidden_states_shape=None,
)
moe_comm_method.finalize.side_effect = lambda hidden_states, **_: hidden_states
layer.quant_method.apply.return_value = FusedExpertsResult(
routed_out=torch.ones(2, 4),
expert_tokens=torch.tensor([4, 6]),
group_list_type=1,
)
monkeypatch.setattr(fused_moe_legacy_module, "get_forward_context", MagicMock(return_value=SimpleNamespace()))
monkeypatch.setattr(
fused_moe_legacy_module,
"_EXTRA_CTX",
SimpleNamespace(
in_profile_run=False,
moe_comm_method=moe_comm_method,
flash_comm_v1_enabled=False,
eplb_heat_collection_status=True,
),
)
layer.forward_impl(torch.zeros(2, 4), torch.zeros(2, 2))
assert torch.equal(layer.moe_load[1], torch.tensor([4, 6], dtype=torch.int32))
assert layer.load_counter.item() == 2
class TestAscendFusedMoESharedExperts:
def test_properties_and_forward_delegate(self, monkeypatch):
layer = AscendFusedMoE.__new__(AscendFusedMoE)
if not hasattr(type(layer), "gate"):
pytest.skip("Current AscendFusedMoE does not expose gate property")
layer.multistream_overlap_shared_expert = False
layer._gate = MagicMock()
layer.use_overlapped = True
assert layer.gate is layer._gate
layer.use_overlapped = False
assert layer.gate is None
assert layer.is_internal_router is False
assert layer.use_dp_chunking is False
monkeypatch.setattr(fused_moe_module.AscendFusedMoE, "forward", MagicMock(return_value="routed"))
layer._shared_experts = None
assert layer.forward(torch.ones(1, 2), torch.ones(1, 2)) == "routed"
fused_moe_module.AscendFusedMoE.forward.return_value = "forwarded"
layer._shared_experts = MagicMock()
assert layer.forward(torch.ones(1, 2), torch.ones(1, 2)) == "forwarded"
def test_shared_experts_split_with_expert_gate(self):
layer = AscendFusedMoE.__new__(AscendFusedMoE)
if not hasattr(layer, "_shared_experts_part1"):
pytest.skip("Current AscendFusedMoE does not split shared experts")
hidden_states = torch.tensor([[1.0, -1.0]])
gate_up = torch.tensor([[2.0, -2.0]])
down_out = torch.tensor([[3.0, 4.0]])
gate_out = torch.tensor([[0.0, 2.0]])
shared_experts = MagicMock()
shared_experts.gate_up_proj.return_value = (gate_up, None)
shared_experts.act_fn.side_effect = lambda tensor: tensor + 1
shared_experts.down_proj.return_value = (down_out, None)
shared_experts.expert_gate.return_value = (gate_out, None)
layer._shared_experts = shared_experts
part1_out = layer._shared_experts_part1(hidden_states)
part2_out = layer._shared_experts_part2(hidden_states, part1_out)
torch.testing.assert_close(part1_out, gate_up)
torch.testing.assert_close(part2_out, F.sigmoid(gate_out) * down_out)
@pytest.mark.parametrize("has_shared_experts", [False, True])
def test_shared_forward_impl_routes_shared_output(self, monkeypatch, has_shared_experts):
layer = AscendFusedMoE.__new__(AscendFusedMoE)
if not hasattr(layer, "shared_forward_impl"):
pytest.skip("Current AscendFusedMoE has no shared_forward_impl")
layer.multistream_overlap_shared_expert = False
layer.shared_multistream_overlap_gate = False
layer.use_overlapped = False
layer._shared_experts = MagicMock() if has_shared_experts else None
hidden_states = torch.randn(2, 4)
router_logits = torch.randn(2, 3)
fused_result = fused_moe_module.FusedMoEResult(
routed_out=torch.ones(2, 4),
before_dispatch_evt=MagicMock(),
before_combine_evt=MagicMock(),
)
monkeypatch.setattr(
fused_moe_module.torch.npu,
"current_stream",
MagicMock(return_value=MagicMock(record_event=MagicMock(return_value=MagicMock()))),
)
monkeypatch.setattr(fused_moe_module.AscendFusedMoE, "forward_impl", MagicMock(return_value=fused_result))
layer._forward_shared_experts = MagicMock(return_value="shared_out")
result = layer.shared_forward_impl(hidden_states, router_logits)
if has_shared_experts:
assert result == ("shared_out", fused_result.routed_out)
layer._forward_shared_experts.assert_called_once()
else:
torch.testing.assert_close(result, fused_result.routed_out)

View File

@@ -0,0 +1,70 @@
#
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
#
# 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 unittest
from unittest import mock
from unittest.mock import patch
import torch
from tests.ut.base import TestBase
from vllm_ascend.ops.fused_moe.gate_linear import AscendGateLinear
class TestAscendGateLinear(TestBase):
def setUp(self):
super().setUp()
self.mock_group = mock.MagicMock()
self.mock_group.world_size = 1
self.mock_group.rank_in_group = 0
self.patches = [
patch(
"vllm.distributed.parallel_state.get_tp_group",
return_value=self.mock_group,
),
]
for p in self.patches:
p.start()
def tearDown(self):
for p in self.patches:
p.stop()
super().tearDown()
def test_forward_keeps_router_logits_fp32(self):
gate = AscendGateLinear(
input_size=16,
output_size=4,
bias=False,
prefix="test.gate",
)
self.assertEqual(gate.weight.dtype, torch.float32)
hidden_states = torch.randn(2, 16, dtype=torch.bfloat16)
output, output_bias = gate(hidden_states)
self.assertEqual(output.dtype, torch.float32)
self.assertIsNone(output_bias)
if __name__ == "__main__":
unittest.main()

View File

@@ -0,0 +1,923 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass
from types import SimpleNamespace
from unittest.mock import patch
import pytest
import torch
from vllm.config.compilation import CUDAGraphMode
from vllm.model_executor.layers.fla.ops import index as _fla_index
from vllm.v1.attention.backend import CommonAttentionMetadata
from vllm.v1.attention.backends.utils import PAD_SLOT_ID
from vllm.v1.kv_cache_interface import MambaSpec
from vllm_ascend.attention.utils import AscendCommonAttentionMetadata
from vllm_ascend.ops import gdn_attn_builder as ascend_gdn_attn_builder
from vllm_ascend.ops.gdn import AscendGatedDeltaNetAttention
from vllm_ascend.ops.gdn_attn_builder import (
AscendGDNAttentionBackend,
AscendGDNAttentionMetadataBuilder,
)
from vllm_ascend.ops.triton.fla import utils as fla_utils
from vllm_ascend.ops.triton.fla.utils import (
prepare_chunk_indices as runtime_prepare_chunk_indices,
)
from vllm_ascend.ops.triton.fla.utils import (
prepare_chunk_offsets as runtime_prepare_chunk_offsets,
)
from vllm_ascend.ops.triton.fla.utils import (
prepare_final_chunk_indices as runtime_prepare_final_chunk_indices,
)
from vllm_ascend.ops.triton.fla.utils import (
prepare_update_chunk_offsets as runtime_prepare_update_chunk_offsets,
)
from vllm_ascend.utils import vllm_version_is
@pytest.fixture(autouse=True)
def _patch_triton_cdiv(monkeypatch):
if not hasattr(_fla_index.triton, "cdiv"):
monkeypatch.setattr(
_fla_index.triton,
"cdiv",
lambda a, b: (a + b - 1) // b,
raising=False,
)
@pytest.fixture(autouse=True)
def _no_pin_memory():
# compute_causal_conv1d_metadata uses np_to_pinned_tensor which reads
# PIN_MEMORY. Without physical NPU, t.pin_memory() raises
# "Please register PrivateUse1HooksInterface first".
with patch("vllm.utils.torch_utils.PIN_MEMORY", False):
if vllm_version_is("0.23.0"):
yield
else:
with patch("vllm.v1.attention.backends.utils.PIN_MEMORY", False):
yield
@dataclass
class BatchSpec:
seq_lens: list[int]
query_lens: list[int]
name: str = "unnamed"
@property
def batch_size(self) -> int:
return len(self.seq_lens)
def create_common_attn_metadata(
batch_spec: BatchSpec,
block_size: int,
device: torch.device,
) -> CommonAttentionMetadata:
query_lens_cpu = torch.tensor(batch_spec.query_lens, dtype=torch.int32)
query_start_loc_cpu = torch.zeros(
batch_spec.batch_size + 1,
dtype=torch.int32,
)
query_start_loc_cpu[1:] = query_lens_cpu.cumsum(0)
query_start_loc = query_start_loc_cpu.to(device=device)
num_tokens = sum(batch_spec.query_lens)
seq_lens_cpu = torch.tensor(batch_spec.seq_lens, dtype=torch.int32)
seq_lens = seq_lens_cpu.to(device=device)
max_seq_len = int(seq_lens_cpu.max())
context_lens = [batch_spec.seq_lens[i] - batch_spec.query_lens[i] for i in range(batch_spec.batch_size)]
num_computed_tokens_cpu = torch.tensor(context_lens, dtype=torch.int32)
# Mirror model_runner: is_prefilling = num_computed < num_prompt_tokens.
# Chunked prefills still have prompt tokens beyond num_computed; decodes do not.
num_prompt_tokens_cpu = torch.tensor(
[
context_lens[i] + batch_spec.query_lens[i] if batch_spec.query_lens[i] > 1 else context_lens[i]
for i in range(batch_spec.batch_size)
],
dtype=torch.int32,
)
is_prefilling = num_computed_tokens_cpu < num_prompt_tokens_cpu
max_blocks = (max(batch_spec.seq_lens) + block_size - 1) // block_size
block_table_tensor = torch.arange(
batch_spec.batch_size * max_blocks,
dtype=torch.int32,
device=device,
).view(batch_spec.batch_size, max_blocks)
slot_mapping = torch.arange(num_tokens, dtype=torch.int64, device=device)
return AscendCommonAttentionMetadata(
query_start_loc=query_start_loc,
query_start_loc_cpu=query_start_loc_cpu,
seq_lens=seq_lens,
_seq_lens_cpu=seq_lens_cpu,
seq_lens_cpu=seq_lens_cpu,
seq_lens_cpu_upper_bound=seq_lens_cpu,
_num_computed_tokens_cpu=num_computed_tokens_cpu,
num_computed_tokens_cpu=num_computed_tokens_cpu,
num_reqs=batch_spec.batch_size,
num_actual_tokens=num_tokens,
max_query_len=max(batch_spec.query_lens),
max_seq_len=max_seq_len,
block_table_tensor=block_table_tensor,
slot_mapping=slot_mapping,
causal=True,
is_prefilling=is_prefilling,
)
def _make_vllm_config(
*,
max_model_len: int = 8192,
max_num_seqs: int = 16,
max_num_batched_tokens: int = 8192,
num_heads: int = 32,
num_speculative_tokens: int = 0,
mamba_cache_mode: str = "none",
cudagraph_mode: CUDAGraphMode = CUDAGraphMode.NONE,
prefill_context_parallel_size: int = 1,
):
speculative_config = None
if num_speculative_tokens > 0:
speculative_config = SimpleNamespace(
num_speculative_tokens=num_speculative_tokens,
parallel_drafting=False,
)
model_config = SimpleNamespace(max_model_len=max_model_len)
model_config.get_num_attention_heads = lambda parallel_config: num_heads
return SimpleNamespace(
cache_config=SimpleNamespace(mamba_cache_mode=mamba_cache_mode),
compilation_config=SimpleNamespace(
cudagraph_mode=cudagraph_mode,
max_cudagraph_capture_size=None,
),
speculative_config=speculative_config,
scheduler_config=SimpleNamespace(
max_num_seqs=max_num_seqs,
max_num_batched_tokens=max_num_batched_tokens,
),
parallel_config=SimpleNamespace(
decode_context_parallel_size=1,
prefill_context_parallel_size=prefill_context_parallel_size,
tensor_parallel_size=1,
),
model_config=model_config,
additional_config=None,
)
def _make_builder(
*,
device: torch.device,
num_heads: int,
num_speculative_tokens: int,
mamba_cache_mode: str = "none",
block_size: int = 16,
num_speculative_blocks: int = 0,
cudagraph_mode: CUDAGraphMode = CUDAGraphMode.NONE,
prefill_context_parallel_size: int = 1,
):
vllm_config = _make_vllm_config(
num_heads=num_heads,
num_speculative_tokens=num_speculative_tokens,
mamba_cache_mode=mamba_cache_mode,
cudagraph_mode=cudagraph_mode,
prefill_context_parallel_size=prefill_context_parallel_size,
)
spec = MambaSpec(
block_size=block_size,
shapes=((1,), (1,)),
dtypes=(torch.float32,),
mamba_cache_mode=mamba_cache_mode,
num_speculative_blocks=num_speculative_blocks,
)
return AscendGDNAttentionMetadataBuilder(spec, ["layer0"], vllm_config, device)
def _build_attn_metadata(
batch_spec: BatchSpec,
*,
num_speculative_tokens: int,
num_decode_draft_tokens_cpu: torch.Tensor | None,
):
device = torch.device("cpu")
common_attn_metadata = create_common_attn_metadata(
batch_spec=batch_spec,
block_size=16,
device=device,
)
builder = _make_builder(
device=device,
num_heads=32,
num_speculative_tokens=num_speculative_tokens,
)
num_accepted_tokens = None
if num_decode_draft_tokens_cpu is not None:
num_accepted_tokens = torch.ones(
batch_spec.batch_size,
dtype=torch.int32,
)
attn_metadata = builder.build(
0,
common_attn_metadata,
num_accepted_tokens=num_accepted_tokens,
num_decode_draft_tokens_cpu=num_decode_draft_tokens_cpu,
)
return builder, common_attn_metadata, attn_metadata
def _assert_chunk_meta_matches_runtime(builder, chunk_meta, cu_seqlens: torch.Tensor) -> None:
hf_text_config = getattr(builder.vllm_config.model_config, "hf_text_config", None)
if hf_text_config is not None and hasattr(hf_text_config, "linear_num_value_heads"):
gdn_num_heads = (
hf_text_config.linear_num_value_heads // builder.vllm_config.parallel_config.tensor_parallel_size
)
else:
gdn_num_heads = builder.vllm_config.model_config.get_num_attention_heads(builder.vllm_config.parallel_config)
cumsum_chunks = max(
1,
ascend_gdn_attn_builder._GDN_CUMSUM_WORKING_SET // (gdn_num_heads * ascend_gdn_attn_builder._GDN_CHUNK_SIZE),
)
cumsum_chunk_size = 1 if cumsum_chunks <= 1 else 1 << (cumsum_chunks - 1).bit_length()
sequence_lengths = cu_seqlens[1:] - cu_seqlens[:-1]
assert chunk_meta.num_decodes == (sequence_lengths == 1).sum().item()
assert torch.equal(
chunk_meta.chunk_indices_chunk64,
runtime_prepare_chunk_indices(cu_seqlens, ascend_gdn_attn_builder._GDN_CHUNK_SIZE),
)
assert torch.equal(
chunk_meta.chunk_offsets_chunk64,
runtime_prepare_chunk_offsets(cu_seqlens, ascend_gdn_attn_builder._GDN_CHUNK_SIZE),
)
assert torch.equal(
chunk_meta.update_chunk_offsets_chunk64,
runtime_prepare_update_chunk_offsets(
cu_seqlens,
ascend_gdn_attn_builder._GDN_CHUNK_SIZE,
),
)
assert torch.equal(
chunk_meta.final_chunk_indices_chunk64,
runtime_prepare_final_chunk_indices(
cu_seqlens,
ascend_gdn_attn_builder._GDN_CHUNK_SIZE,
),
)
assert torch.equal(
chunk_meta.chunk_indices_large_block,
runtime_prepare_chunk_indices(
cu_seqlens,
ascend_gdn_attn_builder._GDN_SOLVE_TRIL_LARGE_BLOCK_SIZE,
),
)
assert torch.equal(
chunk_meta.block_indices_cumsum,
runtime_prepare_chunk_indices(
cu_seqlens,
cumsum_chunk_size,
),
)
def _patch_missing_runtime_cdiv(monkeypatch: pytest.MonkeyPatch) -> None:
if hasattr(fla_utils.triton, "cdiv"):
return
monkeypatch.setattr(
fla_utils.triton,
"cdiv",
lambda x, y: (x + y - 1) // y,
raising=False,
)
def test_ascend_gdn_attention_uses_ascend_backend():
assert AscendGatedDeltaNetAttention.get_attn_backend(object()) is AscendGDNAttentionBackend
assert AscendGDNAttentionBackend.get_builder_cls() is AscendGDNAttentionMetadataBuilder
def test_sequence_index_buffers_cover_spec_decode_when_cudagraph_disabled():
builder = _make_builder(
device=torch.device("cpu"),
num_heads=32,
num_speculative_tokens=3,
)
assert builder.spec_sequence_indices_cpu.numel() >= builder.vllm_config.scheduler_config.max_num_seqs
spec_indices, non_spec_indices = builder._copy_sequence_indices_to_device(
torch.tensor([True], dtype=torch.bool),
num_spec_decodes=1,
)
assert torch.equal(spec_indices, torch.tensor([0]))
assert non_spec_indices.numel() == 0
def _cache_index_first_column(cache_indices: torch.Tensor) -> torch.Tensor:
if cache_indices.dim() == 1:
return cache_indices
return cache_indices[:, 0]
def _assert_non_spec_conv1d_args_match_metadata(attn_metadata) -> None:
conv1d_meta = attn_metadata.non_spec_prefill_metadata.causal_conv1d
assert torch.equal(conv1d_meta.query_start_loc, attn_metadata.non_spec_query_start_loc)
assert torch.equal(
_cache_index_first_column(conv1d_meta.cache_indices),
attn_metadata.non_spec_state_indices_tensor,
)
assert torch.equal(conv1d_meta.initial_state_mode, attn_metadata.has_initial_state)
@pytest.mark.parametrize(
("batch_spec", "num_speculative_tokens", "num_decode_draft_tokens_cpu"),
[
(
BatchSpec(
seq_lens=[8, 12],
query_lens=[4, 8],
name="pure_non_spec_prefill",
),
0,
None,
),
(
BatchSpec(
seq_lens=[8, 4, 0, 12],
query_lens=[4, 4, 0, 8],
name="mixed_spec_non_spec_with_padding",
),
3,
torch.tensor([-1, 3, -1, -1], dtype=torch.int32),
),
(
BatchSpec(
seq_lens=[5, 12, 0, 9],
query_lens=[1, 8, 0, 1],
name="mixed_prefill_decode_without_spec",
),
0,
None,
),
],
ids=lambda case: case.name if isinstance(case, BatchSpec) else None,
)
def test_non_spec_prefill_metadata_matches_original_inputs_and_runtime_helpers(
batch_spec: BatchSpec,
num_speculative_tokens: int,
num_decode_draft_tokens_cpu: torch.Tensor | None,
monkeypatch: pytest.MonkeyPatch,
):
_patch_missing_runtime_cdiv(monkeypatch)
builder, _, attn_metadata = _build_attn_metadata(
batch_spec,
num_speculative_tokens=num_speculative_tokens,
num_decode_draft_tokens_cpu=num_decode_draft_tokens_cpu,
)
prefill_metadata = getattr(attn_metadata, "non_spec_prefill_metadata", None)
assert prefill_metadata is not None
assert prefill_metadata.causal_conv1d is not None
assert prefill_metadata.chunk is not None
_assert_non_spec_conv1d_args_match_metadata(attn_metadata)
_assert_chunk_meta_matches_runtime(
builder,
prefill_metadata.chunk,
attn_metadata.prefill_query_start_loc,
)
def test_non_spec_prefill_metadata_uses_prefill_tail_for_chunk_metadata(
monkeypatch: pytest.MonkeyPatch,
):
_patch_missing_runtime_cdiv(monkeypatch)
batch_spec = BatchSpec(
seq_lens=[5, 12, 9],
query_lens=[1, 8, 4],
name="decode_prefill_without_spec",
)
builder, _, attn_metadata = _build_attn_metadata(
batch_spec,
num_speculative_tokens=0,
num_decode_draft_tokens_cpu=None,
)
assert attn_metadata.num_decodes == 1
assert attn_metadata.num_prefills == 2
assert torch.equal(
attn_metadata.non_spec_query_start_loc,
torch.tensor([0, 1, 9, 13], dtype=torch.int32),
)
assert torch.equal(
attn_metadata.prefill_query_start_loc,
torch.tensor([0, 8, 12], dtype=torch.int32),
)
assert torch.equal(
attn_metadata.non_spec_state_indices_tensor,
torch.tensor([0, 1, 2], dtype=torch.int32),
)
assert torch.equal(
attn_metadata.prefill_state_indices,
torch.tensor([1, 2], dtype=torch.int32),
)
prefill_metadata = getattr(attn_metadata, "non_spec_prefill_metadata", None)
assert prefill_metadata is not None
decode_metadata = getattr(attn_metadata, "non_spec_decode_metadata", None)
assert decode_metadata is not None
assert torch.equal(
decode_metadata.actual_seq_lengths,
torch.tensor([0, 1], dtype=torch.int32),
)
conv1d_meta = prefill_metadata.causal_conv1d
assert torch.equal(conv1d_meta.query_start_loc, torch.tensor([0, 1, 9, 13], dtype=torch.int32))
assert torch.equal(_cache_index_first_column(conv1d_meta.cache_indices), torch.tensor([0, 1, 2], dtype=torch.int32))
assert torch.equal(conv1d_meta.initial_state_mode, torch.tensor([True, True, True]))
assert prefill_metadata.chunk.num_decodes == 0
_assert_chunk_meta_matches_runtime(
builder,
prefill_metadata.chunk,
attn_metadata.prefill_query_start_loc,
)
def test_mixed_spec_prefill_chunk_metadata_preserves_single_token_count(
monkeypatch: pytest.MonkeyPatch,
):
_patch_missing_runtime_cdiv(monkeypatch)
batch_spec = BatchSpec(
seq_lens=[1, 4, 8],
query_lens=[1, 4, 8],
name="mixed_spec_prefill_with_single_token_non_spec",
)
builder, _, attn_metadata = _build_attn_metadata(
batch_spec,
num_speculative_tokens=3,
num_decode_draft_tokens_cpu=torch.tensor([-1, 3, -1], dtype=torch.int32),
)
assert attn_metadata.num_decodes == 0
assert attn_metadata.num_prefills == 2
assert torch.equal(
attn_metadata.prefill_query_start_loc,
torch.tensor([0, 1, 9], dtype=torch.int32),
)
chunk_metadata = attn_metadata.non_spec_prefill_metadata.chunk
assert chunk_metadata.num_decodes == 1
_assert_chunk_meta_matches_runtime(
builder,
chunk_metadata,
attn_metadata.prefill_query_start_loc,
)
def test_spec_conv1d_args_use_device_cache_and_accepted_tokens():
batch_spec = BatchSpec(
seq_lens=[4, 4],
query_lens=[4, 4],
name="spec_only_device_args",
)
common_attn_metadata = create_common_attn_metadata(
batch_spec=batch_spec,
block_size=16,
device=torch.device("cpu"),
)
common_attn_metadata.block_table_tensor = torch.tensor(
[[10, 11, 12, 13], [20, 21, 22, 23]],
dtype=torch.int32,
)
builder = _make_builder(
device=torch.device("cpu"),
num_heads=32,
num_speculative_tokens=3,
)
num_accepted_tokens = torch.tensor([2, 4], dtype=torch.int32)
attn_metadata = builder.build(
0,
common_attn_metadata,
num_accepted_tokens=num_accepted_tokens,
num_decode_draft_tokens_cpu=torch.tensor([3, 3], dtype=torch.int32),
)
spec_conv1d_meta = attn_metadata.spec_decode_metadata.spec_causal_conv1d
query_start_loc = spec_conv1d_meta.query_start_loc
assert torch.equal(query_start_loc, torch.tensor([0, 4, 8], dtype=torch.int32))
assert torch.equal(
spec_conv1d_meta.cache_indices,
torch.tensor([[10, 11, 12, 13], [20, 21, 22, 23]], dtype=torch.int32),
)
assert torch.equal(spec_conv1d_meta.num_accepted_tokens, num_accepted_tokens)
assert torch.equal(
attn_metadata.spec_decode_metadata.actual_seq_lengths,
torch.tensor([0, 4, 4], dtype=torch.int32),
)
def test_full_graph_spec_conv1d_args_keep_request_granularity():
batch_spec = BatchSpec(
seq_lens=[4, 4, 4],
query_lens=[4, 4, 4],
name="full_graph_spec_only_device_args",
)
common_attn_metadata = create_common_attn_metadata(
batch_spec=batch_spec,
block_size=16,
device=torch.device("cpu"),
)
common_attn_metadata.block_table_tensor = torch.tensor(
[[10, 11, 12, 13], [20, 21, 22, 23], [30, 31, 32, 33]],
dtype=torch.int32,
)
builder = _make_builder(
device=torch.device("cpu"),
num_heads=32,
num_speculative_tokens=3,
cudagraph_mode=CUDAGraphMode.FULL_DECODE_ONLY,
)
num_accepted_tokens = torch.tensor([2, 4, 3], dtype=torch.int32)
attn_metadata = builder.build(
0,
common_attn_metadata,
num_accepted_tokens=num_accepted_tokens,
num_decode_draft_tokens_cpu=torch.tensor([3, 3, 3], dtype=torch.int32),
)
spec_conv1d_meta = attn_metadata.spec_decode_metadata.spec_causal_conv1d
query_start_loc = spec_conv1d_meta.query_start_loc
assert torch.equal(query_start_loc, torch.tensor([0, 4, 8, 12], dtype=torch.int32))
assert query_start_loc.numel() == batch_spec.batch_size + 1
assert spec_conv1d_meta.cache_indices.shape == (batch_spec.batch_size, 4)
assert torch.equal(spec_conv1d_meta.cache_indices[:, 0], torch.tensor([10, 20, 30], dtype=torch.int32))
assert torch.equal(spec_conv1d_meta.num_accepted_tokens, num_accepted_tokens)
assert torch.equal(
attn_metadata.spec_decode_metadata.actual_seq_lengths,
torch.tensor([0, 4, 4, 4], dtype=torch.int32),
)
def test_full_graph_spec_actual_seq_lengths_use_padded_builder_buffer():
batch_spec = BatchSpec(
seq_lens=[4, 4],
query_lens=[4, 4],
name="full_graph_padded_spec_actual_seq_lengths",
)
common_attn_metadata = create_common_attn_metadata(
batch_spec=batch_spec,
block_size=16,
device=torch.device("cpu"),
)
common_attn_metadata.num_reqs = 4
common_attn_metadata.block_table_tensor = torch.tensor(
[[10, 11, 12, 13], [20, 21, 22, 23]],
dtype=torch.int32,
)
builder = _make_builder(
device=torch.device("cpu"),
num_heads=32,
num_speculative_tokens=3,
cudagraph_mode=CUDAGraphMode.FULL_DECODE_ONLY,
)
attn_metadata = builder.build(
0,
common_attn_metadata,
num_accepted_tokens=torch.tensor([2, 4], dtype=torch.int32),
num_decode_draft_tokens_cpu=torch.tensor([3, 3], dtype=torch.int32),
)
assert torch.equal(
attn_metadata.spec_query_start_loc,
torch.tensor([0, 4, 8, 8, 8], dtype=torch.int32),
)
assert (
attn_metadata.spec_decode_metadata.actual_seq_lengths.data_ptr() == builder.spec_actual_seq_lengths.data_ptr()
)
assert torch.equal(
attn_metadata.spec_decode_metadata.actual_seq_lengths,
torch.tensor([0, 4, 4, 0, 0], dtype=torch.int32),
)
def test_full_graph_without_runtime_spec_resets_captured_spec_inputs():
capture_batch = BatchSpec(
seq_lens=[4, 4],
query_lens=[4, 4],
name="full_graph_spec_capture",
)
capture_common_metadata = create_common_attn_metadata(
batch_spec=capture_batch,
block_size=16,
device=torch.device("cpu"),
)
capture_common_metadata.num_reqs = 4
capture_common_metadata.block_table_tensor = torch.tensor(
[[10, 11, 12, 13], [20, 21, 22, 23]],
dtype=torch.int32,
)
builder = _make_builder(
device=torch.device("cpu"),
num_heads=32,
num_speculative_tokens=3,
cudagraph_mode=CUDAGraphMode.FULL_DECODE_ONLY,
)
captured_metadata = builder.build(
0,
capture_common_metadata,
num_accepted_tokens=torch.tensor([2, 4], dtype=torch.int32),
num_decode_draft_tokens_cpu=torch.tensor([3, 3], dtype=torch.int32),
)
captured_spec_metadata = captured_metadata.spec_decode_metadata
captured_conv1d_metadata = captured_spec_metadata.spec_causal_conv1d
assert torch.count_nonzero(captured_conv1d_metadata.query_start_loc) > 0
assert torch.count_nonzero(captured_spec_metadata.actual_seq_lengths) > 0
replay_batch = BatchSpec(
seq_lens=[1, 1, 0, 0],
query_lens=[1, 1, 0, 0],
name="full_graph_replay_without_spec",
)
replay_common_metadata = create_common_attn_metadata(
batch_spec=replay_batch,
block_size=16,
device=torch.device("cpu"),
)
replay_metadata = builder.build(
0,
replay_common_metadata,
num_accepted_tokens=torch.ones(4, dtype=torch.int32),
num_decode_draft_tokens_cpu=torch.full((4,), -1, dtype=torch.int32),
)
assert replay_metadata.spec_sequence_masks is None
assert replay_metadata.spec_decode_metadata is None
assert torch.equal(
captured_conv1d_metadata.cache_indices,
torch.full((4, 4), PAD_SLOT_ID, dtype=torch.int32),
)
assert torch.count_nonzero(captured_conv1d_metadata.query_start_loc) == 0
assert torch.count_nonzero(captured_conv1d_metadata.num_accepted_tokens) == 0
assert torch.count_nonzero(captured_spec_metadata.actual_seq_lengths) == 0
@pytest.mark.parametrize(
("num_speculative_tokens", "num_decode_draft_tokens_cpu"),
[
pytest.param(0, None, id="without_mtp"),
pytest.param(
3,
torch.full((4,), -1, dtype=torch.int32),
id="mtp_without_spec_requests",
),
],
)
def test_full_graph_non_spec_metadata_nulls_padded_state_indices(
num_speculative_tokens: int,
num_decode_draft_tokens_cpu: torch.Tensor | None,
):
batch_spec = BatchSpec(
seq_lens=[1, 1, 0, 0],
query_lens=[1, 1, 0, 0],
name="full_graph_padded_non_spec_actual_seq_lengths",
)
common_attn_metadata = create_common_attn_metadata(
batch_spec=batch_spec,
block_size=16,
device=torch.device("cpu"),
)
# PCP leaves padded block-table rows untouched. Model the stale valid
# state slots that can remain there after the preceding decode batch.
common_attn_metadata.block_table_tensor[:, 0] = torch.tensor([10, 11, 98, 99])
builder = _make_builder(
device=torch.device("cpu"),
num_heads=32,
num_speculative_tokens=num_speculative_tokens,
cudagraph_mode=CUDAGraphMode.FULL_DECODE_ONLY,
)
builder.non_spec_state_indices_tensor.fill_(77)
builder.non_spec_query_start_loc.fill_(77)
builder.non_spec_actual_seq_lengths.fill_(77)
attn_metadata = builder.build(
0,
common_attn_metadata,
num_decode_draft_tokens_cpu=num_decode_draft_tokens_cpu,
)
assert attn_metadata.num_decodes == 4
assert attn_metadata.num_decode_tokens == 2
assert torch.equal(
attn_metadata.non_spec_query_start_loc,
torch.tensor([0, 1, 2, 2, 2], dtype=torch.int32),
)
assert torch.equal(
attn_metadata.non_spec_state_indices_tensor,
torch.tensor([10, 11, 0, 0], dtype=torch.int32),
)
decode_metadata = attn_metadata.non_spec_decode_metadata
conv1d_metadata = decode_metadata.causal_conv1d
assert conv1d_metadata.query_start_loc.data_ptr() == attn_metadata.non_spec_query_start_loc.data_ptr()
assert conv1d_metadata.cache_indices.data_ptr() == attn_metadata.non_spec_state_indices_tensor.data_ptr()
assert decode_metadata.actual_seq_lengths.data_ptr() == builder.non_spec_actual_seq_lengths.data_ptr()
assert torch.equal(
decode_metadata.actual_seq_lengths,
torch.tensor([0, 1, 1, 0, 0], dtype=torch.int32),
)
def test_causal_conv1d_cache_indices_use_device_block_table(monkeypatch: pytest.MonkeyPatch):
_patch_missing_runtime_cdiv(monkeypatch)
batch_spec = BatchSpec(
seq_lens=[4, 4],
query_lens=[4, 4],
name="device_block_table_source",
)
common_attn_metadata = create_common_attn_metadata(
batch_spec=batch_spec,
block_size=16,
device=torch.device("cpu"),
)
common_attn_metadata.block_table_tensor = torch.tensor(
[[40], [41]],
dtype=torch.int32,
)
builder = _make_builder(
device=torch.device("cpu"),
num_heads=32,
num_speculative_tokens=0,
)
attn_metadata = builder.build(0, common_attn_metadata)
assert torch.equal(
attn_metadata.non_spec_state_indices_tensor,
torch.tensor([40, 41], dtype=torch.int32),
)
conv1d_meta = attn_metadata.non_spec_prefill_metadata.causal_conv1d
assert torch.equal(conv1d_meta.query_start_loc, torch.tensor([0, 4, 8], dtype=torch.int32))
assert torch.equal(_cache_index_first_column(conv1d_meta.cache_indices), torch.tensor([40, 41], dtype=torch.int32))
assert torch.equal(conv1d_meta.initial_state_mode, torch.tensor([False, False]))
def test_pcp_prefill_initial_state_mode_is_built_in_metadata(monkeypatch: pytest.MonkeyPatch):
_patch_missing_runtime_cdiv(monkeypatch)
batch_spec = BatchSpec(
seq_lens=[1, 4],
query_lens=[1, 4],
name="pcp_decode_prefill",
)
common_attn_metadata = create_common_attn_metadata(
batch_spec=batch_spec,
block_size=16,
device=torch.device("cpu"),
)
builder = _make_builder(
device=torch.device("cpu"),
num_heads=32,
num_speculative_tokens=0,
prefill_context_parallel_size=2,
)
with patch(
"vllm_ascend.ops.gdn_attn_builder.get_pcp_group",
return_value=SimpleNamespace(world_size=2, rank_in_group=1),
):
attn_metadata = builder.build(0, common_attn_metadata)
conv1d_meta = attn_metadata.non_spec_prefill_metadata.causal_conv1d
assert torch.equal(
conv1d_meta.initial_state_mode,
torch.tensor([False, True]),
)
def test_mamba_align_cache_indices_follow_device_seq_lens(monkeypatch: pytest.MonkeyPatch):
_patch_missing_runtime_cdiv(monkeypatch)
batch_spec = BatchSpec(
seq_lens=[1, 9],
query_lens=[1, 1],
name="align_device_seq_lens",
)
common_attn_metadata = create_common_attn_metadata(
batch_spec=batch_spec,
block_size=4,
device=torch.device("cpu"),
)
common_attn_metadata.block_table_tensor = torch.arange(20, dtype=torch.int32).view(2, 10)
common_attn_metadata._seq_lens_cpu = torch.tensor([5, 13], dtype=torch.int32)
builder = _make_builder(
device=torch.device("cpu"),
num_heads=32,
num_speculative_tokens=0,
mamba_cache_mode="align",
block_size=4,
num_speculative_blocks=2,
)
attn_metadata = builder.build(0, common_attn_metadata)
conv1d_meta = attn_metadata.non_spec_decode_metadata.causal_conv1d
assert torch.equal(
_cache_index_first_column(conv1d_meta.cache_indices),
torch.tensor([0, 12], dtype=torch.int32),
)
def test_builder_builds_prebuilt_chunk_metadata_with_prefill_query_start_loc(monkeypatch):
_patch_missing_runtime_cdiv(monkeypatch)
batch_spec = BatchSpec(
seq_lens=[8, 4, 0, 12],
query_lens=[4, 4, 0, 8],
name="mixed_spec_non_spec_with_padding",
)
builder, common_attn_metadata, _ = _build_attn_metadata(
batch_spec,
num_speculative_tokens=3,
num_decode_draft_tokens_cpu=torch.tensor([-1, 3, -1, -1], dtype=torch.int32),
)
attn_metadata = builder.build(
0,
common_attn_metadata,
num_accepted_tokens=torch.ones(batch_spec.batch_size, dtype=torch.int32),
num_decode_draft_tokens_cpu=torch.tensor([-1, 3, -1, -1], dtype=torch.int32),
)
chunk_meta = attn_metadata.non_spec_prefill_metadata.chunk
assert chunk_meta.chunk_indices_chunk64 is attn_metadata.chunk_indices
assert chunk_meta.chunk_offsets_chunk64 is attn_metadata.chunk_offsets
_assert_chunk_meta_matches_runtime(
builder,
chunk_meta,
attn_metadata.prefill_query_start_loc,
)
assert chunk_meta.cu_seqlens_host == tuple(attn_metadata.prefill_query_start_loc.to(torch.int64).tolist())
expected_chunk_indices = runtime_prepare_chunk_indices(
attn_metadata.prefill_query_start_loc,
ascend_gdn_attn_builder._GDN_CHUNK_SIZE,
)
assert chunk_meta.chunk_indices_chunk64_host == tuple(expected_chunk_indices.to(torch.int64).reshape(-1).tolist())
@pytest.mark.parametrize(
"batch_spec",
[
BatchSpec(seq_lens=[1, 1, 1], query_lens=[1, 1, 1], name="decode_only"),
BatchSpec(seq_lens=[4, 4], query_lens=[4, 4], name="spec_only"),
],
)
def test_builder_skips_prebuilt_meta_without_non_spec_prefill(batch_spec: BatchSpec):
builder = _make_builder(
device=torch.device("cpu"),
num_heads=32,
num_speculative_tokens=3 if batch_spec.name == "spec_only" else 0,
)
common_attn_metadata = create_common_attn_metadata(
batch_spec=batch_spec,
block_size=16,
device=torch.device("cpu"),
)
num_accepted_tokens = None
num_decode_draft_tokens_cpu = None
if batch_spec.name == "spec_only":
num_accepted_tokens = torch.ones(
batch_spec.batch_size,
dtype=torch.int32,
)
num_decode_draft_tokens_cpu = torch.full(
(batch_spec.batch_size,),
3,
dtype=torch.int32,
)
attn_metadata = builder.build(
0,
common_attn_metadata,
num_accepted_tokens=num_accepted_tokens,
num_decode_draft_tokens_cpu=num_decode_draft_tokens_cpu,
)
assert getattr(attn_metadata, "non_spec_prefill_metadata", None) is None
if batch_spec.name == "decode_only":
decode_metadata = getattr(attn_metadata, "non_spec_decode_metadata", None)
assert decode_metadata is not None
assert torch.equal(
decode_metadata.actual_seq_lengths,
torch.tensor([0, 1, 1, 1], dtype=torch.int32),
)
else:
spec_decode_metadata = getattr(attn_metadata, "spec_decode_metadata", None)
assert spec_decode_metadata is not None
assert torch.equal(
spec_decode_metadata.actual_seq_lengths,
torch.tensor([0, 4, 4], dtype=torch.int32),
)

View File

@@ -0,0 +1,468 @@
#
# 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 unittest.mock import MagicMock, Mock, patch
import pytest
import torch
from vllm_ascend.ops.layer_shard_linear import (
LayerExternalMetadata,
LayerMetadata,
SeriesMetadata,
ShardWindowMetadata,
_create_forward_wrapper,
dispose_tensor,
is_hidden_layer,
register_layer_to_shard_weight_series,
)
class TestDisposeTensor:
def test_dispose_tensor_replaces_with_empty(self):
original_tensor = torch.randn(10, 10)
original_shape = original_tensor.shape
dispose_tensor(original_tensor)
assert original_tensor.shape == torch.Size([])
assert original_tensor.shape != original_shape
def test_dispose_tensor_preserves_device_and_dtype(self):
original_tensor = torch.randn(5, 5, dtype=torch.float32)
original_dtype = original_tensor.dtype
dispose_tensor(original_tensor)
assert original_tensor.dtype == original_dtype
class TestLayerMetadata:
def test_layer_metadata_creation(self):
layer = MagicMock()
post_method = Mock()
weight = torch.randn(10, 10)
metadata = LayerMetadata(
layer_idx=0,
layer=layer,
post_method=post_method,
weight=weight,
window_idx=0,
)
assert metadata.layer_idx == 0
assert metadata.layer is layer
assert metadata.post_method is post_method
assert metadata.weight is weight
assert metadata.window_idx == 0
class TestShardWindowMetadata:
def test_shard_window_metadata_creation(self):
weight = torch.randn(10, 10)
window = ShardWindowMetadata(
weight=weight,
data_layer_idx=0,
work=None,
)
assert window.weight is weight
assert window.data_layer_idx == 0
assert window.work is None
class TestSeriesMetadata:
@pytest.fixture
def mock_group(self):
group = MagicMock()
group.world_size = 2
group.rank_in_group = 0
group.ranks = [0, 1]
group.device_group = MagicMock()
return group
@pytest.fixture
def series_metadata(self, mock_group):
return SeriesMetadata(
group=mock_group,
start_layer=0,
end_layer=0,
num_layers=0,
prefetch_step=1,
dummy_weight=torch.randn(10, 10),
layers=[],
shard_windows=[],
window_offset=1,
)
def test_is_source_rank_zero(self, series_metadata):
series_metadata.group.rank_in_group = 0
assert series_metadata.is_source(0) is True
assert series_metadata.is_source(1) is False
assert series_metadata.is_source(2) is True
assert series_metadata.is_source(3) is False
def test_is_source_rank_one(self, series_metadata):
series_metadata.group.rank_in_group = 1
assert series_metadata.is_source(0) is False
assert series_metadata.is_source(1) is True
assert series_metadata.is_source(2) is False
assert series_metadata.is_source(3) is True
@patch("torch.distributed.broadcast")
def test_post_process_after_loading_basic(self, mock_broadcast, series_metadata):
layer0 = MagicMock()
layer0.layer_idx = 0
layer0.weight = torch.randn(10, 10)
layer0.post_method = Mock()
layer1 = MagicMock()
layer1.layer_idx = 1
layer1.weight = torch.randn(10, 10)
layer1.post_method = Mock()
series_metadata.layers = [layer0, layer1]
series_metadata.prefetch_step = 0
series_metadata.post_process_after_loading()
assert series_metadata.num_layers == 2
assert series_metadata.start_layer == 0
assert series_metadata.end_layer == 2
assert len(series_metadata.shard_windows) == 1
assert mock_broadcast.call_count == 2
@patch("torch.distributed.broadcast")
def test_post_process_after_loading_with_prefetch(self, mock_broadcast, series_metadata):
layer0 = MagicMock()
layer0.layer_idx = 0
layer0.weight = torch.randn(10, 10)
layer0.post_method = Mock()
layer1 = MagicMock()
layer1.layer_idx = 1
layer1.weight = torch.randn(10, 10)
layer1.post_method = Mock()
layer2 = MagicMock()
layer2.layer_idx = 2
layer2.weight = torch.randn(10, 10)
layer2.post_method = Mock()
series_metadata.layers = [layer0, layer1, layer2]
series_metadata.prefetch_step = 1
series_metadata.post_process_after_loading()
assert series_metadata.num_layers == 3
assert len(series_metadata.shard_windows) == 2
assert mock_broadcast.call_count == 3
def test_post_process_after_loading_already_initialized(self, series_metadata):
series_metadata.shard_windows = [MagicMock()]
result = series_metadata.post_process_after_loading()
assert result is None
def test_post_process_after_loading_empty_layers(self, series_metadata):
series_metadata.layers = []
with pytest.raises(AssertionError, match="No layers in the series"):
series_metadata.post_process_after_loading()
@patch("torch.distributed.broadcast")
def test_reach_layer(self, mock_broadcast, series_metadata):
layer0 = MagicMock()
layer0.layer_idx = 0
layer0.weight = torch.randn(10, 10)
layer0.window_idx = -1
layer1 = MagicMock()
layer1.layer_idx = 1
layer1.weight = torch.randn(10, 10)
layer1.window_idx = -1
series_metadata.layers = [layer0, layer1]
series_metadata.num_layers = 2
series_metadata.start_layer = 0
series_metadata.prefetch_step = 0
series_metadata.window_offset = 0
window = ShardWindowMetadata(
weight=torch.randn(10, 10),
data_layer_idx=-1,
work=None,
)
series_metadata.shard_windows = [window]
mock_work = MagicMock()
mock_broadcast.return_value = mock_work
series_metadata.reach_layer(0)
assert layer0.window_idx == 0
assert layer1.window_idx == -1
assert window.data_layer_idx == 0
assert window.work is not None
mock_broadcast.assert_called_once()
@patch("torch.distributed.broadcast")
def test_wait_weight(self, mock_broadcast, series_metadata):
mock_work = MagicMock()
window = ShardWindowMetadata(
weight=torch.randn(10, 10),
data_layer_idx=0,
work=mock_work,
)
layer0 = MagicMock()
layer0.layer_idx = 0
layer0.window_idx = 0
series_metadata.layers = [layer0]
series_metadata.start_layer = 0
series_metadata.shard_windows = [window]
series_metadata.wait_weight(0)
mock_work.wait.assert_called_once()
assert window.work is None
def test_wait_weight_no_work(self, series_metadata):
window = ShardWindowMetadata(
weight=torch.randn(10, 10),
data_layer_idx=0,
work=None,
)
layer0 = MagicMock()
layer0.layer_idx = 0
layer0.window_idx = 0
series_metadata.layers = [layer0]
series_metadata.start_layer = 0
series_metadata.shard_windows = [window]
series_metadata.wait_weight(0)
assert window.work is None
class TestLayerExternalMetadata:
def test_layer_external_metadata_creation(self):
series = MagicMock()
layer_idx = 5
ext_metadata = LayerExternalMetadata(
series=series,
layer_idx=layer_idx,
)
assert ext_metadata.series is series
assert ext_metadata.layer_idx == layer_idx
class TestCreateForwardWrapper:
def test_create_forward_wrapper_calls_wait_weight(self):
mock_series = MagicMock()
mock_forward = Mock(return_value="output")
layer_idx = 0
wrapped = _create_forward_wrapper(mock_forward, mock_series, layer_idx)
result = wrapped("arg1", "arg2", kwarg1="value1")
mock_series.wait_weight.assert_called_once_with(layer_idx)
mock_forward.assert_called_once_with("arg1", "arg2", kwarg1="value1")
assert result == "output"
def test_create_forward_wrapper_preserves_return_value(self):
mock_series = MagicMock()
expected_output = torch.randn(10, 10)
mock_forward = Mock(return_value=expected_output)
wrapped = _create_forward_wrapper(mock_forward, mock_series, 0)
result = wrapped()
assert result is expected_output
class TestRegisterLayerToShardWeightSeries:
@pytest.fixture
def mock_layer(self):
layer = MagicMock()
layer.weight = torch.randn(10, 10)
layer.prefix = "model.layers.0.mlp.gate_up_proj"
layer.forward = Mock(return_value="forward_output")
quant_method = MagicMock()
quant_method.process_weights_after_loading = Mock()
layer.quant_method = quant_method
return layer
@pytest.fixture
def mock_group(self):
group = MagicMock()
group.world_size = 2
group.rank_in_group = 0
group.ranks = [0, 1]
return group
@patch("vllm_ascend.ops.layer_shard_linear._series_dict", new_callable=dict)
@patch("vllm_ascend.ops.layer_shard_linear._layer_external_dict", new_callable=dict)
@patch("vllm_ascend.ops.layer_shard_linear.extract_layer_index", return_value=0)
def test_register_layer_creates_new_series(
self,
mock_extract_index,
mock_layer_dict,
mock_series_dict,
mock_layer,
mock_group,
):
import vllm_ascend.ops.layer_shard_linear as module
register_layer_to_shard_weight_series(
series_name="test_series",
group=mock_group,
layer=mock_layer,
prefetch_step=1,
)
assert "test_series" in module._series_dict
series = module._series_dict["test_series"]
assert series.group is mock_group
assert series.prefetch_step == 1
assert len(series.layers) == 1
@patch("vllm_ascend.ops.layer_shard_linear._series_dict", new_callable=dict)
@patch("vllm_ascend.ops.layer_shard_linear._layer_external_dict", new_callable=dict)
@patch("vllm_ascend.ops.layer_shard_linear.extract_layer_index", return_value=1)
def test_register_layer_adds_to_existing_series(
self,
mock_extract_index,
mock_layer_dict,
mock_series_dict,
mock_layer,
mock_group,
):
import vllm_ascend.ops.layer_shard_linear as module
existing_series = SeriesMetadata(
group=mock_group,
start_layer=0,
end_layer=0,
num_layers=0,
prefetch_step=1,
dummy_weight=torch.randn(10, 10),
layers=[],
shard_windows=[],
window_offset=1,
)
module._series_dict["test_series"] = existing_series
register_layer_to_shard_weight_series(
series_name="test_series",
group=mock_group,
layer=mock_layer,
prefetch_step=1,
)
assert len(existing_series.layers) == 1
assert existing_series.layers[0].layer_idx == 1
@patch("vllm_ascend.ops.layer_shard_linear._series_dict", new_callable=dict)
@patch("vllm_ascend.ops.layer_shard_linear._layer_external_dict", new_callable=dict)
@patch("vllm_ascend.ops.layer_shard_linear.extract_layer_index", return_value=1)
def test_register_layer_disposes_weight_for_non_source(
self,
mock_extract_index,
mock_layer_dict,
mock_series_dict,
mock_layer,
mock_group,
):
import vllm_ascend.ops.layer_shard_linear as module
mock_group.rank_in_group = 0
register_layer_to_shard_weight_series(
series_name="test_series",
group=mock_group,
layer=mock_layer,
prefetch_step=1,
)
series = module._series_dict["test_series"]
assert series.is_source(1) is False
class TestIsHiddenLayer:
@patch("vllm_ascend.ops.layer_shard_linear.get_current_model_num_hidden_layers")
@patch("vllm_ascend.ops.layer_shard_linear.extract_layer_index")
def test_is_hidden_layer_true(
self,
mock_extract_index,
mock_get_num_layers,
):
mock_get_num_layers.return_value = 32
mock_extract_index.return_value = 10
layer = MagicMock()
layer.prefix = "model.layers.10.mlp"
result = is_hidden_layer(layer)
assert result is True
@patch("vllm_ascend.ops.layer_shard_linear.get_current_model_num_hidden_layers")
@patch("vllm_ascend.ops.layer_shard_linear.extract_layer_index")
def test_is_hidden_layer_false(
self,
mock_extract_index,
mock_get_num_layers,
):
mock_get_num_layers.return_value = 32
mock_extract_index.return_value = 40
layer = MagicMock()
layer.prefix = "model.layers.40.mlp"
result = is_hidden_layer(layer)
assert result is False
@patch("vllm_ascend.ops.layer_shard_linear.get_current_model_num_hidden_layers")
@patch("vllm_ascend.ops.layer_shard_linear.extract_layer_index")
def test_is_hidden_layer_boundary(
self,
mock_extract_index,
mock_get_num_layers,
):
mock_get_num_layers.return_value = 32
mock_extract_index.return_value = 31
layer = MagicMock()
layer.prefix = "model.layers.31.mlp"
result = is_hidden_layer(layer)
assert result is True

View File

@@ -1,18 +1,19 @@
import unittest
from unittest.mock import MagicMock, patch
import pytest
import torch
from pytest_mock import MockerFixture
from vllm.config import set_current_vllm_config
from vllm.model_executor.layers.layernorm import RMSNorm
from tests.ut.base import PytestBase
from vllm_ascend.quantization.w8a8 import AscendW8A8LinearMethod
from vllm_ascend.utils import enable_custom_op
from vllm_ascend.utils import is_310p as is_310p_hw
enable_custom_op()
def mock_maybe_chunk_residual(x, residual):
if x.size(0) != residual.size(0):
return residual[:4]
return residual
@pytest.fixture
def dummy_tensor():
return torch.randn(4, 8, dtype=torch.float16)
def mock_rms_norm(x, weight, eps):
@@ -23,139 +24,61 @@ def mock_add_rms_norm(x, residual, weight, eps):
return 2 * x, None, 2 * residual
def mock_add_rms_norm_quant(x, residual, weight, quant_scale, quant_offset,
epsilon):
x_out = 2 * x
residual_out = 2 * residual
x_out_quant = x_out.to(torch.int8)
residual_out_quant = residual_out.to(torch.int8)
return x_out_quant, None, residual_out_quant
def mock_add_rms_norm_bias(x, residual, weight, bias, eps):
if bias is None:
return 2 * x, None, 2 * residual
else:
return 2 * x + bias, None, 2 * residual
class TestAscendRMSNorm(PytestBase):
@pytest.fixture(autouse=True)
def default_vllm_config():
mock_config = MagicMock()
mock_config.compilation_config.custom_ops = ["all"]
@pytest.fixture(autouse=True)
def context(self, mocker: MockerFixture):
mocker.patch("torch.ops.vllm.maybe_chunk_residual",
side_effect=mock_maybe_chunk_residual)
mocker.patch("torch_npu.npu_rms_norm", side_effect=mock_rms_norm)
mocker.patch("torch_npu.npu_add_rms_norm",
side_effect=mock_add_rms_norm)
mocker.patch("torch_npu.npu_add_rms_norm_quant",
side_effect=mock_add_rms_norm_quant)
mocker.patch("torch.ops.vllm.maybe_wait_prefetch_done",
side_effect=lambda x: None)
# Test case for the most common and basic scenario
@pytest.mark.parametrize(
"residual", [None, torch.randn(4, 8, dtype=torch.float16)])
def test_forward_oot_basic(self, residual):
layer = RMSNorm(hidden_size=8, eps=1e-05)
x = torch.randn(4, 8, dtype=torch.float16)
if residual is not None:
x_out, residual_out = layer.forward_oot(x, residual)
x_out_expected = 2 * x
residual_out_expected = 2 * residual
assert torch.allclose(x_out, x_out_expected)
assert torch.allclose(residual_out, residual_out_expected)
else:
x_out = layer.forward(x, residual)
x_out_expected = x + 1
assert torch.allclose(x_out, x_out_expected)
# Test case for flashcomm_v1 scenario
def test_forward_oot_with_flashcomm_v1(self):
layer = RMSNorm(hidden_size=512, eps=1e-05)
x = torch.randn(4, 512, dtype=torch.bfloat16)
residual = torch.randn(16, 512, dtype=torch.bfloat16)
x_out, residual_out = layer.forward_oot(x, residual)
x_out_expected = 2 * x
residual_out_expected = 2 * residual[:4]
assert residual_out.size(0) == 4
assert torch.allclose(x_out, x_out_expected)
assert torch.allclose(residual_out, residual_out_expected)
# Test case for addrmsnorm + w8a8 quant fusion
def test_forward_oot_with_quant_fusion(self, mocker: MockerFixture):
mock_is_310p = mocker.patch("vllm_ascend.utils.is_310p")
mock_is_310p.return_value = False
mock_get_forward_context = mocker.patch(
"vllm_ascend.ops.layernorm.get_forward_context")
# Simulating a scenario with quant_fusion enabled
mock_forward_context = mocker.MagicMock()
mock_model_instance = mocker.MagicMock()
mock_forward_context.model_instance = mock_model_instance
mock_model_instance.model.layers = [
mocker.MagicMock() for _ in range(2)
]
mock_layer_0 = mock_model_instance.model.layers[0]
mock_layer_0.self_attn.qkv_proj = mocker.MagicMock()
mock_layer_0.mlp.gate_up_proj = mocker.MagicMock()
mock_layer_1 = mock_model_instance.model.layers[1]
mock_layer_1.self_attn.qkv_proj = mocker.MagicMock()
mock_layer_1.mlp.gate_up_proj = mocker.MagicMock()
mock_quant_method_0_qkv = mocker.MagicMock()
mock_quant_method_0_qkv.quant_method = AscendW8A8LinearMethod()
mock_quant_method_0_gate_up = mocker.MagicMock()
mock_quant_method_0_gate_up.quant_method = AscendW8A8LinearMethod()
mock_layer_0.self_attn.qkv_proj.quant_method = mock_quant_method_0_qkv
mock_layer_0.mlp.gate_up_proj.quant_method = mock_quant_method_0_gate_up
mock_quant_method_1_qkv = mocker.MagicMock()
mock_quant_method_1_qkv.quant_method = AscendW8A8LinearMethod()
mock_quant_method_1_gate_up = mocker.MagicMock()
mock_quant_method_1_gate_up.quant_method = AscendW8A8LinearMethod()
mock_layer_1.self_attn.qkv_proj.quant_method = mock_quant_method_1_qkv
mock_layer_1.mlp.gate_up_proj.quant_method = mock_quant_method_1_gate_up
mock_get_forward_context.return_value = mock_forward_context
mock_forward_context.addrmsnorm_quant_fusion_enabled = True
mock_forward_context.prefetch_mlp_enabled = False
mock_forward_context.layer_idx = 0
mock_forward_context.num_hidden_layers = 2
mock_forward_context.fusion_linear = "gate_up_dense"
# Ensure fusion and layer_idx increment are handled correctly
x = torch.randn(4, 8, dtype=torch.float16)
residual = torch.randn(4, 8, dtype=torch.float16)
layer = RMSNorm(hidden_size=8, eps=1e-05)
x_out, residual_out = layer.forward_oot(x, residual)
assert mock_get_forward_context.call_count == 1
assert mock_forward_context.fusion_linear == "qkv_dense"
assert mock_forward_context.layer_idx == 1
x_out, residual_out = layer.forward_oot(x, residual)
assert mock_get_forward_context.call_count == 2
assert mock_forward_context.fusion_linear == "gate_up_dense"
assert mock_forward_context.layer_idx == 1
x_out, residual_out = layer.forward_oot(x, residual)
assert mock_get_forward_context.call_count == 3
assert mock_forward_context.fusion_linear == "qkv_dense"
assert mock_forward_context.layer_idx == 2
x_out, residual_out = layer.forward_oot(x, residual)
assert mock_get_forward_context.call_count == 4
assert mock_forward_context.fusion_linear == "qkv_dense"
assert mock_forward_context.layer_idx == 2
with set_current_vllm_config(mock_config):
yield mock_config
if __name__ == '__main__':
unittest.main()
@pytest.mark.skip("Skip as register_kernels has NPU SocName checking in CANN 8.5.0.")
@pytest.mark.parametrize("residual", [None, torch.randn(4, 8, dtype=torch.float32)])
@patch("torch_npu.npu_rms_norm", side_effect=mock_rms_norm)
@patch("torch_npu.npu_add_rms_norm", side_effect=mock_add_rms_norm)
@patch("torch.ops._C_ascend.npu_add_rms_norm_bias", side_effect=mock_add_rms_norm_bias)
def test_RMSNorm_forward(
mock_add_rms_norm_bias, mock_add_rmsnorm, mock_rmsnorm, residual, dummy_tensor, default_vllm_config
):
layer = RMSNorm(hidden_size=8, eps=1e-05)
if residual is not None:
out_x, out_residual = layer.forward_oot(dummy_tensor, residual)
expected_out_x = 2 * dummy_tensor
expected_out_residual = 2 * residual
mock_add_rms_norm_bias.assert_called_once()
assert torch.allclose(out_x, expected_out_x)
assert torch.allclose(out_residual, expected_out_residual)
else:
out_x = layer.forward_oot(dummy_tensor, residual)
expected_out_x = dummy_tensor + 1
mock_rmsnorm.assert_called_once()
assert torch.allclose(out_x, expected_out_x)
@pytest.mark.skipif(not is_310p_hw(), reason="310P device unittest case.")
@pytest.mark.parametrize("residual", [None, torch.randn(4, 8, dtype=torch.float16)])
@patch("torch_npu.npu_rms_norm", side_effect=mock_rms_norm)
@patch("torch_npu.npu_add_rms_norm", side_effect=mock_add_rms_norm)
def test_RMSNorm_forward_310p(mock_add_rmsnorm, mock_rmsnorm, residual, dummy_tensor, default_vllm_config):
layer = RMSNorm(hidden_size=8, eps=1e-05)
if residual is not None:
out_x, out_residual = layer.forward_oot(dummy_tensor, residual)
expected_out_x = 2 * dummy_tensor
expected_out_residual = 2 * residual
mock_add_rmsnorm.assert_called_once()
assert torch.allclose(out_x, expected_out_x)
assert torch.allclose(out_residual, expected_out_residual)
else:
out_x = layer.forward_oot(dummy_tensor, residual)
expected_out_x = dummy_tensor + 1
mock_rmsnorm.assert_called_once()
assert torch.allclose(out_x, expected_out_x)

View File

@@ -1,18 +1,21 @@
import os
import unittest
from unittest import mock
from unittest.mock import MagicMock, patch
import torch
from tests.ut.base import TestBase
from vllm_ascend import ascend_config
from vllm_ascend.distributed import parallel_state
from vllm_ascend.ops.linear import (AscendMergedColumnParallelLinear,
AscendRowParallelLinear)
from vllm_ascend.ops.linear import (
AscendMergedColumnParallelLinear,
AscendReplicatedLinear,
AscendRowParallelLinear,
AscendUnquantizedLinearMethod,
)
class BaseLinearTest(unittest.TestCase):
def setUp(self):
self.mock_group = mock.MagicMock()
self.mock_group.world_size = 2
@@ -22,23 +25,22 @@ class BaseLinearTest(unittest.TestCase):
parallel_state._OTP = self.mock_group
self.mock_ascend_config = MagicMock()
self.mock_ascend_config.oproj_tensor_parallel_size = 2
self.mock_ascend_config.finegrained_tp_config.oproj_tensor_parallel_size = 2
self.mock_ascend_config.finegrained_tp_config.mlp_tensor_parallel_size = 2
self.patches = [
patch("vllm_ascend.ascend_config.get_ascend_config",
return_value=self.mock_ascend_config),
patch("vllm_ascend.distributed.parallel_state.get_otp_group",
return_value=self.mock_group),
patch("vllm_ascend.distributed.parallel_state.get_mlp_tp_group",
return_value=self.mock_group),
patch("vllm_ascend.ops.linear_op.get_tp_group",
return_value=self.mock_group),
patch("vllm_ascend.ascend_config.get_ascend_config", return_value=self.mock_ascend_config),
patch("vllm_ascend.distributed.parallel_state.get_otp_group", return_value=self.mock_group),
patch("vllm_ascend.distributed.parallel_state.get_mlp_tp_group", return_value=self.mock_group),
patch("vllm_ascend.ops.linear_op.get_tp_group", return_value=self.mock_group),
patch(
"vllm.distributed.parallel_state.get_tp_group",
return_value=self.mock_group,
),
patch("vllm_ascend.utils.mlp_tp_enable", return_value=True),
patch("vllm_ascend.utils.oproj_tp_enable", return_value=True)
patch("vllm_ascend.utils.oproj_tp_enable", return_value=True),
patch("vllm_ascend.ops.linear_op.enable_dsa_cp", return_value=False),
patch("vllm_ascend.ops.linear_op.enable_dsa_cp_with_layer_shard", return_value=False),
]
for p in self.patches:
@@ -49,10 +51,57 @@ class BaseLinearTest(unittest.TestCase):
p.stop()
class TestAscendRowParallelLinear(BaseLinearTest):
class TestAscendUnquantizedLinearMethod(TestBase):
def setUp(self):
self.method = AscendUnquantizedLinearMethod()
self.layer = mock.MagicMock()
mock_dtype = mock.PropertyMock(return_value=torch.float16)
type(self.layer.weight.data).dtype = mock_dtype
mock_is_meta = mock.PropertyMock(return_value=False)
type(self.layer.weight.data).is_meta = mock_is_meta
self.layer.precast_fp32_weight = False
def test_mlp_optimize(self):
os.environ["VLLM_ASCEND_ENABLE_MLP_OPTIMIZE"] = "1"
@patch("vllm_ascend.utils.get_ascend_config")
@mock.patch("torch_npu.npu_format_cast")
def test_process_weights_after_loading_with_nz0(self, mock_format_cast, mock_get_config):
mock_config = MagicMock()
mock_config.weight_nz_mode = 0
mock_get_config.return_value = mock_config
self.method.process_weights_after_loading(self.layer)
mock_format_cast.assert_not_called()
@patch("vllm_ascend.utils.get_ascend_config")
@mock.patch("torch_npu.npu_format_cast")
def test_process_weights_after_loading_with_nz1(self, mock_format_cast, mock_get_config):
mock_config = MagicMock()
mock_config.weight_nz_mode = 1
mock_get_config.return_value = mock_config
self.method.process_weights_after_loading(self.layer)
mock_format_cast.assert_not_called()
@patch("vllm_ascend.utils.get_ascend_config")
@mock.patch("torch_npu.npu_format_cast")
def test_process_weights_after_loading_with_nz2(self, mock_format_cast, mock_get_config):
mock_config = MagicMock()
mock_config.weight_nz_mode = 2
mock_get_config.return_value = mock_config
self.method.process_weights_after_loading(self.layer)
mock_format_cast.assert_called_once()
class TestAscendRowParallelLinear(BaseLinearTest):
@patch("vllm_ascend.ops.linear_op.get_weight_prefetch_method", return_value=MagicMock())
@patch("vllm_ascend.ops.linear.get_current_vllm_config", return_value=MagicMock())
@patch("vllm_ascend.ops.linear.enable_sp", return_value=False)
@patch(
"vllm_ascend.ops.linear.AscendUnquantizedLinearMethod.apply",
new=lambda self, layer, x, bias=None: torch.nn.functional.linear(x, layer.weight, bias),
)
def test_mlp_optimize(self, mock_enable_sp, mock_get_current_vllm_config, mock_get_weight_prefetch_method):
ascend_config._ASCEND_CONFIG = MagicMock()
ascend_config._ASCEND_CONFIG.recompute_scheduler_enable = False
ascend_config._ASCEND_CONFIG.finegrained_tp_config.mlp_tensor_parallel_size = 2
ascend_config._ASCEND_CONFIG.ascend_scheduler_config.enabled = False
linear = AscendRowParallelLinear(
input_size=16,
@@ -64,9 +113,18 @@ class TestAscendRowParallelLinear(BaseLinearTest):
input_tensor = torch.randn(16, 8)
linear(input_tensor)
def test_oproj_tp(self):
@patch("vllm_ascend.ops.linear_op.get_weight_prefetch_method", return_value=MagicMock())
@patch("vllm_ascend.ops.linear.get_current_vllm_config", return_value=MagicMock())
@patch("vllm_ascend.ops.linear.enable_sp", return_value=False)
@patch(
"vllm_ascend.ops.linear.AscendUnquantizedLinearMethod.apply",
new=lambda self, layer, x, bias=None: torch.nn.functional.linear(x, layer.weight, bias),
)
def test_oproj_tp(self, mock_enable_sp, mock_get_current_vllm_config, mock_get_weight_prefetch_method):
ascend_config._ASCEND_CONFIG = MagicMock()
ascend_config._ASCEND_CONFIG.oproj_tensor_parallel_size = 2
ascend_config._ASCEND_CONFIG.recompute_scheduler_enable = False
ascend_config._ASCEND_CONFIG.finegrained_tp_config.oproj_tensor_parallel_size = 2
ascend_config._ASCEND_CONFIG.ascend_scheduler_config.enabled = False
linear = AscendRowParallelLinear(
input_size=16,
@@ -80,9 +138,11 @@ class TestAscendRowParallelLinear(BaseLinearTest):
class TestAscendMergedColumnParallelLinear(BaseLinearTest):
def test_merged_mlp_tp_init(self):
os.environ["VLLM_ASCEND_ENABLE_MLP_OPTIMIZE"] = "1"
ascend_config._ASCEND_CONFIG = MagicMock()
ascend_config._ASCEND_CONFIG.recompute_scheduler_enable = False
ascend_config._ASCEND_CONFIG.finegrained_tp_config.mlp_tensor_parallel_size = 2
ascend_config._ASCEND_CONFIG.ascend_scheduler_config.enabled = False
linear = AscendMergedColumnParallelLinear(
input_size=16,
@@ -92,5 +152,21 @@ class TestAscendMergedColumnParallelLinear(BaseLinearTest):
self.assertEqual(linear.custom_op.comm_group, parallel_state._MLP_TP)
if __name__ == '__main__':
class TestAscendReplicatedLinear(BaseLinearTest):
def test_init_disable_tp(self):
linear = AscendReplicatedLinear(
input_size=16,
output_size=8,
)
self.assertTrue(isinstance(linear.quant_method, AscendUnquantizedLinearMethod))
def test_init_without_disable_tp(self):
linear = AscendReplicatedLinear(
input_size=16,
output_size=8,
)
self.assertTrue(isinstance(linear.quant_method, AscendUnquantizedLinearMethod))
if __name__ == "__main__":
unittest.main()

162
tests/ut/ops/test_mla.py Normal file
View File

@@ -0,0 +1,162 @@
from unittest.mock import MagicMock, patch
import torch
from torch import nn
from vllm.config import CacheConfig, CompilationConfig, VllmConfig
from vllm.forward_context import ForwardContext
from vllm.model_executor.layers.mla import MLAModules
from tests.ut.base import TestBase
from vllm_ascend.ops.mla import AscendMultiHeadLatentAttention, IndexerWrapper
class TestIndexerWrapper(TestBase):
def test_initialization(self):
mock_indexer = MagicMock()
mock_indexer.n_head = 64
mock_indexer.head_dim = 128
mock_indexer.topk_tokens = 2048
mock_indexer.q_lora_rank = 1536
mock_indexer.wq_b = nn.Linear(128, 128)
mock_indexer.wk_weights_proj = nn.Linear(128, 128)
mock_indexer.k_norm = nn.LayerNorm(128)
mock_indexer.softmax_scale = 0.123
mock_indexer.topk_indices_buffer = torch.randn(10)
mock_indexer.k_cache = torch.randn(10)
wrapper = IndexerWrapper(mock_indexer)
self.assertEqual(wrapper.n_head, 64)
self.assertEqual(wrapper.head_dim, 128)
self.assertEqual(wrapper.topk_tokens, 2048)
self.assertEqual(wrapper.q_lora_rank, 1536)
self.assertIs(wrapper.wq_b, mock_indexer.wq_b)
self.assertIs(wrapper.wk_weights_proj, mock_indexer.wk_weights_proj)
self.assertIs(wrapper.k_norm, mock_indexer.k_norm)
self.assertEqual(wrapper.softmax_scale, 0.123)
self.assertIsNone(mock_indexer.topk_indices_buffer)
self.assertIsNone(mock_indexer.k_cache)
def test_forward(self):
mock_indexer = MagicMock()
wrapper = IndexerWrapper(mock_indexer)
result = wrapper.forward()
self.assertIsNone(result)
class TestAscendMultiHeadLatentAttention(TestBase):
def setUp(self):
self.hidden_size = 4096
self.num_heads = 32
self.scale = 0.123
self.qk_nope_head_dim = 64
self.qk_rope_head_dim = 64
self.v_head_dim = 128
self.q_lora_rank = 1536
self.kv_lora_rank = 128
self.prefix = "model.layers.0.mla"
self.mock_mla_modules = MagicMock(spec=MLAModules)
self.mock_mla_modules.indexer = MagicMock()
self.mock_mla_modules.is_sparse = False
self.mock_mla_modules.rotary_emb = MagicMock()
self.mock_mla_modules.fused_qkv_a_proj = MagicMock()
self.mock_mla_modules.q_b_proj = MagicMock()
self.mock_mla_modules.q_a_layernorm = MagicMock()
self.mock_mla_modules.q_proj = MagicMock()
self.mock_mla_modules.kv_a_proj_with_mqa = MagicMock()
self.mock_mla_modules.kv_a_layernorm = MagicMock()
self.mock_mla_modules.kv_b_proj = MagicMock()
self.mock_mla_modules.o_proj = MagicMock()
self.mock_cache_config = MagicMock(spec=CacheConfig)
self.mock_quant_config = MagicMock()
@patch("vllm_ascend.ops.mla.get_current_vllm_config")
@patch("vllm_ascend.ops.mla.get_tensor_model_parallel_world_size")
def test_initialization(self, mock_tp_size, mock_get_vllm_config):
# Create a proper mock for MLAAttention that has the required attributes
mock_mla_attn = MagicMock()
mock_mla_attn.process_weights_after_loading = MagicMock()
mock_mla_attn.impl = MagicMock()
mock_mla_attn.impl.process_weights_after_loading = MagicMock()
with patch("vllm_ascend.ops.mla.MLAAttention", return_value=mock_mla_attn):
mock_tp_size.return_value = 2
mock_vllm_config = MagicMock(spec=VllmConfig)
mock_vllm_config.model_config.hf_text_config = MagicMock(num_hidden_layers=32, first_k_dense_replace=True)
mock_get_vllm_config.return_value = mock_vllm_config
mock_vllm_config.compilation_config = CompilationConfig()
attn = AscendMultiHeadLatentAttention(
hidden_size=self.hidden_size,
num_heads=self.num_heads,
scale=self.scale,
qk_nope_head_dim=self.qk_nope_head_dim,
qk_rope_head_dim=self.qk_rope_head_dim,
v_head_dim=self.v_head_dim,
q_lora_rank=self.q_lora_rank,
kv_lora_rank=self.kv_lora_rank,
mla_modules=self.mock_mla_modules,
cache_config=self.mock_cache_config,
quant_config=self.mock_quant_config,
prefix=self.prefix,
)
self.assertEqual(attn.tp_size, 2)
self.assertIsNotNone(attn.mla_attn)
@patch("vllm_ascend.ops.mla.torch.ops.vllm.mla_forward")
@patch("vllm_ascend.ops.mla.get_current_vllm_config")
@patch("vllm_ascend.ops.mla.get_tensor_model_parallel_world_size")
@patch("vllm_ascend.ops.mla.get_forward_context")
@patch("vllm_ascend.ascend_forward_context.get_forward_context")
def test_forward(
self,
mock_get_forward_context_2,
mock_get_forward_context,
mock_tp_size,
mock_get_vllm_config,
mock_mla_forward,
):
mock_tp_size.return_value = 1
mock_vllm_config = MagicMock(spec=VllmConfig)
mock_vllm_config.model_config.hf_text_config = MagicMock(num_hidden_layers=32, first_k_dense_replace=False)
mock_get_vllm_config.return_value = mock_vllm_config
mock_vllm_config.compilation_config = CompilationConfig()
# Create a proper mock for MLAAttention that has the required attributes
mock_mla_attn = MagicMock()
mock_mla_attn.process_weights_after_loading = MagicMock()
mock_mla_attn.impl = MagicMock()
mock_mla_attn.impl.process_weights_after_loading = MagicMock()
with patch("vllm_ascend.ops.mla.MLAAttention", return_value=mock_mla_attn):
attn = AscendMultiHeadLatentAttention(
hidden_size=self.hidden_size,
num_heads=self.num_heads,
scale=self.scale,
qk_nope_head_dim=self.qk_nope_head_dim,
qk_rope_head_dim=self.qk_rope_head_dim,
v_head_dim=self.v_head_dim,
q_lora_rank=self.q_lora_rank,
kv_lora_rank=self.kv_lora_rank,
mla_modules=self.mock_mla_modules,
cache_config=self.mock_cache_config,
quant_config=self.mock_quant_config,
prefix=self.prefix,
)
positions = torch.tensor([0, 1, 2])
hidden_states = torch.randn(3, self.hidden_size)
mock_forward_context = MagicMock(spec=ForwardContext)
mock_forward_context.flash_comm_v1_enabled = False
mock_get_forward_context.return_value = mock_forward_context
mock_get_forward_context_2.return_value = mock_forward_context
mock_mla_forward.return_value = (3, self.hidden_size)
output = attn.forward(positions, hidden_states)
self.assertEqual(output.shape, (3, self.hidden_size))

View File

@@ -0,0 +1,195 @@
from typing import Any
from unittest.mock import MagicMock, patch
import torch
from vllm.config import CompilationConfig, VllmConfig
from vllm.config.vllm import get_cached_compilation_config
from tests.ut.base import TestBase
from vllm_ascend.ops.mm_encoder_attention import (
MAX_PAD_SIZE,
AscendMMEncoderAttention,
)
from vllm_ascend.worker import encoder_acl_graph
from vllm_ascend.worker.encoder_acl_graph import (
get_encoder_graph_params,
set_encoder_forward_context,
set_encoder_graph_params,
)
class FIAMockMixin(TestBase):
captured: dict[str, Any]
def _install_vllm_config_mock(self):
mock_vllm_config = MagicMock(spec=VllmConfig)
mock_vllm_config.compilation_config = CompilationConfig()
patcher = patch(
"vllm.config.vllm.get_current_vllm_config",
return_value=mock_vllm_config,
)
patcher.start()
self.addCleanup(patcher.stop)
get_cached_compilation_config.cache_clear()
self.addCleanup(get_cached_compilation_config.cache_clear)
def _make_layer(self, num_heads=4, num_kv_heads=4, head_size=72, scale=None):
return AscendMMEncoderAttention(
num_heads=num_heads,
head_size=head_size,
scale=scale,
num_kv_heads=num_kv_heads,
)
def _fake_fia(self, **kwargs):
self.captured = {
"mode": "functional",
"q_shape": kwargs["query"].shape,
"input_layout": kwargs["input_layout"],
"actual_seq_lengths": kwargs["actual_seq_lengths"],
}
return torch.zeros_like(kwargs["query"]), None
def _fake_fia_out(self, *, workspace, out, **kwargs):
self.captured = {"mode": "out", "softmax_lse": out[1]}
out[0].zero_()
def _install_fia_mocks(self, *, capture: bool):
self.captured = {}
mock_fia = MagicMock(side_effect=self._fake_fia)
mock_fia.out = self._fake_fia_out
patch_targets: list[tuple[str, Any]] = [
(
"vllm_ascend.ops.mm_encoder_attention.torch_npu.npu_fused_infer_attention_score",
mock_fia,
),
(
"vllm_ascend.ops.mm_encoder_attention.torch_npu._npu_fused_infer_attention_score_get_max_workspace",
MagicMock(return_value=torch.zeros(1)),
),
]
if capture:
self.mock_graph_begin = MagicMock()
self.mock_graph_end = MagicMock(return_value=42)
mock_event = MagicMock()
patch_targets.extend(
[
(
"vllm_ascend.ops.mm_encoder_attention.weak_ref_tensors",
lambda tensors: tensors,
),
(
"vllm_ascend.ops.mm_encoder_attention.torch_npu.npu.current_stream",
MagicMock(return_value=MagicMock()),
),
(
"vllm_ascend.ops.mm_encoder_attention.torch.npu.ExternalEvent",
MagicMock(return_value=mock_event),
),
(
"vllm_ascend.ops.mm_encoder_attention.torch.npu.graph_task_group_begin",
self.mock_graph_begin,
),
(
"vllm_ascend.ops.mm_encoder_attention.torch.npu.graph_task_group_end",
self.mock_graph_end,
),
]
)
for target, replacement in patch_targets:
patcher = patch(target, replacement)
patcher.start()
self.addCleanup(patcher.stop)
class TestAscendMMEncoderAttentionEager(FIAMockMixin):
def setUp(self):
self._install_vllm_config_mock()
self._install_fia_mocks(capture=False)
def test_forward_oot_basic(self):
layer = self._make_layer(num_heads=4, num_kv_heads=4, head_size=128)
bsz, q_len = 2, 4
query = torch.randn(bsz, q_len, layer.num_heads * layer.head_size)
key = query.clone()
value = query.clone()
cu_seqlens = torch.arange(0, (bsz + 1) * q_len, step=q_len, dtype=torch.int32)
out = layer.forward_oot(query, key, value, cu_seqlens=cu_seqlens)
self.assertEqual(out.shape, (bsz, q_len, layer.num_heads * layer.head_size))
self.assertEqual(self.captured["mode"], "functional")
self.assertEqual(self.captured["input_layout"], "TND")
def test_forward_oot_seqlens(self):
layer = self._make_layer(num_heads=4, num_kv_heads=4, head_size=72)
seq_lens = [3, 7, 2]
cu_seqlens = torch.tensor([0, 3, 10, 12], dtype=torch.int32, device="cpu")
max_q_len = max(seq_lens)
query = torch.randn(len(seq_lens), max_q_len, layer.num_heads, 72, dtype=torch.bfloat16)
key = torch.randn_like(query)
value = torch.randn_like(query)
out = layer.forward_oot(query, key, value, cu_seqlens=cu_seqlens)
self.assertEqual(out.shape, query.shape)
self.assertEqual(self.captured["actual_seq_lengths"], [3, 10, 12])
self.assertEqual(self.captured["q_shape"], (len(seq_lens) * max_q_len, 4, MAX_PAD_SIZE))
class TestAscendMMEncoderAttentionCapture(FIAMockMixin):
def setUp(self):
self._install_vllm_config_mock()
set_encoder_graph_params([2048])
self._install_fia_mocks(capture=True)
def tearDown(self):
encoder_acl_graph._encoder_graph_params = None
encoder_acl_graph._reset_encoder_forward_context()
def test_forward_oot_basic(self):
layer = self._make_layer(num_heads=4, num_kv_heads=4, head_size=72)
bsz, q_len = 2, 4
query = torch.randn(bsz, q_len, layer.num_heads, 72, dtype=torch.bfloat16)
key = torch.randn_like(query)
value = torch.randn_like(query)
cu_seqlens = torch.arange(0, (bsz + 1) * q_len, step=q_len, dtype=torch.int32)
with set_encoder_forward_context(2048, True):
layer.forward_oot(query, key, value, cu_seqlens=cu_seqlens)
params = get_encoder_graph_params()
self.assertIsNotNone(params)
self.assertEqual(len(params.attn_params[2048]), 1)
self.assertEqual(len(params.handles[2048]), 1)
self.assertEqual(self.captured["mode"], "out")
self.mock_graph_begin.assert_called_once()
self.mock_graph_end.assert_called_once()
def test_forward_oot_seqlens(self):
layer = self._make_layer(num_heads=4, num_kv_heads=4, head_size=72)
seq_lens = [3, 7, 2]
cu_seqlens = torch.tensor([0, 3, 10, 12], dtype=torch.int32, device="cpu")
max_q_len = max(seq_lens)
query = torch.randn(len(seq_lens), max_q_len, layer.num_heads, 72, dtype=torch.bfloat16)
key = torch.randn_like(query)
value = torch.randn_like(query)
captured_lengths: list[Any] = []
def capture_workspace(**kwargs):
captured_lengths.append(kwargs.get("actual_seq_lengths"))
return torch.zeros(1)
with (
patch(
"vllm_ascend.ops.mm_encoder_attention.torch_npu._npu_fused_infer_attention_score_get_max_workspace",
side_effect=capture_workspace,
),
set_encoder_forward_context(2048, True),
):
layer.forward_oot(query, key, value, cu_seqlens=cu_seqlens)
self.assertEqual(captured_lengths[-1], [7, 14, 21])

View File

@@ -4,13 +4,38 @@ import torch
from vllm.model_executor.layers.fused_moe import FusedMoEConfig
from tests.ut.base import TestBase
from vllm_ascend.ops.moe.moe_comm_method import (AllGatherCommImpl,
AlltoAllCommImpl, MC2CommImpl)
from vllm_ascend.ops.fused_moe.moe_comm_method import (
AllGatherCommImpl,
AlltoAllCommImpl,
MC2CommImpl,
)
from vllm_ascend.ops.fused_moe.moe_runtime_args import (
MoEAllGatherCombineMetadata,
MoEFusedExpertsInput,
MoEPrepareOutput,
MoEQuantParams,
MoERoutingParams,
MoEWeights,
)
from vllm_ascend.ops.fused_moe.token_dispatcher import MoETokenDispatchOutput
from vllm_ascend.quantization.methods.base import QuantType
class TestMoECommMethod(TestBase):
def setUp(self):
self.mock_ascend_config = MagicMock()
self.mock_ascend_config.ascend_fusion_config.fusion_ops_gmmswigluquant = False
self.mock_ascend_config.enable_fused_mc2 = False
self._patch_get_ascend_config = patch(
"vllm_ascend.ops.fused_moe.moe_comm_method.get_ascend_config",
return_value=self.mock_ascend_config,
)
self._patch_get_ascend_config_module = patch(
"vllm_ascend.ascend_config.get_ascend_config",
return_value=self.mock_ascend_config,
)
self._patch_get_ascend_config.start()
self._patch_get_ascend_config_module.start()
# Mock FusedMoEConfig
self.moe_config = MagicMock(spec=FusedMoEConfig)
self.moe_config.num_experts = 8
@@ -20,23 +45,19 @@ class TestMoECommMethod(TestBase):
self.moe_config.tp_group.device_group = MagicMock()
self.moe_config.dp_size = 1
self.moe_config.tp_size = 1
self.moe_config.pcp_size = 1
self.moe_config.ep_size = 1
self.moe_config.dp_group = MagicMock()
self.moe_config.num_global_redundant_experts = 0
self.moe_config.global_redundant_expert_num = 0
@patch("vllm_ascend.ops.moe.moe_comm_method.get_current_vllm_config")
@patch("vllm_ascend.ops.moe.moe_comm_method.get_forward_context")
@patch(
"vllm_ascend.ops.moe.moe_comm_method.FusedMoEPrepareAndFinalizeWithAllGather"
)
@patch("vllm_ascend.ops.moe.moe_comm_method.TokenDispatcherWithAllGather")
def test_all_gather_comm_impl(self, mock_token_dispatcher,
mock_prepare_finalize,
mock_get_forward_context,
mock_get_current_vllm_config):
# Mock vLLM config
mock_get_current_vllm_config.return_value = MagicMock()
def tearDown(self):
self._patch_get_ascend_config.stop()
self._patch_get_ascend_config_module.stop()
@patch("vllm_ascend.ascend_forward_context.get_forward_context")
@patch("vllm_ascend.ops.fused_moe.moe_comm_method.PrepareAndFinalizeWithAllGather")
@patch("vllm_ascend.ops.fused_moe.moe_comm_method.TokenDispatcherWithAllGather")
def test_all_gather_comm_impl(self, mock_token_dispatcher, mock_prepare_finalize, mock_get_forward_context):
# Mock forward context
mock_context = MagicMock()
mock_context.moe_comm_method = "all_gather"
@@ -44,8 +65,12 @@ class TestMoECommMethod(TestBase):
# Mock prepare finalize
mock_pf_instance = MagicMock()
mock_pf_instance.prepare.return_value = (torch.randn(4, 8),
torch.randn(4, 2), None)
mock_pf_instance.prepare.return_value = MoEPrepareOutput(
hidden_states=torch.randn(4, 8),
router_logits=torch.randn(4, 2),
mc2_mask=None,
padded_hidden_states_shape=None,
)
mock_pf_instance.finalize.return_value = torch.randn(4, 8)
mock_prepare_finalize.return_value = mock_pf_instance
@@ -59,28 +84,21 @@ class TestMoECommMethod(TestBase):
# Test prepare method
hidden_states = torch.randn(3, 8)
router_logits = torch.randn(3, 2)
h_out, r_out = comm_impl.prepare(hidden_states, router_logits)
prepare_output = comm_impl.prepare(hidden_states, router_logits)
h_out = prepare_output.hidden_states
padded_hidden_states_shape = prepare_output.padded_hidden_states_shape
# Verify prepare was called with correct arguments
mock_pf_instance.prepare.assert_called_once_with(
hidden_states, router_logits, False, False, False, None)
mock_pf_instance.prepare.assert_called_once_with(hidden_states, router_logits, False, False, QuantType.NONE)
# Test finalize method
comm_impl.finalize(h_out, reduce_results=True)
mock_pf_instance.finalize.assert_called_once_with(h_out, True)
@patch("vllm_ascend.ops.moe.moe_comm_method.get_current_vllm_config")
@patch("vllm_ascend.ops.moe.moe_comm_method.get_forward_context")
@patch(
"vllm_ascend.ops.moe.moe_comm_method.FusedMoEPrepareAndFinalizeWithMC2"
)
@patch("vllm_ascend.ops.moe.moe_comm_method.TokenDispatcherWithMC2")
def test_mc2_comm_impl(self, mock_token_dispatcher, mock_prepare_finalize,
mock_get_forward_context,
mock_get_current_vllm_config):
# Mock vLLM config
mock_get_current_vllm_config.return_value = MagicMock()
comm_impl.finalize(h_out, reduce_results=True, padded_hidden_states_shape=padded_hidden_states_shape)
mock_pf_instance.finalize.assert_called_once_with(h_out, True, None)
@patch("vllm_ascend.ascend_forward_context.get_forward_context")
@patch("vllm_ascend.ops.fused_moe.moe_comm_method.PrepareAndFinalizeWithMC2")
@patch("vllm_ascend.ops.fused_moe.moe_comm_method.TokenDispatcherWithMC2")
def test_mc2_comm_impl(self, mock_token_dispatcher, mock_prepare_finalize, mock_get_forward_context):
# Mock forward context
mock_context = MagicMock()
mock_context.moe_comm_method = "mc2"
@@ -88,9 +106,12 @@ class TestMoECommMethod(TestBase):
# Mock prepare finalize
mock_pf_instance = MagicMock()
mock_pf_instance.prepare.return_value = (torch.randn(4, 8),
torch.randn(4, 2),
torch.tensor([1, 0, 1, 0]))
mock_pf_instance.prepare.return_value = MoEPrepareOutput(
hidden_states=torch.randn(4, 8),
router_logits=torch.randn(4, 2),
mc2_mask=torch.tensor([1, 0, 1, 0]),
padded_hidden_states_shape=None,
)
mock_pf_instance.finalize.return_value = torch.randn(4, 8)
mock_prepare_finalize.return_value = mock_pf_instance
@@ -104,29 +125,21 @@ class TestMoECommMethod(TestBase):
# Test prepare method
hidden_states = torch.randn(3, 8)
router_logits = torch.randn(3, 2)
h_out, r_out = comm_impl.prepare(hidden_states, router_logits)
prepare_output = comm_impl.prepare(hidden_states, router_logits)
h_out = prepare_output.hidden_states
padded_hidden_states_shape = prepare_output.padded_hidden_states_shape
# Verify prepare was called with correct arguments
mock_pf_instance.prepare.assert_called_once_with(
hidden_states, router_logits, False, False, False, None)
mock_pf_instance.prepare.assert_called_once_with(hidden_states, router_logits, False, False, QuantType.NONE)
# Test finalize method
comm_impl.finalize(h_out, reduce_results=True)
mock_pf_instance.finalize.assert_called_once_with(h_out, True)
@patch("vllm_ascend.ops.moe.moe_comm_method.get_current_vllm_config")
@patch("vllm_ascend.ops.moe.moe_comm_method.get_forward_context")
@patch(
"vllm_ascend.ops.moe.moe_comm_method.FusedMoEPrepareAndFinalizeWithAll2All"
)
@patch("vllm_ascend.ops.moe.moe_comm_method.TokenDispatcherWithAll2AllV")
def test_alltoall_comm_impl(self, mock_token_dispatcher,
mock_prepare_finalize,
mock_get_forward_context,
mock_get_current_vllm_config):
# Mock vLLM config
mock_get_current_vllm_config.return_value = MagicMock()
comm_impl.finalize(h_out, reduce_results=True, padded_hidden_states_shape=padded_hidden_states_shape)
mock_pf_instance.finalize.assert_called_once_with(h_out, True, None)
@patch("vllm_ascend.ascend_forward_context.get_forward_context")
@patch("vllm_ascend.ops.fused_moe.moe_comm_method.PrepareAndFinalizeWithAll2All")
@patch("vllm_ascend.ops.fused_moe.moe_comm_method.TokenDispatcherWithAll2AllV")
def test_alltoall_comm_impl(self, mock_token_dispatcher, mock_prepare_finalize, mock_get_forward_context):
# Mock forward context
mock_context = MagicMock()
mock_context.moe_comm_method = "alltoall"
@@ -134,8 +147,12 @@ class TestMoECommMethod(TestBase):
# Mock prepare finalize
mock_pf_instance = MagicMock()
mock_pf_instance.prepare.return_value = (torch.randn(4, 8),
torch.randn(4, 2), None)
mock_pf_instance.prepare.return_value = MoEPrepareOutput(
hidden_states=torch.randn(4, 8),
router_logits=torch.randn(4, 2),
mc2_mask=None,
padded_hidden_states_shape=None,
)
mock_pf_instance.finalize.return_value = torch.randn(4, 8)
mock_prepare_finalize.return_value = mock_pf_instance
@@ -149,26 +166,19 @@ class TestMoECommMethod(TestBase):
# Test prepare method
hidden_states = torch.randn(3, 8)
router_logits = torch.randn(3, 2)
h_out, r_out = comm_impl.prepare(hidden_states, router_logits)
_ = comm_impl.prepare(hidden_states, router_logits)
# Verify prepare was called with correct arguments
mock_pf_instance.prepare.assert_called_once_with(
hidden_states, router_logits, False, False, False, None)
@patch("vllm_ascend.ops.moe.moe_comm_method.get_current_vllm_config")
@patch("vllm_ascend.ops.moe.moe_comm_method.get_forward_context")
@patch(
"vllm_ascend.ops.moe.moe_comm_method.FusedMoEPrepareAndFinalizeWithAllGather"
)
@patch("vllm_ascend.ops.moe.moe_comm_method.TokenDispatcherWithAllGather")
@patch("vllm_ascend.ops.moe.moe_comm_method.unified_apply_mlp")
def test_fused_experts_method(self, mock_unified_apply_mlp,
mock_token_dispatcher, mock_prepare_finalize,
mock_get_forward_context,
mock_get_current_vllm_config):
# Mock vLLM config
mock_get_current_vllm_config.return_value = MagicMock()
mock_pf_instance.prepare.assert_called_once_with(hidden_states, router_logits, False, False, QuantType.NONE)
@patch("vllm_ascend.ascend_forward_context.get_forward_context")
@patch("vllm_ascend.ops.fused_moe.moe_comm_method.PrepareAndFinalizeWithAllGather")
@patch("vllm_ascend.ops.fused_moe.moe_comm_method.TokenDispatcherWithAllGather")
@patch("vllm_ascend.ops.fused_moe.moe_comm_method.unified_apply_mlp")
@patch("torch.npu.current_stream", MagicMock())
def test_fused_experts_method(
self, mock_unified_apply_mlp, mock_token_dispatcher, mock_prepare_finalize, mock_get_forward_context
):
# Mock forward context
mock_context = MagicMock()
mock_context.moe_comm_method = "all_gather"
@@ -176,23 +186,33 @@ class TestMoECommMethod(TestBase):
# Mock prepare finalize
mock_pf_instance = MagicMock()
mock_pf_instance.prepare.return_value = (torch.randn(4, 8),
torch.randn(4, 2), None)
mock_pf_instance.prepare.return_value = MoEPrepareOutput(
hidden_states=torch.randn(4, 8),
router_logits=torch.randn(4, 2),
mc2_mask=None,
padded_hidden_states_shape=None,
)
mock_pf_instance.finalize.return_value = torch.randn(4, 8)
mock_prepare_finalize.return_value = mock_pf_instance
# Mock token dispatcher
mock_td_instance = MagicMock()
mock_td_instance.token_dispatch.return_value = {
"hidden_states": torch.randn(6, 8),
"group_list": torch.tensor([2, 2, 2]),
"group_list_type": 1
}
dispatch_topk_weights = torch.tensor([[0.5, 0.5], [0.3, 0.7], [0.8, 0.2], [0.6, 0.4]])
mock_td_instance.token_dispatch.return_value = MoETokenDispatchOutput(
hidden_states=torch.randn(6, 8),
group_list=torch.tensor([2, 2, 2]),
group_list_type=1,
combine_metadata=MoEAllGatherCombineMetadata(
topk_weights=dispatch_topk_weights,
expanded_row_idx=torch.arange(8, dtype=torch.int32),
restore_shape=torch.Size([4, 8]),
),
)
mock_td_instance.token_combine.return_value = torch.randn(4, 8)
mock_token_dispatcher.return_value = mock_td_instance
# Mock unified_apply_mlp
mock_unified_apply_mlp.return_value = torch.randn(6, 8)
# Mock unified_apply_mlp returns (tensor, event) tuple
mock_unified_apply_mlp.return_value = (torch.randn(6, 8), MagicMock())
# Create instance
comm_impl = AllGatherCommImpl(self.moe_config)
@@ -201,32 +221,50 @@ class TestMoECommMethod(TestBase):
hidden_states = torch.randn(4, 8).contiguous()
w1 = torch.randn(16, 8).contiguous()
w2 = torch.randn(16, 8).contiguous()
topk_weights = torch.tensor([[0.5, 0.5], [0.3, 0.7], [0.8, 0.2],
[0.6, 0.4]])
topk_weights = dispatch_topk_weights
topk_ids = torch.tensor([[0, 1], [1, 2], [2, 0], [1, 1]])
row_idx = torch.arange(4)
# Make sure tensors are contiguous and have correct strides
hidden_states = hidden_states.contiguous()
w1 = w1.contiguous()
w2 = w2.contiguous()
result = comm_impl.fused_experts(hidden_states=hidden_states,
w1=w1,
w2=w2,
topk_weights=topk_weights,
topk_ids=topk_ids,
row_idx=row_idx,
activation="silu")
result = comm_impl.fused_experts(
fused_experts_input=MoEFusedExpertsInput(
hidden_states=hidden_states,
topk_weights=topk_weights,
topk_ids=topk_ids,
weights=MoEWeights(
w1=[w1],
w2=[w2],
),
routing=MoERoutingParams(
expert_map=None,
global_redundant_expert_num=0,
mc2_mask=None,
apply_router_weight_on_input=False,
),
activation="silu",
need_trans=False,
dynamic_eplb=False,
quant=MoEQuantParams(),
)
)
# Verify result shape
self.assertEqual(result.shape, (4, 8))
self.assertEqual(result.routed_out.shape, (4, 8))
# Verify token_dispatch was called
mock_td_instance.token_dispatch.assert_called_once()
# Verify unified_apply_mlp was called
mock_unified_apply_mlp.assert_called_once()
mlp_compute_input = mock_unified_apply_mlp.call_args.kwargs["mlp_compute_input"]
self.assertFalse(mlp_compute_input.fusion)
self.assertFalse(mlp_compute_input.quant.is_mxfp)
# Verify token_combine was called
mock_td_instance.token_combine.assert_called_once()
mock_td_instance.token_combine.assert_called_once_with(
hidden_states=mock_unified_apply_mlp.return_value[0],
combine_metadata=mock_td_instance.token_dispatch.return_value.combine_metadata,
)

View File

@@ -0,0 +1,253 @@
import unittest
from typing import ClassVar
from unittest.mock import patch
import torch
from vllm.model_executor.layers.fused_moe.activation import MoEActivation
from vllm_ascend.ops.fused_moe.moe_mlp import cumsum_group_list, unified_apply_mlp, unquant_apply_mlp
from vllm_ascend.ops.fused_moe.moe_runtime_args import (
MoEMlpComputeInput,
MoEQuantParams,
MoEWeights,
)
from vllm_ascend.ops.fused_moe.moe_stage_params import MoEMxfpParams
from vllm_ascend.quantization.quant_type import QuantType
MXFP4_TEST_DTYPE = getattr(torch, "float4_e2m1fn_x2", torch.float16)
class TestCumsumGroupList(unittest.TestCase):
glist_dict: ClassVar[dict[int, torch.Tensor]]
@classmethod
def setUpClass(cls):
cls.glist_dict = {
0: torch.tensor([0, 2, 3, 3]),
1: torch.tensor([0, 2, 1, 0]),
2: torch.tensor([[1, 2], [2, 1], [0, 0], [0, 0]]),
}
support_combine = [(0, 0), (1, 0), (0, 1)]
unsupported_combine = [(0, 2), (2, 1), (1, 2)]
def test_cumsum_group_list_supported_conversion(self):
for src_list_type, dst_list_type in self.support_combine:
with self.subTest(src=src_list_type, dst=dst_list_type):
result = cumsum_group_list(self.glist_dict[src_list_type], src_list_type, dst_list_type, expert_num=4)
self.assertTrue(torch.equal(result, self.glist_dict[dst_list_type]))
def test_cumsum_group_list_invalid_type_valueerror(self):
with self.assertRaises(ValueError) as excinfo:
cumsum_group_list(self.glist_dict[0], 4, 0)
self.assertIn("group_list_type should be in [0, 1, 2], but received", str(excinfo.exception))
def test_cumsum_group_list_unsupported_conversion_notimplementederror(self):
for src_list_type, dst_list_type in self.unsupported_combine:
with self.subTest(src=src_list_type, dst=dst_list_type):
with self.assertRaises(NotImplementedError) as excinfo:
cumsum_group_list(self.glist_dict[0], src_list_type, dst_list_type)
self.assertIn("This feature is under development.", str(excinfo.exception))
class TestW4A8RuntimeFlags(unittest.TestCase):
def test_w4a8_per_channel_gmm_swiglu_flag(self):
self.assertTrue(
MoEQuantParams(quant_type=QuantType.W4A8, is_per_channel_weight=True).use_w4a8_per_channel_gmm_swiglu
)
self.assertFalse(
MoEQuantParams(quant_type=QuantType.W4A8, is_per_channel_weight=False).use_w4a8_per_channel_gmm_swiglu
)
self.assertFalse(
MoEQuantParams(quant_type=QuantType.W8A8, is_per_channel_weight=True).use_w4a8_per_channel_gmm_swiglu
)
class TestUnifiedApplyMlpRequest(unittest.TestCase):
def test_unquant_apply_mlp_wraps_tensor_weights_for_grouped_matmul(self):
hidden_states = torch.randn(2, 8)
gate_up_out = torch.randn(2, 16)
expected = torch.randn(2, 8)
w1 = torch.randn(2, 8, 16)
w2 = torch.randn(2, 8, 8)
with (
patch(
"vllm_ascend.ops.fused_moe.moe_mlp.torch_npu.npu_grouped_matmul",
side_effect=[[gate_up_out], [expected]],
create=True,
) as mock_grouped_matmul,
patch(
"vllm_ascend.ops.fused_moe.moe_mlp.torch_npu.npu_swiglu",
return_value=gate_up_out,
create=True,
),
):
output, _ = unquant_apply_mlp(
hidden_states=hidden_states,
w1=w1,
w2=w2,
group_list=torch.tensor([1, 1]),
need_trans=True,
)
self.assertTrue(output is expected)
first_call, second_call = mock_grouped_matmul.call_args_list
self.assertEqual(len(first_call.kwargs["weight"]), 1)
self.assertEqual(len(second_call.kwargs["weight"]), 1)
self.assertEqual(first_call.kwargs["weight"][0].shape, torch.Size([2, 16, 8]))
self.assertEqual(second_call.kwargs["weight"][0].shape, torch.Size([2, 8, 8]))
def test_request_unquant_path(self):
hidden_states = torch.randn(2, 8)
expected = torch.randn(2, 8)
mlp_compute_input = MoEMlpComputeInput(
hidden_states=hidden_states,
group_list=torch.tensor([2, 2], dtype=torch.int64),
group_list_type=1,
dynamic_scale=None,
topk_scales=None,
weights=MoEWeights(
w1=torch.randn(1, 16, 8),
w2=torch.randn(1, 8, 8),
w1_bias=torch.randn(1, 16),
w2_bias=torch.randn(1, 8),
),
quant=MoEQuantParams(quant_type=QuantType.NONE),
fusion=False,
activation="silu",
need_trans=False,
dynamic_eplb=False,
)
with (
patch("vllm_ascend.ops.fused_moe.moe_mlp.unquant_apply_mlp", return_value=expected) as mock_unquant,
patch("vllm_ascend.ops.fused_moe.moe_mlp.quant_apply_mlp") as mock_quant,
):
output = unified_apply_mlp(mlp_compute_input=mlp_compute_input)
self.assertTrue(output is expected)
mock_unquant.assert_called_once()
self.assertEqual(mock_unquant.call_args.kwargs["activation"], "silu")
self.assertFalse(mock_unquant.call_args.kwargs["need_trans"])
mock_quant.assert_not_called()
def test_request_quant_path(self):
for quant_type, mxfp_dtype in (
(QuantType.MXFP8, torch.float8_e4m3fn),
(QuantType.MXFP4, MXFP4_TEST_DTYPE),
):
with self.subTest(quant_type=quant_type):
hidden_states = torch.randn(2, 8)
expected = torch.randn(2, 8)
mlp_compute_input = MoEMlpComputeInput(
hidden_states=hidden_states,
group_list=torch.tensor([2, 2], dtype=torch.int64),
group_list_type=1,
dynamic_scale=torch.randn(2, 1),
topk_scales=None,
weights=MoEWeights(
w1=torch.randn(1, 16, 8),
w2=torch.randn(1, 8, 8),
w1_scale=[torch.randn(1)],
w2_scale=[torch.randn(1)],
),
quant=MoEQuantParams(
quant_type=quant_type,
mxfp=MoEMxfpParams(
act_quant_type=mxfp_dtype,
weight_quant_type=mxfp_dtype,
use_bf16=False,
),
),
fusion=True,
activation="silu",
need_trans=False,
dynamic_eplb=True,
)
with (
patch("vllm_ascend.ops.fused_moe.moe_mlp.quant_apply_mlp", return_value=expected) as mock_quant,
patch("vllm_ascend.ops.fused_moe.moe_mlp.unquant_apply_mlp") as mock_unquant,
):
output = unified_apply_mlp(mlp_compute_input=mlp_compute_input)
self.assertTrue(output is expected)
mock_quant.assert_called_once()
quant_kwargs = mock_quant.call_args.kwargs
self.assertTrue(quant_kwargs["use_mxfp_quant"])
self.assertTrue(quant_kwargs["fusion"])
self.assertTrue(quant_kwargs["dynamic_eplb"])
self.assertEqual(quant_kwargs["act_quant_type"], mxfp_dtype)
self.assertEqual(quant_kwargs["weight_quant_type"], mxfp_dtype)
self.assertFalse(quant_kwargs["use_bf16"])
mock_unquant.assert_not_called()
def test_request_quant_path_passes_w4a8_per_channel_flag(self):
hidden_states = torch.randn(2, 8)
expected = torch.randn(2, 8)
mlp_compute_input = MoEMlpComputeInput(
hidden_states=hidden_states,
group_list=torch.tensor([2, 2], dtype=torch.int64),
group_list_type=1,
dynamic_scale=torch.randn(2, 1),
topk_scales=None,
weights=MoEWeights(
w1=torch.randn(1, 16, 8),
w2=torch.randn(1, 8, 8),
w1_scale=[torch.randn(1, 16)],
w2_scale=[torch.randn(1, 8)],
),
quant=MoEQuantParams(quant_type=QuantType.W4A8, is_per_channel_weight=True),
fusion=False,
activation="silu",
need_trans=False,
dynamic_eplb=False,
)
with (
patch("vllm_ascend.ops.fused_moe.moe_mlp.quant_apply_mlp", return_value=expected) as mock_quant,
patch("vllm_ascend.ops.fused_moe.moe_mlp.unquant_apply_mlp") as mock_unquant,
):
output = unified_apply_mlp(mlp_compute_input=mlp_compute_input)
self.assertTrue(output is expected)
quant_kwargs = mock_quant.call_args.kwargs
self.assertTrue(quant_kwargs["use_w4a8_per_channel_gmm_swiglu"])
mock_unquant.assert_not_called()
def test_request_quant_path_passes_swiglustep_activation(self):
expected = torch.randn(1, 2)
mlp_compute_input = MoEMlpComputeInput(
hidden_states=torch.ones((1, 2), dtype=torch.float32),
group_list=torch.tensor([1], dtype=torch.int64),
group_list_type=1,
dynamic_scale=None,
topk_scales=None,
weights=MoEWeights(
w1=[torch.ones((1, 2, 4), dtype=torch.float32)],
w2=[torch.ones((1, 2, 2), dtype=torch.float32)],
w1_scale=[torch.ones((1,), dtype=torch.float32)],
w2_scale=[torch.ones((1,), dtype=torch.float32)],
),
quant=MoEQuantParams(quant_type=QuantType.W8A8),
fusion=True,
activation=MoEActivation.SWIGLUSTEP,
swiglu_limit=5.0,
)
with (
patch("vllm_ascend.ops.fused_moe.moe_mlp.quant_apply_mlp", return_value=expected) as mock_quant,
patch("vllm_ascend.ops.fused_moe.moe_mlp.unquant_apply_mlp") as mock_unquant,
):
output = unified_apply_mlp(mlp_compute_input=mlp_compute_input)
self.assertTrue(output is expected)
quant_kwargs = mock_quant.call_args.kwargs
self.assertEqual(quant_kwargs["activation"], MoEActivation.SWIGLUSTEP)
self.assertEqual(quant_kwargs["swiglu_limit"], 5.0)
mock_unquant.assert_not_called()
if __name__ == "__main__":
unittest.main(verbosity=2)

View File

@@ -0,0 +1,270 @@
#
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
# Copyright 2023 The vLLM team.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
#
import unittest
import torch
import vllm_ascend.ops.fused_moe.moe_runtime_args as runtime_args
from vllm_ascend.ops.fused_moe.moe_runtime_args import (
MoEAllGatherCombineMetadata,
MoETokenDispatchOutput,
MoEWeights,
build_fused_experts_input,
build_mlp_compute_input,
build_token_dispatch_input,
)
from vllm_ascend.quantization.quant_type import QuantType
MXFP4_TEST_DTYPE = getattr(torch, "float4_e2m1fn_x2", torch.float16)
def _get_test_mxfp_dtype(quant_type: QuantType) -> torch.dtype | None:
if quant_type == QuantType.MXFP8:
return torch.float8_e4m3fn
if quant_type == QuantType.MXFP4:
return MXFP4_TEST_DTYPE
if quant_type == QuantType.W4A8MXFP:
return torch.float8_e4m3fn
return None
class TestMoERuntimeArgs(unittest.TestCase):
def test_runtime_args_facade_exports_public_contracts_and_builders(self):
expected_symbols = [
"MoEAllGatherCombineMetadata",
"MoEAllToAllCombineMetadata",
"MoEFusedExpertsInput",
"MoEMC2CombineMetadata",
"MoEMlpComputeInput",
"MoEPrepareOutput",
"MoEQuantParams",
"MoERoutingParams",
"MoETokenDispatchInput",
"MoETokenDispatchOutput",
"MoEWeights",
"TMoECombineMetadata",
"build_fused_experts_input",
"build_mlp_compute_input",
"build_token_dispatch_input",
]
for symbol in expected_symbols:
with self.subTest(symbol=symbol):
self.assertTrue(hasattr(runtime_args, symbol))
self.assertFalse(hasattr(runtime_args, "MoEMxfpParams"))
def test_build_fused_experts_input_preserves_runtime_semantics(self):
for quant_type in (
QuantType.NONE,
QuantType.W4A16,
QuantType.W4A8,
QuantType.W8A8,
QuantType.MXFP8,
QuantType.MXFP4,
QuantType.W4A8MXFP,
):
with self.subTest(quant_type=quant_type):
hidden_states = torch.randn(4, 8)
topk_weights = torch.randn(4, 2)
topk_ids = torch.randint(0, 4, (4, 2), dtype=torch.int32)
fused_experts_input = build_fused_experts_input(
hidden_states=hidden_states,
topk_weights=topk_weights,
topk_ids=topk_ids,
w1=torch.randn(2, 8, 16),
w2=torch.randn(2, 16, 8),
quant_type=quant_type,
dynamic_eplb=True,
expert_map=torch.tensor([0, 1, 2, 3], dtype=torch.int32),
global_redundant_expert_num=2,
mc2_mask=torch.tensor([True, False, True, False]),
apply_router_weight_on_input=True,
log2phy=torch.tensor([3, 2, 1, 0], dtype=torch.int32),
pertoken_scale=torch.randn(4),
activation="gelu",
mxfp_act_quant_type=_get_test_mxfp_dtype(quant_type),
)
self.assertIs(fused_experts_input.hidden_states, hidden_states)
self.assertIs(fused_experts_input.topk_weights, topk_weights)
self.assertIs(fused_experts_input.topk_ids, topk_ids)
self.assertTrue(fused_experts_input.dynamic_eplb)
self.assertTrue(fused_experts_input.routing.apply_router_weight_on_input)
self.assertEqual(fused_experts_input.routing.global_redundant_expert_num, 2)
self.assertEqual(fused_experts_input.activation, "gelu")
self.assertEqual(fused_experts_input.quant.quant_type, quant_type)
def test_build_fused_experts_input_merges_dense_and_quant_weights(self):
w1 = torch.randn(2, 8, 16)
w2 = torch.randn(2, 16, 8)
w1_scale = [torch.randn(1)]
w2_scale = [torch.randn(1)]
w1_scale_bias = torch.randn(1)
w2_scale_bias = torch.randn(1)
w1_offset = torch.randn(1)
w2_offset = torch.randn(1)
fused_experts_input = build_fused_experts_input(
hidden_states=torch.randn(4, 8),
topk_weights=torch.randn(4, 2),
topk_ids=torch.randint(0, 4, (4, 2), dtype=torch.int32),
w1=w1,
w2=w2,
quant_type=QuantType.W8A8,
dynamic_eplb=False,
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,
)
self.assertIsInstance(fused_experts_input.weights, MoEWeights)
self.assertIs(fused_experts_input.weights.w1, w1)
self.assertIs(fused_experts_input.weights.w2, w2)
self.assertIs(fused_experts_input.weights.w1_scale, w1_scale)
self.assertIs(fused_experts_input.weights.w2_scale, w2_scale)
self.assertIs(fused_experts_input.weights.w1_scale_bias, w1_scale_bias)
self.assertIs(fused_experts_input.weights.w2_scale_bias, w2_scale_bias)
self.assertIs(fused_experts_input.weights.w1_offset, w1_offset)
self.assertIs(fused_experts_input.weights.w2_offset, w2_offset)
def test_build_token_dispatch_input_supports_remapped_topk_ids(self):
fused_experts_input = build_fused_experts_input(
hidden_states=torch.randn(2, 4),
topk_weights=torch.randn(2, 1),
topk_ids=torch.tensor([[0], [1]], dtype=torch.int32),
w1=torch.randn(1, 4, 8),
w2=torch.randn(1, 8, 4),
quant_type=QuantType.NONE,
dynamic_eplb=False,
)
routed_topk_ids = torch.tensor([[3], [2]], dtype=torch.int32)
token_dispatch_input = build_token_dispatch_input(
fused_experts_input=fused_experts_input,
topk_ids=routed_topk_ids,
)
self.assertIs(token_dispatch_input.hidden_states, fused_experts_input.hidden_states)
self.assertIs(token_dispatch_input.topk_weights, fused_experts_input.topk_weights)
self.assertIs(token_dispatch_input.routing, fused_experts_input.routing)
self.assertIs(token_dispatch_input.quant, fused_experts_input.quant)
self.assertIs(token_dispatch_input.topk_ids, routed_topk_ids)
def test_build_fused_experts_input_requires_primitive_mxfp_params_for_mxfp_quant(self):
for quant_type in (QuantType.MXFP8, QuantType.MXFP4, QuantType.W4A8MXFP):
with (
self.subTest(quant_type=quant_type),
self.assertRaisesRegex(ValueError, "primitive MXFP params are required"),
):
build_fused_experts_input(
hidden_states=torch.randn(2, 8),
topk_weights=torch.randn(2, 2),
topk_ids=torch.tensor([[0, 1], [1, 0]], dtype=torch.int32),
w1=torch.randn(2, 8, 16),
w2=torch.randn(2, 16, 8),
quant_type=quant_type,
dynamic_eplb=False,
)
def test_build_mlp_compute_input_derives_fusion_and_preserves_mxfp_params(self):
for quant_type, expected_fusion in (
(QuantType.MXFP8, True),
(QuantType.MXFP4, True),
(QuantType.W4A8MXFP, True),
):
with self.subTest(quant_type=quant_type):
mxfp_dtype = _get_test_mxfp_dtype(quant_type)
fused_experts_input = build_fused_experts_input(
hidden_states=torch.randn(2, 8, dtype=torch.bfloat16),
topk_weights=torch.randn(2, 2),
topk_ids=torch.tensor([[0, 1], [1, 0]], dtype=torch.int32),
w1=torch.randn(2, 8, 16),
w2=torch.randn(2, 16, 8),
quant_type=quant_type,
dynamic_eplb=False,
mxfp_act_quant_type=mxfp_dtype,
mxfp_weight_quant_type=mxfp_dtype,
mxfp_scale_dtype=torch.float32,
mxfp_per_token_scale_dtype=torch.float16,
mxfp_use_bf16=False,
w1_scale=[torch.randn(1)],
w2_scale=[torch.randn(1)],
)
token_dispatch_output = MoETokenDispatchOutput(
hidden_states=torch.randn(4, 8, dtype=torch.bfloat16),
group_list=torch.tensor([2, 2], dtype=torch.int64),
group_list_type=1,
dynamic_scale=torch.randn(4, 1),
combine_metadata=MoEAllGatherCombineMetadata(
topk_weights=fused_experts_input.topk_weights,
expanded_row_idx=torch.arange(4, dtype=torch.int32),
restore_shape=torch.Size([2, 8]),
),
)
mlp_compute_input = build_mlp_compute_input(
fused_experts_input=fused_experts_input,
token_dispatch_output=token_dispatch_output,
use_fusion_ops=True,
)
self.assertIs(mlp_compute_input.hidden_states, token_dispatch_output.hidden_states)
self.assertIs(mlp_compute_input.weights, fused_experts_input.weights)
self.assertIs(mlp_compute_input.weights.w1_scale, fused_experts_input.weights.w1_scale)
self.assertIs(mlp_compute_input.weights.w2_scale, fused_experts_input.weights.w2_scale)
self.assertEqual(mlp_compute_input.fusion, expected_fusion)
self.assertTrue(mlp_compute_input.quant.is_mxfp)
assert mlp_compute_input.quant.mxfp is not None
self.assertEqual(mlp_compute_input.quant.mxfp.act_quant_type, mxfp_dtype)
self.assertEqual(mlp_compute_input.quant.mxfp.weight_quant_type, mxfp_dtype)
self.assertEqual(mlp_compute_input.quant.mxfp.scale_dtype, torch.float32)
self.assertEqual(mlp_compute_input.quant.mxfp.per_token_scale_dtype, torch.float16)
self.assertFalse(mlp_compute_input.quant.mxfp.use_bf16)
def test_build_fused_experts_input_constructs_internal_mxfp_leaf_from_primitives(self):
for quant_type in (QuantType.MXFP8, QuantType.MXFP4):
with self.subTest(quant_type=quant_type):
mxfp_dtype = _get_test_mxfp_dtype(quant_type)
fused_experts_input = build_fused_experts_input(
hidden_states=torch.randn(2, 8, dtype=torch.bfloat16),
topk_weights=torch.randn(2, 2),
topk_ids=torch.tensor([[0, 1], [1, 0]], dtype=torch.int32),
w1=torch.randn(2, 8, 16),
w2=torch.randn(2, 16, 8),
quant_type=quant_type,
dynamic_eplb=False,
mxfp_act_quant_type=mxfp_dtype,
mxfp_weight_quant_type=mxfp_dtype,
mxfp_scale_dtype=torch.float32,
mxfp_per_token_scale_dtype=torch.float16,
mxfp_use_bf16=False,
)
self.assertTrue(fused_experts_input.quant.is_mxfp)
assert fused_experts_input.quant.mxfp is not None
self.assertEqual(fused_experts_input.quant.mxfp.act_quant_type, mxfp_dtype)
self.assertEqual(fused_experts_input.quant.mxfp.weight_quant_type, mxfp_dtype)
self.assertEqual(fused_experts_input.quant.mxfp.scale_dtype, torch.float32)
self.assertEqual(fused_experts_input.quant.mxfp.per_token_scale_dtype, torch.float16)
self.assertFalse(fused_experts_input.quant.mxfp.use_bf16)
if __name__ == "__main__":
unittest.main(verbosity=2)

View File

@@ -0,0 +1,217 @@
import unittest
from unittest.mock import MagicMock, patch
import torch
from vllm.model_executor.layers.fused_moe import FusedMoEConfig
from vllm_ascend.ops.fused_moe.prepare_finalize import (
PrepareAndFinalizeWithAll2All,
PrepareAndFinalizeWithAllGather,
PrepareAndFinalizeWithMC2,
)
class TestPrepareAndFinalize(unittest.TestCase):
def setUp(self):
# Mock FusedMoEConfig
fake_stream = MagicMock()
patcher = patch("torch.npu.Stream", return_value=fake_stream)
patcher.start()
self.addCleanup(patcher.stop)
self.mock_get_config = patch("vllm_ascend.ops.fused_moe.prepare_finalize.get_ascend_config")
mock_config = self.mock_get_config.start()
mock_ascend_config = MagicMock()
mock_ascend_config.multistream_overlap_gate = False
mock_ascend_config.enable_context_parallel = False
mock_ascend_config.enable_flashcomm2_parallel_size = 0
mock_config.return_value = mock_ascend_config
self.addCleanup(self.mock_get_config.stop)
self.mock_get_config_utils = patch("vllm_ascend.utils.get_ascend_config")
mock_config_utils = self.mock_get_config_utils.start()
mock_config_utils.return_value = mock_ascend_config
self.addCleanup(self.mock_get_config_utils.stop)
self.moe_config = MagicMock(spec=FusedMoEConfig)
self.moe_config.tp_group = MagicMock()
self.moe_config.tp_group.device_group = MagicMock()
self.moe_config.dp_size = 1
self.moe_config.tp_size = 1
self.moe_config.pcp_size = 1
self.moe_config.ep_size = 1
self.moe_config.dp_group = MagicMock()
self.moe_config.original_num_experts = 8
@patch("vllm_ascend.ops.fused_moe.prepare_finalize.get_tensor_model_parallel_world_size", return_value=1)
@patch("vllm_ascend.ops.fused_moe.prepare_finalize.get_tensor_model_parallel_rank", return_value=0)
@patch("vllm_ascend.ascend_forward_context.get_forward_context")
def test_mc2_prepare_finalize(self, mock_get_forward_context, mock_tp_rank, mock_tp_size):
mock_context = MagicMock()
mock_context.mc2_mask = torch.tensor([1, 0, 1])
mock_context.padded_num_tokens = 4
mock_get_forward_context.return_value = mock_context
layer = PrepareAndFinalizeWithMC2(self.moe_config)
hidden_states = torch.randn(3, 8)
router_logits = torch.randn(3, 2)
prepare_output = layer.prepare(hidden_states, router_logits)
h_out = prepare_output.hidden_states
r_out = prepare_output.router_logits
mask = prepare_output.mc2_mask
padded_hidden_states_shape = prepare_output.padded_hidden_states_shape
# Check padding and split
self.assertEqual(h_out.shape[0], 4)
self.assertEqual(r_out.shape[0], 4)
self.assertEqual(mask.tolist(), [1, 0, 1])
self.assertEqual(padded_hidden_states_shape, torch.Size([4, 8]))
# Finalize
result = layer.finalize(h_out, reduce_results=False, padded_hidden_states_shape=padded_hidden_states_shape)
self.assertEqual(result.shape[0], 3)
@patch("vllm_ascend.ops.fused_moe.prepare_finalize.get_tensor_model_parallel_world_size", return_value=2)
@patch("vllm_ascend.ops.fused_moe.prepare_finalize.get_tensor_model_parallel_rank", return_value=0)
@patch("vllm_ascend.ascend_forward_context.get_forward_context")
@patch("torch.distributed.all_gather")
def test_mc2_tp_split_allgather(self, mock_all_gather, mock_get_forward_context, mock_tp_rank, mock_tp_size):
mock_context = MagicMock()
mock_context.mc2_mask = torch.tensor([1, 0, 1, 0])
mock_context.padded_num_tokens = 4
mock_get_forward_context.return_value = mock_context
layer = PrepareAndFinalizeWithMC2(self.moe_config)
hidden_states = torch.randn(4, 8)
router_logits = torch.randn(4, 2)
prepare_output = layer.prepare(
hidden_states, router_logits, enable_shared_expert_dp=False, replace_allreduce=False
)
h_out = prepare_output.hidden_states
padded_hidden_states_shape = prepare_output.padded_hidden_states_shape
# With TP=2, should split into 2 parts
self.assertEqual(h_out.shape[0], 2)
self.assertEqual(padded_hidden_states_shape, torch.Size([4, 8]))
# Mock all_gather behavior
def mock_all_gather_func(tensor_list, tensor, group=None):
tensor_list[0] = tensor
tensor_list[1] = tensor.clone()
mock_all_gather.side_effect = mock_all_gather_func
layer.split_hidden_states = [torch.zeros_like(h_out), torch.zeros_like(h_out)]
final_result = layer.finalize(
h_out, reduce_results=False, padded_hidden_states_shape=padded_hidden_states_shape
)
# Should concat back to original size
self.assertEqual(final_result.shape[0], 4)
@patch("vllm_ascend.ops.fused_moe.prepare_finalize.get_tensor_model_parallel_world_size", return_value=1)
@patch("vllm_ascend.ops.fused_moe.prepare_finalize.get_tensor_model_parallel_rank", return_value=0)
def test_all2all_prepare_finalize(self, mock_tp_rank, mock_tp_size):
layer = PrepareAndFinalizeWithAll2All(self.moe_config)
hidden_states = torch.randn(3, 8)
router_logits = torch.randn(3, 2)
prepare_output = layer.prepare(hidden_states, router_logits)
h_out = prepare_output.hidden_states
padded_hidden_states_shape = prepare_output.padded_hidden_states_shape
# Pad to tp_size=1, so no change
self.assertEqual(h_out.shape[0], 3)
self.assertEqual(padded_hidden_states_shape, torch.Size([3, 8]))
result = layer.finalize(h_out, reduce_results=False, padded_hidden_states_shape=padded_hidden_states_shape)
self.assertEqual(result.shape[0], 3)
@patch("vllm_ascend.ops.fused_moe.prepare_finalize.get_tensor_model_parallel_world_size", return_value=2)
@patch("vllm_ascend.ops.fused_moe.prepare_finalize.get_tensor_model_parallel_rank", return_value=0)
@patch("torch.distributed.all_gather")
def test_all2all_tp_split_allgather(self, mock_all_gather, mock_tp_rank, mock_tp_size):
layer = PrepareAndFinalizeWithAll2All(self.moe_config)
hidden_states = torch.randn(2, 8)
router_logits = torch.randn(2, 2)
prepare_output = layer.prepare(
hidden_states, router_logits, enable_shared_expert_dp=False, replace_allreduce=False
)
h_out = prepare_output.hidden_states
padded_hidden_states_shape = prepare_output.padded_hidden_states_shape
# Split due to TP=2
self.assertEqual(h_out.shape[0], 1)
self.assertEqual(padded_hidden_states_shape, torch.Size([2, 8]))
# Mock all_gather
def mock_all_gather_func(tensor_list, tensor, group=None):
tensor_list[0] = tensor
tensor_list[1] = tensor.clone()
mock_all_gather.side_effect = mock_all_gather_func
layer.split_hidden_states = [torch.zeros_like(h_out), torch.zeros_like(h_out)]
final_result = layer.finalize(
h_out, reduce_results=False, padded_hidden_states_shape=padded_hidden_states_shape
)
# Should concat back
self.assertEqual(final_result.shape[0], 2)
@patch("vllm_ascend.ops.fused_moe.prepare_finalize.get_dp_group")
@patch("vllm_ascend.ascend_forward_context.get_forward_context")
@patch("vllm_ascend.ops.fused_moe.prepare_finalize.enable_sp", return_value=False)
@patch("vllm_ascend.ops.fused_moe.prepare_finalize.enable_sp_by_pass", return_value=False)
def test_allgather_prepare_finalize(
self, mock_enable_sp_by_pass, mock_enable_sp, mock_get_forward_context, mock_get_dp_group
):
# Mock forward context
mock_context = MagicMock()
mock_context.max_tokens_across_dp = 6
mock_get_forward_context.return_value = mock_context
# Create a proper mock for DP group with working all_gather
mock_dp_group = MagicMock()
def mock_all_gather_func(tensor, dim):
# Simulate DP=2: repeat the tensor along the specified dimension
return torch.cat([tensor, tensor], dim=dim)
mock_dp_group.all_gather = mock_all_gather_func
mock_get_dp_group.return_value = mock_dp_group
self.moe_config.dp_size = 2
self.moe_config.tp_size = 1
self.moe_config.pcp_size = 1
self.moe_config.ep_size = 1
self.moe_config.dp_group = mock_dp_group
layer = PrepareAndFinalizeWithAllGather(self.moe_config)
hidden_states = torch.randn(3, 8)
router_logits = torch.randn(3, 2)
prepare_output = layer.prepare(hidden_states, router_logits)
h_out = prepare_output.hidden_states
r_out = prepare_output.router_logits
padded_hidden_states_shape = prepare_output.padded_hidden_states_shape
# After all-gather with DP=2, should double the batch size
self.assertEqual(h_out.shape[0], 12)
self.assertEqual(r_out.shape[0], 12)
self.assertIsNone(padded_hidden_states_shape)
# Finalize with reduce_scatter
def mock_reduce_scatter_func(tensor, dim):
# Simulate reduce_scatter: take first half
return tensor[:3]
mock_dp_group.reduce_scatter = mock_reduce_scatter_func
result = layer.finalize(h_out, reduce_results=False, padded_hidden_states_shape=padded_hidden_states_shape)
self.assertEqual(result.shape[0], 3)
result_with_tp = layer.finalize(h_out, reduce_results=True)
self.assertEqual(result_with_tp.shape[0], 3)

View File

@@ -1,378 +1,347 @@
import math
import unittest
from unittest.mock import MagicMock, PropertyMock, patch
#
# 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 inspect
from unittest.mock import MagicMock, patch
import pytest
import torch
from transformers.configuration_utils import PretrainedConfig
from vllm.config import ModelConfig, VllmConfig
from vllm.model_executor.layers.rotary_embedding import (
DeepseekScalingRotaryEmbedding, RotaryEmbedding)
from vllm.model_executor.layers.rotary_embedding import RotaryEmbedding, YaRNScalingRotaryEmbedding
from tests.ut.base import TestBase
from vllm_ascend.ascend_forward_context import set_ascend_forward_context
from vllm_ascend.ops.rotary_embedding import _custom_rotary_embedding_enabled
from vllm_ascend.ops.rotary_embedding import AscendRotaryEmbedding, AscendYaRNRotaryEmbedding
MODEL = "Qwen3-0.6B"
MAX_NUM_BATCHED_TOKEND = 10000
HEAD_SIZE = 64
ROTARY_DIM = 64
MAX_POS = 2048
BASE = 10000.0
DTYPE = torch.bfloat16
SEQ_LEN = 4
NUM_HEADS = 2
class TestCustomRotaryEmbeddingEnabled(unittest.TestCase):
def setUp(self):
# Common setup for tests
self.positions = torch.tensor([1, 2, 3])
self.query = torch.randn(3, 4, dtype=torch.float16)
self.key = torch.randn(3, 4, dtype=torch.float16)
self.head_size = 32
self.cos_sin_cache = torch.randn(3, 4)
# Mock self object for rope_forward_oot
self.mock_self = MagicMock()
self.mock_self.head_size = self.head_size
self.mock_self.cos_sin_cache = self.cos_sin_cache
self.mock_self.is_neox_style = True
self.mock_self.forward_native.return_value = (self.query, self.key)
def test_custom_rotary_embedding_enabled(self):
# Test when all conditions are True
with patch('vllm_ascend.ops.rotary_embedding.enable_custom_op',
return_value=True):
result = _custom_rotary_embedding_enabled(self.query, True,
self.head_size)
self.assertTrue(result)
# Test when dtype is not float16
with patch('vllm_ascend.ops.rotary_embedding.enable_custom_op',
return_value=True):
query = self.query.to(torch.float32)
result = _custom_rotary_embedding_enabled(query, True,
self.head_size)
self.assertFalse(result)
# Test when neox_style is False
with patch('vllm_ascend.ops.rotary_embedding.enable_custom_op',
return_value=True):
result = _custom_rotary_embedding_enabled(self.query, False,
self.head_size)
self.assertFalse(result)
# Test when head_size is not divisible by 32
with patch('vllm_ascend.ops.rotary_embedding.enable_custom_op',
return_value=True):
result = _custom_rotary_embedding_enabled(self.query, True,
self.head_size + 1)
self.assertFalse(result)
# Test when custom op is disabled
with patch('vllm_ascend.ops.rotary_embedding.enable_custom_op',
return_value=False):
result = _custom_rotary_embedding_enabled(self.query, True,
self.head_size)
self.assertFalse(result)
def _make_tensors(seq_len=SEQ_LEN, num_heads=NUM_HEADS, head_size=HEAD_SIZE):
positions = torch.arange(seq_len, dtype=torch.long)
query = torch.randn(seq_len, num_heads * head_size)
key = torch.randn(seq_len, num_heads * head_size)
return positions, query, key
class TestAscendRotaryEmbedding(unittest.TestCase):
def check_parent_init_signature_has_not_changed(parent_func, child_func):
parent_sig = inspect.signature(parent_func)
parent_params = set(parent_sig.parameters) - {"self"}
def setUp(self):
# Common setup for tests
self.positions = torch.tensor([1, 2, 3])
self.query = torch.randn(3, 1, 32, dtype=torch.float16)
self.key = torch.randn(3, 1, 32, dtype=torch.float16)
self.head_size = 32
self.rotary_dim = self.head_size
self.max_position = 16
self.rope_theta = 10000
self.is_neox_style = True
self.cos_sin_cache = torch.randn(3, 1, 32)
self.layer = RotaryEmbedding(self.head_size, self.rotary_dim,
self.max_position, self.rope_theta,
self.is_neox_style, torch.float16)
child_sig = inspect.signature(child_func)
child_params = set(child_sig.parameters) - {"self"}
# Mock self object for rope_forward_oot
self.mock_self = MagicMock()
self.mock_self.head_size = self.head_size
self.mock_self.cos_sin_cache = self.cos_sin_cache
self.mock_self.is_neox_style = self.is_neox_style
added = parent_params - child_params
removed = child_params - parent_params
@patch('torch.ops._C_ascend')
@patch('vllm_ascend.ops.rotary_embedding.is_310p', return_value=False)
@patch('vllm_ascend.ops.rotary_embedding._custom_rotary_embedding_enabled',
return_value=True)
@patch('torch.ops._npu_rotary_embedding')
@patch('vllm.config.ModelConfig.__post_init__', MagicMock())
@patch('vllm.config.VllmConfig.__post_init__', MagicMock())
@patch('vllm.distributed.parallel_state._DP', MagicMock(world_size=1))
@patch('vllm.distributed.parallel_state._TP', MagicMock(world_size=1))
def test_rope_forward_oot_custom_kernel(self, mock_rotary_embedding,
mock_custom_enabled, mock_is_310p,
mock__c):
mock_config = MagicMock()
mock_config.torchair_graph_config.enabled = False
# Setup mock for custom kernel path
mock__c.rotary_embedding.return_value = self.query, self.key
vllm_config = VllmConfig()
model_config = ModelConfig(MODEL,
tokenizer=MODEL,
max_model_len=MAX_NUM_BATCHED_TOKEND)
model_config.hf_config = PretrainedConfig()
vllm_config.model_config = model_config
with set_ascend_forward_context(None, vllm_config):
result_q, result_k = self.layer.forward(self.positions, self.query,
self.key)
mock__c.rotary_embedding.assert_called_once()
self.assertEqual(result_q.shape, self.query.shape)
self.assertEqual(result_k.shape, self.key.shape)
@patch('vllm_ascend.ops.rotary_embedding._custom_rotary_embedding_enabled',
return_value=False)
@patch('torch_npu._npu_rotary_embedding')
@patch('vllm.config.ModelConfig.__post_init__', MagicMock())
@patch('vllm.config.VllmConfig.__post_init__', MagicMock())
@patch('vllm.distributed.parallel_state._DP', MagicMock(world_size=1))
@patch('vllm.distributed.parallel_state._TP', MagicMock(world_size=1))
def test_rope_forward_oot_contiguous(self, mock_npu_rotary,
mock_custom_enabled):
mock_config = MagicMock()
mock_config.torchair_graph_config.enabled = False
# Test contiguous path when custom is disabled
non_contig_query = self.query.transpose(0, 1)
non_contig_key = self.key.transpose(0, 1)
vllm_config = VllmConfig()
model_config = ModelConfig(MODEL,
tokenizer=MODEL,
max_model_len=MAX_NUM_BATCHED_TOKEND)
model_config.hf_config = PretrainedConfig()
vllm_config.model_config = model_config
with set_ascend_forward_context(None, vllm_config):
result_q, result_k = self.layer.forward(self.positions,
non_contig_query,
non_contig_key)
mock_npu_rotary.assert_called_once()
self.assertEqual(result_q.shape, non_contig_query.shape)
self.assertEqual(result_k.shape, non_contig_key.shape)
@patch('vllm.config.ModelConfig.__post_init__', MagicMock())
@patch('vllm.config.VllmConfig.__post_init__', MagicMock())
@patch('vllm.distributed.parallel_state._DP', MagicMock(world_size=1))
@patch('vllm.distributed.parallel_state._TP', MagicMock(world_size=1))
def test_rope_forward_oot_with_offsets(self):
mock_config = MagicMock()
mock_config.torchair_graph_config.enabled = False
# Test that NotImplementedError is raised when offsets is provided
offsets = torch.tensor([1, 2, 3])
with self.assertRaises(NotImplementedError):
vllm_config = VllmConfig()
model_config = ModelConfig(MODEL,
tokenizer=MODEL,
max_model_len=MAX_NUM_BATCHED_TOKEND)
model_config.hf_config = PretrainedConfig()
vllm_config.model_config = model_config
with set_ascend_forward_context(None, vllm_config):
self.layer.forward(self.positions, self.query, self.key,
offsets)
@patch('vllm_ascend.ops.rotary_embedding._custom_rotary_embedding_enabled',
return_value=False)
@patch('torch_npu._npu_rotary_embedding')
@patch('vllm.config.ModelConfig.__post_init__', MagicMock())
@patch('vllm.config.VllmConfig.__post_init__', MagicMock())
@patch('vllm.distributed.parallel_state._DP', MagicMock(world_size=1))
@patch('vllm.distributed.parallel_state._TP', MagicMock(world_size=1))
def test_rope_forward_oot_neox_style_override(self, mock_npu_rotary,
mock_custom_enabled):
mock_config = MagicMock()
mock_config.torchair_graph_config.enabled = False
# Test neox_style override
vllm_config = VllmConfig()
model_config = ModelConfig(MODEL,
tokenizer=MODEL,
max_model_len=MAX_NUM_BATCHED_TOKEND)
model_config.hf_config = PretrainedConfig()
vllm_config.model_config = model_config
with set_ascend_forward_context(None, vllm_config):
result_q, result_k = self.layer.forward(
self.positions,
self.query,
self.key,
is_neox_style_override=False)
# Check that neox_style=False was passed to the NPU function
args, kwargs = mock_npu_rotary.call_args
self.assertFalse(args[-1])
@patch('vllm_ascend.ops.rotary_embedding._custom_rotary_embedding_enabled',
return_value=False)
@patch('torch_npu._npu_rotary_embedding')
@patch('vllm.config.ModelConfig.__post_init__', MagicMock())
@patch('vllm.config.VllmConfig.__post_init__', MagicMock())
@patch('vllm.distributed.parallel_state._DP', MagicMock(world_size=1))
@patch('vllm.distributed.parallel_state._TP', MagicMock(world_size=1))
def test_rope_forward_oot_rotary_dim_less_than_head_size(
self, mock_npu_rotary, mock_custom_enabled):
mock_config = MagicMock()
mock_config.torchair_graph_config.enabled = False
# test case when rotary_dim < head_size
org_rotary_dim = self.layer.rotary_dim
self.layer.rotary_dim = self.layer.head_size // 2
vllm_config = VllmConfig()
model_config = ModelConfig(MODEL,
tokenizer=MODEL,
max_model_len=MAX_NUM_BATCHED_TOKEND)
model_config.hf_config = PretrainedConfig()
vllm_config.model_config = model_config
with set_ascend_forward_context(None, vllm_config):
result_q, result_k = self.layer.forward(self.positions, self.query,
self.key)
mock_npu_rotary.assert_called_once()
self.assertEqual(result_q.shape, self.query.shape)
self.assertEqual(result_k.shape, self.key.shape)
# restore rotary_dim
self.layer.rotary_dim = org_rotary_dim
assert not added, (
f"{parent_func.__name__} added new parameter(s): {added}. "
f"Check whether {child_func.__name__} needs to forward them."
)
assert not removed, (
f"{parent_func.__name__} removed parameter(s): {removed}. "
f"Check whether {child_func.__name__} needs to forward them."
)
class MockRopeModule:
def __init__(self, max_seq_len=2048, is_neox_style=True):
self.max_seq_len = max_seq_len
self.is_neox_style = is_neox_style
self.cos_cached = None
self.sin_cached = None
self.rotary_dim = 1
self.base = 1
@pytest.fixture(autouse=True)
def patch_init_side_effects():
"""
Suppress all side-effects that fire during __init__ so every test starts
from a clean, predictable state without needing real NPU ops or vLLM
global config.
"""
with (
patch("vllm_ascend.ops.rotary_embedding._record_cos_sin_cache"),
patch("vllm_ascend.ops.rotary_embedding._record_cos_and_sin_cache_interleaved"),
patch("vllm_ascend.ops.rotary_embedding.get_current_vllm_config") as mock_cfg,
):
# Default: speculative_config is None → use_mtp = False
mock_cfg.return_value.speculative_config = None
yield mock_cfg
class TestAscendDeepseekScalingRotaryEmbedding(TestBase):
@pytest.fixture()
def make_embedding(patch_init_side_effects):
"""Factory that creates an AscendRotaryEmbedding with controllable use_mtp."""
def setUp(self):
# Common setup for tests
self.positions = torch.tensor([1, 2, 3])
self.query = torch.randn(3, 1, 32, dtype=torch.float16)
self.key = torch.randn(3, 1, 32, dtype=torch.float16)
self.head_size = 32
self.rotary_dim = self.head_size
self.max_position = 16
self.rope_theta = 10000
self.is_neox_style = True
self.scaling_factor = 1
self.layer = None
def _factory(use_mtp: bool = False, is_neox_style: bool = True):
spec_cfg = MagicMock(method="mtp") if use_mtp else None
patch_init_side_effects.return_value.speculative_config = spec_cfg
def _create_layer(self):
self.layer = DeepseekScalingRotaryEmbedding(
self.head_size, self.rotary_dim, self.max_position,
self.rope_theta, self.is_neox_style, self.scaling_factor,
torch.float16)
return self.layer
with patch("vllm_ascend.ops.rotary_embedding.RotaryEmbedding.__init__") as mock_parent_init:
mock_parent_init.return_value = None
from vllm_ascend.ops.rotary_embedding import AscendRotaryEmbedding
@patch("vllm.platforms.current_platform.device_type",
new=torch.device("cpu"))
@patch("vllm_ascend.ops.rotary_embedding.NPUPlatform",
new_callable=PropertyMock)
def test_native_rope_deepseek_forward_base(self, mock_npuplatform):
mock_npuplatform.device_type = torch.device("cpu")
self.layer = self._create_layer()
with patch("vllm_ascend.ops.rotary_embedding._rope_forward_oot",
return_value=(self.query,
self.key)) as mock_rope_forward_oot:
q_pe, k_pe = self.layer.forward(self.positions, self.query,
self.key)
mock_rope_forward_oot.assert_called_once()
assert q_pe.shape == self.query.shape
assert k_pe.shape == self.key.shape
emb = AscendRotaryEmbedding.__new__(AscendRotaryEmbedding)
# Manually set attrs that the real parent would set
emb.head_size = HEAD_SIZE
emb.rotary_dim = ROTARY_DIM
emb.is_neox_style = is_neox_style
emb.cos_sin_cache = torch.zeros(MAX_POS, ROTARY_DIM)
# Call __init__ to exercise our code path
AscendRotaryEmbedding.__init__(emb, HEAD_SIZE, ROTARY_DIM, MAX_POS, BASE, is_neox_style, DTYPE)
return emb
@patch('vllm_ascend.ops.rotary_embedding._rope_forward_oot')
@patch("vllm.platforms.current_platform.device_type",
new=torch.device("cpu"))
@patch("vllm_ascend.ops.rotary_embedding.NPUPlatform",
new_callable=PropertyMock)
def test_native_rope_deepseek_forward_key_reshaping(
self, mock_npuplatform, mock_rope_forward_oot):
mock_npuplatform.device_type = torch.device("cpu")
self.layer = self._create_layer()
return _factory
key = torch.randn(1, 32)
mock_rope_forward_oot.return_value = (self.query, key)
@pytest.fixture()
def make_yarn_embedding(patch_init_side_effects):
"""
Factory for AscendYaRNRotaryEmbedding with parent __init__ suppressed.
patch_init_side_effects is the same autouse fixture as before.
"""
q_pe, k_pe = self.layer.forward(self.positions, self.query, key)
mock_rope_forward_oot.assert_called_once()
assert q_pe.shape == self.query.shape
assert k_pe.shape == key.shape
def _factory(is_neox_style: bool = True):
with patch("vllm_ascend.ops.rotary_embedding.YaRNScalingRotaryEmbedding.__init__") as mock_parent_init:
mock_parent_init.return_value = None
from vllm_ascend.ops.rotary_embedding import AscendYaRNRotaryEmbedding
@patch('vllm_ascend.ops.rotary_embedding._rope_forward_oot')
@patch("vllm.platforms.current_platform.device_type",
new=torch.device("cpu"))
@patch("vllm_ascend.ops.rotary_embedding.NPUPlatform",
new_callable=PropertyMock)
def test_native_rope_deepseek_forward_non_neox_style(
self, mock_npuplatform, mock_rope_forward_oot):
mock_npuplatform.device_type = torch.device("cpu")
self.layer = self._create_layer()
emb = AscendYaRNRotaryEmbedding.__new__(AscendYaRNRotaryEmbedding)
emb.head_size = HEAD_SIZE
emb.rotary_dim = ROTARY_DIM
emb.is_neox_style = is_neox_style
emb.cos_sin_cache = torch.zeros(MAX_POS, ROTARY_DIM)
AscendYaRNRotaryEmbedding.__init__(
emb,
head_size=HEAD_SIZE,
rotary_dim=ROTARY_DIM,
max_position_embeddings=MAX_POS,
base=BASE,
is_neox_style=is_neox_style,
scaling_factor=1.0,
dtype=DTYPE,
)
return emb
mock_rope_forward_oot.return_value = (self.query, self.key)
return _factory
q_pe, k_pe = self.layer.forward(self.positions, self.query, self.key)
mock_rope_forward_oot.assert_called_once()
assert q_pe.shape == self.query.shape
assert k_pe.shape == self.key.shape
class TestAscendEmbeddingForwardOOT:
@patch("torch.ops.vllm.npu_rotary_embedding")
@patch("vllm_ascend.ascend_forward_context.get_forward_context")
def test_basic_call_delegates_to_npu_op(self, mock_get_forward_context, mock_npu_op, make_embedding):
"""forward_oot always calls npu_rotary_embedding and returns its result."""
mock_get_forward_context.return_value = MagicMock()
mock_get_forward_context.return_value.is_draft_model = False
mock_get_forward_context.return_value.flash_comm_v1_enabled = False
expected_output = (torch.randn(SEQ_LEN, NUM_HEADS * HEAD_SIZE),) * 2
mock_npu_op.return_value = expected_output
@patch("vllm.platforms.current_platform.device_type",
new=torch.device("cpu"))
@patch("vllm_ascend.ops.rotary_embedding.NPUPlatform",
new_callable=PropertyMock)
def test_basic_case(self, mock_npuplatform):
# Test with standard values
mock_npuplatform.device_type = torch.device("cpu")
self.layer = self._create_layer()
num_rotations = 100
dim = 512
base = 10000
max_position_embeddings = 2048
emb = make_embedding()
positions, query, key = _make_tensors()
result = self.layer._yarn_find_correction_dim(num_rotations, dim, base,
max_position_embeddings)
result = emb.forward_oot(positions, query, key)
# Calculate expected value manually
expected = (dim * torch.log(
torch.tensor(max_position_embeddings) /
(num_rotations * 2 * torch.pi))) / (2 *
torch.log(torch.tensor(base)))
mock_npu_op.assert_called_once_with(
positions,
query,
key,
emb.cos_sin_cache,
HEAD_SIZE,
ROTARY_DIM,
emb.is_neox_style,
)
assert result is expected_output
self.assertTrue(torch.allclose(result, expected))
@patch("torch.ops.vllm.npu_rotary_embedding")
@patch("vllm_ascend.ascend_forward_context.get_forward_context")
def test_neox_style_override_true(self, mock_get_forward_context, mock_npu_op, make_embedding):
"""is_neox_style_override=True wins over self.is_neox_style=False."""
mock_get_forward_context.return_value = MagicMock()
mock_get_forward_context.return_value.is_draft_model = False
mock_get_forward_context.return_value.flash_comm_v1_enabled = False
mock_npu_op.return_value = MagicMock()
@patch("vllm.platforms.current_platform.device_type",
new=torch.device("cpu"))
@patch("vllm_ascend.ops.rotary_embedding.NPUPlatform",
new_callable=PropertyMock)
def test_yarn_get_mscale(self, mock_npuplatform):
mock_npuplatform.device_type = torch.device("cpu")
self.layer = self._create_layer()
emb = make_embedding(is_neox_style=False)
positions, query, key = _make_tensors()
# test_scale_less_than_or_equal_1
self.assertEqual(self.layer._yarn_get_mscale(scale=0.5), 1.0)
self.assertEqual(self.layer._yarn_get_mscale(scale=1.0), 1.0)
self.assertEqual(self.layer._yarn_get_mscale(scale=0.999), 1.0)
emb.forward_oot(positions, query, key, is_neox_style_override=True)
# test_scale_greater_than_1:
test_cases = [(2.0, 1.0, 1.0 + 0.1 * math.log(2.0)),
(10.0, 1.0, 1.0 + 0.1 * math.log(10.0)),
(5.0, 2.0, 1.0 + 0.2 * math.log(5.0)),
(math.e, 1.0, 1.0 + 0.1)]
_, kwargs = mock_npu_op.call_args
# Verify the override was forwarded correctly
assert mock_npu_op.call_args[0][-1] is True # last positional arg = is_neox_style
for scale, mscale, expected in test_cases:
result = self.layer._yarn_get_mscale(scale, mscale)
self.assertAlmostEqual(
result,
expected,
places=6,
msg=f"Failed for scale={scale}, mscale={mscale}")
@patch("torch.ops.vllm.npu_rotary_embedding")
@patch("vllm_ascend.ascend_forward_context.get_forward_context")
def test_neox_style_override_false(self, mock_get_forward_context, mock_npu_op, make_embedding):
"""is_neox_style_override=False wins over self.is_neox_style=True."""
mock_get_forward_context.return_value = MagicMock()
mock_get_forward_context.return_value.is_draft_model = False
mock_get_forward_context.return_value.flash_comm_v1_enabled = False
mock_npu_op.return_value = MagicMock()
emb = make_embedding(is_neox_style=True)
positions, query, key = _make_tensors()
emb.forward_oot(positions, query, key, is_neox_style_override=False)
assert mock_npu_op.call_args[0][-1] is False
@patch("torch.ops.vllm.npu_rotary_embedding")
@patch("vllm_ascend.ascend_forward_context.get_forward_context")
def test_neox_style_override_none_uses_self(self, mock_get_forward_context, mock_npu_op, make_embedding):
"""When override is None, self.is_neox_style is used unchanged."""
mock_get_forward_context.return_value = MagicMock()
mock_get_forward_context.return_value.is_draft_model = False
mock_get_forward_context.return_value.flash_comm_v1_enabled = False
mock_npu_op.return_value = MagicMock()
emb = make_embedding(is_neox_style=True)
positions, query, key = _make_tensors()
emb.forward_oot(positions, query, key, is_neox_style_override=None)
assert mock_npu_op.call_args[0][-1] is True
@patch("torch.ops.vllm.maybe_all_gather_and_maybe_unpad")
@patch("torch.ops.vllm.npu_rotary_embedding")
@patch("vllm_ascend.ascend_forward_context.get_forward_context")
def test_gather_unpad_called_when_all_conditions_met(
self, mock_get_forward_context, mock_npu_op, mock_gather, make_embedding
):
"""
maybe_all_gather_and_maybe_unpad is called iff:
is_draft_model=True AND use_mtp=True AND flash_comm_v1_enabled=True
"""
mock_get_forward_context.return_value = MagicMock()
mock_get_forward_context.return_value.is_draft_model = True
mock_get_forward_context.return_value.flash_comm_v1_enabled = True
gathered_positions = torch.arange(SEQ_LEN, dtype=torch.long)
mock_gather.return_value = gathered_positions
mock_npu_op.return_value = MagicMock()
emb = make_embedding(use_mtp=True)
positions, query, key = _make_tensors()
emb.forward_oot(positions, query, key)
mock_gather.assert_called_once()
# npu op should receive the gathered positions, not the originals
assert mock_npu_op.call_args[0][0] is gathered_positions
@pytest.mark.parametrize(
"is_draft_model,flash_comm,use_mtp",
[
(False, True, True), # not draft
(True, False, True), # flash_comm disabled
(True, True, False), # use_mtp disabled
],
)
@patch("torch.ops.vllm.maybe_all_gather_and_maybe_unpad")
@patch("torch.ops.vllm.npu_rotary_embedding")
@patch("vllm_ascend.ascend_forward_context.get_forward_context")
def test_gather_unpad_skipped_unless_all_conditions_met(
self,
mock_get_forward_context,
mock_npu_op,
mock_gather,
is_draft_model,
flash_comm,
use_mtp,
make_embedding,
):
"""gather/unpad must NOT fire if any one of the three conditions is False."""
mock_get_forward_context.return_value = MagicMock()
mock_get_forward_context.return_value.is_draft_model = is_draft_model
mock_get_forward_context.return_value.flash_comm_v1_enabled = flash_comm
mock_npu_op.return_value = MagicMock()
emb = make_embedding(use_mtp=use_mtp)
positions, query, key = _make_tensors()
emb.forward_oot(positions, query, key)
mock_gather.assert_not_called()
# Original positions tensor is passed through untouched
assert mock_npu_op.call_args[0][0] is positions
def test_parent_init_signature_has_not_changed(self):
"""
Fail loudly if RotaryEmbedding.__init__ adds, removes, or
renames parameters, so a developer knows to update AscendRotaryEmbedding
accordingly.
"""
check_parent_init_signature_has_not_changed(RotaryEmbedding.__init__, AscendRotaryEmbedding.__init__)
class TestAscendYaRNRotaryEmbeddingForwardOOT:
@patch("vllm_ascend.ops.rotary_embedding.AscendRotaryEmbedding.forward_oot")
def test_delegates_to_ascend_rotary_forward_oot(self, mock_delegate, make_yarn_embedding):
"""forward_oot must delegate to AscendRotaryEmbedding.forward_oot."""
expected = MagicMock()
mock_delegate.return_value = expected
emb = make_yarn_embedding()
positions, query, key = _make_tensors()
result = emb.forward_oot(positions, query, key)
mock_delegate.assert_called_once_with(emb, positions, query, key, None, None)
assert result is expected
@patch("vllm_ascend.ops.rotary_embedding.AscendRotaryEmbedding.forward_oot")
def test_return_value_passed_through(self, mock_delegate, make_yarn_embedding):
"""Return value from the delegate is returned unchanged."""
sentinel = (torch.randn(SEQ_LEN, HEAD_SIZE), torch.randn(SEQ_LEN, HEAD_SIZE))
mock_delegate.return_value = sentinel
emb = make_yarn_embedding()
positions, query, key = _make_tensors()
result = emb.forward_oot(positions, query, key)
assert result is sentinel
@pytest.mark.parametrize("override", [True, False])
@patch("vllm_ascend.ops.rotary_embedding.AscendRotaryEmbedding.forward_oot")
def test_is_neox_style_override_forwarded(self, mock_delegate, override, make_yarn_embedding):
"""is_neox_style_override must be forwarded verbatim, both True and False."""
mock_delegate.return_value = MagicMock()
emb = make_yarn_embedding()
positions, query, key = _make_tensors()
emb.forward_oot(positions, query, key, is_neox_style_override=override)
_, call_args, _ = mock_delegate.mock_calls[0]
assert call_args[5] is override # 6th positional arg
@patch("vllm_ascend.ops.rotary_embedding.AscendRotaryEmbedding.forward_oot")
def test_all_args_forwarded_together(self, mock_delegate, make_yarn_embedding):
"""Smoke test: all args passed simultaneously are all forwarded correctly."""
mock_delegate.return_value = MagicMock()
emb = make_yarn_embedding()
positions, query, key = _make_tensors()
offsets = torch.ones(SEQ_LEN, dtype=torch.long)
emb.forward_oot(positions, query, key, offsets=offsets, is_neox_style_override=False)
mock_delegate.assert_called_once_with(emb, positions, query, key, offsets, False)
def test_parent_init_signature_has_not_changed(self):
"""
Fail loudly if YaRNScalingRotaryEmbedding.__init__ adds, removes, or
renames parameters, so a developer knows to update AscendYaRNRotaryEmbedding
accordingly.
"""
check_parent_init_signature_has_not_changed(
YaRNScalingRotaryEmbedding.__init__, AscendYaRNRotaryEmbedding.__init__
)

View File

@@ -14,39 +14,68 @@
# Adapted from vllm/tests/lora/test_layers.py
import unittest
from unittest import mock
from unittest.mock import MagicMock, patch
import torch
from vllm_ascend.ascend_config import init_ascend_config
from vllm_ascend.distributed import parallel_state
from vllm_ascend.ops.vocab_parallel_embedding import (
AscendLogitsProcessor, AscendParallelLMHead, AscendVocabParallelEmbedding)
AscendLogitsProcessor,
AscendParallelLMHead,
AscendVocabParallelEmbedding,
)
VOCAB_PARALLEL_EMBEDDING_TEST_NUM_RANDOM_SEEDS = 128
class TestCustomVocabParallelEmbedding(unittest.TestCase):
def setUp(self):
self.num_embeddings = 50
self.embedding_dim = 10
self.org_num_embeddings = 40
self.padding_size = 8
self.mock_group = mock.MagicMock()
self.mock_group.world_size = 2
self.mock_group.rank_in_group = 0
parallel_state._MLP_TP = self.mock_group
parallel_state._OTP = self.mock_group
mock_vllm_config = MagicMock()
mock_vllm_config.additional_config = {}
init_ascend_config(mock_vllm_config)
self.mock_ascend_config = MagicMock()
self.mock_ascend_config.finegrained_tp_config.lmhead_tensor_parallel_size = 2
self.mock_ascend_config.finegrained_tp_config.embedding_tensor_parallel_size = 2
self.patches = [
patch("vllm_ascend.utils.get_ascend_config", return_value=self.mock_ascend_config),
patch("vllm_ascend.distributed.parallel_state.get_lmhead_tp_group", return_value=self.mock_group),
patch(
"vllm.distributed.parallel_state.get_tp_group",
return_value=self.mock_group,
),
]
for p in self.patches:
p.start()
def _create_layer(self):
# Patch methods and dependencies for VocabParallelEmbedding
mock_group = MagicMock()
mock_group.world_size = 2
mock_group.rank_in_group = 0
with patch("vllm_ascend.ops.vocab_parallel_embedding.get_tp_group", return_value=mock_group), \
patch("vllm.model_executor.layers.vocab_parallel_embedding.get_tensor_model_parallel_rank", return_value=0), \
patch("vllm.model_executor.layers.vocab_parallel_embedding.get_tensor_model_parallel_world_size", return_value=2), \
patch("vllm.model_executor.layers.vocab_parallel_embedding.pad_vocab_size", side_effect=lambda x, y: x + y), \
patch("vllm.model_executor.layers.vocab_parallel_embedding.divide", side_effect=lambda x, y: x // y):
with (
patch("vllm_ascend.ops.vocab_parallel_embedding.get_tp_group", return_value=mock_group),
patch("vllm.model_executor.layers.vocab_parallel_embedding.get_tensor_model_parallel_rank", return_value=0),
patch(
"vllm.model_executor.layers.vocab_parallel_embedding.get_tensor_model_parallel_world_size",
return_value=2,
),
patch("vllm.model_executor.layers.vocab_parallel_embedding.pad_vocab_size", side_effect=lambda x, y: x + y),
patch("vllm.model_executor.layers.vocab_parallel_embedding.divide", side_effect=lambda x, y: x // y),
):
# Create an instance of VocabParallelEmbedding
layer = AscendVocabParallelEmbedding(
num_embeddings=self.num_embeddings,
@@ -54,7 +83,8 @@ class TestCustomVocabParallelEmbedding(unittest.TestCase):
org_num_embeddings=self.org_num_embeddings,
padding_size=self.padding_size,
quant_config=None, # Mock quantization config
prefix="")
prefix="",
)
layer.shard_indices = MagicMock()
layer.shard_indices.org_vocab_start_index = 10
@@ -65,8 +95,8 @@ class TestCustomVocabParallelEmbedding(unittest.TestCase):
# Mock the quantization method
layer.quant_method.embedding = MagicMock(
side_effect=lambda _, x: torch.randn(x.shape[0], self.
embedding_dim))
side_effect=lambda _, x: torch.randn(x.shape[0], self.embedding_dim)
)
return layer
def test_get_masked_input_and_mask(self):
@@ -81,17 +111,16 @@ class TestCustomVocabParallelEmbedding(unittest.TestCase):
org_vocab_end_index=20,
num_org_vocab_padding=5,
added_vocab_start_index=30,
added_vocab_end_index=40)
added_vocab_end_index=40,
)
expected_mask = torch.tensor([True, False, True, False, True])
self.assertTrue(
torch.equal(mask, expected_mask),
f"Mask mismatch. Expected {expected_mask}, got {mask}")
self.assertTrue(torch.equal(mask, expected_mask), f"Mask mismatch. Expected {expected_mask}, got {mask}")
expected_masked = torch.tensor([0, 5, 0, 20, 0])
self.assertTrue(
torch.equal(masked_input, expected_masked),
f"Masked input mismatch. Expected {expected_masked}, got {masked_input}"
f"Masked input mismatch. Expected {expected_masked}, got {masked_input}",
)
def test_forward_with_tp_size_1(self):
@@ -99,19 +128,15 @@ class TestCustomVocabParallelEmbedding(unittest.TestCase):
# Create a fresh mock embedding with tp_size=1
layer = self._create_layer()
layer.tp_size = 1
layer.quant_method.embedding = MagicMock(
return_value=torch.randn(3, layer.embedding_dim))
layer.quant_method.embedding = MagicMock(return_value=torch.randn(3, layer.embedding_dim))
input_ = torch.tensor([1, 2, 3])
with patch(
"vllm_ascend.ops.vocab_parallel_embedding.tensor_model_parallel_all_reduce",
side_effect=lambda x: x) as mock_reduce_tp1:
with patch("torch.ops.vllm.maybe_pad_and_reduce", side_effect=lambda x: x) as mock_reduce_tp1:
output = layer.forward(input_)
# Should just pass through without masking
layer.quant_method.embedding.assert_called_once_with(
layer, input_.long())
layer.quant_method.embedding.assert_called_once_with(layer, input_.long())
self.assertEqual(output.shape, (3, layer.embedding_dim))
# Verify all_reduce was called once
@@ -123,9 +148,7 @@ class TestCustomVocabParallelEmbedding(unittest.TestCase):
input_ = torch.tensor([15, 35]) # one org vocab, one added vocab
with patch(
"vllm_ascend.ops.vocab_parallel_embedding.tensor_model_parallel_all_reduce",
side_effect=lambda x: x) as mock_reduce_tp:
with patch("torch.ops.vllm.maybe_pad_and_reduce", side_effect=lambda x: x) as mock_reduce_tp:
# Call the forward method
output = layer.forward(input_)
@@ -146,13 +169,10 @@ class TestCustomVocabParallelEmbedding(unittest.TestCase):
input_ = torch.tensor([5, 15, 25, 35, 45]) # includes invalid cases
# Create predictable mock output
mock_output = torch.randn(5, self.embedding_dim)
layer.quant_method.embedding = MagicMock(
return_value=mock_output.clone())
layer.quant_method.embedding = MagicMock(return_value=mock_output.clone())
# Patch tensor_model_parallel_all_reduce to mock its behavior
with patch(
"vllm_ascend.ops.vocab_parallel_embedding.tensor_model_parallel_all_reduce",
side_effect=lambda x: x):
with patch("torch.ops.vllm.maybe_pad_and_reduce", side_effect=lambda x: x):
# Call the forward method
output = layer.forward(input_)
# Check that invalid positions (0, 2, 4) were zeroed out
@@ -176,17 +196,23 @@ class TestCustomVocabParallelEmbedding(unittest.TestCase):
for input_, expected_shape in test_cases:
with self.subTest(input=input_):
with patch(
"vllm_ascend.ops.vocab_parallel_embedding.tensor_model_parallel_all_reduce",
side_effect=lambda x: x):
with patch("torch.ops.vllm.maybe_pad_and_reduce", side_effect=lambda x: x):
# Call the forward method
output = layer.forward(input_)
self.assertEqual(output.shape, expected_shape)
class TestAscendLogitsProcessor(unittest.TestCase):
def setUp(self):
self.mock_vllm_config = MagicMock()
self.mock_vllm_config.compilation_config.custom_ops = ["all"]
from vllm.config.vllm import set_current_vllm_config
set_current_vllm_config(self.mock_vllm_config)
self.config_patch = patch("vllm.config.vllm.get_current_vllm_config", return_value=self.mock_vllm_config)
self.config_patch.start()
self.vocab_size = 50
self.num_embeddings = 50
self.embedding_dim = 10
@@ -198,27 +224,19 @@ class TestAscendLogitsProcessor(unittest.TestCase):
self.mock_group.rank_in_group = 0
self.mock_ascend_config = MagicMock()
self.mock_quant_method = MagicMock()
self.mock_quant_method.apply = MagicMock(
return_value=torch.randn(1, self.vocab_size))
self.mock_quant_method.apply = MagicMock(return_value=torch.randn(1, self.vocab_size))
self.patches = [
patch("vllm_ascend.ascend_config.get_ascend_config",
return_value=self.mock_ascend_config),
patch(
"vllm_ascend.ops.vocab_parallel_embedding.get_lmhead_tp_group",
return_value=self.mock_group),
patch("vllm_ascend.ops.vocab_parallel_embedding.lmhead_tp_enable",
return_value=True),
patch("vllm_ascend.ops.vocab_parallel_embedding.get_ascend_config", return_value=self.mock_ascend_config),
patch("vllm_ascend.ops.vocab_parallel_embedding.get_lmhead_tp_group", return_value=self.mock_group),
patch("vllm_ascend.ops.vocab_parallel_embedding.lmhead_tp_enable", return_value=True),
patch(
"vllm_ascend.ops.vocab_parallel_embedding.get_lmhead_tp_group.all_to_all",
return_value=torch.randn(1, self.vocab_size)),
return_value=torch.randn(1, self.vocab_size),
),
patch(
"vllm_ascend.ops.vocab_parallel_embedding.get_lmhead_tp_group.all_gather",
return_value=torch.randn(1, self.vocab_size)),
patch(
"vllm_ascend.core.schedule_config.AscendSchedulerConfig.initialize_from_config",
return_value=MagicMock(max_num_batched_tokens=1000,
max_model_len=512,
enable_chunked_prefill=False))
return_value=torch.randn(1, self.vocab_size),
),
]
for p in self.patches:
@@ -234,9 +252,9 @@ class TestAscendLogitsProcessor(unittest.TestCase):
def test_get_logits(self):
processor = AscendLogitsProcessor(vocab_size=self.vocab_size)
lmhead = AscendParallelLMHead(num_embeddings=self.num_embeddings,
embedding_dim=self.embedding_dim,
prefix="lm_head")
lmhead = AscendParallelLMHead(
num_embeddings=self.num_embeddings, embedding_dim=self.embedding_dim, prefix="lm_head"
)
lmhead.quant_method = self.mock_quant_method
lmhead.quant_method.apply = self.mock_quant_method.apply
hidden_state = torch.randn(1, self.org_num_embeddings)