Skip to content

vllm.v1.worker.gpu.spec_decode.ngram.speculator

Classes:

NgramGPUSpeculator

Bases: BaseSpeculator

V2-compatible GPU n-gram speculator.

Source code in vllm/v1/worker/gpu/spec_decode/ngram/speculator.py
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]