diff --git a/python/sglang/srt/disaggregation/decode_kvcache_offload_manager.py b/python/sglang/srt/disaggregation/decode_kvcache_offload_manager.py index bccdd809b5c4..a2385411c882 100644 --- a/python/sglang/srt/disaggregation/decode_kvcache_offload_manager.py +++ b/python/sglang/srt/disaggregation/decode_kvcache_offload_manager.py @@ -99,6 +99,8 @@ def __init__( storage_backend=server_args.hicache_storage_backend, model_name=server_args.served_model_name, storage_backend_extra_config=hicache_storage_backend_extra_config, + enable_metrics=server_args.enable_metrics, + extra_metric_labels=server_args.extra_metric_labels, ) self.ongoing_offload = {} @@ -219,9 +221,13 @@ def check_offload_progress(self): def _check_offload_progress(self, finish_count): """Check the progress of offload from device to host.""" while finish_count > 0: - _, finish_event, ack_list = self.cache_controller.ack_write_queue.pop(0) - finish_event.synchronize() - for ack_id in ack_list: + ack = self.cache_controller.ack_write_queue.pop(0) + ack.finish_event.synchronize() + self.cache_controller.record_l1_l2_transfer_complete( + direction="offload", + ack=ack, + ) + for ack_id in ack.node_ids: ( req, host_indices, diff --git a/python/sglang/srt/managers/cache_controller.py b/python/sglang/srt/managers/cache_controller.py index 56057ade66f6..871b300b3d5e 100644 --- a/python/sglang/srt/managers/cache_controller.py +++ b/python/sglang/srt/managers/cache_controller.py @@ -28,6 +28,9 @@ PoolName, PoolTransfer, ) +from sglang.srt.observability.metrics_collector import ( + HiCacheL1L2TransferMetricsCollector, +) if TYPE_CHECKING: from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator @@ -54,8 +57,16 @@ class LayerLoadingEvent: def __init__(self, num_layers: int): self._num_layers = num_layers - self.load_events = [device_module.Event() for _ in range(num_layers)] - self.start_event = device_module.Event() # start event on controller stream + # The last layer's event doubles as finish_event; together with + # start_event it is used for elapsed_time() in transfer metrics, which + # requires both events to be created with enable_timing=True. + self.load_events = [ + device_module.Event(enable_timing=(i == num_layers - 1)) + for i in range(num_layers) + ] + self.start_event = device_module.Event( + enable_timing=True + ) # start event on controller stream def complete(self, layer_index: int): assert 0 <= layer_index < self._num_layers @@ -146,6 +157,18 @@ class HiCacheAck(NamedTuple): finish_event: device_module.Event node_ids: List[int] + # Number of KV token slots moved by this merged operation. + token_count: int + + # Number of KV blocks moved. For HiCache this should be token_count // page_size. + block_count: int + + # Estimated total bytes moved for this operation. + byte_count: int + + # Host-side fallback timer start. Used only if device event elapsed_time is unavailable. + start_time_ns: int + class TransferBuffer: """ @@ -263,6 +286,8 @@ def __init__( pp_rank: int = 0, pp_size: int = 1, enable_storage_metrics: bool = False, + enable_metrics: bool = False, + extra_metric_labels: Optional[dict[str, str]] = None, ): self.tp_group = tp_group self.attn_cp_group = attn_cp_group @@ -286,6 +311,28 @@ def __init__( self.pp_size = pp_size self.enable_storage_metrics = enable_storage_metrics + # init L1/L2 transfer metrics collection (device-host transfers triggered by write/load). + self.enable_l1_l2_transfer_metrics = enable_metrics + self.hicache_l1_l2_transfer_metrics_collector = None + + self.hicache_l1_l2_transfer_totals = { + "offload": { + "events": 0, + "blocks": 0, + "bytes": 0, + "xfer_us": 0, + }, + "onboard": { + "events": 0, + "blocks": 0, + "bytes": 0, + "xfer_us": 0, + }, + } + + if self.enable_l1_l2_transfer_metrics: + self._init_l1_l2_transfer_metrics(extra_metric_labels) + # Draft KV pool support (best-effort piggyback on target L2/L3 ops). self.has_draft = False self.mem_pool_device_draft = None @@ -717,8 +764,13 @@ def start_writing(self) -> None: ) self.write_queue.clear() - start_event = device_module.Event() - finish_event = device_module.Event() + # enable_timing so record_l1_l2_transfer_complete can use + # elapsed_time() for the actual transfer duration. + start_event = device_module.Event(enable_timing=True) + finish_event = device_module.Event(enable_timing=True) + + token_count = int(host_indices.numel()) + start_time_ns = time.perf_counter_ns() start_event.record() with device_module.stream(self.write_stream): @@ -742,7 +794,15 @@ def start_writing(self) -> None: if device_indices.is_cuda: device_indices.record_stream(self.write_stream) - self.ack_write_queue.append(HiCacheAck(start_event, finish_event, op.node_ids)) + self.ack_write_queue.append( + self._make_hicache_ack( + start_event=start_event, + finish_event=finish_event, + node_ids=op.node_ids, + token_count=token_count, + start_time_ns=start_time_ns, + ) + ) def load( self, @@ -794,6 +854,10 @@ def start_loading(self) -> int: ) self.load_queue.clear() producer_event = self.layer_done_counter.events[producer_id] + + token_count = int(host_indices.numel()) + start_time_ns = time.perf_counter_ns() + producer_event.start_event.record() with device_module.stream(self.load_stream): @@ -824,10 +888,12 @@ def start_loading(self) -> int: device_indices.record_stream(self.load_stream) self.ack_load_queue.append( - HiCacheAck( + self._make_hicache_ack( start_event=producer_event.start_event, finish_event=producer_event.finish_event, node_ids=op.node_ids, + token_count=token_count, + start_time_ns=start_time_ns, ) ) return producer_id @@ -1234,3 +1300,188 @@ def backup_thread_func(self): except Empty: continue + + def _init_l1_l2_transfer_metrics( + self, + extra_metric_labels: Optional[dict[str, str]] = None, + ) -> None: + """Initialize Prometheus metrics for L1<->L2 transfer accounting.""" + + from sglang.srt.distributed import ( + get_tensor_model_parallel_rank, + get_tensor_model_parallel_world_size, + ) + from sglang.srt.layers.dp_attention import ( + get_attention_dp_rank, + get_attention_tp_rank, + get_attention_tp_size, + is_dp_attention_enabled, + ) + + if is_dp_attention_enabled(): + tp_rank = get_attention_tp_rank() + tp_size = get_attention_tp_size() + dp_rank = get_attention_dp_rank() + else: + tp_rank = get_tensor_model_parallel_rank() + tp_size = get_tensor_model_parallel_world_size() + dp_rank = 0 + + attn_cp_rank, attn_cp_size = self.get_attn_cp_rank_and_size() + + labels = { + "tp_rank": str(tp_rank), + "tp_size": str(tp_size), + "dp_rank": str(dp_rank), + "pp_rank": str(self.pp_rank), + "pp_size": str(self.pp_size), + "attn_cp_rank": str(attn_cp_rank), + "attn_cp_size": str(attn_cp_size), + "io_backend": str(self.io_backend), + } + + if extra_metric_labels: + labels.update({k: str(v) for k, v in extra_metric_labels.items()}) + + self.hicache_l1_l2_transfer_metrics_collector = ( + HiCacheL1L2TransferMetricsCollector(labels) + ) + + def _host_pool_bytes_per_token(self) -> int: + """Return bytes per token slot for the host-side KV pools. + + Handles both a normal HostKVCache and HostPoolGroup-style wrappers. + """ + if hasattr(self.mem_pool_host, "entries"): + return sum( + int(entry.host_pool.size_per_token) + for entry in self.mem_pool_host.entries + ) + + return int(getattr(self.mem_pool_host, "size_per_token", 0)) + + def _estimate_l1_l2_transfer_bytes(self, token_count: int) -> int: + """Estimate total bytes moved for one L1<->L2 transfer operation.""" + bytes_per_token = self._host_pool_bytes_per_token() + total = int(token_count) * bytes_per_token + + if self.has_draft and self.mem_pool_host_draft is not None: + total += int(token_count) * int( + getattr(self.mem_pool_host_draft, "size_per_token", 0) + ) + + return total + + def _make_hicache_ack( + self, + *, + start_event: device_module.Event, + finish_event: device_module.Event, + node_ids: List[int], + token_count: int, + start_time_ns: int, + ) -> HiCacheAck: + block_count = int(token_count) // int(self.page_size) + byte_count = self._estimate_l1_l2_transfer_bytes(token_count) + + return HiCacheAck( + start_event=start_event, + finish_event=finish_event, + node_ids=node_ids, + token_count=int(token_count), + block_count=block_count, + byte_count=byte_count, + start_time_ns=start_time_ns, + ) + + def _transfer_elapsed_us(self, ack: HiCacheAck) -> int: + """Return transfer duration in microseconds. + + Prefer device event timing. Fall back to host elapsed time for devices/backends + that do not expose elapsed_time(). + """ + try: + return max(0, int(ack.start_event.elapsed_time(ack.finish_event) * 1000)) + except Exception: + return max(0, int((time.perf_counter_ns() - ack.start_time_ns) // 1000)) + + def record_l1_l2_transfer_complete( + self, + *, + direction: str, + ack: HiCacheAck, + ) -> None: + """Record logs and Prometheus metrics after a transfer ack completes. + + direction: + - "offload": L1 -> L2 + - "onboard": L2 -> L1 + """ + should_log = logger.isEnabledFor(logging.DEBUG) + should_record_metrics = ( + self.hicache_l1_l2_transfer_metrics_collector is not None + ) + if not should_log and not should_record_metrics: + return + + if direction == "offload": + action = "Offload" + src = "sglang_hicache::L1" + dst = "sglang_hicache::L2" + elif direction == "onboard": + action = "Onboard" + src = "sglang_hicache::L2" + dst = "sglang_hicache::L1" + else: + raise ValueError(f"Unknown HiCache L1/L2 transfer direction: {direction}") + + xfer_us = self._transfer_elapsed_us(ack) + + if should_log: + ts_us = time.time_ns() // 1000 + logger.debug( + "%s transfer complete ts_us=%d blocks=%d bytes=%d xfer_us=%d " + "bandwidth=%.2fGB/s " + 'src="%s" dst="%s"', + action, + ts_us, + ack.block_count, + ack.byte_count, + xfer_us, + ack.byte_count * 0.001 / xfer_us if xfer_us > 0 else 0, + src, + dst, + ) + + if should_record_metrics: + self.hicache_l1_l2_transfer_metrics_collector.record_transfer( + direction=direction, + src=src, + dst=dst, + blocks=ack.block_count, + bytes_=ack.byte_count, + xfer_us=xfer_us, + ) + + if should_log: + totals = self.hicache_l1_l2_transfer_totals[direction] + totals["events"] += 1 + totals["blocks"] += ack.block_count + totals["bytes"] += ack.byte_count + totals["xfer_us"] += xfer_us + + logger.debug( + '%s transfer cumulative direction="%s" total_events=%d ' + "total_blocks=%d total_bytes=%d total_xfer_us=%d " + "bandwidth=%.2fGB/s cumulative " + 'src="%s" dst="%s"', + action, + direction, + totals["events"], + totals["blocks"], + totals["bytes"], + totals["xfer_us"], + totals["bytes"] * 0.001 / totals["xfer_us"] if totals["xfer_us"] > 0 else 0, + src, + dst, + ) diff --git a/python/sglang/srt/mem_cache/hi_mamba_radix_cache.py b/python/sglang/srt/mem_cache/hi_mamba_radix_cache.py index c3e4c7a80405..5abea09304cc 100644 --- a/python/sglang/srt/mem_cache/hi_mamba_radix_cache.py +++ b/python/sglang/srt/mem_cache/hi_mamba_radix_cache.py @@ -383,9 +383,13 @@ def writing_check(self, write_back=False): if write_back: # blocking till all write back complete while len(self.ongoing_write_through) > 0: - for _, finish_event, ack_list in self.cache_controller.ack_write_queue: - finish_event.synchronize() - for ack_id in ack_list: + for ack in self.cache_controller.ack_write_queue: + ack.finish_event.synchronize() + self.cache_controller.record_l1_l2_transfer_complete( + direction="offload", + ack=ack, + ) + for ack_id in ack.node_ids: backuped_node = self.ongoing_write_through.pop(ack_id) self._record_store_event( backuped_node, medium=StorageMedium.CPU @@ -400,8 +404,8 @@ def writing_check(self, write_back=False): return finish_count = 0 - for _, finish_event, ack_list in self.cache_controller.ack_write_queue: - if not finish_event.query(): + for ack in self.cache_controller.ack_write_queue: + if not ack.finish_event.query(): break finish_count += 1 @@ -415,9 +419,16 @@ def writing_check(self, write_back=False): finish_count = int(queue_size.item()) while finish_count > 0: - _, finish_event, ack_list = self.cache_controller.ack_write_queue.pop(0) - finish_event.synchronize() - for ack_id in ack_list: + ack = self.cache_controller.ack_write_queue.pop(0) + ack.finish_event.synchronize() + + # Record the L1->L2 transfer completion for this write-back operation. + self.cache_controller.record_l1_l2_transfer_complete( + direction="offload", + ack=ack, + ) + + for ack_id in ack.node_ids: backuped_node = self.ongoing_write_through.pop(ack_id) self._record_store_event(backuped_node, medium=StorageMedium.CPU) self.dec_lock_ref(backuped_node) @@ -427,12 +438,20 @@ def writing_check(self, write_back=False): def loading_check(self): finish_count = 0 - for _, finish_event, ack_list in self.cache_controller.ack_load_queue: - if not finish_event.query(): + for ack in self.cache_controller.ack_load_queue: + if not ack.finish_event.query(): # the KV cache loading is still ongoing break + + # ensure completion of loading, before recording transfer completion and updating cache state + ack.finish_event.synchronize() + self.cache_controller.record_l1_l2_transfer_complete( + direction="onboard", + ack=ack, + ) + finish_count += 1 - for ack_id in ack_list: + for ack_id in ack.node_ids: end_node = self.ongoing_load_back.pop(ack_id) self.dec_lock_ref(end_node) diff --git a/python/sglang/srt/mem_cache/hiradix_cache.py b/python/sglang/srt/mem_cache/hiradix_cache.py index aef02c2ca5d7..f3faf53b9570 100644 --- a/python/sglang/srt/mem_cache/hiradix_cache.py +++ b/python/sglang/srt/mem_cache/hiradix_cache.py @@ -154,6 +154,8 @@ def __init__(self, params: CacheInitParams, server_args: ServerArgs): pp_rank=self.pp_rank, pp_size=self.pp_size, enable_storage_metrics=self.enable_storage_metrics, + enable_metrics=params.enable_metrics, + extra_metric_labels=self.extra_metric_labels, ) self._apply_storage_runtime_config( storage_backend=server_args.hicache_storage_backend, @@ -799,9 +801,16 @@ def writing_check(self, write_back=False): if write_back: # blocking till all write back complete while len(self.ongoing_write_through) > 0: - for _, finish_event, ack_list in self.cache_controller.ack_write_queue: - finish_event.synchronize() - for ack_id in ack_list: + for ack in self.cache_controller.ack_write_queue: + ack.finish_event.synchronize() + + # Record the L1->L2 transfer completion for this write-back operation. + self.cache_controller.record_l1_l2_transfer_complete( + direction="offload", + ack=ack, + ) + + for ack_id in ack.node_ids: node, backup_len = self.ongoing_write_through.pop(ack_id) # DMA confirmed -- block is now on host. self._record_store_event(node, medium=StorageMedium.CPU) @@ -816,8 +825,8 @@ def writing_check(self, write_back=False): return finish_count = 0 - for _, finish_event, ack_list in self.cache_controller.ack_write_queue: - if not finish_event.query(): + for ack in self.cache_controller.ack_write_queue: + if not ack.finish_event.query(): break finish_count += 1 queue_size = torch.tensor(finish_count, dtype=torch.int, device="cpu") @@ -826,9 +835,16 @@ def writing_check(self, write_back=False): finish_count = int(queue_size.item()) while finish_count > 0: - _, finish_event, ack_list = self.cache_controller.ack_write_queue.pop(0) - finish_event.synchronize() - for ack_id in ack_list: + ack = self.cache_controller.ack_write_queue.pop(0) + ack.finish_event.synchronize() + + # Record the L1->L2 transfer completion for this write-through operation. + self.cache_controller.record_l1_l2_transfer_complete( + direction="offload", + ack=ack, + ) + + for ack_id in ack.node_ids: node, backup_len = self.ongoing_write_through.pop(ack_id) # DMA confirmed -- block is now on host. self._record_store_event(node, medium=StorageMedium.CPU) @@ -839,13 +855,21 @@ def writing_check(self, write_back=False): def loading_check(self): finish_count = 0 - for _, finish_event, ack_list in self.cache_controller.ack_load_queue: - if not finish_event.query(): + for ack in self.cache_controller.ack_load_queue: + if not ack.finish_event.query(): # the KV cache loading is still ongoing break + + # ensure completion of loading, before recording transfer completion and updating cache state + ack.finish_event.synchronize() + self.cache_controller.record_l1_l2_transfer_complete( + direction="onboard", + ack=ack, + ) + finish_count += 1 # no need to sync across TP workers as batch forwarding is synced - for ack_id in ack_list: + for ack_id in ack.node_ids: end_node = self.ongoing_load_back.pop(ack_id) self.dec_lock_ref(end_node) diff --git a/python/sglang/srt/mem_cache/hybrid_cache/hybrid_cache_controller.py b/python/sglang/srt/mem_cache/hybrid_cache/hybrid_cache_controller.py index cf14fb6c3283..e533140b6acd 100644 --- a/python/sglang/srt/mem_cache/hybrid_cache/hybrid_cache_controller.py +++ b/python/sglang/srt/mem_cache/hybrid_cache/hybrid_cache_controller.py @@ -11,9 +11,6 @@ import torch from sglang.srt.managers.cache_controller import CacheOperation as BaseCacheOperation -from sglang.srt.managers.cache_controller import ( - HiCacheAck, -) from sglang.srt.managers.cache_controller import ( HiCacheController as BaseHiCacheController, ) @@ -172,6 +169,8 @@ def __init__( pp_size: int = 1, transfer_layer_num: Optional[int] = None, enable_storage_metrics: bool = False, + enable_metrics: bool = False, + extra_metric_labels: Optional[dict[str, str]] = None, ): startup_storage_backend = storage_backend self.extra_host_mem_release_queues: dict[PoolName, Queue[torch.Tensor]] = {} @@ -192,6 +191,8 @@ def __init__( pp_rank=pp_rank, pp_size=pp_size, enable_storage_metrics=enable_storage_metrics, + enable_metrics=enable_metrics, + extra_metric_labels=extra_metric_labels, ) # Override layer_num: hybrid models transfer all layers (For example, Linear Model (KV + Mamba)), # not just the full attention layers reported by full_kv_pool. @@ -399,8 +400,14 @@ def start_writing(self) -> None: self.move_hybrid_indices(op) ) self.write_queue.clear() - start_event = device_module.Event() - finish_event = device_module.Event() + # enable_timing so record_l1_l2_transfer_complete can use + # elapsed_time() for the actual transfer duration. + start_event = device_module.Event(enable_timing=True) + finish_event = device_module.Event(enable_timing=True) + + token_count = int(host_indices.numel()) + start_time_ns = time.perf_counter_ns() + start_event.record() with device_module.stream(self.write_stream): start_event.wait(self.write_stream) @@ -418,7 +425,15 @@ def start_writing(self) -> None: device_indices, resolved_pool_transfers, ) - self.ack_write_queue.append(HiCacheAck(start_event, finish_event, op.node_ids)) + self.ack_write_queue.append( + self._make_hicache_ack( + start_event=start_event, + finish_event=finish_event, + node_ids=op.node_ids, + token_count=token_count, + start_time_ns=start_time_ns, + ) + ) def load( self, @@ -473,6 +488,10 @@ def start_loading(self) -> int: ) self.load_queue.clear() producer_event = self.layer_done_counter.events[producer_id] + + token_count = int(host_indices.numel()) + start_time_ns = time.perf_counter_ns() + producer_event.start_event.record() with device_module.stream(self.load_stream): producer_event.start_event.wait(self.load_stream) @@ -493,10 +512,12 @@ def start_loading(self) -> int: resolved_pool_transfers, ) self.ack_load_queue.append( - HiCacheAck( - producer_event.start_event, - producer_event.finish_event, - op.node_ids, + self._make_hicache_ack( + start_event=producer_event.start_event, + finish_event=producer_event.finish_event, + node_ids=op.node_ids, + token_count=token_count, + start_time_ns=start_time_ns, ) ) return producer_id diff --git a/python/sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py b/python/sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py index e12c9d350372..007a785fa3bb 100644 --- a/python/sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py +++ b/python/sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py @@ -152,6 +152,8 @@ def build_kv_only_stack( pp_size=pp_size, transfer_layer_num=transfer_layer_num, enable_storage_metrics=enable_storage_metrics, + enable_metrics=params.enable_metrics, + extra_metric_labels=server_args.extra_metric_labels, ) return host_pool_group, cache_controller @@ -236,6 +238,8 @@ def build_hybrid_swa_stack( pp_size=pp_size, transfer_layer_num=transfer_layer_num, enable_storage_metrics=enable_storage_metrics, + enable_metrics=params.enable_metrics, + extra_metric_labels=server_args.extra_metric_labels, ) return host_pool_group, cache_controller @@ -489,6 +493,8 @@ def build_deepseek_v4_hicache_stack( pp_size=pp_size, transfer_layer_num=transfer_layer_num, enable_storage_metrics=enable_storage_metrics, + enable_metrics=params.enable_metrics, + extra_metric_labels=server_args.extra_metric_labels, ) return host_pool_group, cache_controller @@ -569,6 +575,8 @@ def build_hybrid_mamba_stack( pp_size=pp_size, transfer_layer_num=transfer_layer_num, enable_storage_metrics=enable_storage_metrics, + enable_metrics=params.enable_metrics, + extra_metric_labels=server_args.extra_metric_labels, ) return host_pool_group, cache_controller @@ -641,6 +649,8 @@ def build_anchor_sidecar_stack( pp_size=pp_size, transfer_layer_num=transfer_layer_num, enable_storage_metrics=enable_storage_metrics, + enable_metrics=params.enable_metrics, + extra_metric_labels=server_args.extra_metric_labels, ) return host_pool_group, cache_controller diff --git a/python/sglang/srt/mem_cache/unified_radix_cache.py b/python/sglang/srt/mem_cache/unified_radix_cache.py index c9ba6ec86637..b98016471a47 100644 --- a/python/sglang/srt/mem_cache/unified_radix_cache.py +++ b/python/sglang/srt/mem_cache/unified_radix_cache.py @@ -2138,9 +2138,13 @@ def writing_check(self, write_back: bool = False) -> None: if write_back: # Blocking: wait for all pending write-backs while self.ongoing_write_through: - for _, finish_event, ack_list in cc.ack_write_queue: - finish_event.synchronize() - for ack_id in ack_list: + for ack in cc.ack_write_queue: + ack.finish_event.synchronize() + cc.record_l1_l2_transfer_complete( + direction="offload", + ack=ack, + ) + for ack_id in ack.node_ids: entry = self.ongoing_write_through.pop(ack_id, None) if entry is not None: node, params = entry @@ -2157,8 +2161,8 @@ def writing_check(self, write_back: bool = False) -> None: return finish_count = 0 - for _, finish_event, ack_list in cc.ack_write_queue: - if not finish_event.query(): + for ack in cc.ack_write_queue: + if not ack.finish_event.query(): break finish_count += 1 @@ -2169,9 +2173,16 @@ def writing_check(self, write_back: bool = False) -> None: # Process completed acks while finish_count > 0: - _, finish_event, ack_list = cc.ack_write_queue.pop(0) - finish_event.synchronize() - for ack_id in ack_list: + ack = cc.ack_write_queue.pop(0) + ack.finish_event.synchronize() + + # Post-ack callback for cache controller (e.g. to trigger next write-backs) + self.cache_controller.record_l1_l2_transfer_complete( + direction="offload", + ack=ack, + ) + + for ack_id in ack.node_ids: node, params = self.ongoing_write_through.pop(ack_id) self._record_store_event(node, medium=StorageMedium.CPU) self.dec_lock_ref(node, params) @@ -2185,11 +2196,19 @@ def loading_check(self) -> None: if cc is None or not self.ongoing_load_back: return finish_count = 0 - for _, finish_event, ack_list in cc.ack_load_queue: - if not finish_event.query(): + for ack in cc.ack_load_queue: + if not ack.finish_event.query(): break + + # ensure completion of loading, before recording transfer completion and updating cache state + ack.finish_event.synchronize() + self.cache_controller.record_l1_l2_transfer_complete( + direction="onboard", + ack=ack, + ) + finish_count += 1 - for ack_id in ack_list: + for ack_id in ack.node_ids: node, lock_params = self.ongoing_load_back.pop(ack_id) self.dec_lock_ref(node, lock_params) del cc.ack_load_queue[:finish_count] diff --git a/python/sglang/srt/observability/metrics_collector.py b/python/sglang/srt/observability/metrics_collector.py index 942ad8f5e8df..82cfb97ffe7e 100644 --- a/python/sglang/srt/observability/metrics_collector.py +++ b/python/sglang/srt/observability/metrics_collector.py @@ -1802,6 +1802,76 @@ def log_storage_metrics(self, storage_metrics: Optional[StorageMetrics] = None): self._log_histogram(self.histogram_backup_bandwidth, v) +class HiCacheL1L2TransferMetricsCollector: + """Prometheus metrics for HiCache L1<->L2 KV block transfers. + + L1: device/GPU KV cache. + L2: host/CPU KV cache. + """ + + def __init__(self, labels: Optional[dict[str, str]] = None): + from prometheus_client import Counter, Histogram + + self.labels = labels or {} + labelnames = list(self.labels.keys()) + ["direction", "src", "dst"] + + self.transfer_blocks_total = Counter( + "sglang:hicache_l1_l2_transfer_blocks_total", + "Total number of KV cache blocks transferred between HiCache L1 and L2.", + labelnames=labelnames, + ) + + self.transfer_bytes_total = Counter( + "sglang:hicache_l1_l2_transfer_bytes_total", + "Total number of KV cache bytes transferred between HiCache L1 and L2.", + labelnames=labelnames, + ) + + self.transfer_duration_us = Histogram( + "sglang:hicache_l1_l2_transfer_duration_us", + "Observed duration in microseconds for one completed HiCache L1<->L2 KV block transfer.", + labelnames=labelnames, + buckets=( + 100, + 250, + 500, + 1_000, + 2_500, + 5_000, + 10_000, + 25_000, + 50_000, + 100_000, + 250_000, + 500_000, + 1_000_000, + 2_500_000, + 5_000_000, + ), + ) + + def record_transfer( + self, + *, + direction: str, + src: str, + dst: str, + blocks: int, + bytes_: int, + xfer_us: int, + ) -> None: + metric_labels = { + **self.labels, + "direction": direction, + "src": src, + "dst": dst, + } + + self.transfer_blocks_total.labels(**metric_labels).inc(blocks) + self.transfer_bytes_total.labels(**metric_labels).inc(bytes_) + self.transfer_duration_us.labels(**metric_labels).observe(xfer_us) + + class ExpertDispatchCollector(_StatLoggerDIMixin): def __init__(self, ep_size: int) -> None: from prometheus_client import Histogram as _PromHistogram diff --git a/test/registered/disaggregation/test_specv2_kvcache_offloading.py b/test/registered/disaggregation/test_specv2_kvcache_offloading.py index 0cd5c77bdc64..7351aee00550 100644 --- a/test/registered/disaggregation/test_specv2_kvcache_offloading.py +++ b/test/registered/disaggregation/test_specv2_kvcache_offloading.py @@ -8,6 +8,7 @@ """ import unittest +from types import SimpleNamespace from unittest.mock import MagicMock import torch @@ -90,6 +91,10 @@ def synchronize(self): pass +def _ack(node_id: int): + return SimpleNamespace(finish_event=_FinishedEvent(), node_ids=[node_id]) + + class TestReleaseFinishedReq(unittest.TestCase): """Tests for _release_finished_req overallocation cleanup.""" @@ -292,7 +297,7 @@ def test_unfinished_offload_ack_does_not_free_incremental_slots(self): 8, ) manager.cache_controller = MagicMock() - manager.cache_controller.ack_write_queue = [(None, _FinishedEvent(), [7])] + manager.cache_controller.ack_write_queue = [_ack(7)] manager._trigger_backup = MagicMock(return_value="last_hash") manager._check_offload_progress(1) @@ -326,7 +331,7 @@ def test_offload_kv_cache_tracks_inflight_write_until_ack(self): self.assertEqual(manager.offloaded_state[req.rid].inc_len, 4) manager.cache_controller.write.assert_called_once() - manager.cache_controller.ack_write_queue = [(None, _FinishedEvent(), [1])] + manager.cache_controller.ack_write_queue = [_ack(1)] manager._trigger_backup = MagicMock(return_value="last_hash") manager._check_offload_progress(1) @@ -369,7 +374,7 @@ def test_finished_offload_ack_waits_for_other_inflight_writes(self): 8, ) manager.cache_controller = MagicMock() - manager.cache_controller.ack_write_queue = [(None, _FinishedEvent(), [8])] + manager.cache_controller.ack_write_queue = [_ack(8)] manager._trigger_backup = MagicMock(return_value="last_hash") manager._check_offload_progress(1) @@ -399,7 +404,7 @@ def test_finished_request_releases_all_committed_slots_after_last_offload_ack( 12, ) manager.cache_controller = MagicMock() - manager.cache_controller.ack_write_queue = [(None, _FinishedEvent(), [9])] + manager.cache_controller.ack_write_queue = [_ack(9)] manager._trigger_backup = MagicMock(return_value="last_hash") manager._check_offload_progress(1) diff --git a/test/registered/unit/mem_cache/test_unified_radix_cache_unittest.py b/test/registered/unit/mem_cache/test_unified_radix_cache_unittest.py index 4cb5319658e7..4d2de0f99f87 100644 --- a/test/registered/unit/mem_cache/test_unified_radix_cache_unittest.py +++ b/test/registered/unit/mem_cache/test_unified_radix_cache_unittest.py @@ -390,8 +390,8 @@ def _load_back_node(self, tree, node): self.assertTrue(loaded) producer_id = tree.ready_to_load_host_cache() self.assertNotEqual(producer_id, -1) - for _, finish_event, _ in list(tree.cache_controller.ack_load_queue): - finish_event.synchronize() + for ack in list(tree.cache_controller.ack_load_queue): + ack.finish_event.synchronize() tree.loading_check() def test_kv_events_store_and_remove_full_blocks(self): @@ -1914,8 +1914,8 @@ def _load_back_node(self, tree, node): self.assertTrue(loaded) producer_id = tree.ready_to_load_host_cache() self.assertNotEqual(producer_id, -1) - for _, finish_event, _ in list(tree.cache_controller.ack_load_queue): - finish_event.synchronize() + for ack in list(tree.cache_controller.ack_load_queue): + ack.finish_event.synchronize() tree.loading_check() return node.component_data[ComponentType.FULL].value @@ -2383,8 +2383,8 @@ def _release_ongoing_load_back_locks(self, tree): def _finish_pending_loads(self, tree): producer_id = tree.ready_to_load_host_cache() self.assertNotEqual(producer_id, -1) - for _, finish_event, _ in list(tree.cache_controller.ack_load_queue): - finish_event.synchronize() + for ack in list(tree.cache_controller.ack_load_queue): + ack.finish_event.synchronize() tree.loading_check() def _match_tokens_for_chain(self, chain):