From e03dfa8182b95e1b54ada2b3b9816e2fca1e2fa0 Mon Sep 17 00:00:00 2001 From: Yuzhen Zhou <82826991+zyzshishui@users.noreply.github.com> Date: Wed, 3 Jun 2026 18:45:33 -0700 Subject: [PATCH] [3/N][Sync sglang-miles] TITO Support (#23751) Co-authored-by: Jiajun Li <48857426+guapisolo@users.noreply.github.com> --- .../sglang/srt/entrypoints/openai/protocol.py | 13 +++ .../srt/entrypoints/openai/serving_chat.py | 50 +++++++++-- python/sglang/srt/managers/io_struct.py | 9 ++ .../sglang/srt/managers/tokenizer_manager.py | 21 ++++- .../unit/entrypoints/openai/test_protocol.py | 36 ++++++++ .../entrypoints/openai/test_serving_chat.py | 84 ++++++++++++++++++- .../unit/managers/test_io_struct.py | 13 +++ 7 files changed, 217 insertions(+), 9 deletions(-) diff --git a/python/sglang/srt/entrypoints/openai/protocol.py b/python/sglang/srt/entrypoints/openai/protocol.py index 2404731bf..5f1efbb11 100644 --- a/python/sglang/srt/entrypoints/openai/protocol.py +++ b/python/sglang/srt/entrypoints/openai/protocol.py @@ -675,6 +675,8 @@ class ChatCompletionRequest(BaseModel): return_routed_experts: bool = False routed_experts_start_len: int = 0 return_cached_tokens_details: bool = False + return_prompt_token_ids: bool = False + return_meta_info: bool = False reasoning_effort: Optional[Literal["none", "low", "medium", "high", "max"]] = Field( default=None, description="Constrains effort on reasoning for reasoning models. " @@ -724,6 +726,11 @@ class ChatCompletionRequest(BaseModel): custom_logit_processor: Optional[Union[List[Optional[str]], str]] = None custom_params: Optional[Dict] = None + # Pre-computed prompt token IDs: when provided, bypasses chat template + # tokenization entirely. Messages are still used to derive stop tokens + # and tool_call_constraint. + input_ids: Optional[List[int]] = None + # For request id rid: Optional[Union[List[str], str]] = None # Extra key for classifying the request (e.g. cache_salt) @@ -948,12 +955,18 @@ class ChatCompletionResponseChoice(BaseModel): ] = None matched_stop: Union[None, int, str] = None hidden_states: Optional[object] = None + prompt_token_ids: Optional[List[int]] = None + meta_info: Optional[Dict[str, Any]] = None @model_serializer(mode="wrap") def _serialize(self, handler): data = handler(self) if self.hidden_states is None: data.pop("hidden_states", None) + if self.prompt_token_ids is None: + data.pop("prompt_token_ids", None) + if self.meta_info is None: + data.pop("meta_info", None) return data diff --git a/python/sglang/srt/entrypoints/openai/serving_chat.py b/python/sglang/srt/entrypoints/openai/serving_chat.py index cbd3df776..a24b48f54 100644 --- a/python/sglang/srt/entrypoints/openai/serving_chat.py +++ b/python/sglang/srt/entrypoints/openai/serving_chat.py @@ -471,12 +471,22 @@ class OpenAIServingChat(OpenAIServingBase): if reasoning_effort is not None: request.reasoning_effort = reasoning_effort - """Convert OpenAI chat completion request to internal format""" + if request.stream: + if request.return_prompt_token_ids: + raise ValueError( + "return_prompt_token_ids is not supported with streaming. " + "Please set stream=false when using return_prompt_token_ids=true." + ) + if request.return_meta_info: + raise ValueError( + "return_meta_info is not supported with streaming. " + "Please set stream=false when using return_meta_info=true." + ) + is_multimodal = self.tokenizer_manager.model_config.is_multimodal # Process messages and apply chat template processed_messages = self._process_messages(request, is_multimodal) - # Build sampling parameters sampling_params = request.to_sampling_params( stop=processed_messages.stop, @@ -484,8 +494,9 @@ class OpenAIServingChat(OpenAIServingBase): tool_call_constraint=processed_messages.tool_call_constraint, ) - # Handle single vs multiple requests - if is_multimodal: + if request.input_ids is not None: + prompt_kwargs = {"input_ids": processed_messages.prompt_ids} + elif is_multimodal: prompt_kwargs = {"text": processed_messages.prompt} else: if isinstance(processed_messages.prompt_ids, str): @@ -540,6 +551,7 @@ class OpenAIServingChat(OpenAIServingBase): video_max_dynamic_patch=vid_max_dynamic_patch, max_dynamic_patch=getattr(request, "max_dynamic_patch", None), use_audio_in_video=getattr(request, "use_audio_in_video", False), + return_prompt_token_ids=request.return_prompt_token_ids, ) return adapted_request, request @@ -595,8 +607,19 @@ class OpenAIServingChat(OpenAIServingBase): ) tool_call_constraint = ("json_schema", json_schema) - # Use chat template - if self.template_manager.chat_template_name is None: + # When input_ids are provided, skip template tokenization entirely; + # only stop tokens and tool_call_constraint are needed. + if request.input_ids is not None: + result = MessageProcessingResult( + prompt="", + prompt_ids=request.input_ids, + image_data=None, + audio_data=None, + video_data=None, + modalities=[], + stop=request.stop or [], + ) + elif self.template_manager.chat_template_name is None: result = self._apply_jinja_template(request, tools, is_multimodal) else: result = self._apply_conversation_template(request, is_multimodal) @@ -1245,6 +1268,17 @@ class OpenAIServingChat(OpenAIServingBase): history_tool_calls_cnt, ) + # Extract prompt_token_ids if requested + choice_prompt_token_ids = ( + ret_item.get("prompt_token_ids") + if request.return_prompt_token_ids + else None + ) + + choice_meta_info = ( + ret_item["meta_info"] if request.return_meta_info else None + ) + # NOTE: content should not be None but empty string to make sure retokenize consistency. reasoning_text, tool_calls = self._get_parsed_response_fields( reasoning_text, tool_calls ) @@ -1253,7 +1287,7 @@ class OpenAIServingChat(OpenAIServingBase): index=idx, message=ChatMessage( role="assistant", - content=text if text else None, + content=text if text else "", tool_calls=tool_calls, reasoning_content=reasoning_text if reasoning_text else None, ), @@ -1265,6 +1299,8 @@ class OpenAIServingChat(OpenAIServingBase): else None ), hidden_states=hidden_states, + prompt_token_ids=choice_prompt_token_ids, + meta_info=choice_meta_info, ) choices.append(choice_data) diff --git a/python/sglang/srt/managers/io_struct.py b/python/sglang/srt/managers/io_struct.py index 439908e49..782955944 100644 --- a/python/sglang/srt/managers/io_struct.py +++ b/python/sglang/srt/managers/io_struct.py @@ -259,6 +259,9 @@ class GenerateReqInput(BaseReq): # Whether to return entropy return_entropy: bool = False + # Whether to return prompt token IDs without computing logprobs + return_prompt_token_ids: bool = False + # Propagates trace context via Engine.generate/async_generate external_trace_header: Optional[Dict] = None received_time: Optional[float] = None @@ -712,6 +715,7 @@ class GenerateReqInput(BaseReq): custom_labels=self.custom_labels, return_bytes=self.return_bytes, return_entropy=self.return_entropy, + return_prompt_token_ids=self.return_prompt_token_ids, external_trace_header=self.external_trace_header, http_worker_ipc=self.http_worker_ipc, received_time=self.received_time, @@ -896,6 +900,9 @@ class EmbeddingReqInput(BaseReq): # Whether to return pooled hidden states (pre-head transformer output) return_pooled_hidden_states: bool = False + # Whether to return prompt token IDs without computing logprobs + return_prompt_token_ids: bool = False + # Pre-computed delimiter indices for multi-item scoring. # Batch-level: List[List[int]] (one per request). After __getitem__: List[int]. multi_item_delimiter_indices: Optional[Union[List[List[int]], List[int]]] = None @@ -1002,6 +1009,7 @@ class EmbeddingReqInput(BaseReq): is_cross_encoder_request=True, http_worker_ipc=self.http_worker_ipc, return_pooled_hidden_states=self.return_pooled_hidden_states, + return_prompt_token_ids=self.return_prompt_token_ids, multi_item_delimiter_indices=( self.multi_item_delimiter_indices[i] if self.multi_item_delimiter_indices is not None @@ -1031,6 +1039,7 @@ class EmbeddingReqInput(BaseReq): http_worker_ipc=self.http_worker_ipc, received_time=self.received_time, return_pooled_hidden_states=self.return_pooled_hidden_states, + return_prompt_token_ids=self.return_prompt_token_ids, multi_item_delimiter_indices=( self.multi_item_delimiter_indices[i] if self.multi_item_delimiter_indices is not None diff --git a/python/sglang/srt/managers/tokenizer_manager.py b/python/sglang/srt/managers/tokenizer_manager.py index cb53d1844..efa1dad93 100644 --- a/python/sglang/srt/managers/tokenizer_manager.py +++ b/python/sglang/srt/managers/tokenizer_manager.py @@ -206,6 +206,9 @@ class ReqState: input_token_ids_logprobs: List[Any] = dataclasses.field(default_factory=list) output_token_ids_logprobs: List[Any] = dataclasses.field(default_factory=list) + # For return_prompt_token_ids: stores prompt token IDs captured after tokenization + prompt_token_ids: Optional[List[int]] = None + def _slice_streaming_output_meta_info( meta_info: Dict[Any, Any], @@ -586,6 +589,9 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): # Tokenize the request and send it to the scheduler if obj.is_single: tokenized_obj = await self._tokenize_one_request(obj) + state = self.rid_to_state[obj.rid] + if obj.return_prompt_token_ids: + state.prompt_token_ids = list(tokenized_obj.input_ids) self._send_one_request(tokenized_obj) async for response in self._wait_one_response(obj, request): yield response @@ -1478,6 +1484,9 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): # Set up generators for each request in the batch for i in range(batch_size): tmp_obj = obj[i] + state = self.rid_to_state[tmp_obj.rid] + if tmp_obj.return_prompt_token_ids: + state.prompt_token_ids = list(tokenized_objs[i].input_ids) generators.append(self._wait_one_response(tmp_obj, request)) rids.append(tmp_obj.rid) else: @@ -1490,6 +1499,9 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): for i in range(batch_size): tmp_obj = obj[i] tokenized_obj = await self._tokenize_one_request(tmp_obj) + state = self.rid_to_state[tmp_obj.rid] + if tmp_obj.return_prompt_token_ids: + state.prompt_token_ids = list(tokenized_obj.input_ids) self._send_one_request(tokenized_obj) generators.append(self._wait_one_response(tmp_obj, request)) rids.append(tmp_obj.rid) @@ -1539,7 +1551,10 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): ] tokenized_obj.rid = tmp_obj.regenerate_rid() self._init_req_state(tmp_obj) - tokenized_obj.time_stats = self.rid_to_state[tmp_obj.rid].time_stats + state = self.rid_to_state[tmp_obj.rid] + tokenized_obj.time_stats = state.time_stats + if tmp_obj.return_prompt_token_ids: + state.prompt_token_ids = list(tokenized_objs[i].input_ids) self._send_one_request(tokenized_obj) generators.append(self._wait_one_response(tmp_obj, request)) rids.append(tmp_obj.rid) @@ -1903,6 +1918,8 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): } else: out_dict = None + if out_dict is not None and state.prompt_token_ids is not None: + out_dict["prompt_token_ids"] = state.prompt_token_ids elif isinstance(recv_obj, BatchTokenIDOutput): is_stream = getattr(state.obj, "stream", False) incremental = ( @@ -1938,6 +1955,8 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): } else: out_dict = None + if out_dict is not None and state.prompt_token_ids is not None: + out_dict["prompt_token_ids"] = state.prompt_token_ids else: assert isinstance(recv_obj, BatchEmbeddingOutput) out_dict = { diff --git a/test/registered/unit/entrypoints/openai/test_protocol.py b/test/registered/unit/entrypoints/openai/test_protocol.py index 7ebe87713..19f880455 100644 --- a/test/registered/unit/entrypoints/openai/test_protocol.py +++ b/test/registered/unit/entrypoints/openai/test_protocol.py @@ -179,6 +179,20 @@ class TestChatCompletionRequest(unittest.TestCase): self.assertFalse(request.stream_reasoning) self.assertEqual(request.chat_template_kwargs, {"custom_param": "value"}) + def test_chat_completion_tito_extensions(self): + """Test chat completion with pre-tokenized prompt extensions.""" + messages = [{"role": "user", "content": "Hello"}] + request = ChatCompletionRequest( + model="test-model", + messages=messages, + input_ids=[101, 102, 103], + return_prompt_token_ids=True, + return_meta_info=True, + ) + self.assertEqual(request.input_ids, [101, 102, 103]) + self.assertTrue(request.return_prompt_token_ids) + self.assertTrue(request.return_meta_info) + def test_chat_completion_reasoning_effort(self): """Test chat completion with reasoning effort""" messages = [{"role": "user", "content": "Hello"}] @@ -371,6 +385,28 @@ class TestModelSerialization(unittest.TestCase): self.assertIn("hidden_states", data["choices"][0]) self.assertEqual(data["choices"][0]["hidden_states"], [0.1, 0.2, 0.3]) + def test_prompt_token_ids_and_meta_info_serialization(self): + """Test that prompt_token_ids and meta_info serialize only when set.""" + default_choice = ChatCompletionResponseChoice( + index=0, + message=ChatMessage(role="assistant", content="Hello"), + finish_reason="stop", + ) + default_data = default_choice.model_dump() + self.assertNotIn("prompt_token_ids", default_data) + self.assertNotIn("meta_info", default_data) + + choice = ChatCompletionResponseChoice( + index=0, + message=ChatMessage(role="assistant", content="Hello"), + finish_reason="stop", + prompt_token_ids=[1, 2, 3], + meta_info={"prompt_tokens": 3}, + ) + data = choice.model_dump() + self.assertEqual(data["prompt_token_ids"], [1, 2, 3]) + self.assertEqual(data["meta_info"], {"prompt_tokens": 3}) + class TestFunctionDeferLoading(unittest.TestCase): """Test defer_loading field behavior on Function/Tool.""" diff --git a/test/registered/unit/entrypoints/openai/test_serving_chat.py b/test/registered/unit/entrypoints/openai/test_serving_chat.py index 101402aa5..33df9a611 100644 --- a/test/registered/unit/entrypoints/openai/test_serving_chat.py +++ b/test/registered/unit/entrypoints/openai/test_serving_chat.py @@ -147,6 +147,55 @@ class ServingChatTestCase(unittest.TestCase): self.assertFalse(adapted.stream) self.assertEqual(processed, self.basic_req) + def test_convert_to_internal_request_rejects_stream_return_prompt_token_ids(self): + req = ChatCompletionRequest( + model="x", + messages=[{"role": "user", "content": "Hi?"}], + stream=True, + return_prompt_token_ids=True, + ) + + with self.assertRaisesRegex( + ValueError, "return_prompt_token_ids is not supported with streaming" + ): + self.chat._convert_to_internal_request(req, self.fastapi_request) + + def test_convert_to_internal_request_rejects_stream_return_meta_info(self): + req = ChatCompletionRequest( + model="x", + messages=[{"role": "user", "content": "Hi?"}], + stream=True, + return_meta_info=True, + ) + + with self.assertRaisesRegex( + ValueError, "return_meta_info is not supported with streaming" + ): + self.chat._convert_to_internal_request(req, self.fastapi_request) + + def test_convert_to_internal_request_input_ids_bypasses_template(self): + self.tm.tokenizer = None + req = ChatCompletionRequest( + model="x", + messages=[{"role": "user", "content": "Hi?"}], + input_ids=[101, 102, 103], + stop=["STOP"], + return_prompt_token_ids=True, + ) + + with patch( + "sglang.srt.entrypoints.openai.serving_chat.generate_chat_conv" + ) as conv_mock: + adapted, processed = self.chat._convert_to_internal_request( + req, self.fastapi_request + ) + + self.assertEqual(processed, req) + self.assertEqual(adapted.input_ids, [101, 102, 103]) + self.assertTrue(adapted.return_prompt_token_ids) + self.assertEqual(adapted.sampling_params["stop"], ["STOP"]) + conv_mock.assert_not_called() + def test_kimi_tool_call_keeps_default_reasoning(self): self.template_manager.reasoning_config = ReasoningToggleConfig( toggle_param="thinking", default_enabled=True @@ -1290,6 +1339,39 @@ class ServingChatTestCase(unittest.TestCase): }, ) + def test_non_streaming_chat_response_returns_requested_prompt_ids_and_meta_info( + self, + ): + req = ChatCompletionRequest( + model="x", + messages=[{"role": "user", "content": "Hi?"}], + return_prompt_token_ids=True, + return_meta_info=True, + ) + ret = [ + { + "text": "Answer", + "prompt_token_ids": [11, 12, 13], + "meta_info": { + "id": "chatcmpl-token-ids", + "prompt_tokens": 3, + "completion_tokens": 1, + "cached_tokens": 0, + "finish_reason": {"type": "stop", "matched": None}, + "weight_version": "default", + }, + } + ] + + response = self.chat._build_chat_response(req, ret, created=123) + choice = response.choices[0] + + self.assertEqual(choice.prompt_token_ids, [11, 12, 13]) + self.assertEqual(choice.meta_info, ret[0]["meta_info"]) + dumped_choice = response.model_dump()["choices"][0] + self.assertEqual(dumped_choice["prompt_token_ids"], [11, 12, 13]) + self.assertEqual(dumped_choice["meta_info"], ret[0]["meta_info"]) + def test_streaming_cached_tokens_details_emits_sglext(self): """Test that streaming chat responses emit cached token details in sglext.""" @@ -1709,7 +1791,7 @@ class ServingChatTestCase(unittest.TestCase): response = self.chat._build_chat_response(req, [ret_item], created=0) msg = response.choices[0].message - self.assertIsNone(msg.content) + self.assertEqual(msg.content, "") self.assertEqual(msg.reasoning_content, "42") # --- poolside_v1 (Laguna-XS.2) regression tests --- diff --git a/test/registered/unit/managers/test_io_struct.py b/test/registered/unit/managers/test_io_struct.py index 83e76e4eb..42ed241a1 100644 --- a/test/registered/unit/managers/test_io_struct.py +++ b/test/registered/unit/managers/test_io_struct.py @@ -539,6 +539,19 @@ class TestGenerateReqInputNormalization(CustomTestCase): self.assertEqual(item0.custom_logit_processor, "processor1") self.assertEqual(item0.return_hidden_states, True) + def test_getitem_preserves_return_prompt_token_ids(self): + """Batch subrequests must keep the prompt-token-id return flag.""" + req = GenerateReqInput( + input_ids=[[1, 2, 3], [4, 5, 6]], + sampling_params=[{}, {}], + rid=["id1", "id2"], + return_prompt_token_ids=True, + ) + req.normalize_batch_and_arguments() + + self.assertTrue(req[0].return_prompt_token_ids) + self.assertTrue(req[1].return_prompt_token_ids) + def test_regenerate_rid(self): """Test the regenerate_rid method.""" req = GenerateReqInput(text="Hello")