From 22429d07b35c9f6dc671ae23de12f8d7e07e4f9a Mon Sep 17 00:00:00 2001 From: AlpinDale Date: Tue, 19 Nov 2024 22:42:06 +0000 Subject: [PATCH 1/6] feat: implement top-nsigma sampling method --- aphrodite/common/sampling_params.py | 12 ++++++++++ aphrodite/endpoints/openai/protocol.py | 4 ++++ aphrodite/modeling/layers/sampler.py | 29 ++++++++++++++++++++++++- aphrodite/modeling/sampling_metadata.py | 22 ++++++++++++++----- 4 files changed, 60 insertions(+), 7 deletions(-) diff --git a/aphrodite/common/sampling_params.py b/aphrodite/common/sampling_params.py index 53ef398655..998a9e44ee 100644 --- a/aphrodite/common/sampling_params.py +++ b/aphrodite/common/sampling_params.py @@ -148,6 +148,11 @@ class SamplingParams( above this threshold, consider removing all but the last one. xtc_probability: Probability that the removal will actually happen. 0 disables the sampler, 1 makes it always happen. + nsigma: Number of standard deviations from the maximum logit to use + as a cutoff threshold. Tokens with logits below + (max_logit - nsgima * std_dev) are filtered out. Higher values + (e.g. 3.0) keep more tokens, lower values (e.g. 1.0) are more + selective. Must be positive. 0 to disable. """ n: int = 1 @@ -193,6 +198,7 @@ class SamplingParams( truncate_prompt_tokens: Optional[Annotated[int, msgspec.Meta(ge=1)]] = None xtc_threshold: float = 0.1 xtc_probability: float = 0 + nsigma: float = 0.0 # The below fields are not supposed to be used as an input. # They are set in post_init. @@ -239,6 +245,7 @@ class SamplingParams( "truncate_prompt_tokens": None, "xtc_threshold": 0.1, "xtc_probability": 0, + "nsigma": 0.0, } def __post_init__(self) -> None: @@ -368,6 +375,11 @@ def _verify_args(self) -> None: raise ValueError( "xtc_probability must be in [0, 1], got " f"{self.xtc_probability}.") + if not self.nsigma <= 0.0: + raise ValueError( + "nsigma must be non-negative, got " + f"{self.nsigma}.") + def _verify_beam_search(self) -> None: if self.best_of == 1: diff --git a/aphrodite/endpoints/openai/protocol.py b/aphrodite/endpoints/openai/protocol.py index 8b49167986..ca6ffc52d9 100644 --- a/aphrodite/endpoints/openai/protocol.py +++ b/aphrodite/endpoints/openai/protocol.py @@ -150,6 +150,7 @@ class ChatCompletionRequest(OpenAIBaseModel): dynatemp_min: Optional[float] = 0.0 dynatemp_max: Optional[float] = 0.0 dynatemp_exponent: Optional[float] = 1.0 + nsigma: Optional[float] = 0.0 custom_token_bans: Optional[List[int]] = None # doc: end-chat-completion-sampling-params @@ -293,6 +294,7 @@ def to_sampling_params( dynatemp_min=self.dynatemp_min, dynatemp_max=self.dynatemp_max, dynatemp_exponent=self.dynatemp_exponent, + nsigma=self.nsigma, custom_token_bans=self.custom_token_bans, ) @@ -404,6 +406,7 @@ class CompletionRequest(OpenAIBaseModel): dynatemp_min: Optional[float] = 0.0 dynatemp_max: Optional[float] = 0.0 dynatemp_exponent: Optional[float] = 1.0 + nsigma: Optional[float] = 0.0 custom_token_bans: Optional[List[int]] = None # doc: end-completion-sampling-params @@ -506,6 +509,7 @@ def to_sampling_params( dynatemp_min=self.dynatemp_min, dynatemp_max=self.dynatemp_max, dynatemp_exponent=self.dynatemp_exponent, + nsigma=self.nsigma, custom_token_bans=self.custom_token_bans, ) diff --git a/aphrodite/modeling/layers/sampler.py b/aphrodite/modeling/layers/sampler.py index 2b8042557c..243c7c3d6b 100644 --- a/aphrodite/modeling/layers/sampler.py +++ b/aphrodite/modeling/layers/sampler.py @@ -77,7 +77,7 @@ def _init_sampling_tensors( # Initialize new sampling tensors (sampling_tensors, do_penalties, do_temperatures, do_top_p_top_k, do_top_as, do_min_p, do_tfss, do_eta_cutoffs, do_epsilon_cutoffs, - do_typical_ps, do_quadratic, do_xtc, do_temp_last + do_typical_ps, do_quadratic, do_xtc, do_nsigmas, do_temp_last ) = SamplingTensors.from_sampling_metadata( sampling_metadata, vocab_size, logits.device, logits.dtype) @@ -93,6 +93,7 @@ def _init_sampling_tensors( self._do_typical_ps = do_typical_ps self._do_quadratic = do_quadratic self._do_xtc = do_xtc + self._do_nsgimas = do_nsigmas self._do_temp_last = do_temp_last def forward( @@ -131,6 +132,7 @@ def forward( do_typical_ps = self._do_typical_ps do_quadratic = self._do_quadratic do_xtc = self._do_xtc + do_nsigmas = self._do_nsgimas do_temp_last = self._do_temp_last logits = _apply_min_tokens_penalty(logits, sampling_metadata) @@ -150,6 +152,9 @@ def forward( sampling_tensors.dynatemp_maxs, sampling_tensors.dynatemp_exps) + if do_nsigmas: + logits = _apply_top_nsigma(logits, sampling_tensors.nsigmas) + if do_top_p_top_k: logits = _apply_top_k_top_p(logits, sampling_tensors.top_ps, sampling_tensors.top_ks) @@ -634,6 +639,28 @@ def _apply_xtc_sampling( return logits +def _apply_top_nsigma( + logits: torch.Tensor, + nsigma: torch.Tensor, +) -> torch.Tensor: + """Apply top-nsigma truncation to the logits. + + Reference: https://arxiv.org/abs/2411.07641 + + Args: + logits: Logits of shape (num_tokens, vocab_size) + nsigma: Number of standard deviations to use as threshold + Returns: + Modified logits with values below threshold set to -inf + """ + std = logits.std(dim=-1, keepdim=True) + threshold = (logits.max(dim=-1, keepdim=True).values - + nsigma.unsqueeze(dim=1) * std) + logits[logits < threshold] = float("-inf") + + return logits + + def _greedy_sample( selected_seq_groups: List[SequenceGroupToSample], samples: torch.Tensor, diff --git a/aphrodite/modeling/sampling_metadata.py b/aphrodite/modeling/sampling_metadata.py index 499c93dd27..9e43c5e503 100644 --- a/aphrodite/modeling/sampling_metadata.py +++ b/aphrodite/modeling/sampling_metadata.py @@ -387,6 +387,7 @@ class SamplingTensors: smoothing_curves: torch.Tensor xtc_thresholds: torch.Tensor xtc_probabilities: torch.Tensor + nsigmas: torch.Tensor sampling_seeds: torch.Tensor sample_indices: torch.Tensor extra_seeds: Optional[torch.Tensor] @@ -404,7 +405,7 @@ def from_sampling_metadata( extra_seeds_to_generate: int = 0, extra_entropy: Optional[Tuple[int, ...]] = None ) -> Tuple["SamplingTensors", bool, bool, bool, bool, bool, bool, bool, - bool, bool, bool, bool, bool]: + bool, bool, bool, bool, bool, bool]: """ extra_seeds_to_generate: extra seeds to generate using the user-defined seed for each sequence. @@ -432,6 +433,7 @@ def from_sampling_metadata( smoothing_curves: List[float] = [] xtc_thresholds: List[float] = [] xtc_probabilities: List[float] = [] + nsigmas: List[float] = [] sampling_seeds: List[List[int]] = [] sample_indices: List[int] = [] do_penalties = False @@ -445,6 +447,7 @@ def from_sampling_metadata( do_typical_ps = False do_quadratic = False do_xtc = False + do_nsigmas = False do_temp_last = False if _USE_TRITON_SAMPLER: @@ -487,6 +490,7 @@ def from_sampling_metadata( do_quadratic |= (params.smoothing_factor > _SAMPLING_EPS or params.smoothing_curve > 1.0) do_xtc |= params.xtc_probability > _SAMPLING_EPS + do_nsigmas |= params.nsigma > _SAMPLING_EPS do_temp_last |= params.temperature_last is_prompt = seq_group.is_prompt @@ -521,6 +525,7 @@ def from_sampling_metadata( smoothing_curves += [params.smoothing_curve] * n_seqs xtc_thresholds += [params.xtc_threshold] * n_seqs xtc_probabilities += [params.xtc_probability] * n_seqs + nsigmas += [params.nsigma] * n_seqs if _USE_TRITON_SAMPLER: if is_prompt: @@ -567,13 +572,13 @@ def from_sampling_metadata( temperature_lasts, top_ps, top_ks, top_as, min_ps, presence_penalties, frequency_penalties, repetition_penalties, tfss, eta_cutoffs, epsilon_cutoffs, typical_ps, smoothing_factors, - smoothing_curves, xtc_thresholds, xtc_probabilities, sampling_seeds, - sample_indices, prompt_tokens, output_tokens, vocab_size, - extra_seeds_to_generate, device, dtype) + smoothing_curves, xtc_thresholds, xtc_probabilities, nsigmas, + sampling_seeds, sample_indices, prompt_tokens, output_tokens, + vocab_size, extra_seeds_to_generate, device, dtype) return (sampling_tensors, do_penalties, do_temperatures, do_top_p_top_k, do_top_as, do_min_p, do_tfss, do_eta_cutoffs, do_epsilon_cutoffs, do_typical_ps, do_quadratic, do_xtc, - do_temp_last) + do_nsigmas, do_temp_last) @classmethod def from_lists(cls, temperatures: List[float], dynatemp_mins: List[float], @@ -586,7 +591,7 @@ def from_lists(cls, temperatures: List[float], dynatemp_mins: List[float], eta_cutoffs: List[float], epsilon_cutoffs: List[float], typical_ps: List[float], smoothing_factors: List[float], smoothing_curves: List[float], xtc_thresholds: List[float], - xtc_probabilities: List[float], + xtc_probabilities: List[float], nsigmas: List[float], sampling_seeds: List[List[int]], sample_indices: List[int], prompt_tokens: List[array], output_tokens: List[array], vocab_size: int, @@ -719,6 +724,10 @@ def from_lists(cls, temperatures: List[float], dynatemp_mins: List[float], device="cpu", dtype=dtype, pin_memory=pin_memory) + nsigmas_t = torch.tensor(nsigmas, + device="cpu", + dtype=dtype, + pin_memory=pin_memory) sample_indices_t = torch.tensor( sample_indices, device="cpu", @@ -775,6 +784,7 @@ def from_lists(cls, temperatures: List[float], dynatemp_mins: List[float], non_blocking=True), xtc_probabilities=xtc_probabilities_t.to(device=device, non_blocking=True), + nsigmas=nsigmas_t.to(device=device, non_blocking=True), typical_ps=typical_ps_t.to(device=device, non_blocking=True), prompt_tokens=prompt_t.to(device=device, non_blocking=True), output_tokens=output_t.to(device=device, non_blocking=True), From 22423eff8a565819c4606ee8d6a910123a1b3e62 Mon Sep 17 00:00:00 2001 From: AlpinDale Date: Tue, 19 Nov 2024 23:14:30 +0000 Subject: [PATCH 2/6] fix comparison sign --- aphrodite/common/sampling_params.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/aphrodite/common/sampling_params.py b/aphrodite/common/sampling_params.py index 998a9e44ee..9af5fb9154 100644 --- a/aphrodite/common/sampling_params.py +++ b/aphrodite/common/sampling_params.py @@ -375,7 +375,7 @@ def _verify_args(self) -> None: raise ValueError( "xtc_probability must be in [0, 1], got " f"{self.xtc_probability}.") - if not self.nsigma <= 0.0: + if self.nsigma <= 0.0: raise ValueError( "nsigma must be non-negative, got " f"{self.nsigma}.") From 2242654e0201b71ece7356115f42432264838850 Mon Sep 17 00:00:00 2001 From: AlpinDale Date: Tue, 19 Nov 2024 23:15:19 +0000 Subject: [PATCH 3/6] the above fix: fixed --- aphrodite/common/sampling_params.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/aphrodite/common/sampling_params.py b/aphrodite/common/sampling_params.py index 9af5fb9154..b03dc2d36c 100644 --- a/aphrodite/common/sampling_params.py +++ b/aphrodite/common/sampling_params.py @@ -375,7 +375,7 @@ def _verify_args(self) -> None: raise ValueError( "xtc_probability must be in [0, 1], got " f"{self.xtc_probability}.") - if self.nsigma <= 0.0: + if self.nsigma < 0.0: raise ValueError( "nsigma must be non-negative, got " f"{self.nsigma}.") From 22423892e267547a43782fb5f131bcb0d4987d27 Mon Sep 17 00:00:00 2001 From: AlpinDale Date: Wed, 20 Nov 2024 00:43:51 +0000 Subject: [PATCH 4/6] hijack epsilon cutoff --- aphrodite/endpoints/openai/protocol.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/aphrodite/endpoints/openai/protocol.py b/aphrodite/endpoints/openai/protocol.py index ca6ffc52d9..25acd53a12 100644 --- a/aphrodite/endpoints/openai/protocol.py +++ b/aphrodite/endpoints/openai/protocol.py @@ -482,7 +482,7 @@ def to_sampling_params( top_a=self.top_a, tfs=self.tfs, eta_cutoff=self.eta_cutoff, - epsilon_cutoff=self.epsilon_cutoff, + # epsilon_cutoff=self.epsilon_cutoff, typical_p=self.typical_p, smoothing_factor=self.smoothing_factor, smoothing_curve=self.smoothing_curve, @@ -509,7 +509,7 @@ def to_sampling_params( dynatemp_min=self.dynatemp_min, dynatemp_max=self.dynatemp_max, dynatemp_exponent=self.dynatemp_exponent, - nsigma=self.nsigma, + nsigma=self.epsilon_cutoff, custom_token_bans=self.custom_token_bans, ) From 22421c9aaeb9646d3d4d33809f11607eb41fda05 Mon Sep 17 00:00:00 2001 From: AlpinDale Date: Wed, 20 Nov 2024 02:29:00 +0000 Subject: [PATCH 5/6] Revert "hijack epsilon cutoff " This reverts commit 22423892e267547a43782fb5f131bcb0d4987d27. --- aphrodite/endpoints/openai/protocol.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/aphrodite/endpoints/openai/protocol.py b/aphrodite/endpoints/openai/protocol.py index 25acd53a12..ca6ffc52d9 100644 --- a/aphrodite/endpoints/openai/protocol.py +++ b/aphrodite/endpoints/openai/protocol.py @@ -482,7 +482,7 @@ def to_sampling_params( top_a=self.top_a, tfs=self.tfs, eta_cutoff=self.eta_cutoff, - # epsilon_cutoff=self.epsilon_cutoff, + epsilon_cutoff=self.epsilon_cutoff, typical_p=self.typical_p, smoothing_factor=self.smoothing_factor, smoothing_curve=self.smoothing_curve, @@ -509,7 +509,7 @@ def to_sampling_params( dynatemp_min=self.dynatemp_min, dynatemp_max=self.dynatemp_max, dynatemp_exponent=self.dynatemp_exponent, - nsigma=self.epsilon_cutoff, + nsigma=self.nsigma, custom_token_bans=self.custom_token_bans, ) From 2242029fa9de407c34325bc7ca7eaf55fe3533d2 Mon Sep 17 00:00:00 2001 From: AlpinDale Date: Wed, 20 Nov 2024 03:14:35 +0000 Subject: [PATCH 6/6] add tests --- pytest.ini | 7 ++++++ tests/samplers/test_sampler.py | 46 ++++++++++++++++++++++++++++++++++ 2 files changed, 53 insertions(+) create mode 100644 pytest.ini diff --git a/pytest.ini b/pytest.ini new file mode 100644 index 0000000000..7600b0470e --- /dev/null +++ b/pytest.ini @@ -0,0 +1,7 @@ +[pytest] +filterwarnings = + ignore::DeprecationWarning:pkg_resources.* + ignore::DeprecationWarning:pyairports.* + ignore::pytest_asyncio.plugin.PytestDeprecationWarning +asyncio_mode = strict +asyncio_default_fixture_loop_scope = function \ No newline at end of file diff --git a/tests/samplers/test_sampler.py b/tests/samplers/test_sampler.py index 226877d75c..32fb31a127 100644 --- a/tests/samplers/test_sampler.py +++ b/tests/samplers/test_sampler.py @@ -732,6 +732,52 @@ def test_sampling_params(sampling_params: List[SamplingParams]): assert tokens1[1] == tokens2[0] +@pytest.mark.parametrize("seed", RANDOM_SEEDS) +@pytest.mark.parametrize("device", CUDA_DEVICES) +def test_sampler_nsigma(seed: int, device: str): + """Test that top-nsigma sampling behaves as expected.""" + set_random_seed(seed) + torch.set_default_device(device) + batch_size = random.randint(1, 256) + _, fake_logits, sampler = _prepare_test(batch_size) + + # Create a clear separation in logits for testing + high_logit_indices = {} # Store high logit indices for each batch + for i in range(batch_size): + # Set a few logits significantly higher than others + num_high_logits = random.randint(1, 5) + high_indices = random.sample(range(fake_logits.size(1)), + num_high_logits) + high_logit_indices[i] = set(high_indices) # Store for verification + for idx in high_indices: + fake_logits[i, idx] = 10.0 # Clearly above the mean + + # Test with different nsigma values + for nsigma in [1.5, 2.0, 3.0]: + sampling_params = SamplingParams( + temperature=1.0, + nsigma=nsigma, + seed=random.randint(0, 10000), + ) + + sampler_output = _do_sample(batch_size, fake_logits.clone(), sampler, + sampling_params, device) + + # Verify that sampling only selects from high logits + for batch_idx, sequence_output in enumerate(sampler_output): + for nth_output in sequence_output.samples: + token_id = nth_output.output_token + # The token should come from the high logits region + assert token_id in high_logit_indices[batch_idx], \ + f"Sampled token {token_id} for batch {batch_idx} was not in the high logit set" # noqa + + # Test determinism + second_output = _do_sample(batch_size, fake_logits.clone(), sampler, + sampling_params, device) + assert sampler_output == second_output, \ + "Top-nsigma sampling is not deterministic with same seed" + + @pytest.mark.parametrize("device", CUDA_DEVICES) def test_sampler_include_gpu_probs_tensor(device: str): set_random_seed(42)