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
|
return_routed_experts: bool = False
|
||||||
routed_experts_start_len: int = 0
|
routed_experts_start_len: int = 0
|
||||||
return_cached_tokens_details: bool = False
|
return_cached_tokens_details: bool = False
|
||||||
|
return_spec_tokens_details: bool = False
|
||||||
return_token_ids: bool = False
|
return_token_ids: bool = False
|
||||||
|
|
||||||
# Extra parameters for SRT backend only and will be ignored by OpenAI models.
|
# Extra parameters for SRT backend only and will be ignored by OpenAI models.
|
||||||
@@ -413,6 +414,20 @@ class CompletionRequest(BaseModel):
|
|||||||
return v
|
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):
|
class SglExt(BaseModel):
|
||||||
"""SGLang extension fields for OpenAI-compatible responses.
|
"""SGLang extension fields for OpenAI-compatible responses.
|
||||||
|
|
||||||
@@ -422,6 +437,9 @@ class SglExt(BaseModel):
|
|||||||
|
|
||||||
routed_experts: Optional[str] = None
|
routed_experts: Optional[str] = None
|
||||||
cached_tokens_details: Optional[CachedTokensDetails] = None
|
cached_tokens_details: Optional[CachedTokensDetails] = None
|
||||||
|
spec_tokens_details: Optional[Union[SpecTokensDetails, List[SpecTokensDetails]]] = (
|
||||||
|
None
|
||||||
|
)
|
||||||
|
|
||||||
@model_serializer(mode="wrap")
|
@model_serializer(mode="wrap")
|
||||||
def _serialize(self, handler):
|
def _serialize(self, handler):
|
||||||
@@ -796,6 +814,7 @@ class ChatCompletionRequest(BaseModel):
|
|||||||
return_routed_experts: bool = False
|
return_routed_experts: bool = False
|
||||||
routed_experts_start_len: int = 0
|
routed_experts_start_len: int = 0
|
||||||
return_cached_tokens_details: bool = False
|
return_cached_tokens_details: bool = False
|
||||||
|
return_spec_tokens_details: bool = False
|
||||||
return_prompt_token_ids: bool = False
|
return_prompt_token_ids: bool = False
|
||||||
return_token_ids: bool = False
|
return_token_ids: bool = False
|
||||||
return_meta_info: 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_for_response,
|
||||||
process_hidden_states_from_ret,
|
process_hidden_states_from_ret,
|
||||||
process_routed_experts_from_ret,
|
process_routed_experts_from_ret,
|
||||||
|
process_spec_tokens_details_from_ret,
|
||||||
should_include_usage,
|
should_include_usage,
|
||||||
|
spec_tokens_details_from_meta_info,
|
||||||
to_openai_style_logprobs,
|
to_openai_style_logprobs,
|
||||||
)
|
)
|
||||||
from sglang.srt.entrypoints.request_headers import apply_header_overrides
|
from sglang.srt.entrypoints.request_headers import apply_header_overrides
|
||||||
@@ -1515,6 +1517,7 @@ class OpenAIServingChat(OpenAIServingBase):
|
|||||||
hidden_states = {}
|
hidden_states = {}
|
||||||
routed_experts = {}
|
routed_experts = {}
|
||||||
cached_tokens_details = {}
|
cached_tokens_details = {}
|
||||||
|
spec_tokens_details = {}
|
||||||
image_tokens = {}
|
image_tokens = {}
|
||||||
audio_tokens = {}
|
audio_tokens = {}
|
||||||
video_tokens = {}
|
video_tokens = {}
|
||||||
@@ -1546,6 +1549,10 @@ class OpenAIServingChat(OpenAIServingBase):
|
|||||||
cached_tokens_details[index] = content["meta_info"].get(
|
cached_tokens_details[index] = content["meta_info"].get(
|
||||||
"cached_tokens_details", None
|
"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)
|
image_tokens[index] = content["meta_info"].get("image_tokens", 0)
|
||||||
audio_tokens[index] = content["meta_info"].get("audio_tokens", 0)
|
audio_tokens[index] = content["meta_info"].get("audio_tokens", 0)
|
||||||
video_tokens[index] = content["meta_info"].get("video_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
|
(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:
|
if request.return_cached_tokens_details and cached_tokens_details:
|
||||||
first_details = next(
|
first_details = next(
|
||||||
(v for v in cached_tokens_details.values() if v is not None), None
|
(v for v in cached_tokens_details.values() if v is not None), None
|
||||||
)
|
)
|
||||||
if first_details is not 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(
|
sglext_chunk = ChatCompletionStreamResponse(
|
||||||
id=content["meta_info"]["id"],
|
id=content["meta_info"]["id"],
|
||||||
created=int(time.time()),
|
created=int(time.time()),
|
||||||
@@ -1682,7 +1710,8 @@ class OpenAIServingChat(OpenAIServingBase):
|
|||||||
model=request.model,
|
model=request.model,
|
||||||
sglext=SglExt(
|
sglext=SglExt(
|
||||||
routed_experts=sglext_routed,
|
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"
|
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(
|
cached_tokens_details = process_cached_tokens_details_from_ret(
|
||||||
first_ret, request
|
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
|
response_sglext = None
|
||||||
if routed_experts or cached_tokens_details:
|
if routed_experts or cached_tokens_details or spec_tokens_details:
|
||||||
response_sglext = SglExt(
|
response_sglext = SglExt(
|
||||||
routed_experts=routed_experts,
|
routed_experts=routed_experts,
|
||||||
cached_tokens_details=cached_tokens_details,
|
cached_tokens_details=cached_tokens_details,
|
||||||
|
spec_tokens_details=spec_tokens_details,
|
||||||
)
|
)
|
||||||
|
|
||||||
for idx, ret_item in enumerate(ret):
|
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_for_response,
|
||||||
process_hidden_states_from_ret,
|
process_hidden_states_from_ret,
|
||||||
process_routed_experts_from_ret,
|
process_routed_experts_from_ret,
|
||||||
|
process_spec_tokens_details_from_ret,
|
||||||
should_include_usage,
|
should_include_usage,
|
||||||
|
spec_tokens_details_from_meta_info,
|
||||||
to_openai_style_logprobs,
|
to_openai_style_logprobs,
|
||||||
)
|
)
|
||||||
from sglang.srt.managers.io_struct import GenerateReqInput
|
from sglang.srt.managers.io_struct import GenerateReqInput
|
||||||
@@ -237,6 +239,7 @@ class OpenAIServingCompletion(OpenAIServingBase):
|
|||||||
hidden_states = {}
|
hidden_states = {}
|
||||||
routed_experts = {}
|
routed_experts = {}
|
||||||
cached_tokens_details = {}
|
cached_tokens_details = {}
|
||||||
|
spec_tokens_details = {}
|
||||||
|
|
||||||
stream_started = False
|
stream_started = False
|
||||||
try:
|
try:
|
||||||
@@ -264,6 +267,10 @@ class OpenAIServingCompletion(OpenAIServingBase):
|
|||||||
cached_tokens_details[index] = content["meta_info"].get(
|
cached_tokens_details[index] = content["meta_info"].get(
|
||||||
"cached_tokens_details", None
|
"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
|
is_first_chunk = index not in stream_offsets
|
||||||
offset = stream_offsets.get(index, 0)
|
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
|
(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:
|
if request.return_cached_tokens_details and cached_tokens_details:
|
||||||
first_details = next(
|
first_details = next(
|
||||||
(v for v in cached_tokens_details.values() if v is not None), None
|
(v for v in cached_tokens_details.values() if v is not None), None
|
||||||
)
|
)
|
||||||
if first_details is not 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(
|
sglext_chunk = CompletionStreamResponse(
|
||||||
id=content["meta_info"]["id"],
|
id=content["meta_info"]["id"],
|
||||||
created=created,
|
created=created,
|
||||||
@@ -436,7 +464,8 @@ class OpenAIServingCompletion(OpenAIServingBase):
|
|||||||
model=request.model,
|
model=request.model,
|
||||||
sglext=SglExt(
|
sglext=SglExt(
|
||||||
routed_experts=sglext_routed,
|
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"
|
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(
|
cached_tokens_details = process_cached_tokens_details_from_ret(
|
||||||
first_ret, request
|
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
|
response_sglext = None
|
||||||
if routed_experts or cached_tokens_details:
|
if routed_experts or cached_tokens_details or spec_tokens_details:
|
||||||
response_sglext = SglExt(
|
response_sglext = SglExt(
|
||||||
routed_experts=routed_experts,
|
routed_experts=routed_experts,
|
||||||
cached_tokens_details=cached_tokens_details,
|
cached_tokens_details=cached_tokens_details,
|
||||||
|
spec_tokens_details=spec_tokens_details,
|
||||||
)
|
)
|
||||||
|
|
||||||
for idx, ret_item in enumerate(ret):
|
for idx, ret_item in enumerate(ret):
|
||||||
|
|||||||
@@ -8,6 +8,7 @@ from sglang.srt.entrypoints.openai.protocol import (
|
|||||||
ChatCompletionRequest,
|
ChatCompletionRequest,
|
||||||
CompletionRequest,
|
CompletionRequest,
|
||||||
LogProbs,
|
LogProbs,
|
||||||
|
SpecTokensDetails,
|
||||||
StreamOptions,
|
StreamOptions,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -154,6 +155,53 @@ def process_cached_tokens_details_from_ret(
|
|||||||
return cached_tokens_details_from_dict(details)
|
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(
|
def convert_embeds_to_tensors(
|
||||||
embeds: Optional[Union[List[Optional[List[List[float]]]], List[List[float]]]],
|
embeds: Optional[Union[List[Optional[List[List[float]]]], List[List[float]]]],
|
||||||
) -> Optional[List[Optional[List[torch.Tensor]]]]:
|
) -> 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")
|
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_PREVIEW_ENCODER = 'REASONING_EFFORT_MAX = "preview"\n'
|
||||||
_DSV4_OFFICIAL_ENCODER = (
|
_DSV4_OFFICIAL_ENCODER = (
|
||||||
"REASONING_EFFORT_PROMPTS: Dict[str, str] = "
|
"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(
|
def test_non_streaming_chat_response_returns_requested_token_ids_and_meta_info(
|
||||||
self,
|
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):
|
def _collect_continuous_usage(self, cached_tokens):
|
||||||
content = {
|
content = {
|
||||||
"text": "Hello",
|
"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")
|
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:
|
class _MockTemplateManager:
|
||||||
"""Minimal mock for TemplateManager."""
|
"""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):
|
def test_streaming_cached_tokens_details_emits_sglext(self):
|
||||||
"""Test that streaming completion responses emit cached token details in sglext."""
|
"""Test that streaming completion responses emit cached token details in sglext."""
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user