init v0.23.0

Signed-off-by: Sun Ruoxi <sunruoxi@4paradigm.com>
This commit is contained in:
2026-08-27 15:11:51 +08:00
parent b582a8e7d1
commit 7f8a1b1f7a
2849 changed files with 712887 additions and 22001 deletions

View File

@@ -16,18 +16,40 @@
# This file is a part of the vllm-ascend project.
# Adapted from vllm-project/vllm/vllm/worker/gpu_model_runner.py
#
from vllm_ascend.spec_decode.eagle_proposer import EagleProposer
from vllm_ascend.spec_decode.mtp_proposer import MtpProposer
from vllm_ascend.spec_decode.ngram_proposer import NgramProposer
from vllm_ascend.spec_decode.dflash_proposer import AscendDflashProposer
from vllm_ascend.spec_decode.draft_proposer import AscendDraftModelProposer
from vllm_ascend.spec_decode.eagle_proposer import AscendEagleProposer
from vllm_ascend.spec_decode.extract_hidden_states_proposer import (
AscendExtractHiddenStatesProposer,
)
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
from vllm_ascend.spec_decode.step3p5 import AscendStep3p5MTPProposer
from vllm_ascend.spec_decode.suffix_proposer import AscendSuffixDecodingProposer
def get_spec_decode_method(method, vllm_config, device, runner):
if method == "ngram":
return NgramProposer(vllm_config, device, runner)
elif method in ["eagle", "eagle3"]:
return EagleProposer(vllm_config, device, runner)
elif method == 'deepseek_mtp':
return MtpProposer(vllm_config, device, runner)
return AscendNgramProposer(vllm_config, runner)
elif method == "ngram_gpu":
return AscendNgramProposerNPU(vllm_config, device, runner)
elif method == "suffix":
return AscendSuffixDecodingProposer(vllm_config, runner)
elif method == "medusa":
return AscendMedusaProposer(vllm_config, device)
elif method in ("eagle", "eagle3", "mtp"):
speculative_config = vllm_config.speculative_config
if speculative_config is not None and speculative_config.use_step3p5_mtp():
return AscendStep3p5MTPProposer(vllm_config, device, runner)
return AscendEagleProposer(vllm_config, device, runner)
elif method == "dflash":
return AscendDflashProposer(vllm_config, device, runner)
elif method == "draft_model":
return AscendDraftModelProposer(vllm_config, device, runner)
elif method == "extract_hidden_states":
return AscendExtractHiddenStatesProposer(vllm_config, device, runner)
else:
raise ValueError("Unknown speculative decoding method: "
f"{method}")
raise ValueError(f"Unknown speculative decoding method: {method}")

View File

@@ -0,0 +1,270 @@
from typing import Any
import torch
from vllm.config import CUDAGraphMode, VllmConfig
from vllm.forward_context import get_forward_context
from vllm.v1.attention.backends.utils import CommonAttentionMetadata
from vllm_ascend.ascend_forward_context import _EXTRA_CTX, set_ascend_forward_context
from vllm_ascend.attention.attention_v1 import AscendAttentionState
from vllm_ascend.attention.utils import AscendCommonAttentionMetadata
from vllm_ascend.ops.triton.spec_decode.utils import copy_and_expand_dflash_inputs_kernel_single_grid
from vllm_ascend.spec_decode.eagle_proposer import AscendEagleProposer
class AscendDflashProposer(AscendEagleProposer):
def __init__(
self,
vllm_config: VllmConfig,
device: torch.device,
runner=None,
):
super().__init__(
vllm_config,
device,
runner=runner,
)
self.max_query_tokens = self.max_batch_size * (1 + self.num_speculative_tokens)
self.max_positions = self.max_num_tokens + self.max_query_tokens
self._context_slot_mapping_buffer = torch.zeros(
self.max_num_tokens,
dtype=torch.int32,
device=device,
)
self._slot_mapping_buffer = torch.zeros(
self.max_query_tokens,
dtype=torch.int32,
device=device,
)
self._context_positions_buffer = torch.zeros(
self.max_num_tokens,
dtype=torch.int32,
device=device,
)
self.positions = torch.zeros(
self.max_query_tokens,
dtype=torch.int32,
device=device,
)
self.arange_dflash = torch.arange(self.max_positions + 1, device=device, dtype=torch.int32)
self._dflash_hidden_states = torch.zeros(
(self.max_num_tokens, self.hidden_size), dtype=self.dtype, device=self.device
)
self.parallel_drafting_hidden_state_tensor = None
def set_inputs_first_pass(
self,
target_token_ids: torch.Tensor,
next_token_ids: torch.Tensor,
target_positions: torch.Tensor,
target_hidden_states: torch.Tensor,
token_indices_to_sample: torch.Tensor | None,
cad: CommonAttentionMetadata,
num_rejected_tokens_gpu: torch.Tensor | None,
req_scheduled_tokens=None,
long_seq_metadata=None,
num_prefill_reqs=0,
num_decode_reqs=0,
) -> tuple[int, torch.Tensor, CommonAttentionMetadata, tuple[Any, Any] | None]:
# DFlash cross-attention: context K/V from target hidden states,
# Q from query embeddings (bonus + mask tokens).
batch_size = cad.num_reqs
num_context = target_token_ids.shape[0]
num_query_per_req = 1 + self.num_speculative_tokens
num_query_total = batch_size * num_query_per_req
self._dflash_num_context = num_context
self._dflash_hidden_states[:num_context] = target_hidden_states
token_indices_to_sample = torch.empty(
batch_size * self.num_speculative_tokens,
dtype=torch.int32,
device=self.device,
)
has_num_rejected = num_rejected_tokens_gpu is not None
copy_and_expand_dflash_inputs_kernel_single_grid[1,](
# Inputs
next_token_ids_ptr=next_token_ids,
target_positions_ptr=target_positions,
context_slot_mapping_ptr=cad.slot_mapping,
# Outputs
out_input_ids_ptr=self.input_ids,
out_context_positions_ptr=self._context_positions_buffer,
out_query_positions_ptr=self.positions,
out_context_slot_mapping_ptr=self._context_slot_mapping_buffer,
out_query_slot_mapping_ptr=self._slot_mapping_buffer,
out_token_indices_ptr=token_indices_to_sample,
# Block table
block_table_ptr=cad.block_table_tensor,
block_table_stride=cad.block_table_tensor.stride(0),
# Metadata
query_start_loc_ptr=cad.query_start_loc,
seq_lens_ptr=cad.seq_lens,
num_rejected_tokens_ptr=(num_rejected_tokens_gpu if has_num_rejected else 0),
# Scalars
parallel_drafting_token_id=self.parallel_drafting_token_id,
block_size=self.kernel_block_size,
num_query_per_req=num_query_per_req,
num_speculative_tokens=self.num_speculative_tokens,
total_input_tokens=num_context,
batch_size=batch_size,
HAS_NUM_REJECTED=has_num_rejected,
)
query_slot_mapping = self._slot_mapping_buffer[:num_query_total]
new_query_start_loc = self.arange_dflash[: batch_size + 1] * num_query_per_req
effective_seq_lens = cad.seq_lens
if has_num_rejected:
effective_seq_lens = effective_seq_lens - num_rejected_tokens_gpu
cad.query_start_loc = new_query_start_loc
cad.seq_lens = effective_seq_lens + num_query_per_req
cad.query_start_loc_cpu = (
torch.from_numpy(self.token_arange_np[: batch_size + 1]).clone() * num_query_per_req
).to(torch.int32)
if hasattr(cad, "actual_seq_lengths_q"):
cad.actual_seq_lengths_q = [num_query_per_req] * batch_size
if hasattr(cad, "decode_token_per_req"):
cad.decode_token_per_req = num_query_per_req
cad.num_actual_tokens = num_query_total
cad.max_query_len = num_query_per_req
cad.max_seq_len = cad.max_seq_len + num_query_per_req
cad.slot_mapping = query_slot_mapping
cad.causal = False
cad.attn_mask = None
cad.attn_state = AscendAttentionState.ChunkedPrefill
return num_query_total, token_indices_to_sample, cad, None
@torch.inference_mode()
def dummy_run(
self,
num_tokens: int,
num_reqs: int = 0,
num_tokens_across_dp: torch.Tensor | None = None,
aclgraph_runtime_mode: CUDAGraphMode = CUDAGraphMode.NONE,
batch_descriptor=None,
dummy_compute_logits=lambda hidden_states: None,
is_profile=False,
**kwargs,
) -> None:
num_query_tokens = min(num_tokens, self.max_query_tokens)
(
num_input_tokens,
num_tokens_across_dp,
_,
) = self.runner._sync_metadata_across_dp(num_query_tokens, is_draft_model=True)
if not self.use_cuda_graph:
aclgraph_runtime_mode = CUDAGraphMode.NONE
num_query_per_req = 1 + self.num_speculative_tokens
num_query_total = num_reqs * num_query_per_req
context_positions = self._context_positions_buffer[:num_input_tokens]
context_states = self.hidden_states[:num_input_tokens]
multi_steps_attn_metadata = []
if aclgraph_runtime_mode == CUDAGraphMode.FULL and len(self.runner.attn_groups) > 0:
builder = self.draft_attn_groups[0].get_metadata_builder()
common_attn_metadata = AscendCommonAttentionMetadata(
query_start_loc=self.arange_dflash[: num_reqs + 1] * num_query_per_req,
query_start_loc_cpu=torch.from_numpy(self.token_arange_np[: num_reqs + 1]).clone() * num_query_per_req,
seq_lens_cpu=self.runner.optimistic_seq_lens_cpu,
seq_lens_cpu_upper_bound=self.runner.optimistic_seq_lens_cpu,
seq_lens=self.runner.seq_lens[:num_reqs],
num_reqs=num_reqs,
num_actual_tokens=num_query_tokens,
max_query_len=num_query_per_req,
max_seq_len=0,
slot_mapping=self._slot_mapping_buffer[:num_query_total],
attn_state=AscendAttentionState.ChunkedPrefill,
causal=False,
is_prefilling=torch.zeros(num_reqs, dtype=torch.bool),
block_table_tensor=self.runner.input_batch.block_table[self.kv_cache_gid].get_device_tensor()[
:num_reqs
],
)
attn_metadata_dflash = builder.build_for_graph_capture(
common_attn_metadata,
AscendAttentionState.ChunkedPrefill,
)
attn_metadata_dflash.attn_mask = None
attn_metadata_dflash.attn_state = AscendAttentionState.ChunkedPrefill
per_layer_attn_metadata = dict()
for layer_name in self.attn_layer_names:
per_layer_attn_metadata[layer_name] = attn_metadata_dflash
multi_steps_attn_metadata.append(per_layer_attn_metadata)
self.token_indices_to_sample.fill_(0)
with set_ascend_forward_context(
multi_steps_attn_metadata[0] if multi_steps_attn_metadata else None,
self.vllm_config,
num_tokens=num_input_tokens,
num_tokens_across_dp=num_tokens_across_dp,
num_actual_tokens=num_input_tokens,
in_profile_run=is_profile,
batch_descriptor=batch_descriptor,
aclgraph_runtime_mode=aclgraph_runtime_mode,
is_draft_model=True,
draft_attn_metadatas=multi_steps_attn_metadata,
):
if is_profile:
self.model.precompute_and_store_context_kv(context_states, context_positions)
self.model(
input_ids=self.input_ids[:num_query_total],
positions=self._get_positions(num_query_total),
inputs_embeds=None,
)
else:
self._dflash_num_context = num_input_tokens
self._runnable(
num_input_tokens=num_input_tokens,
batch_size=num_reqs,
token_indices_to_sample=self.token_indices_to_sample[: num_reqs * self.num_speculative_tokens],
target_positions=self._get_positions(num_input_tokens),
inputs_embeds=None,
multi_steps_attn_metadata=multi_steps_attn_metadata,
num_tokens=num_input_tokens,
)
forward_context = get_forward_context()
if forward_context.cudagraph_runtime_mode == CUDAGraphMode.FULL and not _EXTRA_CTX.capturing:
self._update_full_graph_params(forward_context, num_tokens, multi_steps_attn_metadata)
def build_model_inputs_first_pass(
self,
num_input_tokens: int,
) -> dict[str, Any]:
num_context = self._dflash_num_context
self.model.precompute_and_store_context_kv(
self._dflash_hidden_states[:num_context],
self._context_positions_buffer[:num_context],
self._context_slot_mapping_buffer[:num_context],
)
return dict(
input_ids=self.input_ids[:num_input_tokens], positions=self.positions[:num_input_tokens], inputs_embeds=None
)
def _raise_if_multimodal(self):
pass

View File

@@ -0,0 +1,17 @@
import torch
from vllm.config import VllmConfig
from vllm.v1.spec_decode.draft_model import DraftModelProposer
from vllm_ascend.spec_decode.llm_base_proposer import AscendSpecDecodeBaseProposer
class AscendDraftModelProposer(DraftModelProposer, AscendSpecDecodeBaseProposer):
def __init__(
self,
vllm_config: VllmConfig,
device: torch.device,
runner=None,
):
AscendSpecDecodeBaseProposer.__init__(self, vllm_config, device, False, runner=runner)
self._raise_if_vocab_size_mismatch()
self._raise_if_draft_tp_mismatch()

View File

@@ -1,674 +1,12 @@
# SPDX-License-Identifier: Apache-2.0
from typing import Optional
import numpy as np
import torch
import torch.nn as nn
from vllm.attention.layer import Attention
from vllm.config import (CompilationLevel, VllmConfig,
get_layers_from_vllm_config)
from vllm.distributed.parallel_state import get_pp_group
from vllm.logger import logger
from vllm.model_executor.model_loader import get_model
from vllm.model_executor.models import supports_multimodal
from vllm.model_executor.models.llama_eagle3 import Eagle3LlamaForCausalLM
from vllm.v1.core.sched.output import SchedulerOutput
from vllm.v1.sample.metadata import SamplingMetadata
from vllm.v1.spec_decode.metadata import SpecDecodeMetadata
from vllm.config import VllmConfig
from vllm.v1.spec_decode.eagle import EagleProposer
from vllm_ascend.ascend_forward_context import set_ascend_forward_context
from vllm_ascend.attention.attention_mask import AttentionMaskBuilder
from vllm_ascend.attention.attention_v1 import AscendAttentionState
from vllm_ascend.attention.utils import AscendCommonAttentionMetadata
from vllm_ascend.spec_decode.interface import Proposer, SpecDcodeType
from vllm_ascend.utils import vllm_version_is
PADDING_SLOT_ID = -1
from vllm_ascend.spec_decode.llm_base_proposer import AscendSpecDecodeBaseProposer
class EagleProposer(Proposer):
def __init__(self,
vllm_config: VllmConfig,
device: torch.device,
runner=None):
self.name = SpecDcodeType.EAGLE if vllm_config.speculative_config.method == "eagle" else SpecDcodeType.EAGLE3
self.vllm_config = vllm_config
self.device = device
self.runner = runner
self.block_size = vllm_config.cache_config.block_size
# We need to get the hidden size from the draft model config because
# the draft model's hidden size can be different from the target model's
# hidden size (e.g., Llama 3.3 70B).
self.hidden_size = vllm_config.speculative_config.draft_model_config.get_hidden_size(
)
self.use_cuda_graph = (self.vllm_config.compilation_config.level
== CompilationLevel.PIECEWISE and
not self.vllm_config.model_config.enforce_eager)
self.cudagraph_batch_sizes = list(
reversed(
self.vllm_config.compilation_config.cudagraph_capture_sizes))
# persistent buffers for cuda graph
self.input_ids = torch.zeros(
self.vllm_config.scheduler_config.max_num_batched_tokens,
dtype=torch.int32,
device=device)
self.positions = torch.zeros(
self.vllm_config.scheduler_config.max_num_batched_tokens,
dtype=torch.int64,
device=device)
self.hidden_states = torch.zeros(
(self.vllm_config.scheduler_config.max_num_batched_tokens,
self.hidden_size),
dtype=self.vllm_config.model_config.dtype,
device=device)
# We need +1 here because the arange is used to set query_start_loc,
# which has one more element than batch_size.
self.arange = torch.arange(vllm_config.scheduler_config.max_num_seqs +
1,
device=device,
dtype=torch.int32)
attn_mask_len = self.vllm_config.model_config.max_model_len
self.attn_mask_builder = AttentionMaskBuilder(
attn_mask_len, self.vllm_config.model_config.dtype)
def load_model(self, model: nn.Module) -> None:
target_attn_layer_names = set(
get_layers_from_vllm_config(self.vllm_config, Attention).keys())
self.model = get_model(vllm_config=self.vllm_config,
model_config=self.vllm_config.
speculative_config.draft_model_config)
draft_attn_layer_names = (
get_layers_from_vllm_config(self.vllm_config, Attention).keys() -
target_attn_layer_names)
self.attn_layer_name = next(iter(draft_attn_layer_names))
# share embed_tokens with the target model if needed
if get_pp_group().world_size == 1:
logger.info(
"The EAGLE head shares the same vocab embedding" \
" with the target model."
)
self.model.model.embed_tokens = model.model.embed_tokens
else:
logger.info(
"Since PP > 1, the EAGLE head loaded its own vocab embedding" \
" weights instead of sharing them with the target model."
)
# share lm_head with the target model if needed
# some model definition do not define lm_head explicitly
# and reuse embed_tokens for lm_head, e.g., CohereForCausalLM
if self.name == SpecDcodeType.EAGLE and hasattr(model, "lm_head"):
logger.info("Loading EAGLE LM head weights from the target model.")
if supports_multimodal(model):
self.model.lm_head = model.get_language_model().lm_head
else:
self.model.lm_head = model.lm_head
@torch.inference_mode()
def dummy_run(self,
num_tokens: int,
with_prefill: bool = False,
skip_attn: bool = False,
num_reqs: int = 0,
num_tokens_across_dp: Optional[torch.Tensor] = None):
moe_comm_type = self.runner._select_moe_comm_method(
num_tokens, with_prefill)
with set_ascend_forward_context(None,
self.vllm_config,
moe_comm_type=moe_comm_type,
num_tokens=num_tokens):
self.model(
input_ids=self.input_ids[:num_tokens],
positions=self.positions[:num_tokens],
hidden_states=self.hidden_states[:num_tokens],
)
def generate_token_ids(self,
valid_sampled_token_ids: list[list[int]],
sampling_metadata: SamplingMetadata = None,
scheduler_output: SchedulerOutput = None,
spec_decode_metadata: SpecDecodeMetadata = None,
positions: torch.Tensor = None,
num_scheduled_tokens: int = 0,
hidden_states: torch.Tensor = None,
attn_metadata=None,
aux_hidden_states: torch.Tensor = None):
attn_metadata = self._get_eagle_atten_dict(scheduler_output)
next_token_ids: list[int] = []
for i, token_ids in enumerate(valid_sampled_token_ids):
if token_ids:
# Common case.
next_token_id = token_ids[-1]
else:
# Partial prefill (rare case).
# Get the next token id from the request state.
req_id = self.runner.input_batch.req_ids[i]
req_state = self.runner.requests[req_id]
seq_len = (req_state.num_computed_tokens +
scheduler_output.num_scheduled_tokens[req_id])
next_token_id = req_state.get_token_id(seq_len)
next_token_ids.append(next_token_id)
next_token_ids = torch.tensor(next_token_ids,
dtype=torch.int32,
device=self.device)
eagle_attn_metadata = attn_metadata[self.attn_layer_name]
if spec_decode_metadata is None:
# input_ids can be None for multimodal models.
target_token_ids = self.runner.input_ids[:num_scheduled_tokens]
target_positions = positions[:num_scheduled_tokens]
if self.name == SpecDcodeType.EAGLE3:
target_hidden_states = torch.cat(
[h[:num_scheduled_tokens] for h in aux_hidden_states],
dim=-1)
else:
target_hidden_states = hidden_states[:num_scheduled_tokens]
target_slot_mapping = eagle_attn_metadata.slot_mapping
cu_num_tokens = eagle_attn_metadata.query_start_loc
else:
num_draft_tokens = spec_decode_metadata.num_draft_tokens
num_rejected_tokens = [
n + 1 - len(valid_sampled_token_ids[i]) if n > 0 else 0
for i, n in enumerate(num_draft_tokens)
]
num_rejected_tokens = torch.tensor(
num_rejected_tokens,
dtype=torch.int32,
device=self.device,
)
num_tokens = num_scheduled_tokens - sum(num_rejected_tokens)
cu_num_tokens, token_indices = self._prepare_inputs(
eagle_attn_metadata.query_start_loc, num_rejected_tokens,
num_tokens)
target_token_ids = self.runner.input_ids[token_indices]
target_positions = positions[token_indices]
if self.name == SpecDcodeType.EAGLE3:
target_hidden_states = torch.cat(
[h[token_indices] for h in aux_hidden_states], dim=-1)
else:
target_hidden_states = hidden_states[token_indices]
target_slot_mapping = eagle_attn_metadata.slot_mapping[
token_indices]
draft_token_ids = self._propose(
target_token_ids=target_token_ids,
target_positions=target_positions,
target_hidden_states=target_hidden_states,
target_slot_mapping=target_slot_mapping,
next_token_ids=next_token_ids,
cu_num_tokens=cu_num_tokens,
block_table=eagle_attn_metadata.block_tables,
sampling_metadata=sampling_metadata,
)
spec_token_ids = draft_token_ids.tolist()
return spec_token_ids
def _get_eagle_atten_dict(
self,
scheduler_output: "SchedulerOutput",
):
total_num_scheduled_tokens = scheduler_output.total_num_scheduled_tokens
assert total_num_scheduled_tokens > 0
num_reqs = self.runner.input_batch.num_reqs
assert num_reqs > 0
# OPTIMIZATION: Start copying the block table first.
# This way, we can overlap the copy with the following CPU operations.
self.runner.input_batch.block_table.commit_block_table(num_reqs)
# Get the number of scheduled tokens for each request.
req_ids = self.runner.input_batch.req_ids
tokens = [scheduler_output.num_scheduled_tokens[i] for i in req_ids]
num_scheduled_tokens = np.array(tokens, dtype=np.int32)
max_num_scheduled_tokens = max(tokens)
self.runner.query_lens = torch.from_numpy(num_scheduled_tokens)
# Get request indices.
# E.g., [2, 5, 3] -> [0, 0, 1, 1, 1, 1, 1, 2, 2, 2]
req_indices = np.repeat(self.runner.arange_np[:num_reqs],
num_scheduled_tokens)
# cu_num_tokens: [2, 5, 3] -> [2, 7, 10]
# arange: [0, 1, 0, 1, 2, 3, 4, 0, 1, 2]
cu_num_tokens, arange = self._get_cumsum_and_arange(
num_scheduled_tokens)
# Get positions.
positions_np = self.runner.positions_np[:total_num_scheduled_tokens]
np.add(self.runner.input_batch.num_computed_tokens_cpu[req_indices],
arange,
out=positions_np)
# Calculate M-RoPE positions.
# Only relevant for models using M-RoPE (e.g, Qwen2-VL)
if self.runner.uses_mrope:
self.runner._calc_mrope_positions(scheduler_output)
# Get token indices.
# E.g., [0, 1, 0, 1, 2, 3, 4, 0, 1, 2]
# -> [0, 1, M, M + 1, M + 2, M + 3, M + 4, 2 * M, 2 * M + 1, 2 * M + 2]
# where M is the max_model_len.
token_indices = (
positions_np +
req_indices * self.runner.input_batch.token_ids_cpu.shape[1])
# NOTE(woosuk): We use torch.index_select instead of np.take here
# because torch.index_select is much faster than np.take for large
# tensors.
torch.index_select(
self.runner.input_batch.token_ids_cpu_tensor.flatten(),
0,
torch.from_numpy(token_indices),
out=self.runner.input_ids_cpu[:total_num_scheduled_tokens])
# Prepare the attention metadata for each KV cache group and make layers
# in the same group share the same metadata.
# NOTE(Chen): there is exactly one KV cache group that contains all
# attetnion layers in the model for now, so the current logic for
# getting attn_metadata is not related to kv_cache_group information.
# Will extend this part to support multiple KV cache groups later.
for kv_cache_group_id, kv_cache_group_spec in enumerate(
self.runner.kv_cache_config.kv_cache_groups):
block_size = kv_cache_group_spec.kv_cache_spec.block_size
block_table = self.runner.input_batch.block_table[
kv_cache_group_id]
# E.g., [0, 1, 0, 1, 2, 3, 4, 0, 1, 2]
# -> [0, 0, K, K, K + 1, K + 1, K + 2, 2 * K, 2 * K, 2 * K + 1]
# where K is the max_num_blocks_per_req and the block size is 2.
# NOTE(woosuk): We can't simply use `token_indices // block_size`
# here because M (max_model_len) is not necessarily divisible by
# block_size.
block_table_indices = (
req_indices * block_table.max_num_blocks_per_req +
positions_np // block_size)
block_table_cpu = block_table.get_cpu_tensor()
block_numbers = block_table_cpu.flatten(
)[block_table_indices].numpy()
block_offsets = positions_np % block_size
np.add(
block_numbers * block_size,
block_offsets,
out=block_table.slot_mapping_np[:total_num_scheduled_tokens])
# Prepare the attention metadata.
self.runner.query_start_loc_np[0] = 0
self.runner.query_start_loc_np[1:num_reqs + 1] = cu_num_tokens
self.runner.seq_lens_np[:num_reqs] = (
self.runner.input_batch.num_computed_tokens_cpu[:num_reqs] +
num_scheduled_tokens)
# Copy the tensors to the NPU.
self.runner.input_ids[:total_num_scheduled_tokens].copy_(
self.runner.input_ids_cpu[:total_num_scheduled_tokens],
non_blocking=True)
if self.runner.uses_mrope:
# Only relevant for models using M-RoPE (e.g, Qwen2-VL)
self.runner.mrope_positions[:, :total_num_scheduled_tokens].copy_(
self.runner.
mrope_positions_cpu[:, :total_num_scheduled_tokens],
non_blocking=True)
else:
# Common case (1D positions)
self.runner.positions[:total_num_scheduled_tokens].copy_(
self.runner.positions_cpu[:total_num_scheduled_tokens],
non_blocking=True)
self.runner.query_start_loc[:num_reqs + 1].copy_(
self.runner.query_start_loc_cpu[:num_reqs + 1], non_blocking=True)
self.runner.seq_lens[:num_reqs].copy_(
self.runner.seq_lens_cpu[:num_reqs], non_blocking=True)
# Fill unused with -1. Needed for reshape_and_cache
self.runner.seq_lens[num_reqs:].fill_(0)
self.runner.query_start_loc[num_reqs + 1:].fill_(-1)
attn_metadata = {}
# Prepare the attention metadata for each KV cache group and make layers
# in the same group share the same metadata.
for kv_cache_group_id, kv_cache_group_spec in enumerate(
self.runner.kv_cache_config.kv_cache_groups):
common_attn_metadata = AscendCommonAttentionMetadata(
query_start_loc=self.runner.query_start_loc[:num_reqs + 1],
query_start_loc_cpu=self.runner.query_start_loc_cpu[:num_reqs +
1],
seq_lens_cpu=self.runner.seq_lens_cpu,
num_reqs=num_reqs,
max_query_len=max_num_scheduled_tokens,
num_actual_tokens=total_num_scheduled_tokens,
actual_seq_lengths_q=self.runner.actual_seq_lengths_q,
block_table_tensor=self.runner.input_batch.block_table[0].
get_device_tensor(),
slot_mapping=self.runner.slot_mapping,
positions=self.runner.positions,
attn_mask=self.runner.attn_mask,
spec_attn_mask=self.runner.spec_attn_mask,
attn_state=self.runner.attn_state,
decode_token_per_req=self.runner.decode_token_per_req,
num_computed_tokens_cpu=None,
seq_lens=None)
if vllm_version_is("0.10.2"):
builder = self.runner.attn_groups[0][0].metadata_builder
else:
builder = self.runner.attn_groups[0][0].get_metadata_builder()
attn_metadata_i = builder.build(0, common_attn_metadata,
self.runner.get_model())
for layer_name in kv_cache_group_spec.layer_names:
attn_metadata[layer_name] = attn_metadata_i
return attn_metadata
def _get_cumsum_and_arange(
self,
num_tokens: np.ndarray,
cumsum_dtype: Optional[np.dtype] = None,
) -> tuple[np.ndarray, np.ndarray]:
"""Get the cumulative sum and batched arange of the given array.
# E.g., [2, 5, 3] -> ([2, 7, 10], [0, 1, 0, 1, 2, 3, 4, 0, 1, 2])
# Equivalent to but faster than:
# np.concatenate([np.arange(n) for n in num_tokens])
"""
# Step 1. [2, 5, 3] -> [2, 7, 10]
cu_num_tokens = np.cumsum(num_tokens, dtype=cumsum_dtype)
total_num_tokens = cu_num_tokens[-1]
# Step 2. [2, 7, 10] -> [0, 0, 2, 2, 2, 2, 2, 7, 7, 7]
cumsums_offsets = np.repeat(cu_num_tokens - num_tokens, num_tokens)
# Step 3. [0, 1, 0, 1, 2, 3, 4, 0, 1, 2]
arange = self.runner.arange_np[:total_num_tokens] - cumsums_offsets
return cu_num_tokens, arange
def _propose(
self,
# [num_tokens]
target_token_ids: torch.Tensor,
# [num_tokens]
target_positions: torch.Tensor,
# [num_tokens, hidden_size]
target_hidden_states: torch.Tensor,
# [num_tokens]
target_slot_mapping: torch.Tensor,
# [batch_size]
next_token_ids: torch.Tensor,
# [batch_size + 1] starting with 0
cu_num_tokens: torch.Tensor,
# [batch_size, max_num_blocks_per_req]
block_table: torch.Tensor,
sampling_metadata: SamplingMetadata,
) -> torch.Tensor:
device = cu_num_tokens.device
cu_num_tokens = cu_num_tokens.cpu()
block_table = block_table.cpu()
num_tokens = target_token_ids.shape[0]
batch_size = next_token_ids.shape[0]
last_token_indices = cu_num_tokens[1:] - 1
target_positions = target_positions.cpu()
if self.name == SpecDcodeType.EAGLE3:
assert isinstance(self.model, Eagle3LlamaForCausalLM)
target_hidden_states = self.model.combine_hidden_states(
target_hidden_states)
assert target_hidden_states.shape[-1] == self.hidden_size
# Shift the input ids by one token.
# E.g., [a1, b1, b2, c1, c2, c3] -> [b1, b2, c1, c2, c3, c3]
self.input_ids[:num_tokens - 1] = target_token_ids[1:]
# Replace the last token with the next token.
# E.g., [b1, b2, c1, c2, c3, c3] -> [a2, b2, b3, c2, c3, c4]
self.input_ids[last_token_indices] = next_token_ids
seq_lens = (target_positions[last_token_indices] + 1).int()
query_lens = cu_num_tokens[1:] - cu_num_tokens[:-1]
max_query_len = query_lens.max().item()
attn_mask = self.attn_mask_builder.get_splitfuse_attn_mask(
seq_lens, target_positions, self.vllm_config.model_config.dtype,
self.device)
common_attn_metadata = AscendCommonAttentionMetadata(
query_start_loc=cu_num_tokens.to(device),
query_start_loc_cpu=cu_num_tokens,
seq_lens_cpu=seq_lens.cpu(),
max_query_len=max_query_len,
num_reqs=batch_size,
num_actual_tokens=num_tokens,
actual_seq_lengths_q=self.runner.actual_seq_lengths_q,
block_table_tensor=self.runner.input_batch.block_table[0].
get_device_tensor(),
slot_mapping=target_slot_mapping,
positions=target_positions,
attn_mask=attn_mask,
spec_attn_mask=self.runner.spec_attn_mask,
attn_state=self.runner.attn_state,
decode_token_per_req=self.runner.decode_token_per_req,
num_computed_tokens_cpu=None,
seq_lens=None)
# FIXME(woosuk): The below two ops cause synchronization. Optimize.
if vllm_version_is("0.10.2"):
builder = self.runner.attn_groups[0][0].metadata_builder
else:
builder = self.runner.attn_groups[0][0].get_metadata_builder()
attn_metadata = builder.build(0, common_attn_metadata,
self.runner.get_model())
if self.use_cuda_graph and \
num_tokens <= self.cudagraph_batch_sizes[-1]:
num_input_tokens = self.vllm_config.pad_for_cudagraph(num_tokens)
else:
num_input_tokens = num_tokens
with_prefill = attn_metadata.attn_state not in [
AscendAttentionState.DecodeOnly, AscendAttentionState.SpecDecoding
]
moe_comm_type = self.runner._select_moe_comm_method(
num_input_tokens, with_prefill)
# copy inputs to buffer for cudagraph
self.positions[:num_tokens] = target_positions.to(device)
self.hidden_states[:num_tokens] = target_hidden_states
attn_metadata.block_tables = block_table.to(device)
with set_ascend_forward_context(attn_metadata,
self.vllm_config,
moe_comm_type=moe_comm_type,
num_tokens=num_input_tokens):
last_hidden_states, hidden_states = self.model(
input_ids=self.input_ids[:num_input_tokens],
positions=self.positions[:num_input_tokens],
hidden_states=self.hidden_states[:num_input_tokens],
)
sample_hidden_states = last_hidden_states[last_token_indices]
if vllm_version_is("0.10.2"):
logits = self.model.compute_logits(sample_hidden_states, None)
else:
logits = self.model.compute_logits(sample_hidden_states)
draft_token_ids = logits.argmax(dim=-1)
# Early exit if there is only one draft token to be generated.
if self.vllm_config.speculative_config.num_speculative_tokens == 1:
# [batch_size, 1]
return draft_token_ids.view(-1, 1)
# Generate the remaining draft tokens.
draft_token_ids_tensor = torch.zeros(
(self.vllm_config.speculative_config.num_speculative_tokens,
*draft_token_ids.shape),
dtype=draft_token_ids.dtype)
draft_token_ids_tensor[0] = draft_token_ids
positions_cpu = target_positions[last_token_indices].cpu().to(
torch.int64)
hidden_states = hidden_states[last_token_indices]
if self.use_cuda_graph and \
batch_size <= self.cudagraph_batch_sizes[-1]:
input_batch_size = self.vllm_config.pad_for_cudagraph(batch_size)
else:
input_batch_size = batch_size
moe_comm_type = self.runner._select_moe_comm_method(
input_batch_size, False)
attn_metadata.num_actual_tokens = batch_size
attn_metadata.max_query_len = 1
attn_metadata.query_start_loc = self.arange[:batch_size + 1]
query_lens.fill_(1)
attn_metadata.query_lens = query_lens
attn_metadata.attn_state = AscendAttentionState.ChunkedPrefill
for now_speculative in range(
self.vllm_config.speculative_config.num_speculative_tokens -
1):
# Update the inputs.
# cast to int32 is crucial when eagle model is compiled.
# tensor.argmax() returns int64 by default.
input_ids = draft_token_ids_tensor[now_speculative].to(device)
positions_cpu += 1
# NOTE(woosuk): We should handle the case where the draft model
# generates tokens beyond the max model length. Since it is complex
# to remove such requests from the batch, we keep them in the batch
# but adjust the position ids and slot mappings to avoid the
# out-of-range access during the model execution. The draft tokens
# generated with this adjustment should be ignored.
exceeds_max_model_len = positions_cpu >= self.vllm_config.model_config.max_model_len
# Mask out the position ids that exceed the max model length.
# Otherwise, we may get out-of-range error in RoPE.
clamped_positions_cpu = torch.where(exceeds_max_model_len, 0,
positions_cpu)
clamped_positions = clamped_positions_cpu.to(device)
# TODO: Increment the sequence lengths.
attn_metadata.seq_lens += 1
# TODO: Consider max model length.
# attn_metadata.max_seq_len = min(attn_metadata.max_seq_len,
# self.max_model_len)
# For the requests that exceed the max model length, we set the
# TODO: sequence length to 1 to minimize their overheads in attention.
# Compute the slot mapping.
block_numbers = (clamped_positions_cpu // self.block_size)
block_ids = block_table.gather(dim=1,
index=block_numbers.view(-1, 1))
block_ids = block_ids.view(-1)
slot_mapping_cpu = (
block_ids * self.vllm_config.cache_config.block_size +
clamped_positions_cpu % self.block_size)
# Mask out the slot mappings that exceed the max model length.
# Otherwise, the KV cache will be inadvertently updated with the
# padding tokens.
slot_mapping_cpu.masked_fill_(exceeds_max_model_len,
PADDING_SLOT_ID)
# NOTE: ASCEND slot_mapping must on cpu
attn_metadata.slot_mapping = slot_mapping_cpu.to(
torch.int32).to(device)
# copy inputs to buffer for cudagraph
self.input_ids[:batch_size] = input_ids
self.positions[:batch_size] = clamped_positions
self.hidden_states[:batch_size] = hidden_states
attn_mask = self.attn_mask_builder.get_splitfuse_attn_mask(
attn_metadata.seq_lens, positions_cpu,
self.vllm_config.model_config.dtype, self.device)
attn_metadata.attn_mask = attn_mask
attn_metadata.block_tables = block_table.to(device)
# Run the model.
with set_ascend_forward_context(attn_metadata,
self.vllm_config,
moe_comm_type=moe_comm_type,
num_tokens=input_batch_size):
last_hidden_states, hidden_states = self.model(
input_ids=self.input_ids[:input_batch_size],
positions=self.positions[:input_batch_size],
hidden_states=self.hidden_states[:input_batch_size],
)
hidden_states = hidden_states[:batch_size]
if vllm_version_is("0.10.2"):
logits = self.model.compute_logits(
last_hidden_states[:batch_size], None)
else:
logits = self.model.compute_logits(
last_hidden_states[:batch_size])
# TODO(wenlong): get more than one token for tree attention
draft_token_ids = logits.argmax(dim=-1)
draft_token_ids_tensor[now_speculative + 1] = draft_token_ids.cpu()
# [batch_size, num_speculative_tokens]
draft_token_ids = draft_token_ids_tensor.swapaxes(0, 1)
return draft_token_ids
def _prepare_inputs(
self,
# [batch_size + 1]
cu_target_query_lens: torch.Tensor,
# [batch_size]
num_rejected_tokens: torch.Tensor,
num_tokens: int,
) -> tuple[torch.Tensor, torch.Tensor]:
# cu_target_query_lens: [0, a, a + b, a + b + c]
# num_rejected_tokens: [n1, n2, n3]
# num_tokens_per_req: [a - n1, b - n2, c - n3]
# cu_num_tokens: [0, a - n1, a + b - n1 - n2, a + b + c - n1 - n2 - n3]
# token_indices: [0, 1, ..., a - n1 - 1,
# a, a + 1, ..., a + b - n2 - 1,
# a + b, a + b + 1, ..., a + b + c - n3 - 1]
# [0, a, a + b, a + b + c] -> [a, b, c]
query_len_per_req = (cu_target_query_lens[1:] -
cu_target_query_lens[:-1])
# [a, b, c] -> [a - n1, b - n2, c - n3]
num_tokens_per_req = query_len_per_req - num_rejected_tokens
# [a - n1, b - n2, c - n3] ->
# [0, a - n1, a + b - n1 - n2, a + b + c - n1 - n2 - n3]
cu_num_tokens = torch.zeros_like(cu_target_query_lens)
torch.cumsum(num_tokens_per_req, dim=0, out=cu_num_tokens[1:])
token_indices = torch.empty(
num_tokens,
dtype=torch.int32,
device=cu_target_query_lens.device,
)
BLOCK_SIZE = 1024
self._prepare_eagle_input_sequential(
token_indices,
cu_target_query_lens,
cu_num_tokens,
block_size=BLOCK_SIZE,
)
return cu_num_tokens, token_indices
def _prepare_eagle_input_sequential(self, out_tensor: torch.Tensor,
cu_query_lens: torch.Tensor,
cu_num_tokens: torch.Tensor,
block_size: int):
num_programs = len(cu_num_tokens) - 1
for pid in range(num_programs):
start_pos = cu_num_tokens[pid].item()
end_pos = cu_num_tokens[pid + 1].item()
num_tokens = end_pos - start_pos
index_start = cu_query_lens[pid].item()
num_blocks = int(
torch.ceil(torch.tensor(num_tokens / block_size)).item())
for i in range(num_blocks):
offset_tensor = torch.arange(0,
block_size,
dtype=torch.int32,
device=out_tensor.device)
global_start_offset = i * block_size
target_indices = torch.tensor(
start_pos + global_start_offset,
dtype=torch.int32,
device=out_tensor.device) + offset_tensor
values_to_store = torch.tensor(
index_start + global_start_offset,
dtype=torch.int32,
device=out_tensor.device) + offset_tensor
mask = (target_indices >= start_pos) & \
(target_indices < end_pos) & \
(offset_tensor < num_tokens)
out_tensor[target_indices[mask]] = values_to_store[mask]
class AscendEagleProposer(EagleProposer, AscendSpecDecodeBaseProposer):
def __init__(self, vllm_config: VllmConfig, device: torch.device, runner=None):
AscendSpecDecodeBaseProposer.__init__(self, vllm_config, device, True, runner=runner)

View File

@@ -0,0 +1,198 @@
#
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
# Copyright 2023 The vLLM team.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""Ascend adaptation of ExtractHiddenStatesProposer for extracting and caching
hidden states during speculative decoding."""
import torch
from vllm.config import CUDAGraphMode, VllmConfig
from vllm.forward_context import set_forward_context
from vllm.v1.spec_decode.extract_hidden_states import ExtractHiddenStatesProposer
class AscendExtractHiddenStatesProposer(ExtractHiddenStatesProposer):
"""Ascend-adapted ExtractHiddenStatesProposer for NPU devices.
This proposer extracts hidden states from the target model and caches them
in the KV cache without performing actual speculation. It's used with the
ExampleHiddenStatesConnector for KV transfer.
The main differences from the GPU version:
- Uses ACL graphs instead of CUDA graphs
- Implements dummy_run for ACL graph capture with Ascend-specific signature
- Adapts prepare_next_token_ids_padded for Ascend's indices/count pattern
"""
def __init__(self, vllm_config: VllmConfig, device: torch.device, runner=None):
self.runner = runner
super().__init__(vllm_config, device)
@torch.inference_mode()
def _determine_batch_execution_and_padding(
self,
num_tokens: int,
use_cudagraphs: bool = True,
) -> tuple[CUDAGraphMode, int, torch.Tensor | None]:
"""Determine cudagraph mode and padded token count for this proposer step.
Same contract as upstream ``ExtractHiddenStatesProposer`` but on the
Ascend runner path: SP-pad ``num_tokens`` before dispatch and reuse
``runner._sync_metadata_across_dp`` for DP coordination. Upstream's
``coordinate_batch_across_dp`` posts a differently shaped tensor to the
same DP cpu_group as the main runner and breaks the gloo collective.
"""
assert self.runner is not None, (
"AscendExtractHiddenStatesProposer requires a runner reference "
"for _pad_for_sequence_parallelism / _sync_metadata_across_dp"
)
# SP-pad before DP sync, mirroring the main runner. The v2
# NPUModelRunner lacks this hook; raise a clear error instead of an
# opaque AttributeError.
if not hasattr(self.runner, "_pad_for_sequence_parallelism"):
raise NotImplementedError(
"The current model runner does not support sequence "
"parallelism padding (_pad_for_sequence_parallelism) required "
"for AscendExtractHiddenStatesProposer."
)
num_tokens = self.runner._pad_for_sequence_parallelism(num_tokens)
cudagraph_mode, batch_desc = self.cudagraph_dispatcher.dispatch(
num_tokens,
valid_modes=({CUDAGraphMode.NONE} if not use_cudagraphs else None),
)
num_tokens_padded = batch_desc.num_tokens
num_tokens_across_dp = None
if self.vllm_config.parallel_config.data_parallel_size > 1:
# The v2 NPUModelRunner lacks this hook; raise a clear error here
# too.
if not hasattr(self.runner, "_sync_metadata_across_dp"):
raise NotImplementedError(
"The current model runner does not support DP metadata "
"synchronization (_sync_metadata_across_dp) required for "
"data parallel size > 1."
)
# Reuse the runner's DP sync so the collective shape matches the
# main forward. ``is_draft_model=True`` short-circuits the
# all_reduce (cache-only drafter is not MoE); ``dummy_run`` issues
# the identical call to keep busy and idle DP ranks balanced.
(
_max_tokens_across_dp,
num_tokens_across_dp,
synced_cudagraph_mode,
) = self.runner._sync_metadata_across_dp(
num_tokens=num_tokens_padded,
is_draft_model=True,
cudagraph_mode=cudagraph_mode,
allow_dp_padding=use_cudagraphs,
)
if num_tokens_across_dp is not None:
num_tokens_padded = int(num_tokens_across_dp[self.dp_rank].item())
# Re-dispatch with DP-synced padding.
cudagraph_mode, batch_desc = self.cudagraph_dispatcher.dispatch(
num_tokens_padded,
valid_modes={synced_cudagraph_mode},
)
assert batch_desc.num_tokens == num_tokens_padded
return cudagraph_mode, num_tokens_padded, num_tokens_across_dp
@torch.inference_mode()
def dummy_run(
self,
num_tokens,
with_prefill=None,
in_graph_capturing=None,
num_reqs=None,
num_tokens_across_dp=None,
aclgraph_runtime_mode=None,
batch_descriptor=None,
dummy_compute_logits=lambda hidden_states: None,
is_profile=False,
) -> None:
"""Dummy run for ACL graph capture.
Same functional logic as GPU version but with Ascend's parameter signature.
"""
assert self.model is not None, "Model must be initialized before dummy_run"
assert self.runner is not None, (
"AscendExtractHiddenStatesProposer requires a runner reference for _sync_metadata_across_dp"
)
# Idle DP ranks must issue the same drafter DP sync that busy ranks
# issue in _determine_batch_execution_and_padding (mirrors
# llm_base_proposer.dummy_run); otherwise the DP cpu_group collectives
# desynchronize and the group deadlocks.
(
num_tokens,
num_tokens_across_dp,
_,
) = self.runner._sync_metadata_across_dp(num_tokens, is_draft_model=True)
with set_forward_context(
None,
self.vllm_config,
num_tokens=num_tokens,
num_tokens_across_dp=num_tokens_across_dp,
cudagraph_runtime_mode=aclgraph_runtime_mode or CUDAGraphMode.NONE,
slot_mapping={},
):
self.model(
hidden_states=self.hidden_states[:num_tokens],
)
def prepare_next_token_ids_padded(
self,
sampled_token_ids: torch.Tensor,
requests,
gpu_input_batch,
discard_request_indices: torch.Tensor,
num_discarded_requests: int,
) -> tuple[torch.Tensor, torch.Tensor]:
"""Prepare next token IDs for speculative decoding.
Since num_speculative_tokens == 1, sampled_token_ids has shape
(batch_size, 1). For each request we either use the sampled token
(if valid and not discarded) or a backup token from the request state.
This adapts the GPU version for Ascend's indices/count pattern
(discard_request_indices instead of boolean mask).
"""
num_reqs = gpu_input_batch.num_reqs
device = sampled_token_ids.device
# Compute backup tokens for discarded / invalid requests
seq_lens_list = (gpu_input_batch.num_tokens_no_spec[:num_reqs] - 1).tolist()
backup_tokens = torch.tensor(
[requests[gpu_input_batch.req_ids[i]].get_token_id(seq_lens_list[i]) for i in range(num_reqs)],
dtype=torch.int32,
device=device,
)
# Create discard mask from indices (Ascend uses indices/count pattern)
discard_mask = torch.zeros(num_reqs, dtype=torch.bool, device=device)
discard_mask[discard_request_indices[:num_discarded_requests]] = True
# With num_speculative_tokens == 1, there is exactly one token
sampled = sampled_token_ids[:, 0]
is_valid = (sampled >= 0) & (sampled < gpu_input_batch.vocab_size)
valid_sampled_tokens_count = is_valid.to(torch.int32)
use_sampled = is_valid & ~discard_mask
next_token_ids = torch.where(use_sampled, sampled.to(torch.int32), backup_tokens)
return next_token_ids, valid_sampled_tokens_count

File diff suppressed because it is too large Load Diff

View File

@@ -0,0 +1,70 @@
import torch
from vllm.config import CUDAGraphMode
from vllm.v1.sample.metadata import SamplingMetadata
from vllm.v1.spec_decode.medusa import MedusaProposer
from vllm.v1.spec_decode.metadata import SpecDecodeMetadata
from vllm_ascend.ascend_forward_context import set_ascend_forward_context
class AscendMedusaProposer(MedusaProposer):
"""
Medusa proposer class for generating token sequences
"""
@torch.inference_mode()
def dummy_run(
self,
num_tokens: int,
with_prefill: bool = False,
in_graph_capturing: bool = False,
num_reqs: int = 0,
num_tokens_across_dp: torch.Tensor | None = None,
aclgraph_runtime_mode: CUDAGraphMode = CUDAGraphMode.NONE,
batch_descriptor=None,
dummy_compute_logits=lambda hidden_states: None,
is_profile=False,
):
hidden_states = torch.zeros(
(self.max_num_tokens, self.hidden_size),
dtype=self.dtype,
device=self.device,
)
with set_ascend_forward_context(
None,
self.vllm_config,
num_tokens=num_tokens,
num_actual_tokens=0,
in_profile_run=is_profile,
batch_descriptor=batch_descriptor,
aclgraph_runtime_mode=aclgraph_runtime_mode,
is_draft_model=True,
):
self.model(hidden_states)
dummy_compute_logits(hidden_states)
def propose(
self,
valid_sampled_token_ids: list[list[int]],
sampling_metadata: SamplingMetadata,
spec_decode_metadata: SpecDecodeMetadata,
sample_hidden_states: torch.Tensor,
):
if sample_hidden_states.shape[0] == len(valid_sampled_token_ids):
# The input to the target model does not include draft tokens.
hidden_states = sample_hidden_states
else:
num_accepted_tokens = torch.tensor(
[len(t) for t in valid_sampled_token_ids], device=self.device, dtype=torch.long
)
num_draft_tokens = torch.tensor(spec_decode_metadata.num_draft_tokens, device=self.device, dtype=torch.long)
offsets = torch.cumsum(num_draft_tokens + 1, dim=0) - (num_draft_tokens + 1)
indices = offsets + num_accepted_tokens - 1
hidden_states = sample_hidden_states[indices]
spec_token_ids = super().propose(
target_hidden_states=hidden_states,
sampling_metadata=sampling_metadata,
)
return spec_token_ids

View File

@@ -1,65 +1,102 @@
import torch
from vllm.v1.spec_decode.ngram_proposer import \
NgramProposer as VllmNgramProposer
from vllm.v1.spec_decode.ngram_proposer import NgramProposer
from vllm_ascend.spec_decode.interface import Proposer, SpecDcodeType
from vllm_ascend.utils import vllm_version_is
class NgramProposer(VllmNgramProposer, Proposer):
def __init__(self, vllm_config, device, runner):
super().__init__(vllm_config)
self.name = SpecDcodeType.NGRAM
self.device = device
class AscendNgramProposer(NgramProposer):
def __init__(self, vllm_config, runner):
self.runner = runner
super().__init__(vllm_config)
def load_model(self, *args, **kwargs):
# No model to load.
pass
@torch.inference_mode()
def dummy_run(self,
num_tokens,
with_prefill=None,
skip_attn=None,
num_reqs=None,
num_tokens_across_dp=None):
def dummy_run(
self,
num_tokens,
with_prefill=None,
in_graph_capturing=None,
num_reqs=None,
num_tokens_across_dp=None,
aclgraph_runtime_mode=None,
batch_descriptor=None,
dummy_compute_logits=lambda hidden_states: None,
is_profile=False,
):
pass
def generate_token_ids(self,
valid_sampled_token_ids,
sampling_metadata=None,
scheduler_output=None,
spec_decode_metadata=None,
positions=None,
num_scheduled_tokens=None,
hidden_states=None,
attn_metadata=None,
aux_hidden_states=None) -> list[list[int]]:
# TODO(woosuk): Optimize.
draft_token_ids: list[list[int]] = []
for i, sampled_ids in enumerate(valid_sampled_token_ids):
num_sampled_ids = len(sampled_ids)
if not num_sampled_ids:
# Skip speculative decoding.
draft_token_ids.append([])
continue
if vllm_version_is("0.23.0"):
# Skip requests that require top-p, top-k, etc.
req_id = self.runner.input_batch.req_ids[i]
if req_id in self.runner.input_batch.spec_decode_unsupported_reqs:
draft_token_ids.append([])
continue
def propose(
self,
sampled_token_ids: list[list[int]],
num_tokens_no_spec=None,
token_ids_cpu=None,
slot_mappings: dict[str, torch.Tensor] | list[dict[str, torch.Tensor]] | None = None,
) -> list[list[int]]:
input_batch = self.runner.input_batch
valid_ngram_requests = []
for i, sampled_ids in enumerate(sampled_token_ids):
num_sampled_ids = len(sampled_ids)
if not num_sampled_ids:
continue
# Add sampled_token_ids to token_ids_cpu.
start_idx = self.runner.input_batch.num_tokens_no_spec[i]
end_idx = start_idx + num_sampled_ids
self.runner.input_batch.token_ids_cpu[
i, start_idx:end_idx] = sampled_ids
drafter_output = self.propose(
self.runner.input_batch.token_ids_cpu[i, :end_idx])
if drafter_output is None or len(drafter_output) == 0:
draft_token_ids.append([])
else:
draft_token_ids.append(drafter_output.tolist())
return draft_token_ids
req_id = input_batch.req_ids[i]
if req_id in input_batch.spec_decode_unsupported_reqs:
continue
num_tokens = input_batch.num_tokens_no_spec[i]
if num_tokens >= input_batch.max_model_len:
# Skip requests that have already reached the max model length.
continue
start_idx = input_batch.num_tokens_no_spec[i]
end_idx = start_idx + num_sampled_ids
input_batch.token_ids_cpu[i, start_idx:end_idx] = sampled_ids
valid_ngram_requests.append(i)
return self.batch_propose(
len(sampled_token_ids),
valid_ngram_requests,
input_batch.num_tokens_no_spec,
input_batch.token_ids_cpu,
)
else:
def propose( # type: ignore[misc]
self,
num_speculative_tokens: int,
sampled_token_ids: list[list[int]],
num_tokens_no_spec=None,
token_ids_cpu=None,
slot_mappings: dict[str, torch.Tensor] | list[dict[str, torch.Tensor]] | None = None,
) -> list[list[int]]:
assert num_speculative_tokens <= self.k
assert num_tokens_no_spec is not None
assert token_ids_cpu is not None
valid_ngram_requests = []
for i, sampled_ids in enumerate(sampled_token_ids):
num_sampled_ids = len(sampled_ids)
if not num_sampled_ids:
continue
num_tokens = num_tokens_no_spec[i]
if num_tokens >= self.max_model_len:
# Skip requests that have already reached the max model length.
continue
valid_ngram_requests.append(i)
return self.batch_propose(
len(sampled_token_ids),
valid_ngram_requests,
num_tokens_no_spec,
token_ids_cpu,
num_speculative_tokens,
)

View File

@@ -0,0 +1,35 @@
import torch
from vllm.v1.spec_decode.ngram_proposer_gpu import NgramProposerGPU
class AscendNgramProposerNPU(NgramProposerGPU):
def __init__(self, vllm_config, device: torch.device, runner):
super().__init__(vllm_config, device=device)
def load_model(self, *args, **kwargs):
# No model to load.
pass
@torch.inference_mode()
def dummy_run(
self,
num_tokens,
with_prefill=None,
in_graph_capturing=None,
num_reqs=None,
num_tokens_across_dp=None,
aclgraph_runtime_mode=None,
batch_descriptor=None,
dummy_compute_logits=lambda hidden_states: None,
is_profile=False,
):
pass
def propose(
self,
num_tokens_no_spec: torch.Tensor, # [batch_size]
token_ids_gpu: torch.Tensor, # [batch_size, max_len]
valid_sampled_token_ids_gpu: torch.Tensor, # [batch_size, num_spec_tokens + 1]
valid_sampled_tokens_count: torch.Tensor, # [batch_size]
):
pass

View File

@@ -0,0 +1,830 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
from collections.abc import Callable
from functools import partial
from typing import Any
import torch
from vllm.config import CUDAGraphMode, VllmConfig, get_layers_from_vllm_config, replace
from vllm.forward_context import BatchDescriptor, get_forward_context
from vllm.model_executor.layers.attention_layer_base import AttentionLayerBase
from vllm.model_executor.models.utils import get_draft_quant_config
from vllm.v1.attention.backends.utils import CommonAttentionMetadata
from vllm.v1.core.sched.output import SchedulerOutput
from vllm.v1.kv_cache_interface import KVCacheConfig, KVCacheSpec, UniformTypeKVCacheSpecs
from vllm.v1.sample.metadata import SamplingMetadata
from vllm.v1.spec_decode.llm_base_proposer import compute_probs_and_sample_next_token
from vllm.v1.spec_decode.utils import PADDING_SLOT_ID
from vllm.v1.worker.utils import AttentionGroup
from vllm_ascend.ascend_config import get_ascend_config
from vllm_ascend.ascend_forward_context import _EXTRA_CTX, set_ascend_forward_context
from vllm_ascend.attention.attention_v1 import AscendAttentionState
from vllm_ascend.attention.utils import AscendCommonAttentionMetadata
from vllm_ascend.distributed.parallel_state import get_lmhead_tp_group
from vllm_ascend.spec_decode.eagle_proposer import AscendEagleProposer
from vllm_ascend.utils import lmhead_tp_enable
from vllm_ascend.worker.utils import copy_snapshot_to_gpu
class AscendStep3p5MTPProposer(AscendEagleProposer):
"""Step3.5 MTP proposer with per-MTP-layer independent KV cache groups.
Each Step3.5 MTP layer owns a different attention/KV-cache group. The
generic Ascend MTP proposer assumes one draft KV-cache group and builds one
attention metadata object for all draft layers. Step3.5 therefore keeps the
whole propose flow in this subclass: it builds per-group metadata for all MTP
layers, reuses that step0 metadata across layer calls, and refreshes the
metadata tensors after each in-graph input update.
"""
_runnable: Callable
def __init__(
self,
vllm_config: VllmConfig,
device: torch.device,
runner=None,
) -> None:
super().__init__(vllm_config, device, runner=runner)
# Per KV cache group block tables / slot mappings captured from the
# model runner each step (see set_per_group_attn_metadata).
self._per_group_block_tables: dict[int, torch.Tensor] = {}
self._per_group_slot_mappings: dict[int, torch.Tensor] = {}
# Slot-mapping buffers for additional KV cache groups. FULL graph
# capture records tensor addresses for reshape/cache, so every group
# that can appear in the shared step0 metadata needs a persistent
# buffer.
self._group_slot_buffers: dict[int, torch.Tensor] = {}
def set_per_group_attn_metadata(
self,
gid: int,
block_table: torch.Tensor,
slot_mapping: torch.Tensor,
) -> None:
self._per_group_block_tables[gid] = block_table
self._per_group_slot_mappings[gid] = slot_mapping
def _seed_graph_capture_per_group_metadata(
self,
num_reqs: int,
num_tokens: int,
) -> None:
if self.runner is None or not self.draft_attn_groups:
return
block_tables = self.runner.input_batch.block_table
for attn_group in self.draft_attn_groups:
gid = attn_group.kv_cache_group_id
try:
block_table = block_tables[gid]
except (IndexError, KeyError):
continue
block_table_tensor = block_table.get_device_tensor()
if num_reqs > 0:
block_table_tensor = block_table_tensor[:num_reqs]
slot_mapping = block_table.slot_mapping.gpu[:num_tokens]
self.set_per_group_attn_metadata(gid, block_table_tensor, slot_mapping)
def _slot_mapping_buffer_for_group(
self,
attn_group: AttentionGroup,
) -> torch.Tensor:
gid = attn_group.kv_cache_group_id
if gid == self.kv_cache_gid:
return self.slot_mapping_group[0]
buf = self._group_slot_buffers.get(gid)
if buf is None:
buf = torch.zeros_like(self.slot_mapping_group[0])
self._group_slot_buffers[gid] = buf
return buf
def _common_attn_metadata_for_group(
self,
common_attn_metadata: CommonAttentionMetadata,
attn_group: AttentionGroup,
) -> CommonAttentionMetadata:
group_common_attn_metadata = self.shallow_copy_metadata(common_attn_metadata)
gid = attn_group.kv_cache_group_id
block_table = self._per_group_block_tables.get(gid)
if block_table is not None:
group_common_attn_metadata.block_table_tensor = block_table[: group_common_attn_metadata.num_reqs]
slot_mapping = self._per_group_slot_mappings.get(gid)
if slot_mapping is None:
return group_common_attn_metadata
slot_mapping_buffer = self._slot_mapping_buffer_for_group(attn_group)
slot_mapping_len = slot_mapping.shape[0]
if slot_mapping.data_ptr() != slot_mapping_buffer.data_ptr():
slot_mapping_buffer[:slot_mapping_len].copy_(slot_mapping.to(torch.int32))
slot_mapping_buffer[slot_mapping_len:].fill_(PADDING_SLOT_ID)
slot_mapping_view = slot_mapping_buffer[: group_common_attn_metadata.num_actual_tokens]
self._per_group_slot_mappings[gid] = slot_mapping_view
group_common_attn_metadata.slot_mapping = slot_mapping_view
return group_common_attn_metadata
def _build_step_attn_metadatas(
self,
common_attn_metadata: CommonAttentionMetadata,
*,
graph_capture: bool = False,
) -> tuple[list[Any], list[dict[str, Any]]]:
"""Build base-proposer-style per-step metadata for Step3.5 MTP.
The full-window path still reuses the same logical window for every MTP
layer, but graph-param update/replay expects draft metadata to be
indexed by speculative step. Each Step3.5 MTP attention group maps to
one draft step/KV-cache group, so return:
``[{layer0: meta0}, {layer1: meta1}, {layer2: meta2}]``
instead of one dict containing all MTP layers.
"""
per_group_attn_metadata: list[Any] = []
multi_steps_attn_metadata: list[dict[str, Any]] = []
extra_attn_metadata_args: dict[str, Any] = {}
if self.use_compress:
extra_attn_metadata_args = dict(
prefill_ratio_to_sas_metadata=dict(),
decode_ratio_to_sas_metadata=dict(),
common_ratio_to_sas_metadata=dict(),
block_size=self.draft_attn_groups[0].kv_cache_spec.block_size,
)
for attn_group in self.draft_attn_groups:
group_common_attn_metadata = self._common_attn_metadata_for_group(common_attn_metadata, attn_group)
builder = attn_group.get_metadata_builder()
if graph_capture:
attn_metadata = builder.build_for_graph_capture(
group_common_attn_metadata,
AscendAttentionState.SpecDecoding,
)
else:
attn_metadata = builder.build(
0,
group_common_attn_metadata,
self.runner.get_model(),
**extra_attn_metadata_args,
)
if hasattr(attn_metadata, "causal") and not attn_metadata.causal:
attn_metadata.attn_mask = None
per_group_attn_metadata.append(attn_metadata)
per_step_attn_metadata: dict[str, Any] = {}
for layer_name in attn_group.layer_names:
per_step_attn_metadata[layer_name] = attn_metadata
multi_steps_attn_metadata.append(per_step_attn_metadata)
return per_group_attn_metadata, multi_steps_attn_metadata
def _sample_draft_tokens_for_step(
self,
hidden_states: torch.Tensor,
sampling_metadata: SamplingMetadata,
spec_step_idx: int,
num_indices: int,
) -> tuple[torch.Tensor, torch.Tensor | None]:
"""GPU Step3.5 sampling semantics with Ascend TP/reduce-sample paths."""
logits: torch.Tensor | None = None
if get_ascend_config().enable_reduce_sample and self.method == "mtp":
if not hasattr(self.model.model, "compute_logits"):
draft_token_ids = self.compute_draft_token_ids(hidden_states)
if lmhead_tp_enable() and num_indices < draft_token_ids.shape[0]:
draft_token_ids = draft_token_ids[:num_indices]
return draft_token_ids, None
logits = self.model.compute_logits(hidden_states, spec_step_idx=spec_step_idx)
if lmhead_tp_enable():
logits = get_lmhead_tp_group().all_to_all(logits)
else:
logits = self.model.model.logits_processor._gather_logits(logits)
else:
logits = self.model.compute_logits(hidden_states, spec_step_idx=spec_step_idx)
if lmhead_tp_enable() and num_indices < logits.shape[0]:
logits = logits[:num_indices]
if not self._enable_probabilistic_draft_probs or sampling_metadata.all_greedy:
return logits.argmax(dim=-1), None
return compute_probs_and_sample_next_token(logits, sampling_metadata)
@torch.inference_mode()
def dummy_run(
self,
num_tokens: int,
with_prefill: bool = False,
in_graph_capturing: bool = False,
num_reqs: int = 0,
num_tokens_across_dp: torch.Tensor | None = None,
aclgraph_runtime_mode: CUDAGraphMode = CUDAGraphMode.NONE,
batch_descriptor=None,
dummy_compute_logits=lambda hidden_states: None,
is_profile=False,
):
(
num_tokens,
num_tokens_across_dp,
_,
) = self.runner._sync_metadata_across_dp(num_tokens, is_draft_model=True)
multi_steps_attn_metadata: list[dict[str, Any]] = []
if not self.use_cuda_graph:
aclgraph_runtime_mode = CUDAGraphMode.NONE
if (
self.pcp_size * self.dcp_size > 1
and self.use_cuda_graph
and not is_profile
and self.block_table_tensor_clone is None
):
self.block_table_tensor_clone = torch.zeros(
(
self.runner.max_num_tokens + 2 * self.pcp_size * self.runner.max_num_reqs,
self.runner.input_batch.block_table[0].get_device_tensor().shape[1],
),
dtype=torch.int32,
device=self.device,
pin_memory=self.runner.pin_memory,
)
batch_size = max(num_tokens // (self.num_speculative_tokens + 1), 1)
if is_profile:
batch_size = min(batch_size, self.runner.max_num_reqs)
if aclgraph_runtime_mode == CUDAGraphMode.FULL and len(self.runner.attn_groups) > 0:
num_computed_tokens_cpu = self.runner.input_batch.num_computed_tokens_cpu_tensor[:num_reqs]
self.query_start_loc.cpu[: num_reqs + 1].copy_(self.runner.query_start_loc.cpu[: num_reqs + 1])
copy_snapshot_to_gpu(self.query_start_loc)
common_attn_metadata = AscendCommonAttentionMetadata(
query_start_loc=self.query_start_loc.gpu[: num_reqs + 1],
query_start_loc_cpu=self.query_start_loc.cpu[: num_reqs + 1],
seq_lens_cpu=self.runner.optimistic_seq_lens_cpu,
seq_lens_cpu_upper_bound=self.runner.optimistic_seq_lens_cpu,
seq_lens=self.runner.seq_lens[:num_reqs],
num_reqs=num_reqs,
num_actual_tokens=num_tokens,
num_input_tokens=num_tokens,
max_query_len=self.num_speculative_tokens + 1,
num_computed_tokens_cpu=num_computed_tokens_cpu,
actual_seq_lengths_q=self.runner.actual_seq_lengths_q,
block_table_tensor=self.runner.input_batch.block_table[0].get_device_tensor()[:num_reqs],
slot_mapping=self.runner.input_batch.block_table[0].slot_mapping.gpu,
positions=self.runner.positions,
attn_state=self.runner.attn_state,
decode_token_per_req=self.runner.decode_token_per_req,
max_seq_len=0,
)
if self.pcp_size * self.dcp_size > 1:
common_attn_metadata.prefill_context_parallel_metadata = self.runner.pcp_manager.long_seq_metadata
common_attn_metadata = self.shallow_copy_metadata(common_attn_metadata)
common_attn_metadata.slot_mapping = self.slot_mapping_group[0]
common_attn_metadata.seq_lens = self.seq_lens_group[0][:num_reqs]
common_attn_metadata.query_start_loc = self.query_start_loc_group[0][: num_reqs + 1]
self._seed_graph_capture_per_group_metadata(num_reqs, num_tokens)
_, multi_steps_attn_metadata = self._build_step_attn_metadatas(common_attn_metadata, graph_capture=True)
model_positions = self._get_positions(num_tokens)
if self.supports_mm_inputs:
inputs_embeds = self.model.embed_input_ids(self.input_ids[:num_tokens])
self.inputs_embeds[:num_tokens] = inputs_embeds
inputs_embeds = self.inputs_embeds[:num_tokens]
else:
inputs_embeds = None
self.token_indices_to_sample.fill_(0)
with set_ascend_forward_context(
multi_steps_attn_metadata[0] if multi_steps_attn_metadata else None,
self.vllm_config,
num_tokens=num_tokens,
num_tokens_across_dp=num_tokens_across_dp,
num_actual_tokens=0,
in_profile_run=is_profile,
batch_descriptor=batch_descriptor,
aclgraph_runtime_mode=aclgraph_runtime_mode,
is_draft_model=True,
draft_attn_metadatas=multi_steps_attn_metadata,
):
forward_context = get_forward_context()
if forward_context is not None:
forward_context.moe_layer_index = 0
self._runnable(
num_input_tokens=num_tokens,
batch_size=batch_size,
token_indices_to_sample=self.token_indices_to_sample[: batch_size * self.extra_slots_per_request],
target_positions=model_positions,
inputs_embeds=inputs_embeds,
multi_steps_attn_metadata=multi_steps_attn_metadata,
num_tokens=num_tokens,
)
forward_context = get_forward_context()
if forward_context.cudagraph_runtime_mode == CUDAGraphMode.FULL and not _EXTRA_CTX.capturing:
self._update_full_graph_params(forward_context, num_tokens, multi_steps_attn_metadata)
def _propose(
self,
target_token_ids: torch.Tensor,
target_positions: torch.Tensor,
target_hidden_states: torch.Tensor,
next_token_ids: torch.Tensor,
token_indices_to_sample: torch.Tensor | None,
common_attn_metadata: CommonAttentionMetadata,
target_model_batch_desc: BatchDescriptor,
sampling_metadata: SamplingMetadata,
mm_embed_inputs: tuple[list[torch.Tensor], torch.Tensor] | None = None,
req_scheduled_tokens=None,
long_seq_metadata=None,
num_prefill_reqs=0,
num_decode_reqs=0,
scheduler_output: SchedulerOutput = None,
num_scheduled_tokens: int = 0,
num_rejected_tokens_gpu: torch.Tensor | None = None,
) -> torch.Tensor:
self._last_draft_probs = None
batch_size = common_attn_metadata.batch_size()
if token_indices_to_sample is None:
token_indices_to_sample = common_attn_metadata.query_start_loc[1:] - 1
num_tokens, token_indices_to_sample, common_attn_metadata, long_seq_args = self.set_inputs_first_pass(
target_token_ids=target_token_ids,
next_token_ids=next_token_ids,
target_positions=target_positions,
target_hidden_states=target_hidden_states,
token_indices_to_sample=token_indices_to_sample,
cad=common_attn_metadata,
num_rejected_tokens_gpu=num_rejected_tokens_gpu,
req_scheduled_tokens=req_scheduled_tokens,
long_seq_metadata=long_seq_metadata,
num_prefill_reqs=num_prefill_reqs,
num_decode_reqs=num_decode_reqs,
)
if self.pcp_size * self.dcp_size > 1:
assert long_seq_args is not None
assert self.runner is not None
has_lora = len(self.runner.input_batch.lora_id_to_lora_request) > 0
uniform_decode = target_model_batch_desc.uniform
if self.use_cuda_graph:
_, batch_descriptor = self.runner.cudagraph_dispatcher.dispatch(
num_tokens=num_tokens,
uniform_decode=uniform_decode,
has_lora=has_lora,
)
num_input_tokens = batch_descriptor.num_tokens
else:
num_input_tokens = num_tokens
(
num_input_tokens,
num_tokens_across_dp,
_,
) = self.runner._sync_metadata_across_dp(num_input_tokens, is_draft_model=True)
if self.use_cuda_graph:
aclgraph_runtime_mode, batch_descriptor = self.runner.cudagraph_dispatcher.dispatch(
num_tokens=num_input_tokens,
uniform_decode=uniform_decode,
has_lora=has_lora,
)
num_input_tokens = batch_descriptor.num_tokens
else:
aclgraph_runtime_mode = CUDAGraphMode.NONE
batch_descriptor = None
if aclgraph_runtime_mode == CUDAGraphMode.FULL:
num_reqs_padded = self.runner._pad_query_start_loc_for_fia(
num_input_tokens,
batch_descriptor.num_reqs if batch_descriptor.num_reqs is not None else common_attn_metadata.num_reqs,
common_attn_metadata.num_reqs,
aclgraph_runtime_mode,
batch_descriptor.num_reqs,
)
common_attn_metadata.num_reqs = num_reqs_padded
common_attn_metadata.query_start_loc = self.runner.query_start_loc.gpu[: num_reqs_padded + 1]
common_attn_metadata.query_start_loc_cpu = self.runner.query_start_loc.cpu[: num_reqs_padded + 1]
slicing_length = (
num_reqs_padded * self.decode_threshold if self.pcp_size * self.dcp_size > 1 else num_reqs_padded
)
common_attn_metadata.block_table_tensor = self._adjust_tensor(
common_attn_metadata.block_table_tensor, slicing_length
)
common_attn_metadata.seq_lens = self._adjust_tensor(self.runner.seq_lens, num_reqs_padded)
common_attn_metadata.seq_lens_cpu = self._adjust_tensor(
self.runner.optimistic_seq_lens_cpu, num_reqs_padded
)
if common_attn_metadata._seq_lens_cpu is not None:
common_attn_metadata._seq_lens_cpu = common_attn_metadata.seq_lens_cpu.clone()
if common_attn_metadata.num_computed_tokens_cpu is not None:
common_attn_metadata.num_computed_tokens_cpu = self._adjust_tensor(
common_attn_metadata.num_computed_tokens_cpu, num_reqs_padded
)
if self.pcp_size > 1:
pcp_allgather_restore_idx = (
common_attn_metadata.prefill_context_parallel_metadata.pcp_allgather_restore_idx
)
index = torch.arange(
pcp_allgather_restore_idx.shape[0],
device=pcp_allgather_restore_idx.device,
)
mask = (index % (self.pcp_size * self.decode_threshold)) >= self.decode_threshold
pcp_allgather_restore_idx[mask] = 0
self.runner.pcp_manager.pcp_allgather_restore_idx.gpu[: pcp_allgather_restore_idx.shape[0]] = (
pcp_allgather_restore_idx
)
self.runner.pcp_manager.pcp_allgather_restore_idx.gpu[pcp_allgather_restore_idx.shape[0] :] = 0
else:
num_reqs_padded = common_attn_metadata.num_reqs
if not self.vllm_config.model_config.use_mla and self.pcp_size * self.dcp_size == 1:
common_attn_metadata.block_table_tensor = self._adjust_tensor(
common_attn_metadata.block_table_tensor, num_reqs_padded
)
if self.supports_mm_inputs:
inputs_embeds = self.model.embed_input_ids(self.input_ids[:num_tokens])
self.inputs_embeds[:num_tokens] = inputs_embeds
inputs_embeds = self.inputs_embeds[:num_input_tokens]
else:
inputs_embeds = None
slot_mapping_len = common_attn_metadata.slot_mapping.shape[0]
self.slot_mapping_group[0][:slot_mapping_len].copy_(common_attn_metadata.slot_mapping)
self.slot_mapping_group[0][slot_mapping_len:].fill_(PADDING_SLOT_ID)
common_attn_metadata.slot_mapping = self.slot_mapping_group[0]
self._per_group_slot_mappings[self.kv_cache_gid] = self.slot_mapping_group[0][
: common_attn_metadata.num_actual_tokens
]
self.seq_lens_group[0][:num_reqs_padded].copy_(common_attn_metadata.seq_lens)
self.seq_lens_group[0][num_reqs_padded:].fill_(0)
common_attn_metadata.seq_lens = self.seq_lens_group[0][:num_reqs_padded]
self.query_start_loc_group[0][: num_reqs_padded + 1].copy_(common_attn_metadata.query_start_loc)
self.query_start_loc_group[0][num_reqs_padded + 1 :].fill_(0)
common_attn_metadata.query_start_loc = self.query_start_loc_group[0][: num_reqs_padded + 1]
common_attn_metadata.num_input_tokens = num_input_tokens
_, multi_steps_attn_metadata = self._build_step_attn_metadatas(common_attn_metadata)
attn_metadata_i = next(iter(multi_steps_attn_metadata[0].values()))
if not self.use_cuda_graph:
common_attn_metadata.block_table_tensor = common_attn_metadata.block_table_tensor.clone()
token_indices_to_sample_len = token_indices_to_sample.shape[0]
self.token_indices_to_sample[:token_indices_to_sample_len].copy_(token_indices_to_sample)
self.token_indices_to_sample[token_indices_to_sample_len:].fill_(0)
with set_ascend_forward_context(
multi_steps_attn_metadata[0],
self.vllm_config,
num_tokens=num_input_tokens,
num_tokens_across_dp=num_tokens_across_dp,
num_actual_tokens=num_tokens,
batch_descriptor=batch_descriptor,
aclgraph_runtime_mode=aclgraph_runtime_mode,
is_draft_model=True,
draft_attn_metadatas=multi_steps_attn_metadata,
):
forward_context = get_forward_context()
if forward_context is not None:
forward_context.moe_layer_index = 0
model_inputs: dict[str, Any] = {
"num_input_tokens": num_input_tokens,
"batch_size": batch_size,
"token_indices_to_sample": self.token_indices_to_sample[:token_indices_to_sample_len],
"target_positions": target_positions,
"inputs_embeds": inputs_embeds,
"multi_steps_attn_metadata": multi_steps_attn_metadata,
"num_tokens": num_tokens,
"is_prefill": attn_metadata_i.num_prefills,
}
run_draft = partial(self._runnable, **model_inputs)
if self.enable_enpu:
self._update_full_graph_params_if_needed(forward_context, num_input_tokens, multi_steps_attn_metadata)
draft_token_ids = run_draft()
else:
draft_token_ids = run_draft()
self._update_full_graph_params_if_needed(forward_context, num_input_tokens, multi_steps_attn_metadata)
return draft_token_ids
def _run_merged_draft(
self,
num_input_tokens,
batch_size,
token_indices_to_sample,
target_positions,
inputs_embeds,
multi_steps_attn_metadata,
num_tokens,
is_prefill=None,
) -> torch.Tensor:
"""Base MTP execution flow with Step3.5 step-aware layer/head selection."""
self._last_draft_probs = None
sampling_metadata = self.runner.input_batch.sampling_metadata
model_input_ids = self.input_ids[:num_input_tokens]
model_positions = self._get_positions(num_input_tokens)
model_kwargs = {
"input_ids": model_input_ids,
"positions": model_positions,
"inputs_embeds": inputs_embeds,
"spec_step_idx": 0,
}
if self.pass_hidden_states_to_model:
model_hidden_states = self.hidden_states[:num_input_tokens]
model_hidden_states, model_positions = self.maybe_pad_and_reduce(model_hidden_states, model_positions)
model_kwargs["hidden_states"] = model_hidden_states
model_kwargs["positions"] = model_positions
ret_hidden_states = self.model(**model_kwargs)
if not self.model_returns_tuple():
last_hidden_states = ret_hidden_states
hidden_states = last_hidden_states
else:
last_hidden_states, hidden_states = ret_hidden_states
last_hidden_states, model_positions, hidden_states = self.maybe_all_gather_and_unpad(
last_hidden_states, model_positions, hidden_states
)
num_indices = token_indices_to_sample.shape[0]
if lmhead_tp_enable():
max_num_reqs_across_dp = (
self.vllm_config.scheduler_config.max_num_seqs * self.runner.uniform_decode_query_len
)
token_indices_to_sample = torch.nn.functional.pad(
token_indices_to_sample, (0, max_num_reqs_across_dp - num_indices)
)
sample_hidden_states = last_hidden_states[token_indices_to_sample]
draft_token_ids, draft_probs = self._sample_draft_tokens_for_step(
sample_hidden_states,
sampling_metadata,
spec_step_idx=0,
num_indices=num_indices,
)
if self.num_speculative_tokens == 1 or self.parallel_drafting:
if draft_probs is not None:
self._last_draft_probs = draft_probs.view(
-1, self.num_speculative_tokens, draft_probs.shape[-1]
).contiguous()
return draft_token_ids.view(-1, self.num_speculative_tokens)
if self.pcp_size * self.dcp_size > 1 and is_prefill:
draft_token_ids_list = [draft_token_ids for _ in range(self.num_speculative_tokens)]
return torch.stack(draft_token_ids_list, dim=1)
return self._run_window_draft_steps(
first_draft_token_ids=draft_token_ids,
first_draft_probs=draft_probs,
first_hidden_states=hidden_states,
num_input_tokens=num_input_tokens,
batch_size=batch_size,
token_indices_to_sample=token_indices_to_sample,
target_positions=target_positions,
num_tokens=num_tokens,
multi_steps_attn_metadata=multi_steps_attn_metadata,
inputs_embeds=inputs_embeds,
sampling_metadata=sampling_metadata,
)
def _roll_window_inputs_only(
self,
*,
prev_token_ids: torch.Tensor,
next_token_ids: torch.Tensor,
previous_hidden_states: torch.Tensor,
token_indices_to_sample: torch.Tensor,
num_tokens: int,
input_batch_size: int,
) -> torch.Tensor | None:
self.input_ids[: num_tokens - 1] = prev_token_ids[1:]
self.input_ids[token_indices_to_sample] = next_token_ids
self.hidden_states[:num_tokens] = previous_hidden_states.view(num_tokens, -1)
if self.supports_mm_inputs:
self.inputs_embeds[:num_tokens] = self.model.embed_input_ids(self.input_ids[:num_tokens])
return self.inputs_embeds[:input_batch_size]
return None
def _run_window_draft_steps(
self,
*,
first_draft_token_ids: torch.Tensor,
first_draft_probs: torch.Tensor | None,
first_hidden_states: torch.Tensor,
num_input_tokens: int,
batch_size: int,
token_indices_to_sample: torch.Tensor,
target_positions: torch.Tensor,
num_tokens: int,
multi_steps_attn_metadata,
inputs_embeds,
sampling_metadata: SamplingMetadata,
) -> torch.Tensor:
"""Run the Step3.5 MTP full-window layer chain.
This is the canonical Step3.5 MTP draft path, matching the full-window
multi-layer proposer loop: every MTP layer reprocesses the whole draft
window so each independent KV-cache group is seeded over the full
window. Attention metadata is built once before this loop; layer-to-layer
state transition only rolls the model input buffers by shifting the
previous window and placing the newly drafted token at
token_indices_to_sample.
"""
draft_probs_list = None if first_draft_probs is None else [first_draft_probs]
draft_token_ids_list = [first_draft_token_ids]
full_hidden_states = first_hidden_states[:num_tokens]
input_batch_size = num_input_tokens
forward_context = get_forward_context()
_EXTRA_CTX.num_tokens = input_batch_size
_EXTRA_CTX.num_accept_tokens = batch_size
for draft_step in range(self.num_speculative_tokens - 1):
forward_context = get_forward_context()
if forward_context is not None:
forward_context.moe_layer_index = 0
spec_step_idx = draft_step + 1
prev_token_ids = self.input_ids[:num_tokens].clone()
next_token_ids = draft_token_ids_list[-1].int()
inputs_embeds = self._roll_window_inputs_only(
prev_token_ids=prev_token_ids,
next_token_ids=next_token_ids,
previous_hidden_states=full_hidden_states,
token_indices_to_sample=token_indices_to_sample,
num_tokens=num_tokens,
input_batch_size=input_batch_size,
)
model_input_ids = self.input_ids[:input_batch_size]
model_positions = self._get_positions(input_batch_size)
model_hidden_states = self.hidden_states[:input_batch_size]
model_hidden_states, model_positions = self.maybe_pad_and_reduce(
model_hidden_states,
model_positions,
)
if forward_context is not None and multi_steps_attn_metadata:
if spec_step_idx >= len(multi_steps_attn_metadata):
raise AssertionError("Step3.5 MTP metadata must contain one entry per draft step")
forward_context.attn_metadata = multi_steps_attn_metadata[spec_step_idx]
model_kwargs = {
"input_ids": model_input_ids,
"positions": model_positions,
"inputs_embeds": inputs_embeds,
"spec_step_idx": spec_step_idx,
}
if self.pass_hidden_states_to_model:
model_kwargs["hidden_states"] = model_hidden_states
ret_hidden_states = self.model(**model_kwargs)
if not self.model_returns_tuple():
last_hidden_states = ret_hidden_states
hidden_states = ret_hidden_states
else:
last_hidden_states, hidden_states = ret_hidden_states
last_hidden_states, model_positions, hidden_states = self.maybe_all_gather_and_unpad(
last_hidden_states,
model_positions,
hidden_states,
)
num_indices = token_indices_to_sample.shape[0]
sample_hidden_states = last_hidden_states[token_indices_to_sample]
draft_token_ids, draft_probs = self._sample_draft_tokens_for_step(
sample_hidden_states,
sampling_metadata,
spec_step_idx=spec_step_idx,
num_indices=num_indices,
)
if draft_probs is not None:
assert draft_probs_list is not None
draft_probs_list.append(draft_probs)
full_hidden_states = hidden_states[:num_tokens]
draft_token_ids_list.append(draft_token_ids)
draft_token_ids = torch.stack(draft_token_ids_list, dim=1)
if draft_probs_list is not None:
self._last_draft_probs = torch.stack(draft_probs_list, dim=1).contiguous()
return draft_token_ids
# -- overrides matching the GPU Step3p5MTPProposer ----------------------
def _ensure_draft_layer_types_cover_mtp_layers(self) -> None:
hf_config = self.draft_model_config.hf_config
layer_types = getattr(hf_config, "layer_types", None)
num_hidden_layers = getattr(hf_config, "num_hidden_layers", None)
num_mtp_layers = getattr(hf_config, "num_nextn_predict_layers", 0)
if layer_types is None or num_hidden_layers is None or num_mtp_layers is None:
return
needed_num_layer_types = num_hidden_layers + num_mtp_layers
if len(layer_types) >= needed_num_layer_types:
return
hf_config.layer_types = list(layer_types) + ["sliding_attention"] * (needed_num_layer_types - len(layer_types))
def _create_draft_vllm_config(self) -> VllmConfig:
base = super()._create_draft_vllm_config()
# Transformers validates Step3.5 layer_types against num_hidden_layers
# and may leave only the base-layer entries. The MTP draft model builds
# layers at num_hidden_layers + k, so make the draft config cover those
# layer indices before the model is constructed.
self._ensure_draft_layer_types_cover_mtp_layers()
# Ascend ModelSlim keeps quant metadata in quant_model_description.json,
# not in config.json's quantization_config, so get_quant_config falls
# through to the hf_overrides branch. The verifier model_config has
# hf_overrides={}, but the draft_model_config leaves it None, which that
# branch rejects. Normalize to {} so the draft takes the verifier's path.
if not isinstance(getattr(self.draft_model_config, "hf_overrides", None), dict):
self.draft_model_config.hf_overrides = {}
return replace(
base,
model_config=self.draft_model_config,
quant_config=get_draft_quant_config(base),
)
def validate_same_kv_cache_group(self, kv_cache_config: KVCacheConfig) -> None:
"""Step3.5 MTP draft layers may span multiple KV cache groups."""
return
def initialize_attn_backend(
self,
kv_cache_config: KVCacheConfig,
kernel_block_sizes: list[int] | None = None,
) -> None:
all_attn_layers = get_layers_from_vllm_config(
self.vllm_config,
AttentionLayerBase, # type: ignore[type-abstract]
)
layer_to_gid: dict[str, int] = {}
layer_to_spec: dict[str, KVCacheSpec] = {}
for gid, group in enumerate(kv_cache_config.kv_cache_groups):
group_spec = group.kv_cache_spec
for layer_name in group.layer_names:
layer_to_gid[layer_name] = gid
if isinstance(group_spec, UniformTypeKVCacheSpecs):
if layer_name in group_spec.kv_cache_specs:
layer_to_spec[layer_name] = group_spec.kv_cache_specs[layer_name]
else:
target_layer_name = getattr(
all_attn_layers.get(layer_name),
"kv_sharing_target_layer_name",
None,
)
if target_layer_name and target_layer_name in group_spec.kv_cache_specs:
layer_to_spec[layer_name] = group_spec.kv_cache_specs[target_layer_name]
else:
layer_to_spec[layer_name] = group_spec
else:
layer_to_spec[layer_name] = group_spec
attention_groups: dict[tuple[str, int], AttentionGroup] = {}
for layer_name in sorted(self._draft_attn_layer_names):
if layer_name not in layer_to_spec:
continue
attn_layer = all_attn_layers[layer_name]
attn_backend = attn_layer.get_attn_backend()
spec = layer_to_spec[layer_name]
gid = layer_to_gid[layer_name]
group_key = (attn_backend.full_cls_name(), gid)
if group_key not in attention_groups:
kernel_block_size = (
kernel_block_sizes[gid]
if kernel_block_sizes is not None and gid < len(kernel_block_sizes)
else None
)
attn_group = AttentionGroup(
backend=attn_backend,
layer_names=[layer_name],
kv_cache_spec=spec,
kv_cache_group_id=gid,
)
attn_group.create_metadata_builders(
self.vllm_config,
self.device,
kernel_block_size=kernel_block_size,
)
attention_groups[group_key] = attn_group
else:
attention_groups[group_key].layer_names.append(layer_name)
self.draft_attn_groups = list(attention_groups.values())
if self.draft_attn_groups:
self.kv_cache_gid = self.draft_attn_groups[0].kv_cache_group_id
self.block_size = self.draft_attn_groups[0].get_metadata_builder().kv_cache_spec.block_size
else:
self.kv_cache_gid = 0
self.block_size = kv_cache_config.kv_cache_groups[0].kv_cache_spec.block_size

View File

@@ -0,0 +1,42 @@
import torch
from vllm.v1.spec_decode.suffix_decoding import SuffixDecodingProposer
from vllm_ascend.utils import vllm_version_is
class AscendSuffixDecodingProposer(SuffixDecodingProposer):
def __init__(self, vllm_config, runner):
super().__init__(vllm_config)
self.runner = runner
def dummy_run(
self,
num_tokens,
with_prefill=None,
in_graph_capturing=None,
num_reqs=None,
num_tokens_across_dp=None,
aclgraph_runtime_mode=None,
batch_descriptor=None,
dummy_compute_logits=lambda hidden_states: None,
is_profile=False,
):
pass
def propose(
self,
sampled_token_ids: list[list[int]],
num_tokens_no_spec=None,
token_ids_cpu=None,
num_speculative_tokens: int = 0,
slot_mappings: dict[str, torch.Tensor] | list[dict[str, torch.Tensor]] | None = None,
):
if vllm_version_is("0.23.0"):
return super().propose(self.runner.input_batch, sampled_token_ids)
else:
return super().propose(
num_speculative_tokens,
self.runner.input_batch,
sampled_token_ids,
slot_mappings,
)

View File

@@ -0,0 +1,73 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
import numpy as np
import torch
def update_num_computed_tokens_for_batch_change(
num_computed_tokens: torch.Tensor,
num_accepted_tokens: torch.Tensor,
prev_positions: torch.Tensor,
valid_sampled_token_count: torch.Tensor,
prev_num_draft_tokens: torch.Tensor,
cpu_num_computed_tokens: torch.Tensor,
) -> None:
"""Correct num_computed_tokens for async spec decode drift.
Requests that had drafts: corrected = prev_gpu + valid_count.
New requests or non-draft (e.g. prefills): use CPU value directly.
"""
# Clamp because prev_positions can be -1 for new requests
gather_indices = prev_positions.clamp(min=0)
valid_counts = valid_sampled_token_count[gather_indices]
prev_computed = num_computed_tokens[gather_indices]
prev_drafts = prev_num_draft_tokens[gather_indices]
participating = (prev_positions >= 0) & (prev_drafts > 0)
corrected = prev_computed + valid_counts.int()
n = prev_positions.shape[0]
num_computed_tokens[:n].copy_(torch.where(participating, corrected, cpu_num_computed_tokens))
num_accepted_tokens.copy_(torch.where(participating, valid_counts, num_accepted_tokens))
def correct_optimistic_seq_lens_cpu(
optimistic_seq_lens_cpu_np: np.ndarray,
prev_positions_np: np.ndarray,
prev_num_draft_tokens_np: np.ndarray,
valid_sampled_token_count_np: np.ndarray,
num_reqs: int,
) -> None:
"""Correct ``optimistic_seq_lens_cpu`` for async spec decode drift.
The scheduler optimistically advances ``num_computed_tokens_cpu`` by the
full number of tokens scheduled in the previous step (``prev_drafts + 1``
per spec-decode request), assuming all drafts were accepted. The actual
number of valid sampled tokens is ``valid_count = 1 + accepted_drafts``.
The drift, equal to the number of rejected tokens, is therefore::
rejected = prev_drafts + 1 - valid_count
Subtracting this from the optimistic seq_lens recovers the true seq_lens
that ``self.seq_lens`` (GPU) carries for participating requests, without
touching the device. New requests (``prev_positions < 0``) and prefills
(``prev_drafts == 0``) need no correction.
Mirrors ``update_num_computed_tokens_for_batch_change`` on the CPU side.
All arrays are sliced to ``num_reqs``; ``optimistic_seq_lens_cpu_np`` is
modified in place.
"""
prev_positions = prev_positions_np[:num_reqs]
# Clamp negative entries (new requests) to 0; the participating mask zeroes
# out their correction so the gathered values are don't-care.
gather_indices = np.maximum(prev_positions, 0)
prev_drafts = prev_num_draft_tokens_np[gather_indices]
valid_counts = valid_sampled_token_count_np[gather_indices]
participating = (prev_positions >= 0) & (prev_drafts > 0)
# rejected_for_participating == correction; non-participating reqs end up
# at zero via the mask multiply.
correction = (prev_drafts + 1 - valid_counts) * participating
optimistic_seq_lens_cpu_np[:num_reqs] -= correction.astype(optimistic_seq_lens_cpu_np.dtype, copy=False)