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)