1725 lines
79 KiB
Python
1725 lines
79 KiB
Python
#
|
|
# 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 sys
|
|
from unittest.mock import MagicMock, patch
|
|
|
|
import torch
|
|
from vllm.distributed.parallel_state import GroupCoordinator
|
|
|
|
from tests.ut.attention.utils import patch_distributed_groups
|
|
from tests.ut.base import TestBase
|
|
from vllm_ascend.ascend_config import init_ascend_config
|
|
from vllm_ascend.attention.attention_v1 import AscendAttentionState
|
|
|
|
if "torch_npu._inductor" not in sys.modules:
|
|
sys.modules["torch_npu._inductor"] = MagicMock()
|
|
|
|
from vllm_ascend.attention.context_parallel.common_cp import AscendPCPMetadata
|
|
from vllm_ascend.attention.context_parallel.sfa_cp import (
|
|
AscendSFACPImpl,
|
|
AscendSFACPMetadataBuilder,
|
|
AscendSFADCPImpl,
|
|
)
|
|
from vllm_ascend.attention.sfa_v1 import AscendSFAImpl, AscendSFAMetadata, DCPContext, DSACPContext
|
|
|
|
|
|
def _make_indexer_mock():
|
|
indexer = MagicMock()
|
|
indexer.n_head = 64
|
|
indexer.head_dim = 128
|
|
indexer.wq_b = MagicMock()
|
|
indexer.wk_weights_proj = MagicMock()
|
|
indexer.k_norm = MagicMock()
|
|
return indexer
|
|
|
|
|
|
def _make_impl_kwargs(extra=None):
|
|
kv_a_layernorm = MagicMock()
|
|
kv_a_layernorm.weight = torch.randn(96)
|
|
kv_a_layernorm.variance_epsilon = 1e-6
|
|
kwargs = {
|
|
"kv_lora_rank": 32,
|
|
"qk_nope_head_dim": 64,
|
|
"qk_rope_head_dim": 32,
|
|
"qk_head_dim": 96,
|
|
"v_head_dim": 128,
|
|
"q_lora_rank": 64,
|
|
"q_proj": MagicMock(),
|
|
"q_b_proj": MagicMock(),
|
|
"kv_b_proj": MagicMock(),
|
|
"o_proj": MagicMock(),
|
|
"kv_a_proj_with_mqa": MagicMock(),
|
|
"fused_qkv_a_proj": MagicMock(),
|
|
"kv_a_layernorm": kv_a_layernorm,
|
|
"q_a_layernorm": MagicMock(),
|
|
"rotary_emb": MagicMock(),
|
|
"indexer": _make_indexer_mock(),
|
|
"layer_name": "layer_0",
|
|
}
|
|
if extra:
|
|
kwargs.update(extra)
|
|
return kwargs
|
|
|
|
|
|
class TestAscendSFACPMetadataBuilder(TestBase):
|
|
"""Tests for AscendSFACPMetadataBuilder."""
|
|
|
|
@patch("vllm.distributed.parallel_state._TP", new_callable=lambda: MagicMock(spec=GroupCoordinator))
|
|
def setUp(self, mock_tp):
|
|
mock_tp.world_size = 2
|
|
mock_tp.rank_in_group = MagicMock()
|
|
mock_tp.device_group = MagicMock()
|
|
|
|
self.mock_cfg = MagicMock()
|
|
self.mock_cfg.parallel_config = MagicMock()
|
|
self.mock_cfg.parallel_config.tensor_parallel_size = 1
|
|
self.mock_cfg.parallel_config.prefill_context_parallel_size = 1
|
|
self.mock_cfg.parallel_config.decode_context_parallel_size = 1
|
|
|
|
self.mock_cfg.compilation_config = MagicMock()
|
|
self.mock_cfg.compilation_config.pass_config = MagicMock()
|
|
self.mock_cfg.compilation_config.pass_config.enable_sp = False
|
|
|
|
self.mock_cfg.speculative_config.num_speculative_tokens = 0
|
|
|
|
self.patcher = patch("vllm.config.get_current_vllm_config", return_value=self.mock_cfg)
|
|
self.patcher.start()
|
|
|
|
# Mock parent class __init__ to avoid complex initialization,
|
|
# but still set the essential attributes that child class needs.
|
|
def mock_parent_init(
|
|
self, kv_cache_spec, layer_names, vllm_config, device, metadata_cls, supports_dcp_with_varlen
|
|
):
|
|
self.metadata_cls = metadata_cls
|
|
self.kv_cache_spec = kv_cache_spec
|
|
self.model_config = vllm_config.model_config
|
|
self.vllm_config = vllm_config
|
|
self.device = device
|
|
self.chunked_prefill_workspace_size = 128 * 1024
|
|
self.chunked_prefill_workspace = torch.empty(
|
|
(self.chunked_prefill_workspace_size, vllm_config.model_config.get_head_size()),
|
|
dtype=vllm_config.model_config.dtype,
|
|
device=device,
|
|
)
|
|
|
|
self.parent_init_patcher = patch(
|
|
"vllm.model_executor.layers.attention.mla_attention.MLACommonMetadataBuilder.__init__", mock_parent_init
|
|
)
|
|
self.parent_init_patcher.start()
|
|
|
|
def tearDown(self):
|
|
self.patcher.stop()
|
|
self.parent_init_patcher.stop()
|
|
|
|
def _make_vllm_config(self):
|
|
vllm_config = MagicMock()
|
|
vllm_config.cache_config.block_size = 16
|
|
vllm_config.model_config.max_model_len = 1024
|
|
vllm_config.model_config.get_head_size.return_value = 64
|
|
vllm_config.model_config.dtype = torch.float16
|
|
vllm_config.model_config.hf_text_config.qk_rope_head_dim = 64
|
|
vllm_config.model_config.hf_text_config = MagicMock(qk_rope_head_dim=64)
|
|
vllm_config.model_config.hf_config.model_type = "deepseek_v3"
|
|
speculative_config = MagicMock()
|
|
speculative_config.num_speculative_tokens = 0
|
|
vllm_config.speculative_config = speculative_config
|
|
vllm_config.scheduler_config.max_num_seqs = 16
|
|
vllm_config.scheduler_config.max_num_batched_tokens = 256
|
|
vllm_config.parallel_config = MagicMock()
|
|
vllm_config.parallel_config.cp_kv_cache_interleave_size = 1
|
|
vllm_config.parallel_config.tensor_parallel_size = 1
|
|
vllm_config.parallel_config.prefill_context_parallel_size = 2
|
|
vllm_config.parallel_config.decode_context_parallel_size = 2
|
|
vllm_config.kv_transfer_config = None
|
|
return vllm_config
|
|
|
|
def _build_builder(self, pcp_size=2, dcp_size=2):
|
|
kv_cache_spec = MagicMock(block_size=128)
|
|
layer_names = ["layer1", "layer2"]
|
|
vllm_config = self._make_vllm_config()
|
|
device = torch.device("cpu")
|
|
builder = AscendSFACPMetadataBuilder(
|
|
kv_cache_spec=kv_cache_spec,
|
|
layer_names=layer_names,
|
|
vllm_config=vllm_config,
|
|
device=device,
|
|
)
|
|
# The parent mock above sets minimal attributes, set the rest by ourselves
|
|
builder.block_size = 16
|
|
builder.speculative_config = vllm_config.speculative_config
|
|
builder.decode_threshold = 1
|
|
builder.reorder_batch_threshold = 1
|
|
return builder
|
|
|
|
@patch("vllm_ascend.attention.context_parallel.sfa_cp.enabling_mlapo")
|
|
@patch_distributed_groups(dcp_size=2, pcp_size=2, needs_mocks=False)
|
|
def test_init_default(self, mock_enabling_mlapo):
|
|
mock_enabling_mlapo.return_value = False
|
|
builder = self._build_builder(pcp_size=2, dcp_size=2)
|
|
self.assertEqual(builder.pcp_size, 2)
|
|
self.assertEqual(builder.pcp_rank, 0)
|
|
self.assertEqual(builder.dcp_size, 2)
|
|
self.assertEqual(builder.dcp_rank, 0)
|
|
self.assertFalse(builder.enable_mlapo)
|
|
self.assertEqual(builder.cp_local_block_size, 1)
|
|
# cp_virtual_block_size = 1 * 2 * 2 = 4
|
|
self.assertEqual(builder.cp_virtual_block_size, 4)
|
|
# block_size = lcm(16, 4) = 16
|
|
self.assertEqual(builder.block_size, 16)
|
|
self.assertIsNotNone(builder.slot_mapping_buf)
|
|
self.assertEqual(builder.block_arange_buffer.shape[0], 4)
|
|
|
|
@patch("vllm_ascend.attention.context_parallel.sfa_cp.enabling_mlapo")
|
|
@patch_distributed_groups(dcp_size=1, pcp_size=1, needs_mocks=False)
|
|
def test_init_no_cp(self, mock_enabling_mlapo):
|
|
mock_enabling_mlapo.return_value = False
|
|
builder = self._build_builder(pcp_size=1, dcp_size=1)
|
|
self.assertEqual(builder.pcp_size, 1)
|
|
self.assertEqual(builder.pcp_rank, 0)
|
|
self.assertEqual(builder.dcp_size, 1)
|
|
self.assertEqual(builder.dcp_rank, 0)
|
|
self.assertIsNone(builder.pcp_group)
|
|
self.assertIsNone(builder.dcp_group)
|
|
|
|
@patch("vllm_ascend.attention.context_parallel.sfa_cp.enabling_mlapo")
|
|
@patch_distributed_groups(dcp_size=2, pcp_size=2, needs_mocks=False)
|
|
def test_init_with_mlapo_enabled(self, mock_enabling_mlapo):
|
|
mock_enabling_mlapo.return_value = True
|
|
builder = self._build_builder(pcp_size=2, dcp_size=2)
|
|
self.assertTrue(builder.enable_mlapo)
|
|
|
|
@patch("vllm_ascend.attention.context_parallel.sfa_cp.enabling_mlapo", return_value=False)
|
|
@patch_distributed_groups(dcp_size=2, pcp_size=2, needs_mocks=False)
|
|
def test_compact_varlen_decode_slot_mapping_basic(self, mock_enabling_mlapo):
|
|
builder = self._build_builder()
|
|
# pcp_size=2, total tokens=6 (3 per req with pcp gather)
|
|
# decode_query_lens: [2, 1] (2 + 1 = 3 total valid)
|
|
# slot_mapping with pcp expansion: each decode token has pcp_size=2 entries
|
|
# Layout: [t0_p0, t0_p1, t1_p0, t1_p1, t2_p0, t2_p1] = 6 tokens
|
|
# req0 spans 2 tokens => valid_in: [0, 2], req1 spans 1 token => valid_in: [4]
|
|
decode_slot_mapping = torch.tensor([10, 20, 30, 40, 50, 60], dtype=torch.int32)
|
|
decode_query_lens = torch.tensor([2, 1], dtype=torch.int64)
|
|
builder._compact_varlen_decode_slot_mapping(decode_slot_mapping, decode_query_lens)
|
|
# With pcp_size=2:
|
|
# req_spans = [4, 2], req_starts = [0, 4]
|
|
# token_offsets after rebase: [0, 1, 0]
|
|
# valid_in_idx = [0, 2, 4] => slots [10, 30, 50]
|
|
# valid_out_idx = [0, 1, 4]
|
|
# Final: pos 0=10, 1=30, 4=50, others=-1
|
|
self.assertEqual(decode_slot_mapping[0].item(), 10)
|
|
self.assertEqual(decode_slot_mapping[1].item(), 30)
|
|
self.assertEqual(decode_slot_mapping[2].item(), -1)
|
|
self.assertEqual(decode_slot_mapping[3].item(), -1)
|
|
self.assertEqual(decode_slot_mapping[4].item(), 50)
|
|
self.assertEqual(decode_slot_mapping[5].item(), -1)
|
|
|
|
@patch("vllm_ascend.attention.context_parallel.sfa_cp.enabling_mlapo", return_value=False)
|
|
@patch_distributed_groups(dcp_size=2, pcp_size=2, needs_mocks=False)
|
|
def test_compact_varlen_decode_slot_mapping_zero_tokens(self, mock_enabling_mlapo):
|
|
builder = self._build_builder()
|
|
decode_slot_mapping = torch.tensor([10, 20, 30], dtype=torch.int32)
|
|
decode_query_lens = torch.tensor([0, 0], dtype=torch.int64)
|
|
# Should return early without modification
|
|
original = decode_slot_mapping.clone()
|
|
builder._compact_varlen_decode_slot_mapping(decode_slot_mapping, decode_query_lens)
|
|
self.assertTrue(torch.equal(decode_slot_mapping, original))
|
|
|
|
@patch("vllm_ascend.attention.context_parallel.sfa_cp.enabling_mlapo", return_value=False)
|
|
@patch_distributed_groups(dcp_size=2, pcp_size=2, needs_mocks=False)
|
|
def test_build_prefill_compact_block_metadata(self, mock_enabling_mlapo):
|
|
builder = self._build_builder()
|
|
# Make a block_table with 3 reqs, 1 decode, 2 prefills, 4 blocks per req
|
|
block_table = torch.tensor(
|
|
[
|
|
[0, 0, 0, 0],
|
|
[1, 2, 3, 4],
|
|
[5, 6, 1, 2],
|
|
],
|
|
dtype=torch.int32,
|
|
)
|
|
valid_block_ids, block_table_cp = builder.build_prefill_compact_block_metadata(block_table, num_decodes=1)
|
|
# prefill block_table covers reqs 1 and 2: blocks 1, 2, 3, 4, 5, 6, 1, 2 (8 entries)
|
|
# unique: [1, 2, 3, 4, 5, 6]
|
|
self.assertEqual(valid_block_ids.numel(), 6)
|
|
# block_table_cp shape should be (num_prefill_reqs, num_blocks_per_req * pcp*dcp)
|
|
self.assertEqual(block_table_cp.shape[0], 2)
|
|
# 4 blocks per req * (pcp_size * dcp_size) = 4 * 4 = 16
|
|
self.assertEqual(block_table_cp.shape[1], 16)
|
|
|
|
@patch("vllm_ascend.attention.context_parallel.sfa_cp.enabling_mlapo", return_value=False)
|
|
@patch_distributed_groups(dcp_size=2, pcp_size=2, needs_mocks=False)
|
|
def test_build_cp_metadata(self, mock_enabling_mlapo):
|
|
builder = self._build_builder()
|
|
block_arange = builder.block_arange_buffer
|
|
seq_lens = torch.tensor([8, 16], dtype=torch.int32)
|
|
|
|
common_attn_metadata = MagicMock()
|
|
long_seq_metadata = MagicMock()
|
|
long_seq_metadata.q_head_idx_tensor = torch.tensor([0, 1])
|
|
long_seq_metadata.q_tail_idx_tensor = torch.tensor([2, 3])
|
|
long_seq_metadata.q_full_idx = torch.tensor([0, 1, 2, 3])
|
|
long_seq_metadata.pcp_allgather_restore_idx = torch.tensor([0, 1, 2, 3])
|
|
common_attn_metadata.prefill_context_parallel_metadata = long_seq_metadata
|
|
common_attn_metadata.num_computed_tokens_cpu = torch.tensor([0, 0], dtype=torch.int32)
|
|
|
|
result = builder.build_cp_metadata(block_arange, seq_lens, common_attn_metadata)
|
|
self.assertIsInstance(result, AscendPCPMetadata)
|
|
self.assertIs(result.q_head_idx, long_seq_metadata.q_head_idx_tensor)
|
|
self.assertIs(result.q_tail_idx, long_seq_metadata.q_tail_idx_tensor)
|
|
self.assertIsNotNone(result.head_attn_nomask_seqlens)
|
|
self.assertIsNotNone(result.tail_attn_nomask_seqlens)
|
|
|
|
@patch("vllm_ascend.attention.context_parallel.sfa_cp.enabling_mlapo", return_value=False)
|
|
@patch("vllm_ascend.attention.context_parallel.sfa_cp.split_decodes_and_prefills")
|
|
@patch_distributed_groups(dcp_size=1, pcp_size=1, needs_mocks=False)
|
|
def test_build_decode_only_no_pcp(self, mock_split, mock_mlapo):
|
|
# Build path with no prefills, no pcp, simplest case
|
|
builder = self._build_builder(pcp_size=1, dcp_size=1)
|
|
mock_split.return_value = (2, 0, 2, 0) # decodes, prefills, decode_tokens, prefill_tokens
|
|
|
|
common_attn_metadata = MagicMock()
|
|
common_attn_metadata.num_reqs = 2
|
|
common_attn_metadata.num_actual_tokens = 2
|
|
common_attn_metadata.num_input_tokens = 2
|
|
|
|
long_seq_metadata = MagicMock()
|
|
long_seq_metadata.q_head_idx_tensor = torch.tensor([0])
|
|
long_seq_metadata.q_tail_idx_tensor = torch.tensor([1])
|
|
long_seq_metadata.q_full_idx = torch.tensor([0, 1])
|
|
long_seq_metadata.pcp_allgather_restore_idx = torch.tensor([0, 1])
|
|
long_seq_metadata.num_actual_tokens_pcp_padded = 2
|
|
common_attn_metadata.prefill_context_parallel_metadata = long_seq_metadata
|
|
common_attn_metadata.num_computed_tokens_cpu = torch.tensor([0, 0], dtype=torch.int32)
|
|
common_attn_metadata.slot_mapping = torch.arange(8, dtype=torch.int32)
|
|
|
|
# Mock super().build()
|
|
fake_metadata = AscendSFAMetadata(
|
|
num_actual_tokens=2,
|
|
slot_mapping=torch.zeros(2, dtype=torch.int32),
|
|
seq_lens=torch.tensor([4, 4], dtype=torch.int32),
|
|
seq_lens_cpu=torch.tensor([4, 4], dtype=torch.int32),
|
|
cum_query_lens=torch.tensor([1, 2], dtype=torch.int32),
|
|
block_table=torch.zeros((2, 4), dtype=torch.int32),
|
|
sin=torch.randn(2, 32),
|
|
cos=torch.randn(2, 32),
|
|
num_input_tokens=2,
|
|
attn_state=AscendAttentionState.DecodeOnly,
|
|
)
|
|
with patch.object(
|
|
AscendSFACPMetadataBuilder.__bases__[0],
|
|
"build",
|
|
return_value=fake_metadata,
|
|
):
|
|
result = builder.build(common_prefix_len=0, common_attn_metadata=common_attn_metadata)
|
|
self.assertIs(result, fake_metadata)
|
|
self.assertEqual(result.num_decodes, 2)
|
|
self.assertEqual(result.num_decode_tokens, 2)
|
|
self.assertEqual(result.num_prefills, 0)
|
|
# In pcp_size=1 path, sfa_cp_metadata should be set but block_table_cp/valid_block_ids None
|
|
self.assertIsNotNone(result.sfa_cp_metadata)
|
|
self.assertIsNone(result.sfa_cp_metadata.valid_block_ids)
|
|
self.assertIsNone(result.sfa_cp_metadata.block_table_cp)
|
|
|
|
@patch("vllm_ascend.attention.context_parallel.sfa_cp.enabling_mlapo", return_value=False)
|
|
@patch("vllm_ascend.attention.context_parallel.sfa_cp.split_decodes_and_prefills")
|
|
@patch_distributed_groups(dcp_size=2, pcp_size=2, needs_mocks=False)
|
|
def test_build_with_prefills_and_decodes(self, mock_split, mock_mlapo):
|
|
builder = self._build_builder(pcp_size=2, dcp_size=2)
|
|
# 1 decode + 2 prefills, 1 decode token, 6 prefill tokens
|
|
mock_split.return_value = (1, 2, 1, 6)
|
|
|
|
common_attn_metadata = MagicMock()
|
|
common_attn_metadata.num_reqs = 3
|
|
common_attn_metadata.num_actual_tokens = 7
|
|
common_attn_metadata.num_input_tokens = 7
|
|
|
|
long_seq_metadata = MagicMock()
|
|
long_seq_metadata.q_head_idx_tensor = torch.tensor([0, 1])
|
|
long_seq_metadata.q_tail_idx_tensor = torch.tensor([2, 3])
|
|
long_seq_metadata.q_full_idx = torch.tensor([0, 1, 2, 3])
|
|
long_seq_metadata.pcp_allgather_restore_idx = torch.tensor([0, 1, 2, 3])
|
|
long_seq_metadata.num_actual_tokens_pcp_padded = 14
|
|
long_seq_metadata.query_lens_pcp_full_cpu = torch.tensor([1, 3, 3], dtype=torch.int32)
|
|
common_attn_metadata.prefill_context_parallel_metadata = long_seq_metadata
|
|
common_attn_metadata.num_computed_tokens_cpu = torch.tensor([0, 0, 0], dtype=torch.int32)
|
|
common_attn_metadata.slot_mapping = torch.arange(64, dtype=torch.int32)
|
|
|
|
block_table = torch.tensor(
|
|
[
|
|
[0, 0, 0, 0],
|
|
[1, 2, 3, 4],
|
|
[5, 6, 7, 8],
|
|
],
|
|
dtype=torch.int32,
|
|
)
|
|
fake_metadata = AscendSFAMetadata(
|
|
num_actual_tokens=7,
|
|
slot_mapping=torch.zeros(7, dtype=torch.int32),
|
|
seq_lens=torch.tensor([4, 8, 8], dtype=torch.int32),
|
|
seq_lens_cpu=torch.tensor([4, 8, 8], dtype=torch.int32),
|
|
cum_query_lens=torch.tensor([1, 4, 7], dtype=torch.int32),
|
|
block_table=block_table,
|
|
sin=torch.randn(7, 32),
|
|
cos=torch.randn(7, 32),
|
|
num_input_tokens=7,
|
|
attn_state=AscendAttentionState.ChunkedPrefill,
|
|
)
|
|
with patch.object(
|
|
AscendSFACPMetadataBuilder.__bases__[0],
|
|
"build",
|
|
return_value=fake_metadata,
|
|
):
|
|
result = builder.build(common_prefix_len=0, common_attn_metadata=common_attn_metadata)
|
|
self.assertEqual(result.num_decodes, 1)
|
|
self.assertEqual(result.num_prefills, 2)
|
|
self.assertEqual(result.num_decode_tokens, 1)
|
|
self.assertIsNotNone(result.sfa_cp_metadata)
|
|
self.assertIsNotNone(result.sfa_cp_metadata.valid_block_ids)
|
|
self.assertIsNotNone(result.sfa_cp_metadata.block_table_cp)
|
|
self.assertIsNotNone(result.sfa_cp_metadata.prefill_q_cum_seqlens)
|
|
|
|
@patch("vllm_ascend.attention.context_parallel.sfa_cp.enabling_mlapo", return_value=False)
|
|
@patch("vllm_ascend.attention.context_parallel.sfa_cp.split_decodes_and_prefills")
|
|
@patch_distributed_groups(dcp_size=2, pcp_size=2, needs_mocks=False)
|
|
def test_build_prefills_only(self, mock_split, mock_mlapo):
|
|
# Verifies prefill_q_cum_seqlens equals actual_seq_lengths_query when no decodes
|
|
builder = self._build_builder(pcp_size=2, dcp_size=2)
|
|
mock_split.return_value = (0, 2, 0, 6)
|
|
|
|
common_attn_metadata = MagicMock()
|
|
common_attn_metadata.num_reqs = 2
|
|
common_attn_metadata.num_actual_tokens = 6
|
|
common_attn_metadata.num_input_tokens = 6
|
|
|
|
long_seq_metadata = MagicMock()
|
|
long_seq_metadata.q_head_idx_tensor = torch.tensor([0, 1])
|
|
long_seq_metadata.q_tail_idx_tensor = torch.tensor([2, 3])
|
|
long_seq_metadata.q_full_idx = torch.tensor([0, 1, 2, 3])
|
|
long_seq_metadata.pcp_allgather_restore_idx = torch.tensor([0, 1, 2, 3])
|
|
long_seq_metadata.num_actual_tokens_pcp_padded = 12
|
|
long_seq_metadata.query_lens_pcp_full_cpu = torch.tensor([3, 3], dtype=torch.int32)
|
|
common_attn_metadata.prefill_context_parallel_metadata = long_seq_metadata
|
|
common_attn_metadata.num_computed_tokens_cpu = torch.tensor([0, 0], dtype=torch.int32)
|
|
common_attn_metadata.slot_mapping = torch.arange(64, dtype=torch.int32)
|
|
|
|
block_table = torch.tensor([[0, 1], [2, 3]], dtype=torch.int32)
|
|
fake_metadata = AscendSFAMetadata(
|
|
num_actual_tokens=6,
|
|
slot_mapping=torch.zeros(6, dtype=torch.int32),
|
|
seq_lens=torch.tensor([8, 8], dtype=torch.int32),
|
|
seq_lens_cpu=torch.tensor([8, 8], dtype=torch.int32),
|
|
cum_query_lens=torch.tensor([3, 6], dtype=torch.int32),
|
|
block_table=block_table,
|
|
sin=torch.randn(6, 32),
|
|
cos=torch.randn(6, 32),
|
|
num_input_tokens=6,
|
|
attn_state=AscendAttentionState.ChunkedPrefill,
|
|
)
|
|
with patch.object(
|
|
AscendSFACPMetadataBuilder.__bases__[0],
|
|
"build",
|
|
return_value=fake_metadata,
|
|
):
|
|
result = builder.build(common_prefix_len=0, common_attn_metadata=common_attn_metadata)
|
|
self.assertEqual(result.num_decodes, 0)
|
|
self.assertEqual(result.num_prefills, 2)
|
|
self.assertIsNotNone(result.sfa_cp_metadata.prefill_q_cum_seqlens)
|
|
|
|
@patch("vllm_ascend.attention.context_parallel.sfa_cp.enabling_mlapo", return_value=True)
|
|
@patch("vllm_ascend.attention.context_parallel.sfa_cp.split_decodes_and_prefills")
|
|
@patch_distributed_groups(dcp_size=2, pcp_size=2, needs_mocks=False)
|
|
def test_build_with_mlapo_enabled(self, mock_split, mock_mlapo):
|
|
# When mlapo is on: slot_mapping is compacted by pcp_size
|
|
builder = self._build_builder(pcp_size=2, dcp_size=2)
|
|
# 2 decodes, 0 prefills
|
|
mock_split.return_value = (2, 0, 2, 0)
|
|
|
|
common_attn_metadata = MagicMock()
|
|
common_attn_metadata.num_reqs = 2
|
|
common_attn_metadata.num_actual_tokens = 2
|
|
common_attn_metadata.num_input_tokens = 2
|
|
|
|
long_seq_metadata = MagicMock()
|
|
long_seq_metadata.q_head_idx_tensor = torch.tensor([0])
|
|
long_seq_metadata.q_tail_idx_tensor = torch.tensor([1])
|
|
long_seq_metadata.q_full_idx = torch.tensor([0, 1])
|
|
long_seq_metadata.pcp_allgather_restore_idx = torch.tensor([0, 1])
|
|
long_seq_metadata.num_actual_tokens_pcp_padded = 4
|
|
common_attn_metadata.prefill_context_parallel_metadata = long_seq_metadata
|
|
common_attn_metadata.num_computed_tokens_cpu = torch.tensor([0, 0], dtype=torch.int32)
|
|
common_attn_metadata.slot_mapping = torch.tensor(
|
|
[10, 20, 30, 40, 50, 60, 70, 80], dtype=torch.int32
|
|
) # 8 tokens, padded
|
|
|
|
fake_metadata = AscendSFAMetadata(
|
|
num_actual_tokens=2,
|
|
slot_mapping=torch.zeros(2, dtype=torch.int32),
|
|
seq_lens=torch.tensor([4, 4], dtype=torch.int32),
|
|
seq_lens_cpu=torch.tensor([4, 4], dtype=torch.int32),
|
|
cum_query_lens=torch.tensor([1, 2], dtype=torch.int32),
|
|
block_table=torch.zeros((2, 4), dtype=torch.int32),
|
|
sin=torch.randn(2, 32),
|
|
cos=torch.randn(2, 32),
|
|
num_input_tokens=2,
|
|
attn_state=AscendAttentionState.DecodeOnly,
|
|
)
|
|
with patch.object(
|
|
AscendSFACPMetadataBuilder.__bases__[0],
|
|
"build",
|
|
return_value=fake_metadata,
|
|
):
|
|
result = builder.build(common_prefix_len=0, common_attn_metadata=common_attn_metadata)
|
|
# The first num_decode_tokens slot mappings are taken at every pcp_size stride
|
|
self.assertEqual(result.num_decodes, 2)
|
|
self.assertEqual(result.slot_mapping.shape[0], 4)
|
|
|
|
@patch("vllm_ascend.attention.context_parallel.sfa_cp.enabling_mlapo", return_value=False)
|
|
@patch("vllm_ascend.attention.context_parallel.sfa_cp.split_decodes_and_prefills")
|
|
@patch_distributed_groups(dcp_size=2, pcp_size=2, needs_mocks=False)
|
|
def test_build_with_speculative_and_pcp(self, mock_split, mock_mlapo):
|
|
# Tests speculative_config branch (compact_varlen_decode_slot_mapping)
|
|
builder = self._build_builder(pcp_size=2, dcp_size=2)
|
|
builder.speculative_config = MagicMock() # Truthy speculative_config
|
|
# 2 decodes, 0 prefills, num_decode_tokens = 3 (varlen)
|
|
mock_split.return_value = (2, 0, 3, 0)
|
|
|
|
common_attn_metadata = MagicMock()
|
|
common_attn_metadata.num_reqs = 2
|
|
common_attn_metadata.num_actual_tokens = 3
|
|
common_attn_metadata.num_input_tokens = 3
|
|
|
|
long_seq_metadata = MagicMock()
|
|
long_seq_metadata.q_head_idx_tensor = torch.tensor([0])
|
|
long_seq_metadata.q_tail_idx_tensor = torch.tensor([1])
|
|
long_seq_metadata.q_full_idx = torch.tensor([0, 1])
|
|
long_seq_metadata.pcp_allgather_restore_idx = torch.tensor([0, 1])
|
|
long_seq_metadata.num_actual_tokens_pcp_padded = 6
|
|
long_seq_metadata.query_lens_pcp_full_cpu = torch.tensor([2, 1], dtype=torch.int64)
|
|
common_attn_metadata.prefill_context_parallel_metadata = long_seq_metadata
|
|
common_attn_metadata.num_computed_tokens_cpu = torch.tensor([0, 0], dtype=torch.int32)
|
|
common_attn_metadata.slot_mapping = torch.arange(20, dtype=torch.int32)
|
|
|
|
fake_metadata = AscendSFAMetadata(
|
|
num_actual_tokens=3,
|
|
slot_mapping=torch.zeros(3, dtype=torch.int32),
|
|
seq_lens=torch.tensor([4, 4], dtype=torch.int32),
|
|
seq_lens_cpu=torch.tensor([4, 4], dtype=torch.int32),
|
|
cum_query_lens=torch.tensor([2, 3], dtype=torch.int32),
|
|
block_table=torch.zeros((2, 4), dtype=torch.int32),
|
|
sin=torch.randn(3, 32),
|
|
cos=torch.randn(3, 32),
|
|
num_input_tokens=3,
|
|
attn_state=AscendAttentionState.SpecDecoding,
|
|
)
|
|
with patch.object(
|
|
AscendSFACPMetadataBuilder.__bases__[0],
|
|
"build",
|
|
return_value=fake_metadata,
|
|
):
|
|
result = builder.build(common_prefix_len=0, common_attn_metadata=common_attn_metadata)
|
|
self.assertEqual(result.num_decodes, 2)
|
|
self.assertIsNotNone(result.slot_mapping)
|
|
|
|
|
|
class TestAscendSFACPImpl(TestBase):
|
|
"""Tests for AscendSFACPImpl."""
|
|
|
|
@patch("vllm_ascend.attention.context_parallel.sfa_cp.enabling_mlapo", return_value=False)
|
|
@patch("vllm.distributed.parallel_state._TP", new_callable=lambda: MagicMock(spec=GroupCoordinator))
|
|
@patch("vllm_ascend.attention.sfa_v1.enable_dsa_cp_with_o_proj_tp", return_value=False)
|
|
@patch("vllm_ascend.attention.sfa_v1.enable_dsa_cp_with_layer_shard", return_value=False)
|
|
@patch("vllm_ascend.attention.sfa_v1.enable_dsa_cp", return_value=False)
|
|
@patch("vllm_ascend.attention.sfa_v1.get_current_vllm_config")
|
|
@patch_distributed_groups(dcp_size=2, pcp_size=2, needs_mocks=False)
|
|
def setUp(
|
|
self,
|
|
mock_get_current_vllm_config,
|
|
_mock_enable_dsa_cp,
|
|
_mock_enable_dsa_cp_with_layer_shard,
|
|
_mock_enable_dsa_cp_with_o_proj_tp,
|
|
mock_tp,
|
|
_mock_enabling_mlapo,
|
|
):
|
|
mock_tp.world_size = 2
|
|
mock_tp.rank_in_group = MagicMock()
|
|
mock_tp.device_group = MagicMock()
|
|
|
|
vllm_config = MagicMock()
|
|
speculative_config = MagicMock()
|
|
model_config = MagicMock()
|
|
parallel_config = MagicMock()
|
|
parallel_config.prefill_context_parallel_size = 2
|
|
parallel_config.decode_context_parallel_size = 2
|
|
parallel_config.tensor_parallel_size = 2
|
|
speculative_config.num_speculative_tokens = 0
|
|
vllm_config.speculative_config = speculative_config
|
|
model_config.dtype = torch.float16
|
|
model_config.hf_config.model_type = "deepseek_v3"
|
|
vllm_config.model_config = model_config
|
|
vllm_config.kv_transfer_config = None
|
|
vllm_config.additional_config = {"refresh": True}
|
|
vllm_config.parallel_config = parallel_config
|
|
mock_get_current_vllm_config.return_value = vllm_config
|
|
init_ascend_config(vllm_config)
|
|
|
|
self.kwargs = _make_impl_kwargs()
|
|
self.impl = AscendSFACPImpl(
|
|
num_heads=256,
|
|
head_size=1024,
|
|
scale=0.1,
|
|
num_kv_heads=8,
|
|
alibi_slopes=None,
|
|
sliding_window=None,
|
|
kv_cache_dtype="auto",
|
|
logits_soft_cap=None,
|
|
attn_type=None,
|
|
kv_sharing_target_layer_name=None,
|
|
**self.kwargs,
|
|
)
|
|
AscendSFAImpl.o_proj_full_pool = None
|
|
AscendSFAImpl.q_hadamard = None
|
|
AscendSFAImpl.k_hadamard = None
|
|
|
|
def test_init_default(self):
|
|
self.assertEqual(self.impl.pcp_size, 2)
|
|
self.assertEqual(self.impl.dcp_size, 2)
|
|
self.assertEqual(self.impl.pcp_rank, 0)
|
|
self.assertEqual(self.impl.dcp_rank, 0)
|
|
self.assertIsNotNone(self.impl.pcp_group)
|
|
self.assertIsNotNone(self.impl.dcp_group)
|
|
self.assertFalse(self.impl.enable_mlapo)
|
|
|
|
@patch("vllm_ascend.attention.context_parallel.sfa_cp.enabling_mlapo", return_value=False)
|
|
@patch("vllm.distributed.parallel_state._TP", new_callable=lambda: MagicMock(spec=GroupCoordinator))
|
|
@patch("vllm_ascend.attention.sfa_v1.enable_dsa_cp_with_o_proj_tp", return_value=False)
|
|
@patch("vllm_ascend.attention.sfa_v1.enable_dsa_cp_with_layer_shard", return_value=False)
|
|
@patch("vllm_ascend.attention.sfa_v1.enable_dsa_cp", return_value=False)
|
|
@patch("vllm_ascend.attention.sfa_v1.get_current_vllm_config")
|
|
@patch_distributed_groups(dcp_size=1, pcp_size=1, needs_mocks=False)
|
|
def test_init_no_cp(
|
|
self,
|
|
mock_get_current_vllm_config,
|
|
_e_dsa,
|
|
_e_layer_shard,
|
|
_e_o_proj_tp,
|
|
mock_tp,
|
|
_e_mlapo,
|
|
):
|
|
mock_tp.world_size = 1
|
|
mock_tp.rank_in_group = MagicMock()
|
|
vllm_config = MagicMock()
|
|
speculative_config = MagicMock()
|
|
speculative_config.num_speculative_tokens = 0
|
|
vllm_config.speculative_config = speculative_config
|
|
vllm_config.model_config.dtype = torch.float16
|
|
vllm_config.model_config.hf_config.model_type = "deepseek_v3"
|
|
vllm_config.kv_transfer_config = None
|
|
vllm_config.additional_config = {"refresh": True}
|
|
parallel_config = MagicMock()
|
|
parallel_config.prefill_context_parallel_size = 1
|
|
parallel_config.decode_context_parallel_size = 1
|
|
parallel_config.tensor_parallel_size = 1
|
|
vllm_config.parallel_config = parallel_config
|
|
mock_get_current_vllm_config.return_value = vllm_config
|
|
init_ascend_config(vllm_config)
|
|
|
|
impl = AscendSFACPImpl(
|
|
num_heads=4,
|
|
head_size=128,
|
|
scale=0.1,
|
|
num_kv_heads=2,
|
|
alibi_slopes=None,
|
|
sliding_window=None,
|
|
kv_cache_dtype="auto",
|
|
logits_soft_cap=None,
|
|
attn_type=None,
|
|
kv_sharing_target_layer_name=None,
|
|
**_make_impl_kwargs(),
|
|
)
|
|
self.assertEqual(impl.pcp_size, 1)
|
|
self.assertEqual(impl.dcp_size, 1)
|
|
self.assertEqual(impl.pcp_rank, 0)
|
|
self.assertEqual(impl.dcp_rank, 0)
|
|
self.assertIsNone(impl.pcp_group)
|
|
self.assertIsNone(impl.dcp_group)
|
|
|
|
def test_align_to_graph_bucket_tokens_none_input(self):
|
|
self.impl.pcp_size = 2
|
|
result = self.impl._align_to_graph_bucket_tokens(None, MagicMock())
|
|
self.assertIsNone(result)
|
|
|
|
def test_align_to_graph_bucket_tokens_no_pcp(self):
|
|
self.impl.pcp_size = 1
|
|
attn_output = torch.randn(4, 8)
|
|
result = self.impl._align_to_graph_bucket_tokens(attn_output, MagicMock())
|
|
self.assertIs(result, attn_output)
|
|
|
|
@patch("vllm_ascend.attention.context_parallel.sfa_cp.get_forward_context")
|
|
def test_align_to_graph_bucket_tokens_already_aligned(self, mock_get_fc):
|
|
self.impl.pcp_size = 2
|
|
forward_context = MagicMock()
|
|
forward_context.num_tokens = 8
|
|
mock_get_fc.return_value = forward_context
|
|
|
|
attn_metadata = MagicMock()
|
|
attn_metadata.num_input_tokens = 8
|
|
|
|
attn_output = torch.randn(8, 16)
|
|
result = self.impl._align_to_graph_bucket_tokens(attn_output, attn_metadata)
|
|
# Already aligned, returns same tensor
|
|
self.assertIs(result, attn_output)
|
|
|
|
@patch("vllm_ascend.attention.context_parallel.sfa_cp.get_forward_context")
|
|
def test_align_to_graph_bucket_tokens_pad_smaller(self, mock_get_fc):
|
|
self.impl.pcp_size = 2
|
|
forward_context = MagicMock()
|
|
forward_context.num_tokens = 16
|
|
mock_get_fc.return_value = forward_context
|
|
|
|
attn_metadata = MagicMock()
|
|
attn_metadata.num_input_tokens = 8
|
|
|
|
attn_output = torch.randn(8, 16)
|
|
result = self.impl._align_to_graph_bucket_tokens(attn_output, attn_metadata)
|
|
self.assertEqual(result.shape, (16, 16))
|
|
# First 8 rows match input
|
|
self.assertTrue(torch.equal(result[:8], attn_output))
|
|
# Padded rows are zeros
|
|
self.assertTrue(torch.all(result[8:] == 0))
|
|
|
|
@patch("vllm_ascend.attention.context_parallel.sfa_cp.get_forward_context")
|
|
def test_align_to_graph_bucket_tokens_truncate(self, mock_get_fc):
|
|
# Edge: target is smaller than attn output (rare; valid_tokens = min)
|
|
self.impl.pcp_size = 2
|
|
forward_context = MagicMock()
|
|
forward_context.num_tokens = 4
|
|
mock_get_fc.return_value = forward_context
|
|
|
|
attn_metadata = MagicMock()
|
|
attn_metadata.num_input_tokens = 4
|
|
|
|
attn_output = torch.randn(8, 16)
|
|
result = self.impl._align_to_graph_bucket_tokens(attn_output, attn_metadata)
|
|
self.assertEqual(result.shape, (4, 16))
|
|
|
|
@patch("vllm_ascend.attention.context_parallel.sfa_cp.get_forward_context")
|
|
def test_align_to_graph_bucket_tokens_no_forward_context(self, mock_get_fc):
|
|
self.impl.pcp_size = 2
|
|
mock_get_fc.return_value = None
|
|
|
|
attn_metadata = MagicMock()
|
|
attn_metadata.num_input_tokens = 16
|
|
|
|
attn_output = torch.randn(8, 16)
|
|
result = self.impl._align_to_graph_bucket_tokens(attn_output, attn_metadata)
|
|
self.assertEqual(result.shape, (16, 16))
|
|
|
|
def test_execute_sparse_flash_attention(self):
|
|
ql_nope = torch.randn(2, 4, 32)
|
|
q_pe = torch.randn(2, 4, 16)
|
|
kv = torch.randn(2, 4, 1, 32)
|
|
key_rope = torch.randn(2, 4, 1, 16)
|
|
block_table = torch.tensor([[0]], dtype=torch.int32)
|
|
topk_indices = torch.tensor([[0]], dtype=torch.int32)
|
|
actual_seq_lengths_query = torch.tensor([1, 2], dtype=torch.int32)
|
|
actual_seq_lengths_key = torch.tensor([1, 2], dtype=torch.int32)
|
|
|
|
with patch.object(
|
|
torch.ops._C_ascend,
|
|
"npu_sparse_flash_attention",
|
|
create=True,
|
|
return_value=(torch.randn(2, 4, 32), None, None),
|
|
) as mock_sfa:
|
|
result = self.impl._execute_sparse_flash_attention(
|
|
ql_nope, q_pe, kv, key_rope, block_table, topk_indices, actual_seq_lengths_query, actual_seq_lengths_key
|
|
)
|
|
self.assertIsNotNone(result)
|
|
mock_sfa.assert_called_once()
|
|
|
|
@patch_distributed_groups(dcp_size=2, pcp_size=2, needs_mocks=False)
|
|
def test_gather_kv_cross_cp(self):
|
|
self.impl.pcp_size = 2
|
|
self.impl.dcp_size = 2
|
|
kv_cache = torch.randn(8, 4, 1, 16)
|
|
block_tables = torch.tensor([[0, 1], [2, 3]], dtype=torch.int32)
|
|
|
|
result, block_num = self.impl.gather_kv_cross_cp(kv_cache, block_tables)
|
|
# block_num is num blocks selected before all_gather
|
|
self.assertEqual(block_num, 4)
|
|
# After both pcp and dcp all_gather, total blocks = 4 * 2 * 2 = 16
|
|
self.assertEqual(result.shape[0], 16)
|
|
|
|
@patch_distributed_groups(dcp_size=1, pcp_size=1, needs_mocks=False)
|
|
def test_gather_kv_cross_cp_no_cp(self):
|
|
self.impl.pcp_size = 1
|
|
self.impl.dcp_size = 1
|
|
kv_cache = torch.randn(8, 4, 1, 16)
|
|
block_tables = torch.tensor([[0, 1]], dtype=torch.int32)
|
|
|
|
result, block_num = self.impl.gather_kv_cross_cp(kv_cache, block_tables)
|
|
self.assertEqual(block_num, 2)
|
|
self.assertEqual(result.shape[0], 2)
|
|
|
|
@patch_distributed_groups(dcp_size=2, pcp_size=2, needs_mocks=False)
|
|
def test_gather_kv_cross_cp_compact(self):
|
|
self.impl.pcp_size = 2
|
|
self.impl.dcp_size = 2
|
|
kv_cache = torch.randn(8, 4, 1, 16)
|
|
valid_block_ids = torch.tensor([0, 2, 4], dtype=torch.int64)
|
|
|
|
result = self.impl.gather_kv_cross_cp_compact(kv_cache, valid_block_ids)
|
|
# 3 blocks * 2 (dcp) * 2 (pcp) = 12
|
|
self.assertEqual(result.shape[0], 12)
|
|
|
|
@patch_distributed_groups(dcp_size=1, pcp_size=1, needs_mocks=False)
|
|
def test_gather_kv_cross_cp_compact_no_cp(self):
|
|
self.impl.pcp_size = 1
|
|
self.impl.dcp_size = 1
|
|
kv_cache = torch.randn(8, 4, 1, 16)
|
|
valid_block_ids = torch.tensor([0, 2, 4], dtype=torch.int64)
|
|
|
|
result = self.impl.gather_kv_cross_cp_compact(kv_cache, valid_block_ids)
|
|
self.assertEqual(result.shape[0], 3)
|
|
|
|
def test_gather_block_table(self):
|
|
block_num = 4
|
|
block_tables = torch.tensor([[0, 1], [2, 3]], dtype=torch.int32)
|
|
block_arange = torch.arange(4, dtype=torch.int32)
|
|
|
|
result = self.impl.gather_block_table(block_num, block_tables, block_arange)
|
|
# Shape: (num_reqs, num_blocks_per_req * pcp*dcp) = (2, 2*4=8)
|
|
self.assertEqual(result.shape, (2, 8))
|
|
self.assertEqual(result.dtype, block_tables.dtype)
|
|
|
|
def test_execute_indexer_select_torch_npu(self):
|
|
self.impl.use_torch_npu_lightning_indexer = True
|
|
q = torch.randn(2, 64, 128)
|
|
key = torch.randn(2, 1, 1, 128)
|
|
weights = torch.randn(2, 64)
|
|
actual_seq_lengths_query = torch.tensor([1, 2])
|
|
actual_seq_lengths_key = torch.tensor([1, 2])
|
|
block_table = torch.tensor([[0]], dtype=torch.int32)
|
|
|
|
with patch("vllm_ascend.attention.context_parallel.sfa_cp.torch_npu") as mock_torch_npu:
|
|
mock_torch_npu.npu_lightning_indexer.return_value = (torch.tensor([[0]]), None)
|
|
result = self.impl._execute_indexer_select(
|
|
q, key, weights, actual_seq_lengths_query, actual_seq_lengths_key, block_table
|
|
)
|
|
self.assertIsNotNone(result)
|
|
|
|
def test_execute_indexer_select_ascend_op(self):
|
|
self.impl.use_torch_npu_lightning_indexer = False
|
|
q = torch.randn(2, 64, 128)
|
|
key = torch.randn(2, 1, 1, 128)
|
|
weights = torch.randn(2, 64)
|
|
actual_seq_lengths_query = torch.tensor([1, 2])
|
|
actual_seq_lengths_key = torch.tensor([1, 2])
|
|
block_table = torch.tensor([[0]], dtype=torch.int32)
|
|
|
|
with patch.object(
|
|
torch.ops._C_ascend,
|
|
"npu_lightning_indexer",
|
|
create=True,
|
|
return_value=(torch.tensor([[0]]), None),
|
|
) as mock_indexer:
|
|
result = self.impl._execute_indexer_select(
|
|
q, key, weights, actual_seq_lengths_query, actual_seq_lengths_key, block_table
|
|
)
|
|
self.assertIsNotNone(result)
|
|
mock_indexer.assert_called_once()
|
|
|
|
def test_get_full_kv_no_pcp(self):
|
|
self.impl.pcp_size = 1
|
|
k = torch.randn(4, 8, 16)
|
|
result = self.impl._get_full_kv(k, MagicMock())
|
|
self.assertIs(result, k)
|
|
|
|
def test_get_full_kv_mlapo(self):
|
|
self.impl.pcp_size = 2
|
|
self.impl.enable_mlapo = True
|
|
k = torch.randn(4, 8, 16)
|
|
result = self.impl._get_full_kv(k, MagicMock())
|
|
self.assertIs(result, k)
|
|
|
|
@patch_distributed_groups(dcp_size=1, pcp_size=2, needs_mocks=False)
|
|
def test_get_full_kv_with_pcp(self):
|
|
self.impl.pcp_size = 2
|
|
self.impl.enable_mlapo = False
|
|
k = torch.randn(4, 8, 16)
|
|
attn_metadata = MagicMock()
|
|
sfa_cp_metadata = MagicMock()
|
|
sfa_cp_metadata.pcp_allgather_restore_idx = torch.arange(8)
|
|
attn_metadata.sfa_cp_metadata = sfa_cp_metadata
|
|
|
|
result = self.impl._get_full_kv(k, attn_metadata)
|
|
# After all_gather pcp_size=2 -> 8 entries, then index_select with 8 indices
|
|
self.assertEqual(result.shape[0], 8)
|
|
|
|
@patch("vllm_ascend.attention.context_parallel.sfa_cp.torch_npu")
|
|
@patch_distributed_groups(dcp_size=1, pcp_size=1, needs_mocks=False)
|
|
def test_exec_kv_no_pcp(self, mock_torch_npu):
|
|
# When pcp_size==1, simply delegates to super().exec_kv
|
|
self.impl.pcp_size = 1
|
|
with patch.object(AscendSFAImpl, "exec_kv", return_value=("a", "b")) as mock_super:
|
|
result = self.impl.exec_kv(
|
|
kv_no_split=torch.randn(2, 64),
|
|
cos=torch.randn(2, 32),
|
|
sin=torch.randn(2, 32),
|
|
kv_cache=(torch.randn(4, 1, 1, 32), torch.randn(4, 1, 1, 32)),
|
|
slots=torch.tensor([0, 1], dtype=torch.int32),
|
|
attn_metadata=MagicMock(),
|
|
)
|
|
mock_super.assert_called_once()
|
|
self.assertEqual(result, ("a", "b"))
|
|
|
|
@patch("vllm_ascend.attention.context_parallel.sfa_cp.torch_npu")
|
|
@patch_distributed_groups(dcp_size=1, pcp_size=2, needs_mocks=False)
|
|
def test_exec_kv_with_pcp(self, mock_torch_npu):
|
|
self.impl.pcp_size = 2
|
|
# Configure dimensions
|
|
self.impl.kv_lora_rank = 32
|
|
self.impl.qk_rope_head_dim = 16
|
|
self.impl.num_kv_heads = 1
|
|
|
|
kv_a_layernorm = MagicMock()
|
|
kv_a_layernorm.side_effect = lambda x: x
|
|
self.impl.kv_a_layernorm = kv_a_layernorm
|
|
self.impl.rope_single = MagicMock(side_effect=lambda x, cos, sin: x)
|
|
|
|
# 2 input tokens, [num_tokens, kv_lora_rank + qk_rope_head_dim]
|
|
kv_no_split = torch.randn(2, 32 + 16)
|
|
cos = torch.randn(2, 16)
|
|
sin = torch.randn(2, 16)
|
|
kv_cache = (torch.randn(4, 1, 1, 32), torch.randn(4, 1, 1, 16))
|
|
slots = torch.tensor([0, 1, 2, 3], dtype=torch.int32)
|
|
|
|
attn_metadata = MagicMock()
|
|
sfa_cp_metadata = MagicMock()
|
|
sfa_cp_metadata.pcp_allgather_restore_idx = torch.arange(4)
|
|
attn_metadata.sfa_cp_metadata = sfa_cp_metadata
|
|
attn_metadata.slot_mapping = slots
|
|
|
|
result = self.impl.exec_kv(kv_no_split, cos, sin, kv_cache, slots, attn_metadata)
|
|
self.assertEqual(result, (None, None))
|
|
mock_torch_npu._npu_reshape_and_cache.assert_called_once()
|
|
|
|
@patch_distributed_groups(dcp_size=1, pcp_size=1, needs_mocks=False)
|
|
def test_execute_sparse_flash_attention_process_decode_only(self):
|
|
# num_prefills < 1: returns aligned decode output
|
|
self.impl.pcp_size = 1
|
|
self.impl.dcp_size = 1
|
|
ql_nope = torch.randn(2, 4, 32)
|
|
q_pe = torch.randn(2, 4, 16)
|
|
kv_cache = (
|
|
torch.randn(4, 1, 1, 32),
|
|
torch.randn(4, 1, 1, 16),
|
|
torch.randn(4, 1, 1, 32),
|
|
)
|
|
topk_indices = torch.tensor([[0], [0]], dtype=torch.int32)
|
|
attn_metadata = MagicMock()
|
|
attn_metadata.num_decodes = 2
|
|
attn_metadata.num_decode_tokens = 2
|
|
attn_metadata.num_prefills = 0
|
|
attn_metadata.block_table = torch.tensor([[0], [1]], dtype=torch.int32)
|
|
sfa_cp_metadata = MagicMock()
|
|
sfa_cp_metadata.block_arange = torch.tensor([0, 1, 2, 3], dtype=torch.int32)
|
|
attn_metadata.sfa_cp_metadata = sfa_cp_metadata
|
|
|
|
actual_seq_lengths_query = torch.tensor([1, 2], dtype=torch.int32)
|
|
actual_seq_lengths_key = torch.tensor([1, 2], dtype=torch.int32)
|
|
|
|
with patch.object(
|
|
torch.ops._C_ascend,
|
|
"npu_sparse_flash_attention",
|
|
create=True,
|
|
return_value=(torch.randn(2, 4, 32), None, None),
|
|
):
|
|
result = self.impl._execute_sparse_flash_attention_process(
|
|
ql_nope, q_pe, kv_cache, topk_indices, attn_metadata, actual_seq_lengths_query, actual_seq_lengths_key
|
|
)
|
|
self.assertIsNotNone(result)
|
|
|
|
@patch_distributed_groups(dcp_size=1, pcp_size=1, needs_mocks=False)
|
|
def test_execute_sparse_flash_attention_process_prefill_only_no_pcp(self):
|
|
# Case: only prefills, pcp_size==1
|
|
self.impl.pcp_size = 1
|
|
self.impl.dcp_size = 1
|
|
ql_nope = torch.randn(4, 4, 32)
|
|
q_pe = torch.randn(4, 4, 16)
|
|
kv_cache = (
|
|
torch.randn(4, 1, 1, 32),
|
|
torch.randn(4, 1, 1, 16),
|
|
torch.randn(4, 1, 1, 32),
|
|
)
|
|
topk_indices = torch.tensor([[0]] * 4, dtype=torch.int32)
|
|
attn_metadata = MagicMock()
|
|
attn_metadata.num_decodes = 0
|
|
attn_metadata.num_decode_tokens = 0
|
|
attn_metadata.num_prefills = 2
|
|
attn_metadata.block_table = torch.tensor([[0, 1], [2, 3]], dtype=torch.int32)
|
|
|
|
sfa_cp_metadata = MagicMock()
|
|
sfa_cp_metadata.valid_block_ids = torch.tensor([0, 1, 2, 3], dtype=torch.int64)
|
|
sfa_cp_metadata.block_table_cp = torch.tensor([[0, 1], [2, 3]], dtype=torch.int32)
|
|
sfa_cp_metadata.prefill_q_cum_seqlens = torch.tensor([2, 4], dtype=torch.int32)
|
|
sfa_cp_metadata.block_arange = torch.tensor([0], dtype=torch.int32)
|
|
attn_metadata.sfa_cp_metadata = sfa_cp_metadata
|
|
|
|
actual_seq_lengths_query = torch.tensor([2, 4], dtype=torch.int32)
|
|
actual_seq_lengths_key = torch.tensor([4, 8], dtype=torch.int32)
|
|
|
|
with patch.object(
|
|
torch.ops._C_ascend,
|
|
"npu_sparse_flash_attention",
|
|
create=True,
|
|
return_value=(torch.randn(4, 4, 32), None, None),
|
|
):
|
|
result = self.impl._execute_sparse_flash_attention_process(
|
|
ql_nope, q_pe, kv_cache, topk_indices, attn_metadata, actual_seq_lengths_query, actual_seq_lengths_key
|
|
)
|
|
self.assertIsNotNone(result)
|
|
self.assertEqual(result.shape[0], 4)
|
|
|
|
@patch_distributed_groups(dcp_size=2, pcp_size=2, needs_mocks=False)
|
|
def test_execute_sparse_flash_attention_process_prefill_with_pcp(self):
|
|
self.impl.pcp_size = 2
|
|
self.impl.dcp_size = 2
|
|
ql_nope = torch.randn(4, 4, 32)
|
|
q_pe = torch.randn(4, 4, 16)
|
|
kv_cache = (
|
|
torch.randn(4, 1, 1, 32),
|
|
torch.randn(4, 1, 1, 16),
|
|
torch.randn(4, 1, 1, 32),
|
|
)
|
|
topk_indices = torch.tensor([[0]] * 4, dtype=torch.int32)
|
|
attn_metadata = MagicMock()
|
|
attn_metadata.num_decodes = 0
|
|
attn_metadata.num_decode_tokens = 0
|
|
attn_metadata.num_prefills = 2
|
|
attn_metadata.num_input_tokens = 4
|
|
attn_metadata.block_table = torch.tensor([[0, 1], [2, 3]], dtype=torch.int32)
|
|
|
|
sfa_cp_metadata = MagicMock()
|
|
sfa_cp_metadata.valid_block_ids = torch.tensor([0, 1, 2, 3], dtype=torch.int64)
|
|
sfa_cp_metadata.block_table_cp = torch.tensor([[0, 1], [2, 3]], dtype=torch.int32)
|
|
sfa_cp_metadata.prefill_q_cum_seqlens = torch.tensor([2, 4], dtype=torch.int32)
|
|
sfa_cp_metadata.q_head_idx = torch.tensor([0, 1], dtype=torch.int64)
|
|
sfa_cp_metadata.q_tail_idx = torch.tensor([2, 3], dtype=torch.int64)
|
|
sfa_cp_metadata.q_full_idx = torch.tensor([0, 1, 2, 3], dtype=torch.int64)
|
|
sfa_cp_metadata.head_attn_nomask_seqlens = torch.tensor([4, 4], dtype=torch.int32)
|
|
sfa_cp_metadata.tail_attn_nomask_seqlens = torch.tensor([8, 8], dtype=torch.int32)
|
|
sfa_cp_metadata.block_arange = torch.tensor([0, 1, 2, 3], dtype=torch.int32)
|
|
attn_metadata.sfa_cp_metadata = sfa_cp_metadata
|
|
|
|
actual_seq_lengths_query = torch.tensor([2, 4], dtype=torch.int32)
|
|
actual_seq_lengths_key = torch.tensor([4, 8], dtype=torch.int32)
|
|
|
|
with (
|
|
patch.object(
|
|
torch.ops._C_ascend,
|
|
"npu_sparse_flash_attention",
|
|
create=True,
|
|
return_value=(torch.randn(2, 4, 32), None, None),
|
|
),
|
|
patch("vllm_ascend.attention.context_parallel.sfa_cp.get_forward_context") as mock_fc,
|
|
):
|
|
mock_fc.return_value = MagicMock(num_tokens=4)
|
|
result = self.impl._execute_sparse_flash_attention_process(
|
|
ql_nope, q_pe, kv_cache, topk_indices, attn_metadata, actual_seq_lengths_query, actual_seq_lengths_key
|
|
)
|
|
self.assertIsNotNone(result)
|
|
|
|
@patch_distributed_groups(dcp_size=2, pcp_size=2, needs_mocks=False)
|
|
def test_execute_sparse_flash_attention_process_decode_and_prefill_with_pcp(self):
|
|
# Covers final torch.cat([decode_attn_out, attn_output]) (line 326)
|
|
self.impl.pcp_size = 2
|
|
self.impl.dcp_size = 2
|
|
|
|
ql_nope = torch.randn(5, 4, 32)
|
|
q_pe = torch.randn(5, 4, 16)
|
|
kv_cache = (
|
|
torch.randn(4, 1, 1, 32),
|
|
torch.randn(4, 1, 1, 16),
|
|
torch.randn(4, 1, 1, 32),
|
|
)
|
|
topk_indices = torch.tensor([[0]] * 5, dtype=torch.int32)
|
|
attn_metadata = MagicMock()
|
|
attn_metadata.num_decodes = 1
|
|
attn_metadata.num_decode_tokens = 1
|
|
attn_metadata.num_prefills = 2
|
|
attn_metadata.num_input_tokens = 5
|
|
attn_metadata.block_table = torch.tensor([[0], [1], [2]], dtype=torch.int32)
|
|
|
|
sfa_cp_metadata = MagicMock()
|
|
sfa_cp_metadata.valid_block_ids = torch.tensor([0, 1, 2, 3], dtype=torch.int64)
|
|
sfa_cp_metadata.block_table_cp = torch.tensor([[0, 1], [2, 3]], dtype=torch.int32)
|
|
sfa_cp_metadata.prefill_q_cum_seqlens = torch.tensor([2, 4], dtype=torch.int32)
|
|
sfa_cp_metadata.q_head_idx = torch.tensor([0, 1], dtype=torch.int64)
|
|
sfa_cp_metadata.q_tail_idx = torch.tensor([2, 3], dtype=torch.int64)
|
|
sfa_cp_metadata.q_full_idx = torch.tensor([0, 1, 2, 3], dtype=torch.int64)
|
|
sfa_cp_metadata.head_attn_nomask_seqlens = torch.tensor([4, 4, 4], dtype=torch.int32)
|
|
sfa_cp_metadata.tail_attn_nomask_seqlens = torch.tensor([8, 8, 8], dtype=torch.int32)
|
|
sfa_cp_metadata.block_arange = torch.tensor([0, 1, 2, 3], dtype=torch.int32)
|
|
attn_metadata.sfa_cp_metadata = sfa_cp_metadata
|
|
|
|
actual_seq_lengths_query = torch.tensor([1, 3, 5], dtype=torch.int32)
|
|
actual_seq_lengths_key = torch.tensor([4, 8, 8], dtype=torch.int32)
|
|
|
|
def fake_sfa(query, **kwargs):
|
|
return torch.randn(query.shape[0], query.shape[1], query.shape[2]), None, None
|
|
|
|
with (
|
|
patch.object(
|
|
torch.ops._C_ascend,
|
|
"npu_sparse_flash_attention",
|
|
create=True,
|
|
side_effect=fake_sfa,
|
|
),
|
|
patch("vllm_ascend.attention.context_parallel.sfa_cp.get_forward_context") as mock_fc,
|
|
):
|
|
mock_fc.return_value = MagicMock(num_tokens=5)
|
|
result = self.impl._execute_sparse_flash_attention_process(
|
|
ql_nope, q_pe, kv_cache, topk_indices, attn_metadata, actual_seq_lengths_query, actual_seq_lengths_key
|
|
)
|
|
self.assertIsNotNone(result)
|
|
self.assertEqual(result.shape[0], 5)
|
|
|
|
@patch_distributed_groups(dcp_size=1, pcp_size=1, needs_mocks=False)
|
|
def test_execute_sparse_flash_attention_process_decode_and_prefill_no_pcp(self):
|
|
self.impl.pcp_size = 1
|
|
self.impl.dcp_size = 1
|
|
ql_nope = torch.randn(3, 4, 32)
|
|
q_pe = torch.randn(3, 4, 16)
|
|
kv_cache = (
|
|
torch.randn(4, 1, 1, 32),
|
|
torch.randn(4, 1, 1, 16),
|
|
torch.randn(4, 1, 1, 32),
|
|
)
|
|
topk_indices = torch.tensor([[0]] * 3, dtype=torch.int32)
|
|
attn_metadata = MagicMock()
|
|
attn_metadata.num_decodes = 1
|
|
attn_metadata.num_decode_tokens = 1
|
|
attn_metadata.num_prefills = 1
|
|
attn_metadata.block_table = torch.tensor([[0], [1]], dtype=torch.int32)
|
|
|
|
sfa_cp_metadata = MagicMock()
|
|
sfa_cp_metadata.valid_block_ids = torch.tensor([0, 1], dtype=torch.int64)
|
|
sfa_cp_metadata.block_table_cp = torch.tensor([[0, 1]], dtype=torch.int32)
|
|
sfa_cp_metadata.prefill_q_cum_seqlens = torch.tensor([2], dtype=torch.int32)
|
|
sfa_cp_metadata.block_arange = torch.tensor([0], dtype=torch.int32)
|
|
attn_metadata.sfa_cp_metadata = sfa_cp_metadata
|
|
|
|
actual_seq_lengths_query = torch.tensor([1, 3], dtype=torch.int32)
|
|
actual_seq_lengths_key = torch.tensor([4, 8], dtype=torch.int32)
|
|
|
|
# Use side_effect so each call returns attention_out with the q-shape.
|
|
def fake_sfa(query, **kwargs):
|
|
return torch.randn(query.shape[0], query.shape[1], query.shape[2]), None, None
|
|
|
|
with patch.object(
|
|
torch.ops._C_ascend,
|
|
"npu_sparse_flash_attention",
|
|
create=True,
|
|
side_effect=fake_sfa,
|
|
):
|
|
result = self.impl._execute_sparse_flash_attention_process(
|
|
ql_nope, q_pe, kv_cache, topk_indices, attn_metadata, actual_seq_lengths_query, actual_seq_lengths_key
|
|
)
|
|
self.assertIsNotNone(result)
|
|
self.assertEqual(result.shape[0], 3)
|
|
|
|
@patch("vllm_ascend.attention.context_parallel.sfa_cp.HAS_TRITON", True)
|
|
@patch("vllm_ascend.attention.context_parallel.sfa_cp.rope_forward_triton_siso")
|
|
@patch_distributed_groups(dcp_size=1, pcp_size=1, needs_mocks=False)
|
|
def test_indexer_select_post_process_decode_only_simple(self, mock_rope):
|
|
# Case: num_prefills==0, returns decode_topk_indices
|
|
self.impl.pcp_size = 1
|
|
self.impl.dcp_size = 1
|
|
self.impl.use_torch_npu_lightning_indexer = False
|
|
|
|
x = torch.randn(2, self.impl.qk_head_dim)
|
|
q_c = torch.randn(2, self.impl.q_lora_rank)
|
|
kv_cache = (
|
|
torch.randn(4, 1, 1, 32),
|
|
torch.randn(4, 1, 1, 16),
|
|
torch.randn(4, 1, 1, self.impl.head_dim),
|
|
)
|
|
cos = torch.randn(2, self.impl.qk_rope_head_dim)
|
|
sin = torch.randn(2, self.impl.qk_rope_head_dim)
|
|
|
|
kw_out = torch.randn(2, self.impl.head_dim * 2)
|
|
self.impl.wk_weights_proj.return_value = (kw_out, None)
|
|
self.impl.wq_b.return_value = (
|
|
torch.randn(2, self.impl.n_head * self.impl.head_dim),
|
|
None,
|
|
)
|
|
mock_rope.return_value = torch.randn(2, self.impl.n_head, self.impl.head_dim)
|
|
|
|
attn_metadata = MagicMock()
|
|
attn_metadata.num_decodes = 2
|
|
attn_metadata.num_decode_tokens = 2
|
|
attn_metadata.num_prefills = 0
|
|
attn_metadata.block_table = torch.tensor([[0], [1]], dtype=torch.int32)
|
|
|
|
sfa_cp_metadata = MagicMock()
|
|
sfa_cp_metadata.block_arange = torch.tensor([0], dtype=torch.int32)
|
|
attn_metadata.sfa_cp_metadata = sfa_cp_metadata
|
|
|
|
actual_seq_lengths_query = torch.tensor([1, 2], dtype=torch.int32)
|
|
actual_seq_lengths_key = torch.tensor([4, 8], dtype=torch.int32)
|
|
|
|
with patch.object(
|
|
torch.ops._C_ascend,
|
|
"npu_lightning_indexer",
|
|
create=True,
|
|
return_value=(torch.tensor([[0]] * 2), None),
|
|
):
|
|
result = self.impl.indexer_select_post_process(
|
|
x, q_c, kv_cache, attn_metadata, cos, sin, actual_seq_lengths_query, actual_seq_lengths_key
|
|
)
|
|
self.assertIsNotNone(result)
|
|
|
|
@patch("vllm_ascend.attention.context_parallel.sfa_cp.HAS_TRITON", False)
|
|
@patch("vllm_ascend.attention.context_parallel.sfa_cp.torch_npu")
|
|
@patch_distributed_groups(dcp_size=1, pcp_size=1, needs_mocks=False)
|
|
def test_indexer_select_post_process_decode_only_no_triton(self, mock_torch_npu):
|
|
# Test no-triton path
|
|
self.impl.pcp_size = 1
|
|
self.impl.dcp_size = 1
|
|
self.impl.use_torch_npu_lightning_indexer = False
|
|
self.impl.is_rope_neox_style = True
|
|
|
|
x = torch.randn(2, self.impl.qk_head_dim)
|
|
q_c = torch.randn(2, self.impl.q_lora_rank)
|
|
kv_cache = (
|
|
torch.randn(4, 1, 1, 32),
|
|
torch.randn(4, 1, 1, 16),
|
|
torch.randn(4, 1, 1, self.impl.head_dim),
|
|
)
|
|
cos = torch.randn(2, self.impl.qk_rope_head_dim)
|
|
sin = torch.randn(2, self.impl.qk_rope_head_dim)
|
|
|
|
kw_out = torch.randn(2, self.impl.head_dim * 2)
|
|
self.impl.wk_weights_proj.return_value = (kw_out, None)
|
|
self.impl.wq_b.return_value = (
|
|
torch.randn(2, self.impl.n_head * self.impl.head_dim),
|
|
None,
|
|
)
|
|
mock_torch_npu.npu_rotary_mul.return_value = torch.randn(2, self.impl.n_head, 1, self.impl.qk_rope_head_dim)
|
|
|
|
attn_metadata = MagicMock()
|
|
attn_metadata.num_decodes = 2
|
|
attn_metadata.num_decode_tokens = 2
|
|
attn_metadata.num_prefills = 0
|
|
attn_metadata.block_table = torch.tensor([[0], [1]], dtype=torch.int32)
|
|
|
|
sfa_cp_metadata = MagicMock()
|
|
sfa_cp_metadata.block_arange = torch.tensor([0], dtype=torch.int32)
|
|
attn_metadata.sfa_cp_metadata = sfa_cp_metadata
|
|
|
|
actual_seq_lengths_query = torch.tensor([1, 2], dtype=torch.int32)
|
|
actual_seq_lengths_key = torch.tensor([4, 8], dtype=torch.int32)
|
|
|
|
with patch.object(
|
|
torch.ops._C_ascend,
|
|
"npu_lightning_indexer",
|
|
create=True,
|
|
return_value=(torch.tensor([[0]] * 2), None),
|
|
):
|
|
result = self.impl.indexer_select_post_process(
|
|
x, q_c, kv_cache, attn_metadata, cos, sin, actual_seq_lengths_query, actual_seq_lengths_key
|
|
)
|
|
self.assertIsNotNone(result)
|
|
|
|
@patch("vllm_ascend.attention.context_parallel.sfa_cp.HAS_TRITON", True)
|
|
@patch("vllm_ascend.attention.context_parallel.sfa_cp.rope_forward_triton_siso")
|
|
@patch_distributed_groups(dcp_size=1, pcp_size=1, needs_mocks=False)
|
|
def test_indexer_select_post_process_prefill_only_no_pcp(self, mock_rope):
|
|
# Case: only prefills, pcp_size==1
|
|
self.impl.pcp_size = 1
|
|
self.impl.dcp_size = 1
|
|
self.impl.use_torch_npu_lightning_indexer = False
|
|
|
|
x = torch.randn(2, self.impl.qk_head_dim)
|
|
q_c = torch.randn(2, self.impl.q_lora_rank)
|
|
kv_cache = (
|
|
torch.randn(4, 1, 1, 32),
|
|
torch.randn(4, 1, 1, 16),
|
|
torch.randn(4, 1, 1, self.impl.head_dim),
|
|
)
|
|
cos = torch.randn(2, self.impl.qk_rope_head_dim)
|
|
sin = torch.randn(2, self.impl.qk_rope_head_dim)
|
|
|
|
kw_out = torch.randn(2, self.impl.head_dim * 2)
|
|
self.impl.wk_weights_proj.return_value = (kw_out, None)
|
|
self.impl.wq_b.return_value = (
|
|
torch.randn(2, self.impl.n_head * self.impl.head_dim),
|
|
None,
|
|
)
|
|
mock_rope.return_value = torch.randn(2, self.impl.n_head, self.impl.head_dim)
|
|
|
|
attn_metadata = MagicMock()
|
|
attn_metadata.num_decodes = 0
|
|
attn_metadata.num_decode_tokens = 0
|
|
attn_metadata.num_prefills = 1
|
|
attn_metadata.block_table = torch.tensor([[0]], dtype=torch.int32)
|
|
|
|
sfa_cp_metadata = MagicMock()
|
|
sfa_cp_metadata.valid_block_ids = torch.tensor([0, 1], dtype=torch.int64)
|
|
sfa_cp_metadata.block_table_cp = torch.tensor([[0, 1]], dtype=torch.int32)
|
|
sfa_cp_metadata.prefill_q_cum_seqlens = torch.tensor([2], dtype=torch.int32)
|
|
sfa_cp_metadata.block_arange = torch.tensor([0], dtype=torch.int32)
|
|
attn_metadata.sfa_cp_metadata = sfa_cp_metadata
|
|
|
|
actual_seq_lengths_query = torch.tensor([2], dtype=torch.int32)
|
|
actual_seq_lengths_key = torch.tensor([4], dtype=torch.int32)
|
|
|
|
with patch.object(
|
|
torch.ops._C_ascend,
|
|
"npu_lightning_indexer",
|
|
create=True,
|
|
return_value=(torch.tensor([[0]] * 2), None),
|
|
):
|
|
result = self.impl.indexer_select_post_process(
|
|
x, q_c, kv_cache, attn_metadata, cos, sin, actual_seq_lengths_query, actual_seq_lengths_key
|
|
)
|
|
self.assertIsNotNone(result)
|
|
|
|
@patch("vllm_ascend.attention.context_parallel.sfa_cp.HAS_TRITON", True)
|
|
@patch("vllm_ascend.attention.context_parallel.sfa_cp.rope_forward_triton_siso")
|
|
@patch_distributed_groups(dcp_size=1, pcp_size=1, needs_mocks=False)
|
|
def test_indexer_select_post_process_decode_and_prefill_no_pcp(self, mock_rope):
|
|
self.impl.pcp_size = 1
|
|
self.impl.dcp_size = 1
|
|
self.impl.use_torch_npu_lightning_indexer = False
|
|
|
|
x = torch.randn(3, self.impl.qk_head_dim)
|
|
q_c = torch.randn(3, self.impl.q_lora_rank)
|
|
kv_cache = (
|
|
torch.randn(4, 1, 1, 32),
|
|
torch.randn(4, 1, 1, 16),
|
|
torch.randn(4, 1, 1, self.impl.head_dim),
|
|
)
|
|
cos = torch.randn(3, self.impl.qk_rope_head_dim)
|
|
sin = torch.randn(3, self.impl.qk_rope_head_dim)
|
|
|
|
kw_out = torch.randn(3, self.impl.head_dim * 2)
|
|
self.impl.wk_weights_proj.return_value = (kw_out, None)
|
|
self.impl.wq_b.return_value = (
|
|
torch.randn(3, self.impl.n_head * self.impl.head_dim),
|
|
None,
|
|
)
|
|
mock_rope.return_value = torch.randn(3, self.impl.n_head, self.impl.head_dim)
|
|
|
|
attn_metadata = MagicMock()
|
|
attn_metadata.num_decodes = 1
|
|
attn_metadata.num_decode_tokens = 1
|
|
attn_metadata.num_prefills = 1
|
|
attn_metadata.block_table = torch.tensor([[0], [1]], dtype=torch.int32)
|
|
|
|
sfa_cp_metadata = MagicMock()
|
|
sfa_cp_metadata.valid_block_ids = torch.tensor([0, 1], dtype=torch.int64)
|
|
sfa_cp_metadata.block_table_cp = torch.tensor([[0, 1]], dtype=torch.int32)
|
|
sfa_cp_metadata.prefill_q_cum_seqlens = torch.tensor([2], dtype=torch.int32)
|
|
sfa_cp_metadata.block_arange = torch.tensor([0], dtype=torch.int32)
|
|
attn_metadata.sfa_cp_metadata = sfa_cp_metadata
|
|
|
|
actual_seq_lengths_query = torch.tensor([1, 3], dtype=torch.int32)
|
|
actual_seq_lengths_key = torch.tensor([4, 8], dtype=torch.int32)
|
|
|
|
# In each call, returned tensor has rows matching q
|
|
call_counter = [0]
|
|
|
|
def fake_indexer(query, **kwargs):
|
|
call_counter[0] += 1
|
|
return torch.tensor([[0]] * query.shape[0]), None
|
|
|
|
with patch.object(
|
|
torch.ops._C_ascend,
|
|
"npu_lightning_indexer",
|
|
create=True,
|
|
side_effect=fake_indexer,
|
|
):
|
|
result = self.impl.indexer_select_post_process(
|
|
x, q_c, kv_cache, attn_metadata, cos, sin, actual_seq_lengths_query, actual_seq_lengths_key
|
|
)
|
|
self.assertIsNotNone(result)
|
|
self.assertEqual(result.shape[0], 3)
|
|
|
|
@patch("vllm_ascend.attention.context_parallel.sfa_cp.HAS_TRITON", True)
|
|
@patch("vllm_ascend.attention.context_parallel.sfa_cp.rope_forward_triton_siso")
|
|
@patch_distributed_groups(dcp_size=2, pcp_size=2, needs_mocks=False)
|
|
def test_indexer_select_post_process_prefill_with_pcp(self, mock_rope):
|
|
# Case: prefills + pcp head/tail processing
|
|
self.impl.pcp_size = 2
|
|
self.impl.dcp_size = 2
|
|
self.impl.use_torch_npu_lightning_indexer = False
|
|
|
|
# 4 prefill tokens
|
|
x = torch.randn(4, self.impl.qk_head_dim)
|
|
q_c = torch.randn(4, self.impl.q_lora_rank)
|
|
kv_cache = (
|
|
torch.randn(4, 1, 1, 32),
|
|
torch.randn(4, 1, 1, 16),
|
|
torch.randn(4, 1, 1, self.impl.head_dim),
|
|
)
|
|
cos = torch.randn(4, self.impl.qk_rope_head_dim)
|
|
sin = torch.randn(4, self.impl.qk_rope_head_dim)
|
|
|
|
kw_out = torch.randn(4, self.impl.head_dim * 2)
|
|
self.impl.wk_weights_proj.return_value = (kw_out, None)
|
|
self.impl.wq_b.return_value = (
|
|
torch.randn(4, self.impl.n_head * self.impl.head_dim),
|
|
None,
|
|
)
|
|
mock_rope.return_value = torch.randn(4, self.impl.n_head, self.impl.head_dim)
|
|
|
|
attn_metadata = MagicMock()
|
|
attn_metadata.num_decodes = 0
|
|
attn_metadata.num_decode_tokens = 0
|
|
attn_metadata.num_prefills = 2
|
|
attn_metadata.block_table = torch.tensor([[0, 1], [2, 3]], dtype=torch.int32)
|
|
|
|
sfa_cp_metadata = MagicMock()
|
|
sfa_cp_metadata.valid_block_ids = torch.tensor([0, 1, 2, 3], dtype=torch.int64)
|
|
sfa_cp_metadata.block_table_cp = torch.tensor([[0, 1], [2, 3]], dtype=torch.int32)
|
|
sfa_cp_metadata.prefill_q_cum_seqlens = torch.tensor([2, 4], dtype=torch.int32)
|
|
sfa_cp_metadata.q_head_idx = torch.tensor([0, 1], dtype=torch.int64)
|
|
sfa_cp_metadata.q_tail_idx = torch.tensor([2, 3], dtype=torch.int64)
|
|
sfa_cp_metadata.q_full_idx = torch.tensor([0, 1, 2, 3], dtype=torch.int64)
|
|
sfa_cp_metadata.head_attn_nomask_seqlens = torch.tensor([4, 4], dtype=torch.int32)
|
|
sfa_cp_metadata.tail_attn_nomask_seqlens = torch.tensor([8, 8], dtype=torch.int32)
|
|
sfa_cp_metadata.block_arange = torch.tensor([0, 1, 2, 3], dtype=torch.int32)
|
|
attn_metadata.sfa_cp_metadata = sfa_cp_metadata
|
|
|
|
actual_seq_lengths_query = torch.tensor([2, 4], dtype=torch.int32)
|
|
actual_seq_lengths_key = torch.tensor([4, 8], dtype=torch.int32)
|
|
|
|
def fake_indexer(query, **kwargs):
|
|
return torch.tensor([[0]] * query.shape[0]), None
|
|
|
|
with patch.object(
|
|
torch.ops._C_ascend,
|
|
"npu_lightning_indexer",
|
|
create=True,
|
|
side_effect=fake_indexer,
|
|
):
|
|
result = self.impl.indexer_select_post_process(
|
|
x, q_c, kv_cache, attn_metadata, cos, sin, actual_seq_lengths_query, actual_seq_lengths_key
|
|
)
|
|
self.assertIsNotNone(result)
|
|
self.assertEqual(result.shape[0], 4)
|
|
|
|
@patch("vllm_ascend.attention.context_parallel.sfa_cp.HAS_TRITON", True)
|
|
@patch("vllm_ascend.attention.context_parallel.sfa_cp.rope_forward_triton_siso")
|
|
@patch_distributed_groups(dcp_size=2, pcp_size=2, needs_mocks=False)
|
|
def test_indexer_select_post_process_decode_and_prefill_with_pcp(self, mock_rope):
|
|
# Case: decodes + prefills + pcp; covers final torch.cat([decode, attn_output]).
|
|
self.impl.pcp_size = 2
|
|
self.impl.dcp_size = 2
|
|
self.impl.use_torch_npu_lightning_indexer = False
|
|
|
|
# 1 decode + 4 prefill = 5 total
|
|
x = torch.randn(5, self.impl.qk_head_dim)
|
|
q_c = torch.randn(5, self.impl.q_lora_rank)
|
|
kv_cache = (
|
|
torch.randn(4, 1, 1, 32),
|
|
torch.randn(4, 1, 1, 16),
|
|
torch.randn(4, 1, 1, self.impl.head_dim),
|
|
)
|
|
cos = torch.randn(5, self.impl.qk_rope_head_dim)
|
|
sin = torch.randn(5, self.impl.qk_rope_head_dim)
|
|
|
|
kw_out = torch.randn(5, self.impl.head_dim * 2)
|
|
self.impl.wk_weights_proj.return_value = (kw_out, None)
|
|
self.impl.wq_b.return_value = (
|
|
torch.randn(5, self.impl.n_head * self.impl.head_dim),
|
|
None,
|
|
)
|
|
mock_rope.return_value = torch.randn(5, self.impl.n_head, self.impl.head_dim)
|
|
|
|
attn_metadata = MagicMock()
|
|
attn_metadata.num_decodes = 1
|
|
attn_metadata.num_decode_tokens = 1
|
|
attn_metadata.num_prefills = 2
|
|
attn_metadata.block_table = torch.tensor([[0], [1], [2]], dtype=torch.int32)
|
|
|
|
sfa_cp_metadata = MagicMock()
|
|
sfa_cp_metadata.valid_block_ids = torch.tensor([0, 1, 2, 3], dtype=torch.int64)
|
|
sfa_cp_metadata.block_table_cp = torch.tensor([[0, 1], [2, 3]], dtype=torch.int32)
|
|
sfa_cp_metadata.prefill_q_cum_seqlens = torch.tensor([2, 4], dtype=torch.int32)
|
|
sfa_cp_metadata.q_head_idx = torch.tensor([0, 1], dtype=torch.int64)
|
|
sfa_cp_metadata.q_tail_idx = torch.tensor([2, 3], dtype=torch.int64)
|
|
sfa_cp_metadata.q_full_idx = torch.tensor([0, 1, 2, 3], dtype=torch.int64)
|
|
sfa_cp_metadata.head_attn_nomask_seqlens = torch.tensor([4, 4, 4], dtype=torch.int32)
|
|
sfa_cp_metadata.tail_attn_nomask_seqlens = torch.tensor([8, 8, 8], dtype=torch.int32)
|
|
sfa_cp_metadata.block_arange = torch.tensor([0, 1, 2, 3], dtype=torch.int32)
|
|
attn_metadata.sfa_cp_metadata = sfa_cp_metadata
|
|
|
|
actual_seq_lengths_query = torch.tensor([1, 3, 5], dtype=torch.int32)
|
|
actual_seq_lengths_key = torch.tensor([4, 8, 8], dtype=torch.int32)
|
|
|
|
def fake_indexer(query, **kwargs):
|
|
return torch.tensor([[0]] * query.shape[0]), None
|
|
|
|
with patch.object(
|
|
torch.ops._C_ascend,
|
|
"npu_lightning_indexer",
|
|
create=True,
|
|
side_effect=fake_indexer,
|
|
):
|
|
result = self.impl.indexer_select_post_process(
|
|
x, q_c, kv_cache, attn_metadata, cos, sin, actual_seq_lengths_query, actual_seq_lengths_key
|
|
)
|
|
self.assertIsNotNone(result)
|
|
self.assertEqual(result.shape[0], 5)
|
|
|
|
|
|
class TestAscendSFADCPImpl(TestBase):
|
|
@staticmethod
|
|
def _make_dsa_cp_context(
|
|
*,
|
|
num_tokens: int = 4,
|
|
num_tokens_pad: int = 4,
|
|
local_start: int = 0,
|
|
local_end_with_pad: int = 2,
|
|
) -> DSACPContext:
|
|
return DSACPContext(
|
|
num_tokens=num_tokens,
|
|
num_tokens_pad=num_tokens_pad,
|
|
local_start=local_start,
|
|
local_end=min(local_end_with_pad, num_tokens),
|
|
local_end_with_pad=local_end_with_pad,
|
|
slot_mapping_cp=torch.empty(0, dtype=torch.int32),
|
|
actual_seq_lengths_query=torch.empty(0, dtype=torch.int32),
|
|
actual_seq_lengths_key=torch.empty(0, dtype=torch.int32),
|
|
)
|
|
|
|
def test_all_to_all_dcp_tensor_supports_head_and_token_scatter(self):
|
|
impl = AscendSFADCPImpl.__new__(AscendSFADCPImpl)
|
|
impl.dcp_size = 2
|
|
impl.dcp_group = MagicMock()
|
|
tensor = torch.arange(4 * 6 * 2).view(4, 6, 2)
|
|
|
|
def copy_all_to_all(recv, send, *, group):
|
|
self.assertIs(group, impl.dcp_group.device_group)
|
|
recv.copy_(send)
|
|
|
|
with patch(
|
|
"vllm_ascend.attention.context_parallel.sfa_cp.dist.all_to_all_single",
|
|
side_effect=copy_all_to_all,
|
|
):
|
|
for scatter_dim in (0, 1):
|
|
with self.subTest(scatter_dim=scatter_dim):
|
|
result = impl._all_to_all_dcp_tensor(tensor, scatter_dim)
|
|
send = tensor.movedim(scatter_dim, 0).contiguous()
|
|
expected = send.view(impl.dcp_size, send.shape[0] // impl.dcp_size, *send.shape[1:])
|
|
torch.testing.assert_close(result, expected)
|
|
|
|
def test_all_to_all_dcp_tensor_rejects_nondivisible_scatter_size(self):
|
|
impl = AscendSFADCPImpl.__new__(AscendSFADCPImpl)
|
|
impl.dcp_size = 2
|
|
impl.dcp_group = MagicMock()
|
|
|
|
with self.assertRaisesRegex(RuntimeError, "scatter dimension to be divisible"):
|
|
impl._all_to_all_dcp_tensor(torch.empty(3, 4, 2), scatter_dim=0)
|
|
|
|
def test_merge_dcp_outputs_with_torch_matches_head_and_token_layouts(self):
|
|
output_by_head = torch.arange(2 * 2 * 3 * 2, dtype=torch.float16).view(2, 2, 3, 2)
|
|
lse_by_head = torch.tensor(
|
|
[
|
|
[[0.0, 1.0, -1.0], [2.0, 0.0, 1.0]],
|
|
[[1.0, 0.0, -2.0], [0.0, 2.0, 1.0]],
|
|
],
|
|
dtype=torch.float32,
|
|
)
|
|
weights = torch.softmax(lse_by_head, dim=0)
|
|
expected = (output_by_head.float() * weights.unsqueeze(-1)).sum(dim=0).movedim(1, 0).contiguous()
|
|
|
|
native_output = AscendSFADCPImpl._merge_dcp_outputs_with_torch(
|
|
output_by_head,
|
|
lse_by_head,
|
|
token_dim=2,
|
|
)
|
|
dsa_output = AscendSFADCPImpl._merge_dcp_outputs_with_torch(
|
|
output_by_head.movedim(1, 2),
|
|
lse_by_head.movedim(1, 2),
|
|
token_dim=1,
|
|
)
|
|
|
|
torch.testing.assert_close(native_output, expected)
|
|
torch.testing.assert_close(dsa_output, expected)
|
|
|
|
def test_merge_dcp_outputs_selects_native_head_scatter(self):
|
|
impl = AscendSFADCPImpl.__new__(AscendSFADCPImpl)
|
|
impl.dcp_size = 2
|
|
impl.dcp_rank = 0
|
|
impl.dcp_group = MagicMock()
|
|
sfa_output = torch.randn(4, 4, 8)
|
|
softmax_lse = torch.randn(4, 4, 1)
|
|
output_recv = torch.randn(2, 2, 4, 8)
|
|
lse_recv = torch.randn(2, 2, 4, 1)
|
|
expected = torch.randn(4, 2, 8)
|
|
impl._all_to_all_dcp_tensor = MagicMock(side_effect=(output_recv, lse_recv))
|
|
impl._merge_dcp_outputs_with_torch = MagicMock(return_value=expected)
|
|
|
|
result = impl._merge_dcp_outputs(sfa_output, softmax_lse)
|
|
|
|
self.assertIs(result, expected)
|
|
all_to_all_calls = impl._all_to_all_dcp_tensor.call_args_list
|
|
self.assertEqual(len(all_to_all_calls), 2)
|
|
self.assertIs(all_to_all_calls[0].args[0], sfa_output)
|
|
self.assertEqual(all_to_all_calls[0].args[1], 1)
|
|
self.assertIs(all_to_all_calls[1].args[0], softmax_lse)
|
|
self.assertEqual(all_to_all_calls[1].args[1], 1)
|
|
merge_args = impl._merge_dcp_outputs_with_torch.call_args.args
|
|
self.assertIs(merge_args[0], output_recv)
|
|
torch.testing.assert_close(merge_args[1], lse_recv.squeeze(-1))
|
|
self.assertEqual(merge_args[2], 2)
|
|
|
|
def test_merge_dcp_outputs_selects_dsa_token_scatter(self):
|
|
impl = AscendSFADCPImpl.__new__(AscendSFADCPImpl)
|
|
impl.dcp_size = 2
|
|
impl.dcp_rank = 0
|
|
impl.dcp_group = MagicMock()
|
|
sfa_output = torch.randn(4, 4, 8)
|
|
softmax_lse = torch.randn(4, 4, 1)
|
|
output_recv = torch.randn(2, 2, 4, 8)
|
|
lse_recv = torch.randn(2, 2, 4, 1)
|
|
expected = torch.randn(2, 4, 8)
|
|
dsa_cp_context = self._make_dsa_cp_context()
|
|
impl._all_to_all_dcp_tensor = MagicMock(side_effect=(output_recv, lse_recv))
|
|
impl._merge_dcp_outputs_with_torch = MagicMock(return_value=expected)
|
|
|
|
result = impl._merge_dcp_outputs(sfa_output, softmax_lse, dsa_cp_context)
|
|
|
|
self.assertIs(result, expected)
|
|
all_to_all_calls = impl._all_to_all_dcp_tensor.call_args_list
|
|
self.assertEqual(len(all_to_all_calls), 2)
|
|
self.assertIs(all_to_all_calls[0].args[0], sfa_output)
|
|
self.assertEqual(all_to_all_calls[0].args[1], 0)
|
|
self.assertIs(all_to_all_calls[1].args[0], softmax_lse)
|
|
self.assertEqual(all_to_all_calls[1].args[1], 0)
|
|
merge_args = impl._merge_dcp_outputs_with_torch.call_args.args
|
|
self.assertIs(merge_args[0], output_recv)
|
|
torch.testing.assert_close(merge_args[1], lse_recv.squeeze(-1))
|
|
self.assertEqual(merge_args[2], 1)
|
|
|
|
def test_merge_dcp_outputs_rejects_misaligned_dsa_token_shard(self):
|
|
impl = AscendSFADCPImpl.__new__(AscendSFADCPImpl)
|
|
impl.dcp_size = 2
|
|
impl.dcp_rank = 0
|
|
impl.dcp_group = MagicMock()
|
|
dsa_cp_context = self._make_dsa_cp_context(local_start=1, local_end_with_pad=3)
|
|
|
|
with self.assertRaisesRegex(RuntimeError, "token shards must follow DCP rank order"):
|
|
impl._merge_dcp_outputs(
|
|
torch.randn(4, 4, 8),
|
|
torch.randn(4, 4, 1),
|
|
dsa_cp_context,
|
|
)
|
|
|
|
def test_record_dcp_kv_gather_context_c8_only_gathers_packed_kv(self):
|
|
impl = AscendSFADCPImpl.__new__(AscendSFADCPImpl)
|
|
impl.dcp_group = MagicMock()
|
|
impl.enable_sparse_sfa_c8 = True
|
|
gather_context = MagicMock()
|
|
impl._start_dcp_gather = MagicMock(return_value=gather_context)
|
|
|
|
packed_kv = torch.arange(4 * 16, dtype=torch.int8).view(4, 1, 1, 16)
|
|
indexer_k = torch.zeros(8, 1, 1, 8, dtype=torch.int8)
|
|
indexer_scale = torch.zeros(8, 1, 1, 1, dtype=torch.float16)
|
|
valid_block_ids = torch.tensor([1, 3], dtype=torch.int64)
|
|
attn_metadata = MagicMock()
|
|
attn_metadata.num_prefills = 1
|
|
attn_metadata.dcp_context = DCPContext(
|
|
slot_mapping=torch.tensor([0], dtype=torch.int32),
|
|
block_table=torch.tensor([[0]], dtype=torch.int32),
|
|
seq_lens=torch.tensor([1], dtype=torch.int32),
|
|
kv_gather_block_ids=valid_block_ids,
|
|
kv_gather_block_table=torch.tensor([[0, 1]], dtype=torch.int32),
|
|
)
|
|
|
|
impl._record_dcp_kv_gather_context(
|
|
(packed_kv, indexer_k, indexer_scale),
|
|
attn_metadata,
|
|
)
|
|
|
|
gather_input = impl._start_dcp_gather.call_args.args[0]
|
|
self.assertTrue(torch.equal(gather_input, packed_kv.index_select(0, valid_block_ids)))
|
|
self.assertEqual(impl._start_dcp_gather.call_args.kwargs["dim"], 0)
|
|
self.assertEqual(impl._start_dcp_gather.call_args.kwargs["split_sizes"], (16,))
|
|
self.assertIs(attn_metadata.dcp_context.gather_context, gather_context)
|
|
|
|
def test_execute_sparse_flash_attention_process_uses_c8_device_operator_lse(self):
|
|
impl = AscendSFADCPImpl.__new__(AscendSFADCPImpl)
|
|
impl.dcp_group = MagicMock()
|
|
impl.dcp_size = 2
|
|
impl.enable_dsa_cp = False
|
|
impl.enable_sparse_sfa_c8 = True
|
|
impl.qk_rope_head_dim = 4
|
|
impl.sfa_qsfa_tile_size = 128
|
|
impl._remap_sparse_indices = MagicMock(side_effect=lambda x: x)
|
|
merged_output = torch.randn(2, 2, 8)
|
|
impl._merge_dcp_outputs = MagicMock(return_value=merged_output)
|
|
|
|
ql_nope = torch.randn(2, 2, 8)
|
|
q_pe = torch.randn(2, 2, 4)
|
|
impl._finish_dcp_gather = MagicMock(side_effect=lambda _ctx: (ql_nope, q_pe))
|
|
kv_cache = (torch.empty(4, 1, 1, 16, dtype=torch.int8),)
|
|
topk_indices = torch.zeros(2, 1, dtype=torch.int32)
|
|
actual_seq_lengths_query = torch.tensor([2], dtype=torch.int32)
|
|
actual_seq_lengths_key = torch.tensor([4], dtype=torch.int32)
|
|
dcp_seq_lens = torch.tensor([3], dtype=torch.int32)
|
|
dcp_block_table = torch.tensor([[0, 1]], dtype=torch.int32)
|
|
|
|
attn_metadata = MagicMock()
|
|
attn_metadata.num_prefills = 0
|
|
attn_metadata.dsa_cp_context = None
|
|
attn_metadata.dcp_context = DCPContext(
|
|
slot_mapping=torch.tensor([0, 1], dtype=torch.int32),
|
|
block_table=dcp_block_table,
|
|
seq_lens=dcp_seq_lens,
|
|
gather_context=MagicMock(),
|
|
)
|
|
sfa_output = torch.randn(2, 2, 8)
|
|
softmax_lse = torch.randn(2, 2, 1)
|
|
|
|
with patch(
|
|
"vllm_ascend.attention.context_parallel.sfa_cp.DeviceOperator.execute_sparse_flash_attention_process",
|
|
return_value=(sfa_output, softmax_lse),
|
|
) as mock_execute_sfa:
|
|
result = impl._execute_sparse_flash_attention_process(
|
|
ql_nope,
|
|
q_pe,
|
|
kv_cache,
|
|
topk_indices,
|
|
attn_metadata,
|
|
actual_seq_lengths_query,
|
|
actual_seq_lengths_key,
|
|
)
|
|
|
|
self.assertIs(result, merged_output)
|
|
call_args = mock_execute_sfa.call_args.args
|
|
call_kwargs = mock_execute_sfa.call_args.kwargs
|
|
self.assertTrue(call_args[0].enable_sparse_sfa_c8)
|
|
self.assertEqual(len(call_args[3]), 1)
|
|
self.assertEqual(call_args[3][0].dtype, torch.int8)
|
|
self.assertIs(call_args[7], dcp_seq_lens)
|
|
self.assertIs(call_kwargs["block_table"], dcp_block_table)
|
|
self.assertEqual(call_kwargs["sparse_mode"], 0)
|
|
self.assertTrue(call_kwargs["return_lse"])
|
|
impl._merge_dcp_outputs.assert_called_once_with(sfa_output, softmax_lse, None)
|