[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:
Vladislav Nosivskoy
2026-02-03 11:45:35 -08:00
committed by GitHub
co-authored by ishandhanani
parent d48bbe3bed
commit e166ca8758
14 changed files with 333 additions and 74 deletions
@@ -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,
+8
View File
@@ -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
), ),
+37 -1
View File
@@ -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
+11 -1
View File
@@ -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):
+21 -3
View File
@@ -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
+29 -4
View File
@@ -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)