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:
@@ -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]]]]:
|
||||
|
||||
Reference in New Issue
Block a user