Repository navigation
[BugFix] Fix batch-invariant reduce_sum crash on non-last-dim reductions - #16413
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 resolves a critical issue where the batch-invariant reduce_sum operation on NPU devices would crash when attempting to reduce dimensions other than the last one. By introducing dimension normalization and restricting the use of the specialized kernel to genuine last-dimension reductions, the implementation now safely handles unsupported cases by falling back to native PyTorch operations. This change maintains batch invariance for critical paths while significantly improving the robustness of the NPU integration. 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. Tip 💡 Consider Linking a Related Issue or RFCYour PR title contains the [BugFix] tag, indicating a bug fix or new feature. Linking a related issue or RFC in the PR description is strongly encouraged — it gives reviewers helpful context and speeds up the review. You can use any of these keywords:
🙏 Thanks for helping us keep the project well-organized! |
There was a problem hiding this comment.
Code Review
Suggested PR Title:
[Ops][BugFix] Restrict batch-invariant reduce_sum to last dimensionSuggested PR Summary:
### What this PR does / why we need it?
This PR ensures that the batch-invariant `reduce_sum` is only applied when reducing the last dimension, preventing `AclNN_Parameter_Error` on non-last dimensions. It introduces dimension normalization to handle negative indexing and falls back to native `torch_sum` for other dimensions.
Feedback:
- The negative dimension normalization logic can be simplified by directly checking if `dim == -1 or dim == x.dim() - 1`, which avoids manual normalization arithmetic and improves readability.
### Does this PR introduce _any_ user-facing change?
No.
### How was this patch tested?
No tests were added in this PR.
linfeng-yuan
left a comment
There was a problem hiding this comment.
Please remove and isinstance(dim, int)
1039cc7 to
b363b4b
Compare
bedf5ee to
b363b4b
Compare
done |
b363b4b to
a770cbe
Compare
…ductions `vllm_ascend/batch_invariant.py:reduce_sum` forwarded every NPU `torch.sum` to `npu_reduce_sum_batch_invariant`, but the underlying `aclnnReduceSumBatchInvariant` kernel only supports reducing the last dimension and raises `AclNN_Parameter_Error(EZ1001): Provided dim only support last dim` for any other dim. Because `enable_batch_invariant_mode()` monkey-patches `torch.sum` globally, any non-last-dim `.sum()` on NPU kills all workers -- e.g. during `profile_run` on multimodal models (qwen3_vl `pos_embed_interpolate_native` does `embeds.sum(dim=0)`). Only use the batch-invariant kernel when the reduction is over the last dimension (`dim == -1` or `dim == x.dim() - 1`), forwarding the caller's `dim` unchanged; fall back to the saved native `torch.sum` for non-last-dim, tuple dim, full reduction (dim is None), CPU tensors, and unsupported dtypes. Both spellings select the same axis, so previously-working last-dim calls keep their exact behavior and the existing unit test stays valid. Signed-off-by: EliasKaslan <3248569436@qq.com>
a770cbe to
bab98a1
Compare
…ed decode for non-PD serving (#16328) ### 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: 1. **Per-step int64 slot cast**: `exec_kv` re-cast 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 (`_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. 2. **Duplicate indexer GEMM**: the indexer's k path (`forward_k`) and top-k stage (`forward`) both ran the same `wk_weights_proj` GEMM (`[tokens, hidden] x [160, hidden]`) 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**: `npu_rms_norm` returns a contiguous tensor so the copy before the C8 block-quant view was discarded, and `torch.cat` already allocates contiguous outputs for the sparse-attention query concat. 4. **PROLOG_V3 by default, no new switch**: 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. 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_cp` is the prefill/P-node route selector: it routes to `AscendSFADSACPImpl`, which unconditionally disables fused preprocessing, so the two are mutually exclusive by construction (dsa_cp on => prolog off, dsa_cp off => prolog on); - the C8 switches (`enable_sparse_sfa_c8` / `enable_sparse_li_c8`) 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). ### 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`: - per-step int64 slot caching (passthrough / convert-and-cache / re-convert on new step) and `exec_kv` reusing the cached slots across layers; - single `wk_weights_proj` invocation across the indexer k path and top-k stage (plus the fallback recomputation when the tensors differ); - PROLOG_V3 routing matrix for the default gate (non-PD W8A8Dynamic with and without C8 -> PROLOG_V3; non-PD MXFP8 -> PROLOG_V3; non-PD unquantized -> NATIVE; KV producer quantized -> PROLOG_V3, unquantized -> NATIVE); - weight-disposal guard confirming non-consumer workers keep the fallback weights. - `uvx ruff==0.14.0 check` and `format --check` on 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_cp` requires 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: | item | base (main `125924bb2`) | PR (`cb525a869`) | |---|---|---| | **decode preprocessing route** | `enable_dsa_cp=true` | PROLOG_V3 route (now the default; measured with the earlier opt-in build of this PR) | | SP-MoE / sequence parallelism (`VLLM_ASCEND_ENABLE_FLASHCOMM1=1`; auto-enabled by DSA-CP on the base side) | on | on | | `enable_sparse_sfa_c8` + `enable_sparse_li_c8` | on | on | | `enable_balance_scheduling`, `enable_fused_mc2=0` | on | on | | `--enable-expert-parallel` (EP), DP2xTP8 | on | on | | MTP speculative decoding (`deepseek_mtp`, num_speculative_tokens=3, enforce_eager) | on | on | | cudagraph `FULL_DECODE_ONLY`, `--quantization ascend` | on | on | | `multistream_overlap_shared_expert` | off | off (incompatible with FlashComm1 on DP>1, see #16446) | Both sides verified via serve logs: no "Disabling DSA-CP" / sp-MoE active on the base side, `MlaPrologV3` kernels present in the PR-side profile only. Workload: GSM8K test full 1319 prompts, ais-bench stream mode, concurrency 8, temperature 0. | metric | base (DSA-CP route) | PR (PROLOG_V3 route) | delta | |---|---|---|---| | TPOT avg | 27.9 ms | **22.9 ms** | **-17.9%** | | TPOT median | 27.8 ms | 22.8 ms | -18.0% | | E2EL avg | 7,181.6 ms | **5,885.0 ms** | -18.1% | | TTFT avg | 459.9 ms | 395.8 ms | -13.9% | | per-request output throughput | 33.73 tok/s | **40.91 tok/s** | +21.3% | | request throughput (aggregate) | 1.1117 req/s | 1.3564 req/s | +22.0% | | failed requests | 0/1319 | 0/1319 | = | Kernel-level verification (rank0 profile of a 500-token decode request): | kernel | base (DSA-CP) | PR (PROLOG_V3) | note | |---|---|---|---| | MlaPrologV3 | 0 | **13,369** | fused decode path taken over | | InterleaveRope | 28,980 | **1,156** | -96%: NATIVE rope chain leaves decode | | DynamicBlockQuant | 14,490 | **578** | -96%: c8 quant chain leaves decode | | ScatterNdUpdate | 21,035 | 6,967 | -67% | | Slice | 31,556 | 12,211 | -61% count | | QuantBatchMatmulV3 | 73,652 (1,629 ms) | 44,207 (835 ms) | -40% count / -49% time | | KvQuantSparseFlashAttention | 14,190 | 13,703 | ~same (attention body) | | **total decode kernels** | **711,070** | **479,326** | **-32.6%** | **E2E coverage for the PROLOG_V3 default route**: - Following #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-28` and in the E2E CI run of commit `28e9e11` (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)`; the `axis` kwarg comes from #13713 and torch_npu's `reduce_sum` rejects it; fixed on main by #16413). - vLLM main: vllm-project/vllm@84030bb --------- Signed-off-by: huamus <1943805462@qq.com>
What this PR does / why we need it?
vllm_ascend/batch_invariant.py:reduce_sumunconditionally forwards every NPUtorch.sum(x, dim, keepdim)(for supported dtypes) totorch.ops.batch_invariant_ops.npu_reduce_sum_batch_invariant. But theaclnnReduceSumBatchInvariantkernel only supports reducing thelast dimension; for any other
dimit raises:Because
enable_batch_invariant_mode()monkey-patchestorch.sum/torch.Tensor.sumglobally, any non-last-dim.sum()on NPU kills all TPworkers.
To fix it, this PR only takes the batch-invariant path for a genuine
last-dim reduction:
-1orx.dim() - 1, so the guard issimply
dim == -1 or dim == x.dim() - 1, and the caller'sdimisforwarded to the kernel unchanged. No dimension arithmetic is needed,
and equality comparison is safe for every value of
dim: a tuple, forexample, simply fails both comparisons and falls back instead of raising.
The kernel is therefore used only when all of the following hold: NPU
tensor, a genuine last-dim reduction, and dtype in
{fp16, fp32, bf16}.Forwarding
dimunchanged means every previously-working last-dim call keepsthe exact behavior it had before this PR, and the existing unit test stays
valid.
dim, full reduction (dim is Noneon >1-D), CPU tensors, unsupported dtypes - falls back to the saved
native
torch_sumreference captured at module import, so the fallbackcannot recurse into the patched entry points.
Every input that now falls back previously crashed, so the change strictly
widens the set of inputs that run; there is no working behavior being changed.
Why this does not affect batch invariance?
The invariance-critical reductions are all last-dim, and they are
unchanged. The reductions whose results must be batch-invariant (softmax
denominators, RMSNorm variance - reductions over the hidden/vocab
dimension) are last-dim reductions on
[num_tokens, hidden]tensors.Native kernels pick their reduction strategy from the total tensor size,
so the per-row accumulation order can differ between BS=1 and BS=N - that
is exactly what
npu_reduce_sum_batch_invariantprevents with arow-count-independent accumulation order. These calls still take the
batch-invariant kernel; their behavior is byte-for-byte identical to
before.
The newly-fallback paths have no batch-shape-dependent input. The
non-last-dim sums that now fall back (e.g.
embeds.sum(dim=0)inmultimodal position-embedding preprocessing) are per-request computations:
the input tensor is built from one request's own patches, so its shape and
values are identical in a BS=1 run and in a BS=N run. Native
torch.sumis deterministic for a given tensor (same shape ? same kernel and tiling ?
same accumulation order), so identical input yields bitwise-identical
output across batch sizes. Invariance is preserved.
No regression is possible. Every path that now falls back previously
raised EZ1001 (non-last-dim int dims) or a schema type error (tuple dims).
There was no working behavior to preserve.
It follows the established fallback design. [Ops][Feature] Support Ascend 950 and upgrade batch invariant ops to 2.0.0 #15956 introduced the same
pattern for the dtype axis: unsupported dtypes "fall back to the saved
native torch.sum implementation, avoiding recursive calls through the
patched Tensor.sum". This PR extends that pattern from dtype to dim.
Does this PR introduce any user-facing change?
No
How was this patch tested?
vs native, keepdim, 1-D
.sum(), tuple dim, all-dims reduction, fp16,functional
torch.sum, fp32) - all pass.profile_runand generation with no EZ1001.vllm-project/vllm@b2f6858
Signed-off-by: EliasKaslan 3248569436@qq.com