Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion .github/workflows/scripts/estimated_times.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -14,7 +14,7 @@ estimated_times:
tests/e2e/pull_request/one_card/_310p/test_dense_model_310p.py: 1210
tests/e2e/pull_request/one_card/_310p/test_embedding_310p.py: 480
tests/e2e/pull_request/one_card/_310p/test_scoring_310p.py: 330
tests/e2e/pull_request/one_card/_310p/test_spec_decode_mtp_310p.py: 340
tests/e2e/pull_request/one_card/_310p/test_spec_decode_mtp_310p.py: 600
tests/e2e/pull_request/one_card/_310p/test_vl_model_310p.py: 370
tests/e2e/pull_request/one_card/aclgraph/test_aclgraph_batch_invariant.py: 550
tests/e2e/pull_request/one_card/compile/test_graphex_norm_quant_fusion.py: 170
Expand Down
48 changes: 44 additions & 4 deletions tests/e2e/pull_request/one_card/_310p/test_spec_decode_mtp_310p.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,13 +15,26 @@
# limitations under the License.
# This file is a part of the vllm-ascend project.

from tests.e2e.conftest import VllmRunner
"""310P MTP e2e: MRv1 baseline + MRv2 eager smoke (1-card CI safe)."""

import os
from unittest.mock import patch

from tests.e2e.conftest import VllmRunner, wait_until_npu_memory_free

# CI uses hub id; local verify can override (e.g. /home/weights/Qwen3.5-4B-W8A8).
QWEN35_MTP_MODEL = os.environ.get("QWEN35_MTP_MODEL", "Qwen/Qwen3.5-4B")
QWEN35_MTP_QUANTIZATION = os.environ.get("QWEN35_MTP_QUANTIZATION") # e.g. "ascend"


def _quant_kw():
return {"quantization": QWEN35_MTP_QUANTIZATION} if QWEN35_MTP_QUANTIZATION else {}


def test_qwen3_5_mtp_tp1_eager():
example_prompts = ["Hello, my name is"]
"""MRv1 baseline (no V2 runner env)."""
with VllmRunner(
"Qwen/Qwen3.5-4B",
QWEN35_MTP_MODEL,
tensor_parallel_size=1,
enforce_eager=True,
dtype="float16",
Expand All @@ -31,5 +44,32 @@ def test_qwen3_5_mtp_tp1_eager():
"method": "qwen3_5_mtp",
"num_speculative_tokens": 1,
},
**_quant_kw(),
) as vllm_model:
vllm_model.generate_greedy(["Hello, my name is"], max_tokens=8)


@wait_until_npu_memory_free()
@patch.dict(os.environ, {"VLLM_USE_V2_MODEL_RUNNER": "1"})
def test_qwen3_5_mtp_mrv2_tp1_eager():
"""MRv2 MTP eager smoke.

FULL_DECODE_ONLY is covered by local/nightly runs: 1-card PR CI OOMs during
target+draft ACLGraph capture (Engine core init fails with empty Failed core
proc(s)).
"""
with VllmRunner(
QWEN35_MTP_MODEL,
tensor_parallel_size=1,
enforce_eager=True,
dtype="float16",
max_model_len=2048,
mamba_ssm_cache_dtype="float16",
speculative_config={
"method": "mtp",
"num_speculative_tokens": 1,
},
**_quant_kw(),
) as vllm_model:
vllm_model.generate_greedy(example_prompts, max_tokens=8)
assert vllm_model.model.llm_engine.vllm_config.use_v2_model_runner
vllm_model.generate_greedy(["Hello, my name is"], max_tokens=8)
10 changes: 6 additions & 4 deletions tests/ut/_310p/ops/test_gdn_310.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,7 +18,7 @@
from types import SimpleNamespace

import torch
from vllm.v1.attention.backends.utils import NULL_BLOCK_ID
from vllm.v1.attention.backends.utils import PAD_SLOT_ID

from vllm_ascend._310p.ops.fla.gdn_310 import (
AscendGatedDeltaNetAttention310,
Expand Down Expand Up @@ -85,15 +85,17 @@ def test_builder310_pads_spec_decode_metadata_with_dummy_requests():

builder._pad_spec_decode_metadata(attn_metadata, graph_batch_size=4)

# 310P pads with PAD_SLOT_ID (-1), not NULL_BLOCK_ID (0), so FULL replay
# does not write into mamba block 0. Pad accepted tokens stay 1 (not 0).
assert attn_metadata.spec_state_indices_tensor.tolist() == [
[3, 30],
[4, 40],
[NULL_BLOCK_ID, NULL_BLOCK_ID],
[NULL_BLOCK_ID, NULL_BLOCK_ID],
[PAD_SLOT_ID, PAD_SLOT_ID],
[PAD_SLOT_ID, PAD_SLOT_ID],
]
assert attn_metadata.spec_sequence_masks.tolist() == [True, True, False, False]
assert attn_metadata.spec_query_start_loc.tolist() == [0, 4, 8, 8, 8]
assert attn_metadata.num_accepted_tokens.tolist() == [2, 3, 0, 0]
assert attn_metadata.num_accepted_tokens.tolist() == [2, 3, 1, 1]
spec_meta = attn_metadata.spec_decode_metadata.spec_causal_conv1d
assert spec_meta.query_start_loc.data_ptr() == attn_metadata.spec_query_start_loc.data_ptr()
assert spec_meta.cache_indices.data_ptr() == attn_metadata.spec_state_indices_tensor.data_ptr()
Expand Down
119 changes: 119 additions & 0 deletions tests/ut/_310p/spec_decode/test_mtp_mrv2_310.py
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"])
31 changes: 30 additions & 1 deletion tests/ut/_310p/test_model_runner_v2_310p.py
Original file line number Diff line number Diff line change
Expand Up @@ -340,10 +340,19 @@ def test_config_rejects_non_tp_parallelism(setting: str) -> None:
NPUModelRunner310V2._validate_config(config)


def test_config_accepts_mtp_and_rejects_non_mtp() -> None:
"""310P MRv2 allows method=mtp only."""
NPUModelRunner310V2._validate_config(
_make_vllm_config(speculative_config=SimpleNamespace(method="mtp", num_speculative_tokens=1))
)
with pytest.raises(NotImplementedError, match="only supported via MTP"):
NPUModelRunner310V2._validate_config(_make_vllm_config(speculative_config=SimpleNamespace(method="eagle")))


@pytest.mark.parametrize(
("field", "value", "message"),
[
("speculative_config", object(), "Speculative decoding"),
("speculative_config", object(), "only supported via MTP"),
("kv_transfer_config", object(), "KV cache transfer"),
("lora_config", object(), "LoRA"),
],
Expand All @@ -353,6 +362,26 @@ def test_config_rejects_out_of_scope_features(field, value, message) -> None:
NPUModelRunner310V2._validate_config(_make_vllm_config(**{field: value}))


def test_copy_kv_cache_blocks_flattens_mamba_lists() -> None:
"""Prefix-cache CoW must flatten list[Tensor] mamba layers for upstream copy."""
runner = object.__new__(NPUModelRunner310V2)
runner._attn_kv_copy_params = []
t0 = torch.zeros(4, 2)
t1 = torch.zeros(4, 2)
runner.kv_caches = [[t0, t1], torch.zeros(2)] # hybrid: mamba list + other
runner.kv_cache_config = SimpleNamespace(num_blocks=4)
copies = [SimpleNamespace(src_block_id=0, dst_block_id=1)]

with patch.object(model_runner_module, "copy_kv_cache_blocks_inplace") as mock_copy:
NPUModelRunner310V2._copy_kv_cache_blocks_310p(runner, copies)

mock_copy.assert_called_once()
tensors_arg, num_blocks, copies_arg = mock_copy.call_args[0]
assert tensors_arg == [t0, t1]
assert num_blocks == 4
assert copies_arg is copies


def test_sampler_rejects_random_sampling_parameters() -> None:
sampler = Ascend310PSampler()
sampler.add_request(0, 4, SamplingParams(temperature=0))
Expand Down
8 changes: 8 additions & 0 deletions vllm_ascend/_310p/attention/attention_v1.py
Original file line number Diff line number Diff line change
Expand Up @@ -184,6 +184,12 @@ def forward_paged_attention(
Any: The result of the attention operation.
"""
if attn_metadata.seq_lens.device != query.device:
# Pageable H2D is illegal under NPU GLOBAL ACLGraph capture.
if torch.npu.is_current_stream_capturing():
raise RuntimeError(
"310P paged attention: seq_lens must already be on-device before "
"ACLGraph capture; move it outside torch.npu.graph()."
)
attn_metadata.seq_lens = attn_metadata.seq_lens.to(
device=query.device,
non_blocking=True,
Expand Down Expand Up @@ -265,6 +271,8 @@ def forward_chunked_prefill_310(self, query, attn_metadata, output):
block_table = attn_metadata.block_tables

if attn_metadata.seq_lens.device != query.device:
if torch.npu.is_current_stream_capturing():
raise RuntimeError("310P splitfuse: seq_lens must already be on-device before ACLGraph capture.")
attn_metadata.seq_lens = attn_metadata.seq_lens.to(
device=query.device,
non_blocking=True,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -96,10 +96,12 @@ def _run_recurrent_gated_delta_rule(
if seq_len <= 0:
continue

# Match NPU recurrent_gated_delta_rule_v310: num_accepted_tokens only
# selects the resume state slot (accepted-1). Do NOT truncate the
# current query (MTP verify still has 1+K tokens when prior accept=1).
accepted = None
if num_accepted_tokens is not None:
accepted = int(num_accepted_tokens[seq_idx].item())
seq_len = min(seq_len, accepted)
if seq_len <= 0:
continue
Comment on lines 105 to 106

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

The check if seq_len <= 0: is redundant here because seq_len was already checked at the beginning of the loop (lines 96-97) and is no longer modified within this block (since the truncation seq_len = min(seq_len, accepted) was removed).

Copy link
Copy Markdown
Contributor Author

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_len with accepted, the second seq_len <= 0 check 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.


Expand Down
37 changes: 21 additions & 16 deletions vllm_ascend/_310p/ops/fla/gdn_310.py
Original file line number Diff line number Diff line change
Expand Up @@ -64,16 +64,17 @@ def _flatten_state_indices(
return ssm_state_indices[:total_tokens].to(torch.int32).contiguous()

num_seqs = (cu_seqlens[1:] - cu_seqlens[:-1]).shape[0]
seq_lens = cu_seqlens[1 : num_seqs + 1] - cu_seqlens[:num_seqs]
ssm_state_indices = ssm_state_indices[:num_seqs]
q_per_seq = ssm_state_indices.shape[1]

# Uniform spec-decode ACL graph uses fixed q_len per request; reshape avoids
# NPU masked_select which breaks stream capture (aclnnMaskedSelect / 107027).
if _EXTRA_CTX.capturing or (seq_lens.numel() > 0 and torch.all(seq_lens == seq_lens[0])):
q_per_seq = ssm_state_indices.shape[1]
flat = ssm_state_indices[:, :q_per_seq].reshape(-1)
# NPU masked_select and seq_lens scalar reads which break stream capture.
if _EXTRA_CTX.capturing or total_tokens == num_seqs * q_per_seq:
flat = ssm_state_indices[:num_seqs, :q_per_seq].reshape(-1)
return flat[:total_tokens].to(torch.int32).contiguous()

seq_lens = cu_seqlens[1 : num_seqs + 1] - cu_seqlens[:num_seqs]
ssm_state_indices = ssm_state_indices[:num_seqs]

# Eager mixed batches with variable seq_lens: compact on CPU, copy back async.
ssm_cpu = ssm_state_indices.cpu()
seq_lens_cpu = seq_lens.cpu()
Expand Down Expand Up @@ -119,10 +120,8 @@ def npu_recurrent_gated_delta_rule_310(
total_tokens = v.shape[1]
flat_state_indices = _flatten_state_indices(ssm_state_indices, cu_seqlens, total_tokens)
actual_seq_lengths = (cu_seqlens[1:] - cu_seqlens[:-1]).to(torch.int32).contiguous()
flat_state_indices = torch.clamp_min(
flat_state_indices,
0,
).contiguous()
# Do not clamp PAD_SLOT_ID (-1) to 0 — that would write mamba block 0.
# Callers must slice away padded requests before invoking this helper.
accepted_tokens = None
if num_accepted_tokens is not None:
accepted_tokens = _mask_padded_recurrent_accepted_tokens(
Expand Down Expand Up @@ -262,15 +261,20 @@ def _forward_core(
spec_valid_tokens,
token_dim=0,
)
# Always slice to real spec rows (GPU GDN pattern). FULL-graph pad
# tails must not reach npu_causal_conv1d_310 / recurrent kernels —
# even with PAD_SLOT_ID, cu_seqlens / accepted desync under MTP
# has been observed to poison hybrid state and emit wrong tokens.
num_spec = attn_metadata.num_spec_decodes
mixed_qkv_spec = torch.ops._C_ascend.npu_causal_conv1d_310(
mixed_qkv_spec,
conv_weights,
bias=self.conv1d.bias,
conv_states=conv_state,
query_start_loc=spec_query_start_loc_device,
cache_indices=spec_causal_conv1d_meta.cache_indices,
query_start_loc=spec_query_start_loc_device[: num_spec + 1],
cache_indices=spec_causal_conv1d_meta.cache_indices[:num_spec],
initial_state_mode=None,
num_accepted_tokens=spec_causal_conv1d_meta.num_accepted_tokens,
num_accepted_tokens=spec_causal_conv1d_meta.num_accepted_tokens[:num_spec],
activation_mode=activation_num,
pad_slot_id=PAD_SLOT_ID,
run_mode=1,
Expand Down Expand Up @@ -335,16 +339,17 @@ def _forward_core(

# 2.1: Process the multi-query part
if spec_sequence_masks is not None:
num_spec = attn_metadata.num_spec_decodes
core_attn_out_spec = npu_recurrent_gated_delta_rule_310(
q=query_spec,
k=key_spec,
v=value_spec,
g=g_spec,
beta=beta_spec,
state=ssm_state,
cu_seqlens=spec_query_start_loc[: attn_metadata.num_spec_decodes + 1],
ssm_state_indices=spec_state_indices_tensor[: attn_metadata.num_spec_decodes],
num_accepted_tokens=spec_causal_conv1d_meta.num_accepted_tokens,
cu_seqlens=spec_query_start_loc[: num_spec + 1],
ssm_state_indices=spec_state_indices_tensor[:num_spec],
num_accepted_tokens=spec_causal_conv1d_meta.num_accepted_tokens[:num_spec],
use_qk_l2norm_in_kernel=True,
)
else:
Expand Down
Loading
Loading