39
vllm_ascend/worker/v2/README.md
Normal file
39
vllm_ascend/worker/v2/README.md
Normal 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`.
|
||||
0
vllm_ascend/worker/v2/__init__.py
Normal file
0
vllm_ascend/worker/v2/__init__.py
Normal file
163
vllm_ascend/worker/v2/aclgraph_utils.py
Normal file
163
vllm_ascend/worker/v2/aclgraph_utils.py
Normal 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)
|
||||
504
vllm_ascend/worker/v2/attn_utils.py
Normal file
504
vllm_ascend/worker/v2/attn_utils.py
Normal 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
|
||||
164
vllm_ascend/worker/v2/block_table.py
Normal file
164
vllm_ascend/worker/v2/block_table.py
Normal 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)
|
||||
210
vllm_ascend/worker/v2/input_batch.py
Normal file
210
vllm_ascend/worker/v2/input_batch.py
Normal 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,
|
||||
)
|
||||
510
vllm_ascend/worker/v2/model_runner.py
Normal file
510
vllm_ascend/worker/v2/model_runner.py
Normal 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
|
||||
34
vllm_ascend/worker/v2/model_states/__init__.py
Normal file
34
vllm_ascend/worker/v2/model_states/__init__.py
Normal 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)
|
||||
77
vllm_ascend/worker/v2/model_states/default.py
Normal file
77
vllm_ascend/worker/v2/model_states/default.py
Normal 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
|
||||
0
vllm_ascend/worker/v2/sample/__init__.py
Normal file
0
vllm_ascend/worker/v2/sample/__init__.py
Normal file
183
vllm_ascend/worker/v2/sample/bad_words.py
Normal file
183
vllm_ascend/worker/v2/sample/bad_words.py
Normal 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,
|
||||
)
|
||||
205
vllm_ascend/worker/v2/sample/gumbel.py
Normal file
205
vllm_ascend/worker/v2/sample/gumbel.py
Normal 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
|
||||
164
vllm_ascend/worker/v2/sample/logprob.py
Normal file
164
vllm_ascend/worker/v2/sample/logprob.py
Normal 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,
|
||||
)
|
||||
91
vllm_ascend/worker/v2/sample/min_p.py
Normal file
91
vllm_ascend/worker/v2/sample/min_p.py
Normal 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,
|
||||
)
|
||||
215
vllm_ascend/worker/v2/sample/penalties.py
Normal file
215
vllm_ascend/worker/v2/sample/penalties.py
Normal 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,
|
||||
)
|
||||
0
vllm_ascend/worker/v2/spec_decode/__init__.py
Normal file
0
vllm_ascend/worker/v2/spec_decode/__init__.py
Normal file
36
vllm_ascend/worker/v2/spec_decode/eagle/__init__.py
Normal file
36
vllm_ascend/worker/v2/spec_decode/eagle/__init__.py
Normal 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.")
|
||||
231
vllm_ascend/worker/v2/spec_decode/eagle/aclgraph.py
Normal file
231
vllm_ascend/worker/v2/spec_decode/eagle/aclgraph.py
Normal 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()
|
||||
370
vllm_ascend/worker/v2/spec_decode/eagle/speculator.py
Normal file
370
vllm_ascend/worker/v2/spec_decode/eagle/speculator.py
Normal 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)
|
||||
486
vllm_ascend/worker/v2/spec_decode/rejection_sampler_utils.py
Normal file
486
vllm_ascend/worker/v2/spec_decode/rejection_sampler_utils.py
Normal 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
|
||||
71
vllm_ascend/worker/v2/states.py
Normal file
71
vllm_ascend/worker/v2/states.py
Normal 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
|
||||
68
vllm_ascend/worker/v2/structured_outputs.py
Normal file
68
vllm_ascend/worker/v2/structured_outputs.py
Normal 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),
|
||||
)
|
||||
57
vllm_ascend/worker/v2/utils.py
Normal file
57
vllm_ascend/worker/v2/utils.py
Normal 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())
|
||||
Reference in New Issue
Block a user