diff --git a/tests/v1/sample/test_logprobs.py b/tests/v1/sample/test_logprobs.py index fba240fea6a..aa17d2a1004 100644 --- a/tests/v1/sample/test_logprobs.py +++ b/tests/v1/sample/test_logprobs.py @@ -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. diff --git a/vllm/config/model.py b/vllm/config/model.py index e36b672cd82..6b032ae7621 100644 --- a/vllm/config/model.py +++ b/vllm/config/model.py @@ -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 diff --git a/vllm/config/vllm.py b/vllm/config/vllm.py index 2e57cb1ec8e..cca9adbcf58 100644 --- a/vllm/config/vllm.py +++ b/vllm/config/vllm.py @@ -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") diff --git a/vllm/model_executor/models/diffusion_gemma.py b/vllm/model_executor/models/diffusion_gemma.py index 11a10131df1..70566871e09 100644 --- a/vllm/model_executor/models/diffusion_gemma.py +++ b/vllm/model_executor/models/diffusion_gemma.py @@ -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 diff --git a/vllm/v1/sample/rejection_sampler.py b/vllm/v1/sample/rejection_sampler.py index 69309213836..0f56b8a4a2c 100644 --- a/vllm/v1/sample/rejection_sampler.py +++ b/vllm/v1/sample/rejection_sampler.py @@ -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 ( diff --git a/vllm/v1/worker/gpu/model_runner.py b/vllm/v1/worker/gpu/model_runner.py index 2bb52e2fd89..518d12a6b28 100644 --- a/vllm/v1/worker/gpu/model_runner.py +++ b/vllm/v1/worker/gpu/model_runner.py @@ -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, diff --git a/vllm/v1/worker/gpu/sample/logprob.py b/vllm/v1/worker/gpu/sample/logprob.py index cb2cf1a590e..c0f31cdf9f0 100644 --- a/vllm/v1/worker/gpu/sample/logprob.py +++ b/vllm/v1/worker/gpu/sample/logprob.py @@ -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, ) diff --git a/vllm/v1/worker/gpu/sample/prompt_logprob.py b/vllm/v1/worker/gpu/sample/prompt_logprob.py index b89ebac35d9..4d4cc244825 100644 --- a/vllm/v1/worker/gpu/sample/prompt_logprob.py +++ b/vllm/v1/worker/gpu/sample/prompt_logprob.py @@ -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 diff --git a/vllm/v1/worker/gpu/sample/sampler.py b/vllm/v1/worker/gpu/sample/sampler.py index b269de9eaed..f0a83c92efb 100644 --- a/vllm/v1/worker/gpu/sample/sampler.py +++ b/vllm/v1/worker/gpu/sample/sampler.py @@ -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) ) diff --git a/vllm/v1/worker/gpu/spec_decode/rejection_sampler.py b/vllm/v1/worker/gpu/spec_decode/rejection_sampler.py index c56252d55d7..4753d281746 100644 --- a/vllm/v1/worker/gpu/spec_decode/rejection_sampler.py +++ b/vllm/v1/worker/gpu/spec_decode/rejection_sampler.py @@ -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, ) diff --git a/vllm/v1/worker/gpu_model_runner.py b/vllm/v1/worker/gpu_model_runner.py index 2c7adaaf2ae..f3f0afb1301 100644 --- a/vllm/v1/worker/gpu_model_runner.py +++ b/vllm/v1/worker/gpu_model_runner.py @@ -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.