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)
|
||||
Reference in New Issue
Block a user