621 lines
24 KiB
Python
621 lines
24 KiB
Python
"""Unit tests for srt/sampling/sampling_params.py — no server, no model loading."""
|
|
|
|
from sglang.test.ci.ci_register import register_cpu_ci, register_xpu_ci
|
|
|
|
register_cpu_ci(est_time=10, suite="base-a-test-cpu")
|
|
register_cpu_ci(est_time=8, suite="stage-b-test-cpu-intel")
|
|
register_xpu_ci(est_time=10, suite="stage-a-test-1-gpu-xpu")
|
|
|
|
import copy
|
|
import re
|
|
import unittest
|
|
from pathlib import Path
|
|
from unittest.mock import MagicMock
|
|
|
|
import msgspec
|
|
|
|
from sglang.srt.sampling.sampling_params import (
|
|
MAX_LEN,
|
|
MAX_REQUEST_REASONING_END_TOKEN_IDS,
|
|
MAX_STOP_COUNT,
|
|
MAX_STOP_REGEX_COUNT,
|
|
MAX_STOP_REGEX_LEN,
|
|
REQUEST_REASONING_END_TOKEN_IDS_KEY,
|
|
TOP_K_ALL,
|
|
SamplingParams,
|
|
get_max_seq_length,
|
|
)
|
|
from sglang.test.test_utils import CustomTestCase
|
|
|
|
|
|
class TestSamplingParamsInit(CustomTestCase):
|
|
def test_zero_temperature_becomes_greedy(self):
|
|
"""Test greedy conversion when temperature is 0."""
|
|
sp = SamplingParams(temperature=0.0)
|
|
self.assertEqual(sp.top_k, 1)
|
|
self.assertEqual(sp.temperature, 1.0)
|
|
|
|
def test_near_zero_temperature_becomes_greedy(self):
|
|
"""Test greedy conversion when temperature is near zero (1e-7)."""
|
|
sp = SamplingParams(temperature=1e-7)
|
|
self.assertEqual(sp.top_k, 1)
|
|
self.assertEqual(sp.temperature, 1.0)
|
|
|
|
def test_temperature_at_eps_boundary_not_greedy(self):
|
|
"""Test that temperature exactly at 1e-6 does not trigger greedy (strict <)."""
|
|
sp = SamplingParams(temperature=1e-6)
|
|
self.assertEqual(sp.temperature, 1e-6)
|
|
# top_k should remain at TOP_K_ALL (from -1 default)
|
|
self.assertEqual(sp.top_k, TOP_K_ALL)
|
|
|
|
def test_negative_temperature_not_modified(self):
|
|
"""Test that __init__ preserves negative temperature (rejected by verify instead)."""
|
|
sp = SamplingParams(temperature=-1.0)
|
|
self.assertEqual(sp.temperature, -1.0)
|
|
|
|
def test_top_k_minus_one_becomes_top_k_all(self):
|
|
"""Test that top_k=-1 is converted to TOP_K_ALL (whole vocabulary)."""
|
|
sp = SamplingParams(top_k=-1)
|
|
self.assertEqual(sp.top_k, TOP_K_ALL)
|
|
|
|
def test_positive_top_k_preserved(self):
|
|
"""Test that explicit positive top_k is kept as-is."""
|
|
sp = SamplingParams(top_k=50)
|
|
self.assertEqual(sp.top_k, 50)
|
|
|
|
def test_stop_token_ids_stored_as_set(self):
|
|
"""Test that stop_token_ids list is converted to set."""
|
|
sp = SamplingParams(stop_token_ids=[1, 2, 3])
|
|
self.assertIsInstance(sp.stop_token_ids, set)
|
|
self.assertEqual(sp.stop_token_ids, {1, 2, 3})
|
|
|
|
def test_stop_token_ids_none_stays_none(self):
|
|
"""Test that None stop_token_ids stays None."""
|
|
sp = SamplingParams(stop_token_ids=None)
|
|
self.assertIsNone(sp.stop_token_ids)
|
|
|
|
def test_empty_stop_token_ids_becomes_none(self):
|
|
"""Test that empty list is treated as None (falsy in Python)."""
|
|
sp = SamplingParams(stop_token_ids=[])
|
|
self.assertIsNone(sp.stop_token_ids)
|
|
|
|
def test_empty_grammar_constraint_becomes_none(self):
|
|
"""An empty grammar string means "unset", not "constrain to nothing".
|
|
Left as "" it reads as set to the is-not-None checks downstream while
|
|
the constraint selection skips it.
|
|
"""
|
|
for field in ("json_schema", "regex", "ebnf", "structural_tag"):
|
|
with self.subTest(field=field):
|
|
sp = SamplingParams(**{field: ""})
|
|
self.assertIsNone(getattr(sp, field))
|
|
|
|
|
|
class TestSamplingParamsVerify(CustomTestCase):
|
|
VOCAB_SIZE = 32000
|
|
GRAMMAR_VALUES = {
|
|
"json_schema": '{"type":"object"}',
|
|
"regex": "abc",
|
|
"ebnf": 'root ::= "abc"',
|
|
"structural_tag": '{"structures":[],"triggers":[]}',
|
|
}
|
|
|
|
def _make(self, **kwargs):
|
|
"""Helper: create SamplingParams with safe defaults, override with kwargs."""
|
|
defaults = dict(temperature=1.0, top_p=1.0, top_k=10, min_p=0.0)
|
|
defaults.update(kwargs)
|
|
return SamplingParams(**defaults)
|
|
|
|
def test_valid_params_pass(self):
|
|
"""Default valid params should pass verify() without raising."""
|
|
sp = self._make()
|
|
sp.verify(self.VOCAB_SIZE)
|
|
|
|
def test_request_reasoning_end_token_ids_are_vocab_bounded_integers(self):
|
|
self._make(
|
|
custom_params={REQUEST_REASONING_END_TOKEN_IDS_KEY: [17, 18]}
|
|
).verify(self.VOCAB_SIZE)
|
|
|
|
invalid_values = [
|
|
[],
|
|
[-1],
|
|
[True],
|
|
[self.VOCAB_SIZE],
|
|
"17",
|
|
list(range(MAX_REQUEST_REASONING_END_TOKEN_IDS + 1)),
|
|
]
|
|
for value in invalid_values:
|
|
with (
|
|
self.subTest(value=value),
|
|
self.assertRaisesRegex(ValueError, "request reasoning end token IDs"),
|
|
):
|
|
self._make(
|
|
custom_params={REQUEST_REASONING_END_TOKEN_IDS_KEY: value}
|
|
).verify(self.VOCAB_SIZE)
|
|
|
|
def test_negative_temperature_raises(self):
|
|
"""Test that verify() rejects negative temperature (must be >= 0)."""
|
|
sp = self._make(temperature=-0.5)
|
|
with self.assertRaises(ValueError):
|
|
sp.verify(self.VOCAB_SIZE)
|
|
|
|
def test_nan_temperature_raises(self):
|
|
"""verify() must reject NaN temperature; the bare < 0.0 check alone lets it through."""
|
|
sp = self._make(temperature=float("nan"))
|
|
with self.assertRaises(ValueError):
|
|
sp.verify(self.VOCAB_SIZE)
|
|
|
|
def test_inf_temperature_raises(self):
|
|
"""verify() must reject non-finite (inf) temperature."""
|
|
sp = self._make(temperature=float("inf"))
|
|
with self.assertRaises(ValueError):
|
|
sp.verify(self.VOCAB_SIZE)
|
|
|
|
# --- top_p ---
|
|
def test_top_p_negative_raises(self):
|
|
"""Test that verify() rejects negative top_p (valid range is (0, 1])."""
|
|
sp = self._make(top_p=-0.5)
|
|
with self.assertRaises(ValueError):
|
|
sp.verify(self.VOCAB_SIZE)
|
|
|
|
def test_top_p_zero_raises(self):
|
|
"""Test that verify() rejects top_p=0 (not in (0, 1])."""
|
|
sp = self._make(top_p=0.0)
|
|
with self.assertRaises(ValueError):
|
|
sp.verify(self.VOCAB_SIZE)
|
|
|
|
def test_top_p_above_one_raises(self):
|
|
"""Test that verify() rejects top_p > 1.0."""
|
|
sp = self._make(top_p=1.1)
|
|
with self.assertRaises(ValueError):
|
|
sp.verify(self.VOCAB_SIZE)
|
|
|
|
def test_top_p_exactly_one_is_valid(self):
|
|
"""Test that top_p=1.0 is accepted (inclusive upper bound)."""
|
|
sp = self._make(top_p=1.0)
|
|
sp.verify(self.VOCAB_SIZE)
|
|
|
|
def test_top_p_small_positive_is_valid(self):
|
|
"""Test that a small positive top_p (0.01) is accepted."""
|
|
sp = self._make(top_p=0.01)
|
|
sp.verify(self.VOCAB_SIZE)
|
|
|
|
# --- min_p ---
|
|
def test_min_p_negative_raises(self):
|
|
"""Test that verify() rejects negative min_p (valid range is [0, 1])."""
|
|
sp = self._make(min_p=-0.1)
|
|
with self.assertRaises(ValueError):
|
|
sp.verify(self.VOCAB_SIZE)
|
|
|
|
def test_min_p_above_one_raises(self):
|
|
"""Test that verify() rejects min_p > 1.0."""
|
|
sp = self._make(min_p=1.01)
|
|
with self.assertRaises(ValueError):
|
|
sp.verify(self.VOCAB_SIZE)
|
|
|
|
def test_min_p_boundaries_valid(self):
|
|
"""Test that both 0.0 and 1.0 are accepted."""
|
|
self._make(min_p=0.0).verify(self.VOCAB_SIZE)
|
|
self._make(min_p=1.0).verify(self.VOCAB_SIZE)
|
|
|
|
def test_top_k_zero_raises(self):
|
|
"""Test that verify() rejects top_k=0 (must be >=1 or -1 for all)."""
|
|
sp = self._make()
|
|
sp.top_k = 0 # bypass __init__ conversion
|
|
with self.assertRaises(ValueError):
|
|
sp.verify(self.VOCAB_SIZE)
|
|
|
|
def test_top_k_negative_raises(self):
|
|
"""Test that top_k=-2 is rejected (__init__ only converts -1)."""
|
|
sp = self._make()
|
|
sp.top_k = -2 # bypass __init__ conversion
|
|
with self.assertRaises(ValueError):
|
|
sp.verify(self.VOCAB_SIZE)
|
|
|
|
# --- frequency_penalty ---
|
|
def test_frequency_penalty_below_minus_two_raises(self):
|
|
"""Test that verify() rejects frequency_penalty < -2.0."""
|
|
sp = self._make(frequency_penalty=-2.1)
|
|
with self.assertRaises(ValueError):
|
|
sp.verify(self.VOCAB_SIZE)
|
|
|
|
def test_frequency_penalty_above_two_raises(self):
|
|
"""Test that verify() rejects frequency_penalty > 2.0."""
|
|
sp = self._make(frequency_penalty=2.1)
|
|
with self.assertRaises(ValueError):
|
|
sp.verify(self.VOCAB_SIZE)
|
|
|
|
def test_frequency_penalty_boundaries_valid(self):
|
|
"""Test that both -2.0 and 2.0 are accepted."""
|
|
self._make(frequency_penalty=-2.0).verify(self.VOCAB_SIZE)
|
|
self._make(frequency_penalty=2.0).verify(self.VOCAB_SIZE)
|
|
|
|
# --- presence_penalty ---
|
|
def test_presence_penalty_out_of_range_raises(self):
|
|
"""Test that verify() rejects presence_penalty outside [-2, 2]."""
|
|
sp = self._make(presence_penalty=2.5)
|
|
with self.assertRaises(ValueError):
|
|
sp.verify(self.VOCAB_SIZE)
|
|
|
|
# --- repetition_penalty ---
|
|
def test_repetition_penalty_negative_raises(self):
|
|
"""Test that verify() rejects negative repetition_penalty (valid range is (0, 2])."""
|
|
sp = self._make(repetition_penalty=-0.1)
|
|
with self.assertRaises(ValueError):
|
|
sp.verify(self.VOCAB_SIZE)
|
|
|
|
def test_repetition_penalty_zero_raises(self):
|
|
"""Test that verify() rejects repetition_penalty=0.
|
|
|
|
A value of 0 makes the sampling kernel divide logits by 0, producing
|
|
inf/NaN in the probability tensor and crashing every TP rank with a
|
|
device-side assert.
|
|
"""
|
|
sp = self._make(repetition_penalty=0.0)
|
|
with self.assertRaises(ValueError):
|
|
sp.verify(self.VOCAB_SIZE)
|
|
|
|
def test_repetition_penalty_above_two_raises(self):
|
|
"""Test that verify() rejects repetition_penalty > 2.0."""
|
|
sp = self._make(repetition_penalty=2.1)
|
|
with self.assertRaises(ValueError):
|
|
sp.verify(self.VOCAB_SIZE)
|
|
|
|
def test_repetition_penalty_boundary_two_valid(self):
|
|
"""Test that the upper boundary value 2.0 is accepted."""
|
|
self._make(repetition_penalty=2.0).verify(self.VOCAB_SIZE)
|
|
|
|
def test_repetition_penalty_small_positive_valid(self):
|
|
"""Test that a small positive repetition_penalty (e.g. 1e-3) is accepted."""
|
|
self._make(repetition_penalty=1e-3).verify(self.VOCAB_SIZE)
|
|
|
|
# --- min_new_tokens / max_new_tokens ---
|
|
def test_negative_min_new_tokens_raises(self):
|
|
"""Test that verify() rejects negative min_new_tokens."""
|
|
sp = self._make(min_new_tokens=-1)
|
|
with self.assertRaises(ValueError):
|
|
sp.verify(self.VOCAB_SIZE)
|
|
|
|
def test_negative_max_new_tokens_raises(self):
|
|
"""Test that verify() rejects negative max_new_tokens."""
|
|
sp = self._make(max_new_tokens=-1)
|
|
with self.assertRaises(ValueError):
|
|
sp.verify(self.VOCAB_SIZE)
|
|
|
|
def test_min_exceeds_max_new_tokens_raises(self):
|
|
"""Test that verify() rejects min_new_tokens > max_new_tokens."""
|
|
sp = self._make(min_new_tokens=100, max_new_tokens=50)
|
|
with self.assertRaises(ValueError):
|
|
sp.verify(self.VOCAB_SIZE)
|
|
|
|
def test_min_equals_max_new_tokens_valid(self):
|
|
"""Test that min_new_tokens == max_new_tokens is accepted."""
|
|
sp = self._make(min_new_tokens=10, max_new_tokens=10)
|
|
sp.verify(self.VOCAB_SIZE)
|
|
|
|
def test_max_new_tokens_none_skips_validation(self):
|
|
"""Test that max_new_tokens=None skips the min<=max check."""
|
|
sp = self._make(min_new_tokens=9999, max_new_tokens=None)
|
|
sp.verify(self.VOCAB_SIZE) # should not raise
|
|
|
|
# --- logit_bias ---
|
|
def test_logit_bias_token_exceeds_vocab_raises(self):
|
|
"""Test that verify() rejects logit_bias with token_id >= vocab_size."""
|
|
sp = self._make(logit_bias={"99999": 1.0})
|
|
with self.assertRaises(ValueError):
|
|
sp.verify(self.VOCAB_SIZE)
|
|
|
|
def test_logit_bias_negative_token_raises(self):
|
|
"""Test that verify() rejects logit_bias with negative token_id."""
|
|
sp = self._make(logit_bias={"-1": 1.0})
|
|
with self.assertRaises(ValueError):
|
|
sp.verify(self.VOCAB_SIZE)
|
|
|
|
def test_logit_bias_valid_tokens(self):
|
|
"""Test that logit_bias with token_ids within [0, vocab_size) is accepted."""
|
|
sp = self._make(logit_bias={"0": 1.0, "31999": -0.5})
|
|
sp.verify(self.VOCAB_SIZE)
|
|
|
|
def test_multiple_grammars_raises(self):
|
|
"""Reject structural_tag combined with any other grammar constraint.
|
|
|
|
Constraint selection is a fixed if/elif chain, so a constraint left out of
|
|
this check is silently dropped with no error to the caller.
|
|
"""
|
|
for other in ("json_schema", "regex", "ebnf"):
|
|
with self.subTest(other=other):
|
|
sp = self._make(
|
|
structural_tag=self.GRAMMAR_VALUES["structural_tag"],
|
|
**{other: self.GRAMMAR_VALUES[other]},
|
|
)
|
|
with self.assertRaisesRegex(ValueError, "Only one of"):
|
|
sp.verify(self.VOCAB_SIZE)
|
|
|
|
def test_single_grammar_valid(self):
|
|
"""Test that each grammar constraint is valid on its own."""
|
|
for grammar, value in self.GRAMMAR_VALUES.items():
|
|
with self.subTest(grammar=grammar):
|
|
self._make(**{grammar: value}).verify(self.VOCAB_SIZE)
|
|
|
|
|
|
class TestSamplingParamsNormalize(CustomTestCase):
|
|
def _mock_tokenizer(self, encode_map=None):
|
|
"""Create a mock tokenizer that returns predetermined token lists."""
|
|
tokenizer = MagicMock()
|
|
if encode_map:
|
|
tokenizer.encode.side_effect = lambda s, add_special_tokens=False: (
|
|
encode_map.get(s, [1])
|
|
)
|
|
else:
|
|
tokenizer.encode.return_value = [1] # Default: 1 token
|
|
return tokenizer
|
|
|
|
def test_none_stop_strs_becomes_empty_list(self):
|
|
"""Test that normalize() converts None stop to empty list with max_len=0."""
|
|
sp = SamplingParams(stop=None)
|
|
sp.normalize(tokenizer=None)
|
|
self.assertEqual(sp.stop_strs, [])
|
|
self.assertEqual(sp.stop_str_max_len, 0)
|
|
|
|
def test_string_stop_str_wrapped_in_list(self):
|
|
"""Test that normalize() wraps a single stop string into a list."""
|
|
sp = SamplingParams(stop="<|end|>")
|
|
tokenizer = self._mock_tokenizer()
|
|
sp.normalize(tokenizer=tokenizer)
|
|
self.assertEqual(sp.stop_strs, ["<|end|>"])
|
|
|
|
def test_list_stop_strs_unchanged(self):
|
|
"""Test that normalize() preserves a list of stop strings as-is."""
|
|
sp = SamplingParams(stop=["stop1", "stop2"])
|
|
tokenizer = self._mock_tokenizer()
|
|
sp.normalize(tokenizer=tokenizer)
|
|
self.assertEqual(sp.stop_strs, ["stop1", "stop2"])
|
|
|
|
def test_stop_count_limit(self):
|
|
tokenizer = self._mock_tokenizer()
|
|
SamplingParams(stop=["x"] * MAX_STOP_COUNT).normalize(tokenizer)
|
|
|
|
with self.assertRaises(ValueError) as cm:
|
|
SamplingParams(stop=["x"] * (MAX_STOP_COUNT + 1)).normalize(tokenizer)
|
|
self.assertEqual(
|
|
str(cm.exception),
|
|
f"at most {MAX_STOP_COUNT} stop strings are allowed, got {MAX_STOP_COUNT + 1}",
|
|
)
|
|
|
|
def test_stop_str_max_len_uses_encoded_length(self):
|
|
"""Test that max_len is based on encoded token count, not character count."""
|
|
# "ab" encodes to 1 token, "cdef" encodes to 2 tokens
|
|
tokenizer = self._mock_tokenizer(encode_map={"ab": [1], "cdef": [2, 3]})
|
|
sp = SamplingParams(stop=["ab", "cdef"])
|
|
sp.normalize(tokenizer=tokenizer)
|
|
self.assertEqual(sp.stop_str_max_len, 2) # max token count
|
|
|
|
def test_stop_str_max_len_with_tokenizer(self):
|
|
"""Test that with a tokenizer, max_len counts encoded token IDs."""
|
|
tokenizer = MagicMock()
|
|
# "hello" encodes to 2 tokens, "world!!" to 3 tokens
|
|
tokenizer.encode.side_effect = lambda s, add_special_tokens=False: {
|
|
"hello": [101, 102],
|
|
"world!!": [201, 202, 203],
|
|
}[s]
|
|
sp = SamplingParams(stop=["hello", "world!!"])
|
|
sp.normalize(tokenizer=tokenizer)
|
|
self.assertEqual(sp.stop_str_max_len, 3)
|
|
|
|
def test_none_stop_regex_becomes_empty_list(self):
|
|
"""Test that normalize() converts None stop_regex to empty list with max_len=0."""
|
|
sp = SamplingParams(stop_regex=None)
|
|
sp.normalize(tokenizer=None)
|
|
self.assertEqual(sp.stop_regex_strs, [])
|
|
self.assertEqual(sp.stop_regex_max_len, 0)
|
|
|
|
def test_string_stop_regex_wrapped_in_list(self):
|
|
"""Test that normalize() wraps a single stop_regex string into a list."""
|
|
sp = SamplingParams(stop_regex=r"\d+")
|
|
tokenizer = self._mock_tokenizer()
|
|
sp.normalize(tokenizer=tokenizer)
|
|
self.assertEqual(sp.stop_regex_strs, [r"\d+"])
|
|
|
|
def test_stop_regex_max_len_computed(self):
|
|
"""Test that bounded regex computes a finite max length."""
|
|
sp = SamplingParams(stop_regex=r"[a-z]{3}")
|
|
tokenizer = self._mock_tokenizer()
|
|
sp.normalize(tokenizer=tokenizer)
|
|
self.assertEqual(sp.stop_regex_max_len, 3)
|
|
|
|
def test_stop_regex_count_limit(self):
|
|
tokenizer = self._mock_tokenizer()
|
|
SamplingParams(stop_regex=["x"] * MAX_STOP_REGEX_COUNT).normalize(tokenizer)
|
|
|
|
with self.assertRaises(ValueError) as cm:
|
|
SamplingParams(stop_regex=["x"] * (MAX_STOP_REGEX_COUNT + 1)).normalize(
|
|
tokenizer
|
|
)
|
|
self.assertEqual(
|
|
str(cm.exception),
|
|
f"at most {MAX_STOP_REGEX_COUNT} stop_regex patterns are allowed, "
|
|
f"got {MAX_STOP_REGEX_COUNT + 1}",
|
|
)
|
|
|
|
def test_stop_regex_byte_length_limit(self):
|
|
tokenizer = self._mock_tokenizer()
|
|
pattern = "é" * (MAX_STOP_REGEX_LEN // 2)
|
|
SamplingParams(stop_regex=pattern).normalize(tokenizer)
|
|
|
|
with self.assertRaises(ValueError) as cm:
|
|
SamplingParams(stop_regex=pattern + "a").normalize(tokenizer)
|
|
self.assertEqual(
|
|
str(cm.exception),
|
|
f"stop_regex is {MAX_STOP_REGEX_LEN + 1} bytes, over the "
|
|
f"{MAX_STOP_REGEX_LEN}-byte limit",
|
|
)
|
|
|
|
|
|
class TestSamplingParamsMsgspecStruct(CustomTestCase):
|
|
def test_rust_sampling_schema_stays_in_lockstep(self):
|
|
"""Compare Rust fields with the imported Python wire schema."""
|
|
rust_path = (
|
|
Path(__file__).resolve().parents[4]
|
|
/ "rust/sglang-server/src/message/sampling.rs"
|
|
)
|
|
source = rust_path.read_text()
|
|
start = source.index("pub struct SamplingParams {")
|
|
end = source.index("\n}\n\n/// The `/generate`", start)
|
|
rust_fields = tuple(
|
|
re.findall(
|
|
r"^\s*pub ([a-z][a-z0-9_]*):",
|
|
source[start:end],
|
|
re.MULTILINE,
|
|
)
|
|
)
|
|
|
|
self.assertEqual(SamplingParams.__struct_fields__, rust_fields)
|
|
|
|
def test_copy_remains_mutable_and_independent(self):
|
|
sp = SamplingParams(max_new_tokens=8, custom_params={"a": 1})
|
|
|
|
copied = copy.copy(sp)
|
|
copied.max_new_tokens = 16
|
|
copied.custom_params = {"b": 2}
|
|
|
|
self.assertEqual(sp.max_new_tokens, 8)
|
|
self.assertEqual(sp.custom_params, {"a": 1})
|
|
self.assertEqual(copied.max_new_tokens, 16)
|
|
self.assertEqual(copied.custom_params, {"b": 2})
|
|
|
|
def test_none_values_still_use_constructor_defaults(self):
|
|
sp = SamplingParams(
|
|
temperature=None,
|
|
top_p=None,
|
|
top_k=None,
|
|
min_p=None,
|
|
frequency_penalty=None,
|
|
presence_penalty=None,
|
|
repetition_penalty=None,
|
|
min_new_tokens=None,
|
|
n=None,
|
|
ignore_eos=None,
|
|
skip_special_tokens=None,
|
|
spaces_between_special_tokens=None,
|
|
no_stop_trim=None,
|
|
)
|
|
|
|
self.assertEqual(sp.temperature, 1.0)
|
|
self.assertEqual(sp.top_p, 1.0)
|
|
self.assertEqual(sp.top_k, TOP_K_ALL)
|
|
self.assertEqual(sp.min_p, 0.0)
|
|
self.assertEqual(sp.frequency_penalty, 0.0)
|
|
self.assertEqual(sp.presence_penalty, 0.0)
|
|
self.assertEqual(sp.repetition_penalty, 1.0)
|
|
self.assertEqual(sp.min_new_tokens, 0)
|
|
self.assertEqual(sp.n, 1)
|
|
self.assertFalse(sp.ignore_eos)
|
|
self.assertTrue(sp.skip_special_tokens)
|
|
self.assertTrue(sp.spaces_between_special_tokens)
|
|
self.assertFalse(sp.no_stop_trim)
|
|
|
|
def test_msgpack_round_trip_preserves_normalized_state(self):
|
|
tokenizer = MagicMock()
|
|
tokenizer.encode.side_effect = lambda s, add_special_tokens=False: {
|
|
"hello": [101, 102],
|
|
"world": [201],
|
|
}[s]
|
|
sp = SamplingParams(
|
|
stop=["hello", "world"],
|
|
stop_regex=r"[a-z]{3}",
|
|
stop_token_ids=[1, 2],
|
|
temperature=0.5,
|
|
)
|
|
sp.normalize(tokenizer)
|
|
|
|
encoder = msgspec.msgpack.Encoder()
|
|
decoder = msgspec.msgpack.Decoder(SamplingParams)
|
|
rebuilt = decoder.decode(encoder.encode(sp))
|
|
|
|
self.assertIsInstance(rebuilt, SamplingParams)
|
|
self.assertTrue(rebuilt.is_normalized)
|
|
self.assertEqual(rebuilt.stop_strs, ["hello", "world"])
|
|
self.assertEqual(rebuilt.stop_str_max_len, 2)
|
|
self.assertEqual(rebuilt.stop_regex_strs, [r"[a-z]{3}"])
|
|
self.assertEqual(rebuilt.stop_regex_max_len, 3)
|
|
self.assertEqual(rebuilt.stop_token_ids, {1, 2})
|
|
self.assertEqual(rebuilt.temperature, 0.5)
|
|
|
|
|
|
class TestRegexMaxLength(CustomTestCase):
|
|
def test_literal_string(self):
|
|
"""Test that plain string 'abc' gives max length 3."""
|
|
self.assertEqual(get_max_seq_length("abc"), 3)
|
|
|
|
def test_character_class(self):
|
|
"""Test that character class '[a-z]' gives max length 1."""
|
|
self.assertEqual(get_max_seq_length("[a-z]"), 1)
|
|
|
|
def test_dot_any(self):
|
|
"""Test that dot wildcard '.' gives max length 1."""
|
|
self.assertEqual(get_max_seq_length("."), 1)
|
|
|
|
def test_unbounded_star(self):
|
|
"""Test that 'a*' (zero or more, no upper bound) returns MAX_LEN."""
|
|
result = get_max_seq_length("a*")
|
|
self.assertEqual(result, MAX_LEN)
|
|
|
|
def test_unbounded_plus(self):
|
|
"""Test that 'a+' (one or more, no upper bound) returns MAX_LEN."""
|
|
result = get_max_seq_length("a+")
|
|
self.assertEqual(result, MAX_LEN)
|
|
|
|
def test_bounded_repeat(self):
|
|
"""Test that exact repeat 'a{5}' gives max length 5."""
|
|
self.assertEqual(get_max_seq_length("a{5}"), 5)
|
|
|
|
def test_bounded_range_repeat(self):
|
|
"""Test that range repeat 'a{2,4}' uses upper bound, giving max length 4."""
|
|
self.assertEqual(get_max_seq_length("a{2,4}"), 4)
|
|
|
|
def test_branch_takes_max(self):
|
|
"""Test that alternation 'abc|de' takes the longer branch: max(3, 2) = 3."""
|
|
self.assertEqual(get_max_seq_length("abc|de"), 3)
|
|
|
|
def test_subpattern_group(self):
|
|
"""Test that capturing group '(abc)' gives max length 3 from inner content."""
|
|
self.assertEqual(get_max_seq_length("(abc)"), 3)
|
|
|
|
def test_zero_width_assertions_ignored(self):
|
|
"""Test that anchors ^ and $ in '^abc$' add 0, giving max length 3."""
|
|
self.assertEqual(get_max_seq_length("^abc$"), 3)
|
|
|
|
def test_complex_pattern(self):
|
|
"""Test combined pattern '(foo|bar)\\d{2}': branch(3) + repeat(2) = 5."""
|
|
self.assertEqual(get_max_seq_length(r"(foo|bar)\d{2}"), 5)
|
|
|
|
def test_nested_groups(self):
|
|
"""Test that nested groups '((ab))' correctly recurse to give max length 2."""
|
|
self.assertEqual(get_max_seq_length("((ab))"), 2)
|
|
|
|
def test_question_mark_optional(self):
|
|
"""Test that optional 'a?' (equivalent to a{0,1}) gives max length 1."""
|
|
self.assertEqual(get_max_seq_length("a?"), 1)
|
|
|
|
def test_mixed_unbounded_and_bounded(self):
|
|
"""Test that 'ab+c{3}' gives >= MAX_LEN because b+ is unbounded."""
|
|
result = get_max_seq_length("ab+c{3}")
|
|
self.assertGreaterEqual(result, MAX_LEN)
|
|
|
|
def test_empty_regex(self):
|
|
"""Test that empty regex gives max length 0 (no tokens to match)."""
|
|
self.assertEqual(get_max_seq_length(""), 0)
|
|
|
|
def test_lookahead_triggers_unhandled_token(self):
|
|
"""Test that lookahead (?=a) hits the unhandled-token fallback (MAX_LEN)."""
|
|
result = get_max_seq_length("(?=a)b")
|
|
self.assertGreaterEqual(result, MAX_LEN)
|
|
|
|
def test_lookbehind_triggers_unhandled_token(self):
|
|
"""Test that lookbehind (?<=x) hits the unhandled-token fallback (MAX_LEN)."""
|
|
result = get_max_seq_length("(?<=x)y")
|
|
self.assertGreaterEqual(result, MAX_LEN)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|