From 38d4c9ba8841870a14d2556f78f7d3b6ba1c0cf8 Mon Sep 17 00:00:00 2001 From: Ke Bao Date: Tue, 2 Jun 2026 21:17:50 +0800 Subject: [PATCH] Improve type annotations in unified radix cache (#26948) --- .../sglang/srt/managers/cache_controller.py | 6 +-- .../hybrid_cache/hybrid_cache_controller.py | 2 +- .../srt/mem_cache/unified_radix_cache.py | 38 +++++++++++++------ 3 files changed, 31 insertions(+), 15 deletions(-) diff --git a/python/sglang/srt/managers/cache_controller.py b/python/sglang/srt/managers/cache_controller.py index 65ef06ab3..480717aed 100644 --- a/python/sglang/srt/managers/cache_controller.py +++ b/python/sglang/srt/managers/cache_controller.py @@ -400,9 +400,9 @@ class HiCacheController: self.prefetch_queue = Queue() self.backup_queue = Queue() - self.prefetch_revoke_queue = Queue() - self.ack_backup_queue = Queue() - self.host_mem_release_queue = Queue() + self.prefetch_revoke_queue: Queue[str] = Queue() + self.ack_backup_queue: Queue[StorageOperation] = Queue() + self.host_mem_release_queue: Queue[torch.Tensor] = Queue() self.prefetch_thread.start() self.backup_thread.start() 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 1b09a5b96..cf14fb6c3 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 @@ -174,7 +174,7 @@ class HybridCacheController(BaseHiCacheController): enable_storage_metrics: bool = False, ): startup_storage_backend = storage_backend - self.extra_host_mem_release_queues: dict[PoolName, Queue] = {} + self.extra_host_mem_release_queues: dict[PoolName, Queue[torch.Tensor]] = {} super().__init__( token_to_kv_pool_allocator=token_to_kv_pool_allocator, mem_pool_host=mem_pool_host, diff --git a/python/sglang/srt/mem_cache/unified_radix_cache.py b/python/sglang/srt/mem_cache/unified_radix_cache.py index 34e00ae19..4e5339330 100644 --- a/python/sglang/srt/mem_cache/unified_radix_cache.py +++ b/python/sglang/srt/mem_cache/unified_radix_cache.py @@ -7,8 +7,8 @@ import time from array import array from collections import defaultdict from functools import partial -from queue import Empty -from typing import TYPE_CHECKING, Any, Optional +from queue import Empty, Queue +from typing import TYPE_CHECKING, Any, Iterator, Optional, TypeVar import torch @@ -61,9 +61,15 @@ from sglang.srt.session.streaming_session import StreamingSession if TYPE_CHECKING: from sglang.srt.managers.schedule_batch import Req from sglang.srt.mem_cache.cache_init_params import CacheInitParams + from sglang.srt.mem_cache.hybrid_cache.hybrid_cache_controller import ( + PrefetchOperation, + ) from sglang.srt.server_args import ServerArgs +T = TypeVar("T") + + class UnifiedTreeNode: counter = 0 @@ -305,7 +311,7 @@ class UnifiedRadixCache(KVCacheEventMixin, BasePrefixCache): ) # HiCache D↔H defaults (overridden by init_hicache) - self.cache_controller = None + self.cache_controller: Optional[HybridCacheController] = None self.write_through_threshold = 256 self.prefetch_stop_policy = "best_effort" self.prefetch_threshold = 256 @@ -348,8 +354,18 @@ class UnifiedRadixCache(KVCacheEventMixin, BasePrefixCache): self.ongoing_load_back: dict[int, tuple[UnifiedTreeNode, DecLockRefParams]] = {} self.enable_storage = False self.prefetch_loaded_tokens_by_reqid: dict[str, int] = {} - self.ongoing_prefetch: dict = {} - self.ongoing_backup: dict = {} + self.ongoing_prefetch: dict[ + str, + tuple[ + UnifiedTreeNode, + RadixKey, + torch.Tensor, + PrefetchOperation, + DecLockRefParams, + dict[ComponentType, list[PoolTransfer]], + ], + ] = {} + self.ongoing_backup: dict[int, tuple[UnifiedTreeNode, DecLockRefParams]] = {} if self.cache_controller is not None: self.cache_controller.reset() @@ -629,7 +645,7 @@ class UnifiedRadixCache(KVCacheEventMixin, BasePrefixCache): req, is_finished=True, insert_result=result, insert_params=insert_params ) - def cache_unfinished_req(self, req: Req, chunked=False, **kwargs) -> None: + def cache_unfinished_req(self, req: Req, chunked: bool = False, **kwargs) -> None: if self.session.try_cache_unfinished_req(req, chunked=chunked, **kwargs): return @@ -1150,7 +1166,7 @@ class UnifiedRadixCache(KVCacheEventMixin, BasePrefixCache): node: UnifiedTreeNode, comp: TreeComponent, target: EvictLayer = EvictLayer.DEVICE, - tracker: dict[ComponentType, int] = None, + tracker: Optional[dict[ComponentType, int]] = None, ) -> tuple[int, int]: device_freed, host_freed = comp.evict_component(node, target=target) if tracker is not None: @@ -1296,7 +1312,7 @@ class UnifiedRadixCache(KVCacheEventMixin, BasePrefixCache): self.evictable_host_leaves.discard(node) def _evict_to_host( - self, node: UnifiedTreeNode, tracker: dict[ComponentType, int] = None + self, node: UnifiedTreeNode, tracker: Optional[dict[ComponentType, int]] = None ) -> None: """GPU→CPU demotion: release all device resources, node stays in tree.""" assert not node.evicted and node.backuped @@ -1725,14 +1741,14 @@ class UnifiedRadixCache(KVCacheEventMixin, BasePrefixCache): ) self.cache_controller.prefetch_tokens_occupied += len(prefetch_key) - def _prefetch_timeout_check_linear_func(self, operation) -> bool: + def _prefetch_timeout_check_linear_func(self, operation: PrefetchOperation) -> bool: return ( time.monotonic() - operation.start_time > self.prefetch_timeout_base + len(operation.hash_value) * self.prefetch_timeout_per_page ) - def can_terminate_prefetch(self, operation) -> bool: + def can_terminate_prefetch(self, operation: PrefetchOperation) -> bool: if self.prefetch_stop_policy == "best_effort": return True @@ -1888,7 +1904,7 @@ class UnifiedRadixCache(KVCacheEventMixin, BasePrefixCache): ) -> None: cc = self.cache_controller - def _drain_queue(q, limit: Optional[int]): + def _drain_queue(q: Queue[T], limit: Optional[int]) -> Iterator[T]: drained = 0 while limit is None or drained < limit: try: