[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 = {
|
||||
|
||||
Reference in New Issue
Block a user