From b3f1b78d2e0c7c8b96d550815ab7a23ee17e7d7e Mon Sep 17 00:00:00 2001 From: AlpinDale Date: Thu, 23 Jul 2026 02:17:27 +0430 Subject: [PATCH] feat: support min_p with speculative decoding --- aphrodite/sampling_params.py | 6 +-- .../v1/sample/logits_processor/__init__.py | 9 +++- aphrodite/v1/sample/rejection_sampler.py | 23 +++++++- tests/v1/sample/test_rejection_sampler.py | 52 ++++++++++++++++++- 4 files changed, 82 insertions(+), 8 deletions(-) diff --git a/aphrodite/sampling_params.py b/aphrodite/sampling_params.py index 43cc3d6b9b..c65798f41e 100644 --- a/aphrodite/sampling_params.py +++ b/aphrodite/sampling_params.py @@ -1003,10 +1003,8 @@ def _validate_spec_decode( return # Some sampling parameters are not yet compatible with spec decoding. - if self.min_p > _SAMPLING_EPS or self.logit_bias: - raise ValueError( - "The min_p and logit_bias sampling parameters are not yet supported with speculative decoding." - ) + if self.logit_bias: + raise ValueError("The logit_bias sampling parameter is not yet supported with speculative decoding.") def _validate_diffusion(self, model_config: ModelConfig) -> None: if not model_config.is_diffusion: diff --git a/aphrodite/v1/sample/logits_processor/__init__.py b/aphrodite/v1/sample/logits_processor/__init__.py index 7d98128afa..7cd36b2ed6 100644 --- a/aphrodite/v1/sample/logits_processor/__init__.py +++ b/aphrodite/v1/sample/logits_processor/__init__.py @@ -187,8 +187,13 @@ def build_logitsprocs( if aphrodite_config.speculative_config: if custom_logitsprocs: raise ValueError(STR_SPEC_DEC_REJECTS_LOGITSPROCS) - logger.warning("min_p and logit_bias parameters won't work with speculative decoding.") - return LogitsProcessors([MinTokensLogitsProcessor(aphrodite_config, device, is_pin_memory)]) + logger.warning("logit_bias parameter won't work with speculative decoding.") + return LogitsProcessors( + [ + MinTokensLogitsProcessor(aphrodite_config, device, is_pin_memory), + MinPLogitsProcessor(aphrodite_config, device, is_pin_memory), + ] + ) custom_logitsprocs_classes = _load_custom_logitsprocs(custom_logitsprocs) return LogitsProcessors( diff --git a/aphrodite/v1/sample/rejection_sampler.py b/aphrodite/v1/sample/rejection_sampler.py index 107c47c161..d742779704 100644 --- a/aphrodite/v1/sample/rejection_sampler.py +++ b/aphrodite/v1/sample/rejection_sampler.py @@ -12,7 +12,10 @@ from aphrodite.logger import init_logger from aphrodite.triton_utils import tl, triton from aphrodite.v1.outputs import LogprobsLists, LogprobsTensors, SamplerOutput -from aphrodite.v1.sample.logits_processor.builtin import MinTokensLogitsProcessor +from aphrodite.v1.sample.logits_processor.builtin import ( + MinPLogitsProcessor, + MinTokensLogitsProcessor, +) from aphrodite.v1.sample.metadata import SamplingMetadata from aphrodite.v1.sample.ops.bad_words import apply_bad_words_with_drafts from aphrodite.v1.sample.ops.penalties import apply_all_penalties @@ -515,6 +518,24 @@ def apply_sampling_constraints( # NOTE(woosuk): Update `logits` in place to avoid allocating a new tensor. logits.div_(temperature.unsqueeze(-1)) + # Apply min_p after temperature scaling and before top-k/top-p, matching + # where MinPLogitsProcessor runs in the non-spec sampling path. The + # processor's per-request state is expanded to per-token rows here since + # its own apply() assumes one logits row per request. + min_p_processor = next( + (proc for proc in sampling_metadata.logitsprocs.argmax_invariant if isinstance(proc, MinPLogitsProcessor)), + None, + ) + if min_p_processor is not None and min_p_processor.min_p_count: + min_p = expand_batch_to_tokens( + min_p_processor.min_p.squeeze(-1), + cu_num_draft_tokens, + num_tokens, + ) + probs = logits.softmax(dim=-1) + threshold = probs.amax(dim=-1, keepdim=True).mul_(min_p.unsqueeze(-1)) + logits.masked_fill_(probs < threshold, -float("inf")) + # Get expanded top_k and top_p tensors. top_k = None if sampling_metadata.top_k is not None: diff --git a/tests/v1/sample/test_rejection_sampler.py b/tests/v1/sample/test_rejection_sampler.py index 989aab47b2..5382998416 100644 --- a/tests/v1/sample/test_rejection_sampler.py +++ b/tests/v1/sample/test_rejection_sampler.py @@ -76,6 +76,7 @@ def create_sampling_metadata( repetition_penalties: list[float] | None = None, bad_words_token_ids: dict[int, list[list[int]]] | None = None, allowed_token_ids_mask: torch.Tensor | None = None, + logitsprocs: LogitsProcessors | None = None, ) -> SamplingMetadata: """Create a v1 sampling metadata object with all_greedy set to the given value. Either all greedy or all random sampling @@ -119,7 +120,7 @@ def create_sampling_metadata( spec_token_ids=[] if spec_token_ids is None else spec_token_ids, allowed_token_ids_mask=allowed_token_ids_mask, bad_words_token_ids={} if bad_words_token_ids is None else bad_words_token_ids, - logitsprocs=LogitsProcessors(), + logitsprocs=logitsprocs if logitsprocs is not None else LogitsProcessors(), ) @@ -690,6 +691,55 @@ def test_top_p(rejection_sampler, top_p): ) +@pytest.mark.parametrize("min_p", [0.1, 0.5, 0.9]) +def test_min_p(rejection_sampler, min_p): + """Test rejection sampling with min-p sampling""" + from types import SimpleNamespace + + from aphrodite.v1.sample.logits_processor import MinPLogitsProcessor + + vocab_size = 100 + batch_size = 100 + num_draft_tokens = 3 + num_tokens = batch_size * num_draft_tokens + + target_logits = torch.randn((num_tokens, vocab_size), device=DEVICE_TYPE) + temperature = torch.ones(batch_size, dtype=torch.float32, device=DEVICE_TYPE) + + # With temperature=1, min_p thresholds on softmax of the raw logits. + probs = (target_logits / temperature[0]).softmax(dim=-1) + threshold = probs.amax(dim=-1, keepdim=True) * min_p + min_p_indices = [] + for i in range(num_tokens): + min_p_indices.append(torch.nonzero(probs[i] >= threshold[i]).flatten().tolist()) + + # Build a MinPLogitsProcessor with populated per-request state, as + # build_logitsprocs does under spec decode. + fake_config = SimpleNamespace(scheduler_config=SimpleNamespace(max_num_seqs=batch_size)) + min_p_proc = MinPLogitsProcessor(fake_config, torch.device(DEVICE_TYPE), is_pin_memory=False) + min_p_proc.min_p_cpu[:batch_size] = min_p + min_p_proc.min_p_count = batch_size + min_p_proc.min_p = min_p_proc.min_p_device[:batch_size] + min_p_proc.min_p.copy_(min_p_proc.min_p_cpu_tensor[:batch_size]) + min_p_proc.min_p.unsqueeze_(1) + + sampling_metadata = create_sampling_metadata( + all_greedy=False, + temperature=temperature, + logitsprocs=LogitsProcessors([min_p_proc]), + ) + + _test_masked_logits( + rejection_sampler, + batch_size=batch_size, + num_draft_tokens=num_draft_tokens, + vocab_size=vocab_size, + target_logits=target_logits, + unmasked_indices=min_p_indices, + sampling_metadata=sampling_metadata, + ) + + ########################### Tests for Logit Processors ################### def test_frequency_penalties(rejection_sampler): """Test rejection sampling with frequency penalties"""