Skip to content
Draft
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
3 changes: 3 additions & 0 deletions .github/workflows/scripts/test_config.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -467,13 +467,16 @@
source_file_dependencies:
- vllm_ascend/spec_decode/llm_base_proposer.py
- vllm_ascend/spec_decode/eagle_proposer.py
- vllm_ascend/spec_decode/gemma4_proposer.py
- vllm_ascend/patch/platform/patch_speculative_config.py
- vllm_ascend/models/llama_eagle3_vwn.py
tests:
- tests/e2e/pull_request/one_card/spec_decode/test_eagle.py
- tests/e2e/pull_request/one_card/spec_decode/test_mtp_eagle_correctness.py
- tests/e2e/pull_request/two_card/spec_decode/test_spec_decode.py
- tests/e2e/pull_request/four_card/spec_decode/test_mtp_qwen3_next.py
- tests/ut/spec_decode/test_speculators_vwn_eagle3.py
- tests/ut/spec_decode/test_gemma4_proposer.py

- name: spec_decode_dspark
optional: false
Expand Down
18 changes: 18 additions & 0 deletions tests/ut/ops/test_rotary_embedding.py
Original file line number Diff line number Diff line change
Expand Up @@ -160,6 +160,24 @@ def test_basic_call_delegates_to_npu_op(self, mock_get_forward_context, mock_npu
)
assert result is expected_output

@patch("vllm_ascend.ops.rotary_embedding.is_forward_context_available", return_value=False)
@patch("vllm_ascend.ops.rotary_embedding.rope_forward_oot")
def test_q_only_uses_throwaway_key(self, mock_rope, _mock_is_ctx, make_embedding):
emb = make_embedding()
positions, query, _ = _make_tensors()
expected_query = torch.randn_like(query)
mock_rope.return_value = expected_query, torch.empty_like(query)

with patch("vllm_ascend.ops.rotary_embedding.HAS_TRITON", False):
result = emb.forward_oot(positions, query, None)

assert result[0] is expected_query
assert result[1] is None
dummy_key = mock_rope.call_args.args[2]
assert dummy_key.shape == (query.shape[0], HEAD_SIZE)
assert dummy_key.dtype == query.dtype
assert dummy_key.device == query.device

@patch("torch.ops.vllm.npu_rotary_embedding")
@patch("vllm_ascend.ascend_forward_context.get_forward_context")
def test_neox_style_override_true(self, mock_get_forward_context, mock_npu_op, make_embedding):
Expand Down
196 changes: 196 additions & 0 deletions tests/ut/spec_decode/test_gemma4_proposer.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,196 @@
# SPDX-License-Identifier: Apache-2.0

from types import SimpleNamespace
from unittest.mock import MagicMock, patch

import torch

import vllm_ascend.spec_decode as spec_decode
from vllm_ascend.attention.attention_v1 import AscendAttentionState
from vllm_ascend.patch.platform import patch_speculative_config
from vllm_ascend.spec_decode.gemma4_proposer import AscendGemma4Proposer
from vllm_ascend.spec_decode.llm_base_proposer import AscendSpecDecodeBaseProposer


def test_routes_gemma4_mtp_to_ascend_proposer():
speculative_config = MagicMock()
speculative_config.use_gemma4_mtp.return_value = True
vllm_config = SimpleNamespace(speculative_config=speculative_config)
expected = object()

with patch.object(
spec_decode,
"AscendGemma4Proposer",
return_value=expected,
) as proposer_cls:
result = spec_decode.get_spec_decode_method(
"mtp",
vllm_config,
device="npu",
runner=object(),
)

assert result is expected
proposer_cls.assert_called_once()


def test_gemma_config_override_delegates_to_vllm(monkeypatch):
hf_config = SimpleNamespace(
architectures=["Gemma4ForConditionalGeneration"],
model_type="gemma4_assistant",
)
expected = object()
original_override = MagicMock(return_value=expected)
monkeypatch.setattr(
patch_speculative_config,
"_orig_hf_config_override",
original_override,
)

result = patch_speculative_config.hf_config_override(hf_config)

assert result is expected
original_override.assert_called_once_with(hf_config)


def test_sync_kv_sharing_target_to_impl():
proposer = AscendGemma4Proposer.__new__(AscendGemma4Proposer)
proposer.vllm_config = MagicMock()
proposer._draft_attn_layer_names = {"draft.attn"}

impl = SimpleNamespace(kv_sharing_target_layer_name=None)
attn = SimpleNamespace(
impl=impl,
kv_sharing_target_layer_name="target.attn",
)
with patch(
"vllm_ascend.spec_decode.gemma4_proposer.get_layers_from_vllm_config",
return_value={"draft.attn": attn},
):
proposer._sync_kv_sharing_target_to_impl()

assert impl.kv_sharing_target_layer_name == "target.attn"


def test_keeps_draft_lm_head():
proposer = AscendGemma4Proposer.__new__(AscendGemma4Proposer)
draft_lm_head = object()
proposer.model = SimpleNamespace(lm_head=draft_lm_head)
proposer.method = "mtp"
proposer.use_cuda_graph = False
proposer.vllm_config = SimpleNamespace(
model_config=SimpleNamespace(is_deepseek_mla=False),
compilation_config=SimpleNamespace(
cudagraph_mode=SimpleNamespace(
has_full_cudagraphs=lambda: False,
)
),
)

proposer._maybe_share_lm_head(SimpleNamespace(lm_head=object()))

assert proposer.model.lm_head is draft_lm_head


def test_build_draft_attn_metadata_uses_per_group_block_tables():
proposer = AscendGemma4Proposer.__new__(AscendGemma4Proposer)
block_tables = {
0: torch.arange(12).view(3, 4),
1: torch.arange(12, 24).view(3, 4),
}
proposer._per_group_block_tables = block_tables
proposer.runner = SimpleNamespace(get_model=MagicMock(return_value=object()))
metadata = [
SimpleNamespace(attn_state=None, causal=True),
SimpleNamespace(attn_state=None, causal=False, attn_mask=object()),
]
builders = [MagicMock(), MagicMock()]
for builder, group_metadata in zip(builders, metadata):
builder.build.return_value = group_metadata
proposer.draft_attn_groups = [
SimpleNamespace(
kv_cache_group_id=gid,
layer_names=[f"draft.attn.{gid}"],
get_metadata_builder=MagicMock(return_value=builders[gid]),
)
for gid in range(2)
]
common_metadata = SimpleNamespace(
num_reqs=2,
block_table_tensor=torch.zeros(2, 4),
)

multi_steps, first_metadata = proposer.build_draft_attn_metadata(
common_metadata,
num_input_tokens=2,
num_actual_tokens=2,
)

assert first_metadata is metadata[0]
assert multi_steps == [
{
"draft.attn.0": metadata[0],
"draft.attn.1": metadata[1],
}
]
for gid, builder in enumerate(builders):
group_common_metadata = builder.build.call_args.args[1]
assert group_common_metadata is not common_metadata
assert torch.equal(
group_common_metadata.block_table_tensor,
block_tables[gid][:2],
)
assert metadata[gid].attn_state == AscendAttentionState.SpecDecoding
assert metadata[1].attn_mask is None

graph_metadata = [object(), object()]
for builder, graph_item in zip(builders, graph_metadata):
builder.build_for_graph_capture.return_value = graph_item
graph_result = proposer._build_multi_group_graph_capture_metadata(
common_metadata,
draft_index=0,
)

assert graph_result == {
"draft.attn.0": graph_metadata[0],
"draft.attn.1": graph_metadata[1],
}
for gid, builder in enumerate(builders):
call_args = builder.build_for_graph_capture.call_args.args
assert torch.equal(call_args[0].block_table_tensor, block_tables[gid][:2])
assert call_args[1] == AscendAttentionState.SpecDecoding


def test_attn_update_uses_only_active_group_block_table():
proposer = AscendGemma4Proposer.__new__(AscendGemma4Proposer)
proposer._per_group_block_tables = {1: torch.arange(12).view(3, 4)}
common_metadata = SimpleNamespace(
num_reqs=2,
block_table_tensor=torch.zeros(2, 4),
)
attn_group = SimpleNamespace(kv_cache_group_id=1)
expected = object()

with patch.object(
AscendSpecDecodeBaseProposer,
"attn_update_stack_num_spec_norm",
return_value=expected,
) as base_update:
result = proposer.attn_update_stack_num_spec_norm(
1,
common_metadata,
2,
2,
torch.tensor([3, 4]),
"none",
attn_group=attn_group,
)

assert result is expected
group_common_metadata = base_update.call_args.args[1]
assert group_common_metadata is not common_metadata
assert torch.equal(
group_common_metadata.block_table_tensor,
proposer._per_group_block_tables[1][:2],
)
assert base_update.call_args.kwargs["attn_group"] is attn_group
24 changes: 23 additions & 1 deletion vllm_ascend/ops/rotary_embedding.py
Original file line number Diff line number Diff line change
Expand Up @@ -235,7 +235,7 @@ def forward_oot(
self,
positions: torch.Tensor,
query: torch.Tensor,
key: torch.Tensor,
key: torch.Tensor | None,
offsets: torch.Tensor | None = None,
is_neox_style_override: bool | None = None,
):
Expand All @@ -246,6 +246,28 @@ def forward_oot(
flash_comm_v1_enabled = _EXTRA_CTX.flash_comm_v1_enabled if is_forward_context_available() else False
if is_draft_model and self.use_mtp and flash_comm_v1_enabled:
positions = torch.ops.vllm.maybe_all_gather_and_maybe_unpad(positions.contiguous(), True)
if key is None:
# Gemma4 MTP reads K/V from the target cache. Reuse the regular
# rotary implementation with a throwaway key buffer.
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) if query.ndim == 3 else (query.shape[0], self.head_size),
dtype=query.dtype,
device=query.device,
)
)
Comment on lines +252 to +260

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)
                )
            )

query, _ = rope_forward_oot(
positions,
query,
dummy_key,
self.cos_sin_cache,
self.head_size,
self.rotary_dim,
is_neox_style,
)
return query, None
return torch.ops.vllm.npu_rotary_embedding(
positions, query, key, self.cos_sin_cache, self.head_size, self.rotary_dim, is_neox_style
)
Expand Down
3 changes: 3 additions & 0 deletions vllm_ascend/patch/platform/patch_speculative_config.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@
from vllm.utils.import_utils import LazyLoader

_orig_post_init = SpeculativeConfig.__post_init__
_orig_hf_config_override = SpeculativeConfig.hf_config_override

if TYPE_CHECKING:
import vllm.model_executor.layers.quantization as me_quant
Expand All @@ -16,6 +17,8 @@

def hf_config_override(hf_config: PretrainedConfig) -> PretrainedConfig:
initial_architecture = hf_config.architectures[0]
if hf_config.model_type in ("gemma4_assistant", "gemma4_unified_assistant"):
return _orig_hf_config_override(hf_config)
if hf_config.model_type in ("deepseek_v3", "deepseek_v32", "deepseek_v4", "glm_moe_dsa"):
target_model_type = hf_config.model_type
hf_config.model_type = "deepseek_mtp"
Expand Down
3 changes: 3 additions & 0 deletions vllm_ascend/spec_decode/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -25,6 +25,7 @@
from vllm_ascend.spec_decode.extract_hidden_states_proposer import (
AscendExtractHiddenStatesProposer,
)
from vllm_ascend.spec_decode.gemma4_proposer import AscendGemma4Proposer
from vllm_ascend.spec_decode.medusa_proposer import AscendMedusaProposer
from vllm_ascend.spec_decode.ngram_proposer import AscendNgramProposer
from vllm_ascend.spec_decode.ngram_proposer_npu import AscendNgramProposerNPU
Expand All @@ -43,6 +44,8 @@ def get_spec_decode_method(method, vllm_config, device, runner):
return AscendMedusaProposer(vllm_config, device)
elif method == "dspark":
return AscendDSparkProposer(vllm_config, device, runner)
elif method == "mtp" and vllm_config.speculative_config.use_gemma4_mtp():
return AscendGemma4Proposer(vllm_config, device, runner)
elif method in ("eagle", "eagle3", "mtp"):
speculative_config = vllm_config.speculative_config
if speculative_config is not None and speculative_config.use_step3p5_mtp():
Expand Down
Loading
Loading