Repository navigation
[Feature] Add minimal Gemma4 MTP support on Ascend - #13263
Alex-stack-hub wants to merge 1 commit 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 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
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
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
|
There was a problem hiding this comment.
Code Review
Suggested PR Title:
[SpecDecode][Feature] Add Gemma4 MTP proposer support for Ascend NPUsSuggested 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`.| # 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 |
There was a problem hiding this comment.
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.
| # 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 |
| # 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 |
There was a problem hiding this comment.
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.
| # 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 |
| 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) | ||
| ) |
There was a problem hiding this comment.
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)
)
)|
👋 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! |
7160a53 to
f98b96f
Compare
Signed-off-by: XUE TONGYAO <268319265+Alex-stack-hub@users.noreply.github.com>
f98b96f to
3e80309
Compare
|
This pull request has conflicts, please resolve those before we can evaluate the pull request. |
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
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
mainbranch and the vLLM v0.26.0 integration point.Gemma4 MTP differs from existing Ascend proposers in a few required ways:
The implementation reuses vLLM's upstream
Gemma4Proposerfor model-generic behavior and keeps only Ascend-specific adaptation:AscendGemma4Proposer;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 checkandruff format --checkon all changed Python files;.github/workflows/scripts/test_config.yaml;compileallon all changed Python files;git diff --check.Focused unit tests cover:
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.