Repository navigation
[Feature]Gemma MTP speculative decoding support on Ascend NPU A5 - #13045
Conversation
|
👋 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! |
|
This pull request has conflicts, please resolve those before we can evaluate the pull request. |
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 NPU hardware. By implementing a new proposer class that leverages multiple inheritance, the changes integrate seamlessly with existing vLLM infrastructure while addressing platform-specific requirements for multi-group KV cache management, attention head alignment, and eager sampling. The implementation ensures that Gemma4 models can run efficiently on Ascend NPUs without requiring significant rewrites of core vLLM internals. 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:
[Ops/SpecDecode][Feature] Add Gemma4 MTP speculative decoding support for Ascend NPUSuggested PR Summary:
### What this PR does / why we need it?
This PR introduces support for Gemma4 MTP speculative decoding on Ascend NPUs. It adds a Q-only RoPE helper (`gemma4_q_only_rope`) to handle Gemma4's sliding-attention layers, implements the `AscendGemma4Proposer` by combining upstream GPU logic with Ascend-specific NPU initialization, and updates the base proposer and model runner to support multi-group KV cache models with per-group block tables and constant draft positions.
Feedback and improvement opportunities:
- **Shape Mismatch in RoPE**: In `rope_q_only.py`, `gemma4_q_only_rope` returns different shapes (2D vs 3D) depending on whether Triton is used. The original query shape should be saved and restored before returning.
- **Potential AttributeError in Proposer**: In `gemma4_proposer.py`, accessing `layer.self_attn.attn` directly may crash if a layer lacks these attributes. Use `getattr` for safe access.
- **AttributeError in Base Proposer**: In `llm_base_proposer.py`, accessing `self.constant_draft_positions` directly will crash other proposers (e.g., Eagle, Medusa) that do not define this attribute. Use `getattr(self, "constant_draft_positions", False)` 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?
CI and manual verification of speculative decoding with Gemma4 MTP models.da5eaac to
1fc8d2b
Compare
Smoking test script of Acceptancepython3 /home/wumeng/test/gemma4-test/measure_acceptance.py import argparse
import json
import re
import time
import requests
URL = "http://127.0.0.1:8831"
MODEL = "gemma4-31b-it-mtp"
K = 3
# 5 coding prompts used for the 94.8% baseline (15 reqs = 5 × 3 repeats).
_EMBEDDED_PROMPTS = [
(
"Complete the following Python function. Only output the function "
"body, no explanation.\n\n"
"def binary_search(arr, target):\n"
' """Return the index of target in sorted array arr, or -1 if not found."""'
),
(
"Complete the following Python function. Only output the function "
"body, no explanation.\n\n"
"def merge_sort(arr):\n"
' """Sort the array arr in ascending order and return it."""'
),
(
"Complete the following Python function. Only output the function "
"body, no explanation.\n\n"
"def lru_cache_get(cache, key):\n"
' """Return value for key from an OrderedDict-backed LRU cache, '
"or None. Mark as recently used on hit.\"\"\""
),
(
"Complete the following Python function. Only output the function "
"body, no explanation.\n\n"
"def is_balanced(s):\n"
' """Return True if the string s has balanced '
"parentheses/brackets/braces, else False.\"\"\""
),
(
"Complete the following Python function. Only output the function "
"body, no explanation.\n\n"
"def flatten(nested):\n"
' """Yield elements from an arbitrarily nested list, depth-first."""'
),
]
def scrape(url: str, k: int):
"""Return (num_drafts_total, [accepted_per_pos_0, ..., accepted_per_pos_{k-1}])."""
txt = requests.get(f"{url}/metrics", timeout=10).text
def num(pat):
m = re.search(pat, txt)
return float(m.group(1)) if m else 0.0
drafts = num(r'vllm:spec_decode_num_drafts_total\{[^}]*\}\s+([\d.eE+-]+)')
acc = []
for pos in range(k):
acc.append(
num(
r'vllm:spec_decode_num_accepted_tokens_per_pos_total\{[^}]*position="%d"[^}]*\}\s+([\d.eE+-]+)'
% pos
)
)
return drafts, acc
def drive(
prompts: list[str],
url: str,
model: str,
repeats: int = 3,
max_tokens: int = 256,
):
"""Send chat completion requests and return (num_requests, elapsed_seconds)."""
n = 0
t0 = time.time()
for _ in range(repeats):
for q in prompts:
requests.post(
f"{url}/v1/chat/completions",
json={
"model": model,
"messages": [{"role": "user", "content": q}],
"max_tokens": max_tokens,
"temperature": 0,
},
timeout=300,
)
n += 1
return n, time.time() - t0
def main():
parser = argparse.ArgumentParser(description="Measure MTP acceptance rate")
parser.add_argument("--url", default=URL)
parser.add_argument("--model", default=MODEL)
parser.add_argument("--k", type=int, default=K, help="num_speculative_tokens")
parser.add_argument("--repeats", type=int, default=3)
parser.add_argument(
"--prompts",
default=None,
help="path to an optional jsonl prompts file (overrides the built-in coding prompts)",
)
args = parser.parse_args()
# Load prompts: file if given, otherwise the embedded coding-problem set.
if args.prompts:
prompts = [
json.loads(line)["question"]
for line in open(args.prompts)
if line.strip()
]
else:
prompts = list(_EMBEDDED_PROMPTS)
d0, a0 = scrape(args.url, args.k)
print(f"[baseline] drafts={d0:.0f} accepted_per_pos={[round(x) for x in a0]}")
n, dt = drive(prompts, args.url, args.model, args.repeats)
d1, a1 = scrape(args.url, args.k)
dd = d1 - d0
da = [a1[i] - a0[i] for i in range(args.k)]
print(f"[after] drafts={d1:.0f} accepted_per_pos={[round(x) for x in a1]}")
print(
f"[delta] drafts=+{dd:.0f} accepted_per_pos=+{[round(x) for x in da]} "
f"({n} reqs, {dt:.1f}s)"
)
# Per-position acceptance rate = accepted_tokens_delta / num_drafts_delta
rates = [a / dd for a in da] if dd else [0.0] * args.k
print(
f"\n=== gemma4 MTP k={args.k} acceptance "
"(coding prompts, greedy) ==="
)
for i, r in enumerate(rates):
print(f" pos{i}: {r * 100:.1f}%")
agg = sum(rates) / args.k
print(
f" aggregate(avg/pos): {agg * 100:.1f}% "
f"avg accepted {sum(rates):.2f}/{args.k} token"
)
if __name__ == "__main__":
main() |
cd30b64 to
8042de2
Compare
|
This pull request has conflicts, please resolve those before we can evaluate the pull request. |
232ab62 to
cece87a
Compare
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
cece87a to
3d2ca8a
Compare
The rebase onto main lost the _get_attn_metadata_layer_names hook call in the _propose step-update loop, reverting to self.attn_layer_names. With Gemma4's two KV cache groups (sliding + full attention), every group's call overwrote the metadata of ALL draft layers, so steps 1+ read the wrong group's block table while step 0 (build_draft_attn_ metadata) stayed per-group correct -- acceptance degraded only at pos1/pos2 (77/62/52 reference collapsed to 77/~50/~35). Restore the hook call: single-group drafters keep the previous behavior via the default hook (returns self.attn_layer_names); AscendGemma4Proposer returns attn_group.layer_names, matching the original PR 13045 logic. Co-Authored-By: Claude Code <noreply@anthropic.com> Signed-off-by: Wu Meng <wumeng@ascend-debug.local> Co-authored-by: XUE TONGYAO <xuetongyao2001@gmail.com>
|
需要确认Gemma4 MTP是否能复用DCP,如果叠加DCP,现在的代码版本应该存在问题 |
|
第二点就是关于接受率,建议增加接受率的CI测试看护 |
即实际调用Gemma4模型进行看护,例如MoE的那个 |
|
还有就是文档方面也需要同步更新,传递给其他开发者相关信息 |
The failed test_deepseek_v4_dsa_pcp_dspark asserted pos2 acceptance 0.4933 < 0.55 while the same run showed 0.490 and 0.587 across metric windows (+/-10pp variance). Workflow changes blocked a plain rerun. Co-Authored-By: Claude Code <noreply@anthropic.com> Signed-off-by: Wu Meng <wumeng@ascend-debug.local> Co-authored-by: XUE TONGYAO <xuetongyao2001@gmail.com>
a278c24 to
a1640c9
Compare
后面在V2的时候补充E2E用例 |
在V2的时候补充文档 |
DCP暂不支持,v2可以考虑 |
…m-project#13045) ### What this PR does / why we need it? Adds **Gemma4 MTP speculative decoding** support on Ascend NPU (A5 / 950PR). Gemma4 (`gemma-4-31B-it` + `gemma-4-31B-it-assistant`) is a **multi-group KV cache** model: layers are split into a `sliding` group (head_dim=256) and a `full_attention` group (global head_dim=512). The MTP draft reuses the target's architecture and **shares the target's per-layer KV cache** (sliding for layer 58, global for layer 59), so each draft attention group must read K/V from its own physical block layout. Upstream vLLM supports this on CUDA via `Gemma4Proposer`. This PR adds the Ascend path as a **thin wrapper** of the upstream proposer — no rewrite of vLLM internals, no forking of base-class draft-loop methods. | file | change | role | |---|---|---| | `spec_decode/gemma4_proposer.py` | **+new (150)** | `AscendGemma4Proposer` — multiple inheritance of upstream `Gemma4Proposer` + `AscendSpecDecodeBaseProposer`. Overrides only: eager centroids sampling, `_sync_kv_sharing_target_to_impl`, per-group block table swap (`attn_update_stack_num_spec_norm` wrapper), `build_draft_attn_metadata` with FIA SpecDecoding state. | | `spec_decode/llm_base_proposer.py` | mod (+31/-11) | (1) `constant_draft_positions` guards in `_run_merged_draft` and `attn_update_stack_num_spec_norm` — restores upstream `SpecDecodeBaseProposer` semantics (vllm/v1/spec_decode/llm_base_proposer.py:635,693,705), no-op for existing proposers (flag defaults False upstream). (2) `_build_multi_group_graph_capture_metadata` hook (returns None for existing proposers) for per-group ACL graph capture. (3) `_propose` step-update loop iterates `attn_group.layer_names`, matching `build_draft_attn_metadata`'s existing iteration. | | `ops/rotary_embedding.py` | mod (+22/-1) | Q-only RoPE: when `key is None` (Gemma4 MTP sliding layers, K/V from target cache), use a throwaway key buffer with the regular rotary implementation (zero-kv-head dummy on the Triton path). | | `worker/model_runner_v1.py` | mod (+19/-2) | Wire `AscendGemma4Proposer` into the drafter type union, per-group block table capture, and kv-cache / cudagraph-key asserts. | | `spec_decode/__init__.py` | mod (+3) | Route `mtp` + `use_gemma4_mtp()` to `AscendGemma4Proposer`. | | `tests/ut/spec_decode/test_gemma4_proposer.py` | **+new (130)** | Unit tests: proposer routing, KV-sharing sync, per-group metadata. | | `tests/ut/ops/test_rotary_embedding.py` | mod (+18) | Q-only RoPE test. | ### Does this PR introduce _any_ user-facing change? Additive only. Users can now run Gemma4 MTP on Ascend with `--speculative-config '{"method":"mtp",...}'`. No change to existing models. ### How was this patch tested? #### Acceptance Gemma4-31B-it MTP, k=3, greedy (temp=0), `FULL_DECODE_ONLY` graph mode, coding prompt suite (5 prompts × 3 = 15 reqs). Metric: `vllm:spec_decode_num_accepted_tokens_per_pos` ÷ `vllm:spec_decode_num_drafts`, scraped from `/metrics`. | platform | backend | TP | pos0 | pos1 | pos2 | agg | accepted/step | |---|---|---|---|---|---|---|---| | **Ascend A5 (950PR) —TP4DP1** | vllm-ascend | 1 | **98.2%** | **95.4%** | **90.8%** | **94.8%** | **2.85/3** | | **Ascend A5 (950DT) —TP1DP4** | vllm-ascend | 1×4 | 100.0% | 98.8% | 91.9% | 96.9% | 2.91/3 | | NVIDIA L20 — CUDA reference | upstream vLLM | 2 | 98.2% | 95.5% | 90.0% | 94.6% | 2.84/3 | Ascend now matches the CUDA reference. Here is the prompts: ```python #{"question": "Complete the following Python function. Only output the function body, no explanation.\n\ndef binary_search(arr, target):\n \"\"\"Return the index of target in sorted array arr, or -1 if not found.\"\"\"", "max_out_len": 512, "answer": ""} #{"question": "Complete the following Python function. Only output the function body, no explanation.\n\ndef merge_sort(arr):\n \"\"\"Sort the array arr in ascending order and return it.\"\"\"", "max_out_len": 512, "answer": ""} #{"question": "Complete the following Python function. Only output the function body, no explanation.\n\ndef lru_cache_get(cache, key):\n \"\"\"Return value for key from an OrderedDict-backed LRU cache, or None. Mark as recently used on hit.\"\"\"", "max_out_len": 512, "answer": ""} #{"question": "Complete the following Python function. Only output the function body, no explanation.\n\ndef is_balanced(s):\n \"\"\"Return True if the string s has balanced parentheses/brackets/braces, else False.\"\"\"", "max_out_len": 512, "answer": ""} #{"question": "Complete the following Python function. Only output the function body, no explanation.\n\ndef flatten(nested):\n \"\"\"Yield elements from an arbitrarily nested list, depth-first.\"\"\"", "max_out_len": 512, "answer": ""} ``` #### Accuracy - Matches the raw target performance - Tested on the **gpqa_diamond** dataset by using Evalscope 1. Accuracy Metrics:**Score: 0.7778**,Sample Size: 198 samples 2. Average Latency: 15.0612 seconds 3. Average Throughput: 65.59 tokens/second (Avg Thpt) #### Performance (950PR TP=4,DP=1) **Overall Conclusion:** Tested on 4K input/1K output by Evalscope. The model demonstrates excellent scalability and stability. As concurrency increases, the overall throughput grows significantly while maintaining a consistent speculative acceptance rate, proving that the MTP (Multi-Token Prediction) mechanism is functioning effectively. ##### **1. Key Latency & Efficiency Metrics** * **TPOT (Time Per Output Token):** * At low concurrency (**Conc=1**), the average TPOT is very low (**15.0 ms**), indicating extremely fast generation. * As concurrency increases to **64**, the average TPOT rises to **81.0 ms**, which is expected due to increased system load, but remains within an acceptable range for high-throughput serving. * **Speculative Acceptance Rate:** * The acceptance rate remains remarkably stable across all concurrency levels, hovering around **71% - 72%**. This indicates that the MTP model's predictions are consistently accurate regardless of the system load. ##### **2. Throughput and Scalability** * **Completion Throughput (Output Speed):** * There is a clear linear increase in output speed as concurrency rises. * **Conc 1:** 59.46 tok/s $\rightarrow$ **Conc 64:** 345.05 tok/s. * **Total System Throughput:** * The **Overall Total Prompt throughput** scales from **237.85 tok/s** (Conc 1) up to **1380.35 tok/s** (Conc 64). * The **"Last 30s"** peak throughput reaches nearly **3,924 tok/s** at maximum concurrency, showing the model's ability to handle heavy bursts of traffic. ##### **3. Summary Table** | Concurrency | Avg TPOT (ms) | Spec. Accept Rate | Completion Throughput | Total Prompt Throughput | | :--- | :--- | :--- | :--- | :--- | | **1** | 15.0 | 71.5% | 59.46 tok/s | 237.85 tok/s | | **8** | 30.9 | 71.6% | 236.27 tok/s | 945.17 tok/s | | **16** | 43.3 | 71.7% | 299.15 tok/s | 1196.73 tok/s | | **32** | 79.4 | 71.5% | 330.30 tok/s | 1321.32 tok/s | | **64** | 81.0 | 71.6% | 345.05 tok/s | 1380.35 tok/s | #### Additional: Real-world Agent Prompt Validation (950DT TP=1,DP=4) Tested with a ~3.4K-token real tel-sales agent system prompt (300 output tokens, `vllm bench serve`, greedy, ignore-eos, TP=1): | Conc | TPOT p99 (ms) | Output tok/s | Spec accept | |---|---|---|---| | 1 | 9.6 | 102 | 63.4% | | 8 | 10.9 | 732 | 63.6% | | 16 | 12.2 | 1290 | 63.6% | | 32 | 14.3 | 2267 | 63.5% | | 48 | 16.0 | 2981 | 63.6% | | 64 | 17.2 | 3469 | 63.9% | Acceptance stays flat across concurrency; throughput scales linearly to 3469 tok/s at conc=64. Zero errors across the full sweep. Repro (copy-paste-able): ```bash ASCEND_RT_VISIBLE_DEVICES=0 vllm serve <gemma-4-31B-it> \ --served-model-name gemma4-31b-it-mtp \ --tensor-parallel-size 1 \ --speculative-config '{"method":"mtp","model":<gemma-4-31B-it-assistant>,"num_speculative_tokens":3}' \ --max-model-len 32000 --gpu-memory-utilization 0.85 \ --compilation-config '{"cudagraph_mode":"FULL_DECODE_ONLY"}' \ --trust-remote-code --host 0.0.0.0 --port 8831 # then drive the server and scrape /metrics for per-position acceptance ``` - vLLM main: vllm-project/vllm@b2f6858 --------- Signed-off-by: Wu Meng <wumeng@ascend-debug.local> Co-authored-by: Wu Meng <wumeng@ascend-debug.local> Co-authored-by: Claude Code <noreply@anthropic.com> Co-authored-by: XUE TONGYAO <xuetongyao2001@gmail.com>
…m-project#13045) ### What this PR does / why we need it? Adds **Gemma4 MTP speculative decoding** support on Ascend NPU (A5 / 950PR). Gemma4 (`gemma-4-31B-it` + `gemma-4-31B-it-assistant`) is a **multi-group KV cache** model: layers are split into a `sliding` group (head_dim=256) and a `full_attention` group (global head_dim=512). The MTP draft reuses the target's architecture and **shares the target's per-layer KV cache** (sliding for layer 58, global for layer 59), so each draft attention group must read K/V from its own physical block layout. Upstream vLLM supports this on CUDA via `Gemma4Proposer`. This PR adds the Ascend path as a **thin wrapper** of the upstream proposer — no rewrite of vLLM internals, no forking of base-class draft-loop methods. | file | change | role | |---|---|---| | `spec_decode/gemma4_proposer.py` | **+new (150)** | `AscendGemma4Proposer` — multiple inheritance of upstream `Gemma4Proposer` + `AscendSpecDecodeBaseProposer`. Overrides only: eager centroids sampling, `_sync_kv_sharing_target_to_impl`, per-group block table swap (`attn_update_stack_num_spec_norm` wrapper), `build_draft_attn_metadata` with FIA SpecDecoding state. | | `spec_decode/llm_base_proposer.py` | mod (+31/-11) | (1) `constant_draft_positions` guards in `_run_merged_draft` and `attn_update_stack_num_spec_norm` — restores upstream `SpecDecodeBaseProposer` semantics (vllm/v1/spec_decode/llm_base_proposer.py:635,693,705), no-op for existing proposers (flag defaults False upstream). (2) `_build_multi_group_graph_capture_metadata` hook (returns None for existing proposers) for per-group ACL graph capture. (3) `_propose` step-update loop iterates `attn_group.layer_names`, matching `build_draft_attn_metadata`'s existing iteration. | | `ops/rotary_embedding.py` | mod (+22/-1) | Q-only RoPE: when `key is None` (Gemma4 MTP sliding layers, K/V from target cache), use a throwaway key buffer with the regular rotary implementation (zero-kv-head dummy on the Triton path). | | `worker/model_runner_v1.py` | mod (+19/-2) | Wire `AscendGemma4Proposer` into the drafter type union, per-group block table capture, and kv-cache / cudagraph-key asserts. | | `spec_decode/__init__.py` | mod (+3) | Route `mtp` + `use_gemma4_mtp()` to `AscendGemma4Proposer`. | | `tests/ut/spec_decode/test_gemma4_proposer.py` | **+new (130)** | Unit tests: proposer routing, KV-sharing sync, per-group metadata. | | `tests/ut/ops/test_rotary_embedding.py` | mod (+18) | Q-only RoPE test. | ### Does this PR introduce _any_ user-facing change? Additive only. Users can now run Gemma4 MTP on Ascend with `--speculative-config '{"method":"mtp",...}'`. No change to existing models. ### How was this patch tested? #### Acceptance Gemma4-31B-it MTP, k=3, greedy (temp=0), `FULL_DECODE_ONLY` graph mode, coding prompt suite (5 prompts × 3 = 15 reqs). Metric: `vllm:spec_decode_num_accepted_tokens_per_pos` ÷ `vllm:spec_decode_num_drafts`, scraped from `/metrics`. | platform | backend | TP | pos0 | pos1 | pos2 | agg | accepted/step | |---|---|---|---|---|---|---|---| | **Ascend A5 (950PR) —TP4DP1** | vllm-ascend | 1 | **98.2%** | **95.4%** | **90.8%** | **94.8%** | **2.85/3** | | **Ascend A5 (950DT) —TP1DP4** | vllm-ascend | 1×4 | 100.0% | 98.8% | 91.9% | 96.9% | 2.91/3 | | NVIDIA L20 — CUDA reference | upstream vLLM | 2 | 98.2% | 95.5% | 90.0% | 94.6% | 2.84/3 | Ascend now matches the CUDA reference. Here is the prompts: ```python #{"question": "Complete the following Python function. Only output the function body, no explanation.\n\ndef binary_search(arr, target):\n \"\"\"Return the index of target in sorted array arr, or -1 if not found.\"\"\"", "max_out_len": 512, "answer": ""} #{"question": "Complete the following Python function. Only output the function body, no explanation.\n\ndef merge_sort(arr):\n \"\"\"Sort the array arr in ascending order and return it.\"\"\"", "max_out_len": 512, "answer": ""} #{"question": "Complete the following Python function. Only output the function body, no explanation.\n\ndef lru_cache_get(cache, key):\n \"\"\"Return value for key from an OrderedDict-backed LRU cache, or None. Mark as recently used on hit.\"\"\"", "max_out_len": 512, "answer": ""} #{"question": "Complete the following Python function. Only output the function body, no explanation.\n\ndef is_balanced(s):\n \"\"\"Return True if the string s has balanced parentheses/brackets/braces, else False.\"\"\"", "max_out_len": 512, "answer": ""} #{"question": "Complete the following Python function. Only output the function body, no explanation.\n\ndef flatten(nested):\n \"\"\"Yield elements from an arbitrarily nested list, depth-first.\"\"\"", "max_out_len": 512, "answer": ""} ``` #### Accuracy - Matches the raw target performance - Tested on the **gpqa_diamond** dataset by using Evalscope 1. Accuracy Metrics:**Score: 0.7778**,Sample Size: 198 samples 2. Average Latency: 15.0612 seconds 3. Average Throughput: 65.59 tokens/second (Avg Thpt) #### Performance (950PR TP=4,DP=1) **Overall Conclusion:** Tested on 4K input/1K output by Evalscope. The model demonstrates excellent scalability and stability. As concurrency increases, the overall throughput grows significantly while maintaining a consistent speculative acceptance rate, proving that the MTP (Multi-Token Prediction) mechanism is functioning effectively. ##### **1. Key Latency & Efficiency Metrics** * **TPOT (Time Per Output Token):** * At low concurrency (**Conc=1**), the average TPOT is very low (**15.0 ms**), indicating extremely fast generation. * As concurrency increases to **64**, the average TPOT rises to **81.0 ms**, which is expected due to increased system load, but remains within an acceptable range for high-throughput serving. * **Speculative Acceptance Rate:** * The acceptance rate remains remarkably stable across all concurrency levels, hovering around **71% - 72%**. This indicates that the MTP model's predictions are consistently accurate regardless of the system load. ##### **2. Throughput and Scalability** * **Completion Throughput (Output Speed):** * There is a clear linear increase in output speed as concurrency rises. * **Conc 1:** 59.46 tok/s $\rightarrow$ **Conc 64:** 345.05 tok/s. * **Total System Throughput:** * The **Overall Total Prompt throughput** scales from **237.85 tok/s** (Conc 1) up to **1380.35 tok/s** (Conc 64). * The **"Last 30s"** peak throughput reaches nearly **3,924 tok/s** at maximum concurrency, showing the model's ability to handle heavy bursts of traffic. ##### **3. Summary Table** | Concurrency | Avg TPOT (ms) | Spec. Accept Rate | Completion Throughput | Total Prompt Throughput | | :--- | :--- | :--- | :--- | :--- | | **1** | 15.0 | 71.5% | 59.46 tok/s | 237.85 tok/s | | **8** | 30.9 | 71.6% | 236.27 tok/s | 945.17 tok/s | | **16** | 43.3 | 71.7% | 299.15 tok/s | 1196.73 tok/s | | **32** | 79.4 | 71.5% | 330.30 tok/s | 1321.32 tok/s | | **64** | 81.0 | 71.6% | 345.05 tok/s | 1380.35 tok/s | #### Additional: Real-world Agent Prompt Validation (950DT TP=1,DP=4) Tested with a ~3.4K-token real tel-sales agent system prompt (300 output tokens, `vllm bench serve`, greedy, ignore-eos, TP=1): | Conc | TPOT p99 (ms) | Output tok/s | Spec accept | |---|---|---|---| | 1 | 9.6 | 102 | 63.4% | | 8 | 10.9 | 732 | 63.6% | | 16 | 12.2 | 1290 | 63.6% | | 32 | 14.3 | 2267 | 63.5% | | 48 | 16.0 | 2981 | 63.6% | | 64 | 17.2 | 3469 | 63.9% | Acceptance stays flat across concurrency; throughput scales linearly to 3469 tok/s at conc=64. Zero errors across the full sweep. Repro (copy-paste-able): ```bash ASCEND_RT_VISIBLE_DEVICES=0 vllm serve <gemma-4-31B-it> \ --served-model-name gemma4-31b-it-mtp \ --tensor-parallel-size 1 \ --speculative-config '{"method":"mtp","model":<gemma-4-31B-it-assistant>,"num_speculative_tokens":3}' \ --max-model-len 32000 --gpu-memory-utilization 0.85 \ --compilation-config '{"cudagraph_mode":"FULL_DECODE_ONLY"}' \ --trust-remote-code --host 0.0.0.0 --port 8831 # then drive the server and scrape /metrics for per-position acceptance ``` - vLLM main: vllm-project/vllm@b2f6858 --------- Signed-off-by: Wu Meng <wumeng@ascend-debug.local> Co-authored-by: Wu Meng <wumeng@ascend-debug.local> Co-authored-by: Claude Code <noreply@anthropic.com> Co-authored-by: XUE TONGYAO <xuetongyao2001@gmail.com> Signed-off-by: tianming2009 <13246728590@163.com>
…m-project#13045) ### What this PR does / why we need it? Adds **Gemma4 MTP speculative decoding** support on Ascend NPU (A5 / 950PR). Gemma4 (`gemma-4-31B-it` + `gemma-4-31B-it-assistant`) is a **multi-group KV cache** model: layers are split into a `sliding` group (head_dim=256) and a `full_attention` group (global head_dim=512). The MTP draft reuses the target's architecture and **shares the target's per-layer KV cache** (sliding for layer 58, global for layer 59), so each draft attention group must read K/V from its own physical block layout. Upstream vLLM supports this on CUDA via `Gemma4Proposer`. This PR adds the Ascend path as a **thin wrapper** of the upstream proposer — no rewrite of vLLM internals, no forking of base-class draft-loop methods. | file | change | role | |---|---|---| | `spec_decode/gemma4_proposer.py` | **+new (150)** | `AscendGemma4Proposer` — multiple inheritance of upstream `Gemma4Proposer` + `AscendSpecDecodeBaseProposer`. Overrides only: eager centroids sampling, `_sync_kv_sharing_target_to_impl`, per-group block table swap (`attn_update_stack_num_spec_norm` wrapper), `build_draft_attn_metadata` with FIA SpecDecoding state. | | `spec_decode/llm_base_proposer.py` | mod (+31/-11) | (1) `constant_draft_positions` guards in `_run_merged_draft` and `attn_update_stack_num_spec_norm` — restores upstream `SpecDecodeBaseProposer` semantics (vllm/v1/spec_decode/llm_base_proposer.py:635,693,705), no-op for existing proposers (flag defaults False upstream). (2) `_build_multi_group_graph_capture_metadata` hook (returns None for existing proposers) for per-group ACL graph capture. (3) `_propose` step-update loop iterates `attn_group.layer_names`, matching `build_draft_attn_metadata`'s existing iteration. | | `ops/rotary_embedding.py` | mod (+22/-1) | Q-only RoPE: when `key is None` (Gemma4 MTP sliding layers, K/V from target cache), use a throwaway key buffer with the regular rotary implementation (zero-kv-head dummy on the Triton path). | | `worker/model_runner_v1.py` | mod (+19/-2) | Wire `AscendGemma4Proposer` into the drafter type union, per-group block table capture, and kv-cache / cudagraph-key asserts. | | `spec_decode/__init__.py` | mod (+3) | Route `mtp` + `use_gemma4_mtp()` to `AscendGemma4Proposer`. | | `tests/ut/spec_decode/test_gemma4_proposer.py` | **+new (130)** | Unit tests: proposer routing, KV-sharing sync, per-group metadata. | | `tests/ut/ops/test_rotary_embedding.py` | mod (+18) | Q-only RoPE test. | ### Does this PR introduce _any_ user-facing change? Additive only. Users can now run Gemma4 MTP on Ascend with `--speculative-config '{"method":"mtp",...}'`. No change to existing models. ### How was this patch tested? #### Acceptance Gemma4-31B-it MTP, k=3, greedy (temp=0), `FULL_DECODE_ONLY` graph mode, coding prompt suite (5 prompts × 3 = 15 reqs). Metric: `vllm:spec_decode_num_accepted_tokens_per_pos` ÷ `vllm:spec_decode_num_drafts`, scraped from `/metrics`. | platform | backend | TP | pos0 | pos1 | pos2 | agg | accepted/step | |---|---|---|---|---|---|---|---| | **Ascend A5 (950PR) —TP4DP1** | vllm-ascend | 1 | **98.2%** | **95.4%** | **90.8%** | **94.8%** | **2.85/3** | | **Ascend A5 (950DT) —TP1DP4** | vllm-ascend | 1×4 | 100.0% | 98.8% | 91.9% | 96.9% | 2.91/3 | | NVIDIA L20 — CUDA reference | upstream vLLM | 2 | 98.2% | 95.5% | 90.0% | 94.6% | 2.84/3 | Ascend now matches the CUDA reference. Here is the prompts: ```python #{"question": "Complete the following Python function. Only output the function body, no explanation.\n\ndef binary_search(arr, target):\n \"\"\"Return the index of target in sorted array arr, or -1 if not found.\"\"\"", "max_out_len": 512, "answer": ""} #{"question": "Complete the following Python function. Only output the function body, no explanation.\n\ndef merge_sort(arr):\n \"\"\"Sort the array arr in ascending order and return it.\"\"\"", "max_out_len": 512, "answer": ""} #{"question": "Complete the following Python function. Only output the function body, no explanation.\n\ndef lru_cache_get(cache, key):\n \"\"\"Return value for key from an OrderedDict-backed LRU cache, or None. Mark as recently used on hit.\"\"\"", "max_out_len": 512, "answer": ""} #{"question": "Complete the following Python function. Only output the function body, no explanation.\n\ndef is_balanced(s):\n \"\"\"Return True if the string s has balanced parentheses/brackets/braces, else False.\"\"\"", "max_out_len": 512, "answer": ""} #{"question": "Complete the following Python function. Only output the function body, no explanation.\n\ndef flatten(nested):\n \"\"\"Yield elements from an arbitrarily nested list, depth-first.\"\"\"", "max_out_len": 512, "answer": ""} ``` #### Accuracy - Matches the raw target performance - Tested on the **gpqa_diamond** dataset by using Evalscope 1. Accuracy Metrics:**Score: 0.7778**,Sample Size: 198 samples 2. Average Latency: 15.0612 seconds 3. Average Throughput: 65.59 tokens/second (Avg Thpt) #### Performance (950PR TP=4,DP=1) **Overall Conclusion:** Tested on 4K input/1K output by Evalscope. The model demonstrates excellent scalability and stability. As concurrency increases, the overall throughput grows significantly while maintaining a consistent speculative acceptance rate, proving that the MTP (Multi-Token Prediction) mechanism is functioning effectively. ##### **1. Key Latency & Efficiency Metrics** * **TPOT (Time Per Output Token):** * At low concurrency (**Conc=1**), the average TPOT is very low (**15.0 ms**), indicating extremely fast generation. * As concurrency increases to **64**, the average TPOT rises to **81.0 ms**, which is expected due to increased system load, but remains within an acceptable range for high-throughput serving. * **Speculative Acceptance Rate:** * The acceptance rate remains remarkably stable across all concurrency levels, hovering around **71% - 72%**. This indicates that the MTP model's predictions are consistently accurate regardless of the system load. ##### **2. Throughput and Scalability** * **Completion Throughput (Output Speed):** * There is a clear linear increase in output speed as concurrency rises. * **Conc 1:** 59.46 tok/s $\rightarrow$ **Conc 64:** 345.05 tok/s. * **Total System Throughput:** * The **Overall Total Prompt throughput** scales from **237.85 tok/s** (Conc 1) up to **1380.35 tok/s** (Conc 64). * The **"Last 30s"** peak throughput reaches nearly **3,924 tok/s** at maximum concurrency, showing the model's ability to handle heavy bursts of traffic. ##### **3. Summary Table** | Concurrency | Avg TPOT (ms) | Spec. Accept Rate | Completion Throughput | Total Prompt Throughput | | :--- | :--- | :--- | :--- | :--- | | **1** | 15.0 | 71.5% | 59.46 tok/s | 237.85 tok/s | | **8** | 30.9 | 71.6% | 236.27 tok/s | 945.17 tok/s | | **16** | 43.3 | 71.7% | 299.15 tok/s | 1196.73 tok/s | | **32** | 79.4 | 71.5% | 330.30 tok/s | 1321.32 tok/s | | **64** | 81.0 | 71.6% | 345.05 tok/s | 1380.35 tok/s | #### Additional: Real-world Agent Prompt Validation (950DT TP=1,DP=4) Tested with a ~3.4K-token real tel-sales agent system prompt (300 output tokens, `vllm bench serve`, greedy, ignore-eos, TP=1): | Conc | TPOT p99 (ms) | Output tok/s | Spec accept | |---|---|---|---| | 1 | 9.6 | 102 | 63.4% | | 8 | 10.9 | 732 | 63.6% | | 16 | 12.2 | 1290 | 63.6% | | 32 | 14.3 | 2267 | 63.5% | | 48 | 16.0 | 2981 | 63.6% | | 64 | 17.2 | 3469 | 63.9% | Acceptance stays flat across concurrency; throughput scales linearly to 3469 tok/s at conc=64. Zero errors across the full sweep. Repro (copy-paste-able): ```bash ASCEND_RT_VISIBLE_DEVICES=0 vllm serve <gemma-4-31B-it> \ --served-model-name gemma4-31b-it-mtp \ --tensor-parallel-size 1 \ --speculative-config '{"method":"mtp","model":<gemma-4-31B-it-assistant>,"num_speculative_tokens":3}' \ --max-model-len 32000 --gpu-memory-utilization 0.85 \ --compilation-config '{"cudagraph_mode":"FULL_DECODE_ONLY"}' \ --trust-remote-code --host 0.0.0.0 --port 8831 # then drive the server and scrape /metrics for per-position acceptance ``` - vLLM main: vllm-project/vllm@b2f6858 --------- Signed-off-by: Wu Meng <wumeng@ascend-debug.local> Co-authored-by: Wu Meng <wumeng@ascend-debug.local> Co-authored-by: Claude Code <noreply@anthropic.com> Co-authored-by: XUE TONGYAO <xuetongyao2001@gmail.com> Signed-off-by: like-0517 <ithwlike@126.com>
What this PR does / why we need it?
Adds Gemma4 MTP speculative decoding support on Ascend NPU (A5 / 950PR).
Gemma4 (
gemma-4-31B-it+gemma-4-31B-it-assistant) is a multi-group KVcache model: layers are split into a
slidinggroup (head_dim=256) and afull_attentiongroup (global head_dim=512). The MTP draft reuses thetarget's architecture and shares the target's per-layer KV cache (sliding
for layer 58, global for layer 59), so each draft attention group must read
K/V from its own physical block layout.
Upstream vLLM supports this on CUDA via
Gemma4Proposer. This PR adds theAscend path as a thin wrapper of the upstream proposer — no rewrite of
vLLM internals, no forking of base-class draft-loop methods.
spec_decode/gemma4_proposer.pyAscendGemma4Proposer— multiple inheritance of upstreamGemma4Proposer+AscendSpecDecodeBaseProposer. Overrides only: eager centroids sampling,_sync_kv_sharing_target_to_impl, per-group block table swap (attn_update_stack_num_spec_normwrapper),build_draft_attn_metadatawith FIA SpecDecoding state.spec_decode/llm_base_proposer.pyconstant_draft_positionsguards in_run_merged_draftandattn_update_stack_num_spec_norm— restores upstreamSpecDecodeBaseProposersemantics (vllm/v1/spec_decode/llm_base_proposer.py:635,693,705), no-op for existing proposers (flag defaults False upstream). (2)_build_multi_group_graph_capture_metadatahook (returns None for existing proposers) for per-group ACL graph capture. (3)_proposestep-update loop iteratesattn_group.layer_names, matchingbuild_draft_attn_metadata's existing iteration.ops/rotary_embedding.pykey is None(Gemma4 MTP sliding layers, K/V from target cache), use a throwaway key buffer with the regular rotary implementation (zero-kv-head dummy on the Triton path).worker/model_runner_v1.pyAscendGemma4Proposerinto the drafter type union, per-group block table capture, and kv-cache / cudagraph-key asserts.spec_decode/__init__.pymtp+use_gemma4_mtp()toAscendGemma4Proposer.tests/ut/spec_decode/test_gemma4_proposer.pytests/ut/ops/test_rotary_embedding.pyDoes this PR introduce any user-facing change?
Additive only. Users can now run Gemma4 MTP on Ascend with
--speculative-config '{"method":"mtp",...}'. No change to existing models.How was this patch tested?
Acceptance
Gemma4-31B-it MTP, k=3, greedy (temp=0),
FULL_DECODE_ONLYgraph mode, codingprompt suite (5 prompts × 3 = 15 reqs). Metric:
vllm:spec_decode_num_accepted_tokens_per_pos÷vllm:spec_decode_num_drafts,scraped from
/metrics.Ascend now matches the CUDA reference.
Here is the prompts:
Accuracy
Performance (950PR TP=4,DP=1)
Overall Conclusion:
Tested on 4K input/1K output by Evalscope.
The model demonstrates excellent scalability and stability. As concurrency increases, the overall throughput grows significantly while maintaining a consistent speculative acceptance rate, proving that the MTP (Multi-Token Prediction) mechanism is functioning effectively.
1. Key Latency & Efficiency Metrics
2. Throughput and Scalability
3. Summary Table
Additional: Real-world Agent Prompt Validation (950DT TP=1,DP=4)
Tested with a ~3.4K-token real tel-sales agent system prompt (300 output tokens,
vllm bench serve, greedy, ignore-eos, TP=1):Acceptance stays flat across concurrency; throughput scales linearly to
3469 tok/s at conc=64. Zero errors across the full sweep.
Repro (copy-paste-able):