diff --git a/python/sglang/srt/mem_cache/unified_radix_cache.py b/python/sglang/srt/mem_cache/unified_radix_cache.py index a42c5d0c3..8351f519b 100644 --- a/python/sglang/srt/mem_cache/unified_radix_cache.py +++ b/python/sglang/srt/mem_cache/unified_radix_cache.py @@ -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) 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 a5e262a91..3289aa3cc 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 @@ -9,6 +9,11 @@ from unittest import mock import torch from sglang.srt.configs.mamba_utils import Mamba2CacheParams, Mamba2StateShape +from sglang.srt.disaggregation.kv_events import ( + BlockRemoved, + BlockStored, + StorageMedium, +) from sglang.srt.environ import envs from sglang.srt.layers.attention.fla.chunk_delta_h import CHUNK_SIZE as FLA_CHUNK_SIZE from sglang.srt.managers.schedule_batch import Req @@ -113,7 +118,7 @@ class CacheConfig: return "_".join(parts) -def build_fixture(cfg: CacheConfig): +def build_fixture(cfg: CacheConfig, *, enable_kv_cache_events: bool = False): """Create (tree, allocator, req_to_token_pool) from a CacheConfig.""" server_args = ServerArgs(model_path="dummy", page_size=cfg.page_size) # MambaRadixCache reads mamba_cache_chunk_size, whose property otherwise @@ -227,6 +232,7 @@ def build_fixture(cfg: CacheConfig): sliding_window_size=cfg.sliding_window_size, tree_components=cfg.components, enable_mamba_extra_buffer=cfg.enable_mamba_extra_buffer, + enable_kv_cache_events=enable_kv_cache_events, ) tree = UnifiedRadixCache(params=cache_init_params) tree.cache_init_params = cache_init_params @@ -234,6 +240,168 @@ def build_fixture(cfg: CacheConfig): return tree, allocator, req_to_token_pool +class TestUnifiedRadixCacheKVEvents(CustomTestCase): + cfg = CacheConfig(page_size=2, kv_size=64, max_context_len=64) + + def _insert(self, tree, allocator, tokens): + key = RadixKey(array("q", tokens)) + value = allocator.alloc(len(tokens)) + self.assertIsNotNone(value) + return tree.insert(InsertParams(key=key, value=value[: len(key)])) + + def _stored_events(self, tree, medium=None): + events = [e for e in tree.take_events() if isinstance(e, BlockStored)] + if medium is not None: + events = [e for e in events if e.medium == medium] + return events + + def _removed_events(self, tree, medium=None): + events = [e for e in tree.take_events() if isinstance(e, BlockRemoved)] + if medium is not None: + events = [e for e in events if e.medium == medium] + return events + + def _leaf_for(self, tree, tokens): + match = tree.match_prefix(MatchPrefixParams(key=RadixKey(array("q", tokens)))) + self.assertIsNot(match.last_device_node, tree.root_node) + return match.last_device_node + + def _init_hicache(self, tree, *, write_policy: str = "write_through"): + import sglang.srt.mem_cache.hybrid_cache.hybrid_pool_assembler as assembler + + orig_kv_host_pool = assembler.MHATokenToKVPoolHost + + def kv_host_pool_wrapper(*args, **kwargs): + kwargs["pin_memory"] = False + return orig_kv_host_pool(*args, **kwargs) + + patcher = mock.patch.object( + assembler, + "MHATokenToKVPoolHost", + side_effect=kv_host_pool_wrapper, + ) + patcher.start() + self.addCleanup(patcher.stop) + + server_args = ServerArgs( + model_path="dummy", + page_size=self.cfg.page_size, + hicache_io_backend="direct", + hicache_write_policy=write_policy, + ) + set_global_server_args_for_scheduler(server_args) + tree.init_hicache(server_args, tree.cache_init_params) + tree.write_through_threshold = 1 << 30 + tree.load_back_threshold = 0 + + def _backup_node(self, tree, node): + backed_up = tree.write_backup(node, write_back=True) + self.assertGreater(backed_up, 0) + tree.writing_check(write_back=True) + + def _load_back_node(self, tree, node): + loaded = tree.load_back(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() + tree.loading_check() + + def test_kv_events_store_and_remove_full_blocks(self): + tree, allocator, _ = build_fixture(self.cfg, enable_kv_cache_events=True) + tree.take_events() # Clear the reset event. + + seq = [1, 2, 3, 4] + self._insert(tree, allocator, seq) + stored = self._stored_events(tree, StorageMedium.GPU) + self.assertEqual(len(stored), 2) + self.assertEqual([list(e.token_ids) for e in stored], [[1, 2], [3, 4]]) + stored_hashes = [e.block_hashes[0] for e in stored] + + result = tree.evict(EvictParams(num_tokens=len(seq))) + self.assertGreaterEqual(result.num_tokens_evicted, len(seq)) + removed = self._removed_events(tree, StorageMedium.GPU) + self.assertCountEqual([e.block_hashes[0] for e in removed], stored_hashes) + + def test_kv_events_split_preserves_block_hash_parentage(self): + tree, allocator, _ = build_fixture(self.cfg, enable_kv_cache_events=True) + tree.take_events() # Clear the reset event. + + self._insert(tree, allocator, [1, 2, 3, 4]) + first_insert = self._stored_events(tree, StorageMedium.GPU) + self.assertEqual(len(first_insert), 2) + split_parent_hash = first_insert[0].block_hashes[0] + + self._insert(tree, allocator, [1, 2, 5, 6]) + second_insert = self._stored_events(tree, StorageMedium.GPU) + self.assertEqual(len(second_insert), 1) + self.assertEqual(list(second_insert[0].token_ids), [5, 6]) + self.assertEqual(second_insert[0].parent_block_hash, split_parent_hash) + + split_parent = next(iter(tree.root_node.children.values())) + split_child = split_parent.children.get((3, 4)) + self.assertIsNotNone(split_child) + self.assertEqual(len(split_parent.hash_value), 1) + self.assertIsNotNone(split_child.hash_value) + self.assertEqual(len(split_child.hash_value), 1) + + def test_hicache_kv_events_track_gpu_cpu_transitions(self): + tree, allocator, _ = build_fixture(self.cfg, enable_kv_cache_events=True) + self._init_hicache(tree) + tree.take_events() # Clear reset / init events. + + seq = [1, 2, 3, 4] + self._insert(tree, allocator, seq) + stored_gpu = self._stored_events(tree, StorageMedium.GPU) + self.assertEqual(len(stored_gpu), 2) + stored_hashes = [e.block_hashes[0] for e in stored_gpu] + + node = self._leaf_for(tree, seq) + self._backup_node(tree, node) + stored_cpu = self._stored_events(tree, StorageMedium.CPU) + self.assertCountEqual([e.block_hashes[0] for e in stored_cpu], stored_hashes) + + tree.evict(EvictParams(num_tokens=len(seq))) + removed_gpu = self._removed_events(tree, StorageMedium.GPU) + self.assertCountEqual([e.block_hashes[0] for e in removed_gpu], stored_hashes) + + self._load_back_node(tree, node) + restored_gpu = self._stored_events(tree, StorageMedium.GPU) + self.assertCountEqual([e.block_hashes[0] for e in restored_gpu], stored_hashes) + + tree.evict(EvictParams(num_tokens=len(seq))) + self._removed_events(tree, StorageMedium.GPU) + tree.evict_host(len(seq)) + removed_cpu = self._removed_events(tree, StorageMedium.CPU) + self.assertCountEqual([e.block_hashes[0] for e in removed_cpu], stored_hashes) + + def test_hicache_reinsert_evicted_node_emits_gpu_store(self): + tree, allocator, _ = build_fixture(self.cfg, enable_kv_cache_events=True) + self._init_hicache(tree) + tree.take_events() # Clear reset / init events. + + seq = [1, 2, 3, 4] + self._insert(tree, allocator, seq) + stored_gpu = self._stored_events(tree, StorageMedium.GPU) + self.assertEqual(len(stored_gpu), 2) + stored_hashes = [e.block_hashes[0] for e in stored_gpu] + + node = self._leaf_for(tree, seq) + self._backup_node(tree, node) + self._stored_events(tree, StorageMedium.CPU) + + tree.evict(EvictParams(num_tokens=len(seq))) + self._removed_events(tree, StorageMedium.GPU) + self.assertTrue(node.evicted) + self.assertTrue(node.backuped) + + self._insert(tree, allocator, seq) + restored_gpu = self._stored_events(tree, StorageMedium.GPU) + self.assertFalse(node.evicted) + self.assertCountEqual([e.block_hashes[0] for e in restored_gpu], stored_hashes) + + class UnifiedRadixCacheSuite: cfg: CacheConfig