Skip to content

[Feature] Add minimal Gemma4 MTP support on Ascend - #13263

Draft
Alex-stack-hub wants to merge 1 commit into
vllm-project:mainfrom
Alex-stack-hub:refactor/gemma4-mtp-upstream
Draft

Alex-stack-hub wants to merge 1 commit into
vllm-project:mainfrom
Alex-stack-hub:refactor/gemma4-mtp-upstream

Conversation

@Alex-stack-hub

@Alex-stack-hub Alex-stack-hub commented Jul 31, 2026 •

Copy link
Copy Markdown
Contributor

What this PR does / why we need it?

实验pr
This PR provides a minimal, upstreamable Ascend integration for Gemma4 MTP speculative decoding. It refactors #13045 on top of the current main branch and the vLLM v0.26.0 integration point.

Gemma4 MTP differs from existing Ascend proposers in a few required ways:

  • draft layers span multiple KV-cache groups for sliding and full attention;
  • draft attention reads K/V from target-model layers;
  • every draft step uses the same target position and sequence length;
  • some attention layers apply RoPE to Q only;
  • FIA metadata must use the speculative-decoding state;
  • CUDA centroid graphs must not be captured on NPU.

The implementation reuses vLLM's upstream Gemma4Proposer for model-generic behavior and keeps only Ascend-specific adaptation:

  • route Gemma4 MTP to AscendGemma4Proposer;
  • build attention metadata with the correct per-group block table;
  • propagate late-bound KV-sharing targets to the Ascend attention implementation;
  • support Q-only RoPE with a one-head throwaway key;
  • preserve the draft-dimension LM head while installing the ACL graph wrapper;
  • keep model positions and attention metadata constant across Gemma4 draft steps;
  • add focused unit tests and CI test routing.

Compared with #13045, this removes duplicated Gemma4 initialization/configuration logic, keeps the common proposer changes behind default-preserving hooks/flags, and preserves the existing kernel block-size behavior for non-Gemma proposers.

Does this PR introduce any user-facing change?

Yes. It adds Gemma4 MTP speculative decoding support for the Ascend V1 model runner. Existing Eagle, MTP, DFlash, DSpark, and other proposer paths retain their previous defaults.

How was this patch tested?

Local validation:

  • ruff check and ruff format --check on all changed Python files;
  • targeted mypy checks under Python 3.10, 3.11, and 3.12;
  • YAML parsing for .github/workflows/scripts/test_config.yaml;
  • compileall on all changed Python files;
  • selective-test routing verification;
  • git diff --check.

Focused unit tests cover:

  • Gemma4 MTP proposer routing;
  • delegation to the upstream Gemma4 speculative-config override;
  • KV-sharing target propagation;
  • preservation of the draft LM head;
  • FIA speculative-decoding metadata;
  • per-group block-table selection for eager and graph metadata;
  • Q-only RoPE shape, dtype, and device handling.

The current local Windows environment does not provide PyTorch/pytest, torch_npu, or an Ascend device, so runtime tests were not executed locally. Before marking this PR ready for review, rerun the original A5 accuracy and performance cases from #13045, including acceptance-rate and GPQA checks, together with non-Gemma speculative-decoding regressions.

@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 MTP speculative decoding on Ascend NPUs by implementing a dedicated proposer class and adapting existing speculative decoding infrastructure. The changes include necessary adjustments to rotary embedding, KV-cache metadata handling, and model runner routing to ensure compatibility with Gemma4's specific architectural requirements, such as draft-dimension LM head preservation and eager centroid masking.

Highlights

  • Gemma4 MTP Support: Added the AscendGemma4Proposer class to enable Gemma4 MTP speculative decoding on Ascend NPUs.
  • Rotary Embedding Updates: Updated rotary embedding operations to support Q-only RoPE, accommodating Gemma4's specific attention requirements.
  • Metadata Refactoring: Refactored speculative decoding metadata building to support per-group block tables and improved KV-sharing target propagation.
  • Pipeline Integration: Integrated Gemma4 MTP routing into the model runner and speculative configuration patching to ensure seamless model execution.
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.

Ignored Files
  • Ignored by pattern: .github/workflows/** (1)
    • .github/workflows/scripts/test_config.yaml
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. ↩

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

[SpecDecode][Feature] Add Gemma4 MTP proposer support for Ascend NPUs

Suggested PR Summary:

### What this PR does / why we need it?
This pull request introduces support for the Gemma4 Multi-Token Prediction (MTP) proposer on Ascend NPUs. It adds the `AscendGemma4Proposer` class, integrates it into the model runner, and updates the rotary embedding operator to handle cases where the key tensor is `None` by utilizing a dummy key. Additionally, it updates the base proposer to support Gemma4 architectures and handle constant draft positions.

Feedback:
- Direct access to `self.constant_draft_positions` in `llm_base_proposer.py` will raise an `AttributeError` for other proposers (e.g., `AscendEagleProposer`). It is recommended to use `getattr(self, "constant_draft_positions", False)` instead.
- In `rotary_embedding.py`, allocating a full `torch.empty_like(query)` when `HAS_TRITON` is False is inefficient and could lead to OOMs. A smaller dummy key tensor representing a single head should be allocated instead.

### Does this PR introduce _any_ user-facing change?
Yes, it adds support for Gemma4 MTP speculative decoding on Ascend NPUs.

### How was this patch tested?
The changes are covered by new unit tests added in `tests/ut/ops/test_rotary_embedding.py`, `tests/ut/spec_decode/test_gemma4_proposer.py`, and `tests/ut/spec_decode/test_llm_base_proposer.py`.

Comment on lines +1561 to +1564
# MTP (constant_draft_positions): all draft tokens share the same
# target position — skip position and seq_len advancing.
if not self.constant_draft_positions:
used_update_positions += 1

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

Accessing self.constant_draft_positions directly will raise an AttributeError for other speculative decoding proposers (such as AscendEagleProposer) that do not inherit from Gemma4Proposer and thus do not have this attribute defined. Use getattr(self, "constant_draft_positions", False) to safely default to False for non-Gemma models.

Suggested change
# MTP (constant_draft_positions): all draft tokens share the same
# target position — skip position and seq_len advancing.
if not self.constant_draft_positions:
used_update_positions += 1
# MTP (constant_draft_positions): all draft tokens share the same
# target position — skip position and seq_len advancing.
if not getattr(self, "constant_draft_positions", False):
used_update_positions += 1

Comment on lines +1599 to +1618
# MTP (constant_draft_positions): keep target's last seq_lens.
if not self.constant_draft_positions:
common_attn_metadata.seq_lens[:batch_size] += 1
# For the requests that exceed the max model length, we set the
# sequence length to 1 to minimize their overheads in attention.
exceeds_mask = common_attn_metadata.seq_lens[:batch_size] > self.max_model_len
common_attn_metadata.seq_lens[:batch_size].masked_fill_(exceeds_mask, 1)
if common_attn_metadata.seq_lens_cpu is not None:
common_attn_metadata.seq_lens_cpu[:batch_size] = common_attn_metadata.seq_lens_cpu[:batch_size] + 1
if not self.constant_draft_positions:
common_attn_metadata.seq_lens_cpu[:batch_size] = common_attn_metadata.seq_lens_cpu[:batch_size] + 1
exceeds_mask_cpu = common_attn_metadata.seq_lens_cpu[:batch_size] > self.max_model_len
common_attn_metadata.seq_lens_cpu[:batch_size].masked_fill_(exceeds_mask_cpu, 1)
if common_attn_metadata._seq_lens_cpu is not None:
common_attn_metadata._seq_lens_cpu[:batch_size] = common_attn_metadata._seq_lens_cpu[:batch_size] + 1
if not self.constant_draft_positions:
common_attn_metadata._seq_lens_cpu[:batch_size] = common_attn_metadata._seq_lens_cpu[:batch_size] + 1
exceeds_mask_internal_cpu = common_attn_metadata._seq_lens_cpu[:batch_size] > self.max_model_len
common_attn_metadata._seq_lens_cpu[:batch_size].masked_fill_(exceeds_mask_internal_cpu, 1)
if common_attn_metadata.num_computed_tokens_cpu is not None:
common_attn_metadata.num_computed_tokens_cpu[:batch_size] += 1
if not self.constant_draft_positions:
common_attn_metadata.num_computed_tokens_cpu[:batch_size] += 1

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

Accessing self.constant_draft_positions directly will raise an AttributeError for other speculative decoding proposers (such as AscendEagleProposer) that do not inherit from Gemma4Proposer and thus do not have this attribute defined. Use getattr(self, "constant_draft_positions", False) to safely default to False for non-Gemma models.

Suggested change
# MTP (constant_draft_positions): keep target's last seq_lens.
if not self.constant_draft_positions:
common_attn_metadata.seq_lens[:batch_size] += 1
# For the requests that exceed the max model length, we set the
# sequence length to 1 to minimize their overheads in attention.
exceeds_mask = common_attn_metadata.seq_lens[:batch_size] > self.max_model_len
common_attn_metadata.seq_lens[:batch_size].masked_fill_(exceeds_mask, 1)
if common_attn_metadata.seq_lens_cpu is not None:
common_attn_metadata.seq_lens_cpu[:batch_size] = common_attn_metadata.seq_lens_cpu[:batch_size] + 1
if not self.constant_draft_positions:
common_attn_metadata.seq_lens_cpu[:batch_size] = common_attn_metadata.seq_lens_cpu[:batch_size] + 1
exceeds_mask_cpu = common_attn_metadata.seq_lens_cpu[:batch_size] > self.max_model_len
common_attn_metadata.seq_lens_cpu[:batch_size].masked_fill_(exceeds_mask_cpu, 1)
if common_attn_metadata._seq_lens_cpu is not None:
common_attn_metadata._seq_lens_cpu[:batch_size] = common_attn_metadata._seq_lens_cpu[:batch_size] + 1
if not self.constant_draft_positions:
common_attn_metadata._seq_lens_cpu[:batch_size] = common_attn_metadata._seq_lens_cpu[:batch_size] + 1
exceeds_mask_internal_cpu = common_attn_metadata._seq_lens_cpu[:batch_size] > self.max_model_len
common_attn_metadata._seq_lens_cpu[:batch_size].masked_fill_(exceeds_mask_internal_cpu, 1)
if common_attn_metadata.num_computed_tokens_cpu is not None:
common_attn_metadata.num_computed_tokens_cpu[:batch_size] += 1
if not self.constant_draft_positions:
common_attn_metadata.num_computed_tokens_cpu[:batch_size] += 1
# MTP (constant_draft_positions): keep target's last seq_lens.
if not getattr(self, "constant_draft_positions", False):
common_attn_metadata.seq_lens[:batch_size] += 1
# For the requests that exceed the max model length, we set the
# sequence length to 1 to minimize their overheads in attention.
exceeds_mask = common_attn_metadata.seq_lens[:batch_size] > self.max_model_len
common_attn_metadata.seq_lens[:batch_size].masked_fill_(exceeds_mask, 1)
if common_attn_metadata.seq_lens_cpu is not None:
if not getattr(self, "constant_draft_positions", False):
common_attn_metadata.seq_lens_cpu[:batch_size] = common_attn_metadata.seq_lens_cpu[:batch_size] + 1
exceeds_mask_cpu = common_attn_metadata.seq_lens_cpu[:batch_size] > self.max_model_len
common_attn_metadata.seq_lens_cpu[:batch_size].masked_fill_(exceeds_mask_cpu, 1)
if common_attn_metadata._seq_lens_cpu is not None:
if not getattr(self, "constant_draft_positions", False):
common_attn_metadata._seq_lens_cpu[:batch_size] = common_attn_metadata._seq_lens_cpu[:batch_size] + 1
exceeds_mask_internal_cpu = common_attn_metadata._seq_lens_cpu[:batch_size] > self.max_model_len
common_attn_metadata._seq_lens_cpu[:batch_size].masked_fill_(exceeds_mask_internal_cpu, 1)
if common_attn_metadata.num_computed_tokens_cpu is not None:
if not getattr(self, "constant_draft_positions", False):
common_attn_metadata.num_computed_tokens_cpu[:batch_size] += 1

Comment on lines +252 to +256
dummy_key = (
torch.empty(query.shape[0], 0, self.head_size, dtype=query.dtype, device=query.device)
if HAS_TRITON
else torch.empty_like(query)
)

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

Allocating a full torch.empty_like(query) tensor on every forward pass of every layer when HAS_TRITON is False is highly inefficient and can lead to significant memory overhead and potential OOMs for large batch sizes or prefill phases. Since we only need a dummy key to satisfy the operator signature, we can allocate a much smaller tensor representing just a single head (either 2D or 3D depending on the query dimensions).

            dummy_key = (
                torch.empty(query.shape[0], 0, self.head_size, dtype=query.dtype, device=query.device)
                if HAS_TRITON
                else (
                    torch.empty(query.shape[0], 1, self.head_size, dtype=query.dtype, device=query.device)
                    if query.ndim == 3
                    else torch.empty(query.shape[0], self.head_size, dtype=query.dtype, device=query.device)
                )
            )

@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!

@Alex-stack-hub
Alex-stack-hub force-pushed the refactor/gemma4-mtp-upstream branch 3 times, most recently from 7160a53 to f98b96f Compare July 31, 2026 09:50
Signed-off-by: XUE TONGYAO <268319265+Alex-stack-hub@users.noreply.github.com>
@Alex-stack-hub
Alex-stack-hub force-pushed the refactor/gemma4-mtp-upstream branch from f98b96f to 3e80309 Compare July 31, 2026 09:51
@github-actions

github-actions Bot commented Aug 5, 2026

Copy link
Copy Markdown
Contributor

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

swimming2007-doge pushed a commit to swimming2007-doge/vllm-ascend that referenced this pull request Sep 8, 2026
Port Gemma4 MTP proposer onto v0.26.0rc1, integrating vllm-project#13045 (runtime-proven
on A5/950) with the refactored structure of vllm-project#13263:

- AscendGemma4Proposer reusing upstream vLLM Gemma4Proposer: eager centroids
  sampling (no CUDA graphs), KV-sharing target sync to Ascend impls, per-group
  block table metadata, FIA SpecDecoding state
- route mtp+use_gemma4_mtp() to the Ascend proposer; delegate gemma4_assistant
  hf_config_override to upstream vLLM
- Q-only RoPE via throwaway key buffer when key is None (Triton path uses a
  zero-kv-head dummy)
- llm_base_proposer: default-preserving multi-group graph-capture and
  layer-name hooks; constant_draft_positions guards in _run_merged_draft and
  attn_update_stack_num_spec_norm
- model_runner_v1: wire AscendGemma4Proposer (drafter union, per-group block
  table capture, kv-cache / cudagraph-key asserts)
- unit tests for proposer routing, KV-sharing sync, per-group metadata,
  Q-only RoPE; CI test routing
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