Support KV events for UnifiedRadixCache (#26387)

Co-authored-by: Zhangheng <hzh0425@apache.org>
This commit is contained in:
weireweire
2026-05-27 22:10:40 +08:00
committed by GitHub
co-authored by Zhangheng
parent 5f8911183b
commit 034dd39189
2 changed files with 184 additions and 2 deletions
@@ -11,6 +11,7 @@ from typing import TYPE_CHECKING, Any, Optional
import torch
from sglang.srt.disaggregation.kv_events import StorageMedium
from sglang.srt.mem_cache.base_prefix_cache import (
BasePrefixCache,
DecLockRefParams,
@@ -24,6 +25,7 @@ from sglang.srt.mem_cache.base_prefix_cache import (
MatchPrefixParams,
MatchResult,
)
from sglang.srt.mem_cache.events import KVCacheEventMixin
from sglang.srt.mem_cache.hicache_storage import (
PoolName,
PoolTransfer,
@@ -220,7 +222,7 @@ COMPONENT_REGISTRY: dict[ComponentType, type[TreeComponent]] = {
logger = logging.getLogger(__name__)
class UnifiedRadixCache(BasePrefixCache):
class UnifiedRadixCache(KVCacheEventMixin, BasePrefixCache):
def __init__(
self,
params: CacheInitParams,
@@ -230,6 +232,8 @@ class UnifiedRadixCache(BasePrefixCache):
self.page_size = params.page_size
self.disable = params.disable
self.is_eagle = params.is_eagle
self.enable_kv_cache_events = params.enable_kv_cache_events
self.kv_event_queue = []
if self.token_to_kv_pool_allocator:
self.device = self.token_to_kv_pool_allocator.device
@@ -328,6 +332,7 @@ class UnifiedRadixCache(BasePrefixCache):
last_host_node=self.root_node,
best_match_node=self.root_node,
)
self._record_all_cleared_event()
def init_hicache(self, server_args: ServerArgs, params: CacheInitParams) -> None:
"""Initialize HiCache infrastructure."""
@@ -869,6 +874,7 @@ class UnifiedRadixCache(BasePrefixCache):
self._update_evictable_leaf_sets(new_node)
self._update_evictable_leaf_sets(parent)
self._record_store_event(new_node)
return new_node
def _unevict_node_on_insert(
@@ -885,6 +891,7 @@ class UnifiedRadixCache(BasePrefixCache):
self._update_evictable_leaf_sets(node)
if node.parent is not None:
self._update_evictable_leaf_sets(node.parent)
self._record_store_event(node, medium=StorageMedium.GPU)
def _insert_helper(
self,
@@ -1241,6 +1248,7 @@ class UnifiedRadixCache(BasePrefixCache):
node, trigger, target=EvictLayer.DEVICE, tracker=tracker
)
self._cascade_evict(node, trigger, tracker)
self._record_remove_event(node, medium=StorageMedium.GPU)
# after device eviction, insert aux components into host LRU.
self._for_each_component_lru(
@@ -1271,6 +1279,7 @@ class UnifiedRadixCache(BasePrefixCache):
return
else:
# Write-through: node has no backup, delete entirely.
self._record_remove_event(node, medium=StorageMedium.GPU)
for comp in self._components_tuple:
self._evict_component_and_detach_lru(
node, comp, target=EvictLayer.ALL, tracker=tracker
@@ -1291,6 +1300,7 @@ class UnifiedRadixCache(BasePrefixCache):
All freed tokens are accumulated into *tracker*."""
assert self._is_host_leaf(node), f"node {node.id} is not an H-leaf"
self._record_remove_event(node, medium=StorageMedium.CPU)
for comp in self._components_tuple:
_, hf = self._evict_component_and_detach_lru(
node, comp, target=EvictLayer.ALL, tracker=None
@@ -1437,6 +1447,8 @@ class UnifiedRadixCache(BasePrefixCache):
CacheTransferPhase.LOAD_BACK,
[kv_xfer],
)
for node in kv_xfer.nodes_to_load or ():
self._record_store_event(node, medium=StorageMedium.GPU)
for ct, xfers in comp_xfers.items():
self.components[ct].commit_hicache_transfer(
best_match_node,
@@ -2031,6 +2043,7 @@ class UnifiedRadixCache(BasePrefixCache):
entry = self.ongoing_write_through.pop(ack_id, None)
if entry is not None:
node, params = entry
self._record_store_event(node, medium=StorageMedium.CPU)
if params is not None:
self.dec_lock_ref(node, params)
if self.enable_storage:
@@ -2062,6 +2075,7 @@ class UnifiedRadixCache(BasePrefixCache):
finish_event.synchronize()
for ack_id in ack_list:
node, params = self.ongoing_write_through.pop(ack_id)
self._record_store_event(node, medium=StorageMedium.CPU)
self.dec_lock_ref(node, params)
if self.enable_storage:
self.write_backup_storage(node)