Skip to content
Closed
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
23 changes: 23 additions & 0 deletions tests/ut/patch/platform/test_patch_speculative_config_dspark.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,23 @@
from transformers import Qwen3Config
from vllm.config.speculative import SpeculativeConfig

import vllm_ascend.patch.platform.patch_speculative_config # noqa: F401


def test_legacy_qwen3_dspark_config_uses_qwen3_loader():
config = Qwen3Config(
architectures=["DSparkDraftModel"],
block_size=7,
dflash_config={
"mask_token_id": 163824,
"target_layer_ids": [7, 23, 51, 67, 83],
},
)

normalized = SpeculativeConfig.hf_config_override(config)

assert normalized is config
assert normalized.architectures == ["Qwen3DSparkModel"]
assert normalized.mask_token_id == 163824
assert normalized.target_layer_ids == [7, 23, 51, 67, 83]
assert normalized.block_size == 7
239 changes: 134 additions & 105 deletions tests/ut/spec_decode/test_dspark_proposer.py

Large diffs are not rendered by default.

17 changes: 12 additions & 5 deletions tests/ut/spec_decode/test_llm_base_proposer.py
Original file line number Diff line number Diff line change
Expand Up @@ -76,16 +76,22 @@ def test_pixtral_uses_vision_config_image_token_id(self):

assert image_token_index == 789

def test_kimi_uses_media_placeholder_token_id(self):
@pytest.mark.parametrize(
"model_name",
[
"KimiK25ForConditionalGeneration",
"KimiK3ForConditionalGeneration",
"AscendKimiK3ForConditionalGeneration",
],
)
def test_kimi_uses_media_placeholder_token_id(self, model_name: str):
config = SimpleNamespace(
image_token_id=123,
image_token_index=456,
media_placeholder_token_id=789,
)

image_token_index = AscendSpecDecodeBaseProposer._get_multimodal_image_token_index(
"KimiK25ForConditionalGeneration", config
)
image_token_index = AscendSpecDecodeBaseProposer._get_multimodal_image_token_index(model_name, config)

assert image_token_index == 789

Expand All @@ -103,7 +109,8 @@ def test_load_model_reads_validated_draft_window_size():
proposer = AscendSpecDecodeBaseProposer.__new__(AscendSpecDecodeBaseProposer)
proposer.vllm_config = SimpleNamespace(additional_config={"draft_window_size": 64})
proposer.maybe_eager_context = nullcontext()
proposer._get_model = MagicMock(return_value=MagicMock())
draft_model = MagicMock()
proposer._get_model = MagicMock(return_value=draft_model)
proposer.method = "eagle3"
proposer.num_speculative_tokens = 4
proposer.runner = SimpleNamespace(max_num_reqs=8)
Expand Down
28 changes: 28 additions & 0 deletions tests/ut/spec_decode/test_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,10 +8,13 @@

from __future__ import annotations

from types import SimpleNamespace

import numpy as np
import torch

from vllm_ascend.spec_decode.utils import (
SlidingWindowAdapter,
correct_optimistic_seq_lens_cpu,
update_num_computed_tokens_for_batch_change,
)
Expand Down Expand Up @@ -184,3 +187,28 @@ def test_cpu_and_gpu_corrections_agree():
gpu_seq_lens = num_computed_gpu.numpy() + num_scheduled_step_n

np.testing.assert_array_equal(optimistic, gpu_seq_lens)


def test_sliding_window_updates_host_length_mirrors():
adapter = SlidingWindowAdapter(
window_size=128,
block_size=64,
max_num_reqs=1,
future_offset=0,
device=torch.device("cpu"),
)
metadata = SimpleNamespace(
block_table_tensor=torch.arange(10, dtype=torch.int32).view(1, 10),
seq_lens=torch.tensor([300], dtype=torch.int32),
seq_lens_cpu=torch.tensor([300], dtype=torch.int32),
_seq_lens_cpu=torch.tensor([300], dtype=torch.int32),
seq_lens_cpu_upper_bound=torch.tensor([300], dtype=torch.int32),
)

adapter.apply(metadata)

expected = torch.tensor([172], dtype=torch.int32)
torch.testing.assert_close(metadata.seq_lens, expected)
torch.testing.assert_close(metadata.seq_lens_cpu, expected)
torch.testing.assert_close(metadata._seq_lens_cpu, expected)
torch.testing.assert_close(metadata.seq_lens_cpu_upper_bound, expected)
18 changes: 18 additions & 0 deletions vllm_ascend/patch/platform/patch_speculative_config.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,23 @@
from transformers import PretrainedConfig
from vllm.config.speculative import SpeculativeConfig

_orig_post_init = SpeculativeConfig.__post_init__
_orig_hf_config_override = SpeculativeConfig.hf_config_override


def _normalize_legacy_qwen3_dspark_config(hf_config: PretrainedConfig) -> PretrainedConfig:
hf_config = _orig_hf_config_override(hf_config)
architectures = hf_config.architectures or ()
if hf_config.model_type == "qwen3" and "DSparkDraftModel" in architectures:
dflash_config = hf_config.dflash_config
hf_config.update(
Comment thread
maoxx241 marked this conversation as resolved.
{
"architectures": ["Qwen3DSparkModel"],
"mask_token_id": dflash_config["mask_token_id"],
"target_layer_ids": dflash_config["target_layer_ids"],
}
)
return hf_config


def _dspark_post_init(self):
Expand All @@ -16,4 +33,5 @@ def _dspark_post_init(self):
draft_hf_config.ptd_token_id = getattr(draft_hf_config, "mask_token_id", None) # type: ignore


SpeculativeConfig.hf_config_override = staticmethod(_normalize_legacy_qwen3_dspark_config)
SpeculativeConfig.__post_init__ = _dspark_post_init
91 changes: 43 additions & 48 deletions vllm_ascend/spec_decode/dspark_proposer.py
Original file line number Diff line number Diff line change
Expand Up @@ -43,8 +43,8 @@ def __init__(
blk = 1 + self.num_speculative_tokens
self._dspark_draft_buffer = torch.zeros((self.max_batch_size, blk), dtype=torch.int64, device=device)
self._dspark_seed_buffer = torch.zeros(self.max_batch_size, dtype=torch.int64, device=device)
# DSpark is not supported in vllm v1, so related property needs to be reset here.
del self.hidden_size, self.hidden_states, self._dflash_hidden_states # type: ignore[has-type]
# Replace the target-sized DFlash buffers with the draft model's hidden
# size. Assignment releases the old tensors without an explicit del.
self.hidden_size = vllm_config.speculative_config.draft_model_config.get_hidden_size()
self.hidden_states = torch.zeros(
(self.max_num_tokens, self.hidden_size),
Expand Down Expand Up @@ -87,27 +87,20 @@ def __init__(
device=device,
)

# TODO simplify these comments
# block_table / slot_mapping bookkeeping (10 dicts below). v1 self-
# manages per kv_cache_group_id / per layer because it lacks v2's
# BlockTables scaffold; v2 injects a single self.block_tables
# (BlockTables, with .slot_mappings) + build_slot_mappings_by_layer,
# so the speculator holds none of these. P2 refactor target (move to
# runner).

# per-gid block_table from runner (just read)
# The v1 runner owns block tables and slot mappings. Keep per-group
# references here because K3 draft layers can span multiple cache
# groups with different logical block sizes.
self._per_group_block_tables: dict[int, torch.Tensor] = {}
# per-gid slot_mapping from runner (just read)
self._per_group_slot_mappings: dict[int, torch.Tensor] = {}
# Per-gid logical block size used to expand slot mappings. The KV
# manager's physical page can be larger when hybrid cache groups share
# one allocation, so kv_cache_spec.block_size is not interchangeable
# with the attention kernel's block size.
self._per_group_kernel_block_sizes: dict[int, int] = {}

# per-gid block_table (use in proposer)
self._per_group_block_table_buffers: dict[int, torch.Tensor] = {}
# per-gid query slot_mapping buffer
self._per_group_query_slot_mapping_buffers: dict[int, torch.Tensor] = {}
# per-gid context slot_mapping buffer
self._per_group_context_slot_mapping_buffers: dict[int, torch.Tensor] = {}

# per-layer context slot mappings as a flat list
self._context_slot_mapping_buffers: list[torch.Tensor | None] | None = None

def _compute_confidence(
Expand All @@ -129,24 +122,22 @@ def _compute_confidence(
confidence.copy_(conf_raw.reshape(num_reqs, self.num_speculative_tokens))
return confidence

def initialize_attn_backend(self, kv_cache_config, kernel_block_sizes=None) -> None:
def initialize_attn_backend(
self,
kv_cache_config,
kernel_block_sizes: list[int] | None = None,
) -> None:
# Find draft layers (attention layers added by draft model)
all_attn_layers = get_layers_from_vllm_config(
self.vllm_config,
AttentionLayerBase, # type: ignore[type-abstract]
)

attention_groups_list: list[dict[tuple[str, str], AttentionGroup]] = []
# the draft layers have multiple kv_cache_groups
if not hasattr(self.model, "get_draft_kv_cache_layer_names"):
raise RuntimeError(
"DSpark standard-cache path requires the draft model to expose get_draft_kv_cache_layer_names"
)

self._draft_attn_layer_names = set(self.model.get_draft_kv_cache_layer_names())
self.attn_layer_names = list(sorted(self._draft_attn_layer_names))
self._per_group_kernel_block_sizes = {}
self.draft_attn_groups: list[AttentionGroup] = []

# there are many kv groups other than one
for kv_cache_gid, kv_cache_group_spec in enumerate(kv_cache_config.kv_cache_groups):
draft_layer_names_in_group = set(kv_cache_group_spec.layer_names) & self._draft_attn_layer_names
if not draft_layer_names_in_group:
Expand All @@ -162,33 +153,31 @@ def initialize_attn_backend(self, kv_cache_config, kernel_block_sizes=None) -> N
key = (attn_backend.full_cls_name(), layer_kv_cache_spec)

if key not in attention_groups:
kernel_block_size = int(
kernel_block_sizes[kv_cache_gid]
if kernel_block_sizes is not None and kv_cache_gid < len(kernel_block_sizes)
else layer_kv_cache_spec.block_size
)
attn_group = AttentionGroup(
attn_backend,
[layer_name],
layer_kv_cache_spec,
kv_cache_gid,
)
attn_group.create_metadata_builders(self.vllm_config, self.device)
attn_group.create_metadata_builders(
self.vllm_config,
self.device,
kernel_block_size=kernel_block_size,
)
self._per_group_kernel_block_sizes[kv_cache_gid] = kernel_block_size
attention_groups[key] = attn_group
else:
attention_groups[key].layer_names.append(layer_name)

attention_groups_list.append(attention_groups)

self.draft_attn_groups = [
attention_group
for attention_groups in attention_groups_list
for attention_group in attention_groups.values()
]
self.kv_cache_gid = 0
if not self.draft_attn_groups:
raise RuntimeError(
"DSpark standard-cache path requires registered draft attention "
f"groups. Missing layers: {self.attn_layer_names}"
)
self.draft_attn_groups.extend(attention_groups.values())

self.kv_cache_gid = self.draft_attn_groups[0].kv_cache_group_id
self.kernel_block_size = int(self.draft_attn_groups[0].kv_cache_spec.block_size)
self.kernel_block_size = self._per_group_kernel_block_sizes[self.kv_cache_gid]

name_to_gid = {
ln: gid
Expand Down Expand Up @@ -257,13 +246,10 @@ def set_inputs_first_pass(
# Query block: reuse the DFlash inputs kernel logic (host-side ref)
# per kv-cache-group to fill positions / input_ids / query slot_mapping
# / token_indices.
draft_attn_groups = getattr(self, "draft_attn_groups", [])
for attn_group in draft_attn_groups:
for attn_group in self.draft_attn_groups:
gid = attn_group.kv_cache_group_id
gid_block_table = self._per_group_block_table_buffers.get(gid)
if gid_block_table is None:
continue
kv_block_size = int(attn_group.kv_cache_spec.block_size)
gid_block_table = self._per_group_block_table_buffers[gid]
kernel_block_size = self._per_group_kernel_block_sizes[gid]
copy_and_expand_dflash_and_dspark_inputs_kernel[
(_compute_num_programs(self._dflash_num_context, num_query_total),)
](
Expand All @@ -287,7 +273,7 @@ def set_inputs_first_pass(
num_rejected_tokens_ptr=num_rejected_tokens_gpu,
# Scalars
parallel_drafting_token_id=self.parallel_drafting_token_id,
block_size=kv_block_size,
block_size=kernel_block_size,
num_query_per_req=self.num_query_per_req,
num_speculative_tokens=self.num_speculative_tokens,
total_input_tokens=self._dflash_num_context,
Expand All @@ -306,6 +292,15 @@ def set_inputs_first_pass(

cad.query_start_loc = self.arange_dflash[: batch_size + 1] * self.num_query_per_req
cad.seq_lens = effective_seq_lens + self.num_query_per_req
# The model runner has already corrected this canonical host mirror
# with the accepted-token count. Extend it on CPU alongside the device
# lengths, without another reject D2H copy or attention-side wait.
if cad._seq_lens_cpu is not None:
Comment thread
maoxx241 marked this conversation as resolved.
draft_seq_lens_cpu = cad._seq_lens_cpu.clone()
draft_seq_lens_cpu[:batch_size].add_(self.num_query_per_req)
cad._seq_lens_cpu = draft_seq_lens_cpu
if getattr(cad, "seq_lens_cpu", None) is not None:
Comment thread
maoxx241 marked this conversation as resolved.
cad.seq_lens_cpu = draft_seq_lens_cpu
cad.query_start_loc_cpu = (
torch.from_numpy(self.token_arange_np[: batch_size + 1]).clone() * self.num_query_per_req
).to(torch.int32)
Expand Down
Loading
Loading