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]]]]:
@@ -43,6 +43,31 @@ from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=11, suite="base-a-test-cpu")
def _spec_result(index):
return {
"text": f"choice-{index}",
"meta_info": {
"id": "chatcmpl-spec-test",
"prompt_tokens": 10,
"completion_tokens": 2,
"cached_tokens": 0,
"finish_reason": {"type": "stop"},
"weight_version": "default",
"spec_accept_rate": 0.5,
"spec_accept_length": 2.0,
"spec_cap_length": index + 1.0,
"spec_block_accept_length": index + 0.5,
"spec_num_correct_drafts": 1,
"spec_num_proposed_drafts": 2,
"spec_verify_ct": 1,
"spec_correct_drafts_histogram": [0, 1],
"spec_cap_lens_histogram": [index, 1],
},
"index": index,
}
_DSV4_PREVIEW_ENCODER = 'REASONING_EFFORT_MAX = "preview"\n'
_DSV4_OFFICIAL_ENCODER = (
"REASONING_EFFORT_PROMPTS: Dict[str, str] = "
@@ -2494,6 +2519,34 @@ class ServingChatTestCase(unittest.TestCase):
},
)
def test_parallel_sampling_returns_spec_details_per_choice(self):
req = ChatCompletionRequest(
model="x",
messages=[{"role": "user", "content": "Hi?"}],
max_tokens=100,
n=2,
return_spec_tokens_details=True,
)
ret = [_spec_result(index) for index in range(2)]
response = self.chat._build_chat_response(req, ret, 1234567890)
details = response.sglext.spec_tokens_details
self.assertEqual([item.spec_cap_length for item in details], [1.0, 2.0])
self.assertEqual(
[item.spec_cap_lens_histogram for item in details],
[[0, 1], [1, 1]],
)
single_req = req.model_copy(update={"n": 1})
single_response = self.chat._build_chat_response(
single_req, ret[:1], 1234567890
)
self.assertEqual(
single_response.sglext.spec_tokens_details.spec_cap_length,
1.0,
)
def test_non_streaming_chat_response_returns_requested_token_ids_and_meta_info(
self,
):
@@ -2609,6 +2662,31 @@ class ServingChatTestCase(unittest.TestCase):
},
)
def test_streaming_parallel_sampling_orders_spec_details_by_choice(self):
async def mock_generate():
for index in (1, 0):
yield _spec_result(index)
self.tm.generate_request.return_value = mock_generate()
req = ChatCompletionRequest(
model="x",
messages=[{"role": "user", "content": "Hi?"}],
max_tokens=100,
n=2,
stream=True,
return_spec_tokens_details=True,
)
parsed = self._parse_chunks(self._run_chat_stream(Mock(), req))
details = next(chunk["sglext"] for chunk in parsed if "sglext" in chunk)[
"spec_tokens_details"
]
self.assertEqual([item["spec_cap_length"] for item in details], [1.0, 2.0])
self.assertEqual(
[item["spec_cap_lens_histogram"] for item in details],
[[0, 1], [1, 1]],
)
def _collect_continuous_usage(self, cached_tokens):
content = {
"text": "Hello",
@@ -25,6 +25,30 @@ from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=11, suite="base-a-test-cpu")
def _spec_result(index):
return {
"text": f"choice-{index}",
"meta_info": {
"id": "cmpl-spec-test",
"prompt_tokens": 10,
"completion_tokens": 2,
"cached_tokens": 0,
"finish_reason": {"type": "stop"},
"weight_version": "default",
"spec_accept_rate": 0.5,
"spec_accept_length": 2.0,
"spec_cap_length": index + 1.0,
"spec_block_accept_length": index + 0.5,
"spec_num_correct_drafts": 1,
"spec_num_proposed_drafts": 2,
"spec_verify_ct": 1,
"spec_correct_drafts_histogram": [0, 1],
"spec_cap_lens_histogram": [index, 1],
},
"index": index,
}
class _MockTemplateManager:
"""Minimal mock for TemplateManager."""
@@ -400,6 +424,103 @@ class ServingCompletionTestCase(unittest.TestCase):
},
)
def test_parallel_sampling_returns_spec_details_per_choice(self):
req = CompletionRequest(
model="x",
prompt="Hello world",
max_tokens=100,
n=2,
return_spec_tokens_details=True,
)
ret = [_spec_result(index) for index in range(2)]
response = self.sc._build_completion_response(req, ret, 1234567890)
details = response.sglext.spec_tokens_details
self.assertEqual(len(details), 2)
self.assertEqual(details[0].spec_cap_length, 1.0)
self.assertEqual(details[0].spec_block_accept_length, 0.5)
self.assertEqual(details[0].spec_cap_lens_histogram, [0, 1])
self.assertEqual(details[1].spec_cap_length, 2.0)
self.assertEqual(details[1].spec_block_accept_length, 1.5)
self.assertEqual(details[1].spec_cap_lens_histogram, [1, 1])
single_req = req.model_copy(update={"n": 1})
single_response = self.sc._build_completion_response(
single_req, ret[:1], 1234567890
)
self.assertEqual(
single_response.sglext.spec_tokens_details.spec_cap_length,
1.0,
)
disabled_req = single_req.model_copy(
update={"return_spec_tokens_details": False}
)
disabled_response = self.sc._build_completion_response(
disabled_req, ret[:1], 1234567890
)
self.assertIsNone(disabled_response.sglext)
def test_streaming_parallel_sampling_orders_spec_details_by_choice(self):
async def mock_generate(*args, **kwargs):
for index in (1, 0):
yield _spec_result(index)
self.sc.tokenizer_manager.generate_request = mock_generate
req = CompletionRequest(
model="x",
prompt="Hello world",
max_tokens=100,
n=2,
stream=True,
return_spec_tokens_details=True,
)
adapted_request, _ = self.sc._convert_to_internal_request(req)
async def run_stream(request):
return [
chunk
async for chunk in self.sc._generate_completion_stream(
adapted_request, request, self.fastapi_request
)
]
chunks = get_or_create_event_loop().run_until_complete(run_stream(req))
parsed = [
json.loads(chunk[len("data: ") :])
for chunk in chunks
if chunk.startswith("data: ") and chunk.strip() != "data: [DONE]"
]
details = next(chunk["sglext"] for chunk in parsed if "sglext" in chunk)[
"spec_tokens_details"
]
self.assertEqual([item["spec_cap_length"] for item in details], [1.0, 2.0])
self.assertEqual(
[item["spec_cap_lens_histogram"] for item in details],
[[0, 1], [1, 1]],
)
async def mock_single_generate(*args, **kwargs):
async for content in mock_generate():
if content["index"] == 0:
yield content
self.sc.tokenizer_manager.generate_request = mock_single_generate
single_req = req.model_copy(update={"n": 1})
single_chunks = get_or_create_event_loop().run_until_complete(
run_stream(single_req)
)
single_parsed = [
json.loads(chunk[len("data: ") :])
for chunk in single_chunks
if chunk.startswith("data: ") and chunk.strip() != "data: [DONE]"
]
single_details = next(
chunk["sglext"] for chunk in single_parsed if "sglext" in chunk
)["spec_tokens_details"]
self.assertIsInstance(single_details, dict)
def test_streaming_cached_tokens_details_emits_sglext(self):
"""Test that streaming completion responses emit cached token details in sglext."""