init v0.23.0

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

View File

@@ -0,0 +1,36 @@
# Adapt from https://github.com/vllm-project/vllm/blob/main/vllm/v1/worker/gpu/sample/spec_decode/__init__.py
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
#
# 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.
# This file is a part of the vllm-ascend project.
#
import torch
from vllm.config import VllmConfig
def init_speculator(
vllm_config: VllmConfig,
device: torch.device,
):
"""Override GPU init_speculator for Ascend NPUs.
Use AscendEagleSpeculator when eagle is used.
"""
speculative_config = vllm_config.speculative_config
assert speculative_config is not None
if speculative_config.use_eagle():
from vllm_ascend.worker.v2.spec_decode.eagle.speculator import AscendEagleSpeculator
return AscendEagleSpeculator(vllm_config, device)
raise NotImplementedError(f"{speculative_config.method} is not supported yet.")

View File

@@ -0,0 +1,231 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM-Ascend project
from collections.abc import Callable
from contextlib import contextmanager
from typing import Any
import torch
from vllm.config import VllmConfig
from vllm.config.compilation import CUDAGraphMode
from vllm.forward_context import get_forward_context, set_forward_context
from vllm.logger import logger
from vllm.v1.kv_cache_interface import KVCacheConfig
from vllm.v1.worker.gpu.block_table import BlockTables
from vllm.v1.worker.gpu.cudagraph_utils import ( # type: ignore[import-not-found]
AttentionStatePair,
BatchExecutionDescriptor,
)
from vllm.v1.worker.gpu.input_batch import InputBuffers
from vllm.v1.worker.gpu.model_states.interface import ModelState
from vllm.v1.worker.gpu.spec_decode.autoregressive.cudagraph_utils import ( # type: ignore[import-not-found]
DecodeSpeculatorCudaGraphManager as DecodeEagleCudaGraphManager,
)
from vllm.v1.worker.gpu.spec_decode.autoregressive.cudagraph_utils import (
PrefillSpeculatorCudaGraphManager as PrefillEagleCudaGraphManager,
)
from vllm.v1.worker.utils import AttentionGroup
from vllm_ascend.ascend_forward_context import _EXTRA_CTX
from vllm_ascend.compilation.acl_graph import (
set_draft_graph_params,
set_draft_graph_prefill_params,
update_full_graph_params,
)
from vllm_ascend.worker.v2.aclgraph_utils import ModelWithContext
from vllm_ascend.worker.v2.utils import communicator_switch
class PrefillEagleAclGraphManager(PrefillEagleCudaGraphManager):
"""AclGraphManager for Eagle speculative decoding."""
def __init__(
self,
vllm_config: VllmConfig,
device: torch.device,
cudagraph_mode: CUDAGraphMode,
decode_query_len: int,
speculator: Any,
):
super().__init__(vllm_config, device, cudagraph_mode, decode_query_len)
# set speculator attribute, so we can access attributes speculator
# when call `run_fullgraph` method in CudaGraphManager,
# then we don't need to # copy `propose` method in `AscendEagleSpeculator` class.
self.speculator = speculator
# capture_sizes sorts in ascending order.
self.capture_sizes = sorted(self.compilation_config.cudagraph_capture_sizes)
# vllm-ascend need to update draft graph params of attention backend.
# so we need to set draft graph params before capture full graph.
# `prefill` graph and `decodes` graph are different, `decode_query_len` can be used to distinguish them
self.is_draft_model_prefill = decode_query_len > 1
if super().needs_capture():
if self.is_draft_model_prefill:
set_draft_graph_prefill_params(self.capture_sizes)
else:
set_draft_graph_params(self.capture_sizes)
def capture(
self,
forward_fn: Callable,
attn_states: dict[BatchExecutionDescriptor, AttentionStatePair],
progress_bar_desc: str = "Capturing CUDA graphs",
) -> None:
"""Capture ACL graphs for Eagle."""
with communicator_switch(), model_capture_wrapper(self.speculator, self.is_draft_model_prefill):
super().capture(
forward_fn,
attn_states,
progress_bar_desc=progress_bar_desc,
)
def run_fullgraph(self, desc: BatchExecutionDescriptor) -> torch.Tensor | tuple[torch.Tensor, list[torch.Tensor]]:
"""Override run_fullgraph to update full graph params in run_fullgraph."""
num_tokens = desc.num_tokens
if self.is_draft_model_prefill:
logger.info_once("PrefillEagleAclGraphManager: draft prefill run_fullgraph with num_tokens=%s", num_tokens)
else:
logger.info_once("DecodeEagleAclGraphManager: draft run_fullgraph with num_tokens=%s", num_tokens)
draft_attn_metadatas = self.speculator.build_draft_attn_metadatas(desc.num_reqs, self.is_draft_model_prefill)
ret = super().run_fullgraph(desc)
positions = self.speculator.input_buffers.positions[:num_tokens]
# refer to vllm.v1.worker.gpu.dp_utils.sync_cudagraph_and_dp_padding to
# calculate num_tokens_across_dp.
num_tokens_across_dp = torch.full([self.speculator.dp_size], num_tokens, device=self.device)
with set_forward_context(
self.speculator.model_state.attn_metadata,
self.vllm_config,
num_tokens=num_tokens,
cudagraph_runtime_mode=desc.cg_mode,
num_tokens_across_dp=num_tokens_across_dp,
batch_descriptor=None, # Full graph model don't need batch_descriptor
slot_mapping=None,
):
# decide to update draft graph params
_EXTRA_CTX.is_draft_model = True
# decide to run `prefill` graph or `decodes` graph
_EXTRA_CTX.is_draft_model_prefill = self.is_draft_model_prefill
forward_context = get_forward_context()
update_full_graph_params(
# FIXME(Ronald1995): support hybrid attn backend
list(self.speculator.attn_backends.values())[0],
self.speculator.update_stream,
forward_context,
num_tokens,
self.vllm_config,
self.speculator.speculative_config,
positions.shape[0],
draft_attn_metadatas=draft_attn_metadatas,
)
return ret
class DecodeEagleAclGraphManager(DecodeEagleCudaGraphManager):
"""AclGraphManager for Eagle speculative decoding."""
def __init__(
self,
vllm_config: VllmConfig,
device: torch.device,
cudagraph_mode: CUDAGraphMode,
decode_query_len: int,
speculator: Any,
):
super().__init__(vllm_config, device, cudagraph_mode, decode_query_len)
# set speculator attribute, so we can access attributes speculator
# when call `run_fullgraph` method in CudaGraphManager,
# then we don't need to # copy `propose` method in `AscendEagleSpeculator` class.
self.speculator = speculator
# capture_sizes sorts in ascending order.
self.capture_sizes = sorted(self.compilation_config.cudagraph_capture_sizes)
# vllm-ascend need to update draft graph params of attention backend.
# so we need to set draft graph params before capture full graph.
# `prefill` graph and `decodes` graph are different, `decode_query_len` can be used to distinguish them
self.is_draft_model_prefill = decode_query_len > 1
if super().needs_capture():
if self.is_draft_model_prefill:
set_draft_graph_prefill_params(self.capture_sizes)
else:
set_draft_graph_params(self.capture_sizes)
def capture(
self,
forward_fn: Callable,
model_state: ModelState,
input_buffers: InputBuffers,
block_tables: BlockTables,
attn_groups: list[list[AttentionGroup]],
kv_cache_config: KVCacheConfig,
progress_bar_desc: str = "Capturing CUDA graphs",
) -> None:
"""Capture ACL graphs for Eagle."""
with communicator_switch(), model_capture_wrapper(self.speculator, self.is_draft_model_prefill):
super().capture(
forward_fn,
model_state,
input_buffers,
block_tables,
attn_groups,
kv_cache_config,
progress_bar_desc=progress_bar_desc,
)
def run_fullgraph(self, desc: BatchExecutionDescriptor) -> torch.Tensor | tuple[torch.Tensor, list[torch.Tensor]]:
"""Override run_fullgraph to update full graph params in run_fullgraph."""
num_tokens = desc.num_tokens
if self.is_draft_model_prefill:
logger.info_once("PrefillEagleAclGraphManager: draft prefill run_fullgraph with num_tokens=%s", num_tokens)
else:
logger.info_once("DecodeEagleAclGraphManager: draft run_fullgraph with num_tokens=%s", num_tokens)
draft_attn_metadatas = self.speculator.build_draft_attn_metadatas(desc.num_reqs, self.is_draft_model_prefill)
ret = super().run_fullgraph(desc)
positions = self.speculator.input_buffers.positions[:num_tokens]
# refer to vllm.v1.worker.gpu.dp_utils.sync_cudagraph_and_dp_padding to
# calculate num_tokens_across_dp.
num_tokens_across_dp = torch.full([self.speculator.dp_size], num_tokens, device=self.device)
with set_forward_context(
self.speculator.model_state.attn_metadata,
self.vllm_config,
num_tokens=num_tokens,
cudagraph_runtime_mode=desc.cg_mode,
num_tokens_across_dp=num_tokens_across_dp,
batch_descriptor=None, # Full graph model don't need batch_descriptor
slot_mapping=None,
):
# decide to update draft graph params
_EXTRA_CTX.is_draft_model = True
# decide to run `prefill` graph or `decodes` graph
_EXTRA_CTX.is_draft_model_prefill = self.is_draft_model_prefill
forward_context = get_forward_context()
update_full_graph_params(
# FIXME(Ronald1995): support hybrid attn backend
list(self.speculator.attn_backends.values())[0],
self.speculator.update_stream,
forward_context,
num_tokens,
self.vllm_config,
self.speculator.speculative_config,
positions.shape[0],
draft_attn_metadatas=draft_attn_metadatas,
)
return ret
@contextmanager
def model_capture_wrapper(speculator, is_draft_model_prefill):
"""Context manager to override speculator's model for speculator capturing."""
try:
speculator.model = ModelWithContext(speculator.model, True, is_draft_model_prefill)
yield
finally:
speculator.model = speculator.model.get_original_model()

View File

@@ -0,0 +1,370 @@
# Adapt from https://github.com/vllm-project/vllm/blob/main/vllm/v1/worker/gpu/sample/spec_decode/eagle.py
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
#
# 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.
# This file is a part of the vllm-ascend project.
#
from contextlib import contextmanager
from copy import copy
from typing import Any, cast
import torch
import vllm
from vllm.config import VllmConfig, get_layers_from_vllm_config
from vllm.config.compilation import CUDAGraphMode
from vllm.model_executor.layers.attention_layer_base import AttentionLayerBase
from vllm.v1.attention.backend import AttentionBackend
from vllm.v1.kv_cache_interface import KVCacheConfig
from vllm.v1.worker.gpu.block_table import BlockTables
from vllm.v1.worker.gpu.input_batch import InputBatch
from vllm.v1.worker.gpu.model_states.interface import ModelState
from vllm.v1.worker.gpu.spec_decode.autoregressive import ( # type: ignore[import-not-found]
speculator as vllm_speculator_module, # type: ignore[import-not-found]
)
from vllm.v1.worker.gpu.spec_decode.autoregressive.cudagraph_utils import ( # type: ignore[import-not-found]
PrefillSpeculatorCudaGraphManager,
)
from vllm.v1.worker.gpu.spec_decode.eagle.speculator import EagleSpeculator # type: ignore[import-not-found]
from vllm_ascend.attention.attention_v1 import AscendAttentionState
from vllm_ascend.worker.v2.attn_utils import build_attn_metadata
from vllm_ascend.worker.v2.input_batch import AscendInputBuffers
from vllm_ascend.worker.v2.spec_decode.eagle.aclgraph import PrefillEagleAclGraphManager
_BUILD_ATTN_METADATA_MODULE = vllm.v1.worker.gpu.spec_decode.speculator
_PREFILL_CUDAGRAPH_MANAGER_CLS = PrefillSpeculatorCudaGraphManager
class AscendEagleSpeculator(EagleSpeculator):
def __init__(self, vllm_config: VllmConfig, device: torch.device):
"""Override GPU EagleSpeculator.__init__ for Ascend NPUs.
attnention metadata building in Ascend backend needs more information,
such as seq_lens_cpu from input_batch, so we need to override __init__.
"""
super().__init__(vllm_config, device)
del self.input_buffers
# AscendInputBuffers has extra `seq_lens_cpu` attribute.
# so reinitialize input_buffers here.
self.input_buffers: AscendInputBuffers = AscendInputBuffers(
max_num_reqs=self.max_num_reqs,
max_num_tokens=self.max_num_tokens,
device=device,
)
# add more attributes for `input_buffers` in graph mode
cudagraph_mode = self.vllm_config.compilation_config.cudagraph_mode
if cudagraph_mode.decode_mode() == CUDAGraphMode.FULL:
self.input_buffers.draft_seq_lens_cpus = [
torch.zeros(self.max_num_reqs, dtype=torch.int32, device="cpu")
for _ in range(self.num_speculative_steps - 1)
]
# we need to update full graph params in run_fullgraph,
# so create a stream to update full graph params.
if cudagraph_mode.has_full_cudagraphs():
self.update_stream: torch.npu.Stream = torch.npu.Stream()
# when in decode phase of eagle speculator, we need some value in
# draft model's input_batch. so we keep a reference here.
self.input_batch: InputBatch | None = None
def init_cudagraph_manager(self, cudagraph_mode: CUDAGraphMode) -> None:
with graph_manager_wrapper(self):
super().init_cudagraph_manager(cudagraph_mode)
def propose(
self,
input_batch: InputBatch,
attn_metadata: dict[str, Any],
slot_mappings: dict[str, torch.Tensor],
# [num_tokens, hidden_size]
last_hidden_states: torch.Tensor,
# num_layers x [num_tokens, hidden_size]
aux_hidden_states: list[torch.Tensor] | None,
# [num_reqs]
num_sampled: torch.Tensor,
# [num_reqs]
num_rejected: torch.Tensor,
# [max_num_reqs]
last_sampled: torch.Tensor,
# [max_num_reqs]
next_prefill_tokens: torch.Tensor,
# [max_num_reqs]
temperature: torch.Tensor,
# [max_num_reqs]
seeds: torch.Tensor,
num_tokens_across_dp: torch.Tensor | None = None,
dummy_run: bool = False,
skip_attn_for_dummy_run: bool = False,
mm_inputs: tuple[list[torch.Tensor], torch.Tensor] | None = None,
is_profile: Any = None,
):
"""Override GPU EagleSpeculator.propose for Ascend NPUs,
because npu attention metadata needs more information,
we need to cache input_batch, so we can use it later in
generate_draft.
"""
self.input_batch = input_batch
# wrap build_attn_metadata to use Ascend attention metadata building.
# so we can call super().propose() directly.
with build_attn_metadata_wrapper(), torch_gather_wrapper():
return super().propose(
input_batch,
attn_metadata,
slot_mappings,
last_hidden_states,
aux_hidden_states,
num_sampled,
num_rejected,
last_sampled,
next_prefill_tokens,
temperature,
seeds,
num_tokens_across_dp,
dummy_run,
skip_attn_for_dummy_run,
mm_inputs,
is_profile=is_profile,
)
def set_attn(
self,
model_state: ModelState,
kv_cache_config: KVCacheConfig,
block_tables: BlockTables,
) -> None:
super().set_attn(model_state, kv_cache_config, block_tables)
# npu needs attn_backends to update graph params
attn_backends: dict[str, type[AttentionBackend]] = {}
active_layer_names = self.draft_attn_layer_names
for kv_cache_group_id, kv_cache_group_spec in enumerate(kv_cache_config.kv_cache_groups):
layer_names = kv_cache_group_spec.layer_names
if active_layer_names is not None:
layer_names = list(active_layer_names.intersection(layer_names))
layer_type = cast(type[Any], AttentionLayerBase)
attn_layers = get_layers_from_vllm_config(self.vllm_config, layer_type, layer_names)
for layer_name in layer_names:
attn_backend = attn_layers[layer_name].get_attn_backend()
attn_backends[layer_name] = attn_backend
self.attn_backends = attn_backends
def _generate_draft(
self,
num_reqs: int,
num_tokens_padded: int,
attn_metadata: dict[str, Any] | None,
slot_mappings: dict[str, torch.Tensor] | None,
num_tokens_across_dp: torch.Tensor | None,
cudagraph_runtime_mode: CUDAGraphMode = CUDAGraphMode.NONE,
) -> None:
"""Override AutoRegressiveSpeculator._generate_draft for Ascend NPUs."""
self._ascend_prepare_decode_draft(attn_metadata, num_reqs)
super()._generate_draft(
num_reqs,
num_tokens_padded,
attn_metadata,
slot_mappings,
num_tokens_across_dp,
cudagraph_runtime_mode,
)
self._increment_decode_attn_metadata(attn_metadata)
@torch.inference_mode()
def _run_model(
self,
num_tokens: int,
attn_metadata: dict[str, Any] | None,
slot_mappings: dict[str, torch.Tensor] | None,
num_tokens_across_dp: torch.Tensor | None,
cudagraph_runtime_mode: CUDAGraphMode = CUDAGraphMode.NONE,
mm_inputs: tuple[list[torch.Tensor], torch.Tensor] | None = None,
) -> tuple[torch.Tensor, torch.Tensor]:
"""Override AutoRegressiveSpeculator._run_model for Ascend NPUs."""
last_hidden_states, hidden_states = super()._run_model(
num_tokens,
attn_metadata,
slot_mappings,
num_tokens_across_dp,
cudagraph_runtime_mode,
mm_inputs,
)
self._ascend_update_seq_lens(attn_metadata)
return last_hidden_states, hidden_states
def build_draft_attn_metadatas(self, num_reqs_padded, is_draft_model_prefill):
"""Build draft_attn_metadatas for partial-merged draft graph."""
attn_metadata = self.model_state.attn_metadata
attn_metadata = {
name: metadata for name, metadata in attn_metadata.items() if name in self.draft_attn_layer_names
}
if is_draft_model_prefill:
return [attn_metadata]
draft_attn_metadatas = self._init_decode_draft_attn_metadatas(attn_metadata, num_reqs_padded)
for i, per_step_attn_metadata in enumerate(draft_attn_metadatas):
step = i + 1
assert self.input_batch is not None
self._update_decode_attn_metadata(per_step_attn_metadata, step, self.input_batch.num_reqs)
return draft_attn_metadatas
def _ascend_prepare_decode_draft(self, attn_metadata: dict[str, Any] | None, num_reqs: int) -> None:
self._init_decode_attn_metadata(attn_metadata, num_reqs)
self._increment_decode_attn_metadata(attn_metadata)
def _ascend_update_seq_lens(self, attn_metadata: dict[str, Any] | None) -> None:
if attn_metadata is not None:
for attn_meta in attn_metadata.values():
attn_meta.seq_lens = attn_meta.seq_lens + 1
attn_meta.seq_len_list = attn_meta.seq_lens.tolist()
def _init_decode_attn_metadata(self, attn_metadata: dict[str, Any] | None, num_reqs: int):
"""Initialize attention metadata for decode phase on Ascend NPUs."""
if attn_metadata is None:
return
attn_state = AscendAttentionState.DecodeOnly
seq_lens_cpu = self._get_seq_lens_cpu()[:num_reqs]
# attn_metadata is build in vllm's super class.
# We need to update attn_state for each layer's metadata.
for metadata in attn_metadata.values():
metadata.attn_state = attn_state
metadata.seq_lens_cpu = seq_lens_cpu
def _init_decode_draft_attn_metadatas(self, attn_metadata: dict[str, Any] | None, num_reqs_padded: int):
"""Initialize attention metadata for decode phase in graph mode on Ascend NPUs."""
if attn_metadata is None:
return
attn_state = AscendAttentionState.DecodeOnly
draft_attn_metadatas = []
# attn_metadata is build in vllm's super class.
# We need to update attn_state for each layer's metadata.
for seq_lens_cpu in self.input_buffers.draft_seq_lens_cpus:
per_step_attn_metadata = {k: copy(v) for k, v in attn_metadata.items()}
seq_lens_cpu = seq_lens_cpu[:num_reqs_padded]
for metadata in per_step_attn_metadata.values():
metadata.attn_state = attn_state
metadata.seq_lens_cpu = seq_lens_cpu
draft_attn_metadatas.append(per_step_attn_metadata)
return draft_attn_metadatas
def _increment_decode_attn_metadata(self, attn_metadata: dict[str, Any] | None):
"""Increment attention metadata for decode phase on Ascend NPUs."""
# in eager mode, attn_metadata's seq_lens_cpu and input_buffers's seq_lens_cpu shares the memory
self._update_decode_attn_metadata(attn_metadata, 1)
def _update_decode_attn_metadata(
self, attn_metadata: dict[str, Any] | None, step: int, num_reqs: int | None = None
):
"""Update attention metadata for decode phase on Ascend NPUs."""
if attn_metadata is None:
return
num_reqs_padded = next(iter(attn_metadata.values())).seq_lens_cpu.shape[0]
seq_lens_cpu = self._get_seq_lens_cpu()[:num_reqs_padded]
if num_reqs is None:
num_reqs = num_reqs_padded
next_seq_lens_cpu = self._calc_next_seq_lens_cpu(seq_lens_cpu, num_reqs, num_reqs_padded, step)
query_lens_list = [i for i in range(1, num_reqs_padded + 1)]
seq_lens_list = next_seq_lens_cpu.tolist()
# attn_metadata is build in vllm's super class.
# We need to update attn_state for each layer's metadata.
for metadata in attn_metadata.values():
metadata.actual_seq_lengths_q = query_lens_list
metadata.seq_lens_cpu.copy_(next_seq_lens_cpu)
metadata.seq_lens_list = seq_lens_list
def _calc_next_seq_lens_cpu(self, seq_lens_cpu, num_reqs, num_reqs_padded, step):
# NOTE(drslark) to achieve fully alignment with vllm, `num_rejected` should be subtracted from `seq_lens`
# to avoid extra sync overhead, `v2` is currently aligned with NPU `v1` only
# follows the logic in `prepare_eagle_decode` and `update_eagle_inputs`
next_seqs_cpu = torch.clamp(seq_lens_cpu[:num_reqs_padded] + step, max=self.max_model_len)
next_seqs_cpu[num_reqs:].fill_(0)
return next_seqs_cpu
def _get_seq_lens_cpu(self) -> torch.Tensor:
"""Get seq_lens_cpu from input_batch."""
assert self.input_batch is not None
seq_lens_cpu = torch.from_numpy(self.input_batch.seq_lens_np)
return seq_lens_cpu
@contextmanager
def build_attn_metadata_wrapper():
"""Context manager to override attention metadata building for Ascend NPUs."""
original_func = _BUILD_ATTN_METADATA_MODULE.build_attn_metadata
try:
_BUILD_ATTN_METADATA_MODULE.build_attn_metadata = build_attn_metadata
yield
finally:
_BUILD_ATTN_METADATA_MODULE.build_attn_metadata = original_func
# TODO Remove this patch when cann fix the gather bug.
# NOTE(Ronald1995): torch.gather will pollute the cache such as self.input_buffers.positions
# the bug is reported to huawei CANN team, but not fixed yet.
# NOTE(drslark): make a temporary patch only for `torch.gather`
_original_gather = torch.gather
def gather(input, dim, index, *, sparse_grad=False, out=None):
if out is None:
return _original_gather(input, dim, index, sparse_grad=sparse_grad)
out[:] = _original_gather(input, dim, index, sparse_grad=sparse_grad)
return out
@contextmanager
def torch_gather_wrapper():
"""Context manager to override torch.gather for Ascend NPUs."""
original_gather = torch.gather
try:
torch.gather = gather
yield
finally:
torch.gather = original_gather
@contextmanager
def graph_manager_wrapper(speculator):
"""Context manager to override graph manager."""
original_graph_manager = _PREFILL_CUDAGRAPH_MANAGER_CLS
def factory(vllm_config: VllmConfig, device: torch.device, cudagraph_mode: CUDAGraphMode, decode_query_len: int):
return PrefillEagleAclGraphManager(vllm_config, device, cudagraph_mode, decode_query_len, speculator)
manager_attr = "PrefillSpeculatorCudaGraphManager"
try:
setattr(vllm_speculator_module, manager_attr, factory)
yield
finally:
setattr(vllm_speculator_module, manager_attr, original_graph_manager)

View File

@@ -0,0 +1,486 @@
# Adapt from https://github.com/vllm-project/vllm/blob/main/vllm/v1/worker/gpu/spec_decode/rejection_sampler_utils.py
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
#
# Copyright (c) 2025 Huawei Technologies Co., Ltd. All Rights Reserved.
#
# 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.
# This file is a part of the vllm-ascend project.
import torch
from vllm.triton_utils import tl, triton
from vllm.v1.worker.gpu.spec_decode.rejection_sampler_utils import (
_compute_block_stats_kernel,
_compute_global_lse,
_insert_resampled_kernel,
)
@triton.jit
def _npu_gumbel_block_argmax(
logits,
block,
mask,
token_idx,
expanded_idx_mapping_ptr,
temp_ptr,
seeds_ptr,
pos_ptr,
processed_logits_ptr,
processed_logits_stride,
processed_logits_col_ptr,
vocab_size,
APPLY_TEMPERATURE: tl.constexpr,
):
req_state_idx = tl.load(expanded_idx_mapping_ptr + token_idx)
temp = tl.load(temp_ptr + req_state_idx).to(tl.float32)
if temp != 0.0 and APPLY_TEMPERATURE:
logits = logits / temp
if processed_logits_ptr is not None:
if processed_logits_col_ptr is not None:
col = tl.load(processed_logits_col_ptr)
else:
col = 0
tl.store(
processed_logits_ptr + req_state_idx * processed_logits_stride + col * vocab_size + block,
logits,
mask=mask,
)
logits = logits.to(tl.float32)
if temp != 0.0:
seed = tl.load(seeds_ptr + req_state_idx)
# NPU: cast pos to int32 to avoid uint64 in philox (NPU umulhi only
# supports int32/uint32). Position values fit in int32 in practice.
pos = tl.load(pos_ptr + token_idx).to(tl.int32)
gumbel_seed = tl.randint(seed, pos)
# NPU: use tl.rand (float32) instead of tl_rand64 (float64 not supported)
r = tl.rand(gumbel_seed, block).to(tl.float32)
gumbel_noise = -tl.log(-tl.log(r + 1e-20) + 1e-20)
logits = tl.where(mask, logits + gumbel_noise, float("-inf"))
value, idx = tl.max(logits, axis=0, return_indices=True)
return value, idx
@triton.jit
def _resample_kernel(
# [num_reqs, num_blocks]
resampled_local_argmax_ptr,
resampled_local_argmax_stride,
# [num_reqs, num_blocks]
resampled_local_max_ptr,
resampled_local_max_stride,
# [num_logits, V]
target_logits_ptr,
target_logits_stride,
# [num_reqs]
target_rejected_logsumexp_ptr,
# [max_num_reqs, num_speculative_steps, V]
draft_logits_ptr,
draft_logits_stride_0,
draft_logits_stride_1,
# [num_reqs]
draft_rejected_logsumexp_ptr,
# [num_reqs]
rejected_step_ptr,
# [num_reqs + 1]
cu_num_logits_ptr,
# [num_logits]
expanded_idx_mapping_ptr,
# [num_logits]
draft_sampled_ptr,
# [max_num_reqs]
temp_ptr,
# [max_num_reqs]
seed_ptr,
# [num_logits]
pos_ptr,
vocab_size,
BLOCK_SIZE: tl.constexpr,
HAS_DRAFT_LOGITS: tl.constexpr,
):
req_idx = tl.program_id(0)
resample_idx = tl.load(rejected_step_ptr + req_idx)
start_idx = tl.load(cu_num_logits_ptr + req_idx)
end_idx = tl.load(cu_num_logits_ptr + req_idx + 1)
resample_token_idx = start_idx + resample_idx
req_state_idx = tl.load(expanded_idx_mapping_ptr + resample_token_idx)
temp = tl.load(temp_ptr + req_state_idx).to(tl.float32)
is_bonus = resample_token_idx == end_idx - 1
if temp == 0.0 and not is_bonus:
return
block_idx = tl.program_id(1)
block = block_idx * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
mask = block < vocab_size
target_logits = tl.load(
target_logits_ptr + resample_token_idx * target_logits_stride + block,
mask=mask,
other=float("-inf"),
).to(tl.float32)
if is_bonus:
residual_logits = target_logits
elif HAS_DRAFT_LOGITS:
draft_logits = tl.load(
draft_logits_ptr + req_state_idx * draft_logits_stride_0 + resample_idx * draft_logits_stride_1 + block,
mask=mask,
other=float("-inf"),
).to(tl.float32)
target_lse = tl.load(target_rejected_logsumexp_ptr + req_idx)
draft_lse = tl.load(draft_rejected_logsumexp_ptr + req_idx)
target_log_probs = target_logits - target_lse
draft_log_probs = draft_logits - draft_lse
ratio = tl.exp(draft_log_probs - target_log_probs)
residual_logits = tl.where(
ratio < 1.0,
target_log_probs + tl.log(1 - ratio),
float("-inf"),
).to(tl.float32)
else:
rejected_draft_token = tl.load(draft_sampled_ptr + resample_token_idx + 1)
residual_logits = tl.where(
block != rejected_draft_token,
target_logits,
float("-inf"),
).to(tl.float32)
value, idx = _npu_gumbel_block_argmax(
residual_logits,
block,
mask,
resample_token_idx,
expanded_idx_mapping_ptr,
temp_ptr,
seed_ptr,
pos_ptr,
None,
0,
None,
vocab_size,
APPLY_TEMPERATURE=False,
)
token_id = block_idx * BLOCK_SIZE + idx
tl.store(
resampled_local_argmax_ptr + req_idx * resampled_local_argmax_stride + block_idx,
token_id,
)
tl.store(
resampled_local_max_ptr + req_idx * resampled_local_max_stride + block_idx,
value,
)
@triton.jit
def _probabilistic_rejection_kernel(
# [num_reqs, num_speculative_steps + 1]
sampled_ptr,
sampled_stride,
# [num_reqs]
rejected_steps_ptr,
# [num_reqs]
target_rejected_logsumexp_ptr,
# [num_reqs]
draft_rejected_logsumexp_ptr,
# [num_logits, V]
target_logits_ptr,
target_logits_stride,
# [num_logits, num_blocks]
target_local_argmax_ptr,
target_local_argmax_stride,
# [num_logits, num_blocks]
target_local_max_ptr,
target_local_max_stride,
# [num_logits, num_blocks]
target_local_sumexp_ptr,
target_local_sumexp_stride,
# [num_logits]
draft_sampled_ptr,
# [max_num_reqs, num_speculative_steps, V]
draft_logits_ptr,
draft_logits_stride_0,
draft_logits_stride_1,
# [num_logits, num_blocks]
draft_local_max_ptr,
draft_local_max_stride,
# [num_logits, num_blocks]
draft_local_sumexp_ptr,
draft_local_sumexp_stride,
# [num_reqs + 1]
cu_num_logits_ptr,
# [num_reqs]
idx_mapping_ptr,
# [max_num_reqs]
temp_ptr,
# [max_num_reqs]
seed_ptr,
# [num_logits]
pos_ptr,
vocab_num_blocks,
PADDED_VOCAB_NUM_BLOCKS: tl.constexpr,
HAS_DRAFT_LOGITS: tl.constexpr,
):
req_idx = tl.program_id(0)
req_state_idx = tl.load(idx_mapping_ptr + req_idx)
start_idx = tl.load(cu_num_logits_ptr + req_idx)
end_idx = tl.load(cu_num_logits_ptr + req_idx + 1)
num_tokens = end_idx - start_idx
seed = tl.load(seed_ptr + req_state_idx) # noqa: F841
temp = tl.load(temp_ptr + req_state_idx).to(tl.float32)
rejected_step = 0
target_lse = 0.0
draft_lse = 0.0
accepted = True
for i in range(num_tokens - 1):
if accepted:
logit_idx = start_idx + i
draft_sampled = tl.load(draft_sampled_ptr + logit_idx + 1)
if temp == 0.0:
# Greedy sampling. Accept IFF draft matches target argmax.
# NOTE: Target argmax is stored directly so that resampling
# can be skipped upon rejection.
target_blocks = tl.arange(0, PADDED_VOCAB_NUM_BLOCKS)
target_blocks_mask = target_blocks < vocab_num_blocks
target_local_max = tl.load(
target_local_max_ptr + logit_idx * target_local_max_stride + target_blocks,
mask=target_blocks_mask,
other=float("-inf"),
)
max_target_block_idx = tl.argmax(target_local_max, axis=0)
target_argmax = tl.load(
target_local_argmax_ptr + logit_idx * target_local_argmax_stride + max_target_block_idx
)
accepted &= target_argmax == draft_sampled
tl.store(sampled_ptr + req_idx * sampled_stride + i, target_argmax)
else:
target_logit = tl.load(target_logits_ptr + logit_idx * target_logits_stride + draft_sampled).to(
tl.float32
)
target_lse = _compute_global_lse(
target_local_max_ptr,
target_local_max_stride,
target_local_sumexp_ptr,
target_local_sumexp_stride,
logit_idx,
vocab_num_blocks,
PADDED_VOCAB_NUM_BLOCKS,
)
target_log_prob = target_logit - target_lse
# NPU does not support tl_rand64; always accept the draft token.
u = tl.full([], 0.0, dtype=tl.float32)
if HAS_DRAFT_LOGITS:
draft_logit = tl.load(
draft_logits_ptr
+ req_state_idx * draft_logits_stride_0
+ i * draft_logits_stride_1
+ draft_sampled
).to(tl.float32)
draft_lse = _compute_global_lse(
draft_local_max_ptr,
draft_local_max_stride,
draft_local_sumexp_ptr,
draft_local_sumexp_stride,
logit_idx,
vocab_num_blocks,
PADDED_VOCAB_NUM_BLOCKS,
)
draft_log_prob = draft_logit - draft_lse
else:
# One-hot draft: q(draft_token) = 1, log_q = 0.
draft_log_prob = 0
# Probability ratio test: p(x) > u * q(x)
# Equivalent log form: log_p(x) > log(u) + log_q(x)
accepted &= target_log_prob > tl.log(u) + draft_log_prob
tl.store(sampled_ptr + req_idx * sampled_stride + i, draft_sampled)
rejected_step += accepted
tl.store(rejected_steps_ptr + req_idx, rejected_step)
tl.store(target_rejected_logsumexp_ptr + req_idx, target_lse)
tl.store(draft_rejected_logsumexp_ptr + req_idx, draft_lse)
def rejection_sample(
# [num_logits, V]
target_logits: torch.Tensor,
# [max_num_reqs, num_speculative_steps, V]
draft_logits: torch.Tensor | None,
# [num_logits]
draft_sampled: torch.Tensor,
# [num_reqs + 1]
cu_num_logits: torch.Tensor,
# [num_logits]
pos: torch.Tensor,
# [num_reqs]
idx_mapping: torch.Tensor,
# [num_logits]
expanded_idx_mapping: torch.Tensor,
# [num_logits]
expanded_local_pos: torch.Tensor,
# [max_num_reqs]
temperature: torch.Tensor,
# [max_num_reqs]
seed: torch.Tensor,
num_speculative_steps: int,
# [num_speculative_steps]
synthetic_conditional_rates: torch.Tensor | None = None,
use_fp64: bool = False,
# TODO: refactor speculative decoding functionality in a future PR.
# `use_block_verification` is accepted but not yet implemented on NPU;
# wire it up when the block verification path is supported.
use_block_verification: bool = False,
) -> tuple[torch.Tensor, torch.Tensor]:
if use_fp64:
raise NotImplementedError("FP64 rejection sampling is not supported on NPU.")
if synthetic_conditional_rates is not None:
# Synthetic rejection sampling needs tl_rand64, which NPU Triton does
# not support. The greedy fallback below would silently use u=0.0 and
# produce wrong acceptance — refuse loudly instead.
raise NotImplementedError(
"Synthetic rejection sampling is not supported on NPU yet; use rejection_sample_method='standard'."
)
num_reqs = cu_num_logits.shape[0] - 1
num_logits, vocab_size = target_logits.shape
has_draft_logits = draft_logits is not None
if draft_logits is None:
# When draft_logits is None, create a dummy tensor so that Triton
# kernel signatures receive valid pointers/strides. The kernels
# will never read from it when HAS_DRAFT_LOGITS=False.
draft_logits = target_logits.new_empty(1, 1, 1)
# Compute the block-level logits stats, such as target argmax
# (for greedy requests), and target max + softmax exponential
# (for non-greedy requests).
VOCAB_BLOCK_SIZE = 8192
vocab_num_blocks = triton.cdiv(vocab_size, VOCAB_BLOCK_SIZE)
padded_vocab_num_blocks = triton.next_power_of_2(vocab_num_blocks)
target_local_argmax = target_logits.new_empty(num_logits, vocab_num_blocks, dtype=torch.int64)
target_local_max = target_logits.new_empty(num_logits, vocab_num_blocks, dtype=torch.float32)
target_local_sumexp = target_logits.new_empty(num_logits, vocab_num_blocks, dtype=torch.float32)
draft_local_max = target_logits.new_empty(num_logits, vocab_num_blocks, dtype=torch.float32)
draft_local_sumexp = target_logits.new_empty(num_logits, vocab_num_blocks, dtype=torch.float32)
_compute_block_stats_kernel[(num_logits, vocab_num_blocks)](
target_local_argmax,
target_local_argmax.stride(0),
target_local_max,
target_local_max.stride(0),
target_local_sumexp,
target_local_sumexp.stride(0),
draft_local_max,
draft_local_max.stride(0),
draft_local_sumexp,
draft_local_sumexp.stride(0),
target_logits,
target_logits.stride(0),
draft_logits,
draft_logits.stride(0),
draft_logits.stride(1),
expanded_idx_mapping,
expanded_local_pos,
temperature,
vocab_size,
num_speculative_steps,
BLOCK_SIZE=VOCAB_BLOCK_SIZE,
HAS_DRAFT_LOGITS=has_draft_logits,
)
# Sample up until the first rejected/bonus token, and store
# the step.
sampled = draft_sampled.new_empty(num_reqs, num_speculative_steps + 1, dtype=torch.int64)
num_sampled = sampled.new_empty(num_reqs, dtype=torch.int32)
target_rejected_logsumexp = target_logits.new_empty(num_reqs, dtype=torch.float32)
draft_rejected_logsumexp = target_logits.new_empty(num_reqs, dtype=torch.float32)
_probabilistic_rejection_kernel[(num_reqs,)](
sampled,
sampled.stride(0),
num_sampled,
target_rejected_logsumexp,
draft_rejected_logsumexp,
target_logits,
target_logits.stride(0),
target_local_argmax,
target_local_argmax.stride(0),
target_local_max,
target_local_max.stride(0),
target_local_sumexp,
target_local_sumexp.stride(0),
draft_sampled,
draft_logits,
draft_logits.stride(0),
draft_logits.stride(1),
draft_local_max,
draft_local_max.stride(0),
draft_local_sumexp,
draft_local_sumexp.stride(0),
cu_num_logits,
idx_mapping,
temperature,
seed,
pos,
vocab_num_blocks,
PADDED_VOCAB_NUM_BLOCKS=padded_vocab_num_blocks,
HAS_DRAFT_LOGITS=has_draft_logits,
num_warps=1,
)
# Resample the rejected/bonus tokens.
RESAMPLE_BLOCK_SIZE = 1024
resample_num_blocks = triton.cdiv(vocab_size, RESAMPLE_BLOCK_SIZE)
padded_resample_num_blocks = triton.next_power_of_2(resample_num_blocks)
resampled_local_argmax = target_logits.new_empty(num_reqs, resample_num_blocks, dtype=torch.int64)
# NPU does not support float64; use float32 for resampled_local_max.
resampled_local_max = target_logits.new_empty(num_reqs, resample_num_blocks, dtype=torch.float32)
_resample_kernel[(num_reqs, resample_num_blocks)](
resampled_local_argmax,
resampled_local_argmax.stride(0),
resampled_local_max,
resampled_local_max.stride(0),
target_logits,
target_logits.stride(0),
target_rejected_logsumexp,
draft_logits,
draft_logits.stride(0),
draft_logits.stride(1),
draft_rejected_logsumexp,
num_sampled,
cu_num_logits,
expanded_idx_mapping,
draft_sampled,
temperature,
seed,
pos,
vocab_size,
BLOCK_SIZE=RESAMPLE_BLOCK_SIZE,
HAS_DRAFT_LOGITS=has_draft_logits,
)
# Insert the resampled tokens into the output sampled.
_insert_resampled_kernel[(num_reqs,)](
sampled,
sampled.stride(0),
num_sampled,
resampled_local_argmax,
resampled_local_argmax.stride(0),
resampled_local_max,
resampled_local_max.stride(0),
resample_num_blocks,
cu_num_logits,
expanded_idx_mapping,
temperature,
PADDED_RESAMPLE_NUM_BLOCKS=padded_resample_num_blocks,
)
return sampled, num_sampled