[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 = {