diff --git a/python/sglang/srt/entrypoints/openai/protocol.py b/python/sglang/srt/entrypoints/openai/protocol.py index da62e3b0f..4c98638eb 100644 --- a/python/sglang/srt/entrypoints/openai/protocol.py +++ b/python/sglang/srt/entrypoints/openai/protocol.py @@ -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 diff --git a/python/sglang/srt/entrypoints/openai/serving_chat.py b/python/sglang/srt/entrypoints/openai/serving_chat.py index f572e6360..d1842c567 100644 --- a/python/sglang/srt/entrypoints/openai/serving_chat.py +++ b/python/sglang/srt/entrypoints/openai/serving_chat.py @@ -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): diff --git a/python/sglang/srt/entrypoints/openai/serving_completions.py b/python/sglang/srt/entrypoints/openai/serving_completions.py index d5587765b..820f57c04 100644 --- a/python/sglang/srt/entrypoints/openai/serving_completions.py +++ b/python/sglang/srt/entrypoints/openai/serving_completions.py @@ -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): diff --git a/python/sglang/srt/entrypoints/openai/utils.py b/python/sglang/srt/entrypoints/openai/utils.py index 86e6e663b..bc4820076 100644 --- a/python/sglang/srt/entrypoints/openai/utils.py +++ b/python/sglang/srt/entrypoints/openai/utils.py @@ -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]]]]: diff --git a/test/registered/unit/entrypoints/openai/test_serving_chat.py b/test/registered/unit/entrypoints/openai/test_serving_chat.py index 0e7861a13..f9e546f83 100644 --- a/test/registered/unit/entrypoints/openai/test_serving_chat.py +++ b/test/registered/unit/entrypoints/openai/test_serving_chat.py @@ -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", diff --git a/test/registered/unit/entrypoints/openai/test_serving_completions.py b/test/registered/unit/entrypoints/openai/test_serving_completions.py index 3f0b6736f..4b7c73fc5 100644 --- a/test/registered/unit/entrypoints/openai/test_serving_completions.py +++ b/test/registered/unit/entrypoints/openai/test_serving_completions.py @@ -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."""