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,39 @@
# [Experimental] Model Runner V2
This directory contains the new model runner which is under active development.
please see [Model Runner V2](https://github.com/vllm-project/vllm-ascend/issues/5208)
to get specific plans.
## Gaps with vLLM (To Be Addressed)
- [ ] `set_cos_and_sin` & `update_cos_sin`
Why: DeepSeek-like models (mla) still need cos/sin setting and updating in model_runner. These should be removed when mla can solve cos/sin internally.
Location: `NPUModelRunner.__init__`, `NPUModelRunner.prepare_inputs`, `AscendInputBatch.make_dummy`.
- [ ] `_allocate_kv_cache` & `_reshape_kv_cache`
Why: KV cache requires continuous space (thus divided as K cache and V cache separately) and PD disaggregation requires 2M-aligned tensors for KV cache, so custom KV cache initialization is needed. These should be removed when the above 2 requirements are no longer needed.
Location: `attn_utils._get_layer_kv_cache_specs`, `attn_utils._get_attention_kv_cache_dims`, `attn_utils._align_memory`, `attn_utils._allocate_kv_cache`, `attn_utils._reshape_kv_cache`.
- [ ] `torch_npu_graph_wrapper`
Why: FIA ops in FULL mode need explicit workspace allocating, and each workspace corresponding to each graph (a specific batch_size) should be released via `weak_ref_workspaces` when each capturing is exactly completed to avoid OOM, thus that leads us to regard `weak_ref_workspaces` as post-processing in `torch.npu.graph` and patch it. This should be removed when we don't need such special operations.
Location: `utils.torch_cuda_wrapper`, `utils.torch_npu_graph_wrapper`.
- [ ] `model_runner.graph_manager_wrapper`
Why: ModelAclGraphManager needs model_runner's input_buffers and model_state.attn_metadata to update_full_graph_params, so model_runner should be passed into __init__ of ModelAclGraphManager.
Location: `model_runner.NPUModelRunner.initialize_kv_cache`.
- [ ] `speculator.graph_manager_wrapper`
Why: EagleAclGraphManager needs speculator's input_buffers and model_state.attn_metadata to update_full_graph_params, so speculator should be passed into __init
__ of EagleAclGraphManager.
Location: `speculator.AscendEagleSpeculator.init_cudagraph_manager`.

View File

View File

@@ -0,0 +1,163 @@
# Adapt from https://github.com/vllm-project/vllm/blob/main/vllm/v1/worker/gpu/aclgraph_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.
#
from collections.abc import Callable
from typing import Any
import torch
import torch.nn as nn
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.sequence import IntermediateTensors
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 BatchExecutionDescriptor, ModelCudaGraphManager
from vllm.v1.worker.gpu.input_batch import InputBuffers
from vllm.v1.worker.gpu.model_states.interface import ModelState
from vllm.v1.worker.utils import AttentionGroup
from vllm_ascend.ascend_forward_context import _EXTRA_CTX
from vllm_ascend.compilation.acl_graph import set_graph_params, update_full_graph_params
from vllm_ascend.worker.v2.utils import communicator_switch
class ModelAclGraphManager(ModelCudaGraphManager):
"""ACL Model Cuda Graph Manager for Ascend NPUs."""
def __init__(
self,
vllm_config: VllmConfig,
device: torch.device,
cudagraph_mode: CUDAGraphMode,
decode_query_len: int,
model_runner: Any,
lora_capture_cases: list[int] | None = None,
):
super().__init__(
vllm_config,
device,
cudagraph_mode,
decode_query_len,
lora_capture_cases=lora_capture_cases,
)
# set model runner attribute, so we can access attributes model runner
# when call `run_fullgraph` method in CudaGraphManager,
# then we don't need to # copy `execute_model` method in `NPUModelRunner` class.
self.model_runner = model_runner
# capture_sizes sorts in ascending order.
self.capture_sizes = sorted(self.compilation_config.cudagraph_capture_sizes)
# vllm-ascend need to update graph params of attention backend.
# so we need to set graph params before capture full graph.
if super().needs_capture():
set_graph_params(self.capture_sizes)
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
logger.info_once("run_fullgraph with num_tokens=%s", num_tokens)
ret = super().run_fullgraph(desc)
positions = self.model_runner.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.model_runner.dp_size], num_tokens, device=self.device)
with set_forward_context(
self.model_runner.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,
):
forward_context = get_forward_context()
update_full_graph_params(
# FIXME(Ronald1995): support hybrid attn backend
self.model_runner.attn_groups[0][0].backend,
self.model_runner.update_stream,
forward_context,
num_tokens,
self.vllm_config,
self.model_runner.speculative_config,
positions.shape[0],
)
return ret
def capture(
self,
model: nn.Module,
model_state: ModelState,
input_buffers: InputBuffers,
intermediate_tensors: IntermediateTensors | None,
block_tables: BlockTables,
attn_groups: list[list[AttentionGroup]],
kv_cache_config: KVCacheConfig,
has_lora: bool = False,
use_aux_hidden_state_outputs: bool = False,
lora_capture_hook: Callable[[int, int, int], None] | None = None,
progress_bar_desc: str = "Capturing CUDA graphs",
) -> None:
"""Capture CUDA graphs for model forward pass."""
model = ModelWithContext(model)
with communicator_switch():
return super().capture(
model,
model_state,
input_buffers,
intermediate_tensors,
block_tables,
attn_groups,
kv_cache_config,
has_lora=has_lora,
use_aux_hidden_state_outputs=use_aux_hidden_state_outputs,
lora_capture_hook=lora_capture_hook,
progress_bar_desc=progress_bar_desc,
)
class ModelWithContext(nn.Module):
"""Define a wrapper model to inject forward context.
so we can inherit vllm's CudaGraphManager._capture_full_graph.
"""
def __init__(self, original_model, is_draft_model=False, is_draft_model_prefill=False):
super().__init__()
self.original_model = original_model
self.is_draft_model = is_draft_model
self.is_draft_model_prefill = is_draft_model_prefill
def forward(self, *args, **kwargs):
# In warmup phase, capturing=False by default.
# when capturing, we need to set capturing=True in forward context.
if torch.npu.is_current_stream_capturing():
_EXTRA_CTX.capturing = True
if self.is_draft_model:
_EXTRA_CTX.is_draft_model = True
if self.is_draft_model_prefill:
_EXTRA_CTX.is_draft_model_prefill = True
return self.original_model(*args, **kwargs)
def get_original_model(self):
return self.original_model
def compute_logits(self, hidden_states: torch.Tensor):
# draft model has `compute_logits`, which is not in ModelWithContext
return self.original_model.compute_logits(hidden_states)

View File

@@ -0,0 +1,504 @@
# Adapt from https://github.com/vllm-project/vllm/blob/main/vllm/v1/worker/gpu/attn_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.
#
from collections.abc import Sequence
from typing import Any
import numpy as np
import torch
from vllm.config import VllmConfig, get_current_vllm_config, get_layers_from_vllm_config
from vllm.model_executor.layers.attention import Attention
from vllm.model_executor.layers.attention.mla_attention import MLAAttention
from vllm.model_executor.layers.attention_layer_base import AttentionLayerBase
from vllm.v1.attention.backend import AttentionBackend
from vllm.v1.kv_cache_interface import (
AttentionSpec,
EncoderOnlyAttentionSpec,
KVCacheConfig,
KVCacheSpec,
MLAAttentionSpec,
UniformTypeKVCacheSpecs,
)
from vllm.v1.worker.gpu.model_states.interface import ModelSpecificAttnMetadata
from vllm.v1.worker.utils import AttentionGroup
from vllm_ascend.attention.attention_mask import AttentionMaskBuilder
from vllm_ascend.attention.attention_v1 import AscendAttentionState
from vllm_ascend.attention.utils import AscendCommonAttentionMetadata, AscendPrefillContextParallelMetadata
from vllm_ascend.core.kv_cache_interface import AscendMLAAttentionSpec
from vllm_ascend.quantization.utils import enable_fa_quant
from vllm_ascend.utils import calc_split_factor
_ATTENTION_MASK_BUILDER = None
def get_kv_cache_spec(vllm_config: VllmConfig) -> dict[str, KVCacheSpec]:
"""Build Ascend-specific KV cache specs for v2 worker patching."""
kv_cache_spec: dict[str, KVCacheSpec] = {}
layer_type = AttentionLayerBase
attn_layers = get_layers_from_vllm_config(vllm_config, layer_type)
for layer_name, attn_module in attn_layers.items():
if getattr(attn_module, "kv_sharing_target_layer_name", None):
continue
if isinstance(attn_module, Attention):
if spec := attn_module.get_kv_cache_spec(vllm_config):
kv_cache_spec[layer_name] = spec
continue
if isinstance(attn_module, MLAAttention):
spec = attn_module.get_kv_cache_spec(vllm_config)
if spec is None:
continue
if getattr(attn_module.impl, "fa_quant_layer", False):
head_size = attn_module.head_size + attn_module.qk_rope_head_dim
dtype, cache_dtype_str = attn_module.impl.dtype, None
else:
head_size = spec.head_size
dtype = spec.dtype
cache_dtype_str = spec.cache_dtype_str
kv_cache_spec[layer_name] = AscendMLAAttentionSpec(
block_size=spec.block_size,
num_kv_heads=spec.num_kv_heads,
head_size=head_size,
dtype=dtype,
cache_dtype_str=cache_dtype_str,
)
return kv_cache_spec
def get_attn_mask_builder(device: torch.device):
"""Get attention mask builder which only have one instance."""
global _ATTENTION_MASK_BUILDER
if _ATTENTION_MASK_BUILDER is None:
_ATTENTION_MASK_BUILDER = AttentionMaskBuilder(device)
return _ATTENTION_MASK_BUILDER
def build_attn_metadata(
*,
attn_groups: list[list[AttentionGroup]],
num_reqs: int,
num_tokens: int,
query_start_loc_gpu: torch.Tensor,
query_start_loc_cpu: torch.Tensor,
max_query_len: int,
seq_lens: torch.Tensor,
max_seq_len: int,
block_tables: Sequence[torch.Tensor],
slot_mappings: torch.Tensor,
kv_cache_config: KVCacheConfig,
dcp_local_seq_lens: torch.Tensor | None = None,
# extra attributes for ascend npus.
seq_lens_np: np.ndarray | None = None,
num_computed_tokens_cpu: torch.Tensor | None = None,
positions: torch.Tensor | None = None,
attn_state: Any | None = None,
graph_pad_size: int = -1,
num_input_tokens: int = 0,
prefill_context_parallel_metadata: AscendPrefillContextParallelMetadata | None = None,
model_specific_attn_metadata: ModelSpecificAttnMetadata | None = None,
for_cudagraph_capture: bool = False,
causal: bool = True,
) -> dict[str, Any]:
"""Build attention metadata for Ascend NPUs."""
# TODO(Ronald1995): optimize AscendCommonAttentionMetadata.
# seq_lens_np is used for ascend npus, it maybe None in spec_decode case,
# we fill it with max_seq_len in case `attn_metadata_builder.build` raise
# an error.
if seq_lens_np is None:
seq_lens_np = np.full(num_reqs, max_seq_len, dtype=np.int32)
seq_lens_cpu = torch.from_numpy(seq_lens_np)[:num_reqs]
attn_metadata: dict[str, Any] = {}
kv_cache_groups = kv_cache_config.kv_cache_groups
for i, kv_cache_spec in enumerate(kv_cache_groups):
block_table = block_tables[i]
slot_mapping = slot_mappings[i]
common_attn_metadata_extra_kwargs = (
model_specific_attn_metadata.get_extra_common_attn_kwargs(i, num_reqs)
if model_specific_attn_metadata is not None
else {}
)
common_attn_metadata = AscendCommonAttentionMetadata(
query_start_loc=query_start_loc_gpu,
query_start_loc_cpu=query_start_loc_cpu,
seq_lens_cpu=seq_lens_cpu,
seq_lens_cpu_upper_bound=seq_lens_cpu,
seq_lens=seq_lens[:num_reqs],
num_reqs=num_reqs,
num_actual_tokens=num_tokens,
max_query_len=max_query_len,
block_table_tensor=block_table,
slot_mapping=slot_mapping,
positions=positions,
attn_state=attn_state,
graph_pad_size=graph_pad_size,
num_input_tokens=num_input_tokens,
prefill_context_parallel_metadata=prefill_context_parallel_metadata,
max_seq_len=max_seq_len,
causal=causal,
**common_attn_metadata_extra_kwargs,
)
for attn_group in attn_groups[i]:
attn_metadata_builder = attn_group.get_metadata_builder(0)
if for_cudagraph_capture:
metadata = attn_metadata_builder.build_for_cudagraph_capture(common_attn_metadata)
else:
attn_metadata_extra_kwargs = (
model_specific_attn_metadata.get_extra_attn_kwargs(
attn_metadata_builder,
num_reqs,
)
if model_specific_attn_metadata is not None
else {}
)
metadata = attn_metadata_builder.build(
common_prefix_len=0,
common_attn_metadata=common_attn_metadata,
**attn_metadata_extra_kwargs,
)
for layer_name in attn_group.layer_names:
attn_metadata[layer_name] = metadata
return attn_metadata
def build_attn_state(
vllm_config: VllmConfig,
seq_lens_np: np.ndarray,
num_reqs,
num_scheduled_tokens,
num_valid_tokens,
):
"""Build attention state for npu's attention backend."""
if vllm_config.model_config.runner_type == "pooling":
if isinstance(
vllm_config.kv_cache_config.kv_cache_groups[0].kv_cache_spec,
EncoderOnlyAttentionSpec,
):
attn_state = AscendAttentionState.PrefillNoCache
else:
attn_state = AscendAttentionState.PrefillCacheHit
elif np.array_equal(seq_lens_np[:num_reqs], num_scheduled_tokens):
attn_state = AscendAttentionState.PrefillNoCache
# We assume it is the decode stage, where prefill occurs
# but only one token is not hit in cache.
elif np.all(num_scheduled_tokens == 1):
attn_state = AscendAttentionState.DecodeOnly
if vllm_config.speculative_config and vllm_config.speculative_config.method == "mtp":
# SpecDecoding now supports seq_len=1 and seq_len=2
# In Prefilling Decoding Disaggregation scenario, SpecDecoding
# need to supports seq_len=1
attn_state = AscendAttentionState.SpecDecoding
# Speculative decoding.
elif np.all(num_valid_tokens == 1):
if vllm_config.speculative_config and vllm_config.speculative_config.method == "mtp":
attn_state = AscendAttentionState.SpecDecoding
else:
attn_state = AscendAttentionState.ChunkedPrefill
# splitfuse
elif vllm_config.scheduler_config.enable_chunked_prefill:
attn_state = AscendAttentionState.ChunkedPrefill
else:
attn_state = AscendAttentionState.PrefillCacheHit
return attn_state
def _get_layer_kv_cache_specs(kv_cache_config: KVCacheConfig) -> dict[str, KVCacheSpec]:
layer_kv_cache_spec: dict[str, KVCacheSpec] = {}
for group_kv_cache_spec in kv_cache_config.kv_cache_groups:
group_spec = group_kv_cache_spec.kv_cache_spec
for layer_name in group_kv_cache_spec.layer_names:
if isinstance(group_spec, UniformTypeKVCacheSpecs):
layer_kv_cache_spec[layer_name] = group_spec.kv_cache_specs[layer_name]
else:
layer_kv_cache_spec[layer_name] = group_spec
return layer_kv_cache_spec
def _get_attention_kv_cache_dims(layer_name: str, kv_cache_spec: AttentionSpec) -> tuple[int, int]:
if isinstance(kv_cache_spec, AscendMLAAttentionSpec):
attn_layers = get_layers_from_vllm_config(get_current_vllm_config(), AttentionLayerBase, [layer_name])
attn_layer = attn_layers[layer_name]
if not isinstance(attn_layer, MLAAttention):
raise TypeError(f"Expected AscendMLAAttention layer for {layer_name}, got {type(attn_layer).__name__}.")
return attn_layer.kv_lora_rank, attn_layer.qk_rope_head_dim
head_size_v = kv_cache_spec.head_size_v if hasattr(kv_cache_spec, "head_size_v") else kv_cache_spec.head_size
return kv_cache_spec.head_size, head_size_v
def _align_memory(tensor: torch.Tensor, alignment: int) -> torch.Tensor:
data_ptr = tensor.data_ptr()
aligned_addr = (data_ptr + alignment - 1) // alignment * alignment
offset = (aligned_addr - data_ptr) // tensor.element_size()
return tensor[int(offset) :]
def _allocate_kv_cache(
kv_cache_config: KVCacheConfig,
shared_layers: dict[str, str],
device: torch.device,
) -> dict[str, tuple[torch.Tensor, torch.Tensor]]:
"""
Initialize the KV cache buffer with the correct size. The buffer needs to be
reshaped to the desired shape before being used by the models.
NOTE: To support prefill disaggregation, we need to split kvcache tensor
into k_cache and v_cache, and the addr of both are aligned by 2M.
Args:
kv_cache_config: The KV cache config
device: The device
Returns:
dict[str, tuple[torch.Tensor, torch.Tensor]]: A map between layer names
to their corresponding memory buffer for K cache and V cache
"""
vllm_config = get_current_vllm_config()
# init kv cache tensors
kv_cache_raw_tensors: dict[str, tuple[torch.Tensor, torch.Tensor]] = {}
# prefill disaggregation need the addr of cache tensor be aligned with 2M
alignment = 2 * 1024 * 1024
layer_kv_cache_spec = _get_layer_kv_cache_specs(kv_cache_config)
for kv_cache_tensor in kv_cache_config.kv_cache_tensors:
if len(kv_cache_tensor.shared_by) == 0:
continue
# NOTE: We need to init k_cache tensor (nope cache tensor in mla) and
# v_cache tensor (rope cache tensor in mla) separately to support
# prefill disaggregation, as it only supports the 0-dim of kv_cache is
# `num_blocks`.
# For deepseek mla, we need to spilt cache tensor accrodding to the nope
# head dim and rope head dim.
example_layer_name = kv_cache_tensor.shared_by[0]
example_kv_cache_spec = layer_kv_cache_spec[example_layer_name]
assert isinstance(example_kv_cache_spec, AttentionSpec)
k_dim, v_dim = _get_attention_kv_cache_dims(example_layer_name, example_kv_cache_spec)
assert k_dim > 0 and v_dim > 0
kv_head_dim_list = [k_dim, v_dim]
if enable_fa_quant(vllm_config):
k_tensor_split_factor, v_tensor_split_factor = vllm_config.quant_config.get_kv_quant_split_factor(
example_layer_name, kv_head_dim_list
)
else:
k_tensor_split_factor, v_tensor_split_factor = calc_split_factor(kv_head_dim_list)
k_tensor_size = int(kv_cache_tensor.size // k_tensor_split_factor)
v_tensor_size = int(kv_cache_tensor.size // v_tensor_split_factor)
if vllm_config.kv_transfer_config is None:
k_tensor = torch.zeros(k_tensor_size, dtype=torch.int8, device=device)
v_tensor = torch.zeros(v_tensor_size, dtype=torch.int8, device=device)
else:
k_tensor = torch.zeros(k_tensor_size + alignment, dtype=torch.int8, device=device)
v_tensor = torch.zeros(v_tensor_size + alignment, dtype=torch.int8, device=device)
k_tensor = _align_memory(k_tensor, alignment)[:k_tensor_size]
v_tensor = _align_memory(v_tensor, alignment)[:v_tensor_size]
for layer_name in kv_cache_tensor.shared_by:
kv_cache_raw_tensors[layer_name] = (k_tensor, v_tensor)
layer_names = set()
for group in kv_cache_config.kv_cache_groups:
for layer_name in group.layer_names:
layer_names.add(layer_name)
assert layer_names == (kv_cache_raw_tensors.keys() | shared_layers.keys()), (
"Some layers are not correctly initialized"
)
return kv_cache_raw_tensors
def _reshape_kv_cache(
kv_cache_config: KVCacheConfig,
kv_cache_raw_tensors: dict[str, tuple[torch.Tensor, torch.Tensor]],
attn_backends: dict[str, AttentionBackend],
cache_dtype: str,
kernel_block_sizes: list[int] | None = None,
shared_kv_cache_layers: dict[str, str] | None = None,
) -> dict[str, tuple[torch.Tensor, torch.Tensor]]:
"""
Reshape the KV cache tensors to the desired shape and dtype.
Args:
kv_cache_config: The KV cache config
kv_cache_raw_tensors: The KV cache buffer of each layer, with correct
size but uninitialized shape
Returns:
dict[str, tuple[torch.Tensor, torch.Tensor]]: A map between layer names
to their corresponding memory buffer for KV cache
"""
vllm_config = get_current_vllm_config()
kv_caches: dict[str, tuple[torch.Tensor, torch.Tensor]] = {}
kernel_block_sizes = kernel_block_sizes or []
for kv_cache_group_id, kv_cache_group_spec in enumerate(kv_cache_config.kv_cache_groups):
for layer_name in kv_cache_group_spec.layer_names:
if shared_kv_cache_layers and layer_name in shared_kv_cache_layers:
continue
kv_cache_spec = kv_cache_group_spec.kv_cache_spec
if isinstance(kv_cache_spec, UniformTypeKVCacheSpecs):
kv_cache_spec = kv_cache_spec.kv_cache_specs[layer_name]
assert isinstance(kv_cache_spec, AttentionSpec)
if isinstance(kv_cache_spec, AttentionSpec):
raw_k_tensor, raw_v_tensor = kv_cache_raw_tensors[layer_name]
assert raw_k_tensor is not None
assert raw_v_tensor is not None
sum_page_size_bytes = raw_k_tensor.numel() + raw_v_tensor.numel()
assert sum_page_size_bytes % kv_cache_spec.page_size_bytes == 0
num_blocks = sum_page_size_bytes // kv_cache_spec.page_size_bytes
# `num_blocks` is the number of blocks the model runner can use.
# `kv_cache_config.num_blocks` is the number of blocks that
# KVCacheManager may allocate.
# Since different GPUs may have different number of layers and
# different memory capacities, `num_blocks` can be different on
# different GPUs, and `kv_cache_config.num_blocks` is set to
# the min of all `num_blocks`. Verify it here.
assert num_blocks >= kv_cache_config.num_blocks
attn_backend = attn_backends[layer_name]
if kv_cache_group_id < len(kernel_block_sizes):
kernel_block_size = kernel_block_sizes[kv_cache_group_id]
num_blocks *= kv_cache_spec.block_size // kernel_block_size
else:
kernel_block_size = kv_cache_spec.block_size
if kv_cache_spec.storage_block_size != kv_cache_spec.block_size:
shape_block_size = kv_cache_spec.storage_block_size
else:
shape_block_size = kernel_block_size
kv_cache_shape = attn_backend.get_kv_cache_shape(
num_blocks,
shape_block_size,
kv_cache_spec.num_kv_heads,
kv_cache_spec.head_size,
cache_dtype,
)
if not isinstance(kv_cache_spec, AscendMLAAttentionSpec):
k_shape = kv_cache_shape[1:]
if hasattr(kv_cache_spec, "head_size_v"):
v_shape = (*kv_cache_shape[1:-1], kv_cache_spec.head_size_v)
else:
v_shape = k_shape
else:
# k_cache: nope_cache v_cache: rope_cache
mla_num_blocks, mla_block_size, num_kv_heads, _ = kv_cache_shape
k_dim, v_dim = _get_attention_kv_cache_dims(layer_name, kv_cache_spec)
k_shape = (mla_num_blocks, mla_block_size, num_kv_heads, k_dim)
v_shape = (mla_num_blocks, mla_block_size, num_kv_heads, v_dim)
k_cache_dtype = v_cache_dtype = kv_cache_spec.dtype
if enable_fa_quant(vllm_config):
k_cache_dtype, v_cache_dtype = vllm_config.quant_config.get_kv_quant_dtype(
layer_name, kv_cache_spec.dtype, vllm_config.model_config
)
k_cache = raw_k_tensor.view(k_cache_dtype).view(k_shape)
v_cache = raw_v_tensor.view(v_cache_dtype).view(v_shape)
kv_caches[layer_name] = (k_cache, v_cache)
else:
raise ValueError("Unknown KV cache spec type.")
if shared_kv_cache_layers:
for layer_name, target_layer_name in shared_kv_cache_layers.items():
kv_caches[layer_name] = kv_caches[target_layer_name]
return kv_caches
def _reshape_kv_cache_v2(
attn_groups: Sequence[AttentionGroup],
kv_cache_raw_tensors: dict[str, tuple[torch.Tensor, torch.Tensor]],
cache_dtype: str,
kernel_block_sizes: list[int],
shared_kv_cache_layers: dict[str, str],
kv_cache_config: "KVCacheConfig | None" = None,
) -> dict[str, tuple[torch.Tensor, torch.Tensor]]:
vllm_config = get_current_vllm_config()
is_kv_consumer = (
vllm_config.kv_transfer_config.is_kv_consumer if vllm_config.kv_transfer_config is not None else False
)
kv_caches: dict[str, tuple[torch.Tensor, torch.Tensor]] = {}
for group in attn_groups:
if group.kv_cache_group_id >= len(kernel_block_sizes):
continue
kv_cache_spec = group.kv_cache_spec
if kv_cache_spec.storage_block_size != kv_cache_spec.block_size:
kernel_block_size = kv_cache_spec.storage_block_size
else:
kernel_block_size = kernel_block_sizes[group.kv_cache_group_id]
for layer_name in group.layer_names:
if layer_name in shared_kv_cache_layers:
continue
assert isinstance(kv_cache_spec, AttentionSpec)
raw_k_tensor, raw_v_tensor = kv_cache_raw_tensors[layer_name]
assert raw_k_tensor is not None
assert raw_v_tensor is not None
sum_page_size_bytes = raw_k_tensor.numel() + raw_v_tensor.numel()
assert sum_page_size_bytes % kv_cache_spec.page_size_bytes == 0
num_blocks = sum_page_size_bytes // kv_cache_spec.page_size_bytes
num_blocks_per_kv_block = kv_cache_spec.block_size // kernel_block_size
kernel_num_blocks = num_blocks * num_blocks_per_kv_block
kv_cache_shape = group.backend.get_kv_cache_shape(
kernel_num_blocks,
kernel_block_size,
kv_cache_spec.num_kv_heads,
kv_cache_spec.head_size,
cache_dtype,
)
if not isinstance(kv_cache_spec, (AscendMLAAttentionSpec, MLAAttentionSpec)):
k_shape = kv_cache_shape[1:]
if hasattr(kv_cache_spec, "head_size_v"):
v_shape = (*kv_cache_shape[1:-1], kv_cache_spec.head_size_v)
else:
v_shape = k_shape
else:
mla_num_blocks, mla_block_size, num_kv_heads, _ = kv_cache_shape
k_dim, v_dim = _get_attention_kv_cache_dims(layer_name, kv_cache_spec)
k_shape = (mla_num_blocks, mla_block_size, num_kv_heads, k_dim)
v_shape = (mla_num_blocks, mla_block_size, num_kv_heads, v_dim)
k_cache_dtype = v_cache_dtype = kv_cache_spec.dtype
if is_kv_consumer and enable_fa_quant(vllm_config):
k_cache_dtype, v_cache_dtype = vllm_config.quant_config.get_kv_quant_dtype(
layer_name, kv_cache_spec.dtype, vllm_config.model_config
)
k_cache = raw_k_tensor.view(k_cache_dtype).view(k_shape)
v_cache = raw_v_tensor.view(v_cache_dtype).view(v_shape)
kv_caches[layer_name] = (k_cache, v_cache)
for layer_name, target_layer_name in shared_kv_cache_layers.items():
kv_caches[layer_name] = kv_caches[target_layer_name]
return kv_caches

View File

@@ -0,0 +1,164 @@
# Adapt from https://github.com/vllm-project/vllm/blob/main/vllm/v1/worker/gpu/block_table.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.attention.backends.utils import PAD_SLOT_ID
from vllm.v1.worker.gpu.block_table import BlockTables, _load_ptr
class AscendBlockTables(BlockTables):
"""Block table for Ascend NPUs."""
def __init__(
self,
block_sizes: list[int],
max_num_reqs: int,
max_num_batched_tokens: int,
max_num_blocks_per_group: list[int],
device: torch.device,
kernel_block_sizes: list[int] | None = None,
cp_size: int = 1,
cp_rank: int = 0,
cp_interleave: int = 1,
):
if kernel_block_sizes is None:
kernel_block_sizes = block_sizes
super().__init__(
block_sizes,
max_num_reqs,
max_num_batched_tokens,
max_num_blocks_per_group,
device,
kernel_block_sizes,
cp_size,
cp_rank,
cp_interleave,
)
# because we will override these attribute, delete these attribute to
# make sure it's collected by python gc immediately.
del self.slot_mappings
# vllm-ascend' reshape_and_cache function requires slot_mappings to be int32.
# so we need to redefine slot_mappings to be int32.
self.slot_mappings: torch.Tensor = torch.zeros(
self.num_kv_cache_groups,
self.max_num_batched_tokens,
dtype=torch.int32,
device=self.device,
)
def compute_slot_mappings(
self,
idx_mapping: torch.Tensor,
query_start_loc: torch.Tensor,
positions: torch.Tensor,
num_tokens_padded: int,
) -> torch.Tensor:
num_reqs = idx_mapping.shape[0]
num_groups = self.num_kv_cache_groups
_compute_slot_mappings_kernel[(num_groups, num_reqs + 1)](
self.max_num_batched_tokens,
idx_mapping,
query_start_loc,
positions,
self.block_table_ptrs,
self.block_table_strides,
self.block_sizes_tensor,
self.slot_mappings,
self.slot_mappings.stride(0),
self.cp_rank,
CP_SIZE=self.cp_size,
CP_INTERLEAVE=self.cp_interleave,
PAD_ID=PAD_SLOT_ID,
TRITON_BLOCK_SIZE=1024, # type: ignore
TOTAL_BLOCK_SIZE=4096,
)
return self.slot_mappings[:, :num_tokens_padded]
@triton.jit
def _compute_slot_mappings_kernel(
max_num_tokens,
idx_mapping, # [num_reqs]
query_start_loc, # [num_reqs + 1]
pos, # [num_tokens]
block_table_ptrs, # [num_kv_cache_groups]
block_table_strides, # [num_kv_cache_groups]
block_sizes, # [num_kv_cache_groups]
slot_mappings_ptr, # [num_kv_cache_groups, max_num_tokens]
slot_mappings_stride,
cp_rank,
CP_SIZE: tl.constexpr,
CP_INTERLEAVE: tl.constexpr,
PAD_ID: tl.constexpr,
TRITON_BLOCK_SIZE: tl.constexpr,
TOTAL_BLOCK_SIZE: tl.constexpr,
):
# kv cache group id
group_id = tl.program_id(0)
batch_idx = tl.program_id(1)
slot_mapping_ptr = slot_mappings_ptr + group_id * slot_mappings_stride
if batch_idx == tl.num_programs(1) - 1:
actual_num_tokens = tl.load(query_start_loc + batch_idx)
for i in range(actual_num_tokens, max_num_tokens, TRITON_BLOCK_SIZE):
offset = i + tl.arange(0, TRITON_BLOCK_SIZE)
tl.store(slot_mapping_ptr + offset, PAD_ID, mask=offset < max_num_tokens)
return
block_table_ptr = _load_ptr(block_table_ptrs + group_id, tl.int32)
block_table_stride = tl.load(block_table_strides + group_id)
block_size = tl.load(block_sizes + group_id)
req_state_idx = tl.load(idx_mapping + batch_idx)
start_idx = tl.load(query_start_loc + batch_idx)
end_idx = tl.load(query_start_loc + batch_idx + 1)
for i in range(start_idx, end_idx, TRITON_BLOCK_SIZE):
offset = i + tl.arange(0, TRITON_BLOCK_SIZE)
positions = tl.load(pos + offset, mask=offset < end_idx, other=0)
# Type conversion of 'position' to int32 to be compatible with npu
# otherwise, it will degrade to scalar computation
positions = positions.to(tl.int32)
block_indices = positions // (block_size * CP_SIZE)
# block_offset = positions % (block_size * CP_SIZE)
# The % operation on int32 type will degrade to scalar computation
# replace the % operation with sub and mul instead
block_offsets = positions - (block_size * CP_SIZE) * block_indices
# The 'block_indics' variable results in non-contiguous memory assess,
# which triggers degradation toscalar computation.
# Mitigate this by loading the complete data block and extracting the required data with tl.gather
block_numbers = tl.load(block_table_ptr + req_state_idx * block_table_stride + tl.arange(0, TOTAL_BLOCK_SIZE))
block_numbers = block_numbers.to(tl.float32)
block_numbers = tl.gather(block_numbers, block_indices, 0)
if CP_SIZE == 1:
# Common case: Context parallelism is not used.
slot_ids = block_numbers * block_size + block_offsets
else:
# Context parallelism is used.
is_local = block_offsets // CP_INTERLEAVE % CP_SIZE == cp_rank
rounds = block_offsets // (CP_INTERLEAVE * CP_SIZE)
remainder = block_offsets % CP_INTERLEAVE
local_offsets = rounds * CP_INTERLEAVE + remainder
slot_ids = block_numbers * block_size + local_offsets
slot_ids = tl.where(is_local, slot_ids, PAD_ID)
tl.store(slot_mapping_ptr + offset, slot_ids, mask=offset < end_idx)

View File

@@ -0,0 +1,210 @@
# Adapt from https://github.com/vllm-project/vllm/blob/main/vllm/v1/worker/gpu/input_batch.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 dataclasses import asdict, dataclass
import numpy as np
import torch
from vllm.triton_utils import tl, triton
from vllm.v1.worker.gpu.input_batch import InputBatch, InputBuffers
from vllm_ascend.attention.attention_v1 import AscendAttentionState
from vllm_ascend.ops.rotary_embedding import update_cos_sin
from vllm_ascend.ops.triton.triton_utils import get_vectorcore_num
class AscendInputBuffers(InputBuffers):
"""Input buffers for Ascend NPUs."""
def __init__(
self,
max_num_reqs: int,
max_num_tokens: int,
device: torch.device,
):
super().__init__(
max_num_reqs,
max_num_tokens,
device,
)
del self.query_start_loc
# NOTE: For FULL mode we change +1 to +2 to reserve extra space for padding.
# See _pad_query_start_loc_for_fia.
self.query_start_loc: torch.Tensor = torch.zeros(
max_num_reqs + 2,
dtype=torch.int32,
device=device,
)
# Create seq_lens_cpu and seq_lens_np.
# npu's attention backend still needs seq_lens on CPU side.
self.seq_lens_cpu: torch.Tensor = torch.zeros(
max_num_reqs,
dtype=torch.int32,
device="cpu",
)
# seq_len_np and seq_lens_cpu share the same memory.
# define seq_lens_np for easier calculation with numpy.
self.seq_lens_np: np.ndarray = self.seq_lens_cpu.numpy()
@dataclass
class AscendInputBatch(InputBatch):
"""Input batch for Ascend NPUs."""
# Create seq_lens_np.
# npu's attention backend still needs seq_lens on CPU side.
seq_lens_np: np.ndarray
# attn_state is used to build attention metadata.
attn_state: AscendAttentionState | None = None
@classmethod
def make_dummy(
cls,
num_reqs: int,
num_tokens: int,
input_buffers: AscendInputBuffers,
) -> "AscendInputBatch":
"""Override the make_dummy method to calculate seq_lens_np."""
input_batch = InputBatch.make_dummy(
num_reqs,
num_tokens,
input_buffers,
)
# seq_len equals to query_len
input_buffers.seq_lens_np[:num_reqs] = num_tokens // num_reqs
input_buffers.seq_lens_np[num_reqs - 1] += num_tokens % num_reqs
# Pad for full CUDA graph mode.
input_buffers.seq_lens_np[num_reqs:] = 0
seq_lens_np = input_buffers.seq_lens_np[:num_reqs]
input_batch.seq_lens_np = seq_lens_np
# A dummy run for dp or memory profiling.
# When dummy run for dp, num_tokens is set to 1,
# so attn_state is set to DecodeOnly.
# when dummy run for memory profiling,
# attention metadata isn't needed,
# we can also set attn_state to AscendAttentionState.DecodeOnly.
input_batch.attn_state = AscendAttentionState.DecodeOnly
# For mla/sfa, update cos/sin. Here is for _dummy_run.
update_cos_sin(input_batch.positions)
return cls(**asdict(input_batch), seq_lens_np=seq_lens_np)
@triton.jit
def _post_update_kernel(
idx_mapping_ptr,
idx_mapping_stride,
num_computed_tokens_ptr,
last_sampled_tokens_ptr,
output_bin_counts_ptr,
output_bin_counts_stride,
sampled_tokens_ptr,
sampled_tokens_stride,
num_rows,
num_sampled_ptr,
num_rejected_ptr,
query_start_loc_ptr,
all_token_ids_ptr,
all_token_ids_stride,
total_len_ptr,
):
pid = tl.program_id(0)
n_programs = tl.num_programs(0)
rows_per_program = (num_rows + n_programs - 1) // n_programs
start_row = pid * rows_per_program
end_row = tl.minimum(start_row + rows_per_program, num_rows)
for row_idx in range(start_row, end_row):
req_state_idx = tl.load(idx_mapping_ptr + row_idx * idx_mapping_stride)
total_len = tl.load(total_len_ptr + req_state_idx)
num_sampled = tl.load(num_sampled_ptr + row_idx)
if num_sampled > 0:
token_id = tl.load(sampled_tokens_ptr + row_idx * sampled_tokens_stride + num_sampled - 1)
tl.store(last_sampled_tokens_ptr + req_state_idx, token_id)
tl.store(total_len_ptr + req_state_idx, total_len + num_sampled)
for i in range(num_sampled):
token_id = tl.load(sampled_tokens_ptr + row_idx * sampled_tokens_stride + i)
token_ptr = output_bin_counts_ptr + req_state_idx * output_bin_counts_stride + token_id
count = tl.load(token_ptr)
count += 1
tl.store(token_ptr, count)
tl.store(
all_token_ids_ptr + req_state_idx * all_token_ids_stride + total_len + i,
token_id,
)
query_start = tl.load(query_start_loc_ptr + row_idx)
query_end = tl.load(query_start_loc_ptr + row_idx + 1)
query_len = query_end - query_start
num_rejected = tl.load(num_rejected_ptr + row_idx)
num_computed = tl.load(num_computed_tokens_ptr + req_state_idx)
num_computed += query_len - num_rejected
tl.store(num_computed_tokens_ptr + req_state_idx, num_computed)
def post_update(
# [num_reqs]
idx_mapping: torch.Tensor,
# [max_num_reqs]
num_computed_tokens: torch.Tensor,
# [max_num_reqs]
last_sampled_tokens: torch.Tensor,
# [max_num_reqs, vocab_size]
output_bin_counts: torch.Tensor,
# [num_reqs, num_speculative_steps + 1]
sampled_tokens: torch.Tensor,
# [num_reqs]
num_sampled: torch.Tensor,
# [num_reqs]
num_rejected: torch.Tensor,
# [num_reqs + 1]
query_start_loc: torch.Tensor,
# [max_num_reqs, max_model_len]
all_token_ids: torch.Tensor,
# [max_num_reqs]
total_len: torch.Tensor,
) -> None:
num_rows = idx_mapping.shape[0]
core_num = get_vectorcore_num()
grid = (min(num_rows, core_num),)
_post_update_kernel[grid](
idx_mapping,
idx_mapping.stride(0),
num_computed_tokens,
last_sampled_tokens,
output_bin_counts,
output_bin_counts.stride(0),
sampled_tokens,
sampled_tokens.stride(0),
num_rows,
num_sampled,
num_rejected,
query_start_loc,
all_token_ids,
all_token_ids.stride(0),
total_len,
)

View File

@@ -0,0 +1,510 @@
# Adapt from https://github.com/vllm-project/vllm/blob/main/vllm/v1/worker/gpu/model_runner.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
import numpy as np
import torch
from vllm.config import VllmConfig
from vllm.config.compilation import CUDAGraphMode
from vllm.v1.core.sched.output import SchedulerOutput
from vllm.v1.kv_cache_interface import KVCacheConfig
from vllm.v1.worker.gpu import model_runner as vllm_model_runner
from vllm.v1.worker.gpu.buffer_utils import async_copy_to_gpu
from vllm.v1.worker.gpu.cudagraph_utils import BatchExecutionDescriptor
from vllm.v1.worker.gpu.input_batch import (
combine_sampled_and_draft_tokens,
expand_idx_mapping,
prepare_pos_seq_lens,
prepare_prefill_inputs,
)
from vllm.v1.worker.gpu.model_runner import GPUModelRunner
from vllm_ascend.ascend_config import get_ascend_config
from vllm_ascend.ascend_forward_context import (
MoECommType,
get_mc2_tokens_capacity,
override_mrv2_in_profile_run,
select_moe_comm_method,
set_mc2_mask,
set_mc2_tokens_capacity,
)
from vllm_ascend.ops.rotary_embedding import set_cos_and_sin, update_cos_sin
from vllm_ascend.utils import set_weight_prefetch_method
from vllm_ascend.worker.v2.aclgraph_utils import ModelAclGraphManager
from vllm_ascend.worker.v2.attn_utils import build_attn_state
from vllm_ascend.worker.v2.input_batch import AscendInputBatch, AscendInputBuffers
from vllm_ascend.worker.v2.spec_decode.eagle import init_speculator
from vllm_ascend.worker.v2.spec_decode.eagle.speculator import AscendEagleSpeculator
from vllm_ascend.worker.v2.states import AscendRequestState
from vllm_ascend.worker.v2.utils import torch_cuda_wrapper
class NPUModelRunner(GPUModelRunner):
"""Model runner for Ascend NPUs."""
def __init__(self, vllm_config: VllmConfig, device: torch.device):
# Ascend-specific configurations
self.ascend_config = get_ascend_config()
# The following features are not yet supported in Ascend NPU model runner v2:
# - Context parallelism (prefill or decode)
# - Dynamic EPLB
parallel_config = vllm_config.parallel_config
if parallel_config.prefill_context_parallel_size > 1 or parallel_config.decode_context_parallel_size > 1:
raise NotImplementedError("Context parallelism is not supported by Ascend NPU model runner v2.")
if self.ascend_config.eplb_config.dynamic_eplb:
raise NotImplementedError("dynamic_eplb is not supported by Ascend NPU model runner v2.")
with torch_cuda_wrapper():
super().__init__(vllm_config, device)
# because we will override these attribute, delete these attribute to
# make sure it's collected by python gc immediately.
del self.req_states
del self.input_buffers
del self.speculator
# we define AscendEagleSpeculator in vllm_ascend.worker.v2.spec_decode.eagle.speculator
# init_speculator will return AscendEagleSpeculator when eagle is used.
# so here we just call init_speculator to reinitialize speculator.
self.speculator: AscendEagleSpeculator | None = None
if self.speculative_config is not None:
self.speculator = init_speculator(self.vllm_config, self.device)
# AscendRequestState has extra `num_computed_tokens_cpu` attribute.
# so reinitialize req_states here.
self.req_states: AscendRequestState = AscendRequestState(
max_num_reqs=self.max_num_reqs,
max_model_len=self.max_model_len,
max_num_batched_tokens=self.max_num_tokens,
num_speculative_steps=self.num_speculative_steps,
vocab_size=self.vocab_size,
device=self.device,
)
# 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=self.device,
)
# we need to copy num_computed_tokens back to cpu to help
# update actual seq_lens_cpu. gpu attention backend doesn't need these
# attributes, cause their attention backends doesn't use seq_lens_cpu.
# and seq_lens_cpu is deprecated in gpu_model_runner_v2.
self.num_computed_tokens_event = torch.npu.Event()
self.num_computed_tokens_stream = torch.npu.Stream()
self.num_computed_tokens_cpu = torch.empty(
self.max_num_reqs,
dtype=torch.int32,
device="cpu",
pin_memory=True,
)
# set _WEIGHT_PREFETCH_METHOD, _mc2_tokens_capacity and _reserved_mc2_mask which
# is necessary for weight_prfetching function, and MoE communication optimization.
set_weight_prefetch_method(self.ascend_config.weight_prefetch_config)
# TODO: remove set_cos_and_sin (together with update_cos_sin) when mla can properly handle cos/sin internally
self.decode_query_len = self.num_speculative_steps + 1
set_cos_and_sin(vllm_config, self.max_num_reqs, self.decode_query_len, self.dtype, self.device)
set_mc2_tokens_capacity(vllm_config, self.max_num_reqs, self.decode_query_len)
set_mc2_mask(vllm_config, self.device)
# we need to update full graph params in run_fullgraph,
# so create a stream to update full graph params.
if self.compilation_config.cudagraph_mode.has_full_cudagraphs():
self.update_stream: torch.npu.Stream = torch.npu.Stream()
# we need to use return value of `get_cudagraph_and_dp_padding`
# to set forward_context in `run_fullgraph`.
# so we can inherit `execute_model` method.
self.cudagraph_and_dp_padding: tuple[int, torch.Tensor | None, int] | None = None
# we need to use input_batch to set forward_context in run_fullgraph.
# so we can inherit `execute_model` method.
self.input_batch: AscendInputBatch | None = None
def initialize_kv_cache(self, kv_cache_config: KVCacheConfig) -> None:
with graph_manager_wrapper(self):
super().initialize_kv_cache(kv_cache_config)
@torch.inference_mode()
def profile_run(self) -> None:
"""Override GPUModelRunner.profile_run for Ascend NPUs.
When running moe models, we need an extra dummy run with mc2_tokens_capacity tokens to reserve
necessary HCCL buffer for the MC2 operator before standard `profile_run`. Additionally, we set
override_mrv2_in_profile_run to True to force moe load to be balanced when executing `profile_run`
"""
mc2_tokens_capacity = get_mc2_tokens_capacity()
with override_mrv2_in_profile_run(True):
if (
mc2_tokens_capacity is not None
and self.max_num_tokens > mc2_tokens_capacity
and select_moe_comm_method(mc2_tokens_capacity, self.vllm_config)
in {MoECommType.MC2, MoECommType.FUSED_MC2}
):
self._dummy_run(mc2_tokens_capacity, skip_attn=True, is_profile=True)
super().profile_run()
def prepare_inputs(
self,
scheduler_output: SchedulerOutput,
batch_desc: BatchExecutionDescriptor,
) -> AscendInputBatch:
"""Override GPUModelRunner.prepare_inputs for Ascend NPUs.
npu attention backends need seq_lens_cpu to work.
so we need to prepare seq_lens_cpu here.
"""
num_tokens = scheduler_output.total_num_scheduled_tokens
num_tokens_after_padding = batch_desc.num_tokens
assert num_tokens > 0
num_tokens_per_req = scheduler_output.num_scheduled_tokens
num_reqs = len(num_tokens_per_req)
# Decode first, then prefill.
# batch_idx -> req_id
req_ids = sorted(num_tokens_per_req, key=num_tokens_per_req.get) # type: ignore
self._update_seq_lens_cpu(scheduler_output, req_ids)
numtoks_iter = map(num_tokens_per_req.get, req_ids)
num_scheduled_tokens = np.fromiter(numtoks_iter, dtype=np.int32, count=num_reqs)
num_valid_tokens = num_scheduled_tokens
if scheduler_output.scheduled_spec_decode_tokens:
num_valid_tokens = np.array(
[
num_tokens - len(scheduler_output.scheduled_spec_decode_tokens.get(i, []))
for num_tokens, i in zip(num_scheduled_tokens, req_ids)
],
dtype=np.int32,
)
attn_state = build_attn_state(
self.vllm_config,
self.input_buffers.seq_lens_np,
num_reqs,
num_scheduled_tokens,
num_valid_tokens,
)
idx_mapping_iter = map(self.req_states.req_id_to_index.get, req_ids)
idx_mapping_np = np.fromiter(idx_mapping_iter, dtype=np.int32, count=num_reqs)
idx_mapping_cpu = torch.from_numpy(idx_mapping_np)
idx_mapping = async_copy_to_gpu(idx_mapping_cpu, device=self.device)
# Get the number of draft tokens for each request.
draft_tokens = scheduler_output.scheduled_spec_decode_tokens
num_draft_tokens_per_req: np.ndarray | None = None
if not draft_tokens:
# No draft token scheduled (common case).
total_num_draft_tokens = 0
total_num_logits = num_reqs
cu_num_logits_np = np.arange(num_reqs + 1, dtype=np.int32)
cu_num_logits = torch.arange(num_reqs + 1, device=self.device, dtype=torch.int32)
expanded_idx_mapping = idx_mapping
expanded_local_pos = torch.zeros(num_reqs, dtype=torch.int32, device=self.device)
else:
num_draft_tokens_arr = np.array(
[len(draft_tokens.get(req_id, ())) for req_id in req_ids],
dtype=np.int32,
)
num_draft_tokens_per_req = num_draft_tokens_arr
total_num_draft_tokens = int(num_draft_tokens_arr.sum())
total_num_logits = num_reqs + total_num_draft_tokens
num_logits = num_draft_tokens_arr + 1
cu_num_logits_np = np.empty(num_reqs + 1, dtype=np.int32)
cu_num_logits_np[0] = 0
np.cumsum(num_logits, out=cu_num_logits_np[1:])
cu_num_logits = async_copy_to_gpu(cu_num_logits_np, device=self.device)
max_expand_len = self.num_speculative_steps + 1
expanded_idx_mapping, expanded_local_pos = expand_idx_mapping(
idx_mapping, total_num_logits, cu_num_logits, max_expand_len
)
# Get query_start_loc.
# NOTE: For FULL mode we change +1 to +2 to reserve extra space for padding.
# See _pad_query_start_loc_for_fia.
num_reqs_padded = batch_desc.num_reqs or num_reqs
query_start_loc_np = np.empty(self.max_num_reqs + 2, dtype=np.int32)
query_start_loc_np[0] = 0
np.cumsum(num_scheduled_tokens, out=query_start_loc_np[1 : num_reqs + 1])
# Pad for full CUDA graph mode.
# Some attention backends like FA3 require query_start_loc to be non-decreasing.
query_start_loc_np[num_reqs + 1 :] = num_tokens
if batch_desc.cg_mode == CUDAGraphMode.FULL:
# This is only required for vllm-ascend.
query_start_loc_np, num_reqs_padded = self._pad_query_start_loc_for_fia(
num_tokens_after_padding,
num_reqs_padded,
num_reqs,
query_start_loc_np,
batch_desc.cg_mode,
batch_desc.num_reqs,
)
async_copy_to_gpu(query_start_loc_np, out=self.input_buffers.query_start_loc)
query_start_loc_np = query_start_loc_np[: num_reqs_padded + 1]
query_start_loc = self.input_buffers.query_start_loc[: num_reqs + 1]
prefill_len_np = self.req_states.prefill_len.np[idx_mapping_np]
num_computed_prefill_tokens_np = self.req_states.num_computed_prefill_tokens[idx_mapping_np]
is_prefilling_np = num_computed_prefill_tokens_np < prefill_len_np
# Get prefill tokens if any.
if np.any(is_prefilling_np):
prepare_prefill_inputs(
self.input_buffers.input_ids,
self.req_states.next_prefill_tokens,
idx_mapping,
query_start_loc,
self.req_states.all_token_ids.gpu,
self.req_states.prefill_len.gpu,
self.req_states.num_computed_tokens.gpu,
)
# Prepare positions and seq_lens.
prepare_pos_seq_lens(
idx_mapping,
query_start_loc,
self.req_states.num_computed_tokens.gpu,
self.input_buffers.positions,
self.input_buffers.seq_lens,
)
seq_lens = self.input_buffers.seq_lens[:num_reqs]
# Pad for full CUDA graph mode.
self.input_buffers.seq_lens_np[num_reqs_padded:] = 0
# Some input token ids are directly read from the last sampled tokens
# and draft tokens. Also, get the logits indices to sample tokens from.
logits_indices = combine_sampled_and_draft_tokens(
self.input_buffers.input_ids,
idx_mapping,
self.req_states.last_sampled_tokens,
query_start_loc,
seq_lens,
self.req_states.prefill_len.gpu,
self.req_states.draft_tokens,
cu_num_logits,
total_num_logits,
)
input_ids = self.input_buffers.input_ids[:num_tokens_after_padding]
positions = self.input_buffers.positions[:num_tokens_after_padding]
# CPU upper bound on seq_lens (num_computed_tokens + num_scheduled_tokens).
# Added by vLLM PR #40654 to avoid GPU->CPU sync for seq_lens.
seq_lens_cpu_upper_bound_np = np.zeros(num_reqs_padded, dtype=np.int32)
np.add(
self.req_states.num_computed_tokens_np[idx_mapping_np],
num_scheduled_tokens,
out=seq_lens_cpu_upper_bound_np[:num_reqs],
)
seq_lens_cpu_upper_bound = torch.from_numpy(seq_lens_cpu_upper_bound_np)
num_computed_tokens_np = self.req_states.num_computed_tokens_np[idx_mapping_np]
max_seq_len_np = None
if getattr(self, "use_pp", False):
# max_seq_len is only consumed by the PP `compute_need_sampled_mask`.
max_seq_len_np = self.req_states.max_seq_len[idx_mapping_np]
self.input_batch = AscendInputBatch(
req_ids=req_ids,
num_reqs=num_reqs,
num_reqs_after_padding=num_reqs_padded,
idx_mapping=idx_mapping,
idx_mapping_np=idx_mapping_np,
expanded_idx_mapping=expanded_idx_mapping,
expanded_local_pos=expanded_local_pos,
num_scheduled_tokens=num_scheduled_tokens,
num_tokens=num_tokens,
num_tokens_after_padding=num_tokens_after_padding,
num_draft_tokens=total_num_draft_tokens,
num_draft_tokens_per_req=num_draft_tokens_per_req,
query_start_loc=query_start_loc,
query_start_loc_np=query_start_loc_np,
seq_lens=seq_lens,
seq_lens_cpu_upper_bound=seq_lens_cpu_upper_bound,
dcp_local_seq_lens=None, # TODO(Ronald1995): support cp.
is_prefilling_np=is_prefilling_np,
num_computed_tokens_np=num_computed_tokens_np,
prefill_len_np=prefill_len_np,
num_computed_prefill_tokens_np=num_computed_prefill_tokens_np,
max_seq_len_np=max_seq_len_np,
input_ids=input_ids,
positions=positions,
logits_indices=logits_indices,
cu_num_logits=cu_num_logits,
cu_num_logits_np=cu_num_logits_np,
has_structured_output_reqs=scheduler_output.has_structured_output_requests,
# extra attributes for ascend npus.
seq_lens_np=self.input_buffers.seq_lens_np,
attn_state=attn_state,
)
# For mla/sfa, update cos/sin. Here is for execute_model.
update_cos_sin(self.input_batch.positions)
return self.input_batch
def postprocess(
self,
input_batch,
sampled_tokens,
num_sampled,
num_rejected,
):
"""Override GPUModelRunner.postprocess for Ascend NPUs.
npu attention backends need seq_lens_cpu to work.
so we need to copy num_computed_tokens back to cpu here.
"""
super().postprocess(
input_batch,
sampled_tokens,
num_sampled,
num_rejected,
)
self._copy_num_computed_tokens_to_cpu()
def postprocess_sampled(
self,
idx_mapping,
sampled_tokens,
num_sampled,
num_rejected,
query_start_loc=None,
):
"""Override GPUModelRunner.postprocess_sampled for Ascend NPUs."""
super().postprocess_sampled(
idx_mapping,
sampled_tokens,
num_sampled,
num_rejected,
query_start_loc,
)
self._copy_num_computed_tokens_to_cpu()
def _copy_num_computed_tokens_to_cpu(self):
# npu attention backend still need to use seq_lens_cpu,
# we need to copy num_computed_tokens back to cpu.
default_stream = torch.cuda.current_stream()
assert self.num_computed_tokens_stream is not None
assert self.num_computed_tokens_cpu is not None
with torch.npu.stream(self.num_computed_tokens_stream):
self.num_computed_tokens_stream.wait_stream(default_stream)
self.num_computed_tokens_cpu.copy_(
self.req_states.num_computed_tokens.gpu,
non_blocking=True,
)
self.num_computed_tokens_event.record()
def _update_seq_lens_cpu(
self,
scheduler_output: SchedulerOutput,
req_ids: list[str],
):
num_scheduled_tokens = scheduler_output.num_scheduled_tokens
# wait for num_computed_tokens copy to cpu stream to finish.
self.num_computed_tokens_event.synchronize()
for req_id in scheduler_output.scheduled_cached_reqs.req_ids:
req_index = self.req_states.req_id_to_index[req_id]
# num_computed_tokens_cpu has reverted by num_rejected_tokens already.
# in super postprocess method.
self.req_states.num_computed_tokens_cpu[req_index] = self.num_computed_tokens_cpu[req_index]
# update seq_lens_cpu
for i, req_id in enumerate(req_ids): # type: ignore
req_index = self.req_states.req_id_to_index[req_id]
num_computed_tokens = self.req_states.num_computed_tokens_cpu[req_index]
self.input_buffers.seq_lens_cpu[i] = num_computed_tokens + num_scheduled_tokens[req_id]
def eplb_warmup(self):
# TODO(Ronald1995): just define the method in case calling error in
# worker, implement it in the future.
pass
def _pad_query_start_loc_for_fia(
self,
num_tokens_padded: int,
num_reqs_padded: int,
num_reqs: int,
query_start_loc_np: np.ndarray,
cudagraph_runtime_mode: CUDAGraphMode | None = None,
batch_desc_num_reqs: int | None = None,
) -> tuple[np.ndarray, int]:
"""
This function is only designed to satisfied the constraint that when the layout is TND,
the first dimension of `hidden_states` must equal the last element of `actual_seq_lengths_q`.
"""
# TODO: need refactor later, related to vllm PR #34043 this pr delete func
# relax_for_mixed_batch_cudagraphs, num_reqs no longer equals the actual number of requests.
if cudagraph_runtime_mode == CUDAGraphMode.FULL:
num_reqs_padded = num_reqs
else:
num_reqs_padded = batch_desc_num_reqs if batch_desc_num_reqs is not None else num_reqs
if num_tokens_padded == num_reqs_padded * self.decode_query_len:
# Uniform-batch case: num_reqs must be no greater than num_reqs_padded
assert num_reqs <= num_reqs_padded
last_loc = query_start_loc_np[num_reqs]
query_start_loc_np[num_reqs + 1 : num_reqs_padded + 1] = (
np.arange(1, num_reqs_padded + 1 - num_reqs) * self.decode_query_len + last_loc
)
else:
# Mixed-batch case: num_reqs must equal num_reqs_padded
assert num_reqs == num_reqs_padded
# Insert a dummy request instead of setting query_start_loc[num_reqs] = num_tokens_padded directly
query_start_loc_np[num_reqs_padded + 1] = num_tokens_padded
num_reqs_padded = num_reqs_padded + 1
return query_start_loc_np, num_reqs_padded
@contextmanager
def graph_manager_wrapper(model_runner):
"""Context manager to override graph manager."""
original_graph_manager = vllm_model_runner.ModelCudaGraphManager
def factory(
vllm_config: VllmConfig,
device: torch.device,
cudagraph_mode: CUDAGraphMode,
decode_query_len: int,
lora_capture_cases: list[int] | None = None,
):
return ModelAclGraphManager(
vllm_config,
device,
cudagraph_mode,
decode_query_len,
model_runner,
lora_capture_cases=lora_capture_cases,
)
try:
vllm_model_runner.ModelCudaGraphManager = factory
yield
finally:
vllm_model_runner.ModelCudaGraphManager = original_graph_manager

View File

@@ -0,0 +1,34 @@
# Adapt from https://github.com/vllm-project/vllm/blob/main/vllm/v1/worker/gpu/model_states/__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
import torch.nn as nn
from vllm.config import VllmConfig
from vllm.v1.worker.gpu.mm.encoder_cache import EncoderCache
def init_asecnd_model_state(
vllm_config: VllmConfig,
model: nn.Module,
encoder_cache: EncoderCache | None,
device: torch.device,
):
from vllm_ascend.worker.v2.model_states.default import AscendModelState
return AscendModelState(vllm_config, model, encoder_cache, device)

View File

@@ -0,0 +1,77 @@
# Adapt from https://github.com/vllm-project/vllm/blob/main/vllm/v1/worker/gpu/model_states/default.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 typing import Any
import torch
from vllm.config.compilation import CUDAGraphMode
from vllm.v1.kv_cache_interface import KVCacheConfig
from vllm.v1.worker.gpu.model_states.default import DefaultModelState
from vllm.v1.worker.utils import AttentionGroup
from vllm_ascend.worker.v2.attn_utils import build_attn_metadata
from vllm_ascend.worker.v2.input_batch import AscendInputBatch
class AscendModelState(DefaultModelState):
"""Model state for Ascend NPUs."""
def prepare_attn(
self,
input_batch: AscendInputBatch,
cudagraph_mode: CUDAGraphMode,
block_tables: tuple[torch.Tensor, ...],
slot_mappings: torch.Tensor,
attn_groups: list[list[AttentionGroup]],
kv_cache_config: KVCacheConfig,
for_capture: bool = False,
) -> dict[str, Any]:
"""Override prepare_attn method because `build_attn_metadata` is different from vllm."""
if cudagraph_mode == CUDAGraphMode.FULL:
# Use padded sizes - padding is handled by model_runner.prepare_attn.
num_reqs = input_batch.num_reqs_after_padding
num_tokens = input_batch.num_tokens_after_padding
else:
# For piecewise cudagraphs and eager, use unpadded sizes.
num_reqs = input_batch.num_reqs
num_tokens = input_batch.num_tokens
query_start_loc_cpu = torch.from_numpy(input_batch.query_start_loc_np)
max_query_len = input_batch.num_scheduled_tokens.max().item()
# attn_metadata is needed when update_full_graph_params, but no way can get it now.
# Temporarily store it in model_state.
self.attn_metadata = build_attn_metadata(
attn_groups=attn_groups,
num_reqs=num_reqs,
num_tokens=num_tokens,
query_start_loc_gpu=input_batch.query_start_loc,
query_start_loc_cpu=query_start_loc_cpu,
max_query_len=max_query_len,
seq_lens=input_batch.seq_lens,
max_seq_len=self.max_model_len,
block_tables=block_tables,
slot_mappings=slot_mappings,
kv_cache_config=kv_cache_config,
dcp_local_seq_lens=input_batch.dcp_local_seq_lens,
# extra attributes for ascend npus.
seq_lens_np=input_batch.seq_lens_np,
positions=input_batch.positions,
attn_state=input_batch.attn_state,
for_cudagraph_capture=for_capture,
)
return self.attn_metadata

View File

View File

@@ -0,0 +1,183 @@
# Adapt from https://github.com/vllm-project/vllm/blob/main/vllm/v1/worker/gpu/sample/bad_words.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_ascend.ops.triton.triton_utils import get_vectorcore_num
MAX_BAD_WORDS_TOTAL_TOKENS = 1024 # Max total tokens for all bad words per request
MAX_NUM_BAD_WORDS = 128 # Max number of bad words per request
@triton.jit(do_not_specialize=["num_tokens", "max_num_bad_words"])
def _bad_words_kernel(
logits_ptr,
logits_stride,
expanded_idx_mapping_ptr,
bad_word_token_ids_ptr,
bad_word_token_ids_stride,
bad_word_offsets_ptr,
bad_word_offsets_stride,
num_bad_words_ptr,
all_token_ids_ptr,
all_token_ids_stride,
prompt_len_ptr,
total_len_ptr,
input_ids_ptr,
expanded_local_pos_ptr,
num_tokens,
max_num_bad_words,
MAX_PREFIX_LEN: tl.constexpr,
):
"""
Optimized bad words filtering kernel for Ascend NPU.
Key optimizations:
- Optimized memory access patterns
- Reduced redundant calculations
- Enhanced data locality
- Minimized conditional branches
- Improved load balancing
"""
pid = tl.program_id(0)
num_cores = tl.num_programs(0)
# Calculate tokens per core for better load balancing
tokens_per_core = (num_tokens + num_cores - 1) // num_cores
start_token = pid * tokens_per_core
end_token = min(start_token + tokens_per_core, num_tokens)
# Process each token assigned to this core
for token_idx in range(start_token, end_token):
# Load request state index
req_state_idx = tl.load(expanded_idx_mapping_ptr + token_idx)
num_bad_words = tl.load(num_bad_words_ptr + req_state_idx)
# Only process if there are bad words for this request
if num_bad_words > 0:
# Load position information
pos = tl.load(expanded_local_pos_ptr + token_idx)
cur_req_first_pos = token_idx - pos
# Load length information
prompt_len = tl.load(prompt_len_ptr + req_state_idx)
total_len = tl.load(total_len_ptr + req_state_idx)
output_len = total_len - prompt_len
effective_len = output_len + pos
# Precompute base addresses
bd_offsets_base = bad_word_offsets_ptr + req_state_idx * bad_word_offsets_stride
bd_tokens_base = bad_word_token_ids_ptr + req_state_idx * bad_word_token_ids_stride
output_base = all_token_ids_ptr + req_state_idx * all_token_ids_stride + prompt_len
# Process each bad word for this token
for bw_idx in range(max_num_bad_words):
if bw_idx < num_bad_words:
# Load bad word range
start = tl.load(bd_offsets_base + bw_idx)
end = tl.load(bd_offsets_base + bw_idx + 1)
bad_word_len = end - start
prefix_len = bad_word_len - 1
# Check prefix length validity
if prefix_len <= effective_len:
# Load last token
last_token = tl.load(bd_tokens_base + end - 1)
# Match checking with early termination
match = 1
j = 0
while j < prefix_len and match:
# Load expected token
expected = tl.load(bd_tokens_base + start + j)
# Calculate actual position and load actual token
actual_pos = effective_len - prefix_len + j
if actual_pos >= output_len:
spec_offset = actual_pos - output_len
actual = tl.load(input_ids_ptr + cur_req_first_pos + spec_offset)
else:
actual = tl.load(output_base + actual_pos)
# Check for mismatch
if expected != actual:
match = 0
j += 1
# Store result if match found
if match:
tl.store(logits_ptr + token_idx * logits_stride + last_token, -float("inf"))
def apply_bad_words(
logits: torch.Tensor,
expanded_idx_mapping: torch.Tensor,
bad_word_token_ids: torch.Tensor,
bad_word_offsets: torch.Tensor,
num_bad_words: torch.Tensor,
all_token_ids: torch.Tensor,
prompt_len: torch.Tensor,
total_len: torch.Tensor,
input_ids: torch.Tensor,
expanded_local_pos: torch.Tensor,
max_num_bad_words: int,
) -> None:
"""
Apply bad words filtering to logits.
Args:
logits: [num_tokens, vocab_size] - Model output logits
expanded_idx_mapping: [num_tokens] - Token to request mapping
bad_word_token_ids: [max_num_reqs, MAX_BAD_WORDS_TOTAL_TOKENS] - Bad word token IDs
bad_word_offsets: [max_num_reqs, MAX_NUM_BAD_WORDS + 1] - Bad word offsets
num_bad_words: [max_num_reqs] - Number of bad words per request
all_token_ids: [max_num_reqs, max_seq_len] - All token IDs
prompt_len: [max_num_reqs] - Prompt length
total_len: [max_num_reqs] - Total length
input_ids: [num_tokens] - Input IDs
expanded_local_pos: [num_tokens] - Expanded local position
max_num_bad_words: Maximum number of bad words to check
"""
num_tokens = logits.shape[0]
core_num = get_vectorcore_num()
MAX_PREFIX_LEN = 32
_bad_words_kernel[(core_num,)](
logits,
logits.stride(0),
expanded_idx_mapping,
bad_word_token_ids,
bad_word_token_ids.stride(0),
bad_word_offsets,
bad_word_offsets.stride(0),
num_bad_words,
all_token_ids,
all_token_ids.stride(0),
prompt_len,
total_len,
input_ids,
expanded_local_pos,
num_tokens,
max_num_bad_words,
MAX_PREFIX_LEN,
)

View File

@@ -0,0 +1,205 @@
# Adapt from https://github.com/vllm-project/vllm/blob/main/vllm/v1/worker/gpu/sample/gumbel.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
@triton.jit(do_not_specialize=["logits_stride", "vocab_size"])
def _temperature_kernel(
logits_ptr,
logits_stride,
expanded_idx_mapping_ptr,
temperature_ptr,
vocab_size,
BLOCK_SIZE: tl.constexpr,
):
token_idx = tl.program_id(0)
req_state_idx = tl.load(expanded_idx_mapping_ptr + token_idx)
temperature = tl.load(temperature_ptr + req_state_idx).to(tl.float32)
if temperature == 0.0 or temperature == 1.0:
# Early return to avoid loading logits
return
block_idx = tl.program_id(1)
block = block_idx * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
mask = block < vocab_size
logits = tl.load(logits_ptr + token_idx * logits_stride + block, mask=mask)
logits = logits.to(tl.float32)
logits = logits / temperature
tl.store(logits_ptr + token_idx * logits_stride + block, logits, mask=mask)
def apply_temperature(
logits: torch.Tensor,
expanded_idx_mapping: torch.Tensor,
temperature: torch.Tensor,
) -> None:
"""
Args:
logits: Tensor of shape (num_tokens, vocab_size) containing the logits.
expanded_idx_mapping: Tensor containing the mapping from token index
to request index of tensor temperature.
temperature: Tensor containing the temperature value for each request.
"""
num_tokens, vocab_size = logits.shape
BLOCK_SIZE = 44032
num_blocks = triton.cdiv(vocab_size, BLOCK_SIZE)
_temperature_kernel[(num_tokens, num_blocks)](
logits,
logits.stride(0),
expanded_idx_mapping,
temperature,
vocab_size,
BLOCK_SIZE=BLOCK_SIZE,
multibuffer=False,
)
@triton.jit(
do_not_specialize=[
"local_argmax_stride",
"local_max_stride",
"processed_logits_stride",
"logits_stride",
"vocab_size",
]
)
def _gumbel_sample_kernel(
local_argmax_ptr,
local_argmax_stride,
local_max_ptr,
local_max_stride,
processed_logits_ptr,
processed_logits_stride,
processed_logits_col_ptr,
logits_ptr,
logits_stride,
expanded_idx_mapping_ptr,
seeds_ptr,
pos_ptr,
temp_ptr,
vocab_size,
BLOCK_SIZE: tl.constexpr,
APPLY_TEMPERATURE: tl.constexpr,
):
token_idx = tl.program_id(0)
block_idx = tl.program_id(1)
block = block_idx * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
mask = block < vocab_size
logits = tl.load(
logits_ptr + token_idx * logits_stride + block,
mask=mask,
other=float("-inf"),
)
logits = logits.to(tl.float32)
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:
# NOTE(woosuk): Match the behavior of _temperature_kernel.
logits = logits / temp
if processed_logits_ptr is not None:
# Store the temperature-applied logits.
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,
)
if temp != 0.0:
# Calculate the seed for gumbel noise.
seed = tl.load(seeds_ptr + req_state_idx)
# NOTE(Ronald1995): change pos's dtype to tl.int32, because triton-ascend's
# compiler doesn't support uint64 of pos arg.
pos = tl.load(pos_ptr + token_idx).to(tl.int32)
gumbel_seed = tl.randint(seed, pos)
# NOTE(Ronald1995): r is tl.float64 in vllm, change it to tl.float32,
# because triton-ascend's compiler does not support float64.
r = tl.rand(gumbel_seed, block).to(tl.float32)
gumbel_noise = -tl.log(-tl.log(r + 1e-20) + 1e-20)
# Apply gumbel noise.
logits = tl.where(mask, logits + gumbel_noise, float("-inf"))
idx = tl.argmax(logits, axis=0)
token_id = block_idx * BLOCK_SIZE + idx
value = tl.max(logits, axis=0)
tl.store(local_argmax_ptr + token_idx * local_argmax_stride + block_idx, token_id)
tl.store(local_max_ptr + token_idx * local_max_stride + block_idx, value)
def gumbel_sample(
logits: torch.Tensor, # [num_tokens, vocab_size]
expanded_idx_mapping: torch.Tensor, # [num_tokens]
temperature: torch.Tensor, # [max_num_reqs]
seed: torch.Tensor, # [max_num_reqs]
pos: torch.Tensor, # [num_tokens]
apply_temperature: bool,
output_processed_logits: torch.Tensor | None = None,
output_processed_logits_col: torch.Tensor | None = None,
use_fp64: bool = False,
) -> torch.Tensor:
if use_fp64:
raise NotImplementedError("FP64 Gumbel sampling is not supported on NPU.")
num_tokens, vocab_size = logits.shape
BLOCK_SIZE = 1024
num_blocks = triton.cdiv(vocab_size, BLOCK_SIZE)
local_argmax = torch.empty(
num_tokens,
num_blocks,
dtype=torch.int64,
device=logits.device,
)
local_max = torch.empty(
num_tokens,
num_blocks,
dtype=torch.float32,
device=logits.device,
)
_gumbel_sample_kernel[(num_tokens, num_blocks)](
local_argmax,
local_argmax.stride(0),
local_max,
local_max.stride(0),
output_processed_logits,
output_processed_logits.stride(0) if output_processed_logits is not None else 0,
output_processed_logits_col,
logits,
logits.stride(0),
expanded_idx_mapping,
seed,
pos,
temperature,
vocab_size,
BLOCK_SIZE=BLOCK_SIZE,
APPLY_TEMPERATURE=apply_temperature,
)
# NOTE(woosuk): Use int64 for later indexing.
max_block_idx = local_max.argmax(dim=-1, keepdim=True)
sampled = local_argmax.gather(dim=-1, index=max_block_idx).view(-1)
return sampled

View File

@@ -0,0 +1,164 @@
# Adapt from https://github.com/vllm-project/vllm/blob/main/vllm/v1/worker/gpu/sample/logprob.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.outputs import LogprobsTensors
from vllm_ascend.ops.triton.triton_utils import get_vectorcore_num
@triton.jit
def _topk_log_softmax_kernel(
output_ptr,
logits_ptr,
logits_stride,
topk_ids_ptr,
topk,
vocab_size,
BLOCK_SIZE: tl.constexpr,
PADDED_TOPK: tl.constexpr,
):
req_idx = tl.program_id(0)
row_ptr = logits_ptr + req_idx * logits_stride
max_val = float("-inf")
for i in range(0, vocab_size, BLOCK_SIZE):
block = i + tl.arange(0, BLOCK_SIZE)
logits = tl.load(row_ptr + block, mask=block < vocab_size, other=float("-inf"))
max_val = tl.max(tl.maximum(logits, max_val, propagate_nan=tl.PropagateNan.ALL))
max_val = max_val.to(tl.float32) # type: ignore
se = 0.0
for i in range(0, vocab_size, BLOCK_SIZE):
block = i + tl.arange(0, BLOCK_SIZE)
logits = tl.load(row_ptr + block, mask=block < vocab_size, other=float("-inf"))
logits = logits.to(tl.float32)
e = tl.exp(logits - max_val)
se += tl.sum(e)
lse = tl.log(se)
k_offset = tl.arange(0, PADDED_TOPK)
k_mask = k_offset < topk
topk_ids = tl.load(topk_ids_ptr + req_idx * topk + k_offset, mask=k_mask, other=0)
logits = tl.load(row_ptr + topk_ids, mask=k_mask)
logits = logits.to(tl.float32)
o = logits - lse - max_val
tl.store(output_ptr + req_idx * topk + k_offset, o, mask=k_mask)
def compute_token_logprobs(logits: torch.Tensor, token_ids: torch.Tensor) -> torch.Tensor:
batch_size, vocab_size = logits.shape
token_ids = token_ids.to(torch.int64)
num_logprobs = token_ids.shape[1]
logprobs = logits.new_empty((batch_size, num_logprobs), dtype=torch.float32)
_topk_log_softmax_kernel[(batch_size,)](
logprobs,
logits,
logits.stride(0),
token_ids,
num_logprobs,
vocab_size,
BLOCK_SIZE=12944,
PADDED_TOPK=max(triton.next_power_of_2(num_logprobs), 2),
multibuffer=False,
)
return logprobs
@triton.jit(do_not_specialize=["batch_size", "rows_per_core"])
def _ranks_kernel(
output_ptr,
logits_ptr,
logits_stride,
token_ids_ptr,
vocab_size,
batch_size,
rows_per_core,
BLOCK_SIZE: tl.constexpr,
):
core_id = tl.program_id(0)
start_row = core_id * rows_per_core
end_row = start_row + rows_per_core
for req_idx in range(start_row, end_row):
if req_idx < batch_size:
row_ptr = logits_ptr + req_idx * logits_stride
token_id = tl.load(token_ids_ptr + req_idx)
x = tl.load(row_ptr + token_id)
n_vec = tl.zeros([BLOCK_SIZE], dtype=tl.int32)
for i in range(0, vocab_size, BLOCK_SIZE):
block = i + tl.arange(0, BLOCK_SIZE)
logits = tl.load(row_ptr + block, mask=block < vocab_size, other=float("-inf"))
n_vec += (logits > x).to(tl.int32)
n = tl.sum(n_vec)
tl.store(output_ptr + req_idx, n)
def compute_topk_logprobs(
logits: torch.Tensor,
num_logprobs: int,
sampled_token_ids: torch.Tensor,
cu_num_logits: list[int] | None = None,
) -> LogprobsTensors:
assert num_logprobs >= 0
batch_size, vocab_size = logits.shape
logprob_token_ids = sampled_token_ids.unsqueeze(-1)
if num_logprobs > 0:
topk_indices = torch.topk(logits, num_logprobs, dim=-1).indices
logprob_token_ids = torch.cat((sampled_token_ids.unsqueeze(-1), topk_indices), dim=1)
# NOTE(woosuk): Here, to save GPU memory, we do not materialize the full
# logprobs tensor. Instead, we only compute and return the logprobs of
# the topk + 1 tokens.
logprobs = compute_token_logprobs(logits, logprob_token_ids)
token_ranks = torch.empty(
batch_size,
dtype=torch.int64,
device=logits.device,
)
vec_core = get_vectorcore_num()
NUM_CORES = min(batch_size, vec_core)
rows_per_core = triton.cdiv(batch_size, NUM_CORES)
BLOCK_SIZE = 8192
grid = (NUM_CORES,)
_ranks_kernel[grid](
token_ranks,
logits,
logits.stride(0),
sampled_token_ids,
vocab_size,
batch_size,
rows_per_core,
BLOCK_SIZE=BLOCK_SIZE,
multibuffer=False,
)
return LogprobsTensors(
logprob_token_ids=logprob_token_ids,
logprobs=logprobs,
selected_token_ranks=token_ranks,
cu_num_generated_tokens=cu_num_logits,
)

View File

@@ -0,0 +1,91 @@
# Adapt from https://github.com/vllm-project/vllm/blob/main/vllm/v1/worker/gpu/sample/min_p.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_ascend.ops.triton.triton_utils import get_vectorcore_num
@triton.jit(do_not_specialize=["num_tokens"])
def _min_p_kernel(
in_logits_ptr,
out_logits_ptr,
logits_stride,
expanded_idx_mapping_ptr,
min_p_ptr,
vocab_size,
num_tokens,
BLOCK_SIZE: tl.constexpr,
):
pid = tl.program_id(0)
core_num = tl.num_programs(0)
tokens_per_block = (num_tokens + core_num - 1) // core_num
start_token = pid * tokens_per_block
end_token = tl.minimum(start_token + tokens_per_block, num_tokens)
for token_idx in tl.range(start_token, end_token):
req_state_idx = tl.load(expanded_idx_mapping_ptr + token_idx)
min_p = tl.load(min_p_ptr + req_state_idx).to(tl.float32)
if min_p != 0.0:
max_val = float("-inf")
for i in range(0, vocab_size, BLOCK_SIZE):
block = i + tl.arange(0, BLOCK_SIZE)
mask = block < vocab_size
logits = tl.load(
in_logits_ptr + token_idx * logits_stride + block,
mask=mask,
other=float("-inf"),
)
max_val = tl.max(tl.maximum(logits, max_val))
max_val = max_val.to(tl.float32) # type: ignore
threshold = max_val + tl.log(min_p)
for i in range(0, vocab_size, BLOCK_SIZE):
block = i + tl.arange(0, BLOCK_SIZE)
mask = block < vocab_size
logits = tl.load(
in_logits_ptr + token_idx * logits_stride + block,
mask=mask,
other=float("-inf"),
)
logits = tl.where(logits < threshold, float("-inf"), logits)
tl.store(out_logits_ptr + token_idx * logits_stride + block, logits, mask=mask)
def apply_min_p(logits: torch.Tensor, expanded_idx_mapping: torch.Tensor, min_p: torch.Tensor) -> None:
num_tokens, vocab_size = logits.shape
vec_core = get_vectorcore_num()
core_nums = min(num_tokens, vec_core)
BLOCK_SIZE = min(triton.next_power_of_2(vocab_size), 8192)
_min_p_kernel[(core_nums,)](
logits,
logits,
logits.stride(0),
expanded_idx_mapping,
min_p,
vocab_size,
num_tokens,
BLOCK_SIZE=BLOCK_SIZE,
multibuffer=False,
)

View File

@@ -0,0 +1,215 @@
# Adapt from https://github.com/vllm-project/vllm/blob/main/vllm/v1/worker/gpu/sample/penalties.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
@triton.jit
def _penalties_kernel(
logits_ptr,
logits_stride,
expanded_idx_mapping_ptr,
token_ids_ptr,
expanded_local_pos_ptr,
repetition_penalty_ptr,
frequency_penalty_ptr,
presence_penalty_ptr,
prompt_bin_mask_ptr,
prompt_bin_mask_stride,
output_bin_counts_ptr,
output_bin_counts_stride,
vocab_size,
BLOCK_SIZE: tl.constexpr,
):
token_idx = tl.program_id(0)
req_state_idx = tl.load(expanded_idx_mapping_ptr + token_idx)
rep_penalty = tl.load(repetition_penalty_ptr + req_state_idx)
freq_penalty = tl.load(frequency_penalty_ptr + req_state_idx)
pres_penalty = tl.load(presence_penalty_ptr + req_state_idx)
use_rep_penalty = rep_penalty != 1.0
use_freq_penalty = freq_penalty != 0.0
use_pres_penalty = pres_penalty != 0.0
# NPU doesn't support chained 'or' operations like 'A or B or C'
use_penalty = use_rep_penalty or use_freq_penalty
use_penalty = use_penalty or use_pres_penalty
if not use_penalty:
# Early return to avoid loading logits.
return
block_idx = tl.program_id(1)
block = block_idx * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
mask = block < vocab_size
logits = tl.load(logits_ptr + token_idx * logits_stride + block, mask=mask)
logits = logits.to(tl.float32)
base_output_counts = tl.load(
output_bin_counts_ptr + req_state_idx * output_bin_counts_stride + block,
mask=mask,
other=0,
)
# Accumulate draft token counts from previous positions directly into
# output_bin_counts (preserves its native tensor layout, avoiding an
# expensive shared-memory layout conversion after the loop).
pos = tl.load(expanded_local_pos_ptr + token_idx)
start_idx = token_idx - pos
output_bin_counts = base_output_counts
for prev_pos in tl.range(pos):
prev_token = tl.load(token_ids_ptr + start_idx + prev_pos + 1)
token_match = block == prev_token
output_bin_counts = output_bin_counts + token_match.to(tl.int32)
output_bin_mask = output_bin_counts != 0
# Apply repetition penalties.
if use_rep_penalty:
packed_block = block_idx * BLOCK_SIZE // 32 + tl.arange(0, BLOCK_SIZE // 32)
packed_mask = tl.load(
prompt_bin_mask_ptr + req_state_idx * prompt_bin_mask_stride + packed_block,
mask=packed_block < tl.cdiv(vocab_size, 32),
other=0,
)
bit_masks = 1 << tl.arange(0, 32)
bit_masks_expanded = bit_masks[None, :]
packed_expanded = packed_mask[:, None]
bits_matrix = (packed_expanded & bit_masks_expanded) != 0
prompt_bin_mask = bits_matrix.reshape(BLOCK_SIZE)
# If token appears in prompt or output, apply, otherwise use 1.0 for no-op.
scale = tl.where(prompt_bin_mask | output_bin_mask, rep_penalty, 1.0)
# If logits are positive, divide by penalty, otherwise multiply by penalty.
logits *= tl.where(logits > 0, 1.0 / scale, scale)
# Apply frequency penalties.
logits -= freq_penalty * output_bin_counts
# Apply presence penalties.
logits -= pres_penalty * output_bin_mask
# Store back to logits.
tl.store(logits_ptr + token_idx * logits_stride + block, logits, mask=mask)
def apply_penalties(
logits: torch.Tensor,
expanded_idx_mapping: torch.Tensor,
token_ids: torch.Tensor,
expanded_local_pos: torch.Tensor,
repetition_penalty: torch.Tensor,
frequency_penalty: torch.Tensor,
presence_penalty: torch.Tensor,
prompt_bin_mask: torch.Tensor,
output_bin_counts: torch.Tensor,
) -> None:
num_tokens, vocab_size = logits.shape
BLOCK_SIZE = 4096
num_blocks = triton.cdiv(vocab_size, BLOCK_SIZE)
_penalties_kernel[(num_tokens, num_blocks)](
logits,
logits.stride(0),
expanded_idx_mapping,
token_ids,
expanded_local_pos,
repetition_penalty,
frequency_penalty,
presence_penalty,
prompt_bin_mask,
prompt_bin_mask.stride(0),
output_bin_counts,
output_bin_counts.stride(0),
vocab_size,
BLOCK_SIZE=BLOCK_SIZE,
)
@triton.jit
def _bincount_kernel(
expanded_idx_mapping_ptr,
all_token_ids_ptr,
all_token_ids_stride,
prompt_len_ptr,
prefill_len_ptr,
prompt_bin_mask_ptr,
prompt_bin_mask_stride,
output_bin_counts_ptr,
output_bin_counts_stride,
BLOCK_SIZE: tl.constexpr,
):
token_idx = tl.program_id(0)
block_idx = tl.program_id(1)
req_state_idx = tl.load(expanded_idx_mapping_ptr + token_idx)
prefill_len = tl.load(prefill_len_ptr + req_state_idx)
if block_idx * BLOCK_SIZE >= prefill_len:
return
prompt_len = tl.load(prompt_len_ptr + req_state_idx)
block = block_idx * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
if block_idx * BLOCK_SIZE < prompt_len:
mask = block < prompt_len
prompt_tokens = tl.load(all_token_ids_ptr + req_state_idx * all_token_ids_stride + block, mask=mask)
idx = prompt_tokens // 32
bit_idx = prompt_tokens % 32
bit = tl.full((BLOCK_SIZE,), 1, tl.int32) << bit_idx
tl.atomic_or(
prompt_bin_mask_ptr + req_state_idx * prompt_bin_mask_stride + idx,
bit,
mask=mask,
)
if (block_idx + 1) * BLOCK_SIZE >= prompt_len:
mask = block < prefill_len
mask &= block >= prompt_len
output_tokens = tl.load(all_token_ids_ptr + req_state_idx * all_token_ids_stride + block, mask=mask)
tl.atomic_add(
output_bin_counts_ptr + req_state_idx * output_bin_counts_stride + output_tokens,
1,
mask=mask,
)
def bincount(
expanded_idx_mapping: torch.Tensor,
all_token_ids: torch.Tensor,
prompt_len: torch.Tensor,
prefill_len: torch.Tensor,
prompt_bin_mask: torch.Tensor,
output_bin_counts: torch.Tensor,
max_prefill_len: int,
) -> None:
prompt_bin_mask[expanded_idx_mapping] = 0
output_bin_counts[expanded_idx_mapping] = 0
num_tokens = expanded_idx_mapping.shape[0]
BLOCK_SIZE = 1024
num_blocks = triton.cdiv(max_prefill_len, BLOCK_SIZE)
_bincount_kernel[(num_tokens, num_blocks)](
expanded_idx_mapping,
all_token_ids,
all_token_ids.stride(0),
prompt_len,
prefill_len,
prompt_bin_mask,
prompt_bin_mask.stride(0),
output_bin_counts,
output_bin_counts.stride(0),
BLOCK_SIZE=BLOCK_SIZE,
)

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

View File

@@ -0,0 +1,71 @@
# Adapt from https://github.com/vllm-project/vllm/blob/main/vllm/v1/worker/gpu/states.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.v1.worker.gpu.states import RequestState
class AscendRequestState(RequestState):
"""Request state for Ascend NPUs."""
def __init__(
self,
max_num_reqs: int,
max_model_len: int,
max_num_batched_tokens: int,
num_speculative_steps: int,
vocab_size: int,
device: torch.device,
):
super().__init__(
max_num_reqs,
max_model_len,
max_num_batched_tokens,
num_speculative_steps,
vocab_size,
device,
)
# vllm gpu_model_runner_v2 deprecate the seqs_lens_cpu attribute,
# because they think most attention backends do not need it.
# However, Ascend attention backend muse uses seqs_lens_cpu,
# so we keep num_computed_tokens_cpu here, seq_lens_cpu need to be
# calculated by num_computed_tokens_cpu + decode_token_per_req outside.
self.num_computed_tokens_cpu: torch.Tensor = torch.zeros(
self.max_num_reqs,
dtype=torch.int32,
device="cpu",
)
def add_request(
self,
req_id,
prompt_len,
all_token_ids,
num_computed_tokens,
max_tokens=None,
):
super().add_request(
req_id,
prompt_len,
all_token_ids,
num_computed_tokens,
max_tokens=max_tokens,
)
req_idx = self.req_id_to_index[req_id]
self.num_computed_tokens_cpu[req_idx] = num_computed_tokens

View File

@@ -0,0 +1,68 @@
# 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.
#
# NPU-compatible structured output bitmask kernel.
#
# Upstream cannot be used directly on Ascend NPU: `BLOCK_SIZE=8192` overflows
# UB, while a smaller `BLOCK_SIZE` makes the grid unstable. We therefore keep
# `BLOCK_SIZE=8192` and split each block with `BLOCK_SIZE_SUB=1024`.
#
#
from vllm.triton_utils import tl, triton
# Adapted from
# https://github.com/mlc-ai/xgrammar/blob/main/python/xgrammar/kernels/apply_token_bitmask_inplace_triton.py
# Ascend NPU bitmask kernel (BLOCK_SIZE_SUB tiling)
# TODO: Optimize the kernel performance with NPU profiling data.
@triton.jit
def _apply_grammar_bitmask_kernel(
logits_ptr,
logits_stride,
logits_indices_ptr,
bitmask_ptr,
bitmask_stride,
vocab_size,
BLOCK_SIZE: tl.constexpr,
):
BLOCK_SIZE_SUB: tl.constexpr = 1024
bitmask_idx = tl.program_id(0)
block_id = tl.program_id(1)
logits_idx = tl.load(logits_indices_ptr + bitmask_idx)
# Sub-block tiling loop: process BLOCK_SIZE_SUB tokens per iteration
for sub_offset in tl.range(0, BLOCK_SIZE, BLOCK_SIZE_SUB):
global_token_offset = block_id * BLOCK_SIZE + sub_offset
bitmask_word_start = global_token_offset // 32
bitmask_offset = bitmask_word_start + tl.arange(0, BLOCK_SIZE_SUB // 32)
packed_bitmask = tl.load(
bitmask_ptr + bitmask_idx * bitmask_stride + bitmask_offset,
mask=bitmask_offset < bitmask_stride,
other=0,
)
bitmask = ((packed_bitmask[:, None] >> (tl.arange(0, 32)[None, :])) & 1) == 0
bitmask = bitmask.reshape(BLOCK_SIZE_SUB)
# Apply: set blocked positions to -inf
block_offset = global_token_offset + tl.arange(0, BLOCK_SIZE_SUB)
tl.store(
logits_ptr + logits_idx * logits_stride + block_offset,
-float("inf"),
mask=bitmask & (block_offset < vocab_size),
)

View File

@@ -0,0 +1,57 @@
from contextlib import contextmanager
import torch
from vllm.logger import logger
from vllm_ascend.compilation.acl_graph import get_draft_graph_params, get_graph_params, weak_ref_workspaces
@contextmanager
def torch_cuda_wrapper():
try:
torch.cuda.Event = torch.npu.Event
torch.cuda.Stream = torch.npu.Stream
torch.cuda.stream = torch.npu.stream
torch.cuda.default_stream = torch.npu.default_stream
torch.cuda.current_stream = torch.npu.current_stream
torch.cuda.graph_pool_handle = torch.npu.graph_pool_handle
torch.cuda.CUDAGraph = torch.npu.NPUGraph
torch.cuda.graph = torch_npu_graph_wrapper
torch.cuda.synchronize = torch.npu.synchronize
torch.cuda.set_stream = torch.npu.set_stream
torch.cuda.current_device = torch.npu.current_device
torch.cuda.mem_get_info = torch.npu.mem_get_info
logger.info_once("Wrapping torch.cuda with torch.npu.")
yield
finally:
pass
@contextmanager
def communicator_switch():
import vllm.distributed.device_communicators.cuda_communicator
from vllm_ascend.distributed.device_communicators.npu_communicator import NPUCommunicator
CudaCommunicator = vllm.distributed.device_communicators.cuda_communicator.CudaCommunicator
vllm.distributed.device_communicators.cuda_communicator.CudaCommunicator = NPUCommunicator
logger.debug("Switched CudaCommunicator -> NPUCommunicator for graph capture.")
try:
yield
finally:
vllm.distributed.device_communicators.cuda_communicator.CudaCommunicator = CudaCommunicator
logger.debug("Restored CudaCommunicator after graph capture.")
@contextmanager
def torch_npu_graph_wrapper(*args, **kwargs):
# MRV2-specific cleanup hook: intentionally reuse the graph context
# manager's exit to weak-ref graph workspaces after each capture,
# without adding another upstream monkey patch.
try:
with torch.npu.graph(*args, **kwargs):
yield
finally:
weak_ref_workspaces(get_graph_params())
weak_ref_workspaces(get_draft_graph_params())