[3/N][Sync sglang-miles] TITO Support (#23751)

Co-authored-by: Jiajun Li <48857426+guapisolo@users.noreply.github.com>
This commit is contained in:
Yuzhen Zhou
2026-06-03 21:45:33 -04:00
committed by GitHub
co-authored by Jiajun Li
parent 084c6a7e2a
commit e03dfa8182
7 changed files with 217 additions and 9 deletions
@@ -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
@@ -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)
+9
View File
@@ -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
@@ -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 = {
@@ -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."""
@@ -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 ---
@@ -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")