diff --git a/python/sglang/srt/mem_cache/unified_radix_cache.py b/python/sglang/srt/mem_cache/unified_radix_cache.py index 484b2b633..a9c7a0231 100644 --- a/python/sglang/srt/mem_cache/unified_radix_cache.py +++ b/python/sglang/srt/mem_cache/unified_radix_cache.py @@ -55,7 +55,11 @@ from sglang.srt.mem_cache.utils import ( get_eviction_strategy, split_node_hash_value, ) -from sglang.srt.observability.metrics_collector import StorageMetricsCollector +from sglang.srt.observability.metrics_collector import ( + STAT_LOGGER_ROLE_STORAGE, + StorageMetricsCollector, + resolve_collector_class, +) from sglang.srt.session.streaming_session import StreamingSession if TYPE_CHECKING: @@ -1587,6 +1591,8 @@ class UnifiedRadixCache(KVCacheEventMixin, BasePrefixCache): if self.cache_controller is None: return False + start_time = time.perf_counter() + # Build KV transfer kv_xfer = self.components[BASE_COMPONENT_TYPE].build_hicache_transfers( best_match_node, CacheTransferPhase.LOAD_BACK @@ -1665,6 +1671,13 @@ class UnifiedRadixCache(KVCacheEventMixin, BasePrefixCache): best_match_node, self.inc_lock_ref(best_match_node).to_dec_params(), ) + + if self.metrics_collector is not None: + self.metrics_collector.observe_load_back_duration( + time.perf_counter() - start_time + ) + self.metrics_collector.increment_load_back_num_tokens(len(device_indices)) + return True def _build_sidecar_transfers( @@ -2195,7 +2208,14 @@ class UnifiedRadixCache(KVCacheEventMixin, BasePrefixCache): labels.update(extra_metric_labels) existing_collector = self.storage_metrics_collector if existing_collector is None: - self.storage_metrics_collector = StorageMetricsCollector(labels=labels) + from sglang.srt.server_args import get_global_server_args + + storage_cls = resolve_collector_class( + get_global_server_args(), + STAT_LOGGER_ROLE_STORAGE, + StorageMetricsCollector, + ) + self.storage_metrics_collector = storage_cls(labels=labels) elif set(existing_collector.labels.keys()) == set(labels.keys()): existing_collector.labels = labels else: