@@ -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}")
|
||||
|
||||
270
vllm_ascend/spec_decode/dflash_proposer.py
Normal file
270
vllm_ascend/spec_decode/dflash_proposer.py
Normal 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
|
||||
17
vllm_ascend/spec_decode/draft_proposer.py
Normal file
17
vllm_ascend/spec_decode/draft_proposer.py
Normal 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()
|
||||
@@ -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)
|
||||
|
||||
198
vllm_ascend/spec_decode/extract_hidden_states_proposer.py
Normal file
198
vllm_ascend/spec_decode/extract_hidden_states_proposer.py
Normal 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
|
||||
2095
vllm_ascend/spec_decode/llm_base_proposer.py
Normal file
2095
vllm_ascend/spec_decode/llm_base_proposer.py
Normal file
File diff suppressed because it is too large
Load Diff
70
vllm_ascend/spec_decode/medusa_proposer.py
Normal file
70
vllm_ascend/spec_decode/medusa_proposer.py
Normal 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
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
35
vllm_ascend/spec_decode/ngram_proposer_npu.py
Normal file
35
vllm_ascend/spec_decode/ngram_proposer_npu.py
Normal 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
|
||||
830
vllm_ascend/spec_decode/step3p5.py
Normal file
830
vllm_ascend/spec_decode/step3p5.py
Normal 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
|
||||
42
vllm_ascend/spec_decode/suffix_proposer.py
Normal file
42
vllm_ascend/spec_decode/suffix_proposer.py
Normal 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,
|
||||
)
|
||||
73
vllm_ascend/spec_decode/utils.py
Normal file
73
vllm_ascend/spec_decode/utils.py
Normal 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)
|
||||
Reference in New Issue
Block a user