class NgramGPUSpeculator(BaseSpeculator):
"""V2-compatible GPU n-gram speculator."""
supports_mm_inputs = False
draft_logits = None
def __init__(
self,
vllm_config: VllmConfig,
device: torch.device,
req_states: RequestState,
):
if not HAS_TRITON:
raise RuntimeError("ngram_gpu speculative decoding requires Triton.")
spec = vllm_config.speculative_config
assert spec is not None
assert spec.prompt_lookup_min is not None, (
"prompt_lookup_min must be configured for ngram_gpu"
)
assert spec.prompt_lookup_max is not None, (
"prompt_lookup_max must be configured for ngram_gpu"
)
assert 1 <= spec.prompt_lookup_min <= spec.prompt_lookup_max
self.vllm_config = vllm_config
self.device = device
self.req_states = req_states
self.speculative_config = spec
self.num_speculative_steps: int = spec.num_speculative_tokens
self.min_n: int = spec.prompt_lookup_min
self.max_n: int = spec.prompt_lookup_max
self.max_num_reqs: int = vllm_config.scheduler_config.max_num_seqs
self.max_model_len: int = vllm_config.model_config.max_model_len
L = self.max_model_len
if L >= 1024:
self.block_l = 256
elif L >= 256:
self.block_l = 128
elif L >= 64:
self.block_l = 64
else:
self.block_l = max(16, triton.next_power_of_2(max(L, 1)))
self.n_blocks = triton.cdiv(L, self.block_l)
self.scratch = torch.zeros(
(self.max_num_reqs, self.n_blocks), dtype=torch.int64, device=device
)
# Batch-ordered draft output, scattered into RequestState.draft_tokens
# by the model runner (same contract as the model-based speculators).
self.drafts = torch.zeros(
(self.max_num_reqs, self.num_speculative_steps),
dtype=torch.int64,
device=device,
)
def init_cudagraph_manager(self, cudagraph_mode: CUDAGraphMode) -> None:
del cudagraph_mode
def capture(self) -> None:
return None
@torch.inference_mode()
def propose(
self,
input_batch: InputBatch,
attn_metadata: Any,
slot_mappings: Any,
last_hidden_states: torch.Tensor,
aux_hidden_states: list[torch.Tensor] | None,
num_sampled: torch.Tensor,
num_rejected: torch.Tensor,
last_sampled: torch.Tensor,
next_prefill_tokens: torch.Tensor,
temperature: torch.Tensor,
seeds: torch.Tensor,
dp_sync: DPSyncState | 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: bool = False,
num_speculative_tokens: int | None = None,
) -> torch.Tensor:
num_reqs = input_batch.num_reqs
if dummy_run:
return self.drafts[:num_reqs]
req_states = self.req_states
token_ids = req_states.all_token_ids.gpu
idx_mapping = input_batch.idx_mapping
_ngram_scan_kernel[(num_reqs, self.n_blocks)](
token_ids,
token_ids.stride(0),
idx_mapping,
req_states.total_len.gpu,
num_sampled,
self.scratch,
self.scratch.stride(0),
self.max_model_len,
self.min_n,
self.max_n,
max(1, triton.next_power_of_2(self.max_n)),
self.block_l,
num_warps=4,
num_stages=2,
)
_ngram_finalize_kernel[(num_reqs,)](
token_ids,
token_ids.stride(0),
idx_mapping,
req_states.total_len.gpu,
num_sampled,
last_sampled.view(-1),
self.scratch,
self.scratch.stride(0),
self.drafts,
self.max_model_len,
self.n_blocks,
self.num_speculative_steps,
max(1, triton.next_power_of_2(self.num_speculative_steps)),
max(1, triton.next_power_of_2(self.n_blocks)),
num_warps=2,
num_stages=1,
)
return self.drafts[:num_reqs]