Skip to content

Commit b417b95

Browse files
authored
[Feature][Test] Enable "c8_enable_reshape_optim" for A5 and add ut coverage (vllm-project#11451)
### What this PR does / why we need it? This PR addresses two items: 1. **Add unit test for the `c8_enable_reshape_optim = True` code path** When `c8_enable_reshape_optim` is `True`, `AscendSFAMetadataBuilder._build()` reads `slot_mapping_cpu`, invokes `torch.ops._C_ascend.store_kv_block_pre`, and populates the returned `group_len`, `group_key_idx`, and `group_key_cache_idx` into the metadata. This branch previously had no UT coverage. The new test verifies: - `slot_mapping_cpu` is correctly sliced and forwarded to the op - The op is called with the expected arguments (`slot_mapping`, `slot_mapping_cpu.tolist()`, `block_size`) - The return values are correctly propagated to `metadata.group_len` / `group_key_idx` / `group_key_cache_idx` - `block_size=128` behavior is as expected The test runs under a pure TP setup without DCP/PCP group dependencies. 2. **Enable `store_kv_block` custom operator compilation for A5** Adds `store_kv_block` to the Ascend950 (A5) `CUSTOM_OPS_ARRAY` in `csrc/build_aclnn.sh` so the operator compiles correctly on A5 hardware, supporting downstream E2E testing and inference. ### Does this PR introduce _any_ user-facing change? No. Test and build configuration changes only. ### How was this patch tested? - UT: `pytest -sv tests/ut/attention/a2/test_sfa_v1.py::TestAscendSFAMetadataBuilder::test_ascend_sfa_metadata_builder_build_with_c8_reshape_optim` - vLLM version: v0.23.0 - vLLM main: vllm-project/vllm@1f486d9 Signed-off-by: Li Jiahang <216526138+lijiahang226@users.noreply.github.com>
1 parent ee2a613 commit b417b95

2 files changed

Lines changed: 84 additions & 0 deletions

File tree

csrc/build_aclnn.sh

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -217,6 +217,7 @@ elif [[ "$SOC_VERSION" =~ ^ascend950 ]]; then
217217
"recurrent_gated_delta_rule"
218218
"chunk_fwd_o"
219219
"chunk_gated_delta_rule_fwd_h"
220+
"store_kv_block"
220221
)
221222

222223
CUSTOM_OPS=$(IFS=';'; echo "${CUSTOM_OPS_ARRAY[*]}")

tests/ut/attention/a2/test_sfa_v1.py

Lines changed: 83 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -308,3 +308,86 @@ def test_ascend_sfa_metadata_builder_build_for_graph_capture(
308308

309309
assert isinstance(attn_metadata, AscendSFAMetadata)
310310
assert attn_metadata.attn_state == AscendAttentionState.DecodeOnly
311+
312+
@patch("vllm_ascend.attention.sfa_v1.get_current_vllm_config")
313+
@patch("vllm_ascend.attention.sfa_v1.get_cos_and_sin_mla")
314+
@patch("vllm_ascend.attention.sfa_v1.enable_dsa_cp", return_value=False)
315+
@patch("torch.ops._C_ascend.store_kv_block_pre", create=True)
316+
def test_ascend_sfa_metadata_builder_build_with_c8_reshape_optim(
317+
self,
318+
mock_store_kv_block_pre,
319+
mock_enable_dsa_cp,
320+
mock_get_cos_and_sin_mla,
321+
mock_get_current_vllm_config,
322+
):
323+
cfg = MagicMock()
324+
cfg.model_config = MagicMock()
325+
cfg.model_config.hf_text_config = MagicMock()
326+
327+
mock_get_current_vllm_config.return_value = cfg
328+
kv_cache_spec = MagicMock()
329+
layer_names = ["layer1", "layer2"]
330+
vllm_config = MagicMock()
331+
vllm_config.cache_config.block_size = 16
332+
vllm_config.model_config.max_model_len = 1024
333+
vllm_config.model_config.get_head_size.return_value = 64
334+
vllm_config.model_config.dtype = torch.float16
335+
vllm_config.model_config.hf_text_config.qk_rope_head_dim = 64
336+
speculative_config = MagicMock()
337+
speculative_config.num_speculative_tokens = 4
338+
vllm_config.speculative_config = speculative_config
339+
device = torch.device("cpu")
340+
341+
builder = AscendSFAMetadataBuilder(
342+
kv_cache_spec=kv_cache_spec, layer_names=layer_names, vllm_config=vllm_config, device=device
343+
)
344+
345+
slot_mapping_cpu = torch.randint(0, 10000, (100,))
346+
347+
common_attn_metadata = MagicMock()
348+
common_attn_metadata.num_reqs = 10
349+
common_attn_metadata.num_actual_tokens = 100
350+
common_attn_metadata.query_start_loc = torch.tensor([0, 10, 20, 30, 40, 50, 60, 70, 80, 90])
351+
common_attn_metadata.query_start_loc_cpu = torch.tensor([0, 10, 20, 30, 40, 50, 60, 70, 80, 90])
352+
common_attn_metadata.slot_mapping = torch.randn(100, 4, 1024)
353+
common_attn_metadata.slot_mapping_cpu = slot_mapping_cpu
354+
common_attn_metadata.seq_lens_cpu = torch.tensor([2] * 10)
355+
common_attn_metadata.positions = torch.randn(100)
356+
common_attn_metadata.attn_mask = None
357+
common_attn_metadata.attn_state = AscendAttentionState.ChunkedPrefill
358+
common_attn_metadata.block_table_tensor = torch.randn(100, 4)
359+
common_attn_metadata.cos = None
360+
common_attn_metadata.sin = None
361+
common_attn_metadata.num_input_tokens = 100
362+
363+
mock_get_cos_and_sin_mla.return_value = (torch.randn(100), torch.randn(100))
364+
365+
mock_group_len = torch.tensor([1, 2, 3])
366+
mock_group_key_idx = torch.tensor([0, 1, 2])
367+
mock_group_key_cache_idx = torch.tensor([4, 5, 6])
368+
mock_store_kv_block_pre.return_value = (mock_group_len, mock_group_key_idx, mock_group_key_cache_idx)
369+
370+
with patch("vllm_ascend.attention.sfa_v1.get_ascend_config") as mock_get_ascend_config:
371+
mock_ascend_config = MagicMock()
372+
mock_ascend_config.c8_enable_reshape_optim = True
373+
mock_get_ascend_config.return_value = mock_ascend_config
374+
375+
metadata = builder.build(
376+
common_prefix_len=10,
377+
common_attn_metadata=common_attn_metadata,
378+
)
379+
380+
assert isinstance(metadata, AscendSFAMetadata)
381+
assert metadata.num_actual_tokens == common_attn_metadata.num_actual_tokens
382+
assert metadata.slot_mapping.shape == (100, 4, 1024)
383+
384+
mock_store_kv_block_pre.assert_called_once()
385+
actual_args, _ = mock_store_kv_block_pre.call_args
386+
assert torch.equal(actual_args[0], common_attn_metadata.slot_mapping)
387+
assert actual_args[1] == slot_mapping_cpu.tolist()
388+
assert actual_args[2] == 128
389+
390+
assert metadata.block_size == 128
391+
assert metadata.group_len is mock_group_len
392+
assert metadata.group_key_idx is mock_group_key_idx
393+
assert metadata.group_key_cache_idx is mock_group_key_cache_idx

0 commit comments

Comments
 (0)