diff --git a/python/sglang/srt/managers/cache_controller.py b/python/sglang/srt/managers/cache_controller.py index ca276bc11..f6138f6d5 100644 --- a/python/sglang/srt/managers/cache_controller.py +++ b/python/sglang/srt/managers/cache_controller.py @@ -273,6 +273,11 @@ class HiCacheController: ]: raise ValueError(f"Invalid write policy: {write_policy}") + if write_policy == "write_back": + logger.warning( + "write_back policy will be deprecated in future releases; please migrate to write_through_selective with appropriate configuration for better performance and reliability." + ) + # self.write_queue = PriorityQueue[CacheOperation]() self.load_queue: List[CacheOperation] = [] self.write_queue: List[CacheOperation] = [] @@ -463,9 +468,8 @@ class HiCacheController: self.enable_storage = True # todo: threshold policy for prefetching self.prefetch_threshold = max(prefetch_threshold, self.page_size) - self.prefetch_capacity_limit = max( - 0, int(0.8 * (self.mem_pool_host.size - self.mem_pool_device.size)) - ) + # Budget speculative prefetch at half the host pool, leaving the rest for the write-back staging path. + self.prefetch_capacity_limit = int(0.5 * self.mem_pool_host.size) # tracking the number of tokens locked in prefetching, updated by the main scheduler thread self.prefetch_tokens_occupied = 0 diff --git a/python/sglang/srt/mem_cache/hiradix_cache.py b/python/sglang/srt/mem_cache/hiradix_cache.py index 92e842ad5..65e185413 100644 --- a/python/sglang/srt/mem_cache/hiradix_cache.py +++ b/python/sglang/srt/mem_cache/hiradix_cache.py @@ -1035,51 +1035,65 @@ class HiRadixCache(RadixCache): def evict(self, params: EvictParams) -> EvictResult: start_time = time.perf_counter() num_tokens = params.num_tokens - leaves = list(self.evictable_leaves) - eviction_heap = [ - (self.eviction_strategy.get_priority(node), node) for node in leaves - ] - heapq.heapify(eviction_heap) - - num_evicted = 0 - write_back_nodes = [] - while num_evicted < num_tokens and len(eviction_heap): - _priority, x = heapq.heappop(eviction_heap) - - if x.lock_ref > 0: - continue - - if not x.backuped: - if self.cache_controller.write_policy == "write_back": - # write to host if the node is not backuped - written = self.write_backup(x, write_back=True) - num_evicted += written - if written > 0: - write_back_nodes.append(x) - else: - num_evicted += self._evict_regular(x) - else: - num_evicted += self._evict_backuped(x) - - for child in x.parent.children.values(): - if child in write_back_nodes: - continue - if not child.evicted: - break - else: - # all children are evicted or no children - new_priority = self.eviction_strategy.get_priority(x.parent) - heapq.heappush(eviction_heap, (new_priority, x.parent)) - if self.cache_controller.write_policy == "write_back": - self.writing_check(write_back=True) - for node in write_back_nodes: - assert node.backuped - self._evict_backuped(node) - + num_evicted = self._evict_write_back(num_tokens) + else: + num_evicted = self._evict_write_through(num_tokens) self.update_eviction_metrics(num_evicted, start_time) return EvictResult(num_tokens_evicted=num_evicted) + def _make_eviction_heap(self): + heap = [ + (self.eviction_strategy.get_priority(node), node) + for node in self.evictable_leaves + ] + heapq.heapify(heap) + return heap + + def _promote_parent(self, node: TreeNode, heap) -> None: + # Once all of a node's children are evicted, it becomes a device leaf. + p = node.parent + if p is not self.root_node and all(c.evicted for c in p.children.values()): + heapq.heappush(heap, (self.eviction_strategy.get_priority(p), p)) + + def _evict_write_through(self, num_tokens: int) -> int: + """write_through / write_through_selective: drop non-backuped leaves, + demote already-backuped ones. Nothing is staged to host during eviction, + so this is a plain on-the-fly pass. + """ + heap = self._make_eviction_heap() + num_evicted = 0 + while num_evicted < num_tokens and heap: + _priority, x = heapq.heappop(heap) + if x.lock_ref > 0: + continue + if x.backuped: + num_evicted += self._evict_backuped(x) + else: + num_evicted += self._evict_regular(x) + self._promote_parent(x, heap) + return num_evicted + + def _evict_write_back(self, num_tokens: int) -> int: + """eviction for write_back mode: demote already-backuped leaves, stage non-backuped ones to host if possible, otherwise drop them. + note this path will be deprecated in the future. + """ + heap = self._make_eviction_heap() + num_evicted = 0 + while num_evicted < num_tokens and heap: + _priority, x = heapq.heappop(heap) + if x.lock_ref > 0: + continue + if x.backuped: + num_evicted += self._evict_backuped(x) + elif self.write_backup(x, write_back=True) > 0: + self.writing_check(write_back=True) + num_evicted += self._evict_backuped(x) + else: + num_evicted += self._drop_subtree_no_host(x) + self._promote_parent(x, heap) + return num_evicted + def _evict_backuped(self, node: TreeNode): # GPU -> CPU demotion: block moves from device to host. # Emit remove(GPU) so downstream indexers stop scoring it as device-local. @@ -1105,6 +1119,45 @@ class HiRadixCache(RadixCache): self._delete_leaf(node) return num_evicted + def _drop_subtree_no_host(self, root: TreeNode) -> int: + nodes = [] + stack = [root] + while stack: + n = stack.pop() + nodes.append(n) + stack.extend(n.children.values()) + + if any(n.host_ref_counter > 0 for n in nodes): + return 0 + + logger.warning( + "write_back: KV cache on device are dropped without backup due to host memory pressure, subtree root %d, num_nodes %d", + root.id, + len(nodes), + ) + + freed_device = 0 + for n in nodes: + if n.host_value is not None: + self._record_remove_event(n, medium=StorageMedium.CPU) + self.cache_controller.evict_host(n.host_value) + n.host_value = None + if n.value is not None: + self._record_remove_event(n, medium=StorageMedium.GPU) + self.cache_controller.mem_pool_device_allocator.free(n.value) + freed_device += len(n.value) + self.evictable_size_ -= len(n.value) + n.value = None + self.ongoing_write_through.pop(n.id, None) + self.evictable_leaves.discard(n) + self.evictable_host_leaves.discard(n) + + key = root.key.child_key(self.page_size) + root.parent.children.pop(key, None) + self._update_leaf_status(root.parent) + self._update_host_leaf_status(root.parent) + return freed_device + def evict_host(self, num_tokens: int): leaves = list(self.evictable_host_leaves) eviction_heap = [ @@ -1169,6 +1222,10 @@ class HiRadixCache(RadixCache): self.dec_lock_ref(ancester_node) return None + # Protect the nodes being loaded from host eviction. + for n in nodes_to_load: + n.protect_host() + device_indices = self.cache_controller.load( host_indices=host_indices, node_id=last_hit_node.id, @@ -1184,6 +1241,8 @@ class HiRadixCache(RadixCache): self.dec_lock_ref(ancester_node) if device_indices is None: # no sufficient GPU memory to load back KV caches + for n in nodes_to_load: + n.release_host() logger.warning( "load_back: FAILED to load %d tokens for node %d " "even after eviction (evictable_size=%d)", @@ -1193,6 +1252,8 @@ class HiRadixCache(RadixCache): ) return None + for n in nodes_to_load: + n.release_host() self.ongoing_load_back[last_hit_node.id] = last_hit_node offset = 0 for node in nodes_to_load: 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 daaf05e63..1472ab1d0 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 @@ -262,8 +262,8 @@ def _deepseek_v4_num_host_pages( "use --hicache-ratio instead." ) ratio = server_args.hicache_ratio - full_host_pages = max(int(device_full_pages * ratio), device_full_pages + 1) - swa_host_pages = max(int(device_swa_pages * ratio), device_swa_pages + 1) + full_host_pages = int(device_full_pages * ratio) + swa_host_pages = int(device_swa_pages * ratio) return full_host_pages, swa_host_pages diff --git a/python/sglang/srt/mem_cache/memory_pool_host.py b/python/sglang/srt/mem_cache/memory_pool_host.py index 5f0618cc5..00785cf3b 100644 --- a/python/sglang/srt/mem_cache/memory_pool_host.py +++ b/python/sglang/srt/mem_cache/memory_pool_host.py @@ -1445,9 +1445,14 @@ class MambaPoolHost(HostKVCache): self.page_num = self.size // self.page_size + 1 self.size = self.page_num * self.page_size - assert ( - self.size > device_pool.size - ), "The host memory should be larger than the device memory with the current protocol" + if self.size <= device_pool.size: + logger.warning( + "HiCache host KV pool (%d tokens) is smaller than the device pool (%d tokens);" + "L2 cache effectiveness is reduced." + "Consider increasing --hicache-ratio (or --hicache-size) for higher L2 cache hit rate.", + self.size, + device_pool.size, + ) host_mem = psutil.virtual_memory() requested_bytes = self.size * self.size_per_token diff --git a/python/sglang/srt/mem_cache/pool_host/base.py b/python/sglang/srt/mem_cache/pool_host/base.py index 0ca962225..5ef35c25d 100644 --- a/python/sglang/srt/mem_cache/pool_host/base.py +++ b/python/sglang/srt/mem_cache/pool_host/base.py @@ -111,9 +111,14 @@ class HostKVCache(abc.ABC): self.start_layer = device_pool.start_layer self.end_layer = device_pool.end_layer - assert ( - self.size > device_pool.size - ), "The host memory should be larger than the device memory with the current protocol" + if self.size <= device_pool.size: + logger.warning( + "HiCache host KV pool (%d tokens) is smaller than the device pool (%d tokens);" + "L2 cache effectiveness is reduced." + "Consider increasing --hicache-ratio (or --hicache-size) for higher L2 cache hit rate.", + self.size, + device_pool.size, + ) # Verify there is enough available host memory. host_mem = psutil.virtual_memory()