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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
23 changes: 23 additions & 0 deletions docs/source/user_guide/configuration/additional_config.md
Original file line number Diff line number Diff line change
Expand Up @@ -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**
Expand Down
220 changes: 220 additions & 0 deletions tests/ut/attention/test_indexer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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
38 changes: 37 additions & 1 deletion tests/ut/ops/test_mla.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)

Expand All @@ -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):
Expand Down Expand Up @@ -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")
Expand Down
16 changes: 15 additions & 1 deletion tests/ut/patch/worker/test_patch_deepseek_v2.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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)
12 changes: 12 additions & 0 deletions tests/ut/test_ascend_config.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down
1 change: 1 addition & 0 deletions vllm_ascend/ascend_config.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P2 Badge Document the new bypass configuration

When users want to opt into this user-facing optimization, a repo-wide search finds this switch only in implementation and unit-test code, with no entry in docs/source/user_guide/configuration/additional_config.md or the applicable model guides. Add it to the configuration reference and document its single-sequence, 2,048-token, non-C8 limitations and additional device-memory allocation so users can enable it safely.

AGENTS.md reference: AGENTS.md:L367-L371

Useful? React with 👍 / 👎.

Copy link
Copy Markdown
Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Agreed — this is a user-facing opt-in and the current configuration reference is missing an entry.

The entry needs to cover the default False, an --additional-config example, the single-sequence prefill restriction, the 2,048-token limit on cached context plus new tokens, normal-layer permission, decode/MTP/C8 and other unsupported-path fallbacks, and the persistent per-process/device table allocation (16,785,408 bytes, approximately 16.008 MiB), shared across layers rather than allocated per request.

This remains a documentation follow-up at the current PR HEAD; it does not require changing the runtime algorithm.

Copy link
Copy Markdown
Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Addressed in the current rebased head 180baa0795cba9801a9a9953017f5fd3e1ac1e70. docs/source/user_guide/configuration/additional_config.md now documents enable_sfa_full_visible_index_bypass, including the default-off opt-in, single-sequence prefill eligibility, total visible KV <= 2048 (including cached context), MTP/C8/CP/unsupported-mode fallbacks, and the shared per-process/device table memory cost. The Markdown link and Read the Docs checks pass on this head.

# 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
Expand Down
Loading
Loading