diff --git a/test/manual/test_tokenizer_manager.py b/test/manual/test_tokenizer_manager.py index 04115fb99..9525fedbb 100644 --- a/test/manual/test_tokenizer_manager.py +++ b/test/manual/test_tokenizer_manager.py @@ -14,7 +14,7 @@ python3 -m unittest test_tokenizer_manager.TestTokenizerManagerIntegration import unittest from unittest.mock import Mock, patch -from sglang.srt.managers.tokenizer_manager import TokenizerManager +from sglang.srt.managers.tokenizer_manager import InputFormat, TokenizerManager from sglang.srt.server_args import PortArgs, ServerArgs from sglang.test.test_utils import DEFAULT_SMALL_MODEL_NAME_FOR_TEST @@ -42,7 +42,7 @@ class TestInputFormatDetection(unittest.TestCase): result = self.tokenizer_manager._detect_input_format( text, is_cross_encoder=False ) - self.assertEqual(result, "single_string") + self.assertEqual(result, InputFormat.SINGLE_STRING) def test_detect_single_string_cross_encoder_disabled(self): """Test single string with cross_encoder disabled still returns single_string.""" @@ -50,7 +50,7 @@ class TestInputFormatDetection(unittest.TestCase): result = self.tokenizer_manager._detect_input_format( text, is_cross_encoder=True ) - self.assertEqual(result, "single_string") + self.assertEqual(result, InputFormat.SINGLE_STRING) def test_detect_batch_strings(self): """Test detection of batch string inputs.""" @@ -58,7 +58,7 @@ class TestInputFormatDetection(unittest.TestCase): result = self.tokenizer_manager._detect_input_format( texts, is_cross_encoder=False ) - self.assertEqual(result, "batch_strings") + self.assertEqual(result, InputFormat.BATCH_STRINGS) def test_detect_batch_strings_cross_encoder_disabled(self): """Test batch strings with cross_encoder disabled.""" @@ -66,7 +66,7 @@ class TestInputFormatDetection(unittest.TestCase): result = self.tokenizer_manager._detect_input_format( texts, is_cross_encoder=True ) - self.assertEqual(result, "batch_strings") + self.assertEqual(result, InputFormat.BATCH_STRINGS) def test_detect_cross_encoder_single_pair(self): """Test detection of cross-encoder single pair.""" @@ -74,7 +74,7 @@ class TestInputFormatDetection(unittest.TestCase): result = self.tokenizer_manager._detect_input_format( texts, is_cross_encoder=True ) - self.assertEqual(result, "cross_encoder_pairs") + self.assertEqual(result, InputFormat.CROSS_ENCODER_PAIRS) def test_detect_cross_encoder_multiple_pairs(self): """Test detection of cross-encoder multiple pairs.""" @@ -82,7 +82,7 @@ class TestInputFormatDetection(unittest.TestCase): result = self.tokenizer_manager._detect_input_format( texts, is_cross_encoder=True ) - self.assertEqual(result, "cross_encoder_pairs") + self.assertEqual(result, InputFormat.CROSS_ENCODER_PAIRS) def test_detect_cross_encoder_disabled_with_pairs(self): """Test pairs with cross_encoder disabled should return batch_strings.""" @@ -90,7 +90,7 @@ class TestInputFormatDetection(unittest.TestCase): result = self.tokenizer_manager._detect_input_format( texts, is_cross_encoder=False ) - self.assertEqual(result, "batch_strings") + self.assertEqual(result, InputFormat.BATCH_STRINGS) def test_detect_empty_list(self): """Test detection with empty list.""" @@ -98,7 +98,7 @@ class TestInputFormatDetection(unittest.TestCase): result = self.tokenizer_manager._detect_input_format( texts, is_cross_encoder=True ) - self.assertEqual(result, "batch_strings") + self.assertEqual(result, InputFormat.BATCH_STRINGS) def test_detect_malformed_cross_encoder_pairs(self): """Test malformed cross-encoder pairs (not length 2).""" @@ -106,13 +106,13 @@ class TestInputFormatDetection(unittest.TestCase): result = self.tokenizer_manager._detect_input_format( texts, is_cross_encoder=True ) - self.assertEqual(result, "batch_strings") + self.assertEqual(result, InputFormat.BATCH_STRINGS) texts = [["query", "doc", "extra"]] # Three elements, not a pair result = self.tokenizer_manager._detect_input_format( texts, is_cross_encoder=True ) - self.assertEqual(result, "batch_strings") + self.assertEqual(result, InputFormat.BATCH_STRINGS) class TestTokenizerInputPreparation(unittest.TestCase): @@ -135,20 +135,24 @@ class TestTokenizerInputPreparation(unittest.TestCase): def test_prepare_single_string_input(self): """Test preparation of single string input.""" text = "Hello world" - result = self.tokenizer_manager._prepare_tokenizer_input(text, "single_string") + result = self.tokenizer_manager._prepare_tokenizer_input( + text, InputFormat.SINGLE_STRING + ) self.assertEqual(result, ["Hello world"]) def test_prepare_batch_strings_input(self): """Test preparation of batch strings input.""" texts = ["Hello", "World", "Test"] - result = self.tokenizer_manager._prepare_tokenizer_input(texts, "batch_strings") + result = self.tokenizer_manager._prepare_tokenizer_input( + texts, InputFormat.BATCH_STRINGS + ) self.assertEqual(result, ["Hello", "World", "Test"]) def test_prepare_cross_encoder_pairs_input(self): """Test preparation of cross-encoder pairs input.""" texts = [["query1", "doc1"], ["query2", "doc2"]] result = self.tokenizer_manager._prepare_tokenizer_input( - texts, "cross_encoder_pairs" + texts, InputFormat.CROSS_ENCODER_PAIRS ) self.assertEqual(result, [["query1", "doc1"], ["query2", "doc2"]]) @@ -156,15 +160,15 @@ class TestTokenizerInputPreparation(unittest.TestCase): """Test preparation of single cross-encoder pair.""" texts = [["query text", "document text"]] result = self.tokenizer_manager._prepare_tokenizer_input( - texts, "cross_encoder_pairs" + texts, InputFormat.CROSS_ENCODER_PAIRS ) self.assertEqual(result, [["query text", "document text"]]) - def test_prepare_unknown_input_format(self): - """Test preparation with unknown input format falls back to returning as-is.""" + def test_prepare_batch_strings_input_format_passthrough(self): + """Batch strings should pass through unchanged.""" texts = ["test"] result = self.tokenizer_manager._prepare_tokenizer_input( - texts, "unknown_format" + texts, InputFormat.BATCH_STRINGS ) self.assertEqual(result, ["test"]) @@ -193,7 +197,10 @@ class TestTokenizerResultExtraction(unittest.TestCase): result_input_ids, result_token_type_ids = ( self.tokenizer_manager._extract_tokenizer_results( - input_ids, token_type_ids, "single_string", original_batch_size=1 + input_ids, + token_type_ids, + InputFormat.SINGLE_STRING, + original_batch_size=1, ) ) @@ -207,7 +214,10 @@ class TestTokenizerResultExtraction(unittest.TestCase): result_input_ids, result_token_type_ids = ( self.tokenizer_manager._extract_tokenizer_results( - input_ids, token_type_ids, "cross_encoder_pairs", original_batch_size=1 + input_ids, + token_type_ids, + InputFormat.CROSS_ENCODER_PAIRS, + original_batch_size=1, ) ) @@ -221,7 +231,10 @@ class TestTokenizerResultExtraction(unittest.TestCase): result_input_ids, result_token_type_ids = ( self.tokenizer_manager._extract_tokenizer_results( - input_ids, token_type_ids, "batch_strings", original_batch_size=2 + input_ids, + token_type_ids, + InputFormat.BATCH_STRINGS, + original_batch_size=2, ) ) @@ -235,7 +248,10 @@ class TestTokenizerResultExtraction(unittest.TestCase): result_input_ids, result_token_type_ids = ( self.tokenizer_manager._extract_tokenizer_results( - input_ids, token_type_ids, "cross_encoder_pairs", original_batch_size=2 + input_ids, + token_type_ids, + InputFormat.CROSS_ENCODER_PAIRS, + original_batch_size=2, ) ) @@ -251,7 +267,10 @@ class TestTokenizerResultExtraction(unittest.TestCase): result_input_ids, result_token_type_ids = ( self.tokenizer_manager._extract_tokenizer_results( - input_ids, token_type_ids, "single_string", original_batch_size=1 + input_ids, + token_type_ids, + InputFormat.SINGLE_STRING, + original_batch_size=1, ) ) @@ -265,7 +284,10 @@ class TestTokenizerResultExtraction(unittest.TestCase): result_input_ids, result_token_type_ids = ( self.tokenizer_manager._extract_tokenizer_results( - input_ids, token_type_ids, "single_string", original_batch_size=1 + input_ids, + token_type_ids, + InputFormat.SINGLE_STRING, + original_batch_size=1, ) ) @@ -298,7 +320,7 @@ class TestTokenizerManagerIntegration(unittest.TestCase): input_format = self.tokenizer_manager._detect_input_format( text, is_cross_encoder=False ) - self.assertEqual(input_format, "single_string") + self.assertEqual(input_format, InputFormat.SINGLE_STRING) # Step 2: Prepare input tokenizer_input = self.tokenizer_manager._prepare_tokenizer_input( @@ -329,7 +351,7 @@ class TestTokenizerManagerIntegration(unittest.TestCase): input_format = self.tokenizer_manager._detect_input_format( texts, is_cross_encoder=True ) - self.assertEqual(input_format, "cross_encoder_pairs") + self.assertEqual(input_format, InputFormat.CROSS_ENCODER_PAIRS) # Step 2: Prepare input tokenizer_input = self.tokenizer_manager._prepare_tokenizer_input( @@ -358,7 +380,7 @@ class TestTokenizerManagerIntegration(unittest.TestCase): input_format = self.tokenizer_manager._detect_input_format( texts, is_cross_encoder=False ) - self.assertEqual(input_format, "batch_strings") + self.assertEqual(input_format, InputFormat.BATCH_STRINGS) # Step 2: Prepare input tokenizer_input = self.tokenizer_manager._prepare_tokenizer_input(