35
vllm_ascend/spec_decode/ngram_proposer_npu.py
Normal file
35
vllm_ascend/spec_decode/ngram_proposer_npu.py
Normal file
@@ -0,0 +1,35 @@
|
||||
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
|
||||
Reference in New Issue
Block a user