Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions tests/conftest.py
Original file line number Diff line number Diff line change
Expand Up @@ -996,6 +996,8 @@ def generate(
req_sample_output_ids: list[list[int]] = []
req_sample_output_strs: list[str] = []
req_logprobs = []
if req_output.prompt_logprobs:
req_logprobs.extend(req_output.prompt_logprobs)
for sample in req_output.outputs:
output_str = sample.text
output_ids = list(sample.token_ids)
Expand Down
11 changes: 10 additions & 1 deletion tests/v1/e2e/general/test_async_scheduling.py
Original file line number Diff line number Diff line change
Expand Up @@ -57,6 +57,8 @@ def test_without_spec_decoding(
dict(bad_words=["the", " the"]),
dict(logprobs=2),
dict(logprobs=2, frequency_penalty=-1.0),
dict(prompt_logprobs=2),
dict(prompt_logprobs=2, logprobs=2),
dict(structured_outputs=struct_outputs),
dict(
structured_outputs=struct_outputs,
Expand Down Expand Up @@ -126,6 +128,8 @@ def test_with_eagle3_spec_decoding(sample_json_schema, monkeypatch: pytest.Monke
dict(bad_words=["the", " the"]),
dict(logprobs=2),
dict(logprobs=2, frequency_penalty=-1.0),
dict(prompt_logprobs=2),
dict(prompt_logprobs=2, logprobs=2),
dict(structured_outputs=struct_outputs),
dict(
structured_outputs=struct_outputs,
Expand Down Expand Up @@ -413,7 +417,12 @@ def _all_logprobs_match(req_a, req_b) -> bool:
)


def _logprobs_match(lps_a: dict[int, Logprob], lps_b: dict[int, Logprob]) -> bool:
def _logprobs_match(
lps_a: dict[int, Logprob] | None,
lps_b: dict[int, Logprob] | None,
) -> bool:
if lps_a is None or lps_b is None:
return lps_a is lps_b
rel_tol, abs_tol = 1e-3, 1e-6
return (
len(lps_a) == len(lps_b)
Expand Down
4 changes: 1 addition & 3 deletions vllm/v1/worker/gpu/sample/prompt_logprob.py
Original file line number Diff line number Diff line change
Expand Up @@ -55,10 +55,8 @@ def compute_prompt_logprobs(

num_prompt_logprobs = self.num_prompt_logprobs[idx_mapping_np]
prompt_lens = prompt_lens[idx_mapping_np]
# NOTE(woosuk): -1 because the last prompt token's hidden state is not
# needed for prompt logprobs.
computed_prefill = num_computed_prefill_tokens[idx_mapping_np]
includes_prompt = computed_prefill < prompt_lens - 1
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 < prefill_lens[idx_mapping_np]
Expand Down
6 changes: 2 additions & 4 deletions vllm/v1/worker/gpu_input_batch.py
Original file line number Diff line number Diff line change
Expand Up @@ -49,6 +49,8 @@ class CachedRequestState:

lora_request: LoRARequest | None = None
prompt_embeds: torch.Tensor | None = None
# To accumulate prompt logprobs tensor chunks across prefill steps.
in_progress_prompt_logprobs_cpu: LogprobsTensors | None = None

# Per-position mask for mixed-mode inputs (e.g chat completion with
# prompt_embeds content parts). See `Request.prompt_is_token_ids`.
Expand Down Expand Up @@ -255,9 +257,6 @@ def __init__(
# More efficient than num_logprobs=-1 when only a few tokens are needed
self.logprob_token_ids: dict[str, list[int]] = {}

# To accumulate prompt logprobs tensor chunks across prefill steps.
self.in_progress_prompt_logprobs_cpu: dict[str, LogprobsTensors] = {}

# Internal representation of per-step batch state changes, used for
# reordering persistent batch and generating logitsprocs batch state
# updates. Should reset each step.
Expand Down Expand Up @@ -552,7 +551,6 @@ def remove_request(self, req_id: str) -> int | None:
self.generators.pop(req_index, None)
self.num_logprobs.pop(req_id, None)
self.logprob_token_ids.pop(req_id, None)
self.in_progress_prompt_logprobs_cpu.pop(req_id, None)
if self.prev_req_id_to_index is not None:
self.prev_req_id_to_index.pop(req_id, None)

Expand Down
9 changes: 4 additions & 5 deletions vllm/v1/worker/gpu_model_runner.py
Original file line number Diff line number Diff line change
Expand Up @@ -5094,7 +5094,6 @@ def _get_prompt_logprobs_dict(
if not num_prompt_logprobs_dict:
return {}

in_progress_dict = self.input_batch.in_progress_prompt_logprobs_cpu
prompt_logprobs_dict: dict[str, LogprobsTensors | None] = {}

# Since prompt logprobs are a rare feature, prioritize simple,
Expand All @@ -5118,14 +5117,14 @@ def _get_prompt_logprobs_dict(
)

# Set up target LogprobsTensors object.
logprobs_tensors = in_progress_dict.get(req_id)
if not logprobs_tensors:
logprobs_tensors = request.in_progress_prompt_logprobs_cpu
if logprobs_tensors is None:
# Create empty logprobs CPU tensors for the entire prompt.
# If chunked, we'll copy in slice by slice.
logprobs_tensors = LogprobsTensors.empty_cpu(
num_prompt_tokens - 1, num_prompt_logprobs + 1
)
in_progress_dict[req_id] = logprobs_tensors
request.in_progress_prompt_logprobs_cpu = logprobs_tensors

# Determine number of logits to retrieve.
start_idx = request.num_computed_tokens
Expand Down Expand Up @@ -5182,7 +5181,7 @@ def _get_prompt_logprobs_dict(
# num_prompt_logprobs_dict.
for req_id in completed_prefill_reqs:
del num_prompt_logprobs_dict[req_id]
del in_progress_dict[req_id]
self.requests[req_id].in_progress_prompt_logprobs_cpu = None

# Must synchronize the non-blocking GPU->CPU transfers.
if prompt_logprobs_dict:
Expand Down
Loading