fix(openai): recover logprobs token bytes from token_id (UTF-8 fragments) (#38604)

This commit is contained in:
Mind Lab
2026-09-18 20:30:04 -07:00
committed by GitHub
parent f1fbbd17bb
commit 111b905bd1
3 changed files with 333 additions and 155 deletions
@@ -7,6 +7,7 @@ import math
import time import time
import uuid import uuid
from collections import OrderedDict from collections import OrderedDict
from collections.abc import AsyncGenerator
from enum import Enum from enum import Enum
from http import HTTPStatus from http import HTTPStatus
from typing import TYPE_CHECKING, Any, AsyncGenerator, Dict, List, Optional, Union from typing import TYPE_CHECKING, Any, AsyncGenerator, Dict, List, Optional, Union
@@ -61,7 +62,6 @@ from sglang.srt.entrypoints.openai.protocol import (
DeltaMessage, DeltaMessage,
ErrorResponse, ErrorResponse,
FunctionResponse, FunctionResponse,
LogProbs,
MessageProcessingResult, MessageProcessingResult,
PromptTokensDetails, PromptTokensDetails,
ResponseParserProtocol, ResponseParserProtocol,
@@ -84,7 +84,7 @@ from sglang.srt.entrypoints.openai.utils import (
process_spec_tokens_details_from_ret, process_spec_tokens_details_from_ret,
should_include_usage, should_include_usage,
spec_tokens_details_from_meta_info, spec_tokens_details_from_meta_info,
to_openai_style_logprobs, token_id_to_bytes,
) )
from sglang.srt.entrypoints.request_headers import apply_header_overrides from sglang.srt.entrypoints.request_headers import apply_header_overrides
from sglang.srt.environ import envs from sglang.srt.environ import envs
@@ -144,7 +144,7 @@ def normalize_tool_content(role: str, content):
return content return content
def parse_tool_call_arguments(arguments: str) -> Dict[str, Any]: def parse_tool_call_arguments(arguments: str) -> dict[str, Any]:
"""Parse OpenAI tool call arguments for chat templates.""" """Parse OpenAI tool call arguments for chat templates."""
try: try:
parsed_arguments = orjson.loads(arguments) parsed_arguments = orjson.loads(arguments)
@@ -162,7 +162,7 @@ def parse_tool_call_arguments(arguments: str) -> Dict[str, Any]:
def normalize_assistant_tool_call_arguments( def normalize_assistant_tool_call_arguments(
message: Dict[str, Any], *, strict: bool = True message: dict[str, Any], *, strict: bool = True
) -> None: ) -> None:
"""Normalize assistant history tool call arguments in-place.""" """Normalize assistant history tool call arguments in-place."""
if message.get("role") != "assistant" or not isinstance( if message.get("role") != "assistant" or not isinstance(
@@ -229,7 +229,7 @@ def neutralize_kimi_k3_image_placeholder_value(value: Any) -> Any:
return value return value
def _extract_video_question(request: ChatCompletionRequest) -> Optional[str]: def _extract_video_question(request: ChatCompletionRequest) -> str | None:
"""Return text paired with a video in the last user turn.""" """Return text paired with a video in the last user turn."""
for message in reversed(request.messages or []): for message in reversed(request.messages or []):
if not isinstance(message, ChatCompletionMessageUserParam): if not isinstance(message, ChatCompletionMessageUserParam):
@@ -250,7 +250,7 @@ def _extract_video_question(request: ChatCompletionRequest) -> Optional[str]:
return None return None
def _build_video_config(request: ChatCompletionRequest) -> Optional[Dict[str, Any]]: def _build_video_config(request: ChatCompletionRequest) -> dict[str, Any] | None:
"""Build request-scoped video processor config without model-specific fields.""" """Build request-scoped video processor config without model-specific fields."""
config = dict(request.video_config or {}) config = dict(request.video_config or {})
question = _extract_video_question(request) question = _extract_video_question(request)
@@ -339,7 +339,7 @@ class OpenAIServingChat(OpenAIServingBase):
# Resolve the env-configured Inkling effort default once: the env var is # Resolve the env-configured Inkling effort default once: the env var is
# frozen for the server's lifetime, and a misconfigured value should # frozen for the server's lifetime, and a misconfigured value should
# fail at boot, not 400 every request. # fail at boot, not 400 every request.
self._inkling_default_reasoning_effort: Optional[float] = ( self._inkling_default_reasoning_effort: float | None = (
self._get_inkling_default_reasoning_effort() self._get_inkling_default_reasoning_effort()
if self.chat_encoding_spec == "inkling" if self.chat_encoding_spec == "inkling"
else None else None
@@ -353,7 +353,7 @@ class OpenAIServingChat(OpenAIServingBase):
) )
# Per-request response parser for custom decoding (set by _encode_messages) # Per-request response parser for custom decoding (set by _encode_messages)
self._response_parser: Optional[ResponseParserProtocol] = None self._response_parser: ResponseParserProtocol | None = None
# Probe whether ``encode("")`` returns specials. If it does, we must # Probe whether ``encode("")`` returns specials. If it does, we must
# keep ``add_special_tokens=False`` at the chat-template encode site # keep ``add_special_tokens=False`` at the chat-template encode site
@@ -405,9 +405,9 @@ class OpenAIServingChat(OpenAIServingBase):
def _handle_last_assistant_message( def _handle_last_assistant_message(
self, self,
messages: List[Dict[str, Any]], messages: list[dict[str, Any]],
request: ChatCompletionRequest, request: ChatCompletionRequest,
) -> tuple[List[Dict[str, Any]], Optional[str]]: ) -> tuple[list[dict[str, Any]], str | None]:
""" """
Handle continue_final_message feature: separate final assistant message. Handle continue_final_message feature: separate final assistant message.
@@ -442,8 +442,8 @@ class OpenAIServingChat(OpenAIServingBase):
return messages, assistant_prefix return messages, assistant_prefix
def _append_assistant_prefix_to_prompt_ids( def _append_assistant_prefix_to_prompt_ids(
self, prompt_ids: List[int], assistant_prefix: str self, prompt_ids: list[int], assistant_prefix: str
) -> List[int]: ) -> list[int]:
""" """
Append assistant prefix to prompt_ids. Append assistant prefix to prompt_ids.
@@ -459,7 +459,7 @@ class OpenAIServingChat(OpenAIServingBase):
encoded = encoded[1:] encoded = encoded[1:]
return prompt_ids + encoded return prompt_ids + encoded
def _resolve_chat_encoding_spec(self) -> Optional[str]: def _resolve_chat_encoding_spec(self) -> str | None:
"""Determine which chat encoding spec to use. """Determine which chat encoding spec to use.
Override in subclass to add custom encoding specs. Override in subclass to add custom encoding specs.
@@ -473,7 +473,7 @@ class OpenAIServingChat(OpenAIServingBase):
def _request_id_prefix(self) -> str: def _request_id_prefix(self) -> str:
return "chatcmpl-" return "chatcmpl-"
def _effective_tools(self, request: ChatCompletionRequest) -> List[Tool]: def _effective_tools(self, request: ChatCompletionRequest) -> list[Tool]:
tools = list(request.tools or []) tools = list(request.tools or [])
for message in request.messages: for message in request.messages:
if ( if (
@@ -486,9 +486,9 @@ class OpenAIServingChat(OpenAIServingBase):
def _prepare_kimi_k3_messages( def _prepare_kimi_k3_messages(
self, self,
messages: List[Dict[str, Any]], messages: list[dict[str, Any]],
request: ChatCompletionRequest, request: ChatCompletionRequest,
) -> tuple[List[Dict[str, Any]], int, Optional[str]]: ) -> tuple[list[dict[str, Any]], int, str | None]:
image_count = 0 image_count = 0
for index, message in enumerate(messages): for index, message in enumerate(messages):
content = message.get("content") content = message.get("content")
@@ -560,11 +560,11 @@ class OpenAIServingChat(OpenAIServingBase):
def _encode_messages( def _encode_messages(
self, self,
messages: List[Dict[str, Any]], messages: list[dict[str, Any]],
request: ChatCompletionRequest, request: ChatCompletionRequest,
thinking_mode: ThinkingMode, thinking_mode: ThinkingMode,
tools: Optional[List[Dict]] = None, tools: list[dict] | None = None,
) -> Optional[List[int]]: ) -> list[int] | None:
"""Encode messages for custom chat_encoding_spec values. """Encode messages for custom chat_encoding_spec values.
Returns prompt_ids if handled, None to use default encoding. Returns prompt_ids if handled, None to use default encoding.
@@ -676,9 +676,9 @@ class OpenAIServingChat(OpenAIServingBase):
@staticmethod @staticmethod
def _pop_inkling_assistant_prefix( def _pop_inkling_assistant_prefix(
messages: List[Dict[str, Any]], messages: list[dict[str, Any]],
request: ChatCompletionRequest, request: ChatCompletionRequest,
) -> Optional[str]: ) -> str | None:
"""Extract the trailing assistant text for ``continue_final_message``. """Extract the trailing assistant text for ``continue_final_message``.
Only a plain-string assistant message with no tool calls and no Only a plain-string assistant message with no tool calls and no
@@ -700,8 +700,8 @@ class OpenAIServingChat(OpenAIServingBase):
@staticmethod @staticmethod
def _parse_inkling_reasoning_effort( def _parse_inkling_reasoning_effort(
value: Optional[Union[str, float]], value: str | float | None,
) -> Optional[float]: ) -> float | None:
"""Convert an OpenAI-style reasoning_effort to an Inkling float.""" """Convert an OpenAI-style reasoning_effort to an Inkling float."""
if value is None: if value is None:
return None return None
@@ -766,15 +766,15 @@ class OpenAIServingChat(OpenAIServingBase):
) )
return parsed return parsed
def _decode_response(self, ret_item: Dict[str, Any]) -> Union[str, ErrorResponse]: def _decode_response(self, ret_item: dict[str, Any]) -> str | ErrorResponse:
"""Extract text from response.""" """Extract text from response."""
return ret_item["text"] return ret_item["text"]
def _get_parsed_response_fields( def _get_parsed_response_fields(
self, self,
reasoning_text: Optional[str], reasoning_text: str | None,
tool_calls: Optional[List[Dict]], tool_calls: list[dict] | None,
) -> tuple[Optional[str], Optional[List[Dict]]]: ) -> tuple[str | None, list[dict] | None]:
"""Post-process reasoning and tool_calls before building response.""" """Post-process reasoning and tool_calls before building response."""
return reasoning_text, tool_calls return reasoning_text, tool_calls
@@ -787,15 +787,15 @@ class OpenAIServingChat(OpenAIServingBase):
return request.return_output_ids_in_sglext or get_serving().return_output_ids return request.return_output_ids_in_sglext or get_serving().return_output_ids
def _continuous_usage_cached_details( def _continuous_usage_cached_details(
self, content: Dict[str, Any] self, content: dict[str, Any]
) -> Optional[PromptTokensDetails]: ) -> PromptTokensDetails | None:
if not get_serving().enable_cache_report: if not get_serving().enable_cache_report:
return None return None
return UsageProcessor._details_if_cached( return UsageProcessor._details_if_cached(
content["meta_info"].get("cached_tokens", 0) content["meta_info"].get("cached_tokens", 0)
) )
def _reported_prompt_tokens(self, meta_info: Dict[str, Any]) -> int: def _reported_prompt_tokens(self, meta_info: dict[str, Any]) -> int:
prompt_tokens = meta_info.get("prompt_tokens", 0) prompt_tokens = meta_info.get("prompt_tokens", 0)
if self.chat_encoding_spec == "kimi_k3": if self.chat_encoding_spec == "kimi_k3":
# K3's three-token assistant generation stub is model input, but the # K3's three-token assistant generation stub is model input, but the
@@ -805,8 +805,8 @@ class OpenAIServingChat(OpenAIServingBase):
@staticmethod @staticmethod
def _sort_tool_message_run( def _sort_tool_message_run(
run: List[Dict[str, Any]], tool_calls: List[Dict[str, Any]] run: list[dict[str, Any]], tool_calls: list[dict[str, Any]]
) -> List[Dict[str, Any]]: ) -> list[dict[str, Any]]:
"""Order a tool-message run by tool_call position. """Order a tool-message run by tool_call position.
Templates that associate results by tool_call_id render the run in Templates that associate results by tool_call_id render the run in
@@ -842,8 +842,8 @@ class OpenAIServingChat(OpenAIServingBase):
@classmethod @classmethod
def _canonicalize_tool_message_order( def _canonicalize_tool_message_order(
cls, messages: List[Dict[str, Any]] cls, messages: list[dict[str, Any]]
) -> List[Dict[str, Any]]: ) -> list[dict[str, Any]]:
canonical = [] canonical = []
index = 0 index = 0
while index < len(messages): while index < len(messages):
@@ -865,19 +865,19 @@ class OpenAIServingChat(OpenAIServingBase):
async def _generate_stream_content( async def _generate_stream_content(
self, self,
content: Dict[str, Any], content: dict[str, Any],
index: int, index: int,
request: ChatCompletionRequest, request: ChatCompletionRequest,
stream_offsets: Dict[int, int], stream_offsets: dict[int, int],
reasoning_parser_dict: Dict, reasoning_parser_dict: dict,
parser_dict: Dict, parser_dict: dict,
has_tool_calls: Dict[int, bool], has_tool_calls: dict[int, bool],
choice_logprobs: Optional[Dict], choice_logprobs: dict | None,
finish_reason_type: Optional[str], finish_reason_type: str | None,
continuous_usage_stats: bool, continuous_usage_stats: bool,
prompt_tokens: Dict[int, int], prompt_tokens: dict[int, int],
reasoning_tokens: Dict[int, int], reasoning_tokens: dict[int, int],
completion_tokens: Dict[int, int], completion_tokens: dict[int, int],
) -> AsyncGenerator[str, None]: ) -> AsyncGenerator[str, None]:
"""Generate SSE chunks for streaming content.""" """Generate SSE chunks for streaming content."""
offset = stream_offsets.get(index, 0) offset = stream_offsets.get(index, 0)
@@ -1008,7 +1008,7 @@ class OpenAIServingChat(OpenAIServingBase):
and self.tool_call_parser and self.tool_call_parser
) )
def _validate_request(self, request: ChatCompletionRequest) -> Optional[str]: def _validate_request(self, request: ChatCompletionRequest) -> str | None:
"""Validate that the input is valid.""" """Validate that the input is valid."""
if not request.messages: if not request.messages:
return "Messages cannot be empty." return "Messages cannot be empty."
@@ -1087,7 +1087,7 @@ class OpenAIServingChat(OpenAIServingBase):
return None return None
def _validate_media_content(self, request: ChatCompletionRequest) -> Optional[str]: def _validate_media_content(self, request: ChatCompletionRequest) -> str | None:
if self.tokenizer_manager.model_config.is_multimodal: if self.tokenizer_manager.model_config.is_multimodal:
return None return None
@@ -1433,7 +1433,7 @@ class OpenAIServingChat(OpenAIServingBase):
def _apply_jinja_template( def _apply_jinja_template(
self, self,
request: ChatCompletionRequest, request: ChatCompletionRequest,
tools: Optional[List[Dict]], tools: list[dict] | None,
is_multimodal: bool, is_multimodal: bool,
) -> MessageProcessingResult: ) -> MessageProcessingResult:
"""Apply Jinja chat template""" """Apply Jinja chat template"""
@@ -1852,7 +1852,7 @@ class OpenAIServingChat(OpenAIServingBase):
adapted_request: GenerateReqInput, adapted_request: GenerateReqInput,
request: ChatCompletionRequest, request: ChatCompletionRequest,
raw_request: Request, raw_request: Request,
) -> Union[StreamingResponse, ErrorResponse]: ) -> StreamingResponse | ErrorResponse:
"""Handle streaming chat completion request""" """Handle streaming chat completion request"""
generator = self._generate_chat_stream(adapted_request, request, raw_request) generator = self._generate_chat_stream(adapted_request, request, raw_request)
@@ -1905,8 +1905,8 @@ class OpenAIServingChat(OpenAIServingBase):
image_tokens = {} image_tokens = {}
audio_tokens = {} audio_tokens = {}
video_tokens = {} video_tokens = {}
input_ids: Optional[List[int]] = None input_ids: list[int] | None = None
output_ids: Dict[int, List[int]] = {} output_ids: dict[int, list[int]] = {}
stream_started = False stream_started = False
error_aborted = False error_aborted = False
@@ -2211,7 +2211,7 @@ class OpenAIServingChat(OpenAIServingBase):
adapted_request: GenerateReqInput, adapted_request: GenerateReqInput,
request: ChatCompletionRequest, request: ChatCompletionRequest,
raw_request: Request, raw_request: Request,
) -> Union[ChatCompletionResponse, ErrorResponse, ORJSONResponse]: ) -> ChatCompletionResponse | ErrorResponse | ORJSONResponse:
"""Handle non-streaming chat completion request""" """Handle non-streaming chat completion request"""
try: try:
ret = await self.tokenizer_manager.generate_request( ret = await self.tokenizer_manager.generate_request(
@@ -2234,9 +2234,9 @@ class OpenAIServingChat(OpenAIServingBase):
def _build_chat_response( def _build_chat_response(
self, self,
request: ChatCompletionRequest, request: ChatCompletionRequest,
ret: List[Dict[str, Any]], ret: list[dict[str, Any]],
created: int, created: int,
) -> Union[ChatCompletionResponse, ORJSONResponse]: ) -> ChatCompletionResponse | ORJSONResponse:
"""Build chat completion response from generation results""" """Build chat completion response from generation results"""
if self.chat_encoding_spec == "kimi_k3": if self.chat_encoding_spec == "kimi_k3":
ret = [ ret = [
@@ -2424,40 +2424,79 @@ class OpenAIServingChat(OpenAIServingBase):
sglext=response_sglext, sglext=response_sglext,
) )
def _process_logprobs_tokens( def _process_response_logprobs(self, ret_item: dict[str, Any]) -> ChoiceLogprobs:
self, logprobs: LogProbs, use_token_index: bool = False """Process logprobs for non-streaming response"""
) -> List[ChatCompletionTokenLogprob]: output_token_logprobs = ret_item["meta_info"]["output_token_logprobs"]
"""Common helper to process logprobs tokens for both streaming and non-streaming output_top_logprobs = ret_item["meta_info"].get("output_top_logprobs", None)
token_logprobs = self._build_token_logprobs_from_raw(
output_token_logprobs, output_top_logprobs, use_token_index=True
)
return ChoiceLogprobs(content=token_logprobs)
Args: def _build_token_logprobs_from_raw(
logprobs: LogProbs data from model self,
use_token_index: True for non-streaming (use token_idx), False for streaming (use index 0) output_token_logprobs: list[Any],
output_top_logprobs: list[Any] | None,
use_token_index: bool = False,
) -> list[ChatCompletionTokenLogprob]:
"""Build OpenAI ChatCompletionTokenLogprob from the engine's raw
``(logprob, token_id, token_text)`` triples.
The engine always keeps the real token id in the triple; the token text
is a *detokenized display* string that loses fragmentary byte-level
tokens (e.g. the four single-byte BPE pieces of U+20BB7 decode to
U+FFFD). Recovering the OpenAI ``bytes`` field from the display text
therefore corrupts fragments; we recover the true bytes from the token
id through the GPT-2 byte decoder instead.
Only byte-level BPE tokenizers (GPT-2 family) use the byte decoder;
SentencePiece tokenizers (Mistral, Gemma) fall back to the display
text, so multi-byte characters like é are NOT corrupted to [233].
""" """
token_logprobs = [] tokenizer = self.tokenizer_manager.tokenizer
from sglang.srt.entrypoints.openai.utils import _is_byte_level_tokenizer
is_byte_level = _is_byte_level_tokenizer(tokenizer)
token_logprobs: list[ChatCompletionTokenLogprob] = []
for token_idx, item in enumerate(output_token_logprobs):
# item = (logprob, token_id, token_text)
logprob, token_id, token_text = item
if is_byte_level:
token_bytes = token_id_to_bytes(tokenizer, token_id)
else:
token_bytes = None
if token_bytes is None:
token_bytes = list((token_text or "").encode("utf-8"))
top_logprobs: list[TopLogprob] = []
if output_top_logprobs:
# - Non-streaming (use_token_index=True): output_top_logprobs is
# the full per-position list; take the row for this token.
# - Streaming (use_token_index=False): rows are pre-sliced so the
# current chunk holds exactly one row at index 0.
top_row_idx = token_idx if use_token_index else 0
if top_row_idx < len(output_top_logprobs):
top_row = output_top_logprobs[top_row_idx]
if top_row is not None:
for top_logprob, top_id, top_text in top_row:
if is_byte_level:
top_bytes = token_id_to_bytes(tokenizer, top_id)
else:
top_bytes = None
if top_bytes is None:
top_bytes = list((top_text or "").encode("utf-8"))
top_logprobs.append(
TopLogprob(
token=top_text or "",
bytes=top_bytes,
logprob=top_logprob,
)
)
for token_idx, (token, logprob) in enumerate(
zip(logprobs.tokens, logprobs.token_logprobs)
):
token_bytes = list(token.encode("utf-8"))
top_logprobs = []
if logprobs.top_logprobs:
# - Non-streaming (use_token_index=True): uses token_idx for full data
# - Streaming (use_token_index=False): uses index 0 for pre-sliced data
top_logprobs_idx = token_idx if use_token_index else 0
for top_token, top_logprob in logprobs.top_logprobs[
top_logprobs_idx
].items():
top_token_bytes = list(top_token.encode("utf-8"))
top_logprobs.append(
TopLogprob(
token=top_token,
bytes=top_token_bytes,
logprob=top_logprob,
)
)
token_logprobs.append( token_logprobs.append(
ChatCompletionTokenLogprob( ChatCompletionTokenLogprob(
token=token, token=token_text or "",
bytes=token_bytes, bytes=token_bytes,
logprob=logprob, logprob=logprob,
top_logprobs=top_logprobs, top_logprobs=top_logprobs,
@@ -2466,16 +2505,6 @@ class OpenAIServingChat(OpenAIServingBase):
return token_logprobs return token_logprobs
def _process_response_logprobs(self, ret_item: Dict[str, Any]) -> ChoiceLogprobs:
"""Process logprobs for non-streaming response"""
logprobs = to_openai_style_logprobs(
output_token_logprobs=ret_item["meta_info"]["output_token_logprobs"],
output_top_logprobs=ret_item["meta_info"].get("output_top_logprobs", None),
)
token_logprobs = self._process_logprobs_tokens(logprobs, use_token_index=True)
return ChoiceLogprobs(content=token_logprobs)
def _process_tool_call_id( def _process_tool_call_id(
self, self,
call_item: ToolCallItem, call_item: ToolCallItem,
@@ -2500,9 +2529,9 @@ class OpenAIServingChat(OpenAIServingBase):
def _process_tool_calls( def _process_tool_calls(
self, self,
text: str, text: str,
tools: List[Any], tools: list[Any],
finish_reason: Dict[str, Any], finish_reason: dict[str, Any],
tool_choice: Optional[Union[str, ToolChoice]] = None, tool_choice: str | ToolChoice | None = None,
history_tool_calls_cnt: int = 0, history_tool_calls_cnt: int = 0,
) -> ToolCallProcessingResult: ) -> ToolCallProcessingResult:
"""Process tool calls in the response""" """Process tool calls in the response"""
@@ -2626,7 +2655,7 @@ class OpenAIServingChat(OpenAIServingBase):
def _process_streaming_logprobs( def _process_streaming_logprobs(
self, self,
content: Dict[str, Any], content: dict[str, Any],
n_prev_token: int, n_prev_token: int,
total_output_logprobs: int, total_output_logprobs: int,
) -> ChoiceLogprobs: ) -> ChoiceLogprobs:
@@ -2640,23 +2669,20 @@ class OpenAIServingChat(OpenAIServingBase):
output_top_logprobs = output_top_logprobs[ output_top_logprobs = output_top_logprobs[
n_prev_token:total_output_logprobs n_prev_token:total_output_logprobs
] ]
logprobs = to_openai_style_logprobs( token_logprobs = self._build_token_logprobs_from_raw(
output_token_logprobs=output_token_logprobs, output_token_logprobs, output_top_logprobs, use_token_index=False
output_top_logprobs=output_top_logprobs,
) )
token_logprobs = self._process_logprobs_tokens(logprobs, use_token_index=False)
return ChoiceLogprobs(content=token_logprobs) return ChoiceLogprobs(content=token_logprobs)
def _process_reasoning_stream( def _process_reasoning_stream(
self, self,
index: int, index: int,
delta: str, delta: str,
reasoning_parser_dict: Dict[int, ReasoningParser], reasoning_parser_dict: dict[int, ReasoningParser],
content: Dict[str, Any], content: dict[str, Any],
request: ChatCompletionRequest, request: ChatCompletionRequest,
finish_reason_type: Optional[str] = None, finish_reason_type: str | None = None,
) -> tuple[Optional[str], str]: ) -> tuple[str | None, str]:
"""Process reasoning content in streaming response""" """Process reasoning content in streaming response"""
if index not in reasoning_parser_dict: if index not in reasoning_parser_dict:
is_force_reasoning = ( is_force_reasoning = (
@@ -2753,12 +2779,12 @@ class OpenAIServingChat(OpenAIServingBase):
f"{reasoning_text}\n{d.think_end_token}" f"{reasoning_text}\n{d.think_end_token}"
) )
def _reasoning_default_mode(self) -> Optional[str]: def _reasoning_default_mode(self) -> str | None:
if self._reasoning_detector is None: if self._reasoning_detector is None:
return None return None
return self._reasoning_detector.reasoning_default return self._reasoning_detector.reasoning_default
def _get_reasoning_toggle_param(self) -> Optional[str]: def _get_reasoning_toggle_param(self) -> str | None:
"""Resolve the chat-template kwarg that toggles reasoning, if any.""" """Resolve the chat-template kwarg that toggles reasoning, if any."""
config = self.template_manager.reasoning_config config = self.template_manager.reasoning_config
if config is not None: if config is not None:
@@ -2932,10 +2958,10 @@ class OpenAIServingChat(OpenAIServingBase):
self, self,
index: int, index: int,
delta: str, delta: str,
parser_dict: Dict[int, FunctionCallParser], parser_dict: dict[int, FunctionCallParser],
content: Dict[str, Any], content: dict[str, Any],
request: ChatCompletionRequest, request: ChatCompletionRequest,
has_tool_calls: Dict[int, bool], has_tool_calls: dict[int, bool],
continuous_usage_stats: bool = False, continuous_usage_stats: bool = False,
flush: bool = False, flush: bool = False,
): ):
@@ -3073,11 +3099,11 @@ class OpenAIServingChat(OpenAIServingBase):
def _check_for_unstreamed_tool_args( def _check_for_unstreamed_tool_args(
self, self,
parser: Union[FunctionCallParser, JsonArrayParser], parser: FunctionCallParser | JsonArrayParser,
content: Dict[str, Any], content: dict[str, Any],
request: ChatCompletionRequest, request: ChatCompletionRequest,
index: int, index: int,
) -> Optional[str]: ) -> str | None:
""" """
Check for any remaining tool call arguments that need to be streamed Check for any remaining tool call arguments that need to be streamed
when generation finishes. This ensures tool calls are properly completed when generation finishes. This ensures tool calls are properly completed
@@ -2,8 +2,9 @@ from __future__ import annotations
import logging import logging
import time import time
from collections.abc import AsyncGenerator
from http import HTTPStatus from http import HTTPStatus
from typing import TYPE_CHECKING, Any, AsyncGenerator, Dict, List, Optional, Union from typing import TYPE_CHECKING, Any
from fastapi import Request from fastapi import Request
from fastapi.responses import ORJSONResponse, StreamingResponse from fastapi.responses import ORJSONResponse, StreamingResponse
@@ -59,7 +60,7 @@ class OpenAIServingCompletion(OpenAIServingBase):
def _request_id_prefix(self) -> str: def _request_id_prefix(self) -> str:
return "cmpl-" return "cmpl-"
def _validate_request(self, request: CompletionRequest) -> Optional[str]: def _validate_request(self, request: CompletionRequest) -> str | None:
"""Validate that the input is valid.""" """Validate that the input is valid."""
prompt = request.prompt prompt = request.prompt
if not prompt or (isinstance(prompt, list) and all(not p for p in prompt)): if not prompt or (isinstance(prompt, list) and all(not p for p in prompt)):
@@ -143,7 +144,7 @@ class OpenAIServingCompletion(OpenAIServingBase):
return adapted_request, request return adapted_request, request
def _build_sampling_params(self, request: CompletionRequest) -> Dict[str, Any]: def _build_sampling_params(self, request: CompletionRequest) -> dict[str, Any]:
"""Build sampling parameters for the request""" """Build sampling parameters for the request"""
# Start with common parameters # Start with common parameters
sampling_params = { sampling_params = {
@@ -196,7 +197,7 @@ class OpenAIServingCompletion(OpenAIServingBase):
adapted_request: GenerateReqInput, adapted_request: GenerateReqInput,
request: CompletionRequest, request: CompletionRequest,
raw_request: Request, raw_request: Request,
) -> Union[StreamingResponse, ErrorResponse]: ) -> StreamingResponse | ErrorResponse:
"""Handle streaming completion request""" """Handle streaming completion request"""
generator = self._generate_completion_stream( generator = self._generate_completion_stream(
adapted_request, request, raw_request adapted_request, request, raw_request
@@ -326,6 +327,7 @@ class OpenAIServingCompletion(OpenAIServingBase):
input_top_logprobs=input_top_logprobs, input_top_logprobs=input_top_logprobs,
output_token_logprobs=output_token_logprobs, output_token_logprobs=output_token_logprobs,
output_top_logprobs=output_top_logprobs, output_top_logprobs=output_top_logprobs,
tokenizer=self.tokenizer_manager.tokenizer,
) )
n_prev_tokens[index] = total_output_logprobs n_prev_tokens[index] = total_output_logprobs
@@ -504,7 +506,7 @@ class OpenAIServingCompletion(OpenAIServingBase):
adapted_request: GenerateReqInput, adapted_request: GenerateReqInput,
request: CompletionRequest, request: CompletionRequest,
raw_request: Request, raw_request: Request,
) -> Union[CompletionResponse, ErrorResponse, ORJSONResponse]: ) -> CompletionResponse | ErrorResponse | ORJSONResponse:
"""Handle non-streaming completion request""" """Handle non-streaming completion request"""
try: try:
generator = self.tokenizer_manager.generate_request( generator = self.tokenizer_manager.generate_request(
@@ -528,7 +530,7 @@ class OpenAIServingCompletion(OpenAIServingBase):
def _build_completion_response( def _build_completion_response(
self, self,
request: CompletionRequest, request: CompletionRequest,
ret: List[Dict[str, Any]], ret: list[dict[str, Any]],
created: int, created: int,
) -> CompletionResponse: ) -> CompletionResponse:
"""Build completion response from generation results""" """Build completion response from generation results"""
@@ -596,6 +598,7 @@ class OpenAIServingCompletion(OpenAIServingBase):
output_top_logprobs=ret_item["meta_info"].get( output_top_logprobs=ret_item["meta_info"].get(
"output_top_logprobs", [] "output_top_logprobs", []
), ),
tokenizer=self.tokenizer_manager.tokenizer,
) )
# Handle hidden states # Handle hidden states
@@ -665,7 +668,7 @@ class OpenAIServingCompletion(OpenAIServingBase):
) )
return "" return ""
def _prepare_echo_prompts(self, request: CompletionRequest) -> List[str]: def _prepare_echo_prompts(self, request: CompletionRequest) -> list[str]:
"""Prepare echo prompts for non-streaming response""" """Prepare echo prompts for non-streaming response"""
# TODO: handle the case prompt is token ids # TODO: handle the case prompt is token ids
if isinstance(request.prompt, list) and isinstance(request.prompt[0], str): if isinstance(request.prompt, list) and isinstance(request.prompt[0], str):
+184 -35
View File
@@ -1,5 +1,5 @@
import logging import logging
from typing import Any, Dict, List, Literal, Optional, Union from typing import Any, Literal
import torch import torch
@@ -14,17 +14,136 @@ from sglang.srt.entrypoints.openai.protocol import (
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
# GPT-2 style byte-level BPE decoder table (char -> raw byte). Byte-level BPE
# vocab tokens are stored as a printable-char mapping of the raw UTF-8 bytes
# (see openai/gpt-2 bytes_to_unicode); converting a token id back to its raw
# bytes must go through this table, NOT through `token.encode()` on the
# detokenized display string (that loses fragmentary bytes as U+FFFD).
_BYTE_DECODER: dict[str, int] = {}
# Tokenizer-level cache: once verified, we know whether *all* tokens from a
# given tokenizer can safely use the byte decoder. Avoids per-token checks.
_BYTE_LEVEL_TOKENIZERS: set = set()
def _build_byte_decoder() -> dict[str, int]:
bs = (
list(range(ord("!"), ord("~") + 1))
+ list(range(ord(chr(0xA1)), ord(chr(0xAC)) + 1))
+ list(range(ord(chr(0xAE)), ord(chr(0xFF)) + 1))
)
cs = bs[:]
n = 0
for b in range(2**8):
if b not in bs:
bs.append(b)
cs.append(2**8 + n)
n += 1
cs = [chr(c) for c in cs]
return dict(zip(cs, bs))
def _is_byte_level_tokenizer(tokenizer) -> bool:
"""Heuristically determine whether *tokenizer* uses GPT-2 byte-level BPE.
Only GPT-2 family (GPT2Tokenizer, Llama, Qwen, etc.) store vocabulary
tokens as the ``bytes_to_unicode`` printable-char mapping. SentencePiece
tokenizers (Mistral, Gemma, T5) store raw Unicode pieces, so applying the
byte decoder to them corrupts multi-byte characters.
We probe by checking ``is_byte_level`` (HuggingFace fast tokenizers) or
by verifying that a known multi-byte character (é, U+00E9) round-trips:
GPT-2 encodes it as a single byte-level piece ``chr(233)`` whose byte
decoder output [233] does NOT form valid UTF-8, while SentencePiece stores
the full character ``é`` whose UTF-8 is [195, 169].
"""
tid = id(tokenizer)
if tid in _BYTE_LEVEL_TOKENIZERS:
return True
# Fast path: HuggingFace fast tokenizers expose is_byte_level.
is_bl = getattr(tokenizer, "is_byte_level", None)
if isinstance(is_bl, bool):
if is_bl:
_BYTE_LEVEL_TOKENIZERS.add(tid)
return is_bl
# Slow path: probe with a known é token.
global _BYTE_DECODER
if not _BYTE_DECODER:
_BYTE_DECODER = _build_byte_decoder()
try:
vocab_size = len(tokenizer.get_vocab())
# Sample a few tokens to check if all chars are byte-decodable.
sample_ids = [0, 1, 2, 3, vocab_size // 2, vocab_size - 2]
for sid in sample_ids:
if sid < 0 or sid >= vocab_size:
continue
piece = tokenizer.convert_ids_to_tokens(sid)
if piece is None or not piece:
continue
# If any char in the piece is NOT in the byte decoder table,
# this tokenizer does NOT use byte-level encoding.
if any(ch not in _BYTE_DECODER for ch in piece):
return False
# All sampled tokens are byte-decodable → likely byte-level BPE.
_BYTE_LEVEL_TOKENIZERS.add(tid)
return True
except Exception:
return False
def token_id_to_bytes(tokenizer, token_id) -> list[int] | None:
"""Raw bytes for a byte-level-BPE token id.
Returns the token's original bytes via the GPT-2 byte decoder, or None when
the token is not byte-level representable (e.g. special ids / non byte BPE
tokenizers like SentencePiece), so callers can fall back to the detokenized
display string.
"""
if not _is_byte_level_tokenizer(tokenizer):
return None
global _BYTE_DECODER
if not _BYTE_DECODER:
_BYTE_DECODER = _build_byte_decoder()
try:
piece = tokenizer.convert_ids_to_tokens(token_id)
except Exception:
return None
if piece is None:
return None
out = bytearray()
for ch in piece:
b = _BYTE_DECODER.get(ch)
if b is None:
return None
out.append(b)
if not out:
return None
return list(out)
def to_openai_style_logprobs( def to_openai_style_logprobs(
input_token_logprobs=None, input_token_logprobs=None,
output_token_logprobs=None, output_token_logprobs=None,
input_top_logprobs=None, input_top_logprobs=None,
output_top_logprobs=None, output_top_logprobs=None,
tokenizer=None,
): ):
"""Convert engine logprob triples to an OpenAI ``LogProbs`` object.
Each engine logprob item is a ``(logprob, token_id, token_text)`` triple.
``token_text`` is a detokenized *display* string that loses fragmentary
byte-level tokens (a lone byte of a 4-byte char decodes to U+FFFD). The
legacy completions surface has no per-token ``bytes`` field, so when
``tokenizer`` is provided we render fragments losslessly as latin-1
(one char per raw byte), keeping the string channel reversible.
"""
ret_logprobs = LogProbs() ret_logprobs = LogProbs()
def append_token_logprobs(token_logprobs): def append_token_logprobs(token_logprobs):
for logprob, _, token_text in token_logprobs: for logprob, token_id, token_text in token_logprobs:
token_text = _lossless_token_text(tokenizer, token_id, token_text)
ret_logprobs.tokens.append(token_text) ret_logprobs.tokens.append(token_text)
ret_logprobs.token_logprobs.append(logprob) ret_logprobs.token_logprobs.append(logprob)
@@ -35,7 +154,10 @@ def to_openai_style_logprobs(
for tokens in top_logprobs: for tokens in top_logprobs:
if tokens is not None: if tokens is not None:
ret_logprobs.top_logprobs.append( ret_logprobs.top_logprobs.append(
{token[2]: token[0] for token in tokens} {
_lossless_token_text(tokenizer, token_id, token_text): logprob
for logprob, token_id, token_text in tokens
}
) )
else: else:
ret_logprobs.top_logprobs.append(None) ret_logprobs.top_logprobs.append(None)
@@ -52,13 +174,49 @@ def to_openai_style_logprobs(
return ret_logprobs return ret_logprobs
def _lossless_token_text(tokenizer, token_id, token_text):
"""Return a lossless display string for one engine logprob triple.
Fragmentary byte-level tokens decode to U+FFFD in the display string. When
we can recover the true raw bytes from the token id (byte-level BPE), we
validate that the recovered bytes do NOT form valid UTF-8 (a real fragment
never does), then render them as latin-1 so every byte round-trips.
Three safeguards address the reviewer's concerns:
1. Non-byte-level tokenizers (SentencePiece/Mistral) are detected and
skipped, so multi-byte characters like é are NOT corrupted to [233].
2. Legitimate U+FFFD text (e.g. GPT-2 token 4210 = bytes [239,191,189])
round-trips as valid UTF-8, so we keep the original display text.
3. The latin-1 representation is only applied to genuine fragments (bytes
that fail UTF-8 decode), avoiding key collisions in top_logprobs.
"""
if token_text is not None and "\ufffd" not in token_text:
return token_text
if tokenizer is None or token_id is None:
return token_text if token_text is not None else ""
raw = token_id_to_bytes(tokenizer, token_id)
if raw is None:
return token_text if token_text is not None else ""
# Only treat as a fragment if the recovered bytes do NOT form valid UTF-8.
# A complete token whose display text happens to contain U+FFFD (e.g. token
# 4210 = bytes [239,191,189] = valid UTF-8 for U+FFFD) must be left alone.
try:
bytes(raw).decode("utf-8")
# Valid UTF-8 → this is NOT a fragment; keep the original display text.
return token_text if token_text is not None else ""
except UnicodeDecodeError:
pass
# Genuine fragment: render as latin-1 (one char per byte, lossless).
try:
return bytes(raw).decode("latin-1")
except Exception:
return token_text if token_text is not None else ""
def process_hidden_states_from_ret( def process_hidden_states_from_ret(
ret_item: Dict[str, Any], ret_item: dict[str, Any],
request: Union[ request: ChatCompletionRequest | CompletionRequest,
ChatCompletionRequest, ) -> list | None:
CompletionRequest,
],
) -> Optional[List]:
"""Process hidden states from a ret item in non-streaming response. """Process hidden states from a ret item in non-streaming response.
Args: Args:
@@ -78,9 +236,9 @@ def process_hidden_states_from_ret(
def process_hidden_states_for_response( def process_hidden_states_for_response(
hidden_states: Optional[List], hidden_states: list | None,
return_hidden_states: Union[bool, Literal["last"]], return_hidden_states: bool | Literal["last"],
) -> Optional[List]: ) -> list | None:
"""Format scheduler hidden states for OpenAI API responses.""" """Format scheduler hidden states for OpenAI API responses."""
if not return_hidden_states or hidden_states is None: if not return_hidden_states or hidden_states is None:
return None return None
@@ -107,12 +265,9 @@ def should_include_usage(
def process_routed_experts_from_ret( def process_routed_experts_from_ret(
ret_item: Dict[str, Any], ret_item: dict[str, Any],
request: Union[ request: ChatCompletionRequest | CompletionRequest,
ChatCompletionRequest, ) -> str | None:
CompletionRequest,
],
) -> Optional[str]:
"""Process routed experts from a ret item in non-streaming response.""" """Process routed experts from a ret item in non-streaming response."""
if not getattr(request, "return_routed_experts", False): if not getattr(request, "return_routed_experts", False):
return None return None
@@ -120,7 +275,7 @@ def process_routed_experts_from_ret(
def cached_tokens_details_from_dict( def cached_tokens_details_from_dict(
details: Dict[str, Any], details: dict[str, Any],
) -> CachedTokensDetails: ) -> CachedTokensDetails:
"""Convert a raw cached_tokens_details dict to a CachedTokensDetails object.""" """Convert a raw cached_tokens_details dict to a CachedTokensDetails object."""
if "storage" in details: if "storage" in details:
@@ -138,12 +293,9 @@ def cached_tokens_details_from_dict(
def process_cached_tokens_details_from_ret( def process_cached_tokens_details_from_ret(
ret_item: Dict[str, Any], ret_item: dict[str, Any],
request: Union[ request: ChatCompletionRequest | CompletionRequest,
ChatCompletionRequest, ) -> CachedTokensDetails | None:
CompletionRequest,
],
) -> Optional[CachedTokensDetails]:
"""Process cached tokens details from a ret item in non-streaming response.""" """Process cached tokens details from a ret item in non-streaming response."""
if not request.return_cached_tokens_details: if not request.return_cached_tokens_details:
return None return None
@@ -156,8 +308,8 @@ def process_cached_tokens_details_from_ret(
def spec_tokens_details_from_meta_info( def spec_tokens_details_from_meta_info(
meta_info: Dict[str, Any], meta_info: dict[str, Any],
) -> Optional[SpecTokensDetails]: ) -> SpecTokensDetails | None:
"""Build speculative decoding details from canonical or legacy metrics.""" """Build speculative decoding details from canonical or legacy metrics."""
details = dict(meta_info) details = dict(meta_info)
@@ -190,12 +342,9 @@ def spec_tokens_details_from_meta_info(
def process_spec_tokens_details_from_ret( def process_spec_tokens_details_from_ret(
ret_item: Dict[str, Any], ret_item: dict[str, Any],
request: Union[ request: ChatCompletionRequest | CompletionRequest,
ChatCompletionRequest, ) -> SpecTokensDetails | None:
CompletionRequest,
],
) -> Optional[SpecTokensDetails]:
"""Process speculative decoding details from a response item.""" """Process speculative decoding details from a response item."""
if not getattr(request, "return_spec_tokens_details", False): if not getattr(request, "return_spec_tokens_details", False):
return None return None
@@ -203,8 +352,8 @@ def process_spec_tokens_details_from_ret(
def convert_embeds_to_tensors( def convert_embeds_to_tensors(
embeds: Optional[Union[List[Optional[List[List[float]]]], List[List[float]]]], embeds: list[list[list[float]] | None] | list[list[float]] | None,
) -> Optional[List[Optional[List[torch.Tensor]]]]: ) -> list[list[torch.Tensor] | None] | None:
"""Convert nested float lists from the HTTP API to lists of tensors. """Convert nested float lists from the HTTP API to lists of tensors.
Accepts either: Accepts either: