[HiCache] return cached_tokens_details in sglext for streaming responses (#22055)
Signed-off-by: Vladislav Nosivskoy <vladnosiv@gmail.com>
This commit is contained in:
@@ -39,6 +39,7 @@ from sglang.srt.entrypoints.openai.protocol import (
|
||||
from sglang.srt.entrypoints.openai.serving_base import OpenAIServingBase
|
||||
from sglang.srt.entrypoints.openai.usage_processor import UsageProcessor
|
||||
from sglang.srt.entrypoints.openai.utils import (
|
||||
cached_tokens_details_from_dict,
|
||||
process_cached_tokens_details_from_ret,
|
||||
process_hidden_states_from_ret,
|
||||
process_routed_experts_from_ret,
|
||||
@@ -766,6 +767,7 @@ class OpenAIServingChat(OpenAIServingBase):
|
||||
cached_tokens = {}
|
||||
hidden_states = {}
|
||||
routed_experts = {}
|
||||
cached_tokens_details = {}
|
||||
|
||||
stream_started = False
|
||||
try:
|
||||
@@ -789,6 +791,9 @@ class OpenAIServingChat(OpenAIServingBase):
|
||||
cached_tokens[index] = content["meta_info"].get("cached_tokens", 0)
|
||||
hidden_states[index] = content["meta_info"].get("hidden_states", None)
|
||||
routed_experts[index] = content["meta_info"].get("routed_experts", None)
|
||||
cached_tokens_details[index] = content["meta_info"].get(
|
||||
"cached_tokens_details", None
|
||||
)
|
||||
|
||||
# Handle logprobs
|
||||
choice_logprobs = None
|
||||
@@ -963,20 +968,32 @@ class OpenAIServingChat(OpenAIServingBase):
|
||||
)
|
||||
yield f"data: {hidden_states_chunk.model_dump_json()}\n\n"
|
||||
|
||||
sglext_routed = None
|
||||
if request.return_routed_experts and routed_experts:
|
||||
# Get first non-None routed_experts value
|
||||
first_routed_experts = next(
|
||||
sglext_routed = next(
|
||||
(v for v in routed_experts.values() if v is not None), None
|
||||
)
|
||||
if first_routed_experts is not None:
|
||||
routed_experts_chunk = ChatCompletionStreamResponse(
|
||||
id=content["meta_info"]["id"],
|
||||
created=int(time.time()),
|
||||
choices=[], # sglext is at response level
|
||||
model=request.model,
|
||||
sglext=SglExt(routed_experts=first_routed_experts),
|
||||
)
|
||||
yield f"data: {routed_experts_chunk.model_dump_json()}\n\n"
|
||||
|
||||
sglext_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)
|
||||
|
||||
if sglext_routed is not None or sglext_details is not None:
|
||||
sglext_chunk = ChatCompletionStreamResponse(
|
||||
id=content["meta_info"]["id"],
|
||||
created=int(time.time()),
|
||||
choices=[], # sglext is at response level
|
||||
model=request.model,
|
||||
sglext=SglExt(
|
||||
routed_experts=sglext_routed,
|
||||
cached_tokens_details=sglext_details,
|
||||
),
|
||||
)
|
||||
yield f"data: {sglext_chunk.model_dump_json()}\n\n"
|
||||
|
||||
# Additional usage chunk
|
||||
if include_usage:
|
||||
|
||||
@@ -20,6 +20,7 @@ from sglang.srt.entrypoints.openai.protocol import (
|
||||
from sglang.srt.entrypoints.openai.serving_base import OpenAIServingBase
|
||||
from sglang.srt.entrypoints.openai.usage_processor import UsageProcessor
|
||||
from sglang.srt.entrypoints.openai.utils import (
|
||||
cached_tokens_details_from_dict,
|
||||
process_cached_tokens_details_from_ret,
|
||||
process_hidden_states_from_ret,
|
||||
process_routed_experts_from_ret,
|
||||
@@ -224,6 +225,7 @@ class OpenAIServingCompletion(OpenAIServingBase):
|
||||
cached_tokens = {}
|
||||
hidden_states = {}
|
||||
routed_experts = {}
|
||||
cached_tokens_details = {}
|
||||
|
||||
stream_started = False
|
||||
try:
|
||||
@@ -248,6 +250,9 @@ class OpenAIServingCompletion(OpenAIServingBase):
|
||||
cached_tokens[index] = content["meta_info"].get("cached_tokens", 0)
|
||||
hidden_states[index] = content["meta_info"].get("hidden_states", None)
|
||||
routed_experts[index] = content["meta_info"].get("routed_experts", None)
|
||||
cached_tokens_details[index] = content["meta_info"].get(
|
||||
"cached_tokens_details", None
|
||||
)
|
||||
|
||||
is_first_chunk = index not in stream_offsets
|
||||
offset = stream_offsets.get(index, 0)
|
||||
@@ -379,21 +384,33 @@ class OpenAIServingCompletion(OpenAIServingBase):
|
||||
)
|
||||
yield f"data: {hidden_states_chunk.model_dump_json()}\n\n"
|
||||
|
||||
sglext_routed = None
|
||||
if request.return_routed_experts and routed_experts:
|
||||
# Get first non-None routed_experts value
|
||||
first_routed_experts = next(
|
||||
sglext_routed = next(
|
||||
(v for v in routed_experts.values() if v is not None), None
|
||||
)
|
||||
if first_routed_experts is not None:
|
||||
routed_experts_chunk = CompletionStreamResponse(
|
||||
id=content["meta_info"]["id"],
|
||||
created=created,
|
||||
object="text_completion",
|
||||
choices=[], # sglext is at response level
|
||||
model=request.model,
|
||||
sglext=SglExt(routed_experts=first_routed_experts),
|
||||
)
|
||||
yield f"data: {routed_experts_chunk.model_dump_json()}\n\n"
|
||||
|
||||
sglext_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)
|
||||
|
||||
if sglext_routed is not None or sglext_details is not None:
|
||||
sglext_chunk = CompletionStreamResponse(
|
||||
id=content["meta_info"]["id"],
|
||||
created=created,
|
||||
object="text_completion",
|
||||
choices=[], # sglext is at response level
|
||||
model=request.model,
|
||||
sglext=SglExt(
|
||||
routed_experts=sglext_routed,
|
||||
cached_tokens_details=sglext_details,
|
||||
),
|
||||
)
|
||||
yield f"data: {sglext_chunk.model_dump_json()}\n\n"
|
||||
|
||||
# Handle final usage chunk
|
||||
if include_usage:
|
||||
|
||||
@@ -106,22 +106,10 @@ def process_routed_experts_from_ret(
|
||||
return ret_item["meta_info"].get("routed_experts", None)
|
||||
|
||||
|
||||
def process_cached_tokens_details_from_ret(
|
||||
ret_item: Dict[str, Any],
|
||||
request: Union[
|
||||
ChatCompletionRequest,
|
||||
CompletionRequest,
|
||||
],
|
||||
) -> Optional[CachedTokensDetails]:
|
||||
"""Process cached tokens details from a ret item in non-streaming response."""
|
||||
if not getattr(request, "return_cached_tokens_details", False):
|
||||
return None
|
||||
|
||||
details = ret_item["meta_info"].get("cached_tokens_details", None)
|
||||
if details is None:
|
||||
return None
|
||||
|
||||
# Check if L3 storage fields are present
|
||||
def cached_tokens_details_from_dict(
|
||||
details: Dict[str, Any],
|
||||
) -> CachedTokensDetails:
|
||||
"""Convert a raw cached_tokens_details dict to a CachedTokensDetails object."""
|
||||
if "storage" in details:
|
||||
return CachedTokensDetails(
|
||||
device=details.get("device", 0),
|
||||
@@ -136,6 +124,24 @@ def process_cached_tokens_details_from_ret(
|
||||
)
|
||||
|
||||
|
||||
def process_cached_tokens_details_from_ret(
|
||||
ret_item: Dict[str, Any],
|
||||
request: Union[
|
||||
ChatCompletionRequest,
|
||||
CompletionRequest,
|
||||
],
|
||||
) -> Optional[CachedTokensDetails]:
|
||||
"""Process cached tokens details from a ret item in non-streaming response."""
|
||||
if not request.return_cached_tokens_details:
|
||||
return None
|
||||
|
||||
details = ret_item["meta_info"].get("cached_tokens_details", None)
|
||||
if details is None:
|
||||
return None
|
||||
|
||||
return cached_tokens_details_from_dict(details)
|
||||
|
||||
|
||||
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