Skip to content

[BugFix] Fix batch-invariant reduce_sum crash on non-last-dim reductions - #16413

Merged
linfeng-yuan merged 1 commit into
vllm-project:mainfrom
EliasKaslan:fix/batch-invariant-reduce-sum-non-last-dim
Sep 17, 2026
Merged

linfeng-yuan merged 1 commit into
vllm-project:mainfrom
EliasKaslan:fix/batch-invariant-reduce-sum-non-last-dim

Conversation

@EliasKaslan

@EliasKaslan EliasKaslan commented Sep 12, 2026 •

Copy link
Copy Markdown
Contributor

What this PR does / why we need it?

vllm_ascend/batch_invariant.py:reduce_sum unconditionally forwards every NPU
torch.sum(x, dim, keepdim) (for supported dtypes) to
torch.ops.batch_invariant_ops.npu_reduce_sum_batch_invariant. But the
aclnnReduceSumBatchInvariant kernel only supports reducing the
last dimension
; for any other dim it raises:

AclNN_Parameter_Error(EZ1001): Provided dim only support last dim

Because enable_batch_invariant_mode() monkey-patches torch.sum /
torch.Tensor.sum globally, any non-last-dim .sum() on NPU kills all TP
workers.

To fix it, this PR only takes the batch-invariant path for a genuine
last-dim reduction
:

  1. The last dim can only be spelled as -1 or x.dim() - 1, so the guard is
    simply dim == -1 or dim == x.dim() - 1, and the caller's dim is
    forwarded to the kernel unchanged. No dimension arithmetic is needed,
    and equality comparison is safe for every value of dim: a tuple, for
    example, 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 dim unchanged means every previously-working last-dim call keeps
    the exact behavior it had before this PR, and the existing unit test stays
    valid.
  2. Everything else - non-last-dim, tuple dim, full reduction (dim is None
    on >1-D), CPU tensors, unsupported dtypes - falls back to the saved
    native torch_sum reference
    captured at module import, so the fallback
    cannot 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?

  1. 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_invariant prevents with a
    row-count-independent accumulation order. These calls still take the
    batch-invariant kernel; their behavior is byte-for-byte identical to
    before.

  2. 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) in
    multimodal 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.sum
    is 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.

  3. 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.

  4. 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?

  • 10-path NPU unit test on Ascend910 (dim=0 fallback, dim=2 / dim=-1 equivalence
    vs native, keepdim, 1-D .sum(), tuple dim, all-dims reduction, fp16,
    functional torch.sum, fp32) - all pass.
  • End-to-end batch-invariance logprobs test (BS=1 vs BS=N bitwise) runs through
    profile_run and generation with no EZ1001.

vllm-project/vllm@b2f6858


Signed-off-by: EliasKaslan 3248569436@qq.com

@gemini-code-assist

Copy link
Copy Markdown
Contributor

Summary of Changes

Hello, 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

  • Fix for NPU reduce_sum crash: Added dimension normalization and a strict check to ensure only last-dimension reductions are processed by the batch-invariant kernel, preventing crashes on non-last-dim operations.
  • Improved fallback mechanism: Non-last-dim reductions, tuple dimensions, and unsupported dtypes now correctly fall back to the native torch.sum implementation, ensuring stability across all input types.
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 Assist

The 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 /gemini <command> or @gemini-code-assist <command>. Below is a summary of the supported commands on the current page.

Feature Command Description
Code Review /gemini review Performs a code review for the current pull request in its current state.
Pull Request Summary /gemini summary Provides a summary of the current pull request in its current state.
Comment @gemini-code-assist Responds in comments when explicitly tagged, both in pull request comments and review comments.
Help /gemini help Displays a list of available commands.

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 .gemini/ folder in the base of the repository. Detailed instructions can be found here.

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

  1. Review the Privacy Notices, Generative AI Prohibited Use Policy, Terms of Service, and learn how to configure Gemini Code Assist in GitHub here. Gemini can make mistakes, so double check it and use code with caution. ↩

@github-actions

Copy link
Copy Markdown
Contributor

👋 Hi! Thank you for contributing to the vLLM Ascend project. The following points will speed up your PR merge:‌‌

  • A PR should do only one thing, smaller PRs enable faster reviews.
  • Every PR should include unit tests and end-to-end tests ‌to ensure it works and is not broken by other future PRs.
  • Write the commit message by fulfilling the PR description to help reviewer and future developers understand.

If CI fails, you can run linting and testing checks locally according Contributing and Testing.


Tip

💡 Consider Linking a Related Issue or RFC

Your 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:

  • Fixes #<issue_number>
  • Closes #<issue_number>
  • Resolves #<issue_number>
  • Refs #<rfc_or_issue_number> (for RFCs)

🙏 Thanks for helping us keep the project well-organized!

@gemini-code-assist gemini-code-assist Bot left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

Code Review

Suggested PR Title:

[Ops][BugFix] Restrict batch-invariant reduce_sum to last dimension

Suggested 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.

Comment thread vllm_ascend/batch_invariant.py Outdated
linfeng-yuan
linfeng-yuan previously approved these changes Sep 15, 2026

@linfeng-yuan linfeng-yuan left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

Please remove and isinstance(dim, int)

@EliasKaslan
EliasKaslan force-pushed the fix/batch-invariant-reduce-sum-non-last-dim branch 2 times, most recently from 1039cc7 to b363b4b Compare September 15, 2026 08:30
@EliasKaslan
EliasKaslan force-pushed the fix/batch-invariant-reduce-sum-non-last-dim branch from bedf5ee to b363b4b Compare September 16, 2026 02:48
@linfeng-yuan linfeng-yuan added the ready-precise run selected e2e test for pr label Sep 16, 2026
@EliasKaslan

Copy link
Copy Markdown
Contributor Author

Please remove and isinstance(dim, int)

done

@EliasKaslan
EliasKaslan force-pushed the fix/batch-invariant-reduce-sum-non-last-dim branch from b363b4b to a770cbe Compare September 17, 2026 01:13
…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>
@EliasKaslan
EliasKaslan force-pushed the fix/batch-invariant-reduce-sum-non-last-dim branch from a770cbe to bab98a1 Compare September 17, 2026 01:34
@linfeng-yuan
linfeng-yuan merged commit 7c0747c into vllm-project:main Sep 17, 2026
14 checks passed
ZT-AIA pushed a commit that referenced this pull request Sep 18, 2026
…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>
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

module:core ready-precise run selected e2e test for pr

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants