36 lines
1.0 KiB
Python
36 lines
1.0 KiB
Python
import torch
|
|
from vllm.v1.spec_decode.ngram_proposer_gpu import NgramProposerGPU
|
|
|
|
|
|
class AscendNgramProposerNPU(NgramProposerGPU):
|
|
def __init__(self, vllm_config, device: torch.device, runner):
|
|
super().__init__(vllm_config, device=device)
|
|
|
|
def load_model(self, *args, **kwargs):
|
|
# No model to load.
|
|
pass
|
|
|
|
@torch.inference_mode()
|
|
def dummy_run(
|
|
self,
|
|
num_tokens,
|
|
with_prefill=None,
|
|
in_graph_capturing=None,
|
|
num_reqs=None,
|
|
num_tokens_across_dp=None,
|
|
aclgraph_runtime_mode=None,
|
|
batch_descriptor=None,
|
|
dummy_compute_logits=lambda hidden_states: None,
|
|
is_profile=False,
|
|
):
|
|
pass
|
|
|
|
def propose(
|
|
self,
|
|
num_tokens_no_spec: torch.Tensor, # [batch_size]
|
|
token_ids_gpu: torch.Tensor, # [batch_size, max_len]
|
|
valid_sampled_token_ids_gpu: torch.Tensor, # [batch_size, num_spec_tokens + 1]
|
|
valid_sampled_tokens_count: torch.Tensor, # [batch_size]
|
|
):
|
|
pass
|