[Performance][Attention] Fuse SFA K-path ops and enable PROLOG_V3 fused decode for non-PD serving - #16328
Conversation
|
👋 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. |
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 several performance optimizations for the Sparse Flash Attention (SFA) K-path and expands the availability of the PROLOG_V3 fused decode preprocessing. By caching metadata conversions, reusing GEMM outputs, and eliminating redundant memory copies, the changes reduce kernel overhead per layer. Additionally, a new configuration option allows non-PD serving workers to utilize the PROLOG_V3 fused path, trading a small amount of memory for significant decode kernel savings. 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
|
There was a problem hiding this comment.
Code Review
This pull request introduces SFA PROLOG_V3 fused decode preprocessing support for plain serving via a new configuration option, alongside several K-path optimizations such as caching int64 slot conversions and reusing indexer weights to avoid redundant GEMM operations. Redundant .contiguous() calls have also been cleaned up.
I have no feedback to provide as there are no review comments.
Suggested PR Title:
[Attention][Feature] Enable SFA PROLOG_V3 fused decode preprocessing in plain serving and optimize K-path fusionsSuggested PR Summary:
### What this PR does / why we need it?
This PR introduces several optimizations and features for the Ascend SFA (Sparse Flash Attention) attention backend:
1. **SFA PROLOG_V3 Opt-in for Plain Serving**: Enables the `PROLOG_V3` fused decode preprocessing path (using a single `npu_mla_prolog_v3` op) outside of PD-disaggregated KV-consumer workers (i.e., in plain serving) via the new `enable_sfa_prolog_v3` configuration option.
2. **K-Path Fusions & Optimizations**:
- Caches the `int64` converted KV slot mapping once per scheduling step (`_int64_kv_slots`) to avoid redundant Cast kernels across layers.
- Reuses the non-K tail of the `wk_weights_proj` GEMM output (`indexer_weights`) in the indexer's top-k stage when the hidden states match, avoiding duplicate GEMM computations.
3. **Redundancy Cleanup**: Removes redundant `.contiguous()` calls after `torch.cat` and `npu_rms_norm` since these operations already return contiguous tensors.
### Does this PR introduce _any_ user-facing change?
Yes, it introduces a new configuration option `enable_sfa_prolog_v3` (boolean, default `False`) to allow plain serving to opt into the SFA PROLOG_V3 fused decode preprocessing path.
### How was this patch tested?
Added new unit tests in `tests/ut/attention/test_sfa_v1.py` covering:
- `_int64_kv_slots` caching and passthrough behavior.
- Reuse of `int64` slots across layers in `exec_kv`.
- Reuse of `wk_weights_proj` weights in the indexer.
- Path resolution logic for non-PD opt-in configurations.|
Thanks for the contribution. Could you provide a performance comparison between the MLA Prolog v3 path and DSACP under prefill / mixed-deployment (hybrid) scenarios? These two features currently conflict with each other. |
4422dc5 to
cb525a8
Compare
The cpu-ut CI on PR vllm-project#16328 failed on three tests: 1. tests/ut/ops/test_mla.py (2 tests): the forward_k mocks still return 2-tuples while forward now unpacks 3 values (k, scale, weights tail) - update the mocks to return 3-tuples. 2. tests/ut/attention/test_sfa_v1.py::test_indexer_forward_reuses_ wk_weights_proj: the test emulated the pre-squash calling pattern (forward_k called externally before forward). forward now calls forward_k internally, so the same sequence counts one extra GEMM. Rework the assertions for the current architecture: reset the GEMM mock before forward, expect exactly one call when both stages share the same hidden-states tensor (reuse), expect three calls (one in forward_k + one fallback) for a distinct top-k input, and check the weights tail passed to indexer_select_post_process by value instead of identity. Signed-off-by: huamus <1943805462@qq.com>
cb525a8 to
1ae8f00
Compare
The cpu-ut CI on PR vllm-project#16328 failed on three tests: 1. tests/ut/ops/test_mla.py (2 tests): the forward_k mocks still return 2-tuples while forward now unpacks 3 values (k, scale, weights tail) - update the mocks to return 3-tuples. 2. tests/ut/attention/test_sfa_v1.py::test_indexer_forward_reuses_ wk_weights_proj: the test emulated the pre-squash calling pattern (forward_k called externally before forward). forward now calls forward_k internally, so the same sequence counts one extra GEMM. Rework the assertions for the current architecture: reset the GEMM mock before forward, expect exactly one call when both stages share the same hidden-states tensor (reuse), expect three calls (one in forward_k + one fallback) for a distinct top-k input, and check the weights tail passed to indexer_select_post_process by value instead of identity. Signed-off-by: huamus <1943805462@qq.com>
|
This pull request has conflicts, please resolve those before we can evaluate the pull request. |
The cpu-ut CI on PR vllm-project#16328 failed on three tests: 1. tests/ut/ops/test_mla.py (2 tests): the forward_k mocks still return 2-tuples while forward now unpacks 3 values (k, scale, weights tail) - update the mocks to return 3-tuples. 2. tests/ut/attention/test_sfa_v1.py::test_indexer_forward_reuses_ wk_weights_proj: the test emulated the pre-squash calling pattern (forward_k called externally before forward). forward now calls forward_k internally, so the same sequence counts one extra GEMM. Rework the assertions for the current architecture: reset the GEMM mock before forward, expect exactly one call when both stages share the same hidden-states tensor (reuse), expect three calls (one in forward_k + one fallback) for a distinct top-k input, and check the weights tail passed to indexer_select_post_process by value instead of identity. Signed-off-by: huamus <1943805462@qq.com>
e54da6c to
aa90382
Compare
|
/rerun Rerun (failed jobs only):
|
|
This pull request has conflicts, please resolve those before we can evaluate the pull request. |
aa90382 to
28e9e11
Compare
The cpu-ut CI on PR vllm-project#16328 failed on three tests: 1. tests/ut/ops/test_mla.py (2 tests): the forward_k mocks still return 2-tuples while forward now unpacks 3 values (k, scale, weights tail) - update the mocks to return 3-tuples. 2. tests/ut/attention/test_sfa_v1.py::test_indexer_forward_reuses_ wk_weights_proj: the test emulated the pre-squash calling pattern (forward_k called externally before forward). forward now calls forward_k internally, so the same sequence counts one extra GEMM. Rework the assertions for the current architecture: reset the GEMM mock before forward, expect exactly one call when both stages share the same hidden-states tensor (reuse), expect three calls (one in forward_k + one fallback) for a distinct top-k input, and check the weights tail passed to indexer_select_post_process by value instead of identity. Signed-off-by: huamus <1943805462@qq.com>
|
This pull request has conflicts, please resolve those before we can evaluate the pull request. |
28e9e11 to
f645c50
Compare
The cpu-ut CI on PR vllm-project#16328 failed on three tests: 1. tests/ut/ops/test_mla.py (2 tests): the forward_k mocks still return 2-tuples while forward now unpacks 3 values (k, scale, weights tail) - update the mocks to return 3-tuples. 2. tests/ut/attention/test_sfa_v1.py::test_indexer_forward_reuses_ wk_weights_proj: the test emulated the pre-squash calling pattern (forward_k called externally before forward). forward now calls forward_k internally, so the same sequence counts one extra GEMM. Rework the assertions for the current architecture: reset the GEMM mock before forward, expect exactly one call when both stages share the same hidden-states tensor (reuse), expect three calls (one in forward_k + one fallback) for a distinct top-k input, and check the weights tail passed to indexer_select_post_process by value instead of identity. Signed-off-by: huamus <1943805462@qq.com>
ca572ed to
b0c4779
Compare
|
This pull request has conflicts, please resolve those before we can evaluate the pull request. |
b0c4779 to
73833d1
Compare
73833d1 to
29dddec
Compare
|
This pull request has conflicts, please resolve those before we can evaluate the pull request. |
GLM5.2 SFA profiling shows removable ops in the K processing path of every layer; PROLOG_V3 (the npu_mla_prolog_v3 single fused op covering qkv proj + norm + rope + q up-proj + C8 quantize/pack + direct cache write) was previously gated on is_kv_consumer, i.e. only PD-disaggregated decode workers could take it. 1. exec_kv re-casts the shared slot mapping to int64 for npu_kv_rmsnorm_rope_cache in every layer, although all layers of a scheduling step receive the same int32 slot tensor. The conversion is now cached on the attention metadata, so one Cast kernel runs per step instead of one per layer; the PROLOG_V3 fused preprocess reuses the same cached conversion for its int64 cache indices. 2. The indexer k path (forward_k) and top-k stage (forward) both run the same wk_weights_proj GEMM on the same hidden states, once for the indexer K and once for the lightning-indexer weights. forward_k now returns the non-K tail of the GEMM output and forward reuses it, removing one GEMM plus its slice copy per indexer layer per step (falling back to the GEMM only when the two stages are handed different tensors). 3. Redundant .contiguous() copies removed: npu_rms_norm returns a contiguous tensor so the copy before the C8 block-quant view is discarded, and torch.cat already allocates contiguous outputs for the sparse-attention query concat. 4. PROLOG_V3 becomes the default fused preprocessing for quantized SFA layers in every deployment (plain serving, PD KV producers and KV consumers) and serves every attention state: prefill and decode steps both take the fused path, and the per-step attention-state fallback to NATIVE is gone (only MLAPO keeps its token-count limit). enable_dsa_cp remains the prefill/P-node route selector: it routes to AscendSFADSACPImpl, which unconditionally disables fused preprocessing, so the two are mutually exclusive by construction. The C8 switches only select the KV cache layout and are orthogonal to this choice; W8A8Dynamic layers no longer require enable_sparse_sfa_c8 to take the fused path. Unquantized layers keep the NATIVE chain outside KV consumers because the unquantized weight preparation transposes fused_qkv_a_proj.weight in place, which the NATIVE fallback still consumes; dispose_layer stays gated on is_kv_consumer so producers and plain-serving workers keep the fallback weights (the cost is the extra PROLOG_V3 weight copies: memory, not correctness). Signed-off-by: huamus <1943805462@qq.com>
81d5656 to
af5688a
Compare
af5688a to
9eb069b
Compare
|
Could you please double-check whether we really need this many UTs for this change? |
- routing matrix for the default gate: non-PD W8A8Dynamic (with and without C8) and MXFP8 resolve to PROLOG_V3, unquantized stays NATIVE, KV producers take the fused path too (quantized) or stay NATIVE (unquantized); the KV-consumer cases and the weight-disposal guard are unchanged; - per-step int64 slot-conversion caching (passthrough / convert-and-cache / re-convert on new step) and exec_kv reusing the cached slots across layers; - a single wk_weights_proj invocation across the indexer k path and the top-k stage, plus the fallback recomputation when the two stages are handed different tensors; - adapt the forward_k mocks in test_mla.py to the 3-tuple return. Signed-off-by: huamus <1943805462@qq.com>
9eb069b to
481f6e4
Compare
### What this PR does / why we need it? The A5 `GLM5_1_W4A4_A5.yaml` nightly fails during graph warmup with `aclnnMlaPrologV3WeightNz` error 561002: `Parameter KvQuantMode of MlaPrologV3 has incorrect value 3. It should be 1, 0.` See [the failing job](https://github.com/vllm-project/vllm-ascend/actions/runs/35381103773/job/105718954194). The K3 custom operator introduced in #14454 shares its ACLNN symbols and OPP operator type with CANN's MLAPO v3. Its restricted template set (#16516) rejects the per-tile C8 configuration used by SFA's `torch_npu.npu_mla_prolog_v3` call. #16328 exposed that fused path to ordinary quantized serving. Give the K3 implementation its own identity throughout the build and runtime: - Rename its Torch schema to `npu_mla_prolog_v3_k3`, OPP type to `MlaPrologV3K3`, kernel to `mla_prolog_v3_k3`, and public/inner ACLNN symbols to their `K3` variants. - Route K3's no-RoPE MLA path to the renamed custom operator. Keep ordinary MLA and SFA C8 per-tile mode 3 on `torch_npu.npu_mla_prolog_v3`, backed by CANN. - Update build metadata, bindings, documentation and tests without changing the kernel math, supported K3 modes or cache layouts. ### Does this PR introduce _any_ user-facing change? GLM5 C8 serving can use CANN's per-tile MLAPO v3 while K3 custom operators are installed. Direct users of the internal `_C_ascend.npu_mla_prolog_v3` entry point must use `_C_ascend.npu_mla_prolog_v3_k3`. Rebuild/reinstall the complete custom operator package to remove the old unsuffixed registration; no compatibility alias is retained because it would recreate the collision. ### How was this patch tested? Validated on Ascend950DT with CANN 9.1.0 and torch_npu 2.10.0.post4: - `pytest -q tests/ut/attention/test_mla_v1.py tests/ut/attention/test_sfa_v1.py`: **182 passed**, including native/custom dispatch and MXFP8 C8 mode 3 routing regressions. - `cd csrc && bash build.sh --pkg --ops=mla_prolog_v3_k3 --soc=ascend950 -j256`: passed. Full `vllm_ascend_C` Torch extension CMake build also passed with `SOC_VERSION=ascend950dt_9582` and parallelism 256. - Ran all **8 existing K3 NPU cases** with the freshly built extension and OPP package, covering head96 BF16, no-RoPE, query norm, BSND, PA_BSND, PA_NZ and noncontiguous cache layouts: passed. - Added and ran `test_cann_mla_prolog_v3_mxfp8_per_tile_with_k3_registered`: the K3 no-RoPE operator executes first, followed by native MXFP8 weight mode 3 / KV mode 3; outputs are finite and only the selected KV slots are written. - Compared a separate pure-CANN process with native mode 3 after executing K3: query, query RoPE and packed KV cache are **bitwise identical**. Native mode 3 graph capture/replay with K3 loaded also matches that baseline bitwise. - Audited the fresh `libcust_opapi.so`: only `aclnnMlaPrologV3K3WeightNz*` and `aclnnInnerMlaPrologV3K3*` are exported for this operator; there are no unsuffixed MLAPO v3 exports. Generated OPP metadata uses `MlaPrologV3K3`. - `bash format.sh ci`: all hooks passed (Ruff, codespell, typos, Markdown, actionlint, secret scan, shellcheck and repository checks). The NPU checks load the fresh extension and operator package directly to avoid importing the image's older vLLM-Ascend binaries. The maintainer-triggered [GLM5_1_W4A4_A5 nightly job](https://github.com/vllm-project/vllm-ascend/actions/runs/35426937149/job/105855093279) subsequently passed on PR commit `a49f62d1ec61b80b729f8f09925605ceff1b889f`: `1 passed, 14 warnings in 405.38s`. The log confirms successful server startup, ACL graph replay and a completion request under TP8/DP1/EP, C8, prefix caching, async scheduling and MTP3. This is the configured functional test; no accuracy/performance benchmark artifacts were produced. The workflow aggregate reports failure and its artifact-merge check reports `No artifacts found matching pattern nightly-test-benchmark-results-*`; the GLM5 test job itself is successful. - vLLM main: vllm-project/vllm@84030bb Signed-off-by: MQ <maomaoyu870@gmail.com>
…t#16924) ### What this PR does / why we need it? The A5 `GLM5_1_W4A4_A5.yaml` nightly fails during graph warmup with `aclnnMlaPrologV3WeightNz` error 561002: `Parameter KvQuantMode of MlaPrologV3 has incorrect value 3. It should be 1, 0.` See [the failing job](https://github.com/vllm-project/vllm-ascend/actions/runs/35381103773/job/105718954194). The K3 custom operator introduced in vllm-project#14454 shares its ACLNN symbols and OPP operator type with CANN's MLAPO v3. Its restricted template set (vllm-project#16516) rejects the per-tile C8 configuration used by SFA's `torch_npu.npu_mla_prolog_v3` call. vllm-project#16328 exposed that fused path to ordinary quantized serving. Give the K3 implementation its own identity throughout the build and runtime: - Rename its Torch schema to `npu_mla_prolog_v3_k3`, OPP type to `MlaPrologV3K3`, kernel to `mla_prolog_v3_k3`, and public/inner ACLNN symbols to their `K3` variants. - Route K3's no-RoPE MLA path to the renamed custom operator. Keep ordinary MLA and SFA C8 per-tile mode 3 on `torch_npu.npu_mla_prolog_v3`, backed by CANN. - Update build metadata, bindings, documentation and tests without changing the kernel math, supported K3 modes or cache layouts. ### Does this PR introduce _any_ user-facing change? GLM5 C8 serving can use CANN's per-tile MLAPO v3 while K3 custom operators are installed. Direct users of the internal `_C_ascend.npu_mla_prolog_v3` entry point must use `_C_ascend.npu_mla_prolog_v3_k3`. Rebuild/reinstall the complete custom operator package to remove the old unsuffixed registration; no compatibility alias is retained because it would recreate the collision. ### How was this patch tested? Validated on Ascend950DT with CANN 9.1.0 and torch_npu 2.10.0.post4: - `pytest -q tests/ut/attention/test_mla_v1.py tests/ut/attention/test_sfa_v1.py`: **182 passed**, including native/custom dispatch and MXFP8 C8 mode 3 routing regressions. - `cd csrc && bash build.sh --pkg --ops=mla_prolog_v3_k3 --soc=ascend950 -j256`: passed. Full `vllm_ascend_C` Torch extension CMake build also passed with `SOC_VERSION=ascend950dt_9582` and parallelism 256. - Ran all **8 existing K3 NPU cases** with the freshly built extension and OPP package, covering head96 BF16, no-RoPE, query norm, BSND, PA_BSND, PA_NZ and noncontiguous cache layouts: passed. - Added and ran `test_cann_mla_prolog_v3_mxfp8_per_tile_with_k3_registered`: the K3 no-RoPE operator executes first, followed by native MXFP8 weight mode 3 / KV mode 3; outputs are finite and only the selected KV slots are written. - Compared a separate pure-CANN process with native mode 3 after executing K3: query, query RoPE and packed KV cache are **bitwise identical**. Native mode 3 graph capture/replay with K3 loaded also matches that baseline bitwise. - Audited the fresh `libcust_opapi.so`: only `aclnnMlaPrologV3K3WeightNz*` and `aclnnInnerMlaPrologV3K3*` are exported for this operator; there are no unsuffixed MLAPO v3 exports. Generated OPP metadata uses `MlaPrologV3K3`. - `bash format.sh ci`: all hooks passed (Ruff, codespell, typos, Markdown, actionlint, secret scan, shellcheck and repository checks). The NPU checks load the fresh extension and operator package directly to avoid importing the image's older vLLM-Ascend binaries. The maintainer-triggered [GLM5_1_W4A4_A5 nightly job](https://github.com/vllm-project/vllm-ascend/actions/runs/35426937149/job/105855093279) subsequently passed on PR commit `a49f62d1ec61b80b729f8f09925605ceff1b889f`: `1 passed, 14 warnings in 405.38s`. The log confirms successful server startup, ACL graph replay and a completion request under TP8/DP1/EP, C8, prefix caching, async scheduling and MTP3. This is the configured functional test; no accuracy/performance benchmark artifacts were produced. The workflow aggregate reports failure and its artifact-merge check reports `No artifacts found matching pattern nightly-test-benchmark-results-*`; the GLM5 test job itself is successful. - vLLM main: vllm-project/vllm@84030bb Signed-off-by: MQ <maomaoyu870@gmail.com>
…t#16924) ### What this PR does / why we need it? The A5 `GLM5_1_W4A4_A5.yaml` nightly fails during graph warmup with `aclnnMlaPrologV3WeightNz` error 561002: `Parameter KvQuantMode of MlaPrologV3 has incorrect value 3. It should be 1, 0.` See [the failing job](https://github.com/vllm-project/vllm-ascend/actions/runs/35381103773/job/105718954194). The K3 custom operator introduced in vllm-project#14454 shares its ACLNN symbols and OPP operator type with CANN's MLAPO v3. Its restricted template set (vllm-project#16516) rejects the per-tile C8 configuration used by SFA's `torch_npu.npu_mla_prolog_v3` call. vllm-project#16328 exposed that fused path to ordinary quantized serving. Give the K3 implementation its own identity throughout the build and runtime: - Rename its Torch schema to `npu_mla_prolog_v3_k3`, OPP type to `MlaPrologV3K3`, kernel to `mla_prolog_v3_k3`, and public/inner ACLNN symbols to their `K3` variants. - Route K3's no-RoPE MLA path to the renamed custom operator. Keep ordinary MLA and SFA C8 per-tile mode 3 on `torch_npu.npu_mla_prolog_v3`, backed by CANN. - Update build metadata, bindings, documentation and tests without changing the kernel math, supported K3 modes or cache layouts. ### Does this PR introduce _any_ user-facing change? GLM5 C8 serving can use CANN's per-tile MLAPO v3 while K3 custom operators are installed. Direct users of the internal `_C_ascend.npu_mla_prolog_v3` entry point must use `_C_ascend.npu_mla_prolog_v3_k3`. Rebuild/reinstall the complete custom operator package to remove the old unsuffixed registration; no compatibility alias is retained because it would recreate the collision. ### How was this patch tested? Validated on Ascend950DT with CANN 9.1.0 and torch_npu 2.10.0.post4: - `pytest -q tests/ut/attention/test_mla_v1.py tests/ut/attention/test_sfa_v1.py`: **182 passed**, including native/custom dispatch and MXFP8 C8 mode 3 routing regressions. - `cd csrc && bash build.sh --pkg --ops=mla_prolog_v3_k3 --soc=ascend950 -j256`: passed. Full `vllm_ascend_C` Torch extension CMake build also passed with `SOC_VERSION=ascend950dt_9582` and parallelism 256. - Ran all **8 existing K3 NPU cases** with the freshly built extension and OPP package, covering head96 BF16, no-RoPE, query norm, BSND, PA_BSND, PA_NZ and noncontiguous cache layouts: passed. - Added and ran `test_cann_mla_prolog_v3_mxfp8_per_tile_with_k3_registered`: the K3 no-RoPE operator executes first, followed by native MXFP8 weight mode 3 / KV mode 3; outputs are finite and only the selected KV slots are written. - Compared a separate pure-CANN process with native mode 3 after executing K3: query, query RoPE and packed KV cache are **bitwise identical**. Native mode 3 graph capture/replay with K3 loaded also matches that baseline bitwise. - Audited the fresh `libcust_opapi.so`: only `aclnnMlaPrologV3K3WeightNz*` and `aclnnInnerMlaPrologV3K3*` are exported for this operator; there are no unsuffixed MLAPO v3 exports. Generated OPP metadata uses `MlaPrologV3K3`. - `bash format.sh ci`: all hooks passed (Ruff, codespell, typos, Markdown, actionlint, secret scan, shellcheck and repository checks). The NPU checks load the fresh extension and operator package directly to avoid importing the image's older vLLM-Ascend binaries. The maintainer-triggered [GLM5_1_W4A4_A5 nightly job](https://github.com/vllm-project/vllm-ascend/actions/runs/35426937149/job/105855093279) subsequently passed on PR commit `a49f62d1ec61b80b729f8f09925605ceff1b889f`: `1 passed, 14 warnings in 405.38s`. The log confirms successful server startup, ACL graph replay and a completion request under TP8/DP1/EP, C8, prefix caching, async scheduling and MTP3. This is the configured functional test; no accuracy/performance benchmark artifacts were produced. The workflow aggregate reports failure and its artifact-merge check reports `No artifacts found matching pattern nightly-test-benchmark-results-*`; the GLM5 test job itself is successful. - vLLM main: vllm-project/vllm@84030bb Signed-off-by: MQ <maomaoyu870@gmail.com>
What this PR does / why we need it?
GLM5.2 SFA profiling shows several removable ops in the K processing path of every layer. This PR removes them and makes PROLOG_V3 the default fused preprocessing for quantized SFA layers in every deployment:
Per-step int64 slot cast:
exec_kvre-cast the shared slot mapping to int64 fornpu_kv_rmsnorm_rope_cachein every layer, although all layers of a scheduling step receive the same int32 slot tensor. The conversion is now cached on the attention metadata (_int64_kv_slots), so one Cast kernel runs per step instead of one per layer (~5us x num_layers per step). The PROLOG_V3 fused preprocess reuses the same cached conversion for its int64 cache indices.Duplicate indexer GEMM: the indexer's k path (
forward_k) and top-k stage (forward) both ran the samewk_weights_projGEMM ([tokens, hidden] x [160, hidden]) on the same hidden states, once for the indexer K and once for the lightning-indexer weights.forward_know returns the non-K tail of the GEMM output andforwardreuses it, removing one GEMM plus its slice copy per indexer layer per step (falling back to the GEMM only when the two stages are handed different tensors).Redundant
.contiguous()copies:npu_rms_normreturns a contiguous tensor so the copy before the C8 block-quant view was discarded, andtorch.catalready allocates contiguous outputs for the sparse-attention query concat.PROLOG_V3 by default, no new switch: PROLOG_V3 (the
npu_mla_prolog_v3single fused op covering qkv proj + norm + rope + q up-proj + C8 quantize/pack + direct cache write) was previously gated onis_kv_consumer, i.e. only PD-disaggregated decode workers could take it. It is now the default fused preprocessing for quantized SFA layers in every deployment (plain serving, PD KV producers and KV consumers) and serves every attention state: prefill and decode steps both take the fused path (the per-step attention-state fallback to NATIVE is gone; only MLAPO keeps its token-count limit). Switch convergence:enable_dsa_cpis the prefill/P-node route selector: it routes toAscendSFADSACPImpl, which unconditionally disables fused preprocessing, so the two are mutually exclusive by construction (dsa_cp on => prolog off, dsa_cp off => prolog on);enable_sparse_sfa_c8/enable_sparse_li_c8) only select the KV cache layout and are orthogonal to this choice; W8A8Dynamic layers no longer requireenable_sparse_sfa_c8to take the fused path;fused_qkv_a_proj.weightin place, which the NATIVE fallback still consumes;dispose_layerstays gated onis_kv_consumerso producers and plain-serving workers keep the fallback weights (the cost is the extra PROLOG_V3 weight copies: memory, not correctness).Does this PR introduce any user-facing change?
Yes, a default behavior change: quantized (W8A8Dynamic / W8A8MXFP8) SFA deployments now take the PROLOG_V3 fused preprocessing for both prefill and decode steps by default, without any additional-config option. Deployments on
enable_dsa_cp(prefill/P-node CP route) and unquantized (bf16) layers are unaffected. The default trades extra NPU weight memory (the retained qkv_a/q_b fallback weights on producers and plain-serving workers) for kernel savings.How was this patch tested?
Unit tests added/updated in
tests/ut/attention/test_sfa_v1.py:exec_kvreusing the cached slots across layers;wk_weights_projinvocation across the indexer k path and top-k stage (plus the fallback recomputation when the tensors differ);uvx ruff==0.14.0 checkandformat --checkon all touched Python files.CI cpu-ut green (4431+ tests).
End-to-end A/B benchmark: DSA-CP route (main) vs PROLOG_V3 route (this PR) (Atlas 800 A3, 16x 910B, CANN 9.1.0, vllm 0.28.0, GLM-5.2-w4a8c8, DP2xTP8 + EP):
The comparison is between the two decode preprocessing routes, each in its best usable configuration.
enable_dsa_cprequires SP-MoE and unconditionally disables the fused preprocessing path, so the two options are mutually exclusive by construction; everything else is identical on both sides:125924bb2)cb525a869)enable_dsa_cp=trueVLLM_ASCEND_ENABLE_FLASHCOMM1=1; auto-enabled by DSA-CP on the base side)enable_sparse_sfa_c8+enable_sparse_li_c8enable_balance_scheduling,enable_fused_mc2=0--enable-expert-parallel(EP), DP2xTP8deepseek_mtp, num_speculative_tokens=3, enforce_eager)FULL_DECODE_ONLY,--quantization ascendmultistream_overlap_shared_expertBoth sides verified via serve logs: no "Disabling DSA-CP" / sp-MoE active on the base side,
MlaPrologV3kernels present in the PR-side profile only. Workload: GSM8K test full 1319 prompts, ais-bench stream mode, concurrency 8, temperature 0.Kernel-level verification (rank0 profile of a 500-token decode request):
E2E coverage for the PROLOG_V3 default route:
Following [CI] Remove GLM-5.2 spec decode CI tests #16630 (GLM-5.2 spec-decode PR-level CI tests are removed as nightly-covered), the fused path is exercised by the nightly-covered GLM-5.2 spec-decode e2e suites together with the one-shot hardware verification below; no dedicated per-push full-forward UT is carried in this PR.
Verified on real hardware (Atlas 800 A3, 8x 910B, CANN 9.1.0, vllm-ascend 0.28.0): the PROLOG_V3-route e2e case PASSED (acceptance length within 3.06 +/- 8%) when run as the eight_card job on branch
perf/glm-kpath-fusion-28and in the E2E CI run of commit28e9e11(all 24 jobs green), confirming the fused route serves plain (non-PD) decode correctly under the full serving stack.Note: serving the w4a8c8 checkpoint on torch_npu 2.10 additionally requires a one-word out-of-tree fix in
vllm_ascend/quantization/methods/w4a8/w4a8.py(.sum(axis=1)->.sum(dim=1); theaxiskwarg comes from [Refactor][quantization]Remove deprecated W4A8Linear/W8A8PDMixMoE schemes and sync W4A8 MoE refactor #13713 and torch_npu'sreduce_sumrejects it; fixed on main by [BugFix] Fix batch-invariant reduce_sum crash on non-last-dim reductions #16413).vLLM main: vllm-project/vllm@84030bb