Skip to content
Closed
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
12 changes: 12 additions & 0 deletions aphrodite/common/sampling_params.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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.
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -368,6 +375,11 @@ def _verify_args(self) -> None:
raise ValueError(
"xtc_probability must be in [0, 1], got "
f"{self.xtc_probability}.")
if 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:
Expand Down
4 changes: 4 additions & 0 deletions aphrodite/endpoints/openai/protocol.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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,
)

Expand Down Expand Up @@ -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

Expand Down Expand Up @@ -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,
)

Expand Down
29 changes: 28 additions & 1 deletion aphrodite/modeling/layers/sampler.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)

Expand All @@ -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(
Expand Down Expand Up @@ -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)
Expand All @@ -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)
Expand Down Expand Up @@ -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,
Expand Down
22 changes: 16 additions & 6 deletions aphrodite/modeling/sampling_metadata.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]
Expand All @@ -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.
Expand Down Expand Up @@ -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
Expand All @@ -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:
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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],
Expand All @@ -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,
Expand Down Expand Up @@ -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",
Expand Down Expand Up @@ -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),
Expand Down
7 changes: 7 additions & 0 deletions pytest.ini
Original file line number Diff line number Diff line change
@@ -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
46 changes: 46 additions & 0 deletions tests/samplers/test_sampler.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down