Skip to content

vllm.v1.worker.gpu.spec_decode.rejection_sampler

Functions:

_iter_request_chunks(cu_num_logits, max_chunk_logits)

Yield maximally packed request ranges without splitting requests.

Source code in vllm/v1/worker/gpu/spec_decode/rejection_sampler.py
def _iter_request_chunks(
    cu_num_logits: np.ndarray, max_chunk_logits: int
) -> Iterator[tuple[int, int]]:
    """Yield maximally packed request ranges without splitting requests."""
    assert max_chunk_logits > 0
    num_reqs = cu_num_logits.size - 1
    start = 0
    while start < num_reqs:
        max_logit = int(cu_num_logits[start]) + max_chunk_logits
        end = int(np.searchsorted(cu_num_logits, max_logit, side="right") - 1)
        end = min(num_reqs, max(start + 1, end))
        yield start, end
        start = end

gather_draft_sampled(input_ids, positions, logits_indices, expanded_idx_mapping, expanded_local_pos, prefill_len)

Gather the input token and position of each logits row.

Draft rows of requests that have not yet sampled past their prefill are set to -1 so that the rejection kernels reject them.

Source code in vllm/v1/worker/gpu/spec_decode/rejection_sampler.py
def gather_draft_sampled(
    input_ids: torch.Tensor,
    positions: torch.Tensor,
    logits_indices: torch.Tensor,
    expanded_idx_mapping: torch.Tensor,
    expanded_local_pos: torch.Tensor,
    prefill_len: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor]:
    """Gather the input token and position of each logits row.

    Draft rows of requests that have not yet sampled past their prefill are
    set to -1 so that the rejection kernels reject them.
    """
    num_logits = logits_indices.shape[0]
    draft_sampled = input_ids.new_empty(num_logits)
    pos = positions.new_empty(num_logits)
    BLOCK_SIZE = 1024
    _gather_draft_sampled_kernel[(triton.cdiv(num_logits, BLOCK_SIZE),)](
        draft_sampled,
        pos,
        input_ids,
        positions,
        logits_indices,
        expanded_idx_mapping,
        expanded_local_pos,
        prefill_len,
        num_logits,
        BLOCK_SIZE=BLOCK_SIZE,
    )
    return draft_sampled, pos

get_max_chunk_logits(vocab_size)

Largest number of logits rows one verification chunk may hold.

Source code in vllm/v1/worker/gpu/spec_decode/rejection_sampler.py
def get_max_chunk_logits(vocab_size: int) -> int:
    """Largest number of logits rows one verification chunk may hold."""
    return max(1, MAX_CHUNK_BYTES // (vocab_size * _FP32_BYTES))