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