[Bugfix][V1/V2] Fix prompt_logprobs to respect logprobs_mode (#47680)

Signed-off-by: Wojciech Wais <wojciech.wais@gmail.com>
Signed-off-by: Federico Kamelhar <209537060+fede-kamel@users.noreply.github.com>
Signed-off-by: Allen Shen <aoshen@inferact.ai>
Co-authored-by: Wojciech Wais <wojciech.wais@gmail.com>
Co-authored-by: Federico Kamelhar <209537060+fede-kamel@users.noreply.github.com>
This commit is contained in:
aoshen02
2026-07-17 21:58:59 +01:00
committed by GitHub
co-authored by Wojciech Wais Federico Kamelhar
parent 088c0be268
commit 41ea2dd44a
11 changed files with 106 additions and 40 deletions
+38
View File
@@ -564,6 +564,44 @@ def test_logprobs_mode(logprobs_mode: LogprobsMode):
cleanup_dist_env_and_memory()
def test_prompt_logprobs_mode():
"""prompt_logprobs must respect logprobs_mode: *_logits and *_logprobs
must return different values. Prompt tokens skip sampling processors,
so processed_* == raw_* on the prompt side."""
from vllm import LLM
values: dict[str, float] = {}
for mode in get_args(LogprobsMode):
llm = LLM(
"facebook/opt-125m",
enable_prefix_caching=False,
gpu_memory_utilization=0.05,
max_model_len=16,
logprobs_mode=mode,
)
try:
results = llm.generate(
["Hello world"],
sampling_params=SamplingParams(
max_tokens=1, prompt_logprobs=0, temperature=0
),
)
assert results[0].prompt_logprobs is not None
assert results[0].prompt_logprobs[1] is not None
tok_id = results[0].prompt_token_ids[1]
values[mode] = results[0].prompt_logprobs[1][tok_id].logprob
finally:
del llm
torch.accelerator.empty_cache()
cleanup_dist_env_and_memory()
assert values["raw_logprobs"] <= 0
assert values["processed_logprobs"] <= 0
assert values["raw_logits"] != values["raw_logprobs"]
assert values["processed_logits"] == values["raw_logits"]
assert values["processed_logprobs"] == values["raw_logprobs"]
class TestCorrectDecodedToken:
"""Unit tests for _correct_decoded_token method in LogprobsProcessor.
+2
View File
@@ -233,6 +233,8 @@ class ModelConfig:
Raw means the values before applying any logit processors, like bad words.
Processed means the values after applying all processors, including
temperature and top_k/top_p.
Note: for prompt_logprobs, processed_* and raw_* yield identical results
because prompt tokens do not go through sampling processors.
"""
use_fp64_gumbel: bool = False
"""Whether to use FP64 (instead of FP32) random noise for Gumbel-max and
-7
View File
@@ -2202,13 +2202,6 @@ class VllmConfig:
if model_config is not None and model_config.enable_prompt_embeds:
unsupported.append("prompt embeds")
if (
model_config is not None
and model_config.runner_type == "generate"
and model_config.logprobs_mode in ("raw_logits", "processed_logits")
):
unsupported.append(f"logprobs mode '{model_config.logprobs_mode}'")
if self.cache_config.kv_sharing_fast_prefill:
# Will be added by https://github.com/vllm-project/vllm/pull/35045
unsupported.append("KV sharing fast prefill")
@@ -53,7 +53,7 @@ from vllm.v1.worker.gpu.attn_utils import build_attn_metadata
from vllm.v1.worker.gpu.buffer_utils import UvaBackedTensor, async_copy_to_gpu
from vllm.v1.worker.gpu.input_batch import InputBatch
from vllm.v1.worker.gpu.model_states.interface import ModelState
from vllm.v1.worker.gpu.sample.logprob import compute_topk_logprobs
from vllm.v1.worker.gpu.sample.logprob import compute_topk_scores
from vllm.v1.worker.gpu.sample.output import SamplerOutput
from vllm.v1.worker.gpu.sample.penalties import use_penalty
from vllm.v1.worker.gpu.states import RequestState
@@ -1064,6 +1064,7 @@ class DiffusionSampler:
):
self.sampling_states = sampler.sampling_states
self.req_states = sampler.req_states
self.logits_mode = sampler.logprobs_mode in ("raw_logits", "processed_logits")
# Self-conditioning soft embed = probs @ embed_weight * normalizer,
# computed in the sampler (see _compiled_sample_step). ``embed_weight``
# is the vocab-parallel shard; [sc_vocab_start, sc_vocab_end) is this
@@ -1359,10 +1360,11 @@ class DiffusionSampler:
# positions are never emitted.
k_i = int(valid_canvas_len_np[start_req + li])
pos = li * CL
self._pending_logprobs[slot.item()] = compute_topk_logprobs(
self._pending_logprobs[slot.item()] = compute_topk_scores(
flat_logits[pos : pos + k_i],
max_num_logprobs,
argmax_tokens[local_idx][:k_i],
logits_mode=self.logits_mode,
)
# Commit steps: is_committing was True at entry. Reassemble previously
+8 -2
View File
@@ -67,8 +67,14 @@ class RejectionSampler(nn.Module):
self.sampler = sampler
self.use_fp64_gumbel = getattr(sampler, "use_fp64_gumbel", False)
logprobs_mode = self.sampler.logprobs_mode
self.is_processed_logprobs_mode = logprobs_mode.startswith("processed")
self.is_logits_logprobs_mode = logprobs_mode.endswith("logits")
self.is_processed_logprobs_mode = logprobs_mode in (
"processed_logprobs",
"processed_logits",
)
self.is_logits_logprobs_mode = logprobs_mode in (
"raw_logits",
"processed_logits",
)
self.synthetic_conditional_rates: torch.Tensor | None = None
if (
+4 -1
View File
@@ -341,7 +341,10 @@ class GPUModelRunner(LoRAModelRunnerMixin):
self.speculative_config,
self.device,
)
self.prompt_logprobs_worker = PromptLogprobsWorker(self.max_num_reqs)
self.prompt_logprobs_worker = PromptLogprobsWorker(
self.max_num_reqs,
logprobs_mode=self.model_config.logprobs_mode,
)
self.structured_outputs_worker = StructuredOutputsWorker(
max_num_logits=self.max_num_reqs * self.decode_query_len,
vocab_size=self.vocab_size,
+12 -5
View File
@@ -106,7 +106,7 @@ def compute_token_logprobs(
return logprobs
def compute_topk_logprobs(
def compute_topk_scores(
logits: torch.Tensor,
num_logprobs: int,
sampled_token_ids: torch.Tensor,
@@ -114,6 +114,7 @@ def compute_topk_logprobs(
logprob_token_ids_state: "LogprobTokenIdsState | None" = None,
expanded_idx_mapping: torch.Tensor | None = None,
max_per_req_token_ids: int = 0,
logits_mode: bool = False,
) -> LogprobsTensors:
assert num_logprobs >= 0
batch_size, vocab_size = logits.shape
@@ -124,7 +125,10 @@ def compute_topk_logprobs(
if num_logprobs > 0:
topk_indices = torch.topk(logits, num_logprobs, dim=-1).indices
logprob_token_ids = torch.cat((logprob_token_ids, topk_indices), dim=1)
logprobs = compute_token_logprobs(logits, logprob_token_ids)
if logits_mode:
scores = logits.gather(-1, logprob_token_ids).to(torch.float32)
else:
scores = compute_token_logprobs(logits, logprob_token_ids)
else:
# Some requests specified logprob_token_ids. Build the [batch_size,
# 1 + max_cols] token_ids matrix and validity mask on the GPU via a
@@ -158,8 +162,11 @@ def compute_topk_logprobs(
NUM_TOPK=num_logprobs,
PADDED_COLS=triton.next_power_of_2(num_cols),
)
logprobs = compute_token_logprobs(logits, logprob_token_ids)
logprobs = logprobs.masked_fill(~valid_mask, float("-inf"))
if logits_mode:
scores = logits.gather(-1, logprob_token_ids).to(torch.float32)
else:
scores = compute_token_logprobs(logits, logprob_token_ids)
scores = scores.masked_fill(~valid_mask, float("-inf"))
token_ranks = torch.empty(batch_size, dtype=torch.int64, device=logits.device)
_ranks_kernel[(batch_size,)](
@@ -172,7 +179,7 @@ def compute_topk_logprobs(
)
return LogprobsTensors(
logprob_token_ids=logprob_token_ids,
logprobs=logprobs,
logprobs=scores,
selected_token_ranks=token_ranks,
cu_num_generated_tokens=cu_num_logits,
)
+17 -11
View File
@@ -5,16 +5,18 @@ from collections.abc import Callable
import numpy as np
import torch
from vllm.config.model import LogprobsMode
from vllm.sampling_params import SamplingParams
from vllm.triton_utils import tl, triton
from vllm.v1.outputs import LogprobsTensors
from vllm.v1.worker.gpu.input_batch import InputBatch
from vllm.v1.worker.gpu.sample.logprob import compute_topk_logprobs
from vllm.v1.worker.gpu.sample.logprob import compute_topk_scores
class PromptLogprobsWorker:
def __init__(self, max_num_reqs: int):
def __init__(self, max_num_reqs: int, logprobs_mode: LogprobsMode = "raw_logprobs"):
self.max_num_reqs = max_num_reqs
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)
@@ -82,6 +84,7 @@ class PromptLogprobsWorker:
hidden_states[: input_batch.num_tokens],
logits_fn,
max_num_prompt_logprobs,
self.logprobs_mode,
)
)
@@ -206,33 +209,36 @@ def compute_prompt_logprobs_with_chunking(
prompt_hidden_states: torch.Tensor,
logits_fn: Callable[[torch.Tensor], torch.Tensor],
num_prompt_logprobs: int,
logprobs_mode: LogprobsMode = "raw_logprobs",
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
# Since materializing the full prompt logits can take too much memory,
# we compute it in chunks.
CHUNK_SIZE = 1024
token_ids = []
logprobs = []
scores = []
ranks = []
logits_mode = logprobs_mode in ("raw_logits", "processed_logits")
prompt_token_ids = prompt_token_ids.to(torch.int64)
for start_idx in range(0, prompt_token_ids.shape[0], CHUNK_SIZE):
end_idx = start_idx + CHUNK_SIZE
# NOTE(woosuk): logits_fn can be slow because it involves all-gather.
prompt_logits = logits_fn(prompt_hidden_states[start_idx:end_idx])
requested_num_prompt_logprobs = (
requested_num = (
prompt_logits.shape[-1]
if num_prompt_logprobs == -1
else num_prompt_logprobs
)
prompt_logprobs = compute_topk_logprobs(
result = compute_topk_scores(
prompt_logits,
requested_num_prompt_logprobs,
requested_num,
prompt_token_ids[start_idx:end_idx],
logits_mode=logits_mode,
)
token_ids.append(prompt_logprobs.logprob_token_ids)
logprobs.append(prompt_logprobs.logprobs)
ranks.append(prompt_logprobs.selected_token_ranks)
token_ids.append(result.logprob_token_ids)
scores.append(result.logprobs)
ranks.append(result.selected_token_ranks)
token_ids = torch.cat(token_ids, dim=0) if len(token_ids) > 1 else token_ids[0]
logprobs = torch.cat(logprobs, dim=0) if len(logprobs) > 1 else logprobs[0]
scores = torch.cat(scores, dim=0) if len(scores) > 1 else scores[0]
ranks = torch.cat(ranks, dim=0) if len(ranks) > 1 else ranks[0]
return token_ids, logprobs, ranks
return token_ids, scores, ranks
+8 -6
View File
@@ -19,7 +19,7 @@ from vllm.v1.worker.gpu.sample.gumbel import gumbel_sample
from vllm.v1.worker.gpu.sample.logit_bias import LogitBiasState
from vllm.v1.worker.gpu.sample.logprob import (
LogprobTokenIdsState,
compute_topk_logprobs,
compute_topk_scores,
)
from vllm.v1.worker.gpu.sample.output import SamplerOutput
from vllm.v1.worker.gpu.sample.penalties import PenaltiesState
@@ -38,8 +38,6 @@ class Sampler:
num_speculative_tokens: int = 1,
use_fp64_gumbel: bool = False,
):
if logprobs_mode not in ("processed_logprobs", "raw_logprobs"):
raise NotImplementedError(f"Unsupported logprobs_mode: {logprobs_mode}")
self.logprobs_mode = logprobs_mode
self.compute_nans = envs.VLLM_COMPUTE_NANS_IN_LOGITS # False by default.
self.use_fp64_gumbel = use_fp64_gumbel
@@ -102,12 +100,12 @@ class Sampler:
)
if return_logprobs:
if self.logprobs_mode == "processed_logprobs":
if self.logprobs_mode in ("processed_logprobs", "processed_logits"):
logits = processed_logits
expanded_logits = logits.shape[0] != idx_mapping_np.shape[0]
cu_num_logits = cu_num_logits_np.tolist() if expanded_logits else None
num_logprobs = max_num_logprobs if max_num_logprobs != NO_LOGPROBS else 0
logprobs_tensors = compute_topk_logprobs(
logprobs_tensors = compute_topk_scores(
logits,
num_logprobs,
sampled,
@@ -115,6 +113,7 @@ class Sampler:
logprob_token_ids_state=self.logprob_token_ids_state,
expanded_idx_mapping=input_batch.expanded_idx_mapping,
max_per_req_token_ids=max_per_req_token_ids,
logits_mode=self.logprobs_mode in ("raw_logits", "processed_logits"),
)
else:
logprobs_tensors = None
@@ -222,7 +221,10 @@ class Sampler:
# any greedy requests or per-request seeds, or if post-processed
# logprobs need to be returned for any requests.
(top_k is None and top_p is None)
or (return_logprobs and self.logprobs_mode == "processed_logprobs")
or (
return_logprobs
and self.logprobs_mode in ("processed_logprobs", "processed_logits")
)
or self.sampling_states.any_greedy(idx_mapping_np)
or self.sampling_states.any_explicit_seed(idx_mapping_np)
)
@@ -11,7 +11,7 @@ from vllm.v1.worker.gpu.input_batch import (
get_num_sampled_and_rejected,
)
from vllm.v1.worker.gpu.metrics.logits import get_num_nans
from vllm.v1.worker.gpu.sample.logprob import compute_topk_logprobs
from vllm.v1.worker.gpu.sample.logprob import compute_topk_scores
from vllm.v1.worker.gpu.sample.output import SamplerOutput
from vllm.v1.worker.gpu.sample.sampler import Sampler
from vllm.v1.worker.gpu.sample.states import NO_LOGPROBS
@@ -91,11 +91,13 @@ class RejectionSampler:
num_warps=1,
)
expanded_logits = num_logits != input_batch.idx_mapping.shape[0]
return compute_topk_logprobs(
return compute_topk_scores(
logits,
max_num_logprobs,
flat_sampled,
input_batch.cu_num_logits_np.tolist() if expanded_logits else None,
logits_mode=self.sampler.logprobs_mode
in ("raw_logits", "processed_logits"),
)
def __call__(
@@ -139,7 +141,7 @@ class RejectionSampler:
sampled,
num_sampled,
processed_logits
if self.sampler.logprobs_mode == "processed_logprobs"
if self.sampler.logprobs_mode in ("processed_logprobs", "processed_logits")
else logits,
)
+8 -3
View File
@@ -5635,10 +5635,15 @@ class GPUModelRunner(
# to gather the logprob for.
tgt_token_ids = prompt_token_ids[start_tok : start_tok + num_logits]
# Compute prompt logprobs.
logprobs = self.sampler.compute_logprobs(logits)
# Compute prompt scores respecting logprobs_mode.
# NOTE: prompt tokens skip sampling processors, so
# processed_* and raw_* yield the same scores here.
if self.model_config.logprobs_mode in ("raw_logits", "processed_logits"):
scores = logits.to(torch.float32)
else:
scores = self.sampler.compute_logprobs(logits)
token_ids, logprobs, ranks, _ = self.sampler.gather_logprobs(
logprobs, num_prompt_logprobs, tgt_token_ids
scores, num_prompt_logprobs, tgt_token_ids
)
# Transfer GPU->CPU async.