diff --git a/test/registered/unit/sampling/test_custom_logit_processor.py b/test/registered/unit/sampling/test_custom_logit_processor.py new file mode 100644 index 000000000..bd2df14f3 --- /dev/null +++ b/test/registered/unit/sampling/test_custom_logit_processor.py @@ -0,0 +1,445 @@ +"""Unit tests for srt/sampling/custom_logit_processor.py — no server, no model loading.""" + +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=5, suite="stage-a-cpu-only") + +import json +import unittest +from unittest.mock import MagicMock + +import torch + +from sglang.srt.sampling.custom_logit_processor import ( + CustomLogitProcessor, + DeepseekOCRNoRepeatNGramLogitProcessor, + DeepSeekR1ThinkingBudgetLogitProcessor, + DisallowedTokensLogitsProcessor, + Qwen3ThinkingBudgetLogitProcessor, + _cache_from_str, +) +from sglang.test.test_utils import CustomTestCase + + +# Helper: mock a Req object (used by ThinkingBudget and NGram processors) +def _make_req(origin_input_ids=None, output_ids=None): + req = MagicMock() + req.origin_input_ids = origin_input_ids or [] + req.output_ids = output_ids or [] + return req + + +# Serialization round-trip +class TestCustomLogitProcessorSerialization(CustomTestCase): + + def test_to_str_produces_valid_json(self): + """Test that to_str() produces valid JSON with a 'callable' field.""" + s = DisallowedTokensLogitsProcessor.to_str() + data = json.loads(s) + self.assertIn("callable", data) + self.assertIsInstance(data["callable"], str) + + def test_round_trip_serialization(self): + """Test serialize then deserialize produces a usable processor.""" + s = DisallowedTokensLogitsProcessor.to_str() + processor = CustomLogitProcessor.from_str(s) + self.assertIsInstance(processor, DisallowedTokensLogitsProcessor) + + def test_from_str_is_cached(self): + """Test that from_str uses LRU cache for repeated calls.""" + _cache_from_str.cache_clear() + s = DisallowedTokensLogitsProcessor.to_str() + cls1 = _cache_from_str(s) + cls2 = _cache_from_str(s) + self.assertIs(cls1, cls2) + + +# DisallowedTokensLogitsProcessor +class TestDisallowedTokensLogitsProcessor(CustomTestCase): + def setUp(self): + self.processor = DisallowedTokensLogitsProcessor() + + def test_disallowed_tokens_set_to_neg_inf(self): + """Test that disallowed token positions are set to -inf for all batch items.""" + logits = torch.zeros(2, 10) + params = [{"token_ids": [2, 5]}, {"token_ids": [2, 5]}] + result = self.processor(logits, params) + self.assertTrue(torch.isinf(result[0, 2]) and result[0, 2] < 0) + self.assertTrue(torch.isinf(result[0, 5]) and result[0, 5] < 0) + self.assertTrue(torch.isinf(result[1, 2]) and result[1, 2] < 0) + + def test_allowed_tokens_unchanged(self): + """Test that non-disallowed tokens keep their original logit values.""" + logits = torch.ones(1, 10) + params = [{"token_ids": [3]}] + result = self.processor(logits, params) + self.assertEqual(result[0, 0].item(), 1.0) + self.assertEqual(result[0, 4].item(), 1.0) + self.assertTrue(torch.isinf(result[0, 3]) and result[0, 3] < 0) + + def test_mismatched_params_raises(self): + """Test that mismatched token_ids across batch items raises AssertionError.""" + logits = torch.zeros(2, 10) + params = [{"token_ids": [1, 2]}, {"token_ids": [3, 4]}] + with self.assertRaises(AssertionError): + self.processor(logits, params) + + +# ThinkingBudgetLogitProcessor (using Qwen3 variant) +class TestThinkingBudgetLogitProcessor(CustomTestCase): + """Test thinking budget enforcement using Qwen3 token IDs. + + Qwen3 tokens: + THINKING_START = 151667 + THINKING_END = 151668 + NEW_LINE = 198 + """ + + START = Qwen3ThinkingBudgetLogitProcessor.THINKING_START_TOKEN_ID + END = Qwen3ThinkingBudgetLogitProcessor.THINKING_END_TOKEN_ID + NL = Qwen3ThinkingBudgetLogitProcessor.NEW_LINE_TOKEN_ID + VOCAB = 200000 + + def setUp(self): + self.processor = Qwen3ThinkingBudgetLogitProcessor() + + def _logits(self, batch_size=1): + return torch.zeros(batch_size, self.VOCAB) + + def test_budget_not_exceeded_no_change(self): + """Test no modification when thinking tokens are within budget.""" + req = _make_req( + origin_input_ids=[self.START], + output_ids=[100, 101], # 2 tokens after start + ) + params = [{"thinking_budget": 10, "__req__": req}] + logits = self._logits() + result = self.processor(logits, params) + self.assertEqual(result[0, 0].item(), 0.0) # unchanged + + def test_budget_exceeded_forces_newline_first(self): + """Test forcing newline when budget exceeded and last token is not newline.""" + req = _make_req( + origin_input_ids=[self.START], + output_ids=[100] * 5, # 5 tokens, budget=5 → exceeded + ) + params = [{"thinking_budget": 5, "__req__": req}] + logits = self._logits() + result = self.processor(logits, params) + # newline should be the only non-neg-inf token + self.assertEqual(result[0, self.NL].item(), 0.0) + self.assertTrue(torch.isinf(result[0, 0]) and result[0, 0] < 0) + + def test_budget_exceeded_with_newline_forces_end_token(self): + """Test forcing end token when budget exceeded and last token is newline.""" + req = _make_req( + origin_input_ids=[self.START], + output_ids=[100] * 5 + [self.NL], # 6 tokens, last is newline + ) + params = [{"thinking_budget": 5, "__req__": req}] + logits = self._logits() + result = self.processor(logits, params) + self.assertEqual(result[0, self.END].item(), 0.0) + self.assertTrue(torch.isinf(result[0, 0]) and result[0, 0] < 0) + + def test_skips_when_not_in_thinking(self): + """Test skip when THINKING_START is absent (no thinking phase).""" + req = _make_req(origin_input_ids=[100, 101], output_ids=[102]) + params = [{"thinking_budget": 0, "__req__": req}] + logits = self._logits() + original = logits.clone() + result = self.processor(logits, params) + self.assertTrue(torch.equal(result, original)) + + def test_skips_when_thinking_already_ended(self): + """Test skip when THINKING_END already appeared.""" + req = _make_req( + origin_input_ids=[self.START], + output_ids=[100, self.END, 200], + ) + params = [{"thinking_budget": 0, "__req__": req}] + logits = self._logits() + original = logits.clone() + result = self.processor(logits, params) + self.assertTrue(torch.equal(result, original)) + + def test_skips_when_budget_is_none(self): + """Test that thinking_budget=None is ignored even during thinking phase.""" + req = _make_req(origin_input_ids=[self.START], output_ids=[100] * 10) + params = [{"thinking_budget": None, "__req__": req}] + logits = self._logits() + original = logits.clone() + result = self.processor(logits, params) + self.assertTrue(torch.equal(result, original)) + + def test_skips_when_budget_is_negative(self): + """Test that negative thinking_budget is treated as disabled (no enforcement).""" + req = _make_req(origin_input_ids=[self.START], output_ids=[100] * 10) + params = [{"thinking_budget": -1, "__req__": req}] + logits = self._logits() + original = logits.clone() + result = self.processor(logits, params) + self.assertTrue(torch.equal(result, original)) + + def test_none_params_returns_unchanged(self): + """Test that passing None as param list returns logits unchanged.""" + logits = self._logits() + original = logits.clone() + result = self.processor(logits, None) + self.assertTrue(torch.equal(result, original)) + + def test_empty_params_returns_unchanged(self): + """Test that passing empty param list returns logits unchanged.""" + logits = self._logits() + original = logits.clone() + result = self.processor(logits, []) + self.assertTrue(torch.equal(result, original)) + + def test_budget_zero_forces_immediate_end(self): + """Test that budget=0 forces thinking to end immediately.""" + req = _make_req( + origin_input_ids=[self.START], + output_ids=[100], # 1 token after start > budget=0 + ) + params = [{"thinking_budget": 0, "__req__": req}] + logits = self._logits() + result = self.processor(logits, params) + # Should force newline since last token (100) is not newline + self.assertEqual(result[0, self.NL].item(), 0.0) + + def test_none_param_dict_in_list_skipped(self): + """Test that None entry in param list is skipped gracefully.""" + req = _make_req( + origin_input_ids=[self.START], + output_ids=[100] * 10, + ) + params = [None, {"thinking_budget": 0, "__req__": req}] + logits = self._logits(batch_size=2) + result = self.processor(logits, params) + # Batch 0 (None param) should be unchanged + self.assertEqual(result[0, 0].item(), 0.0) + # Batch 1 should have been modified (budget exceeded) + self.assertEqual(result[1, self.NL].item(), 0.0) + self.assertTrue(torch.isinf(result[1, 0]) and result[1, 0] < 0) + + def test_multiple_thinking_start_counts_from_first(self): + """Test that budget counts from the first THINKING_START occurrence.""" + req = _make_req( + origin_input_ids=[self.START, 100, 101], + output_ids=[self.START, 200, 201], # second START in output + ) + # cur_ids = [START, 100, 101, START, 200, 201] + # First START at index 0, tokens_after_start = 5 + # Budget=10 → 5 < 10 → no modification + params = [{"thinking_budget": 10, "__req__": req}] + logits = self._logits() + original = logits.clone() + result = self.processor(logits, params) + self.assertTrue(torch.equal(result, original)) + + def test_deepseek_r1_variant_forces_end(self): + """Test DeepSeekR1 variant with its own token IDs.""" + proc = DeepSeekR1ThinkingBudgetLogitProcessor() + START = proc.THINKING_START_TOKEN_ID # 128798 + NL = proc.NEW_LINE_TOKEN_ID # 201 + VOCAB = 200000 + + req = _make_req(origin_input_ids=[START], output_ids=[100] * 5) + params = [{"thinking_budget": 5, "__req__": req}] + logits = torch.zeros(1, VOCAB) + result = proc(logits, params) + # Budget exceeded, last token (100) is not newline → force newline + self.assertEqual(result[0, NL].item(), 0.0) + self.assertTrue(torch.isinf(result[0, 0]) and result[0, 0] < 0) + + +# DeepseekOCRNoRepeatNGramLogitProcessor +class TestDeepseekOCRNoRepeatNGramLogitProcessor(CustomTestCase): + VOCAB = 100 + + def setUp(self): + self.processor = DeepseekOCRNoRepeatNGramLogitProcessor() + + def _logits(self, batch_size=1): + return torch.zeros(batch_size, self.VOCAB) + + def test_bans_repeated_bigrams(self): + """Test banning token that completes a repeated bigram.""" + req = _make_req(origin_input_ids=[1, 2, 3, 1, 2]) + params = [ + { + "__req__": req, + "ngram_size": 2, + "window_size": 100, + } + ] + logits = self._logits() + result = self.processor(logits, params) + self.assertTrue(torch.isinf(result[0, 3]) and result[0, 3] < 0) + + def test_non_repeated_tokens_unchanged(self): + """Test that tokens not completing a repeated ngram are unchanged.""" + req = _make_req(origin_input_ids=[1, 2, 3, 1, 2]) + params = [{"__req__": req, "ngram_size": 2, "window_size": 100}] + logits = self._logits() + result = self.processor(logits, params) + # Token 1 is not banned (prefix (2) was followed by 3, not 1) + self.assertEqual(result[0, 1].item(), 0.0) + + def test_window_size_limits_search(self): + """Test that ngrams outside the window are not considered.""" + # Sequence: [1,2,3,...,1,2] but window only covers the last 3 tokens + req = _make_req(origin_input_ids=[1, 2, 3, 4, 5, 1, 2]) + params = [{"__req__": req, "ngram_size": 2, "window_size": 3}] + logits = self._logits() + result = self.processor(logits, params) + # Window covers [5, 1, 2]. The bigram (1,2) from index 0-1 is outside. + # Within window: bigrams are (5,1), (1,2). Current prefix is (2). + # No bigram starting with prefix (2) in window → nothing banned. + self.assertEqual(result[0, 3].item(), 0.0) + + def test_whitelist_protects_tokens(self): + """Test that whitelisted tokens are not banned despite repeated ngrams.""" + req = _make_req(origin_input_ids=[1, 2, 3, 1, 2]) + params = [ + { + "__req__": req, + "ngram_size": 2, + "window_size": 100, + "whitelist_token_ids": [3], + } + ] + logits = self._logits() + result = self.processor(logits, params) + # Token 3 would be banned but is whitelisted + self.assertEqual(result[0, 3].item(), 0.0) + + def test_ngram_size_zero_skips(self): + """ngram_size=0 is invalid and should be skipped (no modification).""" + req = _make_req(origin_input_ids=[1, 2, 1, 2]) + params = [{"__req__": req, "ngram_size": 0, "window_size": 100}] + logits = self._logits() + original = logits.clone() + result = self.processor(logits, params) + self.assertTrue(torch.equal(result, original)) + + def test_window_size_zero_skips(self): + """Test that window_size=0 disables ngram checking (no modification).""" + req = _make_req(origin_input_ids=[1, 2, 1, 2]) + params = [{"__req__": req, "ngram_size": 2, "window_size": 0}] + logits = self._logits() + original = logits.clone() + result = self.processor(logits, params) + self.assertTrue(torch.equal(result, original)) + + def test_empty_params_returns_unchanged(self): + """Test that None param list returns logits unchanged (early return).""" + logits = self._logits() + original = logits.clone() + result = self.processor(logits, None) + self.assertTrue(torch.equal(result, original)) + + def test_short_sequence_skips(self): + """Sequence shorter than ngram_size should be skipped.""" + req = _make_req(origin_input_ids=[1]) + params = [{"__req__": req, "ngram_size": 3, "window_size": 100}] + logits = self._logits() + original = logits.clone() + result = self.processor(logits, params) + self.assertTrue(torch.equal(result, original)) + + def test_unigram_mode(self): + """ngram_size=1 bans any token already seen in the window.""" + req = _make_req(origin_input_ids=[5, 10, 15]) + params = [{"__req__": req, "ngram_size": 1, "window_size": 100}] + logits = self._logits() + result = self.processor(logits, params) + # All tokens in [5, 10, 15] should be banned + self.assertTrue(torch.isinf(result[0, 5]) and result[0, 5] < 0) + self.assertTrue(torch.isinf(result[0, 10]) and result[0, 10] < 0) + self.assertTrue(torch.isinf(result[0, 15]) and result[0, 15] < 0) + # Other tokens should be fine + self.assertEqual(result[0, 0].item(), 0.0) + + def test_none_req_skips(self): + """If __req__ is missing, the batch item should be skipped.""" + params = [{"ngram_size": 2, "window_size": 100}] + logits = self._logits() + original = logits.clone() + result = self.processor(logits, params) + self.assertTrue(torch.equal(result, original)) + + def test_invalid_ngram_size_type_skips(self): + """Non-numeric ngram_size should be handled gracefully.""" + req = _make_req(origin_input_ids=[1, 2, 1, 2]) + params = [{"__req__": req, "ngram_size": "invalid", "window_size": 100}] + logits = self._logits() + original = logits.clone() + result = self.processor(logits, params) + self.assertTrue(torch.equal(result, original)) + + def test_falsy_params_in_list_skipped(self): + """A falsy entry (None, {}, 0) in param list should be skipped.""" + req = _make_req(origin_input_ids=[1, 2, 1, 2]) + params = [None, {"__req__": req, "ngram_size": 2, "window_size": 100}] + logits = self._logits(batch_size=2) + result = self.processor(logits, params) + # Batch 0 (None) unchanged + self.assertEqual(result[0, 0].item(), 0.0) + # Batch 1 has ban applied + self.assertTrue(torch.isinf(result[1, 1]) and result[1, 1] < 0) + + def test_search_end_leq_search_start_skips(self): + """Test skip when window is too small for the ngram_size.""" + # sequence length=4, ngram_size=3, window_size=2 + # search_start = max(0, 4-2) = 2 + # search_end = 4 - 3 + 1 = 2 + # search_end (2) <= search_start (2) → skip + req = _make_req(origin_input_ids=[1, 2, 3, 4]) + params = [{"__req__": req, "ngram_size": 3, "window_size": 2}] + logits = self._logits() + original = logits.clone() + result = self.processor(logits, params) + self.assertTrue(torch.equal(result, original)) + + def test_invalid_whitelist_type_handled(self): + """Test graceful handling of non-iterable whitelist_token_ids.""" + req = _make_req(origin_input_ids=[1, 2, 1, 2]) + params = [ + { + "__req__": req, + "ngram_size": 2, + "window_size": 100, + "whitelist_token_ids": 999, # int, not iterable + } + ] + logits = self._logits() + result = self.processor(logits, params) + # Should still ban token 1 (whitelist parse fails, falls back to empty set) + self.assertTrue(torch.isinf(result[0, 1]) and result[0, 1] < 0) + + def test_batch_processing(self): + """Test that multiple batch items are processed independently.""" + req1 = _make_req( + origin_input_ids=[1, 2, 1, 2] + ) # will ban token 2 (bigram repeat) + req2 = _make_req(origin_input_ids=[3, 4, 5]) # no repeat + params = [ + {"__req__": req1, "ngram_size": 2, "window_size": 100}, + {"__req__": req2, "ngram_size": 2, "window_size": 100}, + ] + logits = self._logits(batch_size=2) + result = self.processor(logits, params) + # Batch 0: bigram (1,2) appeared, prefix is (2) → ban token that followed (2) = 1 + # Also (2,1) appeared, prefix is (2) → already covered + # Actually: sequence is [1,2,1,2], prefix is last (ngram_size-1)=1 token = (2) + # Scanning: index 0: (1,2) prefix=(1); index 1: (2,1) prefix=(2)→bans 1; index 2: (1,2) prefix=(1) + # So prefix (2) appeared at index 1, followed by token 1. Ban token 1. + self.assertTrue(torch.isinf(result[0, 1]) and result[0, 1] < 0) + # Batch 1: prefix is (5), no matching prefix in window → no bans + self.assertEqual(result[1, 3].item(), 0.0) + self.assertEqual(result[1, 4].item(), 0.0) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/sampling/test_penaltylib.py b/test/registered/unit/sampling/test_penaltylib.py new file mode 100644 index 000000000..3b62524ba --- /dev/null +++ b/test/registered/unit/sampling/test_penaltylib.py @@ -0,0 +1,520 @@ +"""Unit tests for srt/sampling/penaltylib/ — no server, no model loading.""" + +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=5, suite="stage-a-cpu-only") + +import unittest +from unittest.mock import MagicMock + +import torch + +from sglang.srt.sampling.penaltylib.frequency_penalty import ( + BatchedFrequencyPenalizer, +) +from sglang.srt.sampling.penaltylib.min_new_tokens import ( + BatchedMinNewTokensPenalizer, +) +from sglang.srt.sampling.penaltylib.orchestrator import ( + BatchedPenalizerOrchestrator, +) +from sglang.srt.sampling.penaltylib.presence_penalty import ( + BatchedPresencePenalizer, +) +from sglang.test.test_utils import CustomTestCase + +VOCAB_SIZE = 32 +DEVICE = "cpu" + + +# Helpers: mock Req and ScheduleBatch +def _make_req(freq=0.0, presence=0.0, min_tokens=0, stop_ids=None, eos_id=2): + """Create a mock request with sampling params.""" + req = MagicMock() + req.sampling_params.frequency_penalty = freq + req.sampling_params.presence_penalty = presence + req.sampling_params.min_new_tokens = min_tokens + req.sampling_params.stop_token_ids = stop_ids + req.tokenizer.additional_stop_token_ids = None + req.tokenizer.eos_token_id = eos_id + return req + + +def _make_batch(reqs): + """Create a mock ScheduleBatch. + Note: orchestrator accesses batch.reqs as an attribute (not a method call).""" + batch = MagicMock() + batch.reqs = reqs + batch.device = DEVICE + return batch + + +# BatchedPenalizerOrchestrator +class TestBatchedPenalizerOrchestrator(CustomTestCase): + + def test_init_detects_required_penalizers(self): + """Test that orchestrator marks is_required=True when any request has nonzero penalty.""" + reqs = [_make_req(freq=1.0)] + batch = _make_batch(reqs) + orch = BatchedPenalizerOrchestrator( + VOCAB_SIZE, batch, {BatchedFrequencyPenalizer} + ) + self.assertTrue(orch.is_required) + + def test_init_not_required_when_no_penalties(self): + """Test that orchestrator marks is_required=False when all penalties are zero.""" + reqs = [_make_req()] # all defaults (0.0) + batch = _make_batch(reqs) + orch = BatchedPenalizerOrchestrator( + VOCAB_SIZE, batch, {BatchedFrequencyPenalizer} + ) + self.assertFalse(orch.is_required) + + def test_batch_property_via_weakref(self): + """Test that batch property returns the original batch via weakref.""" + reqs = [_make_req()] + batch = _make_batch(reqs) + orch = BatchedPenalizerOrchestrator(VOCAB_SIZE, batch, set()) + self.assertIs(orch.batch, batch) + + def test_batch_setter_none(self): + """Test that setting batch to None breaks the weakref cleanly.""" + reqs = [_make_req()] + batch = _make_batch(reqs) + orch = BatchedPenalizerOrchestrator(VOCAB_SIZE, batch, set()) + orch.batch = None + self.assertIsNone(orch.batch) + + def test_batch_setter_new_batch(self): + """Test that batch can be reassigned to a different ScheduleBatch.""" + reqs = [_make_req()] + batch1 = _make_batch(reqs) + batch2 = _make_batch(reqs) + orch = BatchedPenalizerOrchestrator(VOCAB_SIZE, batch1, set()) + orch.batch = batch2 + self.assertIs(orch.batch, batch2) + + def test_context_manager_releases(self): + """Test that exiting the context manager releases all penalizers.""" + reqs = [_make_req(freq=1.0)] + batch = _make_batch(reqs) + with BatchedPenalizerOrchestrator( + VOCAB_SIZE, batch, {BatchedFrequencyPenalizer} + ) as orch: + self.assertTrue(orch.is_required) + self.assertFalse(orch.is_required) + self.assertEqual(len(orch.penalizers), 0) + + def test_filter_empty_indices_releases(self): + """Test that filtering with no indices left fully releases the orchestrator.""" + reqs = [_make_req(freq=1.0)] + batch = _make_batch(reqs) + orch = BatchedPenalizerOrchestrator( + VOCAB_SIZE, batch, {BatchedFrequencyPenalizer} + ) + orch.filter(torch.tensor([], dtype=torch.long)) + self.assertFalse(orch.is_required) + + def test_filter_not_required_is_noop(self): + """Test that filter on a not-required orchestrator does nothing.""" + reqs = [_make_req()] + batch = _make_batch(reqs) + orch = BatchedPenalizerOrchestrator( + VOCAB_SIZE, batch, {BatchedFrequencyPenalizer} + ) + self.assertFalse(orch.is_required) + orch.filter(torch.tensor([0])) # should not raise + + def test_merge_both_not_required_is_noop(self): + """Test that merging two not-required orchestrators stays not-required.""" + reqs = [_make_req()] + batch = _make_batch(reqs) + orch1 = BatchedPenalizerOrchestrator( + VOCAB_SIZE, batch, {BatchedFrequencyPenalizer} + ) + orch2 = BatchedPenalizerOrchestrator( + VOCAB_SIZE, batch, {BatchedFrequencyPenalizer} + ) + orch1.merge(orch2) # should not raise + self.assertFalse(orch1.is_required) + + +# BatchedFrequencyPenalizer +class TestBatchedFrequencyPenalizer(CustomTestCase): + + def _setup(self, freq_values): + reqs = [_make_req(freq=f) for f in freq_values] + batch = _make_batch(reqs) + orch = BatchedPenalizerOrchestrator( + VOCAB_SIZE, batch, {BatchedFrequencyPenalizer} + ) + pen = orch.penalizers[BatchedFrequencyPenalizer] + return orch, pen + + def test_is_required_with_nonzero_penalty(self): + """Test that nonzero frequency_penalty makes the penalizer required.""" + _, pen = self._setup([1.5]) + self.assertTrue(pen.is_required()) + + def test_is_not_required_with_zero_penalty(self): + """Test that zero frequency_penalty makes the penalizer not required.""" + _, pen = self._setup([0.0]) + self.assertFalse(pen.is_required()) + + def test_cumulate_and_apply(self): + """Test that cumulating a token applies frequency penalty to its logit.""" + orch, pen = self._setup([2.0]) + output_ids = torch.tensor([5]) + pen.cumulate_output_tokens(output_ids) + + logits = torch.zeros(1, VOCAB_SIZE) + pen.apply(logits) + self.assertAlmostEqual(logits[0, 5].item(), -2.0, places=5) + # Other tokens unaffected + self.assertAlmostEqual(logits[0, 0].item(), 0.0, places=5) + + def test_cumulate_twice_doubles_penalty(self): + """Test that frequency penalty scales linearly with occurrence count.""" + orch, pen = self._setup([1.0]) + pen.cumulate_output_tokens(torch.tensor([3])) + pen.cumulate_output_tokens(torch.tensor([3])) + + logits = torch.zeros(1, VOCAB_SIZE) + pen.apply(logits) + self.assertAlmostEqual(logits[0, 3].item(), -2.0, places=5) + + def test_filter_keeps_subset(self): + """Test that filter retains only the selected batch indices.""" + orch, pen = self._setup([1.0, 2.0]) + keep = torch.tensor([1]) + pen.filter(keep) + self.assertEqual(pen.frequency_penalties.shape[0], 1) + self.assertAlmostEqual(pen.frequency_penalties[0, 0].item(), 2.0, places=5) + + def test_merge_concatenates(self): + """Test that merge concatenates penalty tensors from two penalizers.""" + _, pen1 = self._setup([1.0]) + _, pen2 = self._setup([2.0]) + pen1.merge(pen2) + self.assertEqual(pen1.frequency_penalties.shape[0], 2) + + def test_teardown_cleans_attributes(self): + """Test that teardown deletes internal tensors and resets prepared state.""" + _, pen = self._setup([1.0]) + pen.teardown() + self.assertFalse(hasattr(pen, "frequency_penalties")) + self.assertFalse(hasattr(pen, "cumulated_frequency_penalties")) + self.assertFalse(pen.is_prepared()) + + def test_cumulate_when_not_prepared_is_noop(self): + """Test that cumulate before prepare does not crash.""" + _, pen = self._setup([0.0]) + # pen is not prepared (is_required=False) + pen.cumulate_output_tokens(torch.tensor([1])) # should not raise + + def test_apply_when_not_prepared_is_noop(self): + """Test that apply on an unprepared penalizer leaves logits unchanged.""" + _, pen = self._setup([0.0]) + logits = torch.zeros(1, VOCAB_SIZE) + original = logits.clone() + pen.apply(logits) + self.assertTrue(torch.equal(logits, original)) + + +# BatchedPresencePenalizer +class TestBatchedPresencePenalizer(CustomTestCase): + + def _setup(self, presence_values): + reqs = [_make_req(presence=p) for p in presence_values] + batch = _make_batch(reqs) + orch = BatchedPenalizerOrchestrator( + VOCAB_SIZE, batch, {BatchedPresencePenalizer} + ) + pen = orch.penalizers[BatchedPresencePenalizer] + return orch, pen + + def test_is_required_with_nonzero_penalty(self): + """Test that nonzero presence_penalty makes the penalizer required.""" + _, pen = self._setup([0.5]) + self.assertTrue(pen.is_required()) + + def test_presence_penalty_does_not_scale(self): + """Test that presence penalty is flat (same value regardless of count).""" + orch, pen = self._setup([1.0]) + pen.cumulate_output_tokens(torch.tensor([7])) + pen.cumulate_output_tokens(torch.tensor([7])) # same token again + + logits = torch.zeros(1, VOCAB_SIZE) + pen.apply(logits) + # scatter_ overwrites (not adds), so penalty should be 1.0, not 2.0 + self.assertAlmostEqual(logits[0, 7].item(), -1.0, places=5) + + def test_filter_keeps_subset(self): + """Test that filter retains the first request's presence penalty.""" + orch, pen = self._setup([1.0, 2.0]) + keep = torch.tensor([0]) + pen.filter(keep) + self.assertEqual(pen.presence_penalties.shape[0], 1) + self.assertAlmostEqual(pen.presence_penalties[0, 0].item(), 1.0, places=5) + + def test_merge_concatenates(self): + """Test that merge concatenates presence penalty tensors.""" + _, pen1 = self._setup([1.0]) + _, pen2 = self._setup([2.0]) + pen1.merge(pen2) + self.assertEqual(pen1.presence_penalties.shape[0], 2) + + def test_teardown_cleans_attributes(self): + """Test that teardown removes the presence_penalties tensor.""" + _, pen = self._setup([1.0]) + pen.teardown() + self.assertFalse(hasattr(pen, "presence_penalties")) + + +# BatchedMinNewTokensPenalizer +class TestBatchedMinNewTokensPenalizer(CustomTestCase): + + def _setup(self, configs): + """configs: list of (min_tokens, stop_ids, eos_id).""" + reqs = [_make_req(min_tokens=c[0], stop_ids=c[1], eos_id=c[2]) for c in configs] + batch = _make_batch(reqs) + orch = BatchedPenalizerOrchestrator( + VOCAB_SIZE, batch, {BatchedMinNewTokensPenalizer} + ) + pen = orch.penalizers[BatchedMinNewTokensPenalizer] + return orch, pen + + def test_is_required_with_positive_min_tokens(self): + """Test that positive min_new_tokens makes the penalizer required.""" + _, pen = self._setup([(5, None, 2)]) + self.assertTrue(pen.is_required()) + + def test_is_not_required_with_zero_min_tokens(self): + """Test that min_new_tokens=0 makes the penalizer not required.""" + _, pen = self._setup([(0, None, 2)]) + self.assertFalse(pen.is_required()) + + def test_blocks_eos_before_min_tokens(self): + """Test that EOS token is blocked before min_new_tokens is reached.""" + orch, pen = self._setup([(3, None, 2)]) + # Before any output: len=0 < min=3 → block EOS (token 2) + logits = torch.zeros(1, VOCAB_SIZE) + pen.apply(logits) + self.assertTrue(torch.isinf(logits[0, 2]) and logits[0, 2] < 0) + # Non-stop tokens should be fine + self.assertEqual(logits[0, 0].item(), 0.0) + + def test_allows_eos_after_min_tokens(self): + """Test that EOS is allowed after generating min_new_tokens.""" + orch, pen = self._setup([(2, None, 2)]) + # Generate 2 tokens + pen.cumulate_output_tokens(torch.tensor([10])) + pen.cumulate_output_tokens(torch.tensor([11])) + # Now len=2 >= min=2 → EOS should NOT be blocked + logits = torch.zeros(1, VOCAB_SIZE) + pen.apply(logits) + self.assertEqual(logits[0, 2].item(), 0.0) + + def test_blocks_custom_stop_tokens(self): + """Test that custom stop_token_ids are also blocked before min_new_tokens.""" + orch, pen = self._setup([(3, {5, 10}, 2)]) + logits = torch.zeros(1, VOCAB_SIZE) + pen.apply(logits) + # EOS (2), stop token 5, stop token 10 should all be blocked + self.assertTrue(torch.isinf(logits[0, 2]) and logits[0, 2] < 0) + self.assertTrue(torch.isinf(logits[0, 5]) and logits[0, 5] < 0) + self.assertTrue(torch.isinf(logits[0, 10]) and logits[0, 10] < 0) + + def test_blocks_additional_stop_tokens(self): + """Test that tokenizer's additional_stop_token_ids are also blocked.""" + req = _make_req(min_tokens=3, stop_ids=None, eos_id=2) + req.tokenizer.additional_stop_token_ids = {7, 8} + batch = _make_batch([req]) + orch = BatchedPenalizerOrchestrator( + VOCAB_SIZE, batch, {BatchedMinNewTokensPenalizer} + ) + pen = orch.penalizers[BatchedMinNewTokensPenalizer] + + logits = torch.zeros(1, VOCAB_SIZE) + pen.apply(logits) + # EOS (2) + additional stops (7, 8) should all be blocked + for tok in [2, 7, 8]: + self.assertTrue( + torch.isinf(logits[0, tok]) and logits[0, tok] < 0, + f"token {tok} should be blocked before min_new_tokens", + ) + # Non-stop tokens should be fine + self.assertEqual(logits[0, 0].item(), 0.0) + + def test_filter_keeps_subset(self): + """Test that filter keeps the second request (min_tokens=5) and drops the first.""" + orch, pen = self._setup([(3, None, 2), (5, None, 2)]) + keep = torch.tensor([1]) + pen.filter(keep) + self.assertEqual(pen.min_new_tokens.shape[0], 1) + self.assertEqual(pen.min_new_tokens[0, 0].item(), 5) + + def test_merge_concatenates(self): + """Test that merge combines min_new_tokens tensors from two penalizers.""" + _, pen1 = self._setup([(3, None, 2)]) + _, pen2 = self._setup([(5, None, 2)]) + pen1.merge(pen2) + self.assertEqual(pen1.min_new_tokens.shape[0], 2) + + def test_teardown_cleans_attributes(self): + """Test that teardown removes min_new_tokens, stop_token_penalties, and len_output_tokens.""" + _, pen = self._setup([(3, None, 2)]) + pen.teardown() + self.assertFalse(hasattr(pen, "min_new_tokens")) + self.assertFalse(hasattr(pen, "stop_token_penalties")) + self.assertFalse(hasattr(pen, "len_output_tokens")) + + +# _BatchedPenalizer base class edge cases +class TestBatchedPenalizerBase(CustomTestCase): + + def test_filter_when_not_prepared_is_noop(self): + """Test that filter on an unprepared penalizer does not crash.""" + reqs = [_make_req()] + batch = _make_batch(reqs) + orch = BatchedPenalizerOrchestrator( + VOCAB_SIZE, batch, {BatchedFrequencyPenalizer} + ) + pen = orch.penalizers[BatchedFrequencyPenalizer] + # pen is not prepared (frequency_penalty=0 → not required) + pen.filter(torch.tensor([0])) # should not raise + + def test_merge_prepares_both_if_needed(self): + """Test that merge prepares unprepared side before concatenating.""" + reqs_a = [_make_req(freq=0.0)] # not required + reqs_b = [_make_req(freq=1.0)] # required + batch_a = _make_batch(reqs_a) + batch_b = _make_batch(reqs_b) + orch_a = BatchedPenalizerOrchestrator( + VOCAB_SIZE, batch_a, {BatchedFrequencyPenalizer} + ) + orch_b = BatchedPenalizerOrchestrator( + VOCAB_SIZE, batch_b, {BatchedFrequencyPenalizer} + ) + pen_a = orch_a.penalizers[BatchedFrequencyPenalizer] + pen_b = orch_b.penalizers[BatchedFrequencyPenalizer] + self.assertFalse(pen_a.is_prepared()) + self.assertTrue(pen_b.is_prepared()) + # Merge should prepare pen_a first + pen_a.merge(pen_b) + self.assertTrue(pen_a.is_prepared()) + self.assertEqual(pen_a.frequency_penalties.shape[0], 2) + + def test_merge_both_unprepared_is_noop(self): + """Test that merging two unprepared penalizers keeps them unprepared.""" + reqs = [_make_req()] + batch = _make_batch(reqs) + orch1 = BatchedPenalizerOrchestrator( + VOCAB_SIZE, batch, {BatchedFrequencyPenalizer} + ) + orch2 = BatchedPenalizerOrchestrator( + VOCAB_SIZE, batch, {BatchedFrequencyPenalizer} + ) + pen1 = orch1.penalizers[BatchedFrequencyPenalizer] + pen2 = orch2.penalizers[BatchedFrequencyPenalizer] + pen1.merge(pen2) # both not prepared → noop + self.assertFalse(pen1.is_prepared()) + + def test_prepare_is_idempotent(self): + """Test that calling prepare() multiple times does not crash.""" + reqs = [_make_req(freq=1.0)] + batch = _make_batch(reqs) + orch = BatchedPenalizerOrchestrator( + VOCAB_SIZE, batch, {BatchedFrequencyPenalizer} + ) + pen = orch.penalizers[BatchedFrequencyPenalizer] + self.assertTrue(pen.is_prepared()) + # Calling prepare again should not crash or reinitialize + pen.prepare() + self.assertTrue(pen.is_prepared()) + + +# Orchestrator with multiple penalizer types +class TestOrchestratorMultiplePenalizers(CustomTestCase): + + def test_all_three_penalizers(self): + """Test orchestrator managing frequency, presence, and min_new_tokens together.""" + reqs = [_make_req(freq=1.0, presence=0.5, min_tokens=2, eos_id=2)] + batch = _make_batch(reqs) + orch = BatchedPenalizerOrchestrator( + VOCAB_SIZE, + batch, + { + BatchedFrequencyPenalizer, + BatchedPresencePenalizer, + BatchedMinNewTokensPenalizer, + }, + ) + self.assertTrue(orch.is_required) + + # Cumulate one token + output_ids = torch.tensor([5]) + orch.cumulate_output_tokens(output_ids) + + # Apply all penalties + logits = torch.zeros(1, VOCAB_SIZE) + orch.apply(logits) + + # Token 5: freq_penalty=1.0 (cumulated once) + pres_penalty=0.5 + self.assertAlmostEqual(logits[0, 5].item(), -1.5, places=4) + # EOS (token 2): blocked by min_new_tokens (len=1 < min=2) + self.assertTrue(torch.isinf(logits[0, 2]) and logits[0, 2] < 0) + + def test_filter_with_penalizer_no_longer_required(self): + """Test that penalizer is torn down when no longer required after filter.""" + reqs = [_make_req(freq=0.0), _make_req(freq=1.0)] + batch = _make_batch(reqs) + orch = BatchedPenalizerOrchestrator( + VOCAB_SIZE, batch, {BatchedFrequencyPenalizer} + ) + self.assertTrue(orch.is_required) + + # Keep only the request with freq=0 (index 0) + batch.reqs = [reqs[0]] + orch.filter(torch.tensor([0])) + + pen = orch.penalizers[BatchedFrequencyPenalizer] + # After filter, only req with freq=0 remains → penalizer not required + self.assertFalse(pen.is_required()) + + def test_filter_keeps_required_penalizer(self): + """Test that filter keeps penalizer active when still required.""" + reqs = [_make_req(freq=1.0), _make_req(freq=2.0)] + batch = _make_batch(reqs) + orch = BatchedPenalizerOrchestrator( + VOCAB_SIZE, batch, {BatchedFrequencyPenalizer} + ) + self.assertTrue(orch.is_required) + + batch.reqs = [reqs[1]] + orch.filter(torch.tensor([1])) + self.assertTrue(orch.is_required) + + def test_merge_one_required(self): + """Test that merge marks orchestrator as required when one side is.""" + reqs_a = [_make_req(freq=0.0)] + reqs_b = [_make_req(freq=1.0)] + batch_a = _make_batch(reqs_a) + batch_b = _make_batch(reqs_b) + orch_a = BatchedPenalizerOrchestrator( + VOCAB_SIZE, batch_a, {BatchedFrequencyPenalizer} + ) + orch_b = BatchedPenalizerOrchestrator( + VOCAB_SIZE, batch_b, {BatchedFrequencyPenalizer} + ) + self.assertFalse(orch_a.is_required) + self.assertTrue(orch_b.is_required) + + orch_a.merge(orch_b) + self.assertTrue(orch_a.is_required) + pen = orch_a.penalizers[BatchedFrequencyPenalizer] + self.assertEqual(pen.frequency_penalties.shape[0], 2) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/sampling/test_sampling_batch_info.py b/test/registered/unit/sampling/test_sampling_batch_info.py new file mode 100644 index 000000000..f6f48089a --- /dev/null +++ b/test/registered/unit/sampling/test_sampling_batch_info.py @@ -0,0 +1,577 @@ +"""Unit tests for srt/sampling/sampling_batch_info.py — no server, no model loading.""" + +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=5, suite="stage-a-cpu-only") + +import unittest +from unittest.mock import MagicMock, patch + +import torch + +from sglang.srt.sampling.sampling_batch_info import ( + SamplingBatchInfo, + merge_bias_tensor, +) +from sglang.srt.sampling.sampling_params import TOP_K_ALL +from sglang.test.test_utils import CustomTestCase + +VOCAB_SIZE = 32 +DEVICE = "cpu" + + +# Helper: construct a minimal SamplingBatchInfo +def _make_info(batch_size=2, **overrides): + """Create a SamplingBatchInfo with sane defaults for testing.""" + defaults = dict( + temperatures=torch.ones(batch_size, 1), + top_ps=torch.ones(batch_size), + top_ks=torch.full((batch_size,), TOP_K_ALL, dtype=torch.int32), + min_ps=torch.zeros(batch_size), + is_all_greedy=False, + need_top_p_sampling=False, + need_top_k_sampling=False, + need_min_p_sampling=False, + vocab_size=VOCAB_SIZE, + device=DEVICE, + penalizer_orchestrator=MagicMock(is_required=False), + ) + defaults.update(overrides) + return SamplingBatchInfo(**defaults) + + +class TestMergeBiasTensor(CustomTestCase): + + def test_both_none_returns_none(self): + """Test that merging two None tensors returns None.""" + result = merge_bias_tensor(None, None, 2, 3, DEVICE, 0.0) + self.assertIsNone(result) + + def test_both_present_concatenates(self): + """Test that two present tensors are concatenated along batch dim.""" + lhs = torch.ones(2, VOCAB_SIZE) + rhs = torch.zeros(3, VOCAB_SIZE) + result = merge_bias_tensor(lhs, rhs, 2, 3, DEVICE, 0.0) + self.assertEqual(result.shape, (5, VOCAB_SIZE)) + self.assertEqual(result[0, 0].item(), 1.0) + self.assertEqual(result[3, 0].item(), 0.0) + + def test_lhs_none_fills_default(self): + """Test that missing lhs is filled with default value before concatenation.""" + rhs = torch.ones(3, VOCAB_SIZE) + result = merge_bias_tensor(None, rhs, 2, 3, DEVICE, 0.0) + self.assertEqual(result.shape, (5, VOCAB_SIZE)) + # First 2 rows filled with default (0.0) + self.assertEqual(result[0, 0].item(), 0.0) + # Last 3 rows from rhs + self.assertEqual(result[2, 0].item(), 1.0) + + def test_rhs_none_fills_default(self): + """Test that missing rhs is filled with default value before concatenation.""" + lhs = torch.ones(2, VOCAB_SIZE) + result = merge_bias_tensor(lhs, None, 2, 3, DEVICE, 0.0) + self.assertEqual(result.shape, (5, VOCAB_SIZE)) + self.assertEqual(result[0, 0].item(), 1.0) + # Last 3 rows filled with default (0.0) + self.assertEqual(result[3, 0].item(), 0.0) + + def test_custom_default_value(self): + """Test that a custom default (-1.0) fills the missing lhs rows.""" + rhs = torch.ones(1, VOCAB_SIZE) + result = merge_bias_tensor(None, rhs, 2, 1, DEVICE, -1.0) + self.assertEqual(result[0, 0].item(), -1.0) + self.assertEqual(result[1, 0].item(), -1.0) + self.assertEqual(result[2, 0].item(), 1.0) + + +# SamplingBatchInfo.__len__ +class TestSamplingBatchInfoLen(CustomTestCase): + + def test_len_matches_batch_size(self): + """Test that __len__ returns batch size (number of temperature rows).""" + info = _make_info(batch_size=5) + self.assertEqual(len(info), 5) + + +class TestMergeCustomLogitProcessor(CustomTestCase): + + def test_both_none_returns_none(self): + """Test that merging two None processor dicts returns None.""" + result = SamplingBatchInfo.merge_custom_logit_processor( + None, None, 2, 3, DEVICE + ) + self.assertIsNone(result) + + def test_same_key_merges_masks(self): + """Test that same processor key concatenates the boolean masks.""" + proc = MagicMock() + lhs = {42: (proc, torch.tensor([True, False]))} + rhs = {42: (proc, torch.tensor([False, True, True]))} + result = SamplingBatchInfo.merge_custom_logit_processor(lhs, rhs, 2, 3, DEVICE) + self.assertIn(42, result) + self.assertEqual(result[42][1].shape[0], 5) + self.assertTrue(result[42][1][0].item()) # from lhs + self.assertFalse(result[42][1][1].item()) # from lhs + self.assertTrue(result[42][1][3].item()) # from rhs + + def test_disjoint_keys(self): + """Test that disjoint processor keys are merged with zero-filled padding.""" + proc_a = MagicMock() + proc_b = MagicMock() + lhs = {1: (proc_a, torch.tensor([True, False]))} + rhs = {2: (proc_b, torch.tensor([True]))} + result = SamplingBatchInfo.merge_custom_logit_processor(lhs, rhs, 2, 1, DEVICE) + # Key 1: lhs mask [True, False] + zero-filled rhs [False] + self.assertEqual(result[1][1].shape[0], 3) + self.assertTrue(result[1][1][0].item()) + self.assertFalse(result[1][1][2].item()) + # Key 2: zero-filled lhs [False, False] + rhs mask [True] + self.assertEqual(result[2][1].shape[0], 3) + self.assertFalse(result[2][1][0].item()) + self.assertTrue(result[2][1][2].item()) + + def test_lhs_none_rhs_present(self): + """Test that None lhs is treated as empty dict and rhs mask is padded.""" + proc = MagicMock() + rhs = {10: (proc, torch.tensor([True]))} + result = SamplingBatchInfo.merge_custom_logit_processor(None, rhs, 2, 1, DEVICE) + self.assertIn(10, result) + self.assertEqual(result[10][1].shape[0], 3) + + +# apply_logits_bias +class TestApplyLogitsBias(CustomTestCase): + + def test_applies_linear_penalties(self): + """Test that pre-accumulated linear penalties are added to logits.""" + info = _make_info(batch_size=1) + info.acc_linear_penalties = torch.tensor([[-1.0] * VOCAB_SIZE]) + logits = torch.zeros(1, VOCAB_SIZE) + info.apply_logits_bias(logits) + self.assertAlmostEqual(logits[0, 0].item(), -1.0, places=5) + + def test_applies_logit_bias(self): + """Test that per-token logit_bias is added to logits.""" + info = _make_info(batch_size=1) + bias = torch.zeros(1, VOCAB_SIZE) + bias[0, 5] = 10.0 + info.logit_bias = bias + logits = torch.zeros(1, VOCAB_SIZE) + info.apply_logits_bias(logits) + self.assertAlmostEqual(logits[0, 5].item(), 10.0, places=5) + self.assertAlmostEqual(logits[0, 0].item(), 0.0, places=5) + + def test_applies_vocab_mask(self): + """Test that vocab_mask triggers the apply_mask_func callback.""" + info = _make_info(batch_size=1) + info.vocab_mask = torch.ones(1, VOCAB_SIZE) + info.apply_mask_func = MagicMock() + logits = torch.zeros(1, VOCAB_SIZE) + info.apply_logits_bias(logits) + info.apply_mask_func.assert_called_once() + + def test_applies_penalizer_orchestrator(self): + """Test that a required orchestrator's apply() is called on logits.""" + orch = MagicMock(is_required=True) + info = _make_info(batch_size=1, penalizer_orchestrator=orch) + logits = torch.zeros(1, VOCAB_SIZE) + info.apply_logits_bias(logits) + orch.apply.assert_called_once_with(logits) + + def test_no_bias_no_change(self): + """Test that logits stay unchanged when no bias sources are set.""" + info = _make_info(batch_size=1) + info.acc_linear_penalties = None + info.logit_bias = None + info.vocab_mask = None + logits = torch.zeros(1, VOCAB_SIZE) + original = logits.clone() + info.apply_logits_bias(logits) + self.assertTrue(torch.equal(logits, original)) + + +# update_penalties +class TestUpdatePenalties(CustomTestCase): + + def test_required_creates_penalties_tensor(self): + """Test that update_penalties allocates a zero tensor and calls orchestrator.apply.""" + orch = MagicMock(is_required=True) + info = _make_info(batch_size=2, penalizer_orchestrator=orch) + info.update_penalties() + self.assertIsNotNone(info.acc_linear_penalties) + self.assertEqual(info.acc_linear_penalties.shape, (2, VOCAB_SIZE)) + orch.apply.assert_called_once() + + def test_not_required_sets_none(self): + """Test that update_penalties sets acc_linear_penalties to None when not required.""" + orch = MagicMock(is_required=False) + info = _make_info(batch_size=2, penalizer_orchestrator=orch) + info.update_penalties() + self.assertIsNone(info.acc_linear_penalties) + + +# update_regex_vocab_mask +class TestUpdateRegexVocabMask(CustomTestCase): + + def test_no_grammars_clears_mask(self): + """Test that None grammars clears both vocab_mask and apply_mask_func.""" + info = _make_info(batch_size=1) + info.grammars = None + info.update_regex_vocab_mask() + self.assertIsNone(info.vocab_mask) + self.assertIsNone(info.apply_mask_func) + + def test_empty_grammars_clears_mask(self): + """Test that empty grammars list clears vocab_mask.""" + info = _make_info(batch_size=1) + info.grammars = [] + info.update_regex_vocab_mask() + self.assertIsNone(info.vocab_mask) + + def test_with_grammars_allocates_and_fills(self): + """Test that an active grammar gets allocate, fill, and move called.""" + grammar = MagicMock() + grammar.finished = False + grammar.is_terminated.return_value = False + grammar.allocate_vocab_mask.return_value = torch.zeros(1, VOCAB_SIZE) + grammar.move_vocab_mask.return_value = torch.zeros(1, VOCAB_SIZE) + info = _make_info(batch_size=1) + info.grammars = [grammar] + info.update_regex_vocab_mask() + grammar.allocate_vocab_mask.assert_called_once() + grammar.fill_vocab_mask.assert_called_once() + grammar.move_vocab_mask.assert_called_once() + + def test_mixed_grammars_only_active_fills(self): + """Test that finished, terminated, and None grammars are skipped.""" + active = MagicMock() + active.finished = False + active.is_terminated.return_value = False + active.allocate_vocab_mask.return_value = torch.zeros(3, VOCAB_SIZE) + active.move_vocab_mask.return_value = torch.zeros(3, VOCAB_SIZE) + + finished = MagicMock() + finished.finished = True + + terminated = MagicMock() + terminated.finished = False + terminated.is_terminated.return_value = True + + info = _make_info(batch_size=3) + info.grammars = [active, finished, terminated] + info.update_regex_vocab_mask() + + active.fill_vocab_mask.assert_called_once() + finished.fill_vocab_mask.assert_not_called() + terminated.fill_vocab_mask.assert_not_called() + + +# filter_batch +class TestFilterBatch(CustomTestCase): + + def test_filter_keeps_correct_indices(self): + """Test that filter retains rows at indices 0 and 2, dropping index 1.""" + info = _make_info(batch_size=3) + info.temperatures = torch.tensor([[1.0], [2.0], [3.0]]) + info.top_ps = torch.tensor([0.9, 0.8, 0.7]) + info.top_ks = torch.tensor([10, 20, 30], dtype=torch.int32) + info.min_ps = torch.tensor([0.0, 0.1, 0.2]) + info.logit_bias = torch.ones(3, VOCAB_SIZE) + keep = torch.tensor([0, 2]) + info.filter_batch([0, 2], keep) + self.assertEqual(len(info), 2) + self.assertAlmostEqual(info.temperatures[0, 0].item(), 1.0) + self.assertAlmostEqual(info.temperatures[1, 0].item(), 3.0) + self.assertAlmostEqual(info.top_ps[1].item(), 0.7) + # logit_bias should also be filtered + self.assertEqual(info.logit_bias.shape, (2, VOCAB_SIZE)) + + def test_filter_with_custom_logit_processor(self): + """Test that filter updates both custom_params list and processor mask.""" + proc = MagicMock() + info = _make_info(batch_size=3) + info.has_custom_logit_processor = True + info.custom_logit_processor = {42: (proc, torch.tensor([True, False, True]))} + info.custom_params = [{"a": 1}, {"b": 2}, {"c": 3}] + keep = torch.tensor([0, 2]) + info.filter_batch([0, 2], keep) + self.assertEqual(info.custom_params, [{"a": 1}, {"c": 3}]) + mask = info.custom_logit_processor[42][1] + self.assertEqual(mask.shape[0], 2) + + def test_filter_removes_all_custom_processors(self): + """Test cleanup when filter removes all requests using a processor.""" + proc = MagicMock() + info = _make_info(batch_size=3) + info.has_custom_logit_processor = True + info.custom_logit_processor = {42: (proc, torch.tensor([False, True, False]))} + info.custom_params = [None, {"x": 1}, None] + # Keep only index 0 and 2 — processor 42's mask becomes [False, False] + keep = torch.tensor([0, 2]) + info.filter_batch([0, 2], keep) + self.assertFalse(info.has_custom_logit_processor) + self.assertIsNone(info.custom_logit_processor) + + def test_filter_with_none_sampling_seed(self): + """Test that filter preserves None sampling_seed without error.""" + info = _make_info(batch_size=3) + info.sampling_seed = None + keep = torch.tensor([1]) + info.filter_batch([1], keep) + self.assertIsNone(info.sampling_seed) + + +# merge_batch +class TestMergeBatch(CustomTestCase): + + def test_merge_concatenates_tensors(self): + """Test that merge concatenates temperature tensors from both batches.""" + info1 = _make_info(batch_size=2) + info1.temperatures = torch.tensor([[1.0], [2.0]]) + info2 = _make_info(batch_size=1) + info2.temperatures = torch.tensor([[3.0]]) + info1.merge_batch(info2) + self.assertEqual(len(info1), 3) + self.assertAlmostEqual(info1.temperatures[2, 0].item(), 3.0) + + def test_merge_combines_flags(self): + """Test that merge ANDs is_all_greedy and ORs need_*_sampling flags.""" + info1 = _make_info( + is_all_greedy=True, + need_top_p_sampling=False, + need_top_k_sampling=False, + need_min_p_sampling=False, + ) + info2 = _make_info( + is_all_greedy=False, + need_top_p_sampling=True, + need_top_k_sampling=True, + need_min_p_sampling=True, + ) + info1.merge_batch(info2) + self.assertFalse(info1.is_all_greedy) # AND semantics + self.assertTrue(info1.need_top_p_sampling) # OR semantics + self.assertTrue(info1.need_top_k_sampling) # OR semantics + self.assertTrue(info1.need_min_p_sampling) # OR semantics + + def test_merge_with_logit_bias(self): + """Test that merge pads missing logit_bias with zeros before concatenation.""" + info1 = _make_info(batch_size=1) + info1.logit_bias = torch.ones(1, VOCAB_SIZE) + info2 = _make_info(batch_size=1) + info2.logit_bias = None + info1.merge_batch(info2) + self.assertEqual(info1.logit_bias.shape, (2, VOCAB_SIZE)) + + def test_merge_with_custom_logit_processor(self): + """Test that merge combines processors when only one side has them.""" + proc = MagicMock() + info1 = _make_info(batch_size=1) + info1.has_custom_logit_processor = True + info1.custom_logit_processor = {1: (proc, torch.tensor([True]))} + info1.custom_params = [{"a": 1}] + info2 = _make_info(batch_size=1) + info2.has_custom_logit_processor = False + info2.custom_logit_processor = None + info2.custom_params = None + info1.merge_batch(info2) + self.assertTrue(info1.has_custom_logit_processor) + self.assertEqual(len(info1.custom_params), 2) + + def test_merge_with_none_sampling_seed(self): + """Test that merge preserves None when both sampling_seeds are None.""" + info1 = _make_info(batch_size=1) + info1.sampling_seed = None + info2 = _make_info(batch_size=1) + info2.sampling_seed = None + info1.merge_batch(info2) + self.assertIsNone(info1.sampling_seed) + + def test_merge_with_both_sampling_seeds(self): + """Test that merge concatenates both sampling_seed tensors.""" + info1 = _make_info(batch_size=2) + info1.sampling_seed = torch.tensor([10, 20], dtype=torch.int64) + info2 = _make_info(batch_size=1) + info2.sampling_seed = torch.tensor([30], dtype=torch.int64) + info1.merge_batch(info2) + self.assertEqual(info1.sampling_seed.shape[0], 3) + self.assertEqual(info1.sampling_seed[0].item(), 10) + self.assertEqual(info1.sampling_seed[1].item(), 20) + self.assertEqual(info1.sampling_seed[2].item(), 30) + + +# copy_for_forward +class TestCopyForForward(CustomTestCase): + + def test_returns_copy_without_orchestrator(self): + """Test that copy_for_forward returns a copy with orchestrator set to None.""" + orch = MagicMock(is_required=False) + info = _make_info(batch_size=1, penalizer_orchestrator=orch) + copied = info.copy_for_forward() + self.assertIsNone(copied.penalizer_orchestrator) + # Original should still have orchestrator + self.assertIsNotNone(info.penalizer_orchestrator) + + +# from_schedule_batch +class TestFromScheduleBatch(CustomTestCase): + + def _make_req( + self, + temp=1.0, + top_p=1.0, + top_k=-1, + min_p=0.0, + freq=0.0, + presence=0.0, + min_tokens=0, + logit_bias=None, + seed=None, + stop_ids=None, + eos_id=2, + ): + req = MagicMock() + req.sampling_params.temperature = temp + req.sampling_params.top_p = top_p + req.sampling_params.top_k = top_k + req.sampling_params.min_p = min_p + req.sampling_params.frequency_penalty = freq + req.sampling_params.presence_penalty = presence + req.sampling_params.min_new_tokens = min_tokens + req.sampling_params.logit_bias = logit_bias + req.sampling_params.sampling_seed = seed + req.sampling_params.stop_token_ids = stop_ids + req.sampling_params.custom_params = None + req.custom_logit_processor = None + req.tokenizer.additional_stop_token_ids = None + req.tokenizer.eos_token_id = eos_id + return req + + @patch("sglang.srt.sampling.sampling_batch_info.get_global_server_args") + def test_basic_construction(self, mock_server_args): + """Test that from_schedule_batch correctly extracts sampling params from requests.""" + mock_server_args.return_value.enable_deterministic_inference = False + mock_server_args.return_value.enable_custom_logit_processor = False + + reqs = [self._make_req(temp=0.8, top_p=0.9, top_k=50, min_p=0.1)] + batch = MagicMock() + batch.reqs = reqs + batch.device = DEVICE + + info = SamplingBatchInfo.from_schedule_batch(batch, VOCAB_SIZE) + self.assertEqual(len(info), 1) + self.assertAlmostEqual(info.temperatures[0, 0].item(), 0.8, places=5) + self.assertAlmostEqual(info.top_ps[0].item(), 0.9, places=5) + self.assertEqual(info.top_ks[0].item(), 50) + + @patch("sglang.srt.sampling.sampling_batch_info.get_global_server_args") + def test_greedy_detection(self, mock_server_args): + """Test that top_k=1 sets is_all_greedy=True.""" + mock_server_args.return_value.enable_deterministic_inference = False + mock_server_args.return_value.enable_custom_logit_processor = False + + reqs = [self._make_req(top_k=1)] + batch = MagicMock() + batch.reqs = reqs + batch.device = DEVICE + info = SamplingBatchInfo.from_schedule_batch(batch, VOCAB_SIZE) + self.assertTrue(info.is_all_greedy) + + @patch("sglang.srt.sampling.sampling_batch_info.get_global_server_args") + def test_logit_bias_construction(self, mock_server_args): + """Test that logit_bias dict is converted to a tensor with correct values.""" + mock_server_args.return_value.enable_deterministic_inference = False + mock_server_args.return_value.enable_custom_logit_processor = False + + reqs = [self._make_req(logit_bias={"5": 2.0, "10": -1.0})] + batch = MagicMock() + batch.reqs = reqs + batch.device = DEVICE + info = SamplingBatchInfo.from_schedule_batch(batch, VOCAB_SIZE) + self.assertIsNotNone(info.logit_bias) + self.assertAlmostEqual(info.logit_bias[0, 5].item(), 2.0) + self.assertAlmostEqual(info.logit_bias[0, 10].item(), -1.0) + self.assertAlmostEqual(info.logit_bias[0, 0].item(), 0.0) + + @patch("sglang.srt.sampling.sampling_batch_info.get_global_server_args") + def test_deterministic_seed(self, mock_server_args): + """Test that explicit seed=123 is kept and missing seed defaults to 42.""" + mock_server_args.return_value.enable_deterministic_inference = True + mock_server_args.return_value.enable_custom_logit_processor = False + + reqs = [self._make_req(seed=123), self._make_req(seed=None)] + batch = MagicMock() + batch.reqs = reqs + batch.device = DEVICE + info = SamplingBatchInfo.from_schedule_batch(batch, VOCAB_SIZE) + self.assertIsNotNone(info.sampling_seed) + self.assertEqual(info.sampling_seed[0].item(), 123) + self.assertEqual(info.sampling_seed[1].item(), 42) # default + + @patch("sglang.srt.sampling.sampling_batch_info.get_global_server_args") + def test_from_schedule_batch_sampling_flags(self, mock_server_args): + """Test that sampling flags (need_top_p/top_k/min_p) are set correctly.""" + mock_server_args.return_value.enable_deterministic_inference = False + mock_server_args.return_value.enable_custom_logit_processor = False + + reqs = [self._make_req(top_p=0.9, top_k=50, min_p=0.1)] + batch = MagicMock() + batch.reqs = reqs + batch.device = DEVICE + info = SamplingBatchInfo.from_schedule_batch(batch, VOCAB_SIZE) + self.assertTrue(info.need_top_p_sampling) # 0.9 != 1.0 + self.assertTrue(info.need_top_k_sampling) # 50 != TOP_K_ALL + self.assertTrue(info.need_min_p_sampling) # 0.1 > 0 + self.assertFalse(info.is_all_greedy) # top_k=50 > 1 + + @patch("sglang.srt.sampling.sampling_batch_info.get_global_server_args") + def test_no_logit_bias_when_all_none(self, mock_server_args): + """Test that logit_bias stays None when no request has logit_bias set.""" + mock_server_args.return_value.enable_deterministic_inference = False + mock_server_args.return_value.enable_custom_logit_processor = False + + reqs = [self._make_req(), self._make_req()] + batch = MagicMock() + batch.reqs = reqs + batch.device = DEVICE + info = SamplingBatchInfo.from_schedule_batch(batch, VOCAB_SIZE) + self.assertIsNone(info.logit_bias) + + @patch("sglang.srt.sampling.sampling_batch_info.get_global_server_args") + def test_custom_logit_processor_merging(self, mock_server_args): + """Test deserialization and merging of custom logit processors.""" + from sglang.srt.sampling.custom_logit_processor import ( + DisallowedTokensLogitsProcessor, + ) + + mock_server_args.return_value.enable_deterministic_inference = False + mock_server_args.return_value.enable_custom_logit_processor = True + + proc_str = DisallowedTokensLogitsProcessor.to_str() + req1 = self._make_req() + req1.custom_logit_processor = proc_str + req1.sampling_params.custom_params = {"token_ids": [1]} + req2 = self._make_req() + req2.custom_logit_processor = None # no processor + req2.sampling_params.custom_params = None + + batch = MagicMock() + batch.reqs = [req1, req2] + batch.device = DEVICE + info = SamplingBatchInfo.from_schedule_batch(batch, VOCAB_SIZE) + + self.assertTrue(info.has_custom_logit_processor) + self.assertIsNotNone(info.custom_logit_processor) + self.assertEqual(len(info.custom_logit_processor), 1) + # Check the mask: req1 has processor (True), req2 doesn't (False) + key = list(info.custom_logit_processor.keys())[0] + proc, mask = info.custom_logit_processor[key] + self.assertIsInstance(proc, DisallowedTokensLogitsProcessor) + self.assertTrue(mask[0].item()) + self.assertFalse(mask[1].item()) + # custom_params should be collected for all reqs + self.assertEqual(len(info.custom_params), 2) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/sampling/test_sampling_params.py b/test/registered/unit/sampling/test_sampling_params.py new file mode 100644 index 000000000..bb16da796 --- /dev/null +++ b/test/registered/unit/sampling/test_sampling_params.py @@ -0,0 +1,398 @@ +"""Unit tests for srt/sampling/sampling_params.py — no server, no model loading.""" + +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=5, suite="stage-a-cpu-only") + +import unittest +from unittest.mock import MagicMock + +from sglang.srt.sampling.sampling_params import ( + MAX_LEN, + 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) + + +class TestSamplingParamsVerify(CustomTestCase): + + VOCAB_SIZE = 32000 + + 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_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) + + # --- 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_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_boundaries_valid(self): + """Test that boundary values 0.0 and 2.0 are both accepted.""" + self._make(repetition_penalty=0.0).verify(self.VOCAB_SIZE) + self._make(repetition_penalty=2.0).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): + """Test that verify() rejects setting both json_schema and regex (mutually exclusive).""" + sp = self._make(json_schema='{"type":"object"}', regex="abc") + with self.assertRaises(ValueError): + sp.verify(self.VOCAB_SIZE) + + def test_single_grammar_valid(self): + """Test that setting only one grammar type is accepted.""" + sp = self._make(json_schema='{"type":"object"}') + sp.verify(self.VOCAB_SIZE) + + def test_all_three_grammars_set_raises(self): + """Test that verify() rejects setting json_schema, regex, and ebnf together.""" + sp = self._make(json_schema="{}", regex="a", ebnf="rule") + with self.assertRaises(ValueError): + sp.verify(self.VOCAB_SIZE) + + +class TestSamplingParamsNormalize(CustomTestCase): + + 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|>") + sp.normalize(tokenizer=None) + 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"]) + sp.normalize(tokenizer=None) + self.assertEqual(sp.stop_strs, ["stop1", "stop2"]) + + def test_stop_str_max_len_without_tokenizer(self): + """Test that without a tokenizer, max_len is the raw string character count.""" + sp = SamplingParams(stop=["ab", "cdef"]) + sp.normalize(tokenizer=None) + self.assertEqual(sp.stop_str_max_len, 4) # len("cdef") + + 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+") + sp.normalize(tokenizer=None) + 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}") + sp.normalize(tokenizer=None) + self.assertEqual(sp.stop_regex_max_len, 3) + + +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()