[Feature][Ops] Add Triton KeyPool compression and pooled indexing - #16243
Conversation
Summary of ChangesHello, I'm Gemini Code Assist1! I'm currently reviewing this pull request and will post my feedback shortly. In the meantime, here's a summary to help you and other reviewers quickly get up to speed! This pull request implements high-performance Triton operators to handle state compression and pooled indexing for the GLM-5.3-Flash model. By moving these operations into Triton, the implementation avoids the overhead of multiple small torch operations. Additionally, the PR enhances the IndexerWrapper to allow for flexible, model-provided backend selection, ensuring that parameter names remain consistent across different model architectures. Highlights
New Features🧠 You can now enable Memory (public preview) to help Gemini Code Assist learn from your team's feedback. This makes future code reviews more consistent and personalized to your project's style. Click here to enable Memory in your admin console. Using Gemini Code AssistThe full guide for Gemini Code Assist can be found on our documentation page, here are some quick tips. Invoking Gemini You can request assistance from Gemini at any point by creating a comment using either
Customization To customize the Gemini Code Assist for GitHub experience, repository maintainers can create a configuration file and/or provide a custom code review style guide (such as PEP-8 for Python) by creating and adding files to a Limitations & Feedback Gemini Code Assist may make mistakes. Please leave feedback on any instances where its feedback is incorrect or counterproductive. You can react with 👍 and 👎 on @gemini-code-assist comments. If you're interested in giving your feedback about your experience with Gemini Code Assist for GitHub and other Google products, sign up here. Footnotes
|
|
👋 Hi! Thank you for contributing to the vLLM Ascend project. The following points will speed up your PR merge:
If CI fails, you can run linting and testing checks locally according Contributing and Testing. |
There was a problem hiding this comment.
Code Review
Suggested PR Title:
[Ops][Feature] Add Triton kernels and backend wrapper for GLM5 Next KPool indexerSuggested PR Summary:
### What this PR does / why we need it?
This pull request introduces Triton fast paths for the GLM5 Next KPool indexer. Specifically, it adds:
1. `glm5_next_kpool_state_compress_and_write_cache_triton` for the pre-compress sequence, which gathers and compresses state windows using softmax and writes them to the paged indexer cache.
2. `glm5_next_lightning_indexer_triton` for scoring compressed pools against head-weighted queries.
3. Updates to `IndexerWrapper` in `vllm_ascend/ops/mla.py` to dynamically select the appropriate backend and preserve checkpoint-visible parameters.
4. Corresponding unit and end-to-end tests.
Feedback:
The review comments identify critical issues in the Triton kernels where out-of-bounds address calculations (e.g., loading from `cum_query_lens_ptr + req_offsets` when `req_offsets >= num_reqs`, or indexing `state_block_table_ptr` and `indexer_block_table_ptr` with out-of-bounds pages) can cause GPU page faults or undefined behavior, even when masked. Clamping these offsets/indices using `tl.where` before address calculation is highly recommended.
### Does this PR introduce _any_ user-facing change?
No, these are internal performance optimizations and backend support for GLM5 Next.
### How was this patch tested?
New unit tests (`test_glm5next_indexer_backend.py`) and end-to-end Triton kernel tests (`test_glm5next_kpool_triton.py`, `test_glm5next_pool_key_indexer_triton.py`) have been added.| req_offsets = tl.arange(0, REQ_POW2) | ||
| query_ends = tl.load( | ||
| cum_query_lens_ptr + req_offsets, | ||
| mask=req_offsets < num_reqs, | ||
| other=2147483647, | ||
| ) |
There was a problem hiding this comment.
In Triton, loading from an out-of-bounds address can cause GPU page faults or undefined behavior, even if the load is masked. Since cum_query_lens is not padded to REQ_POW2, cum_query_lens_ptr + req_offsets will access out-of-bounds memory when req_offsets >= num_reqs. We should clamp req_offsets to a safe index (e.g., 0) using tl.where before performing the load.
| req_offsets = tl.arange(0, REQ_POW2) | |
| query_ends = tl.load( | |
| cum_query_lens_ptr + req_offsets, | |
| mask=req_offsets < num_reqs, | |
| other=2147483647, | |
| ) | |
| req_offsets = tl.arange(0, REQ_POW2) | |
| safe_req_offsets = tl.where(req_offsets < num_reqs, req_offsets, 0) | |
| query_ends = tl.load( | |
| cum_query_lens_ptr + safe_req_offsets, | |
| mask=req_offsets < num_reqs, | |
| other=2147483647, | |
| ) |
| physical = tl.load( | ||
| state_block_table_ptr + req_id * state_block_table_stride_req + page * state_block_table_stride_page, | ||
| mask=history_valid, | ||
| other=-1, | ||
| ).to(tl.int64) |
There was a problem hiding this comment.
In Triton, even if a load is masked, the address calculation must not point to invalid memory to avoid page faults. When page >= state_max_pages, the address state_block_table_ptr + ... + page * state_block_table_stride_page is out of bounds. We should clamp page to a safe index (e.g., 0) using tl.where before the address calculation.
| physical = tl.load( | |
| state_block_table_ptr + req_id * state_block_table_stride_req + page * state_block_table_stride_page, | |
| mask=history_valid, | |
| other=-1, | |
| ).to(tl.int64) | |
| safe_page = tl.where(history_valid, page, 0) | |
| physical = tl.load( | |
| state_block_table_ptr + req_id * state_block_table_stride_req + safe_page * state_block_table_stride_page, | |
| mask=history_valid, | |
| other=-1, | |
| ).to(tl.int64) |
| req_offsets = tl.arange(0, REQ_POW2) | ||
| query_ends = tl.load(cum_query_lens_ptr + req_offsets, mask=req_offsets < num_reqs, other=2147483647) |
There was a problem hiding this comment.
In Triton, loading from an out-of-bounds address can cause GPU page faults or undefined behavior, even if the load is masked. Since cum_query_lens is not padded to REQ_POW2, cum_query_lens_ptr + req_offsets will access out-of-bounds memory when req_offsets >= num_reqs. We should clamp req_offsets to a safe index (e.g., 0) using tl.where before performing the load.
| req_offsets = tl.arange(0, REQ_POW2) | |
| query_ends = tl.load(cum_query_lens_ptr + req_offsets, mask=req_offsets < num_reqs, other=2147483647) | |
| req_offsets = tl.arange(0, REQ_POW2) | |
| safe_req_offsets = tl.where(req_offsets < num_reqs, req_offsets, 0) | |
| query_ends = tl.load(cum_query_lens_ptr + safe_req_offsets, mask=req_offsets < num_reqs, other=2147483647) |
| physical_blocks = tl.load( | ||
| indexer_block_table_ptr + req_id * block_table_stride_req + logical_pages * block_table_stride_page, | ||
| mask=in_range, | ||
| other=0, | ||
| ).to(tl.int64) |
There was a problem hiding this comment.
In Triton, even if a load is masked, the address calculation must not point to invalid memory to avoid page faults. When logical_pages is out of bounds (e.g., when in_range is False), the address calculation for indexer_block_table_ptr can access out-of-bounds memory. We should clamp logical_pages to a safe index (e.g., 0) using tl.where before the address calculation.
| physical_blocks = tl.load( | |
| indexer_block_table_ptr + req_id * block_table_stride_req + logical_pages * block_table_stride_page, | |
| mask=in_range, | |
| other=0, | |
| ).to(tl.int64) | |
| safe_logical_pages = tl.where(in_range, logical_pages, 0) | |
| physical_blocks = tl.load( | |
| indexer_block_table_ptr + req_id * block_table_stride_req + safe_logical_pages * block_table_stride_page, | |
| mask=in_range, | |
| other=0, | |
| ).to(tl.int64) |
acf66c1 to
73c034d
Compare
ZT-AIA
left a comment
There was a problem hiding this comment.
Please supplement the single-operator use cases and the operator documentation.
6bb0374 to
bcaf366
Compare
Implement KeyPool state compression and pooled-key selection using the existing paged cache slots. Reuse the shared next_power_of_2 helper and preserve checkpoint parameter names when selecting a model backend. Add single-operator accuracy tests for paging, historical windows, rollback, causal tails, and graph replay, plus operator contracts and documentation. Signed-off-by: Li Jiahang <216526138+lijiahang226@users.noreply.github.com>
bcaf366 to
1b5ceb9
Compare
Operation ducumentation has been added. |
…16253) ### What this PR does / why we need it? Refs #15665 Connect the GLM-5.3-Flash pooled indexer to `IndexerWrapper` and the shared NoPE SFA backend. Keep cache metadata and the model-side indexer backend together in `vllm_ascend/attention/indexer_kpool.py`. Load KPool execution dependencies when the backend is instantiated to avoid circular imports during metadata-only startup. Keep hardware and PCP/DCP capability checks in backend initialization. Pass complete compressed physical pages and compressor-state metadata to the Triton operators while preserving causal-tail visibility and the cache allocation/grouping contract from #15913. Dependencies: - #16243 provides the Triton operators and indexer backend-selection hook. - #16252 provides shared NoPE SFA execution and metadata. Prerequisite commits are combined only for integration validation and are excluded from this branch. ### Does this PR introduce _any_ user-facing change? Yes. GLM-5.3-Flash routes pooled sparse attention through the shared SFA backend and Triton indexer. Unsupported hardware and PCP/DCP configurations are rejected when the backend is initialized. ### How was this patch tested? At PR revision `1ff19ad48`, combined with its integration prerequisites: - Fresh-process imports of the consolidated metadata module and model backend factory passed. Regression tests cover importing metadata without generic SFA indexing or KPool execution dependencies. - 120 unit tests passed, 1 skipped, covering `test_glm5next_kpool_model_backend.py`, `test_glm5next_kv_cache.py`, `test_glm5next_sfa_routing.py`, `test_glm5next_indexer_backend.py`, `test_platform.py`, and `attention/test_indexer.py`. - Six native NPU comparisons against the pre-consolidation backend produced identical indices and compressed/state caches: prefill tail, decode tail, pool-boundary decode, prefill top-k selection, FULL-mode padding, and cache-only updates. - GitHub CI passed pre-commit, CPU unit tests, selected device tests, and the CI gate. Earlier full-model validation at PR revision `0bff0eb46`, combined with integration prerequisites, passed 320 unit tests (1 skipped) and 8/8 GLM-5.3-Flash-w8a8 inference cases with TP8/EP8, FULL_DECODE_ONLY, and MTP=3. Those results apply to the pre-consolidation revision; the six native comparisons above do not constitute a new full-model or ACLGraph capture/replay run. - vLLM main: vllm-project/vllm@b2f6858 --------- Signed-off-by: Li Jiahang <216526138+lijiahang226@users.noreply.github.com>
…hado/vllm-ascend into main_fix_mrv2_eagle3_mamba * 'main_fix_mrv2_eagle3_mamba' of https://github.com/windshado/vllm-ascend: (42 commits) Update vllm_ascend/worker/v2/model_states/mamba_hybrid.py [Feature][Kimi K3 DSPark] Enable TP for context_proj (vllm-project#16344) [BugFix][SpecDecode] Refresh replicated PCP draft graph cache mappings (vllm-project#16300) [Feature][Model] Integrate Triton KeyPool indexing for GLM-5.3-Flash (vllm-project#16253) [BugFix][Offloader] Re-bind params to NZ static buffers after npu_format_cast (vllm-project#15415) [Feature][Model] Integrate AscendC KDA and causal convolution for GLM-5.3-Flash (vllm-project#16251) [Performance][Communicator] Replace per-layer F.pad with cat of a persistent zero block in MoE prepare (vllm-project#16343) [Feature][Operator] Add DeepSeek V4.1 sparse attention operators (vllm-project#16422) [Doc][Misc] Document batch invariance scheduling limitations (vllm-project#16232) [CI][MRV2] Enable mrv2 dspark e2e test (vllm-project#16319) [BugFix] Precast MoE gate weight_fp32 to avoid aclop Cast (vllm-project#16189) [Feature][MRV2][310P] MRv2 adapting MTP on the 310P for Qwen3.5 (vllm-project#16043) [Revert] Revert "[Feature][MRV1][MRV2] Refactor Host-Side Parameter Updates for ACL Graph Replay." (vllm-project#15908) (vllm-project#16409) [Feature][Ops] Add Triton KeyPool compression and pooled indexing (vllm-project#16243) [Feature][Attention] Support NoPE in the shared SFA backend (vllm-project#16252) [Performance][Model] Reuse fused mHC operators for GLM-5.3-Flash (vllm-project#16321) [Feature][Model] Enable MiniMax-M3 FP8 MSA index score on A5 (vllm-project#15918) [Performance][KDA] Reduce preprocessing copies and redundant output masks (vllm-project#16067) [Feature][Model][MTP] Support speculative decoding for GLM-5.3-Flash (vllm-project#16214) [BugFix][Model] Skip unused hash-router bias when loading DeepSeek-V4 weights (vllm-project#16259) ...
…lm-project#16243) ### What this PR does / why we need it? Refs vllm-project#15665 Add Triton KeyPool state compression, paged cache writes, and pooled-key selection for GLM-5.3-Flash. Allow `IndexerWrapper` to select a model-provided backend while preserving checkpoint parameter names. Both operators reuse the shared `next_power_of_2` helper. Include single-operator accuracy tests and documentation covering formulas, parameters, layout constraints, graph replay, and test commands. Model-specific backend integration is handled separately. ### Does this PR introduce _any_ user-facing change? Yes. Provides the Triton pooled-indexing operators and a backend-selection hook. This PR does not enable the model backend by itself or change cache allocation and grouping. ### How was this patch tested? Validated on Atlas A3 with PyTorch 2.10.0, torch-npu 2.10.0.post4, and Triton-Ascend 3.2.0: - 10 tests passed in `tests/e2e/nightly/single_node/ops/singlecard_ops/triton/test_glm5next_kpool_triton.py`: CPU-reference accuracy, historical windows, rollback, paging, noncontiguous storage, invalid slots, graph padding, empty inputs, and eager/graph replay. - 9 tests passed in `tests/e2e/nightly/single_node/ops/singlecard_ops/triton/test_glm5next_pool_key_indexer_triton.py`: multiple requests, pool capacities, paged caches, token chunking, causal tails, and eager/graph replay with changing inputs. - 16 tests passed across `tests/ut/models/test_glm5next_indexer_backend.py` and `tests/ut/ops/test_mla.py`. The CI follow-up adds a type annotation to the test reference data without changing test behavior. Hardware coverage is limited to Atlas A3; these tests establish operator accuracy, not model-level throughput. - vLLM main: vllm-project/vllm@a97dacb Signed-off-by: Li Jiahang <216526138+lijiahang226@users.noreply.github.com> Signed-off-by: tianming2009 <13246728590@163.com>
…llm-project#16253) ### What this PR does / why we need it? Refs vllm-project#15665 Connect the GLM-5.3-Flash pooled indexer to `IndexerWrapper` and the shared NoPE SFA backend. Keep cache metadata and the model-side indexer backend together in `vllm_ascend/attention/indexer_kpool.py`. Load KPool execution dependencies when the backend is instantiated to avoid circular imports during metadata-only startup. Keep hardware and PCP/DCP capability checks in backend initialization. Pass complete compressed physical pages and compressor-state metadata to the Triton operators while preserving causal-tail visibility and the cache allocation/grouping contract from vllm-project#15913. Dependencies: - vllm-project#16243 provides the Triton operators and indexer backend-selection hook. - vllm-project#16252 provides shared NoPE SFA execution and metadata. Prerequisite commits are combined only for integration validation and are excluded from this branch. ### Does this PR introduce _any_ user-facing change? Yes. GLM-5.3-Flash routes pooled sparse attention through the shared SFA backend and Triton indexer. Unsupported hardware and PCP/DCP configurations are rejected when the backend is initialized. ### How was this patch tested? At PR revision `1ff19ad48`, combined with its integration prerequisites: - Fresh-process imports of the consolidated metadata module and model backend factory passed. Regression tests cover importing metadata without generic SFA indexing or KPool execution dependencies. - 120 unit tests passed, 1 skipped, covering `test_glm5next_kpool_model_backend.py`, `test_glm5next_kv_cache.py`, `test_glm5next_sfa_routing.py`, `test_glm5next_indexer_backend.py`, `test_platform.py`, and `attention/test_indexer.py`. - Six native NPU comparisons against the pre-consolidation backend produced identical indices and compressed/state caches: prefill tail, decode tail, pool-boundary decode, prefill top-k selection, FULL-mode padding, and cache-only updates. - GitHub CI passed pre-commit, CPU unit tests, selected device tests, and the CI gate. Earlier full-model validation at PR revision `0bff0eb46`, combined with integration prerequisites, passed 320 unit tests (1 skipped) and 8/8 GLM-5.3-Flash-w8a8 inference cases with TP8/EP8, FULL_DECODE_ONLY, and MTP=3. Those results apply to the pre-consolidation revision; the six native comparisons above do not constitute a new full-model or ACLGraph capture/replay run. - vLLM main: vllm-project/vllm@b2f6858 --------- Signed-off-by: Li Jiahang <216526138+lijiahang226@users.noreply.github.com> Signed-off-by: tianming2009 <13246728590@163.com>
…lm-project#16243) ### What this PR does / why we need it? Refs vllm-project#15665 Add Triton KeyPool state compression, paged cache writes, and pooled-key selection for GLM-5.3-Flash. Allow `IndexerWrapper` to select a model-provided backend while preserving checkpoint parameter names. Both operators reuse the shared `next_power_of_2` helper. Include single-operator accuracy tests and documentation covering formulas, parameters, layout constraints, graph replay, and test commands. Model-specific backend integration is handled separately. ### Does this PR introduce _any_ user-facing change? Yes. Provides the Triton pooled-indexing operators and a backend-selection hook. This PR does not enable the model backend by itself or change cache allocation and grouping. ### How was this patch tested? Validated on Atlas A3 with PyTorch 2.10.0, torch-npu 2.10.0.post4, and Triton-Ascend 3.2.0: - 10 tests passed in `tests/e2e/nightly/single_node/ops/singlecard_ops/triton/test_glm5next_kpool_triton.py`: CPU-reference accuracy, historical windows, rollback, paging, noncontiguous storage, invalid slots, graph padding, empty inputs, and eager/graph replay. - 9 tests passed in `tests/e2e/nightly/single_node/ops/singlecard_ops/triton/test_glm5next_pool_key_indexer_triton.py`: multiple requests, pool capacities, paged caches, token chunking, causal tails, and eager/graph replay with changing inputs. - 16 tests passed across `tests/ut/models/test_glm5next_indexer_backend.py` and `tests/ut/ops/test_mla.py`. The CI follow-up adds a type annotation to the test reference data without changing test behavior. Hardware coverage is limited to Atlas A3; these tests establish operator accuracy, not model-level throughput. - vLLM main: vllm-project/vllm@a97dacb Signed-off-by: Li Jiahang <216526138+lijiahang226@users.noreply.github.com> Signed-off-by: like-0517 <ithwlike@126.com>
…llm-project#16253) ### What this PR does / why we need it? Refs vllm-project#15665 Connect the GLM-5.3-Flash pooled indexer to `IndexerWrapper` and the shared NoPE SFA backend. Keep cache metadata and the model-side indexer backend together in `vllm_ascend/attention/indexer_kpool.py`. Load KPool execution dependencies when the backend is instantiated to avoid circular imports during metadata-only startup. Keep hardware and PCP/DCP capability checks in backend initialization. Pass complete compressed physical pages and compressor-state metadata to the Triton operators while preserving causal-tail visibility and the cache allocation/grouping contract from vllm-project#15913. Dependencies: - vllm-project#16243 provides the Triton operators and indexer backend-selection hook. - vllm-project#16252 provides shared NoPE SFA execution and metadata. Prerequisite commits are combined only for integration validation and are excluded from this branch. ### Does this PR introduce _any_ user-facing change? Yes. GLM-5.3-Flash routes pooled sparse attention through the shared SFA backend and Triton indexer. Unsupported hardware and PCP/DCP configurations are rejected when the backend is initialized. ### How was this patch tested? At PR revision `1ff19ad48`, combined with its integration prerequisites: - Fresh-process imports of the consolidated metadata module and model backend factory passed. Regression tests cover importing metadata without generic SFA indexing or KPool execution dependencies. - 120 unit tests passed, 1 skipped, covering `test_glm5next_kpool_model_backend.py`, `test_glm5next_kv_cache.py`, `test_glm5next_sfa_routing.py`, `test_glm5next_indexer_backend.py`, `test_platform.py`, and `attention/test_indexer.py`. - Six native NPU comparisons against the pre-consolidation backend produced identical indices and compressed/state caches: prefill tail, decode tail, pool-boundary decode, prefill top-k selection, FULL-mode padding, and cache-only updates. - GitHub CI passed pre-commit, CPU unit tests, selected device tests, and the CI gate. Earlier full-model validation at PR revision `0bff0eb46`, combined with integration prerequisites, passed 320 unit tests (1 skipped) and 8/8 GLM-5.3-Flash-w8a8 inference cases with TP8/EP8, FULL_DECODE_ONLY, and MTP=3. Those results apply to the pre-consolidation revision; the six native comparisons above do not constitute a new full-model or ACLGraph capture/replay run. - vLLM main: vllm-project/vllm@b2f6858 --------- Signed-off-by: Li Jiahang <216526138+lijiahang226@users.noreply.github.com> Signed-off-by: like-0517 <ithwlike@126.com>
What this PR does / why we need it?
Refs #15665
Add Triton KeyPool state compression, paged cache writes, and pooled-key selection for GLM-5.3-Flash. Allow
IndexerWrapperto select a model-provided backend while preserving checkpoint parameter names. Both operators reuse the sharednext_power_of_2helper.Include single-operator accuracy tests and documentation covering formulas, parameters, layout constraints, graph replay, and test commands. Model-specific backend integration is handled separately.
Does this PR introduce any user-facing change?
Yes. Provides the Triton pooled-indexing operators and a backend-selection hook. This PR does not enable the model backend by itself or change cache allocation and grouping.
How was this patch tested?
Validated on Atlas A3 with PyTorch 2.10.0, torch-npu 2.10.0.post4, and Triton-Ascend 3.2.0:
tests/e2e/nightly/single_node/ops/singlecard_ops/triton/test_glm5next_kpool_triton.py: CPU-reference accuracy, historical windows, rollback, paging, noncontiguous storage, invalid slots, graph padding, empty inputs, and eager/graph replay.tests/e2e/nightly/single_node/ops/singlecard_ops/triton/test_glm5next_pool_key_indexer_triton.py: multiple requests, pool capacities, paged caches, token chunking, causal tails, and eager/graph replay with changing inputs.tests/ut/models/test_glm5next_indexer_backend.pyandtests/ut/ops/test_mla.py.The CI follow-up adds a type annotation to the test reference data without changing test behavior. Hardware coverage is limited to Atlas A3; these tests establish operator accuracy, not model-level throughput.