-
Notifications
You must be signed in to change notification settings - Fork 2.4k
[Feature][MRV2][310P] MRv2 adapting MTP on the 310P for Qwen3.5 #16043
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Merged
Tflowers-0129
merged 11 commits into
vllm-project:main
from
Thiagor2002:mrv2_310p_mtp_910
Sep 12, 2026
Merged
Changes from all commits
Commits
Show all changes
11 commits
Select commit
Hold shift + click to select a range
52f3976
[Feature][MRV2][310P] MRv2 adapting MTP on the 310P for Qwen3.5
Thiagor2002 4476729
MRv2 now support MTP+eager with aligned precision
Thiagor2002 daa7ae7
[Feature][MRV2][310P] MTP+ACLGraph, MRv2 with correct precision, left…
Thiagor2002 abe6099
[Feature][MRV2][310P] MTP+ACLGraph, adding ACLGraph support, left dra…
Thiagor2002 e0d1516
MTP+ACLGraph, adding target+draft both FULL_Graph, (K=1) aligned with…
Thiagor2002 c233c9a
fix ci cpu issue
Thiagor2002 6a55f44
[BugFix][MRV2][310P] fix Qwen3.5-35B-A3B vllm serve bug, now support …
Thiagor2002 4f1200c
[BugFix][MRV2][310P] K>1 Draft FULL fixed
Thiagor2002 f148f2f
add test cases for MRV2 adapting MTP
Thiagor2002 aee654e
test_spec_decode_mtp_310p, use eager for running CI stably
Thiagor2002 744fffd
solve ci rebase issue
Thiagor2002 File filter
Filter by extension
Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
There are no files selected for viewing
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -0,0 +1,119 @@ | ||
| # SPDX-License-Identifier: Apache-2.0 | ||
| # Copyright (c) 2026 Huawei Technologies Co., Ltd. All Rights Reserved. | ||
| # | ||
| # Lean UTs for 310P MRv2 MTP (rejection offset, capture-safe draft step, RoPE flag). | ||
|
|
||
| from types import SimpleNamespace | ||
| from unittest.mock import patch | ||
|
|
||
| import torch | ||
| from vllm.config.compilation import CUDAGraphMode | ||
|
|
||
| from tests.ut.base import TestBase | ||
| from vllm_ascend._310p.ops.rotary_embedding import AscendRotaryEmbedding310 | ||
| from vllm_ascend._310p.worker.v2.spec_decode.aclgraph import AutoRegressiveAclGraphManager310 | ||
| from vllm_ascend._310p.worker.v2.spec_decode.mtp_speculator import AscendMTPSpeculator310 | ||
| from vllm_ascend._310p.worker.v2.spec_utils import ( | ||
| greedy_rejection_sample_cpu, | ||
| set_draft_step_host, | ||
| update_draft_inputs_cpu, | ||
| ) | ||
|
|
||
|
|
||
| class TestMRv2Mtp310(TestBase): | ||
| def test_greedy_rejection_uses_logit_idx_plus_one(self): | ||
| # draft_sampled[logit_idx+1] must match argmax(logits[logit_idx]) to accept. | ||
| target_logits = torch.tensor( | ||
| [ | ||
| [0.1, 0.9, 0.0], # predicts token 1 | ||
| [0.0, 0.2, 0.8], # bonus → token 2 | ||
| ], | ||
| dtype=torch.float32, | ||
| ) | ||
| draft_sampled = torch.tensor([7, 1], dtype=torch.int32) | ||
| cu_num_logits = torch.tensor([0, 2], dtype=torch.int32) | ||
|
|
||
| sampled, num_sampled = greedy_rejection_sample_cpu( | ||
| target_logits, draft_sampled, cu_num_logits, num_speculative_steps=1 | ||
| ) | ||
|
|
||
| self.assertEqual(num_sampled.tolist(), [2]) | ||
| self.assertEqual(sampled[0, :2].tolist(), [1, 2]) | ||
|
|
||
| def test_update_draft_inputs_uses_host_step_under_capture(self): | ||
| num_reqs = 2 | ||
| draft_tokens = torch.tensor([11, 22], dtype=torch.int32) | ||
| current_draft_step = torch.tensor([99], dtype=torch.int64) # must not .item() under capture | ||
| hidden_states = torch.randn(num_reqs, 4) | ||
| output_draft_tokens = torch.full((num_reqs, 2), -1, dtype=torch.int32) | ||
| next_input_hidden_states = torch.zeros(num_reqs, 4) | ||
| input_buffers = SimpleNamespace( | ||
| input_ids=torch.zeros(num_reqs, dtype=torch.int32), | ||
| positions=torch.tensor([3, 5], dtype=torch.int64), | ||
| seq_lens=torch.tensor([4, 6], dtype=torch.int32), | ||
| ) | ||
| set_draft_step_host(0) | ||
|
|
||
| with patch("torch.npu.is_current_stream_capturing", return_value=True): | ||
| update_draft_inputs_cpu( | ||
| draft_tokens=draft_tokens, | ||
| current_draft_step=current_draft_step, | ||
| hidden_states=hidden_states, | ||
| output_draft_tokens=output_draft_tokens, | ||
| next_input_hidden_states=next_input_hidden_states, | ||
| input_buffers=input_buffers, | ||
| num_reqs=num_reqs, | ||
| max_model_len=128, | ||
| num_speculative_steps=2, | ||
| advance_draft_positions=True, | ||
| ) | ||
|
|
||
| self.assertEqual(output_draft_tokens[:, 0].tolist(), [11, 22]) | ||
| self.assertEqual(input_buffers.positions.tolist(), [4, 6]) | ||
|
|
||
| def test_run_model_sets_rope_flag(self): | ||
| flag_states: list[bool] = [] | ||
|
|
||
| def mock_parent_run(self, *args, **kwargs): | ||
| del self, args, kwargs | ||
| flag_states.append(AscendRotaryEmbedding310._is_drafting_update_enabled) | ||
| return torch.zeros(1), torch.zeros(1) | ||
|
|
||
| speculator = object.__new__(AscendMTPSpeculator310) | ||
| with patch( | ||
| "vllm_ascend.worker.v2.spec_decode.autoregressive.speculator.AscendAutoRegressiveSpeculator._run_model", | ||
| mock_parent_run, | ||
| ): | ||
| AscendMTPSpeculator310._run_model( | ||
| speculator, | ||
| num_tokens=1, | ||
| attn_metadata=None, | ||
| slot_mappings=None, | ||
| num_tokens_across_dp=None, | ||
| cudagraph_runtime_mode=CUDAGraphMode.NONE, | ||
| ) | ||
|
|
||
| self.assertEqual(flag_states, [True]) | ||
| self.assertFalse(AscendRotaryEmbedding310._is_drafting_update_enabled) | ||
|
|
||
| def test_decode_capture_routes_to_per_step(self): | ||
| manager = object.__new__(AutoRegressiveAclGraphManager310) | ||
| manager.is_draft_model_prefill = False | ||
| called = {"per_step": False} | ||
|
|
||
| def fake_per_step(*args, **kwargs): | ||
| del args, kwargs | ||
| called["per_step"] = True | ||
|
|
||
| with patch.object(AutoRegressiveAclGraphManager310, "_capture_decode_per_step", fake_per_step): | ||
| AutoRegressiveAclGraphManager310.capture( | ||
| manager, | ||
| forward_fn=lambda: None, | ||
| model_state=object(), | ||
| input_buffers=object(), | ||
| block_tables=object(), | ||
| attn_groups=[], | ||
| kv_cache_config=object(), | ||
| ) | ||
|
|
||
| self.assertTrue(called["per_step"]) |
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Oops, something went wrong.
Oops, something went wrong.
Add this suggestion to a batch that can be applied as a single commit.
This suggestion is invalid because no changes were made to the code.
Suggestions cannot be applied while the pull request is closed.
Suggestions cannot be applied while viewing a subset of changes.
Only one suggestion per line can be applied in a batch.
Add this suggestion to a batch that can be applied as a single commit.
Applying suggestions on deleted lines is not supported.
You must change the existing code in this line in order to create a valid suggestion.
Outdated suggestions cannot be applied.
This suggestion has been applied or marked resolved.
Suggestions cannot be applied from pending reviews.
Suggestions cannot be applied on multi-line comments.
Suggestions cannot be applied while the pull request is queued to merge.
Suggestion cannot be applied right now. Please check back later.
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
The check
if seq_len <= 0:is redundant here becauseseq_lenwas already checked at the beginning of the loop (lines 96-97) and is no longer modified within this block (since the truncationseq_len = min(seq_len, accepted)was removed).There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
Fair observation: after we stopped truncating
seq_lenwithaccepted, the secondseq_len <= 0check is unreachable. We’ll leave it as-is for this PR — it is harmless and unrelated to the MTP / ACLGraph behavior under review. it is not necessary here.