feat(api): add sglext_spec (#33518)

Signed-off-by: Muqi Li <muqi1029@gmail.com>
Co-authored-by: Codex <noreply@openai.com>
This commit is contained in:
Muqi Li
2026-08-19 11:54:18 +08:00
committed by GitHub
co-authored by Codex
parent 3b065a56b0
commit 88f6074392
6 changed files with 360 additions and 10 deletions
@@ -352,6 +352,7 @@ class CompletionRequest(BaseModel):
return_routed_experts: bool = False
routed_experts_start_len: int = 0
return_cached_tokens_details: bool = False
return_spec_tokens_details: bool = False
return_token_ids: bool = False
# Extra parameters for SRT backend only and will be ignored by OpenAI models.
@@ -413,6 +414,20 @@ class CompletionRequest(BaseModel):
return v
class SpecTokensDetails(BaseModel):
"""Per-request speculative decoding statistics."""
spec_accept_rate: float = 0.0
spec_accept_length: float = 0.0
spec_cap_length: float = 0.0
spec_block_accept_length: float = 0.0
spec_num_correct_drafts: int = 0
spec_num_proposed_drafts: int = 0
spec_verify_ct: int = 0
spec_correct_drafts_histogram: List[int] = Field(default_factory=list)
spec_cap_lens_histogram: List[int] = Field(default_factory=list)
class SglExt(BaseModel):
"""SGLang extension fields for OpenAI-compatible responses.
@@ -422,6 +437,9 @@ class SglExt(BaseModel):
routed_experts: Optional[str] = None
cached_tokens_details: Optional[CachedTokensDetails] = None
spec_tokens_details: Optional[Union[SpecTokensDetails, List[SpecTokensDetails]]] = (
None
)
@model_serializer(mode="wrap")
def _serialize(self, handler):
@@ -796,6 +814,7 @@ class ChatCompletionRequest(BaseModel):
return_routed_experts: bool = False
routed_experts_start_len: int = 0
return_cached_tokens_details: bool = False
return_spec_tokens_details: bool = False
return_prompt_token_ids: bool = False
return_token_ids: bool = False
return_meta_info: bool = False
@@ -58,7 +58,9 @@ from sglang.srt.entrypoints.openai.utils import (
process_hidden_states_for_response,
process_hidden_states_from_ret,
process_routed_experts_from_ret,
process_spec_tokens_details_from_ret,
should_include_usage,
spec_tokens_details_from_meta_info,
to_openai_style_logprobs,
)
from sglang.srt.entrypoints.request_headers import apply_header_overrides
@@ -1515,6 +1517,7 @@ class OpenAIServingChat(OpenAIServingBase):
hidden_states = {}
routed_experts = {}
cached_tokens_details = {}
spec_tokens_details = {}
image_tokens = {}
audio_tokens = {}
video_tokens = {}
@@ -1546,6 +1549,10 @@ class OpenAIServingChat(OpenAIServingBase):
cached_tokens_details[index] = content["meta_info"].get(
"cached_tokens_details", None
)
if request.return_spec_tokens_details:
spec_tokens_details[index] = spec_tokens_details_from_meta_info(
content["meta_info"]
)
image_tokens[index] = content["meta_info"].get("image_tokens", 0)
audio_tokens[index] = content["meta_info"].get("audio_tokens", 0)
video_tokens[index] = content["meta_info"].get("video_tokens", 0)
@@ -1666,15 +1673,36 @@ class OpenAIServingChat(OpenAIServingBase):
(v for v in routed_experts.values() if v is not None), None
)
sglext_details = None
sglext_cached_tokens_details = None
if request.return_cached_tokens_details and cached_tokens_details:
first_details = next(
(v for v in cached_tokens_details.values() if v is not None), None
)
if first_details is not None:
sglext_details = cached_tokens_details_from_dict(first_details)
sglext_cached_tokens_details = cached_tokens_details_from_dict(
first_details
)
if sglext_routed is not None or sglext_details is not None:
sglext_spec_tokens_details = None
if request.return_spec_tokens_details and spec_tokens_details:
spec_details = [
spec_tokens_details[index]
for index in sorted(spec_tokens_details)
if spec_tokens_details[index] is not None
]
if spec_details:
sglext_spec_tokens_details = (
spec_details if request.n > 1 else spec_details[0]
)
if any(
obj is not None
for obj in [
sglext_routed,
sglext_cached_tokens_details,
sglext_spec_tokens_details,
]
):
sglext_chunk = ChatCompletionStreamResponse(
id=content["meta_info"]["id"],
created=int(time.time()),
@@ -1682,7 +1710,8 @@ class OpenAIServingChat(OpenAIServingBase):
model=request.model,
sglext=SglExt(
routed_experts=sglext_routed,
cached_tokens_details=sglext_details,
cached_tokens_details=sglext_cached_tokens_details,
spec_tokens_details=sglext_spec_tokens_details,
),
)
yield f"data: {sglext_chunk.model_dump_json()}\n\n"
@@ -1782,11 +1811,24 @@ class OpenAIServingChat(OpenAIServingBase):
cached_tokens_details = process_cached_tokens_details_from_ret(
first_ret, request
)
spec_details = [
detail
for detail in (
process_spec_tokens_details_from_ret(item, request) for item in ret
)
if detail is not None
]
spec_tokens_details = (
spec_details
if request.n > 1
else (spec_details[0] if spec_details else None)
)
response_sglext = None
if routed_experts or cached_tokens_details:
if routed_experts or cached_tokens_details or spec_tokens_details:
response_sglext = SglExt(
routed_experts=routed_experts,
cached_tokens_details=cached_tokens_details,
spec_tokens_details=spec_tokens_details,
)
for idx, ret_item in enumerate(ret):
@@ -25,7 +25,9 @@ from sglang.srt.entrypoints.openai.utils import (
process_hidden_states_for_response,
process_hidden_states_from_ret,
process_routed_experts_from_ret,
process_spec_tokens_details_from_ret,
should_include_usage,
spec_tokens_details_from_meta_info,
to_openai_style_logprobs,
)
from sglang.srt.managers.io_struct import GenerateReqInput
@@ -237,6 +239,7 @@ class OpenAIServingCompletion(OpenAIServingBase):
hidden_states = {}
routed_experts = {}
cached_tokens_details = {}
spec_tokens_details = {}
stream_started = False
try:
@@ -264,6 +267,10 @@ class OpenAIServingCompletion(OpenAIServingBase):
cached_tokens_details[index] = content["meta_info"].get(
"cached_tokens_details", None
)
if request.return_spec_tokens_details:
spec_tokens_details[index] = spec_tokens_details_from_meta_info(
content["meta_info"]
)
is_first_chunk = index not in stream_offsets
offset = stream_offsets.get(index, 0)
@@ -419,15 +426,36 @@ class OpenAIServingCompletion(OpenAIServingBase):
(v for v in routed_experts.values() if v is not None), None
)
sglext_details = None
sglext_cached_tokens_details = None
if request.return_cached_tokens_details and cached_tokens_details:
first_details = next(
(v for v in cached_tokens_details.values() if v is not None), None
)
if first_details is not None:
sglext_details = cached_tokens_details_from_dict(first_details)
sglext_cached_tokens_details = cached_tokens_details_from_dict(
first_details
)
if sglext_routed is not None or sglext_details is not None:
sglext_spec_tokens_details = None
if request.return_spec_tokens_details and spec_tokens_details:
spec_details = [
spec_tokens_details[index]
for index in sorted(spec_tokens_details)
if spec_tokens_details[index] is not None
]
if spec_details:
sglext_spec_tokens_details = (
spec_details if request.n > 1 else spec_details[0]
)
if any(
obj is not None
for obj in [
sglext_routed,
sglext_cached_tokens_details,
sglext_spec_tokens_details,
]
):
sglext_chunk = CompletionStreamResponse(
id=content["meta_info"]["id"],
created=created,
@@ -436,7 +464,8 @@ class OpenAIServingCompletion(OpenAIServingBase):
model=request.model,
sglext=SglExt(
routed_experts=sglext_routed,
cached_tokens_details=sglext_details,
cached_tokens_details=sglext_cached_tokens_details,
spec_tokens_details=sglext_spec_tokens_details,
),
)
yield f"data: {sglext_chunk.model_dump_json()}\n\n"
@@ -517,11 +546,24 @@ class OpenAIServingCompletion(OpenAIServingBase):
cached_tokens_details = process_cached_tokens_details_from_ret(
first_ret, request
)
spec_details = [
detail
for detail in (
process_spec_tokens_details_from_ret(item, request) for item in ret
)
if detail is not None
]
spec_tokens_details = (
spec_details
if request.n > 1
else (spec_details[0] if spec_details else None)
)
response_sglext = None
if routed_experts or cached_tokens_details:
if routed_experts or cached_tokens_details or spec_tokens_details:
response_sglext = SglExt(
routed_experts=routed_experts,
cached_tokens_details=cached_tokens_details,
spec_tokens_details=spec_tokens_details,
)
for idx, ret_item in enumerate(ret):
@@ -8,6 +8,7 @@ from sglang.srt.entrypoints.openai.protocol import (
ChatCompletionRequest,
CompletionRequest,
LogProbs,
SpecTokensDetails,
StreamOptions,
)
@@ -154,6 +155,53 @@ def process_cached_tokens_details_from_ret(
return cached_tokens_details_from_dict(details)
def spec_tokens_details_from_meta_info(
meta_info: Dict[str, Any],
) -> Optional[SpecTokensDetails]:
"""Build speculative decoding details from canonical or legacy metrics."""
details = dict(meta_info)
metric_keys = (
"spec_accept_rate",
"spec_accept_length",
"spec_cap_length",
"spec_block_accept_length",
"spec_num_correct_drafts",
"spec_num_proposed_drafts",
"spec_verify_ct",
"spec_correct_drafts_histogram",
"spec_cap_lens_histogram",
)
if not any(key in details for key in metric_keys):
return None
return SpecTokensDetails(
spec_accept_rate=details.get("spec_accept_rate") or 0.0,
spec_accept_length=details.get("spec_accept_length") or 0.0,
spec_cap_length=details.get("spec_cap_length") or 0.0,
spec_block_accept_length=details.get("spec_block_accept_length") or 0.0,
spec_num_correct_drafts=details.get("spec_num_correct_drafts") or 0,
spec_num_proposed_drafts=details.get("spec_num_proposed_drafts") or 0,
spec_verify_ct=details.get("spec_verify_ct") or 0,
spec_correct_drafts_histogram=details.get("spec_correct_drafts_histogram")
or [],
spec_cap_lens_histogram=details.get("spec_cap_lens_histogram") or [],
)
def process_spec_tokens_details_from_ret(
ret_item: Dict[str, Any],
request: Union[
ChatCompletionRequest,
CompletionRequest,
],
) -> Optional[SpecTokensDetails]:
"""Process speculative decoding details from a response item."""
if not getattr(request, "return_spec_tokens_details", False):
return None
return spec_tokens_details_from_meta_info(ret_item["meta_info"])
def convert_embeds_to_tensors(
embeds: Optional[Union[List[Optional[List[List[float]]]], List[List[float]]]],
) -> Optional[List[Optional[List[torch.Tensor]]]]: