Repository navigation
[BugFix][Attention] Keep Gemma4 512-head MTP decode in ACL graphs - #14670
Liuchenbing-2026 wants to merge 5 commits into
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 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
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 [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:
🙏 Thanks for helping us keep the project well-organized! |
There was a problem hiding this comment.
Code Review
Suggested PR Title:
[Attention][Feature] Support Gemma4 MTP and Speculative Decoding on Ascend NPUSuggested 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.| 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 |
There was a problem hiding this comment.
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.
| 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 |
|
|
||
| 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: |
There was a problem hiding this comment.
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.
| 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) |
There was a problem hiding this comment.
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.
| dummy_input_ids = torch.tensor([[1]], device=self.input_ids.device) | |
| dummy_input_ids = torch.tensor([[1]], device=self.device) |
| batch_size = attn_metadata.seq_lens.shape[0] | ||
| num_tokens = query.shape[0] | ||
| if num_tokens % batch_size != 0: |
There was a problem hiding this comment.
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:| block_size = key_cache.shape[1] | ||
| max_seq_len = max(seq_lens) |
There was a problem hiding this comment.
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.
| 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) |
4312a0c to
c77f5d8
Compare
|
This pull request has conflicts, please resolve those before we can evaluate the pull request. |
0e39f6a to
ae0a6c0
Compare
ae0a6c0 to
9bbf7dd
Compare
9bbf7dd to
1f81f2f
Compare
|
This pull request has conflicts, please resolve those before we can evaluate the pull request. |
1f81f2f to
11ffe1a
Compare
6a8cb31 to
fd62099
Compare
|
This pull request has conflicts, please resolve those before we can evaluate the pull request. |
fd62099 to
2d4dd0c
Compare
|
This pull request has conflicts, please resolve those before we can evaluate the pull request. |
d4a6d48 to
1e5c799
Compare
1e5c799 to
977d55b
Compare
|
This pull request has conflicts, please resolve those before we can evaluate the pull request. |
bf48081 to
988a284
Compare
988a284 to
6c8003d
Compare
|
This pull request has conflicts, please resolve those before we can evaluate the pull request. |
|
This pull request has conflicts, please resolve those before we can evaluate the pull request. |
6c8003d to
aec68a3
Compare
|
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>
62d56fa to
ff8d17f
Compare
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:
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-Ascenda8fcedb03d93e60efceddbfc912406f7fa491d57, CANN 9.1.0, torch 2.10.0 and torch-npu 2.10.0.post4, Ascend 910B.Automated regressions, run locally:
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:
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.