Skip to content

[BugFix][Attention] Keep Gemma4 512-head MTP decode in ACL graphs - #14670

Open
Liuchenbing-2026 wants to merge 5 commits into
vllm-project:mainfrom
Liuchenbing-2026:gemma4_mtp_main
Open

Liuchenbing-2026 wants to merge 5 commits into
vllm-project:mainfrom
Liuchenbing-2026:gemma4_mtp_main

Conversation

@Liuchenbing-2026

@Liuchenbing-2026 Liuchenbing-2026 commented Aug 20, 2026 •

Copy link
Copy Markdown
Contributor

What this PR does / why we need it?

Fix ACL graph capture/replay for Gemma4's 512-dimensional global attention heads with MTP on Ascend 910B. On the tested baseline, FULL_DECODE_ONLY fails during graph capture with error 561002 in the TND attention path; ATB paged attention cannot serve as the captured decode fallback.

The final change is confined to attention and its regression tests:

  • Use BNSD fused attention for single-token large-head decode in the full graph. Resolve per-layer block tables and KV lengths into stable captured buffers on each replay, clearing padded rows and rejecting block tables wider than the captured allocation.
  • For multi-token target verification, gather paged KV into BNSD and construct a causal mask from device-side sequence lengths. Keep the sequence-length device buffer address stable while refreshing its contents before replay.
  • Preserve paged attention selection outside FULL_DECODE_ONLY and avoid copying an output tensor onto itself.
  • Add numerical regression coverage for causal visibility, grouped-query attention, dynamic sequence lengths/block tables, and replay buffer addresses.

Does this PR introduce any user-facing change?

Gemma4-31B-it with its assistant checkpoint can capture and replay the tested FULL_DECODE_ONLY + MTP configuration on Ascend 910B. No new environment variable, model definition, or model-runner implementation is introduced.

The validated model configuration is BF16, TP=2, DP=2, MTP=3, max model length 5500, max sequences 16, and max batched tokens 8192. Capture sizes are 4 through 48 in increments of 4. Validation is text-only.

How was this patch tested?

Baseline: vLLM v0.30.0 (ced6857afa0ea7b2e3f0846a62e1394e90f15607) with vLLM-Ascend a8fcedb03d93e60efceddbfc912406f7fa491d57, CANN 9.1.0, torch 2.10.0 and torch-npu 2.10.0.post4, Ascend 910B.

Automated regressions, run locally:

python -m pytest -q \
  tests/ut/attention/test_attention_v1.py \
  tests/ut/attention/test_attention_utils.py \
  tests/ut/attention/test_large_head_graph.py
# 58 passed

python -m pytest -q \
  tests/e2e/pull_request/one_card/aclgraph/test_large_head_verify.py
# 2 passed: FP16 and BF16, real NPU graph capture/replay

bash format.sh ci
# Passed, including the staged secret scan

CPU numerical tests execute the production gather/layout/mask code with a software replacement only for the NPU kernel and compare against an independent per-query visible-prefix reference. They also test decode task argument binding, in-place provider updates, stale-row clearing and oversized block-table rejection.

The one-card NPU test requires no model download. It runs the real verification kernel, captures an ACL graph, then changes query values, KV values, sequence lengths and block tables at fixed addresses before replay. Outputs are compared with CPU FP32 SDPA over each query's visible prefix (rtol=atol=0.015). NPU coverage is necessary because CPU mocks cannot validate kernel capture or device replay semantics. CPU tests are covered by the existing tests/ut CI jobs; the NPU test is under the existing one-card test selection path. Local passes are not a claim that the new commit's remote CI has completed.

Model-level validation:

  • Unmodified baseline reproduces the capture failure. With the fix, both DP groups capture all 12 shapes and replay ACL graphs successfully.
  • 32 varied requests complete successfully; 29/32 outputs and complete token sequences match the eager baseline, and 16/16 arithmetic answers are correct in both runs.
  • C-Eval validation, 1346 questions: eager without MTP 1035/1346 (76.8945%); full graph + MTP 1037/1346 (77.0431%). Same zero-shot chat/greedy protocol, seed 62821, max output 16 tokens and concurrency 8; invalid answers count as incorrect. This is not a claim of exact token equivalence or compliance with an unspecified official accuracy threshold.

Limits: This round does not include a fresh throughput/TTFT/TPOT benchmark. Other hardware, Model Runner V2, KV quantization and additional attention-feature combinations are not covered by these results. Verification materializes dense KV buffers, so its performance and memory cost should be measured before extending the supported envelope.

@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 introduces support for Gemma4 Multi-Token Prediction (MTP) speculative decoding on Ascend NPU hardware. By adapting the upstream Gemma4Proposer and integrating it with the existing Ascend speculative decoding infrastructure, this change enables efficient multi-group KV cache usage and draft model execution. The implementation includes necessary refinements to attention mechanisms, RoPE operations, and configuration handling to ensure compatibility with Gemma4's architecture on Ascend.

Highlights

  • Gemma4 MTP Support: Added AscendGemma4Proposer to enable Gemma4 Multi-Token Prediction (MTP) speculative decoding on Ascend NPU.
  • KV Cache and Metadata Management: Implemented multi-group KV cache support and per-group attention metadata building to accommodate Gemma4's architecture.
  • Attention and RoPE Enhancements: Added large-head attention fallback and draft-only RoPE operations to ensure compatibility with Ascend execution constraints.
  • Configuration and Utilities: Introduced utility functions for multimodal token handling and backported Gemma4 assistant configuration registration.
  • Testing: Added comprehensive unit tests for Gemma4 MTP, RoPE operations, and transformer utility functions.
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 [Feature] 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:

[Attention][Feature] Support Gemma4 MTP and Speculative Decoding on Ascend NPU

Suggested PR Summary:

### What this PR does / why we need it?
This PR adds support for Gemma4 Multi-Token Prediction (MTP) and speculative decoding on Ascend NPUs. It introduces the `AscendGemma4Proposer` and updates the attention mechanism to handle Gemma4's 512-dim global attention heads using FlashAttention fallback pathways (`_forward_large_head_prefill_attention` and `_forward_large_head_graph_verify_attention`). Additionally, it implements query-only RoPE support for cross-layer KV sharing, backports Gemma4 assistant config registration, and updates speculative decoding metadata builders to support per-group slot mapping and block tables.

Several critical issues were identified in the review:
- A potential `ImportError` in `transformers_utils.py` due to the use of non-existent `strict` from `huggingface_hub.dataclasses`.
- A `NameError` in `attention_v1.py` where `layer` is referenced instead of `self`.
- An `AttributeError` in `llm_base_proposer.py` where `self.input_ids` is accessed before initialization.
- Potential runtime errors (`ZeroDivisionError` and `ValueError`) in the new attention fallback functions when handling empty batches or sequence lengths.

### Does this PR introduce _any_ user-facing change?
Yes, it enables speculative decoding and MTP support for Gemma4 models on Ascend NPU platforms.

### How was this patch tested?
New unit tests have been added under `tests/ut/` covering RoPE, Gemma4 MTP, Gemma4 vLLM compatibility, and Transformers utility configurations.

Comment thread vllm_ascend/transformers_utils.py Outdated
Comment on lines +17 to +23
from huggingface_hub.dataclasses import strict
from transformers import AutoConfig, PretrainedConfig
from transformers.models.auto.configuration_auto import CONFIG_MAPPING
from transformers.models.gemma4.configuration_gemma4 import Gemma4TextConfig


@strict

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.

critical

The import from huggingface_hub.dataclasses import strict and the @strict decorator do not exist in standard versions of huggingface_hub and will cause a critical ImportError on startup when register_gemma4_assistant_config is called. Please remove them as they are not required for PretrainedConfig subclasses.

Suggested change
from huggingface_hub.dataclasses import strict
from transformers import AutoConfig, PretrainedConfig
from transformers.models.auto.configuration_auto import CONFIG_MAPPING
from transformers.models.gemma4.configuration_gemma4 import Gemma4TextConfig
@strict
from transformers import AutoConfig, PretrainedConfig
from transformers.models.auto.configuration_auto import CONFIG_MAPPING
from transformers.models.gemma4.configuration_gemma4 import Gemma4TextConfig

Comment thread vllm_ascend/attention/attention_v1.py Outdated

output_padded = None
if key is not None and value is not None:
if key is not None and value is not None and getattr(layer, "kv_sharing_target_layer_name", None) is None:

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.

critical

The variable layer is not defined in the scope of the forward method, which will cause a NameError during execution. Since kv_sharing_target_layer_name is an attribute of the attention backend implementation, you should use self instead of layer.

Suggested change
if key is not None and value is not None and getattr(layer, "kv_sharing_target_layer_name", None) is None:
if key is not None and value is not None and getattr(self, "kv_sharing_target_layer_name", None) is None:

if self.supports_mm_inputs:
# A multimodal target may use a text-only assistant, as Gemma4 does.
try:
dummy_input_ids = torch.tensor([[1]], device=self.input_ids.device)

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.

high

During load_model, self.input_ids is not guaranteed to be initialized and may be None, which would cause an AttributeError when accessing self.input_ids.device. Use self.device instead, which is already initialized and available on the proposer.

Suggested change
dummy_input_ids = torch.tensor([[1]], device=self.input_ids.device)
dummy_input_ids = torch.tensor([[1]], device=self.device)

Comment on lines +1672 to +1674
batch_size = attn_metadata.seq_lens.shape[0]
num_tokens = query.shape[0]
if num_tokens % batch_size != 0:

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.

high

If batch_size is 0 (e.g., during empty steps or idle requests), num_tokens % batch_size will raise a ZeroDivisionError. Add a guard to return early when batch_size is 0.

        batch_size = attn_metadata.seq_lens.shape[0]
        if batch_size == 0:
            return output
        num_tokens = query.shape[0]
        if num_tokens % batch_size != 0:

Comment thread vllm_ascend/attention/attention_v1.py Outdated
Comment on lines +1759 to +1760
block_size = key_cache.shape[1]
max_seq_len = max(seq_lens)

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.

high

If seq_lens is empty, calling max(seq_lens) will raise a ValueError. Add a defensive check to return empty tensors when seq_lens is empty.

Suggested change
block_size = key_cache.shape[1]
max_seq_len = max(seq_lens)
block_size = key_cache.shape[1]
if not seq_lens:
return (
torch.empty(0, dtype=key_cache.dtype, device=key_cache.device),
torch.empty(0, dtype=value_cache.dtype, device=value_cache.device),
)
max_seq_len = max(seq_lens)

@github-actions

Copy link
Copy Markdown
Contributor

This pull request has conflicts, please resolve those before we can evaluate the pull request.

@github-actions

Copy link
Copy Markdown
Contributor

This pull request has conflicts, please resolve those before we can evaluate the pull request.

@github-actions

github-actions Bot commented Sep 3, 2026

Copy link
Copy Markdown
Contributor

This pull request has conflicts, please resolve those before we can evaluate the pull request.

@github-actions

github-actions Bot commented Sep 9, 2026

Copy link
Copy Markdown
Contributor

This pull request has conflicts, please resolve those before we can evaluate the pull request.

@github-actions

Copy link
Copy Markdown
Contributor

This pull request has conflicts, please resolve those before we can evaluate the pull request.

@github-actions

Copy link
Copy Markdown
Contributor

This pull request has conflicts, please resolve those before we can evaluate the pull request.

@github-actions

Copy link
Copy Markdown
Contributor

This pull request has conflicts, please resolve those before we can evaluate the pull request.

@Liuchenbing-2026 Liuchenbing-2026 changed the title [Feature][SpecDecode] Support Gemma4 MTP on Ascend [BugFix][Attention] Keep Gemma4 512-head MTP decode in ACL graphs Oct 1, 2026
@github-actions

github-actions Bot commented Oct 6, 2026

Copy link
Copy Markdown
Contributor

This pull request has conflicts, please resolve those before we can evaluate the pull request.

Add Gemma4 assistant configuration, multi-group speculative metadata, hybrid attention graph replay, large-head attention fallback, KV-sharing support, and regression coverage for clean upstream vLLM.

Signed-off-by: Liuchenbing-2026 <Liuchenbing-2026@users.noreply.github.com>
… graph

Signed-off-by: Liuchenbing-2026 <Liuchenbing-2026@users.noreply.github.com>
Signed-off-by: Liuchenbing-2026 <Liuchenbing-2026@users.noreply.github.com>
Use BNSD attention for 512-head decode and MTP verification, refreshing
captured KV metadata in place and applying a causal verification mask.
Include regression coverage for the large-head graph paths.

Port the final implementation of PR vllm-project#14670 onto Ascend main a8fcedb.
The complete source tree matches the tested rebased head 4cca74cd.

Validated with vLLM v0.30.0 (ced6857), four Ascend 910B4-1 devices,
TP2/DP2, BF16, MTP3, and FULL_DECODE_ONLY:
- clean main reproduces capture error 561002 for TND headDim=512;
- both DP groups complete 12 capture sizes and execute graph replay;
- 52 attention unit tests and 32 varied inference requests pass;
- C-Eval custom zero-shot validation: 1037/1346 versus eager 1035/1346,
  with zero request errors in both modes.

Per-question and token-sequence differences remain; throughput and
latency benchmarks were not rerun in this validation.

Signed-off-by: liuchenbing <chenliumail@163.com>
Add CPU numerical regressions for verification masks, decode parameter binding, and stable replay buffers. Add a model-free NPU graph regression against CPU FP32 SDPA for FP16 and BF16 with changing inputs and metadata. Validated 58 CPU tests, 2 NPU tests, and full format checks.

Signed-off-by: liuchenbing <chenliumail@163.com>
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant