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