[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]
|
||||
|
||||
|
||||
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):
|
||||
prompt_tokens: int = 0
|
||||
total_tokens: int = 0
|
||||
completion_tokens: Optional[int] = 0
|
||||
# only used to return cached tokens when --enable-cache-report is set
|
||||
prompt_tokens_details: Optional[Dict[str, int]] = None
|
||||
# Used to return cached tokens info when --enable-cache-report is set
|
||||
prompt_tokens_details: Optional[PromptTokensDetails] = None
|
||||
reasoning_tokens: Optional[int] = 0
|
||||
|
||||
|
||||
@@ -233,6 +259,7 @@ class CompletionRequest(BaseModel):
|
||||
user: Optional[str] = None
|
||||
return_hidden_states: 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.
|
||||
top_k: int = -1
|
||||
@@ -289,6 +316,7 @@ class SglExt(BaseModel):
|
||||
"""
|
||||
|
||||
routed_experts: Optional[str] = None
|
||||
cached_tokens_details: Optional[CachedTokensDetails] = None
|
||||
|
||||
@model_serializer(mode="wrap")
|
||||
def _serialize(self, handler):
|
||||
@@ -304,15 +332,12 @@ class CompletionResponseChoice(BaseModel):
|
||||
finish_reason: Optional[Literal["stop", "length", "content_filter", "abort"]] = None
|
||||
matched_stop: Union[None, int, str] = None
|
||||
hidden_states: Optional[object] = None
|
||||
sgl_ext: Optional[SglExt] = None
|
||||
|
||||
@model_serializer(mode="wrap")
|
||||
def _serialize(self, handler):
|
||||
data = handler(self)
|
||||
if self.hidden_states is None:
|
||||
data.pop("hidden_states", None)
|
||||
if self.sgl_ext is None:
|
||||
data.pop("sgl_ext", None)
|
||||
return data
|
||||
|
||||
|
||||
@@ -324,6 +349,14 @@ class CompletionResponse(BaseModel):
|
||||
choices: List[CompletionResponseChoice]
|
||||
usage: UsageInfo
|
||||
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):
|
||||
@@ -333,15 +366,12 @@ class CompletionResponseStreamChoice(BaseModel):
|
||||
finish_reason: Optional[Literal["stop", "length", "content_filter", "abort"]] = None
|
||||
matched_stop: Union[None, int, str] = None
|
||||
hidden_states: Optional[object] = None
|
||||
sgl_ext: Optional[SglExt] = None
|
||||
|
||||
@model_serializer(mode="wrap")
|
||||
def _serialize(self, handler):
|
||||
data = handler(self)
|
||||
if self.hidden_states is None:
|
||||
data.pop("hidden_states", None)
|
||||
if self.sgl_ext is None:
|
||||
data.pop("sgl_ext", None)
|
||||
return data
|
||||
|
||||
|
||||
@@ -352,6 +382,14 @@ class CompletionStreamResponse(BaseModel):
|
||||
model: str
|
||||
choices: List[CompletionResponseStreamChoice]
|
||||
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):
|
||||
@@ -526,6 +564,7 @@ class ChatCompletionRequest(BaseModel):
|
||||
) # noqa
|
||||
return_hidden_states: bool = False
|
||||
return_routed_experts: bool = False
|
||||
return_cached_tokens_details: bool = False
|
||||
reasoning_effort: Optional[Literal["low", "medium", "high"]] = Field(
|
||||
default="medium",
|
||||
description="Constrains effort on reasoning for reasoning models. "
|
||||
@@ -755,15 +794,12 @@ class ChatCompletionResponseChoice(BaseModel):
|
||||
] = None
|
||||
matched_stop: Union[None, int, str] = None
|
||||
hidden_states: Optional[object] = None
|
||||
sgl_ext: Optional[SglExt] = None
|
||||
|
||||
@model_serializer(mode="wrap")
|
||||
def _serialize(self, handler):
|
||||
data = handler(self)
|
||||
if self.hidden_states is None:
|
||||
data.pop("hidden_states", None)
|
||||
if self.sgl_ext is None:
|
||||
data.pop("sgl_ext", None)
|
||||
return data
|
||||
|
||||
|
||||
@@ -775,6 +811,14 @@ class ChatCompletionResponse(BaseModel):
|
||||
choices: List[ChatCompletionResponseChoice]
|
||||
usage: UsageInfo
|
||||
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):
|
||||
@@ -783,15 +827,12 @@ class DeltaMessage(BaseModel):
|
||||
reasoning_content: Optional[str] = None
|
||||
tool_calls: Optional[List[ToolCall]] = Field(default=None, examples=[None])
|
||||
hidden_states: Optional[object] = None
|
||||
sgl_ext: Optional[SglExt] = None
|
||||
|
||||
@model_serializer(mode="wrap")
|
||||
def _serialize(self, handler):
|
||||
data = handler(self)
|
||||
if self.hidden_states is None:
|
||||
data.pop("hidden_states", None)
|
||||
if self.sgl_ext is None:
|
||||
data.pop("sgl_ext", None)
|
||||
return data
|
||||
|
||||
|
||||
@@ -814,6 +855,14 @@ class ChatCompletionStreamResponse(BaseModel):
|
||||
model: str
|
||||
choices: List[ChatCompletionResponseStreamChoice]
|
||||
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):
|
||||
|
||||
@@ -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.usage_processor import UsageProcessor
|
||||
from sglang.srt.entrypoints.openai.utils import (
|
||||
process_cached_tokens_details_from_ret,
|
||||
process_hidden_states_from_ret,
|
||||
process_routed_experts_from_ret,
|
||||
to_openai_style_logprobs,
|
||||
@@ -813,25 +814,19 @@ class OpenAIServingChat(OpenAIServingBase):
|
||||
yield f"data: {hidden_states_chunk.model_dump_json()}\n\n"
|
||||
|
||||
if request.return_routed_experts and routed_experts:
|
||||
for index, choice_routed_experts in routed_experts.items():
|
||||
if choice_routed_experts is not None:
|
||||
routed_experts_chunk = ChatCompletionStreamResponse(
|
||||
id=content["meta_info"]["id"],
|
||||
created=int(time.time()),
|
||||
choices=[
|
||||
ChatCompletionResponseStreamChoice(
|
||||
index=index,
|
||||
delta=DeltaMessage(
|
||||
sgl_ext=SglExt(
|
||||
routed_experts=choice_routed_experts
|
||||
)
|
||||
),
|
||||
finish_reason=None,
|
||||
)
|
||||
],
|
||||
model=request.model,
|
||||
)
|
||||
yield (f"data: {routed_experts_chunk.model_dump_json()}\n\n")
|
||||
# Get first non-None routed_experts value
|
||||
first_routed_experts = 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"
|
||||
|
||||
# Additional usage chunk
|
||||
if request.stream_options and request.stream_options.include_usage:
|
||||
@@ -891,6 +886,19 @@ class OpenAIServingChat(OpenAIServingBase):
|
||||
"""Build chat completion response from generation results"""
|
||||
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):
|
||||
# Process logprobs
|
||||
choice_logprobs = None
|
||||
@@ -899,7 +907,6 @@ class OpenAIServingChat(OpenAIServingBase):
|
||||
|
||||
# Handle hidden states
|
||||
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"]
|
||||
text = ret_item["text"]
|
||||
@@ -960,9 +967,6 @@ class OpenAIServingChat(OpenAIServingBase):
|
||||
else None
|
||||
),
|
||||
hidden_states=hidden_states,
|
||||
sgl_ext=(
|
||||
SglExt(routed_experts=routed_experts) if routed_experts else None
|
||||
),
|
||||
)
|
||||
choices.append(choice_data)
|
||||
|
||||
@@ -980,6 +984,7 @@ class OpenAIServingChat(OpenAIServingBase):
|
||||
choices=choices,
|
||||
usage=usage,
|
||||
metadata={"weight_version": ret[0]["meta_info"]["weight_version"]},
|
||||
sglext=response_sglext,
|
||||
)
|
||||
|
||||
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.usage_processor import UsageProcessor
|
||||
from sglang.srt.entrypoints.openai.utils import (
|
||||
process_cached_tokens_details_from_ret,
|
||||
process_hidden_states_from_ret,
|
||||
process_routed_experts_from_ret,
|
||||
to_openai_style_logprobs,
|
||||
@@ -332,25 +333,20 @@ class OpenAIServingCompletion(OpenAIServingBase):
|
||||
yield f"data: {hidden_states_chunk.model_dump_json()}\n\n"
|
||||
|
||||
if request.return_routed_experts and routed_experts:
|
||||
for index, choice_routed_experts in routed_experts.items():
|
||||
if choice_routed_experts is not None:
|
||||
routed_experts_chunk = CompletionStreamResponse(
|
||||
id=content["meta_info"]["id"],
|
||||
created=created,
|
||||
object="text_completion",
|
||||
choices=[
|
||||
CompletionResponseStreamChoice(
|
||||
index=index,
|
||||
text="",
|
||||
sgl_ext=SglExt(
|
||||
routed_experts=choice_routed_experts
|
||||
),
|
||||
finish_reason=None,
|
||||
)
|
||||
],
|
||||
model=request.model,
|
||||
)
|
||||
yield (f"data: {routed_experts_chunk.model_dump_json()}\n\n")
|
||||
# Get first non-None routed_experts value
|
||||
first_routed_experts = 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"
|
||||
|
||||
# Handle final usage chunk
|
||||
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 = 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):
|
||||
text = ret_item["text"]
|
||||
|
||||
@@ -450,7 +459,6 @@ class OpenAIServingCompletion(OpenAIServingBase):
|
||||
|
||||
# Handle hidden states
|
||||
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"]
|
||||
|
||||
@@ -465,9 +473,6 @@ class OpenAIServingCompletion(OpenAIServingBase):
|
||||
else None
|
||||
),
|
||||
hidden_states=hidden_states,
|
||||
sgl_ext=(
|
||||
SglExt(routed_experts=routed_experts) if routed_experts else None
|
||||
),
|
||||
)
|
||||
choices.append(choice_data)
|
||||
|
||||
@@ -484,6 +489,7 @@ class OpenAIServingCompletion(OpenAIServingBase):
|
||||
choices=choices,
|
||||
usage=usage,
|
||||
metadata={"weight_version": ret[0]["meta_info"]["weight_version"]},
|
||||
sglext=response_sglext,
|
||||
)
|
||||
|
||||
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 sglang.srt.entrypoints.openai.protocol import UsageInfo
|
||||
from sglang.srt.entrypoints.openai.protocol import PromptTokensDetails, UsageInfo
|
||||
|
||||
|
||||
@final
|
||||
@@ -10,9 +10,9 @@ class UsageProcessor:
|
||||
"""Stateless helpers that turn raw token counts into a UsageInfo."""
|
||||
|
||||
@staticmethod
|
||||
def _details_if_cached(count: int) -> Optional[Dict[str, int]]:
|
||||
"""Return {"cached_tokens": N} only when N > 0 (keeps JSON slim)."""
|
||||
return {"cached_tokens": count} if count > 0 else None
|
||||
def _details_if_cached(count: int) -> Optional[PromptTokensDetails]:
|
||||
"""Return PromptTokensDetails only when count > 0 (keeps JSON slim)."""
|
||||
return PromptTokensDetails(cached_tokens=count) if count > 0 else None
|
||||
|
||||
@staticmethod
|
||||
def calculate_response_usage(
|
||||
@@ -73,7 +73,7 @@ class UsageProcessor:
|
||||
def calculate_token_usage(
|
||||
prompt_tokens: int,
|
||||
completion_tokens: int,
|
||||
cached_tokens: Optional[Dict[str, int]] = None,
|
||||
cached_tokens: Optional[PromptTokensDetails] = None,
|
||||
) -> UsageInfo:
|
||||
"""Calculate token usage information"""
|
||||
return UsageInfo(
|
||||
|
||||
@@ -2,6 +2,7 @@ import logging
|
||||
from typing import Any, Dict, List, Optional, Union
|
||||
|
||||
from sglang.srt.entrypoints.openai.protocol import (
|
||||
CachedTokensDetails,
|
||||
ChatCompletionRequest,
|
||||
CompletionRequest,
|
||||
LogProbs,
|
||||
@@ -83,3 +84,33 @@ def process_routed_experts_from_ret(
|
||||
if not getattr(request, "return_routed_experts", False):
|
||||
return 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,
|
||||
completion_tokens=recv_obj.completion_tokens,
|
||||
cached_tokens=recv_obj.cached_tokens,
|
||||
cached_tokens_details=recv_obj.cached_tokens_details,
|
||||
spec_verify_ct=recv_obj.spec_verify_ct,
|
||||
spec_accepted_tokens=recv_obj.spec_accepted_tokens,
|
||||
input_token_logprobs_val=recv_obj.input_token_logprobs_val,
|
||||
|
||||
@@ -1005,6 +1005,8 @@ class BatchTokenIDOutput(
|
||||
load: GetLoadReqOutput = None
|
||||
# Customized info
|
||||
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
|
||||
@@ -1094,6 +1096,8 @@ class BatchStrOutput(
|
||||
|
||||
# Customized info
|
||||
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
|
||||
@@ -1119,6 +1123,8 @@ class BatchMultimodalOutput(BaseBatchReq):
|
||||
placeholder_tokens_val: List[Optional[List[int]]]
|
||||
|
||||
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
|
||||
@@ -1136,6 +1142,8 @@ class BatchEmbeddingOutput(BaseBatchReq, RequestTimingMetricsMixin):
|
||||
|
||||
# Number of times each request was retracted.
|
||||
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
|
||||
|
||||
@@ -154,6 +154,9 @@ def _handle_output_by_index(output, i):
|
||||
prompt_tokens=_extract_field_by_index(output, "prompt_tokens", i),
|
||||
completion_tokens=_extract_field_by_index(output, "completion_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(
|
||||
output, "input_token_logprobs_val", i, check_length=False
|
||||
),
|
||||
|
||||
@@ -660,6 +660,8 @@ class Req:
|
||||
self.last_node: Any = None
|
||||
self.last_host_node: Any = None
|
||||
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
|
||||
self.swa_uuid_for_lock: Optional[int] = None
|
||||
# The prefix length that is inserted into the tree cache
|
||||
@@ -750,6 +752,14 @@ class Req:
|
||||
self.cached_tokens = 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.
|
||||
# This is used to compute the average acceptance length per request.
|
||||
self.spec_verify_ct = 0
|
||||
@@ -1577,7 +1587,33 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
||||
# Only calculate cached_tokens once. Once retracted, the 'retracted_stain'
|
||||
# flag will always True
|
||||
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.is_retracted = False
|
||||
|
||||
|
||||
@@ -1677,7 +1677,10 @@ class Scheduler(
|
||||
direction * recv_req.priority < direction * candidate_req.priority
|
||||
)
|
||||
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.waiting_queue.pop(idx)
|
||||
req_to_abort = candidate_req
|
||||
@@ -1705,6 +1708,9 @@ class Scheduler(
|
||||
for req in self.waiting_queue:
|
||||
entry_time = req.time_stats.wait_queue_entry_time
|
||||
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(
|
||||
AbortReq(
|
||||
finished_reason={
|
||||
@@ -2024,6 +2030,10 @@ class Scheduler(
|
||||
if not prefetch_done:
|
||||
# skip staging requests that are ongoing prefetch
|
||||
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)
|
||||
res = adder.add_one_req(
|
||||
|
||||
@@ -46,6 +46,45 @@ class SchedulerOutputProcessorMixin:
|
||||
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):
|
||||
assert self.disaggregation_mode == DisaggregationMode.DECODE
|
||||
for req in batch.reqs:
|
||||
@@ -893,6 +932,7 @@ class SchedulerOutputProcessorMixin:
|
||||
prompt_tokens = []
|
||||
completion_tokens = []
|
||||
cached_tokens = []
|
||||
cached_tokens_details = [] # Detailed breakdown by cache source
|
||||
spec_verify_ct = []
|
||||
spec_accepted_tokens = []
|
||||
retraction_counts = []
|
||||
@@ -1005,6 +1045,10 @@ class SchedulerOutputProcessorMixin:
|
||||
prompt_tokens.append(len(req.origin_input_ids))
|
||||
completion_tokens.append(len(output_ids_))
|
||||
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)
|
||||
|
||||
queue_times.append(req.time_stats.get_queueing_time())
|
||||
@@ -1138,6 +1182,7 @@ class SchedulerOutputProcessorMixin:
|
||||
prompt_tokens=prompt_tokens,
|
||||
completion_tokens=completion_tokens,
|
||||
cached_tokens=cached_tokens,
|
||||
cached_tokens_details=cached_tokens_details,
|
||||
input_token_logprobs_val=input_token_logprobs_val,
|
||||
input_token_logprobs_idx=input_token_logprobs_idx,
|
||||
output_token_logprobs_val=output_token_logprobs_val,
|
||||
@@ -1169,6 +1214,7 @@ class SchedulerOutputProcessorMixin:
|
||||
embeddings = []
|
||||
prompt_tokens = []
|
||||
cached_tokens = []
|
||||
cached_tokens_details = [] # Detailed breakdown by cache source
|
||||
queue_times = []
|
||||
forward_entry_times = []
|
||||
prefill_launch_delays = []
|
||||
@@ -1184,6 +1230,9 @@ class SchedulerOutputProcessorMixin:
|
||||
prompt_tokens.append(len(req.origin_input_ids))
|
||||
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())
|
||||
forward_entry_times.append(req.time_stats.forward_entry_time)
|
||||
|
||||
@@ -1208,6 +1257,7 @@ class SchedulerOutputProcessorMixin:
|
||||
embeddings=embeddings,
|
||||
prompt_tokens=prompt_tokens,
|
||||
cached_tokens=cached_tokens,
|
||||
cached_tokens_details=cached_tokens_details,
|
||||
placeholder_tokens_idx=None,
|
||||
placeholder_tokens_val=None,
|
||||
retraction_counts=retraction_counts,
|
||||
|
||||
@@ -1536,6 +1536,14 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
|
||||
"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):
|
||||
meta_info["hidden_states"] = recv_obj.output_hidden_states[i]
|
||||
@@ -1974,6 +1982,14 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
|
||||
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(
|
||||
labels,
|
||||
recv_obj.prompt_tokens[i],
|
||||
@@ -1982,6 +1998,7 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerMultiItemMixi
|
||||
state.finished_time - state.created_time,
|
||||
self._request_has_grammar(state.obj),
|
||||
retraction_count,
|
||||
cached_tokens_details,
|
||||
)
|
||||
|
||||
def dump_requests(self, state: ReqState, out_dict: dict):
|
||||
|
||||
@@ -151,6 +151,9 @@ class HiRadixCache(RadixCache):
|
||||
# record the ongoing prefetch requests
|
||||
self.ongoing_prefetch = {}
|
||||
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
|
||||
self.write_through_threshold = (
|
||||
1 if server_args.hicache_write_policy == "write_through" else 2
|
||||
@@ -572,6 +575,8 @@ class HiRadixCache(RadixCache):
|
||||
TreeNode.counter = 0
|
||||
self.cache_controller.reset()
|
||||
self.token_to_kv_pool_host.clear()
|
||||
# Clear per-request tracking dicts
|
||||
self.prefetch_loaded_tokens_by_reqid.clear()
|
||||
self.evictable_host_leaves.clear()
|
||||
super().reset()
|
||||
|
||||
@@ -1089,10 +1094,12 @@ class HiRadixCache(RadixCache):
|
||||
del self.ongoing_prefetch[req_id]
|
||||
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:
|
||||
self.storage_metrics_collector.log_prefetched_tokens(
|
||||
min_completed_tokens - matched_length
|
||||
)
|
||||
self.storage_metrics_collector.log_prefetched_tokens(loaded_from_storage)
|
||||
|
||||
return True
|
||||
|
||||
@@ -1105,6 +1112,14 @@ class HiRadixCache(RadixCache):
|
||||
return
|
||||
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):
|
||||
key = params.key
|
||||
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)
|
||||
|
||||
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:
|
||||
return
|
||||
|
||||
|
||||
@@ -17,7 +17,7 @@ import logging
|
||||
import os
|
||||
import time
|
||||
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.environ import envs
|
||||
@@ -1148,8 +1148,8 @@ class TokenizerMetricsCollector:
|
||||
|
||||
self.cached_tokens_total = Counter(
|
||||
name="sglang:cached_tokens_total",
|
||||
documentation="Number of cached prompt tokens.",
|
||||
labelnames=labels.keys(),
|
||||
documentation="Number of cached prompt tokens by source (device/host/storage).",
|
||||
labelnames=list(labels.keys()) + ["cache_source"],
|
||||
)
|
||||
|
||||
self.num_requests_total = Counter(
|
||||
@@ -1303,11 +1303,36 @@ class TokenizerMetricsCollector:
|
||||
e2e_latency: float,
|
||||
has_grammar: bool,
|
||||
retraction_count: int,
|
||||
cached_tokens_details: Optional[Dict[str, Any]] = None,
|
||||
):
|
||||
self.prompt_tokens_total.labels(**labels).inc(prompt_tokens)
|
||||
self.generation_tokens_total.labels(**labels).inc(generation_tokens)
|
||||
|
||||
# Report cached tokens with detailed source breakdown
|
||||
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)
|
||||
if has_grammar:
|
||||
self.num_so_requests_total.labels(**labels).inc(1)
|
||||
|
||||
Reference in New Issue
Block a user