From 84c10e84ac92db129e3fec5dfa32e79d3c5c4809 Mon Sep 17 00:00:00 2001 From: Ruiqiu Zheng <191817791+Ruiqiu-Zheng@users.noreply.github.com> Date: Wed, 9 Sep 2026 07:48:22 +0800 Subject: [PATCH] [Performance][SFA] Rebase full-visible index bypass onto decoupled indexer backend Signed-off-by: Ruiqiu Zheng <191817791+Ruiqiu-Zheng@users.noreply.github.com> --- .../configuration/additional_config.md | 23 ++ tests/ut/attention/test_indexer.py | 220 ++++++++++++++++++ tests/ut/ops/test_mla.py | 38 ++- .../ut/patch/worker/test_patch_deepseek_v2.py | 16 +- tests/ut/test_ascend_config.py | 12 + vllm_ascend/ascend_config.py | 1 + vllm_ascend/attention/indexer.py | 117 +++++++++- vllm_ascend/models/glm5next/attention.py | 2 + vllm_ascend/models/glm5next/model.py | 1 + vllm_ascend/ops/mla.py | 25 +- vllm_ascend/patch/worker/patch_deepseek_v2.py | 24 +- 11 files changed, 461 insertions(+), 18 deletions(-) diff --git a/docs/source/user_guide/configuration/additional_config.md b/docs/source/user_guide/configuration/additional_config.md index f7526f47091..72798fbf7f0 100644 --- a/docs/source/user_guide/configuration/additional_config.md +++ b/docs/source/user_guide/configuration/additional_config.md @@ -83,6 +83,29 @@ The following table lists additional configuration options available in vLLM Asc | `enable_reduce_sample` | bool | `False` | Whether to enable reduce sample optimization to reduce communication and computation overheads in the tensor parallelism scenario. When enabled, logits are kept partitioned across TP ranks and only the small set of top-k candidate values/indices is communicated, instead of performing a full-vocabulary all-to-all/all-gather. **Note**: This is an experimental feature. **Limitations**: (1) Not supported on PD-disaggregated scenario. (2) Must be disabled when sampling logprobs are requested. When reduce sample is enabled, logprobs are silently computed over partitioned logits instead of the full vocabulary, producing incorrect logprob values and top-k rankings. (3) Cannot be enabled together with lmhead TP.| | `combine_quant_mode` | int | `0` | Fused MC2 configuration. This configuration will be passed as the `comm_quant_mode` argument for the `torch_npu.npu_moe_distribute_combine_v2` operator. Please refer to the operator documentation for the valid value range. | +**Short-prefill full-visible indexer bypass** + +`enable_sfa_full_visible_index_bypass` is an opt-in boolean, default `False`: + +```bash +vllm serve /path/to/model --additional-config '{"enable_sfa_full_visible_index_bypass": true}' +``` + +For eligible SFA layers, the indexer writes its key cache and returns all +causally visible indices without LightningIndexer scoring. This applies only +when the model grants scoring-skip permission, to a single sequence in +`PrefillNoCache` or `PrefillCacheHit`, with at most 2048 total visible KV tokens +(including cached context), top-k size 2048, and block size 128 on NPU. +MTP layers, speculative decoding, C8 SFA or indexer caches, PCP/DSA context +parallelism, chunked prefill, decode, multiple sequences, and unsupported +metadata, devices or cache geometry retain the existing scoring path. +Top-k reuse layers retain their existing cache-write and reuse behavior. + +The read-only round-robin index table is `[2049, 2048]` with `int32` entries: +approximately **16.008 MiB per process/device**, shared across layers and +requests, not allocated per layer. Eligible requests receive zero-copy views. +The default setting leaves incumbent scoring behavior unchanged. + The details of each configuration option are as follows: **xlite_graph_config** diff --git a/tests/ut/attention/test_indexer.py b/tests/ut/attention/test_indexer.py index 01d0324e27e..29dd384d54f 100644 --- a/tests/ut/attention/test_indexer.py +++ b/tests/ut/attention/test_indexer.py @@ -4,10 +4,12 @@ from types import SimpleNamespace from unittest.mock import MagicMock, patch +import pytest import torch from vllm.v1.attention.backend import AttentionCGSupport from vllm.v1.kv_cache_interface import FullAttentionSpec +from vllm_ascend.attention.attention_v1 import AscendAttentionState from vllm_ascend.attention.indexer import ( AscendSFAIndexerBackend, AscendSFAIndexerMetadata, @@ -138,3 +140,221 @@ def test_sfa_indexer_metadata_builder_primes_reshape_optim( common.group_key_cache_idx, _KERNEL_BLOCK_SIZE, ) + + +@pytest.mark.parametrize("source", ["_seq_lens_cpu", "seq_lens_cpu", None]) +def test_metadata_cpu_state_bridge_without_device_copy(source): + common = _make_common_metadata() + common.attn_state = AscendAttentionState.PrefillCacheHit + common._seq_lens_cpu = None + common.seq_lens_cpu = None + cpu_lengths = torch.tensor([5, 6, 7]) + if source is not None: + setattr(common, source, cpu_lengths) + with ( + patch("vllm_ascend.attention.indexer.get_ascend_config") as config, + patch("vllm_ascend.attention.indexer.get_cos_and_sin_mla") as rope, + ): + config.return_value.c8_reshape_optim_enabled = False + rope.return_value = (torch.zeros(5, 8), torch.zeros(5, 8)) + metadata = _make_builder().build(0, common) + assert metadata.attn_state is common.attn_state + if source is None: + assert metadata.seq_lens_cpu is None + else: + assert metadata.seq_lens_cpu.tolist() == [5, 6] + assert metadata.seq_lens_cpu.data_ptr() == cpu_lengths.data_ptr() + + +class _DeviceTable: + """Mock only device placement; slice real CPU storage for host assertions.""" + + def __init__(self, table, device): + self.table = table + self.device = device + + def __getitem__(self, key): + return self.table[key] + + +@pytest.fixture +def full_visible_indexer(): + config = SimpleNamespace(enable_sfa_full_visible_index_bypass=True, enable_sparse_sfa_c8=False) + with ( + patch("vllm_ascend.attention.indexer.get_ascend_config", return_value=config), + patch.object(AscendSFAIndexerBackend, "_full_visible_index_tables", {}), + ): + impl = AscendSFAIndexerBackend.__new__(AscendSFAIndexerBackend) + impl.allow_short_prefill_indexer_scoring_skip = True + impl.enable_sparse_li_c8 = False + impl._pcp_active = False + impl._dsa_cp_active = False + impl._speculative_active = False + impl.topk_tokens = 2048 + device = torch.device("npu", 0) + table = impl._get_or_create_full_visible_index_table(torch.device("cpu")) + impl._full_visible_index_table = _DeviceTable(table, device) + yield impl, config, table, device + + +def _full_visible_metadata(total=192, query=64): + return SimpleNamespace( + attn_state=AscendAttentionState.PrefillCacheHit, + seq_lens_cpu=torch.tensor([total]), + seq_lens=torch.tensor([total]), + cum_query_lens=torch.tensor([query]), + num_actual_tokens=query, + num_decode_tokens=0, + block_size=128, + block_table=torch.zeros(1, 16, dtype=torch.int32), + slot_mapping=torch.arange(total - query, total), + actual_seq_lengths_query=torch.tensor([query]), + actual_seq_lengths_key=torch.tensor([total]), + ) + + +def test_exact_rr_table_and_shared_storage(full_visible_indexer): + impl, _, table, _ = full_visible_indexer + assert table.shape == (2049, 2048) + assert table.dtype == torch.int32 + assert table.untyped_storage().nbytes() == 2049 * 2048 * 4 + assert torch.all(table[0] == -1) + assert table is impl._get_or_create_full_visible_index_table(torch.device("cpu")) + for visible in range(1, 2049): + # Independent scalar block/offset oracle, including partial blocks. + expected = [ + block + offset for offset in range(128) for block in range(0, visible, 128) if block + offset < visible + ] + row = table[visible].tolist() + assert row[:visible] == expected + assert set(row[:visible]) == set(range(visible)) + assert row[visible:] == [-1] * (2048 - visible) + + +@pytest.mark.parametrize("total,query", [(1, 1), (64, 64), (192, 192), (2048, 2048), (2048, 64)]) +def test_eligible_views_alias_table(full_visible_indexer, total, query): + impl, _, table, device = full_visible_indexer + metadata = _full_visible_metadata(total, query) + if total == query: + metadata.attn_state = AscendAttentionState.PrefillNoCache + before = table.clone() + rows = impl._get_full_visible_topk_indices(metadata, query, device) + assert rows.shape == (query, 1, 2048) + assert rows.untyped_storage().data_ptr() == table.untyped_storage().data_ptr() + assert rows.storage_offset() == (total - query + 1) * 2048 + destination = torch.empty(query, 2048, dtype=torch.int32) + destination.copy_(rows.squeeze(1)) + assert destination.untyped_storage().data_ptr() != table.untyped_storage().data_ptr() + assert torch.equal(table, before) + assert torch.equal(destination, table[total - query + 1 : total + 1]) + + +@pytest.mark.parametrize( + "target,field,value", + [ + ("config", "enable_sfa_full_visible_index_bypass", False), + ("config", "enable_sparse_sfa_c8", True), + ("impl", "allow_short_prefill_indexer_scoring_skip", False), + ("impl", "enable_sparse_li_c8", True), + ("impl", "_pcp_active", True), + ("impl", "_dsa_cp_active", True), + ("impl", "_speculative_active", True), + ("impl", "topk_tokens", 1024), + ("impl", "_full_visible_index_table", None), + ("metadata", "attn_state", AscendAttentionState.DecodeOnly), + ("metadata", "attn_state", AscendAttentionState.SpecDecoding), + ("metadata", "attn_state", AscendAttentionState.ChunkedPrefill), + ("metadata", "attn_state", None), + ("metadata", "attn_state", AscendAttentionState.PrefillNoCache), + ("metadata", "num_decode_tokens", 1), + ("metadata", "block_size", 64), + ("metadata", "seq_lens_cpu", None), + ("metadata", "seq_lens_cpu", torch.tensor([2049])), + ("metadata", "seq_lens_cpu", torch.tensor([63])), + ("metadata", "seq_lens_cpu", torch.tensor([192.5])), + ("metadata", "seq_lens_cpu", torch.tensor([192, 192])), + ("metadata", "seq_lens_cpu", torch.empty(1, device="meta", dtype=torch.int64)), + ("metadata", "block_table", torch.zeros(2, 16)), + ("metadata", "block_table", torch.zeros(1, 1)), + ("metadata", "seq_lens", torch.tensor([192, 192])), + ("metadata", "cum_query_lens", torch.tensor([32, 64])), + ("metadata", "num_actual_tokens", 65), + ], +) +def test_narrow_fallbacks(full_visible_indexer, target, field, value): + impl, config, _, device = full_visible_indexer + metadata = _full_visible_metadata() + setattr({"impl": impl, "config": config, "metadata": metadata}[target], field, value) + assert impl._get_full_visible_topk_indices(metadata, 64, device) is None + + +def test_unsupported_device_and_empty_query(full_visible_indexer): + impl, _, _, device = full_visible_indexer + metadata = _full_visible_metadata() + assert impl._get_full_visible_topk_indices(metadata, 64, torch.device("cpu")) is None + assert impl._get_full_visible_topk_indices(metadata, 64, torch.device("npu", 1)) is None + metadata.num_actual_tokens = 0 + assert impl._get_full_visible_topk_indices(metadata, 0, device) is None + + +@pytest.mark.parametrize("enabled,compute_topk", [(True, True), (False, True), (True, False)]) +def test_forward_cache_write_before_bypass_or_scorer(full_visible_indexer, enabled, compute_topk): + impl, config, table, device = full_visible_indexer + config.enable_sfa_full_visible_index_bypass = enabled + metadata = _full_visible_metadata() + events = [] + keys = torch.ones(64, 128) + impl.forward_k = MagicMock(return_value=(keys, None)) + impl.write_cache = MagicMock(side_effect=lambda *a, **kw: events.append("write")) + original = impl._get_full_visible_topk_indices + + def eligible(*args): + events.append("eligibility") + return original(*args) + + impl._get_full_visible_topk_indices = MagicMock(side_effect=eligible) + impl.head_dim = 128 + impl.n_head = 1 + impl.qk_rope_head_dim = 64 + impl.is_rope_neox_style = True + impl.use_torch_npu_lightning_indexer = False + impl.k_cache = SimpleNamespace(kv_cache=(MagicMock(),)) + impl.wk_weights_proj = MagicMock(return_value=(torch.zeros(64, 129), None)) + impl.wq_b = MagicMock(return_value=(torch.zeros(64, 128), None)) + hidden = SimpleNamespace(shape=(64, 128), device=device, dtype=torch.float32) + scored = torch.full((64, 1, 2048), -2, dtype=torch.int32) + + def score(*args): + events.append("score") + return scored + + with ( + patch("vllm_ascend.attention.indexer.HAS_TRITON", True), + patch("vllm_ascend.attention.indexer.rope_forward_triton_siso", side_effect=lambda x, *a, **kw: x), + patch("vllm_ascend.attention.indexer.DeviceOperator.indexer_select_post_process", side_effect=score), + ): + result = impl.forward(hidden, torch.zeros(64, 128), None, None, keys, metadata, compute_topk) + impl.write_cache.assert_called_once_with(keys, None, metadata.slot_mapping, indexer_attn_metadata=metadata) + if not compute_topk: + assert result is None + assert events == ["write"] + elif enabled: + assert events == ["write", "eligibility"] + assert result.untyped_storage().data_ptr() == table.untyped_storage().data_ptr() + impl.wk_weights_proj.assert_not_called() + impl.wq_b.assert_not_called() + else: + assert events == ["write", "eligibility", "score"] + assert result is scored + + +def test_specialized_backend_falls_back(full_visible_indexer): + _, _, table, device = full_visible_indexer + + class SpecializedIndexer(AscendSFAIndexerBackend): + pass + + impl = SpecializedIndexer.__new__(SpecializedIndexer) + impl.allow_short_prefill_indexer_scoring_skip = True + impl._full_visible_index_table = _DeviceTable(table, device) + assert impl._get_full_visible_topk_indices(_full_visible_metadata(), 64, device) is None diff --git a/tests/ut/ops/test_mla.py b/tests/ut/ops/test_mla.py index 3b23549195b..022f8b84252 100644 --- a/tests/ut/ops/test_mla.py +++ b/tests/ut/ops/test_mla.py @@ -323,7 +323,7 @@ def test_constructs_backend_and_delegates_sfa_interface(self, mock_backend_cls): vllm_indexer = MagicMock(name="vllm_indexer") wrapper = IndexerWrapper(vllm_indexer, qk_rope_head_dim=64) - mock_backend_cls.assert_called_once_with(vllm_indexer, 64) + mock_backend_cls.assert_called_once_with(vllm_indexer, 64, allow_short_prefill_indexer_scoring_skip=False) self.assertIs(wrapper.impl, mock_backend_cls.return_value) self.assertIs(wrapper.k_cache, wrapper.impl.k_cache) @@ -333,6 +333,12 @@ def test_constructs_backend_and_delegates_sfa_interface(self, mock_backend_cls): wrapper.process_weights_after_loading() wrapper.impl.process_weights_after_loading.assert_called_once_with() + @patch("vllm_ascend.ops.mla.AscendSFAIndexerBackend") + def test_bypass_permission_reaches_backend(self, mock_backend): + indexer = MagicMock() + IndexerWrapper(indexer, 64, allow_short_prefill_indexer_scoring_skip=True) + mock_backend.assert_called_once_with(indexer, 64, allow_short_prefill_indexer_scoring_skip=True) + class TestAscendMultiHeadLatentAttention(TestBase): def setUp(self): @@ -394,9 +400,39 @@ def test_initialization(self, mock_tp_size, mock_get_vllm_config, mock_indexer_c prefix=self.prefix, ) + mock_indexer_cls.assert_called_once_with( + self.mock_mla_modules.indexer, + self.qk_rope_head_dim, + allow_short_prefill_indexer_scoring_skip=False, + ) self.assertEqual(attn.tp_size, 2) self.assertIsNotNone(attn.mla_attn) + @patch("vllm_ascend.ops.mla.MLAAttention") + @patch("vllm_ascend.ops.mla.IndexerWrapper") + @patch("vllm_ascend.ops.mla.get_current_vllm_config") + @patch("vllm_ascend.ops.mla.get_tensor_model_parallel_world_size", return_value=1) + def test_bypass_permission_reaches_indexer(self, mock_tp, mock_config, mock_indexer, mock_attention): + AscendMultiHeadLatentAttention( + self.hidden_size, + self.num_heads, + self.scale, + self.qk_nope_head_dim, + self.qk_rope_head_dim, + self.v_head_dim, + self.q_lora_rank, + self.kv_lora_rank, + self.mock_mla_modules, + prefix=self.prefix, + allow_short_prefill_indexer_scoring_skip=True, + ) + mock_indexer.assert_called_once_with( + self.mock_mla_modules.indexer, + self.qk_rope_head_dim, + allow_short_prefill_indexer_scoring_skip=True, + ) + self.assertNotIn("allow_short_prefill_indexer_scoring_skip", mock_attention.call_args.kwargs) + @patch("vllm_ascend.ops.mla.IndexerWrapper") @patch("vllm_ascend.ops.mla.torch.ops.vllm.mla_forward") @patch("vllm_ascend.ops.mla.get_current_vllm_config") diff --git a/tests/ut/patch/worker/test_patch_deepseek_v2.py b/tests/ut/patch/worker/test_patch_deepseek_v2.py index 6aeb5d98001..f02539a389e 100644 --- a/tests/ut/patch/worker/test_patch_deepseek_v2.py +++ b/tests/ut/patch/worker/test_patch_deepseek_v2.py @@ -2,7 +2,11 @@ from types import SimpleNamespace -from vllm_ascend.patch.worker.patch_deepseek_v2 import _should_skip_indexer_init +from vllm_ascend.patch.worker.patch_deepseek_v2 import ( + _is_mtp_layer, + _resolve_mtp_indexer_permissions, + _should_skip_indexer_init, +) def _config(**overrides) -> SimpleNamespace: @@ -34,3 +38,13 @@ def test_mtp_layer_keeps_indexer(): "model.layers.80.self_attn", skip_topk=True, ) + + +def test_mtp_layer_detection_from_config_and_prefix(): + assert _is_mtp_layer(_config(), "model.layers.79.self_attn") is False + assert _is_mtp_layer(_config(), "model.layers.80.self_attn") is True + + +def test_mtp_disables_short_prefill_bypass_and_topk_reuse(): + assert _resolve_mtp_indexer_permissions(True, False) == (True, True) + assert _resolve_mtp_indexer_permissions(True, True) == (False, False) diff --git a/tests/ut/test_ascend_config.py b/tests/ut/test_ascend_config.py index 0ca1a742454..a1fe0158702 100644 --- a/tests/ut/test_ascend_config.py +++ b/tests/ut/test_ascend_config.py @@ -1052,6 +1052,18 @@ def test_enable_prefill_mc2_string_false_disables(self, mock_fix): vc.additional_config = {"enable_prefill_mc2": "false"} self.assertFalse(init_ascend_config(vc).enable_prefill_mc2) + @_clean_up + @patch("vllm_ascend.platform.NPUPlatform.check_and_update_config") + def test_sfa_full_visible_index_bypass_defaults_false(self, mock_fix): + self.assertFalse(init_ascend_config(VllmConfig()).enable_sfa_full_visible_index_bypass) + + @_clean_up + @patch("vllm_ascend.platform.NPUPlatform.check_and_update_config") + def test_sfa_full_visible_index_bypass_opt_in(self, mock_fix): + vc = VllmConfig() + vc.additional_config = {"enable_sfa_full_visible_index_bypass": True} + self.assertTrue(init_ascend_config(vc).enable_sfa_full_visible_index_bypass) + @_clean_up @patch("vllm_ascend.platform.NPUPlatform.check_and_update_config") def test_a_family_additional_config_gets_typed_validation(self, mock_fix): diff --git a/vllm_ascend/ascend_config.py b/vllm_ascend/ascend_config.py index a82543c7b81..b3b907266b1 100644 --- a/vllm_ascend/ascend_config.py +++ b/vllm_ascend/ascend_config.py @@ -479,6 +479,7 @@ class AscendConfig: # ---- A-family (envs fallback): default = envs module value, before-validator injects ---- enable_fused_mc2: int = 0 enable_mlapo: bool = True + enable_sfa_full_visible_index_bypass: bool = False # When True, keep MLAPO prefill weights on NPU instead of freeing them # on kv_consumer D nodes. Trades NPU memory for stability — D nodes have # normal local-prefill paths (recompute / fallback / preempt) that crash diff --git a/vllm_ascend/attention/indexer.py b/vllm_ascend/attention/indexer.py index de9fcdf8e49..94e3d97c1ad 100644 --- a/vllm_ascend/attention/indexer.py +++ b/vllm_ascend/attention/indexer.py @@ -1,6 +1,7 @@ from dataclasses import dataclass from typing import Any +import numpy as np import scipy # type: ignore import torch import torch_npu @@ -18,6 +19,7 @@ from vllm.v1.worker.utils import select_common_block_size from vllm_ascend.ascend_config import get_ascend_config +from vllm_ascend.attention.attention_v1 import AscendAttentionState from vllm_ascend.device.device_op import DeviceOperator from vllm_ascend.device.hardware_profile import HardwareCapability, get_current_hardware_profile from vllm_ascend.distributed.utils import all_gather_async @@ -34,6 +36,8 @@ # tuple (the scale slot exists only when LI C8 is enabled). INDEXER_K_CACHE_SLOT = 0 INDEXER_SCALE_CACHE_SLOT = 1 +SFA_INDEXER_SPARSE_COUNT = 2048 +SFA_FULL_VISIBLE_TEMPLATE_BLOCK_SIZE = 128 @dataclass @@ -74,6 +78,9 @@ class AscendSFAIndexerMetadata: # gather splits the local prefill region on it (all-decode batches skip # the gather). num_decode_tokens: int = 0 + # CPU-only eligibility bridge; never copy device sequence lengths to host. + seq_lens_cpu: torch.Tensor | None = None + attn_state: AscendAttentionState | None = None class AscendSFAIndexerBackend(nn.Module, AttentionBackend): @@ -114,6 +121,11 @@ def topk_output_width(self) -> int: def get_topk_lengths(self, positions: torch.Tensor) -> torch.Tensor: return (positions + 1).clamp(min=0, max=self.topk_tokens) + # Read-only after creation; shared across layers/requests in this process. + _full_visible_index_tables: dict[torch.device, torch.Tensor] = {} + _full_visible_index_table: torch.Tensor | None = None + allow_short_prefill_indexer_scoring_skip: bool = False + # q_hadamard and k_hadamard tensor shared when dsa c8 enabled q_hadamard: torch.Tensor | None = None k_hadamard: torch.Tensor | None = None @@ -150,7 +162,12 @@ def get_supported_kernel_block_sizes() -> list[int]: # ---- model-side impl interface (per-layer instance) ---- - def __init__(self, vllm_indexer: nn.Module, qk_rope_head_dim: int) -> None: + def __init__( + self, + vllm_indexer: nn.Module, + qk_rope_head_dim: int, + allow_short_prefill_indexer_scoring_skip: bool = False, + ) -> None: super().__init__() self.n_head: int = vllm_indexer.n_head # 64 @@ -189,6 +206,91 @@ def __init__(self, vllm_indexer: nn.Module, qk_rope_head_dim: int) -> None: parallel_config = get_current_vllm_config().parallel_config self._pcp_active = parallel_config.prefill_context_parallel_size > 1 self._dsa_cp_active = enable_dsa_cp() + self.allow_short_prefill_indexer_scoring_skip = allow_short_prefill_indexer_scoring_skip + self._speculative_active = get_current_vllm_config().speculative_config is not None + self._full_visible_index_table = None + if ( + type(self) is AscendSFAIndexerBackend + and get_ascend_config().enable_sfa_full_visible_index_bypass + and self.allow_short_prefill_indexer_scoring_skip + and not get_ascend_config().enable_sparse_sfa_c8 + and not self.enable_sparse_li_c8 + and not self._pcp_active + and not self._dsa_cp_active + and not self._speculative_active + and self.topk_tokens == SFA_INDEXER_SPARSE_COUNT + and torch.npu.is_available() + ): + device = torch.device("npu", torch.npu.current_device()) + self._full_visible_index_table = self._get_or_create_full_visible_index_table(device) + + @classmethod + def _get_or_create_full_visible_index_table(cls, device: torch.device) -> torch.Tensor: + table = cls._full_visible_index_tables.get(device) + if table is None: + order = ( + np.arange(SFA_INDEXER_SPARSE_COUNT, dtype=np.int32) + .reshape(-1, SFA_FULL_VISIBLE_TEMPLATE_BLOCK_SIZE) + .T.reshape(-1) + ) + host_table = np.full( + (SFA_INDEXER_SPARSE_COUNT + 1, SFA_INDEXER_SPARSE_COUNT), + -1, + dtype=np.int32, + ) + for visible_len in range(1, SFA_INDEXER_SPARSE_COUNT + 1): + visible = order[order < visible_len] + host_table[visible_len, :visible_len] = visible + table = torch.from_numpy(host_table).to(device=device) + cls._full_visible_index_tables[device] = table + return table + + def _get_full_visible_topk_indices( + self, + metadata: AscendSFAIndexerMetadata, + num_tokens: int, + device: torch.device, + ) -> torch.Tensor | None: + if ( + type(self) is not AscendSFAIndexerBackend + or not self.allow_short_prefill_indexer_scoring_skip + or not get_ascend_config().enable_sfa_full_visible_index_bypass + or get_ascend_config().enable_sparse_sfa_c8 + or self.enable_sparse_li_c8 + or self._pcp_active + or self._dsa_cp_active + or self._speculative_active + or self.topk_tokens != SFA_INDEXER_SPARSE_COUNT + or device.type != "npu" + or metadata.attn_state not in (AscendAttentionState.PrefillNoCache, AscendAttentionState.PrefillCacheHit) + or metadata.num_decode_tokens != 0 + or metadata.block_size != SFA_FULL_VISIBLE_TEMPLATE_BLOCK_SIZE + or metadata.seq_lens_cpu is None + or metadata.seq_lens_cpu.device.type != "cpu" + or metadata.seq_lens_cpu.dtype not in (torch.int32, torch.int64) + or metadata.seq_lens_cpu.ndim != 1 + or metadata.seq_lens_cpu.numel() != 1 + or metadata.block_table.ndim != 2 + or metadata.block_table.shape[0] != 1 + or metadata.seq_lens.numel() != 1 + or metadata.cum_query_lens.numel() != 1 + or metadata.num_actual_tokens != num_tokens + or num_tokens <= 0 + or self._full_visible_index_table is None + or self._full_visible_index_table.device != device + ): + return None + kv_length = int(metadata.seq_lens_cpu[0]) + context_length = kv_length - num_tokens + if context_length < 0 or kv_length > SFA_INDEXER_SPARSE_COUNT: + return None + if metadata.block_table.shape[1] * metadata.block_size < kv_length: + return None + if metadata.attn_state == AscendAttentionState.PrefillNoCache and context_length != 0: + return None + # Zero-copy, read-only view. SFA consumes final topk; cache destinations + # must remain separately allocated and must not mutate this table. + return self._full_visible_index_table[context_length + 1 : kv_length + 1].unsqueeze(1) def process_weights_after_loading(self) -> None: if self.enable_sparse_li_c8 and AscendSFAIndexerBackend.q_hadamard is None: @@ -383,6 +485,11 @@ def forward( self.write_cache(k_li, k_li_scale, slot_mapping, indexer_attn_metadata=indexer_metadata) if not compute_topk: return None + topk_indices = self._get_full_visible_topk_indices( + indexer_metadata, hidden_states.shape[0], hidden_states.device + ) + if topk_indices is not None: + return topk_indices assert self.wk_weights_proj is not None assert self.wq_b is not None @@ -519,7 +626,15 @@ def build( block_size, ) + seq_lens_cpu = getattr(common_attn_metadata, "_seq_lens_cpu", None) + if seq_lens_cpu is None: + seq_lens_cpu = getattr(common_attn_metadata, "seq_lens_cpu", None) + if seq_lens_cpu is not None: + seq_lens_cpu = seq_lens_cpu[:num_reqs] + return AscendSFAIndexerMetadata( + seq_lens_cpu=seq_lens_cpu, + attn_state=getattr(common_attn_metadata, "attn_state", None), num_actual_tokens=common_attn_metadata.num_actual_tokens, slot_mapping=slot_mapping, seq_lens=common_attn_metadata.seq_lens[:num_reqs], diff --git a/vllm_ascend/models/glm5next/attention.py b/vllm_ascend/models/glm5next/attention.py index 1cb56e16382..60e69e1be6d 100644 --- a/vllm_ascend/models/glm5next/attention.py +++ b/vllm_ascend/models/glm5next/attention.py @@ -223,6 +223,7 @@ def __init__( cache_config: CacheConfig | None = None, quant_config: QuantizationConfig | None = None, prefix: str = "", + allow_short_prefill_indexer_scoring_skip: bool = False, topk_indices_buffer: torch.Tensor | None = None, input_size: int | None = None, skip_rope: bool | None = False, @@ -382,6 +383,7 @@ def __init__( quant_config, prefix, skip_topk=False, + allow_short_prefill_indexer_scoring_skip=allow_short_prefill_indexer_scoring_skip, fuse_qkv_rmsnorm=True, ) # The pluggable MLA wrapper owns the AttentionLayerBase registered in diff --git a/vllm_ascend/models/glm5next/model.py b/vllm_ascend/models/glm5next/model.py index 6bb93a5d342..e91458c7855 100644 --- a/vllm_ascend/models/glm5next/model.py +++ b/vllm_ascend/models/glm5next/model.py @@ -317,6 +317,7 @@ def __init__( cache_config=cache_config, quant_config=None, # MLA projections are BF16 in checkpoint prefix=f"{prefix}.self_attn", + allow_short_prefill_indexer_scoring_skip=not is_mtp_layer, topk_indices_buffer=topk_indices_buffer, skip_rope=config.mla_nope, ) diff --git a/vllm_ascend/ops/mla.py b/vllm_ascend/ops/mla.py index 465996d5c45..ea03dc9fd67 100644 --- a/vllm_ascend/ops/mla.py +++ b/vllm_ascend/ops/mla.py @@ -44,7 +44,12 @@ class IndexerWrapper(nn.Module): ``AscendSFAIndexerBackend`` instance it owns. """ - def __init__(self, vllm_indexer: nn.Module, qk_rope_head_dim: int) -> None: + def __init__( + self, + vllm_indexer: nn.Module, + qk_rope_head_dim: int, + allow_short_prefill_indexer_scoring_skip: bool = False, + ) -> None: super().__init__() # Register the indexer weights directly on the wrapper so module-tree # paths keep the pre-backend layout ("...indexer.") that weight @@ -66,7 +71,14 @@ def __init__(self, vllm_indexer: nn.Module, qk_rope_head_dim: int) -> None: backend_factory = getattr(type(vllm_indexer), "get_ascend_indexer_backend_cls", None) backend_cls = backend_factory(vllm_indexer) if backend_factory is not None else AscendSFAIndexerBackend - self.impl = backend_cls(vllm_indexer, qk_rope_head_dim) + if backend_cls is AscendSFAIndexerBackend: + self.impl = backend_cls( + vllm_indexer, + qk_rope_head_dim, + allow_short_prefill_indexer_scoring_skip=allow_short_prefill_indexer_scoring_skip, + ) + else: + self.impl = backend_cls(vllm_indexer, qk_rope_head_dim) # Interface consumed by the SFA impl - delegated to the backend impl. @property @@ -177,14 +189,15 @@ def __init__( # so only the backing value is stored here. MLAAttention below receives # the same value and initializes the impl consistently. self.skip_topk = skip_topk - # This is an upstream CUDA indexer hint. Ascend accepts it to preserve - # constructor compatibility, but its indexer does not consume it. - del allow_short_prefill_indexer_scoring_skip hf_config = get_current_vllm_config().model_config.hf_text_config self.tp_size = get_tensor_model_parallel_world_size() self.layers = hf_config.num_hidden_layers if mla_modules.indexer is not None: - ascend_indexer = IndexerWrapper(mla_modules.indexer, self.qk_rope_head_dim) + ascend_indexer = IndexerWrapper( + mla_modules.indexer, + self.qk_rope_head_dim, + allow_short_prefill_indexer_scoring_skip=allow_short_prefill_indexer_scoring_skip, + ) else: ascend_indexer = None self.mla_attn = MLAAttention( diff --git a/vllm_ascend/patch/worker/patch_deepseek_v2.py b/vllm_ascend/patch/worker/patch_deepseek_v2.py index 8eed43d6617..e0460a2e391 100644 --- a/vllm_ascend/patch/worker/patch_deepseek_v2.py +++ b/vllm_ascend/patch/worker/patch_deepseek_v2.py @@ -54,6 +54,16 @@ def _should_skip_indexer_init( return isinstance(indexer_type, str) and indexer_type.lower() == "shared" +def _is_mtp_layer(config: DeepseekV2Config | DeepseekV3Config, prefix: str) -> bool: + layer_id = extract_layer_index(prefix) + num_hidden_layers = getattr(config, "num_hidden_layers", None) + return num_hidden_layers is not None and layer_id >= num_hidden_layers + + +def _resolve_mtp_indexer_permissions(skip_topk: bool, is_mtp_layer: bool) -> tuple[bool, bool]: + return skip_topk and not is_mtp_layer, not is_mtp_layer + + def _deepseek_v2_mla_attention_init( self, vllm_config: VllmConfig, @@ -214,6 +224,8 @@ def _deepseek_v2_mla_attention_init( layer_id = extract_layer_index(prefix) + is_mtp_layer = _is_mtp_layer(config, prefix) + _skip_topk = False if _index_topk_pattern is None: _skip_topk = ( max( @@ -226,14 +238,7 @@ def _deepseek_v2_mla_attention_init( elif 0 <= layer_id < len(_index_topk_pattern): _skip_topk = _index_topk_pattern[layer_id] == "S" - # The skip pattern only governs backbone layers. MTP/nextn layers - # (layer_id >= num_hidden_layers) must never start with skip_topk=True: - # they compute their own indices at draft step 0 and toggle at runtime - # via set_skip_topk (index_share_for_mtp_iteration). Matches upstream - # deepseek_v2.py behavior. - num_hidden_layers = getattr(config, "num_hidden_layers", None) - is_mtp_layer = num_hidden_layers is not None and layer_id >= num_hidden_layers - + _skip_topk, allow_short_prefill_indexer_scoring_skip = _resolve_mtp_indexer_permissions(_skip_topk, is_mtp_layer) skip_indexer_init = _should_skip_indexer_init(config, prefix, _skip_topk) if self.is_v32 and not skip_indexer_init: self.indexer_rope_emb = get_rope( @@ -291,7 +296,8 @@ def _deepseek_v2_mla_attention_init( cache_config, quant_config, prefix, - skip_topk=_skip_topk and not is_mtp_layer, + skip_topk=_skip_topk, + allow_short_prefill_indexer_scoring_skip=allow_short_prefill_indexer_scoring_skip, )