Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
11 changes: 9 additions & 2 deletions python/sglang/srt/entrypoints/openai/protocol.py
Original file line number Diff line number Diff line change
Expand Up @@ -81,7 +81,6 @@ class LogProbs(BaseModel):
text_offset: List[int] = Field(default_factory=list)
token_logprobs: List[Optional[float]] = Field(default_factory=list)
tokens: List[str] = Field(default_factory=list)
token_ids: List[int] = Field(default_factory=list)
top_logprobs: List[Optional[Dict[str, float]]] = Field(default_factory=list)


Expand All @@ -93,7 +92,6 @@ class TopLogprob(BaseModel):

class ChatCompletionTokenLogprob(BaseModel):
token: str
token_id: int
bytes: List[int]
logprob: float
top_logprobs: List[TopLogprob]
Expand Down Expand Up @@ -568,6 +566,7 @@ class ChatCompletionRequest(BaseModel):
return_routed_experts: bool = False
return_cached_tokens_details: bool = False
return_prompt_token_ids: bool = False
return_meta_info: bool = False
reasoning_effort: Optional[Literal["low", "medium", "high"]] = Field(
default="medium",
description="Constrains effort on reasoning for reasoning models. "
Expand Down Expand Up @@ -603,6 +602,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)
Expand Down Expand Up @@ -806,6 +810,7 @@ class ChatCompletionResponseChoice(BaseModel):
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):
Expand All @@ -814,6 +819,8 @@ def _serialize(self, handler):
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


Expand Down
29 changes: 22 additions & 7 deletions python/sglang/srt/entrypoints/openai/serving_chat.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@
import copy
import json
import logging
from sglang.srt.entrypoints.openai.protocol import ChatCompletionMessageParam
import time
import uuid
from typing import TYPE_CHECKING, Any, AsyncGenerator, Dict, List, Optional, Union
Expand Down Expand Up @@ -261,7 +262,6 @@ def _convert_to_internal_request(

# 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,
Expand Down Expand Up @@ -356,8 +356,19 @@ def _process_messages(
)
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=self.tokenizer_manager.tokenizer.decode(request.input_ids),
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)
Expand Down Expand Up @@ -975,11 +986,15 @@ def _build_chat_response(
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 consistancy.
choice_data = ChatCompletionResponseChoice(
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,
),
Expand All @@ -992,6 +1007,7 @@ def _build_chat_response(
),
hidden_states=hidden_states,
prompt_token_ids=choice_prompt_token_ids,
meta_info=choice_meta_info,
)
choices.append(choice_data)

Expand Down Expand Up @@ -1023,8 +1039,8 @@ def _process_logprobs_tokens(
"""
token_logprobs = []

for token_idx, (token, token_id, logprob) in enumerate(
zip(logprobs.tokens, logprobs.token_ids, logprobs.token_logprobs)
for token_idx, (token, logprob) in enumerate(
zip(logprobs.tokens, logprobs.token_logprobs)
):
token_bytes = list(token.encode("utf-8"))
top_logprobs = []
Expand All @@ -1046,7 +1062,6 @@ def _process_logprobs_tokens(
token_logprobs.append(
ChatCompletionTokenLogprob(
token=token,
token_id=token_id,
bytes=token_bytes,
logprob=logprob,
top_logprobs=top_logprobs,
Expand Down
3 changes: 1 addition & 2 deletions python/sglang/srt/entrypoints/openai/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,9 +20,8 @@ def to_openai_style_logprobs(
ret_logprobs = LogProbs()

def append_token_logprobs(token_logprobs):
for logprob, token_id, token_text in token_logprobs:
for logprob, _, token_text in token_logprobs:
ret_logprobs.tokens.append(token_text)
ret_logprobs.token_ids.append(token_id)
ret_logprobs.token_logprobs.append(logprob)

# Not supported yet
Expand Down
6 changes: 5 additions & 1 deletion python/sglang/srt/managers/tokenizer_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -1813,7 +1813,11 @@ def detokenize_logprob_tokens(
]
else:
assert self.tokenizer is not None
token_texts = self.tokenizer.batch_decode(token_logprobs_idx)
# Wrap each token ID in its own list for batch_decode to decode them separately
# batch_decode([1, 2, 3]) concatenates tokens, batch_decode([[1], [2], [3]]) decodes separately
token_texts = self.tokenizer.batch_decode(
[[idx] for idx in token_logprobs_idx]
)
return list(zip(token_logprobs_val, token_logprobs_idx, token_texts))

def detokenize_top_logprobs_tokens(
Expand Down
Loading
Loading