[3/N][Sync sglang-miles] TITO Support (#23751)
Co-authored-by: Jiajun Li <48857426+guapisolo@users.noreply.github.com>
This commit is contained in:
@@ -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)
|
||||
|
||||
|
||||
@@ -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")
|
||||
|
||||
Reference in New Issue
Block a user