[HiCache] feat: Add detailed cache hit breakdown for HiCache in sglext and Prometheus metrics (#17648)
Signed-off-by: Vladislav Nosivskoy <vladnosiv@gmail.com> Co-authored-by: ishandhanani <82981111+ishandhanani@users.noreply.github.com>
This commit is contained in:
co-authored by
ishandhanani
parent
d48bbe3bed
commit
e166ca8758
@@ -102,12 +102,38 @@ class ChoiceLogprobs(BaseModel):
|
|||||||
content: List[ChatCompletionTokenLogprob]
|
content: List[ChatCompletionTokenLogprob]
|
||||||
|
|
||||||
|
|
||||||
|
class CachedTokensDetails(BaseModel):
|
||||||
|
"""Detailed breakdown of cached tokens by cache source."""
|
||||||
|
|
||||||
|
device: int = 0 # Tokens from device cache (GPU)
|
||||||
|
host: int = 0 # Tokens from host cache (CPU memory)
|
||||||
|
# L3 storage fields are only present when storage backend is enabled
|
||||||
|
storage: Optional[int] = None # Tokens from L3 storage backend
|
||||||
|
storage_backend: Optional[str] = None # Type of storage backend used
|
||||||
|
|
||||||
|
@model_serializer(mode="wrap")
|
||||||
|
def _serialize(self, handler):
|
||||||
|
data = handler(self)
|
||||||
|
# Remove None fields so they don't appear in response when L3 is disabled
|
||||||
|
if self.storage is None:
|
||||||
|
data.pop("storage", None)
|
||||||
|
if self.storage_backend is None:
|
||||||
|
data.pop("storage_backend", None)
|
||||||
|
return data
|
||||||
|
|
||||||
|
|
||||||
|
class PromptTokensDetails(BaseModel):
|
||||||
|
"""Details about prompt tokens."""
|
||||||
|
|
||||||
|
cached_tokens: int = 0
|
||||||
|
|
||||||
|
|
||||||
class UsageInfo(BaseModel):
|
class UsageInfo(BaseModel):
|
||||||
prompt_tokens: int = 0
|
prompt_tokens: int = 0
|
||||||
total_tokens: int = 0
|
total_tokens: int = 0
|
||||||
completion_tokens: Optional[int] = 0
|
completion_tokens: Optional[int] = 0
|
||||||
# only used to return cached tokens when --enable-cache-report is set
|
# Used to return cached tokens info when --enable-cache-report is set
|
||||||
prompt_tokens_details: Optional[Dict[str, int]] = None
|
prompt_tokens_details: Optional[PromptTokensDetails] = None
|
||||||
reasoning_tokens: Optional[int] = 0
|
reasoning_tokens: Optional[int] = 0
|
||||||
|
|
||||||
|
|
||||||
@@ -233,6 +259,7 @@ class CompletionRequest(BaseModel):
|
|||||||
user: Optional[str] = None
|
user: Optional[str] = None
|
||||||
return_hidden_states: bool = False
|
return_hidden_states: bool = False
|
||||||
return_routed_experts: bool = False
|
return_routed_experts: bool = False
|
||||||
|
return_cached_tokens_details: 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.
|
||||||
top_k: int = -1
|
top_k: int = -1
|
||||||
@@ -289,6 +316,7 @@ class SglExt(BaseModel):
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
routed_experts: Optional[str] = None
|
routed_experts: Optional[str] = None
|
||||||
|
cached_tokens_details: Optional[CachedTokensDetails] = None
|
||||||
|
|
||||||
@model_serializer(mode="wrap")
|
@model_serializer(mode="wrap")
|
||||||
def _serialize(self, handler):
|
def _serialize(self, handler):
|
||||||
@@ -304,15 +332,12 @@ class CompletionResponseChoice(BaseModel):
|
|||||||
finish_reason: Optional[Literal["stop", "length", "content_filter", "abort"]] = None
|
finish_reason: Optional[Literal["stop", "length", "content_filter", "abort"]] = None
|
||||||
matched_stop: Union[None, int, str] = None
|
matched_stop: Union[None, int, str] = None
|
||||||
hidden_states: Optional[object] = None
|
hidden_states: Optional[object] = None
|
||||||
sgl_ext: Optional[SglExt] = None
|
|
||||||
|
|
||||||
@model_serializer(mode="wrap")
|
@model_serializer(mode="wrap")
|
||||||
def _serialize(self, handler):
|
def _serialize(self, handler):
|
||||||
data = handler(self)
|
data = handler(self)
|
||||||
if self.hidden_states is None:
|
if self.hidden_states is None:
|
||||||
data.pop("hidden_states", None)
|
data.pop("hidden_states", None)
|
||||||
if self.sgl_ext is None:
|
|
||||||
data.pop("sgl_ext", None)
|
|
||||||
return data
|
return data
|
||||||
|
|
||||||
|
|
||||||
@@ -324,6 +349,14 @@ class CompletionResponse(BaseModel):
|
|||||||
choices: List[CompletionResponseChoice]
|
choices: List[CompletionResponseChoice]
|
||||||
usage: UsageInfo
|
usage: UsageInfo
|
||||||
metadata: Optional[Dict[str, Any]] = None
|
metadata: Optional[Dict[str, Any]] = None
|
||||||
|
sglext: Optional[SglExt] = None
|
||||||
|
|
||||||
|
@model_serializer(mode="wrap")
|
||||||
|
def _serialize(self, handler):
|
||||||
|
data = handler(self)
|
||||||
|
if self.sglext is None:
|
||||||
|
data.pop("sglext", None)
|
||||||
|
return data
|
||||||
|
|
||||||
|
|
||||||
class CompletionResponseStreamChoice(BaseModel):
|
class CompletionResponseStreamChoice(BaseModel):
|
||||||
@@ -333,15 +366,12 @@ class CompletionResponseStreamChoice(BaseModel):
|
|||||||
finish_reason: Optional[Literal["stop", "length", "content_filter", "abort"]] = None
|
finish_reason: Optional[Literal["stop", "length", "content_filter", "abort"]] = None
|
||||||
matched_stop: Union[None, int, str] = None
|
matched_stop: Union[None, int, str] = None
|
||||||
hidden_states: Optional[object] = None
|
hidden_states: Optional[object] = None
|
||||||
sgl_ext: Optional[SglExt] = None
|
|
||||||
|
|
||||||
@model_serializer(mode="wrap")
|
@model_serializer(mode="wrap")
|
||||||
def _serialize(self, handler):
|
def _serialize(self, handler):
|
||||||
data = handler(self)
|
data = handler(self)
|
||||||
if self.hidden_states is None:
|
if self.hidden_states is None:
|
||||||
data.pop("hidden_states", None)
|
data.pop("hidden_states", None)
|
||||||
if self.sgl_ext is None:
|
|
||||||
data.pop("sgl_ext", None)
|
|
||||||
return data
|
return data
|
||||||
|
|
||||||
|
|
||||||
@@ -352,6 +382,14 @@ class CompletionStreamResponse(BaseModel):
|
|||||||
model: str
|
model: str
|
||||||
choices: List[CompletionResponseStreamChoice]
|
choices: List[CompletionResponseStreamChoice]
|
||||||
usage: Optional[UsageInfo] = None
|
usage: Optional[UsageInfo] = None
|
||||||
|
sglext: Optional[SglExt] = None
|
||||||
|
|
||||||
|
@model_serializer(mode="wrap")
|
||||||
|
def _serialize(self, handler):
|
||||||
|
data = handler(self)
|
||||||
|
if self.sglext is None:
|
||||||
|
data.pop("sglext", None)
|
||||||
|
return data
|
||||||
|
|
||||||
|
|
||||||
class ChatCompletionMessageContentTextPart(BaseModel):
|
class ChatCompletionMessageContentTextPart(BaseModel):
|
||||||
@@ -526,6 +564,7 @@ class ChatCompletionRequest(BaseModel):
|
|||||||
) # noqa
|
) # noqa
|
||||||
return_hidden_states: bool = False
|
return_hidden_states: bool = False
|
||||||
return_routed_experts: bool = False
|
return_routed_experts: bool = False
|
||||||
|
return_cached_tokens_details: bool = False
|
||||||
reasoning_effort: Optional[Literal["low", "medium", "high"]] = Field(
|
reasoning_effort: Optional[Literal["low", "medium", "high"]] = Field(
|
||||||
default="medium",
|
default="medium",
|
||||||
description="Constrains effort on reasoning for reasoning models. "
|
description="Constrains effort on reasoning for reasoning models. "
|
||||||
@@ -755,15 +794,12 @@ class ChatCompletionResponseChoice(BaseModel):
|
|||||||
] = None
|
] = None
|
||||||
matched_stop: Union[None, int, str] = None
|
matched_stop: Union[None, int, str] = None
|
||||||
hidden_states: Optional[object] = None
|
hidden_states: Optional[object] = None
|
||||||
sgl_ext: Optional[SglExt] = None
|
|
||||||
|
|
||||||
@model_serializer(mode="wrap")
|
@model_serializer(mode="wrap")
|
||||||
def _serialize(self, handler):
|
def _serialize(self, handler):
|
||||||
data = handler(self)
|
data = handler(self)
|
||||||
if self.hidden_states is None:
|
if self.hidden_states is None:
|
||||||
data.pop("hidden_states", None)
|
data.pop("hidden_states", None)
|
||||||
if self.sgl_ext is None:
|
|
||||||
data.pop("sgl_ext", None)
|
|
||||||
return data
|
return data
|
||||||
|
|
||||||
|
|
||||||
@@ -775,6 +811,14 @@ class ChatCompletionResponse(BaseModel):
|
|||||||
choices: List[ChatCompletionResponseChoice]
|
choices: List[ChatCompletionResponseChoice]
|
||||||
usage: UsageInfo
|
usage: UsageInfo
|
||||||
metadata: Optional[Dict[str, Any]] = None
|
metadata: Optional[Dict[str, Any]] = None
|
||||||
|
sglext: Optional[SglExt] = None
|
||||||
|
|
||||||
|
@model_serializer(mode="wrap")
|
||||||
|
def _serialize(self, handler):
|
||||||
|
data = handler(self)
|
||||||
|
if self.sglext is None:
|
||||||
|
data.pop("sglext", None)
|
||||||
|
return data
|
||||||
|
|
||||||
|
|
||||||
class DeltaMessage(BaseModel):
|
class DeltaMessage(BaseModel):
|
||||||
@@ -783,15 +827,12 @@ class DeltaMessage(BaseModel):
|
|||||||
reasoning_content: Optional[str] = None
|
reasoning_content: Optional[str] = None
|
||||||
tool_calls: Optional[List[ToolCall]] = Field(default=None, examples=[None])
|
tool_calls: Optional[List[ToolCall]] = Field(default=None, examples=[None])
|
||||||
hidden_states: Optional[object] = None
|
hidden_states: Optional[object] = None
|
||||||
sgl_ext: Optional[SglExt] = None
|
|
||||||
|
|
||||||
@model_serializer(mode="wrap")
|
@model_serializer(mode="wrap")
|
||||||
def _serialize(self, handler):
|
def _serialize(self, handler):
|
||||||
data = handler(self)
|
data = handler(self)
|
||||||
if self.hidden_states is None:
|
if self.hidden_states is None:
|
||||||
data.pop("hidden_states", None)
|
data.pop("hidden_states", None)
|
||||||
if self.sgl_ext is None:
|
|
||||||
data.pop("sgl_ext", None)
|
|
||||||
return data
|
return data
|
||||||
|
|
||||||
|
|
||||||
@@ -814,6 +855,14 @@ class ChatCompletionStreamResponse(BaseModel):
|
|||||||
model: str
|
model: str
|
||||||
choices: List[ChatCompletionResponseStreamChoice]
|
choices: List[ChatCompletionResponseStreamChoice]
|
||||||
usage: Optional[UsageInfo] = None
|
usage: Optional[UsageInfo] = None
|
||||||
|
sglext: Optional[SglExt] = None
|
||||||
|
|
||||||
|
@model_serializer(mode="wrap")
|
||||||
|
def _serialize(self, handler):
|
||||||
|
data = handler(self)
|
||||||
|
if self.sglext is None:
|
||||||
|
data.pop("sglext", None)
|
||||||
|
return data
|
||||||
|
|
||||||
|
|
||||||
class MultimodalEmbeddingInput(BaseModel):
|
class MultimodalEmbeddingInput(BaseModel):
|
||||||
|
|||||||
@@ -37,6 +37,7 @@ from sglang.srt.entrypoints.openai.protocol import (
|
|||||||
from sglang.srt.entrypoints.openai.serving_base import OpenAIServingBase
|
from sglang.srt.entrypoints.openai.serving_base import OpenAIServingBase
|
||||||
from sglang.srt.entrypoints.openai.usage_processor import UsageProcessor
|
from sglang.srt.entrypoints.openai.usage_processor import UsageProcessor
|
||||||
from sglang.srt.entrypoints.openai.utils import (
|
from sglang.srt.entrypoints.openai.utils import (
|
||||||
|
process_cached_tokens_details_from_ret,
|
||||||
process_hidden_states_from_ret,
|
process_hidden_states_from_ret,
|
||||||
process_routed_experts_from_ret,
|
process_routed_experts_from_ret,
|
||||||
to_openai_style_logprobs,
|
to_openai_style_logprobs,
|
||||||
@@ -813,25 +814,19 @@ class OpenAIServingChat(OpenAIServingBase):
|
|||||||
yield f"data: {hidden_states_chunk.model_dump_json()}\n\n"
|
yield f"data: {hidden_states_chunk.model_dump_json()}\n\n"
|
||||||
|
|
||||||
if request.return_routed_experts and routed_experts:
|
if request.return_routed_experts and routed_experts:
|
||||||
for index, choice_routed_experts in routed_experts.items():
|
# Get first non-None routed_experts value
|
||||||
if choice_routed_experts is not None:
|
first_routed_experts = next(
|
||||||
routed_experts_chunk = ChatCompletionStreamResponse(
|
(v for v in routed_experts.values() if v is not None), None
|
||||||
id=content["meta_info"]["id"],
|
)
|
||||||
created=int(time.time()),
|
if first_routed_experts is not None:
|
||||||
choices=[
|
routed_experts_chunk = ChatCompletionStreamResponse(
|
||||||
ChatCompletionResponseStreamChoice(
|
id=content["meta_info"]["id"],
|
||||||
index=index,
|
created=int(time.time()),
|
||||||
delta=DeltaMessage(
|
choices=[], # sglext is at response level
|
||||||
sgl_ext=SglExt(
|
model=request.model,
|
||||||
routed_experts=choice_routed_experts
|
sglext=SglExt(routed_experts=first_routed_experts),
|
||||||
)
|
)
|
||||||
),
|
yield f"data: {routed_experts_chunk.model_dump_json()}\n\n"
|
||||||
finish_reason=None,
|
|
||||||
)
|
|
||||||
],
|
|
||||||
model=request.model,
|
|
||||||
)
|
|
||||||
yield (f"data: {routed_experts_chunk.model_dump_json()}\n\n")
|
|
||||||
|
|
||||||
# Additional usage chunk
|
# Additional usage chunk
|
||||||
if request.stream_options and request.stream_options.include_usage:
|
if request.stream_options and request.stream_options.include_usage:
|
||||||
@@ -891,6 +886,19 @@ class OpenAIServingChat(OpenAIServingBase):
|
|||||||
"""Build chat completion response from generation results"""
|
"""Build chat completion response from generation results"""
|
||||||
choices = []
|
choices = []
|
||||||
|
|
||||||
|
# Build sglext at response level (from first ret_item, as these are per-request)
|
||||||
|
first_ret = ret[0]
|
||||||
|
routed_experts = process_routed_experts_from_ret(first_ret, request)
|
||||||
|
cached_tokens_details = process_cached_tokens_details_from_ret(
|
||||||
|
first_ret, request
|
||||||
|
)
|
||||||
|
response_sglext = None
|
||||||
|
if routed_experts or cached_tokens_details:
|
||||||
|
response_sglext = SglExt(
|
||||||
|
routed_experts=routed_experts,
|
||||||
|
cached_tokens_details=cached_tokens_details,
|
||||||
|
)
|
||||||
|
|
||||||
for idx, ret_item in enumerate(ret):
|
for idx, ret_item in enumerate(ret):
|
||||||
# Process logprobs
|
# Process logprobs
|
||||||
choice_logprobs = None
|
choice_logprobs = None
|
||||||
@@ -899,7 +907,6 @@ class OpenAIServingChat(OpenAIServingBase):
|
|||||||
|
|
||||||
# Handle hidden states
|
# Handle hidden states
|
||||||
hidden_states = process_hidden_states_from_ret(ret_item, request)
|
hidden_states = process_hidden_states_from_ret(ret_item, request)
|
||||||
routed_experts = process_routed_experts_from_ret(ret_item, request)
|
|
||||||
|
|
||||||
finish_reason = ret_item["meta_info"]["finish_reason"]
|
finish_reason = ret_item["meta_info"]["finish_reason"]
|
||||||
text = ret_item["text"]
|
text = ret_item["text"]
|
||||||
@@ -960,9 +967,6 @@ class OpenAIServingChat(OpenAIServingBase):
|
|||||||
else None
|
else None
|
||||||
),
|
),
|
||||||
hidden_states=hidden_states,
|
hidden_states=hidden_states,
|
||||||
sgl_ext=(
|
|
||||||
SglExt(routed_experts=routed_experts) if routed_experts else None
|
|
||||||
),
|
|
||||||
)
|
)
|
||||||
choices.append(choice_data)
|
choices.append(choice_data)
|
||||||
|
|
||||||
@@ -980,6 +984,7 @@ class OpenAIServingChat(OpenAIServingBase):
|
|||||||
choices=choices,
|
choices=choices,
|
||||||
usage=usage,
|
usage=usage,
|
||||||
metadata={"weight_version": ret[0]["meta_info"]["weight_version"]},
|
metadata={"weight_version": ret[0]["meta_info"]["weight_version"]},
|
||||||
|
sglext=response_sglext,
|
||||||
)
|
)
|
||||||
|
|
||||||
def _process_logprobs_tokens(
|
def _process_logprobs_tokens(
|
||||||
|
|||||||
@@ -19,6 +19,7 @@ from sglang.srt.entrypoints.openai.protocol import (
|
|||||||
from sglang.srt.entrypoints.openai.serving_base import OpenAIServingBase
|
from sglang.srt.entrypoints.openai.serving_base import OpenAIServingBase
|
||||||
from sglang.srt.entrypoints.openai.usage_processor import UsageProcessor
|
from sglang.srt.entrypoints.openai.usage_processor import UsageProcessor
|
||||||
from sglang.srt.entrypoints.openai.utils import (
|
from sglang.srt.entrypoints.openai.utils import (
|
||||||
|
process_cached_tokens_details_from_ret,
|
||||||
process_hidden_states_from_ret,
|
process_hidden_states_from_ret,
|
||||||
process_routed_experts_from_ret,
|
process_routed_experts_from_ret,
|
||||||
to_openai_style_logprobs,
|
to_openai_style_logprobs,
|
||||||
@@ -332,25 +333,20 @@ class OpenAIServingCompletion(OpenAIServingBase):
|
|||||||
yield f"data: {hidden_states_chunk.model_dump_json()}\n\n"
|
yield f"data: {hidden_states_chunk.model_dump_json()}\n\n"
|
||||||
|
|
||||||
if request.return_routed_experts and routed_experts:
|
if request.return_routed_experts and routed_experts:
|
||||||
for index, choice_routed_experts in routed_experts.items():
|
# Get first non-None routed_experts value
|
||||||
if choice_routed_experts is not None:
|
first_routed_experts = next(
|
||||||
routed_experts_chunk = CompletionStreamResponse(
|
(v for v in routed_experts.values() if v is not None), None
|
||||||
id=content["meta_info"]["id"],
|
)
|
||||||
created=created,
|
if first_routed_experts is not None:
|
||||||
object="text_completion",
|
routed_experts_chunk = CompletionStreamResponse(
|
||||||
choices=[
|
id=content["meta_info"]["id"],
|
||||||
CompletionResponseStreamChoice(
|
created=created,
|
||||||
index=index,
|
object="text_completion",
|
||||||
text="",
|
choices=[], # sglext is at response level
|
||||||
sgl_ext=SglExt(
|
model=request.model,
|
||||||
routed_experts=choice_routed_experts
|
sglext=SglExt(routed_experts=first_routed_experts),
|
||||||
),
|
)
|
||||||
finish_reason=None,
|
yield f"data: {routed_experts_chunk.model_dump_json()}\n\n"
|
||||||
)
|
|
||||||
],
|
|
||||||
model=request.model,
|
|
||||||
)
|
|
||||||
yield (f"data: {routed_experts_chunk.model_dump_json()}\n\n")
|
|
||||||
|
|
||||||
# Handle final usage chunk
|
# Handle final usage chunk
|
||||||
if request.stream_options and request.stream_options.include_usage:
|
if request.stream_options and request.stream_options.include_usage:
|
||||||
@@ -419,6 +415,19 @@ class OpenAIServingCompletion(OpenAIServingBase):
|
|||||||
echo_prompts = self._prepare_echo_prompts(request)
|
echo_prompts = self._prepare_echo_prompts(request)
|
||||||
echo = True
|
echo = True
|
||||||
|
|
||||||
|
# Build sglext at response level (from first ret_item, as these are per-request)
|
||||||
|
first_ret = ret[0]
|
||||||
|
routed_experts = process_routed_experts_from_ret(first_ret, request)
|
||||||
|
cached_tokens_details = process_cached_tokens_details_from_ret(
|
||||||
|
first_ret, request
|
||||||
|
)
|
||||||
|
response_sglext = None
|
||||||
|
if routed_experts or cached_tokens_details:
|
||||||
|
response_sglext = SglExt(
|
||||||
|
routed_experts=routed_experts,
|
||||||
|
cached_tokens_details=cached_tokens_details,
|
||||||
|
)
|
||||||
|
|
||||||
for idx, ret_item in enumerate(ret):
|
for idx, ret_item in enumerate(ret):
|
||||||
text = ret_item["text"]
|
text = ret_item["text"]
|
||||||
|
|
||||||
@@ -450,7 +459,6 @@ class OpenAIServingCompletion(OpenAIServingBase):
|
|||||||
|
|
||||||
# Handle hidden states
|
# Handle hidden states
|
||||||
hidden_states = process_hidden_states_from_ret(ret_item, request)
|
hidden_states = process_hidden_states_from_ret(ret_item, request)
|
||||||
routed_experts = process_routed_experts_from_ret(ret_item, request)
|
|
||||||
|
|
||||||
finish_reason = ret_item["meta_info"]["finish_reason"]
|
finish_reason = ret_item["meta_info"]["finish_reason"]
|
||||||
|
|
||||||
@@ -465,9 +473,6 @@ class OpenAIServingCompletion(OpenAIServingBase):
|
|||||||
else None
|
else None
|
||||||
),
|
),
|
||||||
hidden_states=hidden_states,
|
hidden_states=hidden_states,
|
||||||
sgl_ext=(
|
|
||||||
SglExt(routed_experts=routed_experts) if routed_experts else None
|
|
||||||
),
|
|
||||||
)
|
)
|
||||||
choices.append(choice_data)
|
choices.append(choice_data)
|
||||||
|
|
||||||
@@ -484,6 +489,7 @@ class OpenAIServingCompletion(OpenAIServingBase):
|
|||||||
choices=choices,
|
choices=choices,
|
||||||
usage=usage,
|
usage=usage,
|
||||||
metadata={"weight_version": ret[0]["meta_info"]["weight_version"]},
|
metadata={"weight_version": ret[0]["meta_info"]["weight_version"]},
|
||||||
|
sglext=response_sglext,
|
||||||
)
|
)
|
||||||
|
|
||||||
def _get_echo_text(self, request: CompletionRequest, index: int) -> str:
|
def _get_echo_text(self, request: CompletionRequest, index: int) -> str:
|
||||||
|
|||||||
@@ -2,7 +2,7 @@ from __future__ import annotations
|
|||||||
|
|
||||||
from typing import Any, Dict, List, Mapping, Optional, final
|
from typing import Any, Dict, List, Mapping, Optional, final
|
||||||
|
|
||||||
from sglang.srt.entrypoints.openai.protocol import UsageInfo
|
from sglang.srt.entrypoints.openai.protocol import PromptTokensDetails, UsageInfo
|
||||||
|
|
||||||
|
|
||||||
@final
|
@final
|
||||||
@@ -10,9 +10,9 @@ class UsageProcessor:
|
|||||||
"""Stateless helpers that turn raw token counts into a UsageInfo."""
|
"""Stateless helpers that turn raw token counts into a UsageInfo."""
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _details_if_cached(count: int) -> Optional[Dict[str, int]]:
|
def _details_if_cached(count: int) -> Optional[PromptTokensDetails]:
|
||||||
"""Return {"cached_tokens": N} only when N > 0 (keeps JSON slim)."""
|
"""Return PromptTokensDetails only when count > 0 (keeps JSON slim)."""
|
||||||
return {"cached_tokens": count} if count > 0 else None
|
return PromptTokensDetails(cached_tokens=count) if count > 0 else None
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def calculate_response_usage(
|
def calculate_response_usage(
|
||||||
@@ -73,7 +73,7 @@ class UsageProcessor:
|
|||||||
def calculate_token_usage(
|
def calculate_token_usage(
|
||||||
prompt_tokens: int,
|
prompt_tokens: int,
|
||||||
completion_tokens: int,
|
completion_tokens: int,
|
||||||
cached_tokens: Optional[Dict[str, int]] = None,
|
cached_tokens: Optional[PromptTokensDetails] = None,
|
||||||
) -> UsageInfo:
|
) -> UsageInfo:
|
||||||
"""Calculate token usage information"""
|
"""Calculate token usage information"""
|
||||||
return UsageInfo(
|
return UsageInfo(
|
||||||
|
|||||||
@@ -2,6 +2,7 @@ import logging
|
|||||||
from typing import Any, Dict, List, Optional, Union
|
from typing import Any, Dict, List, Optional, Union
|
||||||
|
|
||||||
from sglang.srt.entrypoints.openai.protocol import (
|
from sglang.srt.entrypoints.openai.protocol import (
|
||||||
|
CachedTokensDetails,
|
||||||
ChatCompletionRequest,
|
ChatCompletionRequest,
|
||||||
CompletionRequest,
|
CompletionRequest,
|
||||||
LogProbs,
|
LogProbs,
|
||||||
@@ -83,3 +84,33 @@ def process_routed_experts_from_ret(
|
|||||||
if not getattr(request, "return_routed_experts", False):
|
if not getattr(request, "return_routed_experts", False):
|
||||||
return None
|
return None
|
||||||
return ret_item["meta_info"].get("routed_experts", None)
|
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
|
||||||
|
if "storage" in details:
|
||||||
|
return CachedTokensDetails(
|
||||||
|
device=details.get("device", 0),
|
||||||
|
host=details.get("host", 0),
|
||||||
|
storage=details.get("storage", 0),
|
||||||
|
storage_backend=details.get("storage_backend"),
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
return CachedTokensDetails(
|
||||||
|
device=details.get("device", 0),
|
||||||
|
host=details.get("host", 0),
|
||||||
|
)
|
||||||
|
|||||||
@@ -375,6 +375,7 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin):
|
|||||||
prompt_tokens=recv_obj.prompt_tokens,
|
prompt_tokens=recv_obj.prompt_tokens,
|
||||||
completion_tokens=recv_obj.completion_tokens,
|
completion_tokens=recv_obj.completion_tokens,
|
||||||
cached_tokens=recv_obj.cached_tokens,
|
cached_tokens=recv_obj.cached_tokens,
|
||||||
|
cached_tokens_details=recv_obj.cached_tokens_details,
|
||||||
spec_verify_ct=recv_obj.spec_verify_ct,
|
spec_verify_ct=recv_obj.spec_verify_ct,
|
||||||
spec_accepted_tokens=recv_obj.spec_accepted_tokens,
|
spec_accepted_tokens=recv_obj.spec_accepted_tokens,
|
||||||
input_token_logprobs_val=recv_obj.input_token_logprobs_val,
|
input_token_logprobs_val=recv_obj.input_token_logprobs_val,
|
||||||
|
|||||||
@@ -1005,6 +1005,8 @@ class BatchTokenIDOutput(
|
|||||||
load: GetLoadReqOutput = None
|
load: GetLoadReqOutput = None
|
||||||
# Customized info
|
# Customized info
|
||||||
customized_info: Optional[Dict[str, List[Any]]] = None
|
customized_info: Optional[Dict[str, List[Any]]] = None
|
||||||
|
# Detailed breakdown of cached tokens by source (device/host/storage)
|
||||||
|
cached_tokens_details: Optional[List[Optional[Dict[str, Any]]]] = None
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
@@ -1094,6 +1096,8 @@ class BatchStrOutput(
|
|||||||
|
|
||||||
# Customized info
|
# Customized info
|
||||||
customized_info: Optional[Dict[str, List[Any]]] = None
|
customized_info: Optional[Dict[str, List[Any]]] = None
|
||||||
|
# Detailed breakdown of cached tokens by source (device/host/storage)
|
||||||
|
cached_tokens_details: Optional[List[Optional[Dict[str, Any]]]] = None
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
@@ -1119,6 +1123,8 @@ class BatchMultimodalOutput(BaseBatchReq):
|
|||||||
placeholder_tokens_val: List[Optional[List[int]]]
|
placeholder_tokens_val: List[Optional[List[int]]]
|
||||||
|
|
||||||
return_bytes: List[bool]
|
return_bytes: List[bool]
|
||||||
|
# Detailed breakdown of cached tokens by source (device/host/storage)
|
||||||
|
cached_tokens_details: Optional[List[Optional[Dict[str, Any]]]] = None
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
@@ -1136,6 +1142,8 @@ class BatchEmbeddingOutput(BaseBatchReq, RequestTimingMetricsMixin):
|
|||||||
|
|
||||||
# Number of times each request was retracted.
|
# Number of times each request was retracted.
|
||||||
retraction_counts: List[int]
|
retraction_counts: List[int]
|
||||||
|
# Detailed breakdown of cached tokens by source (device/host/storage)
|
||||||
|
cached_tokens_details: Optional[List[Optional[Dict[str, Any]]]] = None
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
|
|||||||
@@ -154,6 +154,9 @@ def _handle_output_by_index(output, i):
|
|||||||
prompt_tokens=_extract_field_by_index(output, "prompt_tokens", i),
|
prompt_tokens=_extract_field_by_index(output, "prompt_tokens", i),
|
||||||
completion_tokens=_extract_field_by_index(output, "completion_tokens", i),
|
completion_tokens=_extract_field_by_index(output, "completion_tokens", i),
|
||||||
cached_tokens=_extract_field_by_index(output, "cached_tokens", i),
|
cached_tokens=_extract_field_by_index(output, "cached_tokens", i),
|
||||||
|
cached_tokens_details=_extract_field_by_index(
|
||||||
|
output, "cached_tokens_details", i
|
||||||
|
),
|
||||||
input_token_logprobs_val=_extract_field_by_index(
|
input_token_logprobs_val=_extract_field_by_index(
|
||||||
output, "input_token_logprobs_val", i, check_length=False
|
output, "input_token_logprobs_val", i, check_length=False
|
||||||
),
|
),
|
||||||
|
|||||||
@@ -660,6 +660,8 @@ class Req:
|
|||||||
self.last_node: Any = None
|
self.last_node: Any = None
|
||||||
self.last_host_node: Any = None
|
self.last_host_node: Any = None
|
||||||
self.host_hit_length = 0
|
self.host_hit_length = 0
|
||||||
|
# Tokens loaded from storage backend (L3) during prefetch for this request
|
||||||
|
self.storage_hit_length = 0
|
||||||
# The node to lock until for swa radix tree lock ref
|
# The node to lock until for swa radix tree lock ref
|
||||||
self.swa_uuid_for_lock: Optional[int] = None
|
self.swa_uuid_for_lock: Optional[int] = None
|
||||||
# The prefix length that is inserted into the tree cache
|
# The prefix length that is inserted into the tree cache
|
||||||
@@ -750,6 +752,14 @@ class Req:
|
|||||||
self.cached_tokens = 0
|
self.cached_tokens = 0
|
||||||
self.already_computed = 0
|
self.already_computed = 0
|
||||||
|
|
||||||
|
# Detailed breakdown of cached tokens by source (for HiCache)
|
||||||
|
self.cached_tokens_device = 0 # Tokens from device cache (GPU)
|
||||||
|
self.cached_tokens_host = 0 # Tokens from host cache (CPU memory)
|
||||||
|
self.cached_tokens_storage = 0 # Tokens from L3 storage backend
|
||||||
|
self._cache_breakdown_computed = (
|
||||||
|
False # Track if breakdown was already computed
|
||||||
|
)
|
||||||
|
|
||||||
# The number of verification forward passes in the speculative decoding.
|
# The number of verification forward passes in the speculative decoding.
|
||||||
# This is used to compute the average acceptance length per request.
|
# This is used to compute the average acceptance length per request.
|
||||||
self.spec_verify_ct = 0
|
self.spec_verify_ct = 0
|
||||||
@@ -1577,7 +1587,33 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
|||||||
# Only calculate cached_tokens once. Once retracted, the 'retracted_stain'
|
# Only calculate cached_tokens once. Once retracted, the 'retracted_stain'
|
||||||
# flag will always True
|
# flag will always True
|
||||||
if not req.retracted_stain:
|
if not req.retracted_stain:
|
||||||
req.cached_tokens += pre_len - req.already_computed
|
new_cached = pre_len - req.already_computed
|
||||||
|
req.cached_tokens += new_cached
|
||||||
|
|
||||||
|
# Calculate detailed breakdown of cached tokens by source (for HiCache)
|
||||||
|
# Only compute once on FIRST chunk - subsequent chunks in chunked prefill
|
||||||
|
# would incorrectly count previously computed tokens as cache hits.
|
||||||
|
if not req._cache_breakdown_computed:
|
||||||
|
# At this point, prefix_indices has been extended with host data
|
||||||
|
# via init_load_back in schedule_policy, so:
|
||||||
|
# - len(prefix_indices) = device_original + host_loaded
|
||||||
|
# - host_hit_length = total tokens from host cache (including storage-prefetched)
|
||||||
|
# - storage_hit_length = tokens loaded from storage backend (L3 hits)
|
||||||
|
# - device_portion = len(prefix_indices) - host_hit_length
|
||||||
|
#
|
||||||
|
# Storage hits are now tracked via scheduler after prefetch completes.
|
||||||
|
# storage_hit_length is set by scheduler.pop_prefetch_loaded_tokens()
|
||||||
|
host_total = req.host_hit_length
|
||||||
|
# Clamp storage to host_total to handle edge cases
|
||||||
|
storage_portion = min(host_total, req.storage_hit_length)
|
||||||
|
host_portion = host_total - storage_portion
|
||||||
|
device_portion = max(0, len(req.prefix_indices) - host_total)
|
||||||
|
|
||||||
|
req.cached_tokens_device = device_portion
|
||||||
|
req.cached_tokens_host = host_portion
|
||||||
|
req.cached_tokens_storage = storage_portion
|
||||||
|
req._cache_breakdown_computed = True
|
||||||
|
|
||||||
req.already_computed = seq_len
|
req.already_computed = seq_len
|
||||||
req.is_retracted = False
|
req.is_retracted = False
|
||||||
|
|
||||||
|
|||||||
@@ -1677,7 +1677,10 @@ class Scheduler(
|
|||||||
direction * recv_req.priority < direction * candidate_req.priority
|
direction * recv_req.priority < direction * candidate_req.priority
|
||||||
)
|
)
|
||||||
if abort_existing_req:
|
if abort_existing_req:
|
||||||
if self.enable_hierarchical_cache:
|
if self.enable_hicache_storage:
|
||||||
|
# Release prefetch events associated with the request
|
||||||
|
self.tree_cache.release_aborted_request(candidate_req.rid)
|
||||||
|
elif self.enable_hierarchical_cache:
|
||||||
self.tree_cache.terminate_prefetch(candidate_req.rid)
|
self.tree_cache.terminate_prefetch(candidate_req.rid)
|
||||||
self.waiting_queue.pop(idx)
|
self.waiting_queue.pop(idx)
|
||||||
req_to_abort = candidate_req
|
req_to_abort = candidate_req
|
||||||
@@ -1705,6 +1708,9 @@ class Scheduler(
|
|||||||
for req in self.waiting_queue:
|
for req in self.waiting_queue:
|
||||||
entry_time = req.time_stats.wait_queue_entry_time
|
entry_time = req.time_stats.wait_queue_entry_time
|
||||||
if 0 < entry_time < deadline:
|
if 0 < entry_time < deadline:
|
||||||
|
if self.enable_hicache_storage:
|
||||||
|
# Release prefetch events associated with the request
|
||||||
|
self.tree_cache.release_aborted_request(req.rid)
|
||||||
self.send_to_tokenizer.send_output(
|
self.send_to_tokenizer.send_output(
|
||||||
AbortReq(
|
AbortReq(
|
||||||
finished_reason={
|
finished_reason={
|
||||||
@@ -2024,6 +2030,10 @@ class Scheduler(
|
|||||||
if not prefetch_done:
|
if not prefetch_done:
|
||||||
# skip staging requests that are ongoing prefetch
|
# skip staging requests that are ongoing prefetch
|
||||||
continue
|
continue
|
||||||
|
# Pop the number of tokens loaded from storage (L3 hits)
|
||||||
|
req.storage_hit_length = self.tree_cache.pop_prefetch_loaded_tokens(
|
||||||
|
req.rid
|
||||||
|
)
|
||||||
|
|
||||||
req.init_next_round_input(self.tree_cache)
|
req.init_next_round_input(self.tree_cache)
|
||||||
res = adder.add_one_req(
|
res = adder.add_one_req(
|
||||||
|
|||||||
@@ -46,6 +46,45 @@ class SchedulerOutputProcessorMixin:
|
|||||||
We put them into a separate file to make the `scheduler.py` shorter.
|
We put them into a separate file to make the `scheduler.py` shorter.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
def _get_storage_backend_type(self) -> str:
|
||||||
|
"""Get storage backend type from tree_cache."""
|
||||||
|
storage_backend_type = "none"
|
||||||
|
cache_controller = getattr(self.tree_cache, "cache_controller", None)
|
||||||
|
if cache_controller and hasattr(cache_controller, "storage_backend"):
|
||||||
|
storage_backend = cache_controller.storage_backend
|
||||||
|
if storage_backend is not None:
|
||||||
|
storage_backend_type = type(storage_backend).__name__
|
||||||
|
return storage_backend_type
|
||||||
|
|
||||||
|
def _get_cached_tokens_details(self, req: Req) -> Optional[dict]:
|
||||||
|
"""Get detailed cache breakdown for a request, if available.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
- None if HiCache is not enabled
|
||||||
|
- {"device": X, "host": Y} if HiCache enabled but L3 storage is not
|
||||||
|
- {"device": X, "host": Y, "storage": Z, "storage_backend": "..."} if L3 enabled
|
||||||
|
"""
|
||||||
|
# Only show details if HiCache is enabled
|
||||||
|
if not getattr(self, "enable_hierarchical_cache", False):
|
||||||
|
return None
|
||||||
|
|
||||||
|
# Only show if there are any cached tokens
|
||||||
|
if (
|
||||||
|
req.cached_tokens_device > 0
|
||||||
|
or req.cached_tokens_host > 0
|
||||||
|
or req.cached_tokens_storage > 0
|
||||||
|
):
|
||||||
|
details = {
|
||||||
|
"device": req.cached_tokens_device,
|
||||||
|
"host": req.cached_tokens_host,
|
||||||
|
}
|
||||||
|
# Only include storage fields if L3 storage is enabled
|
||||||
|
if getattr(self, "enable_hicache_storage", False):
|
||||||
|
details["storage"] = req.cached_tokens_storage
|
||||||
|
details["storage_backend"] = self._get_storage_backend_type()
|
||||||
|
return details
|
||||||
|
return None
|
||||||
|
|
||||||
def process_batch_result_prebuilt(self: Scheduler, batch: ScheduleBatch):
|
def process_batch_result_prebuilt(self: Scheduler, batch: ScheduleBatch):
|
||||||
assert self.disaggregation_mode == DisaggregationMode.DECODE
|
assert self.disaggregation_mode == DisaggregationMode.DECODE
|
||||||
for req in batch.reqs:
|
for req in batch.reqs:
|
||||||
@@ -893,6 +932,7 @@ class SchedulerOutputProcessorMixin:
|
|||||||
prompt_tokens = []
|
prompt_tokens = []
|
||||||
completion_tokens = []
|
completion_tokens = []
|
||||||
cached_tokens = []
|
cached_tokens = []
|
||||||
|
cached_tokens_details = [] # Detailed breakdown by cache source
|
||||||
spec_verify_ct = []
|
spec_verify_ct = []
|
||||||
spec_accepted_tokens = []
|
spec_accepted_tokens = []
|
||||||
retraction_counts = []
|
retraction_counts = []
|
||||||
@@ -1005,6 +1045,10 @@ class SchedulerOutputProcessorMixin:
|
|||||||
prompt_tokens.append(len(req.origin_input_ids))
|
prompt_tokens.append(len(req.origin_input_ids))
|
||||||
completion_tokens.append(len(output_ids_))
|
completion_tokens.append(len(output_ids_))
|
||||||
cached_tokens.append(req.cached_tokens)
|
cached_tokens.append(req.cached_tokens)
|
||||||
|
|
||||||
|
# Collect detailed cache breakdown if available
|
||||||
|
cached_tokens_details.append(self._get_cached_tokens_details(req))
|
||||||
|
|
||||||
retraction_counts.append(req.retraction_count)
|
retraction_counts.append(req.retraction_count)
|
||||||
|
|
||||||
queue_times.append(req.time_stats.get_queueing_time())
|
queue_times.append(req.time_stats.get_queueing_time())
|
||||||
@@ -1138,6 +1182,7 @@ class SchedulerOutputProcessorMixin:
|
|||||||
prompt_tokens=prompt_tokens,
|
prompt_tokens=prompt_tokens,
|
||||||
completion_tokens=completion_tokens,
|
completion_tokens=completion_tokens,
|
||||||
cached_tokens=cached_tokens,
|
cached_tokens=cached_tokens,
|
||||||
|
cached_tokens_details=cached_tokens_details,
|
||||||
input_token_logprobs_val=input_token_logprobs_val,
|
input_token_logprobs_val=input_token_logprobs_val,
|
||||||
input_token_logprobs_idx=input_token_logprobs_idx,
|
input_token_logprobs_idx=input_token_logprobs_idx,
|
||||||
output_token_logprobs_val=output_token_logprobs_val,
|
output_token_logprobs_val=output_token_logprobs_val,
|
||||||
@@ -1169,6 +1214,7 @@ class SchedulerOutputProcessorMixin:
|
|||||||
embeddings = []
|
embeddings = []
|
||||||
prompt_tokens = []
|
prompt_tokens = []
|
||||||
cached_tokens = []
|
cached_tokens = []
|
||||||
|
cached_tokens_details = [] # Detailed breakdown by cache source
|
||||||
queue_times = []
|
queue_times = []
|
||||||
forward_entry_times = []
|
forward_entry_times = []
|
||||||
prefill_launch_delays = []
|
prefill_launch_delays = []
|
||||||
@@ -1184,6 +1230,9 @@ class SchedulerOutputProcessorMixin:
|
|||||||
prompt_tokens.append(len(req.origin_input_ids))
|
prompt_tokens.append(len(req.origin_input_ids))
|
||||||
cached_tokens.append(req.cached_tokens)
|
cached_tokens.append(req.cached_tokens)
|
||||||
|
|
||||||
|
# Collect detailed cache breakdown if available
|
||||||
|
cached_tokens_details.append(self._get_cached_tokens_details(req))
|
||||||
|
|
||||||
queue_times.append(req.time_stats.get_queueing_time())
|
queue_times.append(req.time_stats.get_queueing_time())
|
||||||
forward_entry_times.append(req.time_stats.forward_entry_time)
|
forward_entry_times.append(req.time_stats.forward_entry_time)
|
||||||
|
|
||||||
@@ -1208,6 +1257,7 @@ class SchedulerOutputProcessorMixin:
|
|||||||
embeddings=embeddings,
|
embeddings=embeddings,
|
||||||
prompt_tokens=prompt_tokens,
|
prompt_tokens=prompt_tokens,
|
||||||
cached_tokens=cached_tokens,
|
cached_tokens=cached_tokens,
|
||||||
|
cached_tokens_details=cached_tokens_details,
|
||||||
placeholder_tokens_idx=None,
|
placeholder_tokens_idx=None,
|
||||||
placeholder_tokens_val=None,
|
placeholder_tokens_val=None,
|
||||||
retraction_counts=retraction_counts,
|
retraction_counts=retraction_counts,
|
||||||
|
|||||||
@@ -1536,6 +1536,14 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
|
|||||||
"cached_tokens": recv_obj.cached_tokens[i],
|
"cached_tokens": recv_obj.cached_tokens[i],
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
|
# Add detailed cache breakdown if available
|
||||||
|
if (
|
||||||
|
hasattr(recv_obj, "cached_tokens_details")
|
||||||
|
and recv_obj.cached_tokens_details
|
||||||
|
):
|
||||||
|
meta_info["cached_tokens_details"] = recv_obj.cached_tokens_details[
|
||||||
|
i
|
||||||
|
]
|
||||||
|
|
||||||
if getattr(recv_obj, "output_hidden_states", None):
|
if getattr(recv_obj, "output_hidden_states", None):
|
||||||
meta_info["hidden_states"] = recv_obj.output_hidden_states[i]
|
meta_info["hidden_states"] = recv_obj.output_hidden_states[i]
|
||||||
@@ -1974,6 +1982,14 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
|
|||||||
else 0
|
else 0
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# Get detailed cache breakdown if available
|
||||||
|
cached_tokens_details = None
|
||||||
|
if (
|
||||||
|
hasattr(recv_obj, "cached_tokens_details")
|
||||||
|
and recv_obj.cached_tokens_details
|
||||||
|
):
|
||||||
|
cached_tokens_details = recv_obj.cached_tokens_details[i]
|
||||||
|
|
||||||
self.metrics_collector.observe_one_finished_request(
|
self.metrics_collector.observe_one_finished_request(
|
||||||
labels,
|
labels,
|
||||||
recv_obj.prompt_tokens[i],
|
recv_obj.prompt_tokens[i],
|
||||||
@@ -1982,6 +1998,7 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
|
|||||||
state.finished_time - state.created_time,
|
state.finished_time - state.created_time,
|
||||||
self._request_has_grammar(state.obj),
|
self._request_has_grammar(state.obj),
|
||||||
retraction_count,
|
retraction_count,
|
||||||
|
cached_tokens_details,
|
||||||
)
|
)
|
||||||
|
|
||||||
def dump_requests(self, state: ReqState, out_dict: dict):
|
def dump_requests(self, state: ReqState, out_dict: dict):
|
||||||
|
|||||||
@@ -151,6 +151,9 @@ class HiRadixCache(RadixCache):
|
|||||||
# record the ongoing prefetch requests
|
# record the ongoing prefetch requests
|
||||||
self.ongoing_prefetch = {}
|
self.ongoing_prefetch = {}
|
||||||
self.ongoing_backup = {}
|
self.ongoing_backup = {}
|
||||||
|
# track per-request tokens loaded from storage (L3 hits)
|
||||||
|
# key: request_id, value: number of tokens actually loaded from storage
|
||||||
|
self.prefetch_loaded_tokens_by_reqid: dict[str, int] = {}
|
||||||
# todo: dynamically adjust the threshold
|
# todo: dynamically adjust the threshold
|
||||||
self.write_through_threshold = (
|
self.write_through_threshold = (
|
||||||
1 if server_args.hicache_write_policy == "write_through" else 2
|
1 if server_args.hicache_write_policy == "write_through" else 2
|
||||||
@@ -572,6 +575,8 @@ class HiRadixCache(RadixCache):
|
|||||||
TreeNode.counter = 0
|
TreeNode.counter = 0
|
||||||
self.cache_controller.reset()
|
self.cache_controller.reset()
|
||||||
self.token_to_kv_pool_host.clear()
|
self.token_to_kv_pool_host.clear()
|
||||||
|
# Clear per-request tracking dicts
|
||||||
|
self.prefetch_loaded_tokens_by_reqid.clear()
|
||||||
self.evictable_host_leaves.clear()
|
self.evictable_host_leaves.clear()
|
||||||
super().reset()
|
super().reset()
|
||||||
|
|
||||||
@@ -1089,10 +1094,12 @@ class HiRadixCache(RadixCache):
|
|||||||
del self.ongoing_prefetch[req_id]
|
del self.ongoing_prefetch[req_id]
|
||||||
self.cache_controller.prefetch_tokens_occupied -= len(token_ids)
|
self.cache_controller.prefetch_tokens_occupied -= len(token_ids)
|
||||||
|
|
||||||
|
# Track tokens actually loaded from storage for this request (L3 hits)
|
||||||
|
loaded_from_storage = min_completed_tokens - matched_length
|
||||||
|
self.prefetch_loaded_tokens_by_reqid[req_id] = loaded_from_storage
|
||||||
|
|
||||||
if self.enable_storage_metrics:
|
if self.enable_storage_metrics:
|
||||||
self.storage_metrics_collector.log_prefetched_tokens(
|
self.storage_metrics_collector.log_prefetched_tokens(loaded_from_storage)
|
||||||
min_completed_tokens - matched_length
|
|
||||||
)
|
|
||||||
|
|
||||||
return True
|
return True
|
||||||
|
|
||||||
@@ -1105,6 +1112,14 @@ class HiRadixCache(RadixCache):
|
|||||||
return
|
return
|
||||||
operation.mark_terminate()
|
operation.mark_terminate()
|
||||||
|
|
||||||
|
def pop_prefetch_loaded_tokens(self, req_id: str) -> int:
|
||||||
|
"""
|
||||||
|
Pop and return the number of tokens loaded from storage for a request.
|
||||||
|
Returns 0 if no prefetch was done or was revoked.
|
||||||
|
This should be called after check_prefetch_progress() returns True.
|
||||||
|
"""
|
||||||
|
return self.prefetch_loaded_tokens_by_reqid.pop(req_id, 0)
|
||||||
|
|
||||||
def match_prefix(self, params: MatchPrefixParams):
|
def match_prefix(self, params: MatchPrefixParams):
|
||||||
key = params.key
|
key = params.key
|
||||||
empty_value = torch.empty((0,), dtype=torch.int64, device=self.device)
|
empty_value = torch.empty((0,), dtype=torch.int64, device=self.device)
|
||||||
@@ -1357,6 +1372,9 @@ class HiRadixCache(RadixCache):
|
|||||||
return InsertResult(prefix_len=total_prefix_length)
|
return InsertResult(prefix_len=total_prefix_length)
|
||||||
|
|
||||||
def release_aborted_request(self, rid: str):
|
def release_aborted_request(self, rid: str):
|
||||||
|
# Clean up storage hit tracking for aborted request
|
||||||
|
self.prefetch_loaded_tokens_by_reqid.pop(rid, None)
|
||||||
|
|
||||||
if rid not in self.ongoing_prefetch:
|
if rid not in self.ongoing_prefetch:
|
||||||
return
|
return
|
||||||
|
|
||||||
|
|||||||
@@ -17,7 +17,7 @@ import logging
|
|||||||
import os
|
import os
|
||||||
import time
|
import time
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
from typing import Dict, List, Optional, Union
|
from typing import Any, Dict, List, Optional, Union
|
||||||
|
|
||||||
from sglang.srt.disaggregation.utils import DisaggregationMode
|
from sglang.srt.disaggregation.utils import DisaggregationMode
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
@@ -1148,8 +1148,8 @@ class TokenizerMetricsCollector:
|
|||||||
|
|
||||||
self.cached_tokens_total = Counter(
|
self.cached_tokens_total = Counter(
|
||||||
name="sglang:cached_tokens_total",
|
name="sglang:cached_tokens_total",
|
||||||
documentation="Number of cached prompt tokens.",
|
documentation="Number of cached prompt tokens by source (device/host/storage).",
|
||||||
labelnames=labels.keys(),
|
labelnames=list(labels.keys()) + ["cache_source"],
|
||||||
)
|
)
|
||||||
|
|
||||||
self.num_requests_total = Counter(
|
self.num_requests_total = Counter(
|
||||||
@@ -1303,11 +1303,36 @@ class TokenizerMetricsCollector:
|
|||||||
e2e_latency: float,
|
e2e_latency: float,
|
||||||
has_grammar: bool,
|
has_grammar: bool,
|
||||||
retraction_count: int,
|
retraction_count: int,
|
||||||
|
cached_tokens_details: Optional[Dict[str, Any]] = None,
|
||||||
):
|
):
|
||||||
self.prompt_tokens_total.labels(**labels).inc(prompt_tokens)
|
self.prompt_tokens_total.labels(**labels).inc(prompt_tokens)
|
||||||
self.generation_tokens_total.labels(**labels).inc(generation_tokens)
|
self.generation_tokens_total.labels(**labels).inc(generation_tokens)
|
||||||
|
|
||||||
|
# Report cached tokens with detailed source breakdown
|
||||||
if cached_tokens > 0:
|
if cached_tokens > 0:
|
||||||
self.cached_tokens_total.labels(**labels).inc(cached_tokens)
|
if cached_tokens_details:
|
||||||
|
# Report by cache source (device/host, and storage if L3 enabled)
|
||||||
|
def report_cache_source(source: str, value: int):
|
||||||
|
if value > 0:
|
||||||
|
source_labels = {**labels, "cache_source": source}
|
||||||
|
self.cached_tokens_total.labels(**source_labels).inc(value)
|
||||||
|
|
||||||
|
report_cache_source("device", cached_tokens_details.get("device", 0))
|
||||||
|
report_cache_source("host", cached_tokens_details.get("host", 0))
|
||||||
|
|
||||||
|
# Storage fields are only present when L3 storage backend is enabled
|
||||||
|
if "storage" in cached_tokens_details:
|
||||||
|
storage_tokens = cached_tokens_details.get("storage", 0)
|
||||||
|
if storage_tokens > 0:
|
||||||
|
backend = (
|
||||||
|
cached_tokens_details.get("storage_backend") or "unknown"
|
||||||
|
)
|
||||||
|
report_cache_source(f"storage_{backend}", storage_tokens)
|
||||||
|
else:
|
||||||
|
# Fallback for backward compatibility
|
||||||
|
labels_total = {**labels, "cache_source": "total"}
|
||||||
|
self.cached_tokens_total.labels(**labels_total).inc(cached_tokens)
|
||||||
|
|
||||||
self.num_requests_total.labels(**labels).inc(1)
|
self.num_requests_total.labels(**labels).inc(1)
|
||||||
if has_grammar:
|
if has_grammar:
|
||||||
self.num_so_requests_total.labels(**labels).inc(1)
|
self.num_so_requests_total.labels(**labels).inc(1)
|
||||||
|
|||||||
Reference in New Issue
Block a user