0
tests/ut/ops/__init__.py
Normal file
0
tests/ut/ops/__init__.py
Normal file
0
tests/ut/ops/a2/__init__.py
Normal file
0
tests/ut/ops/a2/__init__.py
Normal file
439
tests/ut/ops/a2/test_gdn_chunk_meta.py
Normal file
439
tests/ut/ops/a2/test_gdn_chunk_meta.py
Normal 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)
|
||||
827
tests/ut/ops/a2/test_token_dispatcher.py
Normal file
827
tests/ut/ops/a2/test_token_dispatcher.py
Normal 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)
|
||||
458
tests/ut/ops/a2/test_weight_prefetch.py
Normal file
458
tests/ut/ops/a2/test_weight_prefetch.py
Normal 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
|
||||
0
tests/ut/ops/a3_2/__init__.py
Normal file
0
tests/ut/ops/a3_2/__init__.py
Normal file
343
tests/ut/ops/a3_2/test_activation.py
Normal file
343
tests/ut/ops/a3_2/test_activation.py
Normal 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)
|
||||
710
tests/ut/ops/a3_2/test_select_experts.py
Normal file
710
tests/ut/ops/a3_2/test_select_experts.py
Normal 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)
|
||||
@@ -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)
|
||||
|
||||
487
tests/ut/ops/test_flashcomm2_oshard_manager.py
Normal file
487
tests/ut/ops/test_flashcomm2_oshard_manager.py
Normal 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()
|
||||
950
tests/ut/ops/test_fused_moe.py
Normal file
950
tests/ut/ops/test_fused_moe.py
Normal 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)
|
||||
70
tests/ut/ops/test_gate_linear.py
Normal file
70
tests/ut/ops/test_gate_linear.py
Normal 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()
|
||||
923
tests/ut/ops/test_gdn_attn_builder.py
Normal file
923
tests/ut/ops/test_gdn_attn_builder.py
Normal 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),
|
||||
)
|
||||
468
tests/ut/ops/test_layer_shard_linear.py
Normal file
468
tests/ut/ops/test_layer_shard_linear.py
Normal 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
|
||||
@@ -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)
|
||||
|
||||
@@ -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
162
tests/ut/ops/test_mla.py
Normal 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))
|
||||
195
tests/ut/ops/test_mm_encoder_attention.py
Normal file
195
tests/ut/ops/test_mm_encoder_attention.py
Normal 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])
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
253
tests/ut/ops/test_moe_mlp.py
Normal file
253
tests/ut/ops/test_moe_mlp.py
Normal 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)
|
||||
270
tests/ut/ops/test_moe_runtime_args.py
Normal file
270
tests/ut/ops/test_moe_runtime_args.py
Normal file
@@ -0,0 +1,270 @@
|
||||
#
|
||||
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
|
||||
# Copyright 2023 The vLLM team.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
#
|
||||
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)
|
||||
217
tests/ut/ops/test_prepare_finalize.py
Normal file
217
tests/ut/ops/test_prepare_finalize.py
Normal 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)
|
||||
@@ -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__
|
||||
)
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user