Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
121 changes: 0 additions & 121 deletions tests/ut/attention/test_attention_v1.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,7 +19,6 @@
)
from vllm_ascend.attention.utils import (
AscendCommonAttentionMetadata,
PagedAttentionGraphParam,
cache_graph_workspace,
needs_layer_aware_fia_graph_replay,
using_paged_attention,
Expand Down Expand Up @@ -781,123 +780,3 @@ def test_forward_decode_only_swa_seq_len_mismatch(
mock_reshape_and_cache.assert_called_once()

assert output.shape == (10, 8, 64)

@patch("vllm_ascend.attention.attention_v1.torch.npu.stream")
@patch("vllm_ascend.attention.attention_v1.torch.npu.graph_task_update_begin")
@patch("vllm_ascend.attention.attention_v1.torch.npu.graph_task_update_end")
@patch("torch_npu.npu_fused_infer_attention_score")
@patch("vllm_ascend.attention.attention_v1.get_graph_params")
@patch("vllm_ascend.attention.attention_v1._EXTRA_CTX")
@patch("vllm_ascend.attention.attention_v1.using_paged_attention", return_value=False)
@patch("vllm_ascend.attention.attention_v1.needs_layer_aware_fia_graph_replay", return_value=False)
@patch("vllm_ascend.attention.attention_v1._ATTN_KEYS_BUFFER", new=[])
def test_update_graph_params(
self,
mock_needs_layer_aware_fia_graph_replay,
mock_using_paged_attention,
mock_EXTRA_CTX,
mock_get_graph_params,
mock_fia,
mock_graph_task_update_end,
mock_graph_task_update_begin,
mock_stream,
):
"""Test behavior when _ATTN_KEYS_BUFFER is [] after dummy_run."""

mock_EXTRA_CTX.sinks = False
mock_EXTRA_CTX.is_draft_model = False

param: list[MagicMock | None] = [MagicMock()] * 22
param[16] = None # sliding_window
param[17] = None # c8_k_aq_scale
param[21] = None # layer_name

mock_get_graph_params.return_value.attn_params = {1: [tuple(param)] * 3}
mock_get_graph_params.return_value.handles = {1: [MagicMock()] * 3}
mock_get_graph_params.return_value.events = {1: [MagicMock()] * 3}

attn_metadata_keys = [
"model.layers.10.self_attn.attn",
"model.layers.2.self_attn.attn",
"model.layers.5.self_attn.attn",
]
forward_context = MagicMock()
forward_context.attn_metadata = {key: MagicMock() for key in attn_metadata_keys}
# breakpoint()
self.impl.update_graph_params(self.mock_stream, forward_context, 1, self.mock_vllm_config)

expected = [
"model.layers.2.self_attn.attn",
"model.layers.5.self_attn.attn",
"model.layers.10.self_attn.attn",
]
self.assertEqual(attn_module._ATTN_KEYS_BUFFER, expected)
self.assertEqual(mock_fia.out.call_count, 3)

@patch("vllm_ascend.attention.attention_v1.torch.npu.stream")
@patch("vllm_ascend.attention.attention_v1.torch.npu.graph_task_update_begin")
@patch("vllm_ascend.attention.attention_v1.torch.npu.graph_task_update_end")
@patch("vllm_ascend.attention.attention_v1.torch_npu._npu_paged_attention")
@patch("vllm_ascend.attention.attention_v1.torch_npu._npu_paged_attention_get_workspace", return_value=MagicMock())
@patch("vllm_ascend.attention.attention_v1.get_graph_params")
@patch("vllm_ascend.attention.attention_v1._EXTRA_CTX")
@patch("vllm_ascend.attention.attention_v1.using_paged_attention", return_value=True)
@patch("vllm_ascend.attention.attention_v1.needs_layer_aware_fia_graph_replay", return_value=False)
@patch("vllm_ascend.attention.attention_v1._ATTN_KEYS_BUFFER", new=[])
def test_update_graph_params_handles_captured_paged_attention_params(
self,
mock_needs_layer_aware_fia_graph_replay,
mock_using_paged_attention,
mock_EXTRA_CTX,
mock_get_graph_params,
mock_get_workspace,
mock_paged_attention,
mock_graph_task_update_end,
mock_graph_task_update_begin,
mock_stream,
):
mock_EXTRA_CTX.sinks = False
mock_EXTRA_CTX.is_draft_model = False

query = MagicMock()
key_cache = MagicMock()
value_cache = MagicMock()
block_table = MagicMock()
output = MagicMock()
captured_seq_lens = MagicMock()
current_seq_lens = MagicMock()
pa_param = PagedAttentionGraphParam(
(
query,
key_cache,
value_cache,
8,
8,
1.0,
block_table,
captured_seq_lens,
output,
),
"model.layers.0.self_attn.attn",
)

mock_get_graph_params.return_value.attn_params = {1: [pa_param]}
mock_get_graph_params.return_value.handles = {1: [MagicMock()]}
mock_get_graph_params.return_value.events = {1: [MagicMock()]}

forward_context = MagicMock()
forward_context.attn_metadata = {
"model.layers.0.self_attn.attn": MagicMock(
seq_lens=current_seq_lens,
block_tables=block_table,
seq_lens_list=[10],
),
}

self.impl.update_graph_params(self.mock_stream, forward_context, 1, self.mock_vllm_config)

mock_get_workspace.assert_called_once()
mock_paged_attention.assert_called_once()
self.assertEqual(mock_paged_attention.call_args.kwargs["context_lens"], current_seq_lens)
mock_graph_task_update_begin.assert_called_once()
mock_graph_task_update_end.assert_called_once()
34 changes: 21 additions & 13 deletions tests/ut/compilation/test_acl_graph.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,7 @@
# limitations under the License.
# This file is a part of the vllm-ascend project.
#
import contextlib
import weakref
from unittest.mock import MagicMock, Mock, patch

Expand Down Expand Up @@ -47,7 +48,8 @@
from vllm_ascend.device_allocator.sleep_mem_optimized import AclGraphSleepWakeupManager


def test_update_full_graph_params_dispatches_draft_metadata_by_keyword():
@patch("vllm_ascend.compilation.acl_graph.use_updatable_graph", return_value=False)
def test_update_full_graph_params_dispatches_draft_metadata_by_keyword(mock_use_updatable):
impl_cls = MagicMock()
attn_backend = MagicMock()
attn_backend.get_impl_cls.return_value = impl_cls
Expand Down Expand Up @@ -123,6 +125,12 @@ def setUp(self):
self.addCleanup(self.get_ascend_config_patcher.stop)
self.mock_get_ascend_config.return_value.ascend_compilation_config.enable_super_kernel = False

self.exit_stack = contextlib.ExitStack()
self.addCleanup(self.exit_stack.close)
self.mock_updatable_graph = self.exit_stack.enter_context(
patch("vllm_ascend.compilation.acl_graph.UpdatableGraph")
)

# Mock VllmConfig
self.mock_vllm_config = MagicMock(spec=VllmConfig)
self.mock_vllm_config.compilation_config = MagicMock()
Expand Down Expand Up @@ -284,7 +292,7 @@ def test_call_capture_graph_first_time(

# Mock torch.npu.NPUGraph
mock_npu_graph = MagicMock()
mock_torch.npu.NPUGraph.return_value = mock_npu_graph
self.mock_updatable_graph.return_value = mock_npu_graph

# Mock torch.npu.graph context manager
mock_graph_context = MagicMock()
Expand Down Expand Up @@ -316,7 +324,7 @@ def test_call_capture_graph_first_time(

# Verify graph capture happened
mock_validate_cudagraph_capturing_enabled.assert_called_once()
mock_torch.npu.NPUGraph.assert_called_once()
self.mock_updatable_graph.assert_called_once()
mock_torch.npu.graph.assert_called_once_with(mock_npu_graph, pool=self.mock_graph_pool)
self.mock_runnable.assert_called_once_with(test_tensor, "arg2")

Expand Down Expand Up @@ -367,7 +375,7 @@ def test_capture_respects_super_kernel_setting(
self.mock_get_ascend_config.return_value.ascend_compilation_config.enable_super_kernel = enabled

mock_npu_graph = MagicMock()
mock_torch.npu.NPUGraph.return_value = mock_npu_graph
self.mock_updatable_graph.return_value = mock_npu_graph
mock_graph_context = MagicMock()
mock_torch.npu.graph.return_value = mock_graph_context
mock_graph_context.__enter__ = Mock(return_value=None)
Expand Down Expand Up @@ -422,7 +430,7 @@ def test_call_replay_graph(

# Mock torch.npu.NPUGraph
mock_npu_graph = MagicMock()
mock_torch.npu.NPUGraph.return_value = mock_npu_graph
self.mock_updatable_graph.return_value = mock_npu_graph

# Mock torch.npu.graph context manager
mock_graph_context = MagicMock()
Expand Down Expand Up @@ -454,7 +462,7 @@ def test_call_replay_graph(

# Verify graph capture happened during first call
mock_validate_cudagraph_capturing_enabled.assert_called_once()
mock_torch.npu.NPUGraph.assert_called_once()
self.mock_updatable_graph.assert_called_once()
mock_torch.npu.graph.assert_called_once()

# Reset mock to track second call
Expand Down Expand Up @@ -501,7 +509,7 @@ def test_call_with_debug_mode_input_address_check(

# Mock torch.npu.NPUGraph
mock_npu_graph = MagicMock()
mock_torch.npu.NPUGraph.return_value = mock_npu_graph
self.mock_updatable_graph.return_value = mock_npu_graph

# Mock torch.npu.graph context manager
mock_graph_context = MagicMock()
Expand Down Expand Up @@ -562,7 +570,7 @@ def test_call_with_debug_mode_input_address_mismatch(

# Mock torch.npu.NPUGraph
mock_npu_graph = MagicMock()
mock_torch.npu.NPUGraph.return_value = mock_npu_graph
self.mock_updatable_graph.return_value = mock_npu_graph

# Mock torch.npu.graph context manager
mock_graph_context = MagicMock()
Expand Down Expand Up @@ -633,7 +641,7 @@ def test_call_capture_graph_with_gc_disable(

# Mock torch.npu.NPUGraph
mock_npu_graph = MagicMock()
mock_torch.npu.NPUGraph.return_value = mock_npu_graph
self.mock_updatable_graph.return_value = mock_npu_graph

# Mock torch.npu.graph context manager
mock_graph_context = MagicMock()
Expand Down Expand Up @@ -675,7 +683,7 @@ def test_call_capture_graph_with_gc_disable(

# Verify graph capture happened
mock_validate_cudagraph_capturing_enabled.assert_called_once()
mock_torch.npu.NPUGraph.assert_called_once()
self.mock_updatable_graph.assert_called_once()
mock_torch.npu.graph.assert_called_once_with(mock_npu_graph, pool=self.mock_graph_pool)

# Should return the original output (not weak ref) since weak_ref_output is not enabled
Expand Down Expand Up @@ -712,7 +720,7 @@ def test_call_capture_graph_with_weak_ref_output(

# Mock torch.npu.NPUGraph
mock_npu_graph = MagicMock()
mock_torch.npu.NPUGraph.return_value = mock_npu_graph
self.mock_updatable_graph.return_value = mock_npu_graph

# Mock torch.npu.graph context manager
mock_graph_context = MagicMock()
Expand Down Expand Up @@ -749,7 +757,7 @@ def test_call_capture_graph_with_weak_ref_output(

# Verify graph capture happened
mock_validate_cudagraph_capturing_enabled.assert_called_once()
mock_torch.npu.NPUGraph.assert_called_once()
self.mock_updatable_graph.assert_called_once()
mock_torch.npu.graph.assert_called_once_with(mock_npu_graph, pool=self.mock_graph_pool)

# Should return the weak ref output when weak_ref_output option is enabled
Expand Down Expand Up @@ -778,7 +786,7 @@ def test_call_capture_graph_with_debug_log(
with patch("vllm_ascend.compilation.acl_graph.torch") as mock_torch:
# Mock torch.npu.NPUGraph
mock_npu_graph = MagicMock()
mock_torch.npu.NPUGraph.return_value = mock_npu_graph
self.mock_updatable_graph.return_value = mock_npu_graph

# Mock torch.npu.graph context manager
mock_graph_context = MagicMock()
Expand Down
2 changes: 2 additions & 0 deletions tests/ut/spec_decode/test_dspark_proposer.py
Original file line number Diff line number Diff line change
Expand Up @@ -154,7 +154,9 @@ def release():
proposer.parallel_drafting = True
proposer.token_indices_to_sample = torch.zeros(2, dtype=torch.int32)
proposer.enable_enpu = False
proposer.draft_attn_groups = [MagicMock()]
proposer._update_full_graph_params_if_needed = MagicMock()
proposer._maybe_update_metadata = MagicMock()
proposer.set_inputs_first_pass = MagicMock()
proposer.build_draft_attn_metadata = MagicMock()

Expand Down
14 changes: 14 additions & 0 deletions tests/ut/spec_decode/test_eagle_proposer.py
Original file line number Diff line number Diff line change
Expand Up @@ -592,19 +592,26 @@ def setUp(self):
self.mock_dp_group = patch("vllm_ascend.ascend_forward_context.get_dp_group", return_value=mock_dp_group)
self.mock_dp_group.start()

self.mock_use_updatable_graph = patch(
"vllm_ascend.spec_decode.llm_base_proposer.use_updatable_graph", return_value=False
)
self.mock_use_updatable_graph.start()

# Set the current vllm config
set_current_vllm_config(self.vllm_config)
self.proposer = AscendEagleProposer(vllm_config=self.vllm_config, device=self.device, runner=self.runner)
self.proposer.model = MagicMock()
self.proposer._runnable = MagicMock()
self.proposer.update_stream = MagicMock()
self.proposer.draft_attn_groups = [MagicMock()]

def tearDown(self):
self.mock_get_ascend_config.stop()
self.mock_cpugpubuffer.stop()
self.mock_supports_multimodal_inputs.stop()
self.mock_tp_world_size.stop()
self.mock_dp_group.stop()
self.mock_use_updatable_graph.stop()
# Clear the current vllm config
set_current_vllm_config(None)

Expand Down Expand Up @@ -649,6 +656,7 @@ def test_dummy_run_in_graph_capture(
mock_get_context.return_value = mock_return_context
mock_get_context_2.return_value = mock_return_context
self.proposer.use_cuda_graph = True
self.proposer.draft_attn_groups = [MagicMock()]
# cpu does not support `torch.ops.vllm.maybe_pad_and_reduce`
with set_current_vllm_config(self.vllm_config):
self.proposer.dummy_run(num_tokens=64, in_graph_capturing=True, aclgraph_runtime_mode=CUDAGraphMode.FULL)
Expand Down Expand Up @@ -843,13 +851,19 @@ def setUp_and_tearDown(self):
set_current_vllm_config(self.vllm_config)
self.proposer = AscendEagleProposer(vllm_config=self.vllm_config, device=self.device, runner=self.runner)

self.mock_use_updatable_graph = patch(
"vllm_ascend.spec_decode.llm_base_proposer.use_updatable_graph", return_value=False
)
self.mock_use_updatable_graph.start()

yield

self.mock_cpugpubuffer.stop()
self.mock_supports_multimodal_inputs.stop()
self.mock_tp_world_size.stop()
self.mock_dp_group.stop()
self.mock_get_ascend_config.stop()
self.mock_use_updatable_graph.stop()
# Clear the current vllm config
set_current_vllm_config(None)
clear_ascend_config()
Expand Down
Loading