[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]
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,
+8
View File
@@ -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
),
+37 -1
View File
@@ -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
+11 -1
View File
@@ -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):
+21 -3
View File
@@ -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
+29 -4
View File
@@ -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)