Skip to content

vllm.v1.worker.gpu.sample.prompt_logprob

Classes:

PromptLogprobsWorker

Methods:

Source code in vllm/v1/worker/gpu/sample/prompt_logprob.py
class PromptLogprobsWorker:
    def __init__(
        self,
        max_num_reqs: int,
        device: torch.device,
        logprobs_mode: LogprobsMode = "raw_logprobs",
    ):
        self.max_num_reqs = max_num_reqs
        self.device = device
        self.logprobs_mode = logprobs_mode

        self.uses_prompt_logprobs = np.zeros(self.max_num_reqs, dtype=bool)
        self.num_prompt_logprobs = np.zeros(self.max_num_reqs, dtype=np.int32)
        # req_idx -> list of in-progress LogprobsTensors
        self.in_progress_prompt_logprobs: dict[str, list[LogprobsTensors]] = {}
        # req_id -> fixed-ID scoring state
        self.token_id_scores: dict[str, _TokenIdScores] = {}

    def add_request(self, req_id: str, req_idx: int, sampling_params: SamplingParams):
        uses_prompt_logprobs = sampling_params.prompt_logprobs is not None
        self.uses_prompt_logprobs[req_idx] = uses_prompt_logprobs
        self.num_prompt_logprobs[req_idx] = sampling_params.prompt_logprobs or 0
        if uses_prompt_logprobs:
            self.in_progress_prompt_logprobs[req_id] = []
        if sampling_params.prompt_logprob_token_ids is not None:
            self.token_id_scores[req_id] = _TokenIdScores(
                async_tensor_h2d(
                    sampling_params.prompt_logprob_token_ids, self.device, torch.int64
                ),
                sampling_params.prompt_logprob_start or 0,
            )

    def remove_request(self, req_id: str) -> None:
        self.in_progress_prompt_logprobs.pop(req_id, None)
        self.token_id_scores.pop(req_id, None)

    def compute_prompt_token_id_logprobs(
        self,
        logits_fn: Callable[[torch.Tensor], torch.Tensor],
        hidden_states: torch.Tensor,
        input_batch: InputBatch,
        prompt_lens: np.ndarray,
    ) -> dict[str, torch.Tensor]:
        """Fill scores row (prompt row - req.start) per chunk; emit on the last."""
        if not self.token_id_scores:
            return {}
        logits_mode = self.logprobs_mode in ("raw_logits", "processed_logits")
        out: dict[str, torch.Tensor] = {}
        for i, req_id in enumerate(input_batch.req_ids):
            req = self.token_id_scores.get(req_id)
            if req is None:
                continue
            prompt_len = int(prompt_lens[input_batch.idx_mapping_np[i]])
            chunk_start = int(input_batch.num_computed_prefill_tokens_np[i])
            chunk_end = chunk_start + int(input_batch.num_scheduled_tokens[i])
            # Skip decode steps and, as compute_prompt_logprobs does, requests
            # resumed after preemption: their scores were already emitted.
            if chunk_start >= prompt_len or prompt_len < input_batch.prefill_len_np[i]:
                continue
            if req.scores is None:
                if chunk_start > req.start:
                    # This prefill starts past the first scored row, so those
                    # rows will never be written; drop the request instead of
                    # emitting a buffer with unwritten rows.
                    self.token_id_scores.pop(req_id, None)
                    continue
                req.scores = hidden_states.new_empty(
                    (max(prompt_len - 1 - req.start, 0), len(req.token_ids)),
                    dtype=torch.float32,
                )
            # The last prompt row predicts the first decode token; skip it.
            lo = max(req.start, chunk_start)
            hi = min(chunk_end, prompt_len - 1)
            base = int(input_batch.query_start_loc_np[i]) - chunk_start
            for a in range(lo, hi, CHUNK_SIZE):
                b = min(a + CHUNK_SIZE, hi)
                logits = logits_fn(hidden_states[base + a : base + b])
                ids = req.token_ids.expand(b - a, -1)
                req.scores[a - req.start : b - req.start] = (
                    logits.gather(-1, ids)
                    if logits_mode
                    else compute_token_logprobs(logits, ids)
                )
            if chunk_end >= prompt_len:
                out[req_id] = req.scores
                req.scores = None
        return out

    def compute_prompt_logprobs(
        self,
        logits_fn: Callable[[torch.Tensor], torch.Tensor],
        hidden_states: torch.Tensor,
        input_batch: InputBatch,
        # [max_num_reqs, max_model_len]
        all_token_ids: torch.Tensor,
        # [max_num_reqs]
        num_computed_tokens: torch.Tensor,
        # [max_num_reqs]
        prompt_lens: np.ndarray,
    ) -> dict[str, LogprobsTensors]:
        idx_mapping_np = input_batch.idx_mapping_np
        needs_prompt_logprobs = self.uses_prompt_logprobs[idx_mapping_np]
        if not np.any(needs_prompt_logprobs):
            # Common case: No request asks for prompt logprobs.
            return {}

        num_prompt_logprobs = self.num_prompt_logprobs[idx_mapping_np]
        prompt_lens = prompt_lens[idx_mapping_np]
        computed_prefill = input_batch.num_computed_prefill_tokens_np
        includes_prompt = computed_prefill < prompt_lens
        # NOTE(woosuk): If the request was resumed after preemption, its prompt
        # logprobs must have been computed before preemption. Skip.
        resumed_after_prompt = prompt_lens < input_batch.prefill_len_np
        needs_prompt_logprobs &= includes_prompt & ~resumed_after_prompt
        if not np.any(needs_prompt_logprobs):
            return {}

        # get the maximum number in this batch
        requested_num_prompt_logprobs = num_prompt_logprobs[needs_prompt_logprobs]
        max_num_prompt_logprobs = (
            -1
            if np.any(requested_num_prompt_logprobs == -1)
            else int(requested_num_prompt_logprobs.max())
        )

        # Get the prompt logprobs token_ids.
        prompt_logprobs_token_ids = get_prompt_logprobs_token_ids(
            input_batch.num_tokens,
            input_batch.query_start_loc,
            input_batch.idx_mapping,
            num_computed_tokens,
            all_token_ids,
        )
        prompt_token_ids, prompt_logprobs, prompt_ranks = (
            compute_prompt_logprobs_with_chunking(
                prompt_logprobs_token_ids,
                hidden_states[: input_batch.num_tokens],
                logits_fn,
                max_num_prompt_logprobs,
                self.logprobs_mode,
            )
        )

        pos_after_step = computed_prefill + input_batch.num_scheduled_tokens
        is_prompt_chunked = pos_after_step < prompt_lens

        query_start_loc_np = input_batch.query_start_loc_np
        prompt_logprobs_dict: dict[str, LogprobsTensors] = {}
        for i, req_id in enumerate(input_batch.req_ids):
            if not needs_prompt_logprobs[i]:
                continue

            req_is_prompt_chunked = is_prompt_chunked[i]
            req_num_prompt_logprobs = int(num_prompt_logprobs[i])
            start_idx = query_start_loc_np[i]
            end_idx = query_start_loc_np[i + 1]
            assert start_idx < end_idx, (
                f"start_idx ({start_idx}) >= end_idx ({end_idx})"
            )
            if not req_is_prompt_chunked:
                end_idx -= 1

            width = (
                prompt_logprobs.shape[1]
                if req_num_prompt_logprobs == -1
                else req_num_prompt_logprobs + 1
            )
            # no logprobs if start_idx >= end_idx
            logprobs = (
                None
                if start_idx >= end_idx
                else LogprobsTensors(
                    logprob_token_ids=prompt_token_ids[start_idx:end_idx, :width],
                    logprobs=prompt_logprobs[start_idx:end_idx, :width],
                    selected_token_ranks=prompt_ranks[start_idx:end_idx],
                )
            )

            prompt_logprobs_list = self.in_progress_prompt_logprobs[req_id]
            if logprobs is not None and (req_is_prompt_chunked or prompt_logprobs_list):
                prompt_logprobs_list.append(logprobs)
            if req_is_prompt_chunked:
                # Prompt is chunked. Do not return the logprobs yet.
                continue

            if prompt_logprobs_list:
                # Merge the in-progress logprobs.
                logprobs = LogprobsTensors.cat(prompt_logprobs_list)
                prompt_logprobs_list.clear()

            if logprobs is None:
                continue

            prompt_logprobs_dict[req_id] = logprobs
        return prompt_logprobs_dict

compute_prompt_token_id_logprobs(logits_fn, hidden_states, input_batch, prompt_lens)

Fill scores row (prompt row - req.start) per chunk; emit on the last.

Source code in vllm/v1/worker/gpu/sample/prompt_logprob.py
def compute_prompt_token_id_logprobs(
    self,
    logits_fn: Callable[[torch.Tensor], torch.Tensor],
    hidden_states: torch.Tensor,
    input_batch: InputBatch,
    prompt_lens: np.ndarray,
) -> dict[str, torch.Tensor]:
    """Fill scores row (prompt row - req.start) per chunk; emit on the last."""
    if not self.token_id_scores:
        return {}
    logits_mode = self.logprobs_mode in ("raw_logits", "processed_logits")
    out: dict[str, torch.Tensor] = {}
    for i, req_id in enumerate(input_batch.req_ids):
        req = self.token_id_scores.get(req_id)
        if req is None:
            continue
        prompt_len = int(prompt_lens[input_batch.idx_mapping_np[i]])
        chunk_start = int(input_batch.num_computed_prefill_tokens_np[i])
        chunk_end = chunk_start + int(input_batch.num_scheduled_tokens[i])
        # Skip decode steps and, as compute_prompt_logprobs does, requests
        # resumed after preemption: their scores were already emitted.
        if chunk_start >= prompt_len or prompt_len < input_batch.prefill_len_np[i]:
            continue
        if req.scores is None:
            if chunk_start > req.start:
                # This prefill starts past the first scored row, so those
                # rows will never be written; drop the request instead of
                # emitting a buffer with unwritten rows.
                self.token_id_scores.pop(req_id, None)
                continue
            req.scores = hidden_states.new_empty(
                (max(prompt_len - 1 - req.start, 0), len(req.token_ids)),
                dtype=torch.float32,
            )
        # The last prompt row predicts the first decode token; skip it.
        lo = max(req.start, chunk_start)
        hi = min(chunk_end, prompt_len - 1)
        base = int(input_batch.query_start_loc_np[i]) - chunk_start
        for a in range(lo, hi, CHUNK_SIZE):
            b = min(a + CHUNK_SIZE, hi)
            logits = logits_fn(hidden_states[base + a : base + b])
            ids = req.token_ids.expand(b - a, -1)
            req.scores[a - req.start : b - req.start] = (
                logits.gather(-1, ids)
                if logits_mode
                else compute_token_logprobs(logits, ids)
            )
        if chunk_end >= prompt_len:
            out[req_id] = req.scores
            req.scores = None
    return out