[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
|
return_routed_experts: bool = False
|
||||||
routed_experts_start_len: int = 0
|
routed_experts_start_len: int = 0
|
||||||
return_cached_tokens_details: bool = False
|
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(
|
reasoning_effort: Optional[Literal["none", "low", "medium", "high", "max"]] = Field(
|
||||||
default=None,
|
default=None,
|
||||||
description="Constrains effort on reasoning for reasoning models. "
|
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_logit_processor: Optional[Union[List[Optional[str]], str]] = None
|
||||||
custom_params: Optional[Dict] = 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
|
# For request id
|
||||||
rid: Optional[Union[List[str], str]] = None
|
rid: Optional[Union[List[str], str]] = None
|
||||||
# Extra key for classifying the request (e.g. cache_salt)
|
# Extra key for classifying the request (e.g. cache_salt)
|
||||||
@@ -948,12 +955,18 @@ class ChatCompletionResponseChoice(BaseModel):
|
|||||||
] = None
|
] = None
|
||||||
matched_stop: Union[None, int, str] = None
|
matched_stop: Union[None, int, str] = None
|
||||||
hidden_states: Optional[object] = None
|
hidden_states: Optional[object] = None
|
||||||
|
prompt_token_ids: Optional[List[int]] = None
|
||||||
|
meta_info: Optional[Dict[str, Any]] = None
|
||||||
|
|
||||||
@model_serializer(mode="wrap")
|
@model_serializer(mode="wrap")
|
||||||
def _serialize(self, handler):
|
def _serialize(self, handler):
|
||||||
data = handler(self)
|
data = handler(self)
|
||||||
if self.hidden_states is None:
|
if self.hidden_states is None:
|
||||||
data.pop("hidden_states", 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
|
return data
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -471,12 +471,22 @@ class OpenAIServingChat(OpenAIServingBase):
|
|||||||
if reasoning_effort is not None:
|
if reasoning_effort is not None:
|
||||||
request.reasoning_effort = reasoning_effort
|
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
|
is_multimodal = self.tokenizer_manager.model_config.is_multimodal
|
||||||
|
|
||||||
# Process messages and apply chat template
|
# Process messages and apply chat template
|
||||||
processed_messages = self._process_messages(request, is_multimodal)
|
processed_messages = self._process_messages(request, is_multimodal)
|
||||||
|
|
||||||
# Build sampling parameters
|
# Build sampling parameters
|
||||||
sampling_params = request.to_sampling_params(
|
sampling_params = request.to_sampling_params(
|
||||||
stop=processed_messages.stop,
|
stop=processed_messages.stop,
|
||||||
@@ -484,8 +494,9 @@ class OpenAIServingChat(OpenAIServingBase):
|
|||||||
tool_call_constraint=processed_messages.tool_call_constraint,
|
tool_call_constraint=processed_messages.tool_call_constraint,
|
||||||
)
|
)
|
||||||
|
|
||||||
# Handle single vs multiple requests
|
if request.input_ids is not None:
|
||||||
if is_multimodal:
|
prompt_kwargs = {"input_ids": processed_messages.prompt_ids}
|
||||||
|
elif is_multimodal:
|
||||||
prompt_kwargs = {"text": processed_messages.prompt}
|
prompt_kwargs = {"text": processed_messages.prompt}
|
||||||
else:
|
else:
|
||||||
if isinstance(processed_messages.prompt_ids, str):
|
if isinstance(processed_messages.prompt_ids, str):
|
||||||
@@ -540,6 +551,7 @@ class OpenAIServingChat(OpenAIServingBase):
|
|||||||
video_max_dynamic_patch=vid_max_dynamic_patch,
|
video_max_dynamic_patch=vid_max_dynamic_patch,
|
||||||
max_dynamic_patch=getattr(request, "max_dynamic_patch", None),
|
max_dynamic_patch=getattr(request, "max_dynamic_patch", None),
|
||||||
use_audio_in_video=getattr(request, "use_audio_in_video", False),
|
use_audio_in_video=getattr(request, "use_audio_in_video", False),
|
||||||
|
return_prompt_token_ids=request.return_prompt_token_ids,
|
||||||
)
|
)
|
||||||
|
|
||||||
return adapted_request, request
|
return adapted_request, request
|
||||||
@@ -595,8 +607,19 @@ class OpenAIServingChat(OpenAIServingBase):
|
|||||||
)
|
)
|
||||||
tool_call_constraint = ("json_schema", json_schema)
|
tool_call_constraint = ("json_schema", json_schema)
|
||||||
|
|
||||||
# Use chat template
|
# When input_ids are provided, skip template tokenization entirely;
|
||||||
if self.template_manager.chat_template_name is None:
|
# 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)
|
result = self._apply_jinja_template(request, tools, is_multimodal)
|
||||||
else:
|
else:
|
||||||
result = self._apply_conversation_template(request, is_multimodal)
|
result = self._apply_conversation_template(request, is_multimodal)
|
||||||
@@ -1245,6 +1268,17 @@ class OpenAIServingChat(OpenAIServingBase):
|
|||||||
history_tool_calls_cnt,
|
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 = self._get_parsed_response_fields(
|
||||||
reasoning_text, tool_calls
|
reasoning_text, tool_calls
|
||||||
)
|
)
|
||||||
@@ -1253,7 +1287,7 @@ class OpenAIServingChat(OpenAIServingBase):
|
|||||||
index=idx,
|
index=idx,
|
||||||
message=ChatMessage(
|
message=ChatMessage(
|
||||||
role="assistant",
|
role="assistant",
|
||||||
content=text if text else None,
|
content=text if text else "",
|
||||||
tool_calls=tool_calls,
|
tool_calls=tool_calls,
|
||||||
reasoning_content=reasoning_text if reasoning_text else None,
|
reasoning_content=reasoning_text if reasoning_text else None,
|
||||||
),
|
),
|
||||||
@@ -1265,6 +1299,8 @@ class OpenAIServingChat(OpenAIServingBase):
|
|||||||
else None
|
else None
|
||||||
),
|
),
|
||||||
hidden_states=hidden_states,
|
hidden_states=hidden_states,
|
||||||
|
prompt_token_ids=choice_prompt_token_ids,
|
||||||
|
meta_info=choice_meta_info,
|
||||||
)
|
)
|
||||||
choices.append(choice_data)
|
choices.append(choice_data)
|
||||||
|
|
||||||
|
|||||||
@@ -259,6 +259,9 @@ class GenerateReqInput(BaseReq):
|
|||||||
# Whether to return entropy
|
# Whether to return entropy
|
||||||
return_entropy: bool = False
|
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
|
# Propagates trace context via Engine.generate/async_generate
|
||||||
external_trace_header: Optional[Dict] = None
|
external_trace_header: Optional[Dict] = None
|
||||||
received_time: Optional[float] = None
|
received_time: Optional[float] = None
|
||||||
@@ -712,6 +715,7 @@ class GenerateReqInput(BaseReq):
|
|||||||
custom_labels=self.custom_labels,
|
custom_labels=self.custom_labels,
|
||||||
return_bytes=self.return_bytes,
|
return_bytes=self.return_bytes,
|
||||||
return_entropy=self.return_entropy,
|
return_entropy=self.return_entropy,
|
||||||
|
return_prompt_token_ids=self.return_prompt_token_ids,
|
||||||
external_trace_header=self.external_trace_header,
|
external_trace_header=self.external_trace_header,
|
||||||
http_worker_ipc=self.http_worker_ipc,
|
http_worker_ipc=self.http_worker_ipc,
|
||||||
received_time=self.received_time,
|
received_time=self.received_time,
|
||||||
@@ -896,6 +900,9 @@ class EmbeddingReqInput(BaseReq):
|
|||||||
# Whether to return pooled hidden states (pre-head transformer output)
|
# Whether to return pooled hidden states (pre-head transformer output)
|
||||||
return_pooled_hidden_states: bool = False
|
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.
|
# Pre-computed delimiter indices for multi-item scoring.
|
||||||
# Batch-level: List[List[int]] (one per request). After __getitem__: List[int].
|
# Batch-level: List[List[int]] (one per request). After __getitem__: List[int].
|
||||||
multi_item_delimiter_indices: Optional[Union[List[List[int]], List[int]]] = None
|
multi_item_delimiter_indices: Optional[Union[List[List[int]], List[int]]] = None
|
||||||
@@ -1002,6 +1009,7 @@ class EmbeddingReqInput(BaseReq):
|
|||||||
is_cross_encoder_request=True,
|
is_cross_encoder_request=True,
|
||||||
http_worker_ipc=self.http_worker_ipc,
|
http_worker_ipc=self.http_worker_ipc,
|
||||||
return_pooled_hidden_states=self.return_pooled_hidden_states,
|
return_pooled_hidden_states=self.return_pooled_hidden_states,
|
||||||
|
return_prompt_token_ids=self.return_prompt_token_ids,
|
||||||
multi_item_delimiter_indices=(
|
multi_item_delimiter_indices=(
|
||||||
self.multi_item_delimiter_indices[i]
|
self.multi_item_delimiter_indices[i]
|
||||||
if self.multi_item_delimiter_indices is not None
|
if self.multi_item_delimiter_indices is not None
|
||||||
@@ -1031,6 +1039,7 @@ class EmbeddingReqInput(BaseReq):
|
|||||||
http_worker_ipc=self.http_worker_ipc,
|
http_worker_ipc=self.http_worker_ipc,
|
||||||
received_time=self.received_time,
|
received_time=self.received_time,
|
||||||
return_pooled_hidden_states=self.return_pooled_hidden_states,
|
return_pooled_hidden_states=self.return_pooled_hidden_states,
|
||||||
|
return_prompt_token_ids=self.return_prompt_token_ids,
|
||||||
multi_item_delimiter_indices=(
|
multi_item_delimiter_indices=(
|
||||||
self.multi_item_delimiter_indices[i]
|
self.multi_item_delimiter_indices[i]
|
||||||
if self.multi_item_delimiter_indices is not None
|
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)
|
input_token_ids_logprobs: List[Any] = dataclasses.field(default_factory=list)
|
||||||
output_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(
|
def _slice_streaming_output_meta_info(
|
||||||
meta_info: Dict[Any, Any],
|
meta_info: Dict[Any, Any],
|
||||||
@@ -586,6 +589,9 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
|||||||
# Tokenize the request and send it to the scheduler
|
# Tokenize the request and send it to the scheduler
|
||||||
if obj.is_single:
|
if obj.is_single:
|
||||||
tokenized_obj = await self._tokenize_one_request(obj)
|
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)
|
self._send_one_request(tokenized_obj)
|
||||||
async for response in self._wait_one_response(obj, request):
|
async for response in self._wait_one_response(obj, request):
|
||||||
yield response
|
yield response
|
||||||
@@ -1478,6 +1484,9 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
|||||||
# Set up generators for each request in the batch
|
# Set up generators for each request in the batch
|
||||||
for i in range(batch_size):
|
for i in range(batch_size):
|
||||||
tmp_obj = obj[i]
|
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))
|
generators.append(self._wait_one_response(tmp_obj, request))
|
||||||
rids.append(tmp_obj.rid)
|
rids.append(tmp_obj.rid)
|
||||||
else:
|
else:
|
||||||
@@ -1490,6 +1499,9 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
|||||||
for i in range(batch_size):
|
for i in range(batch_size):
|
||||||
tmp_obj = obj[i]
|
tmp_obj = obj[i]
|
||||||
tokenized_obj = await self._tokenize_one_request(tmp_obj)
|
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)
|
self._send_one_request(tokenized_obj)
|
||||||
generators.append(self._wait_one_response(tmp_obj, request))
|
generators.append(self._wait_one_response(tmp_obj, request))
|
||||||
rids.append(tmp_obj.rid)
|
rids.append(tmp_obj.rid)
|
||||||
@@ -1539,7 +1551,10 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
|||||||
]
|
]
|
||||||
tokenized_obj.rid = tmp_obj.regenerate_rid()
|
tokenized_obj.rid = tmp_obj.regenerate_rid()
|
||||||
self._init_req_state(tmp_obj)
|
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)
|
self._send_one_request(tokenized_obj)
|
||||||
generators.append(self._wait_one_response(tmp_obj, request))
|
generators.append(self._wait_one_response(tmp_obj, request))
|
||||||
rids.append(tmp_obj.rid)
|
rids.append(tmp_obj.rid)
|
||||||
@@ -1903,6 +1918,8 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
|||||||
}
|
}
|
||||||
else:
|
else:
|
||||||
out_dict = None
|
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):
|
elif isinstance(recv_obj, BatchTokenIDOutput):
|
||||||
is_stream = getattr(state.obj, "stream", False)
|
is_stream = getattr(state.obj, "stream", False)
|
||||||
incremental = (
|
incremental = (
|
||||||
@@ -1938,6 +1955,8 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
|||||||
}
|
}
|
||||||
else:
|
else:
|
||||||
out_dict = None
|
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:
|
else:
|
||||||
assert isinstance(recv_obj, BatchEmbeddingOutput)
|
assert isinstance(recv_obj, BatchEmbeddingOutput)
|
||||||
out_dict = {
|
out_dict = {
|
||||||
|
|||||||
@@ -179,6 +179,20 @@ class TestChatCompletionRequest(unittest.TestCase):
|
|||||||
self.assertFalse(request.stream_reasoning)
|
self.assertFalse(request.stream_reasoning)
|
||||||
self.assertEqual(request.chat_template_kwargs, {"custom_param": "value"})
|
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):
|
def test_chat_completion_reasoning_effort(self):
|
||||||
"""Test chat completion with reasoning effort"""
|
"""Test chat completion with reasoning effort"""
|
||||||
messages = [{"role": "user", "content": "Hello"}]
|
messages = [{"role": "user", "content": "Hello"}]
|
||||||
@@ -371,6 +385,28 @@ class TestModelSerialization(unittest.TestCase):
|
|||||||
self.assertIn("hidden_states", data["choices"][0])
|
self.assertIn("hidden_states", data["choices"][0])
|
||||||
self.assertEqual(data["choices"][0]["hidden_states"], [0.1, 0.2, 0.3])
|
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):
|
class TestFunctionDeferLoading(unittest.TestCase):
|
||||||
"""Test defer_loading field behavior on Function/Tool."""
|
"""Test defer_loading field behavior on Function/Tool."""
|
||||||
|
|||||||
@@ -147,6 +147,55 @@ class ServingChatTestCase(unittest.TestCase):
|
|||||||
self.assertFalse(adapted.stream)
|
self.assertFalse(adapted.stream)
|
||||||
self.assertEqual(processed, self.basic_req)
|
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):
|
def test_kimi_tool_call_keeps_default_reasoning(self):
|
||||||
self.template_manager.reasoning_config = ReasoningToggleConfig(
|
self.template_manager.reasoning_config = ReasoningToggleConfig(
|
||||||
toggle_param="thinking", default_enabled=True
|
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):
|
def test_streaming_cached_tokens_details_emits_sglext(self):
|
||||||
"""Test that streaming chat responses emit cached token details in sglext."""
|
"""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)
|
response = self.chat._build_chat_response(req, [ret_item], created=0)
|
||||||
msg = response.choices[0].message
|
msg = response.choices[0].message
|
||||||
self.assertIsNone(msg.content)
|
self.assertEqual(msg.content, "")
|
||||||
self.assertEqual(msg.reasoning_content, "42")
|
self.assertEqual(msg.reasoning_content, "42")
|
||||||
|
|
||||||
# --- poolside_v1 (Laguna-XS.2) regression tests ---
|
# --- poolside_v1 (Laguna-XS.2) regression tests ---
|
||||||
|
|||||||
@@ -539,6 +539,19 @@ class TestGenerateReqInputNormalization(CustomTestCase):
|
|||||||
self.assertEqual(item0.custom_logit_processor, "processor1")
|
self.assertEqual(item0.custom_logit_processor, "processor1")
|
||||||
self.assertEqual(item0.return_hidden_states, True)
|
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):
|
def test_regenerate_rid(self):
|
||||||
"""Test the regenerate_rid method."""
|
"""Test the regenerate_rid method."""
|
||||||
req = GenerateReqInput(text="Hello")
|
req = GenerateReqInput(text="Hello")
|
||||||
|
|||||||
Reference in New Issue
Block a user