From e166ca87584368699ac35ccb5e518211631849cf Mon Sep 17 00:00:00 2001 From: Vladislav Nosivskoy Date: Tue, 3 Feb 2026 22:45:35 +0300 Subject: [PATCH] [HiCache] feat: Add detailed cache hit breakdown for HiCache in `sglext` and Prometheus metrics (#17648) Signed-off-by: Vladislav Nosivskoy Co-authored-by: ishandhanani <82981111+ishandhanani@users.noreply.github.com> --- .../sglang/srt/entrypoints/openai/protocol.py | 77 +++++++++++++++---- .../srt/entrypoints/openai/serving_chat.py | 51 ++++++------ .../entrypoints/openai/serving_completions.py | 52 +++++++------ .../srt/entrypoints/openai/usage_processor.py | 10 +-- python/sglang/srt/entrypoints/openai/utils.py | 31 ++++++++ .../srt/managers/detokenizer_manager.py | 1 + python/sglang/srt/managers/io_struct.py | 8 ++ .../srt/managers/multi_tokenizer_mixin.py | 3 + python/sglang/srt/managers/schedule_batch.py | 38 ++++++++- python/sglang/srt/managers/scheduler.py | 12 ++- .../scheduler_output_processor_mixin.py | 50 ++++++++++++ .../sglang/srt/managers/tokenizer_manager.py | 17 ++++ python/sglang/srt/mem_cache/hiradix_cache.py | 24 +++++- python/sglang/srt/metrics/collector.py | 33 +++++++- 14 files changed, 333 insertions(+), 74 deletions(-) diff --git a/python/sglang/srt/entrypoints/openai/protocol.py b/python/sglang/srt/entrypoints/openai/protocol.py index 40ad9f3fb..998ff0b08 100644 --- a/python/sglang/srt/entrypoints/openai/protocol.py +++ b/python/sglang/srt/entrypoints/openai/protocol.py @@ -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): diff --git a/python/sglang/srt/entrypoints/openai/serving_chat.py b/python/sglang/srt/entrypoints/openai/serving_chat.py index ab4d6c7fb..0ddb36dc4 100644 --- a/python/sglang/srt/entrypoints/openai/serving_chat.py +++ b/python/sglang/srt/entrypoints/openai/serving_chat.py @@ -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( diff --git a/python/sglang/srt/entrypoints/openai/serving_completions.py b/python/sglang/srt/entrypoints/openai/serving_completions.py index 6f437e4e2..dee22ef18 100644 --- a/python/sglang/srt/entrypoints/openai/serving_completions.py +++ b/python/sglang/srt/entrypoints/openai/serving_completions.py @@ -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: diff --git a/python/sglang/srt/entrypoints/openai/usage_processor.py b/python/sglang/srt/entrypoints/openai/usage_processor.py index cd7a18b80..f02b858eb 100644 --- a/python/sglang/srt/entrypoints/openai/usage_processor.py +++ b/python/sglang/srt/entrypoints/openai/usage_processor.py @@ -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( diff --git a/python/sglang/srt/entrypoints/openai/utils.py b/python/sglang/srt/entrypoints/openai/utils.py index c4c9baede..796f8f59b 100644 --- a/python/sglang/srt/entrypoints/openai/utils.py +++ b/python/sglang/srt/entrypoints/openai/utils.py @@ -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), + ) diff --git a/python/sglang/srt/managers/detokenizer_manager.py b/python/sglang/srt/managers/detokenizer_manager.py index a65b0dd28..450e5d169 100644 --- a/python/sglang/srt/managers/detokenizer_manager.py +++ b/python/sglang/srt/managers/detokenizer_manager.py @@ -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, diff --git a/python/sglang/srt/managers/io_struct.py b/python/sglang/srt/managers/io_struct.py index 624812186..c8074fa86 100644 --- a/python/sglang/srt/managers/io_struct.py +++ b/python/sglang/srt/managers/io_struct.py @@ -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 diff --git a/python/sglang/srt/managers/multi_tokenizer_mixin.py b/python/sglang/srt/managers/multi_tokenizer_mixin.py index 4f2e3fb19..8b0bc6ffd 100644 --- a/python/sglang/srt/managers/multi_tokenizer_mixin.py +++ b/python/sglang/srt/managers/multi_tokenizer_mixin.py @@ -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 ), diff --git a/python/sglang/srt/managers/schedule_batch.py b/python/sglang/srt/managers/schedule_batch.py index 5aa8538ec..5768b28d2 100644 --- a/python/sglang/srt/managers/schedule_batch.py +++ b/python/sglang/srt/managers/schedule_batch.py @@ -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 diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index b7eb45c4a..e59db607e 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -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( diff --git a/python/sglang/srt/managers/scheduler_output_processor_mixin.py b/python/sglang/srt/managers/scheduler_output_processor_mixin.py index d79f929d3..ae2459708 100644 --- a/python/sglang/srt/managers/scheduler_output_processor_mixin.py +++ b/python/sglang/srt/managers/scheduler_output_processor_mixin.py @@ -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, diff --git a/python/sglang/srt/managers/tokenizer_manager.py b/python/sglang/srt/managers/tokenizer_manager.py index c457a8871..0b6983c76 100644 --- a/python/sglang/srt/managers/tokenizer_manager.py +++ b/python/sglang/srt/managers/tokenizer_manager.py @@ -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): diff --git a/python/sglang/srt/mem_cache/hiradix_cache.py b/python/sglang/srt/mem_cache/hiradix_cache.py index 793eca10a..6e0fac3b5 100644 --- a/python/sglang/srt/mem_cache/hiradix_cache.py +++ b/python/sglang/srt/mem_cache/hiradix_cache.py @@ -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 diff --git a/python/sglang/srt/metrics/collector.py b/python/sglang/srt/metrics/collector.py index 5b0c47b86..7511454f6 100644 --- a/python/sglang/srt/metrics/collector.py +++ b/python/sglang/srt/metrics/collector.py @@ -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)