From 903439044a58a61965449804ec5d21cf45a061e9 Mon Sep 17 00:00:00 2001 From: jthomson04 Date: Thu, 13 Aug 2026 13:36:53 -0700 Subject: [PATCH] perf(kv-events): coalesce cache events (#31479) --- python/sglang/srt/mem_cache/events.py | 44 +++++- test/manual/test_kv_events.py | 23 ++- .../unit/mem_cache/test_hiradix_cache_unit.py | 10 +- .../unit/mem_cache/test_mamba_unittest.py | 24 +-- .../unit/mem_cache/test_radix_cache_unit.py | 149 +++++++++++++++--- .../unit/mem_cache/test_swa_unittest.py | 20 +-- .../test_unified_radix_cache_unittest.py | 32 ++-- 7 files changed, 226 insertions(+), 76 deletions(-) diff --git a/python/sglang/srt/mem_cache/events.py b/python/sglang/srt/mem_cache/events.py index 15da2d40b..4f473d46b 100644 --- a/python/sglang/srt/mem_cache/events.py +++ b/python/sglang/srt/mem_cache/events.py @@ -35,6 +35,44 @@ from sglang.srt.mem_cache.utils import ( class KVCacheEventMixin: + def _enqueue_kv_event(self, event): + """Append an event, coalescing it with a compatible queue tail. + + KV event batches already support multiple block hashes. Combining them + here avoids emitting one event per page while preserving ordering and + the parent-linked store chains consumers use to rebuild the cache tree. + """ + if self.kv_event_queue: + tail = self.kv_event_queue[-1] + + if isinstance(tail, BlockRemoved) and isinstance(event, BlockRemoved): + if tail.medium == event.medium: + tail.block_hashes.extend(event.block_hashes) + return + + elif isinstance(tail, BlockStored) and isinstance(event, BlockStored): + tail_metadata = ( + tail.metadata if isinstance(tail, BlockStoredWithMetadata) else None + ) + event_metadata = ( + event.metadata + if isinstance(event, BlockStoredWithMetadata) + else None + ) + if ( + tail.medium == event.medium + and tail.lora_id == event.lora_id + and tail.block_size == event.block_size + and tail_metadata == event_metadata + and tail.block_hashes + and event.parent_block_hash == tail.block_hashes[-1] + ): + tail.block_hashes.extend(event.block_hashes) + tail.token_ids.extend(event.token_ids) + return + + self.kv_event_queue.append(event) + def _record_store_event(self, node: Any, medium=None): # One BlockStored per ``page_size`` chunk. # ``medium`` defaults to StorageMedium.GPU but callers may override @@ -94,7 +132,7 @@ class KVCacheEventMixin: **event_args, metadata=BlockStoredMetadata(cache_salt=node.key.cache_salt), ) - self.kv_event_queue.append(event) + self._enqueue_kv_event(event) parent_block_hash = block_hash page_index += 1 @@ -128,13 +166,13 @@ class KVCacheEventMixin: page_index += 1 if block_hashes: - self.kv_event_queue.append( + self._enqueue_kv_event( BlockRemoved(block_hashes=block_hashes, medium=medium) ) def _record_all_cleared_event(self): if self.enable_kv_cache_events: - self.kv_event_queue.append(AllBlocksCleared()) + self._enqueue_kv_event(AllBlocksCleared()) def take_events(self): """Atomically takes all events and clears the queue. diff --git a/test/manual/test_kv_events.py b/test/manual/test_kv_events.py index 80a1f3e07..d4b9dbc8c 100644 --- a/test/manual/test_kv_events.py +++ b/test/manual/test_kv_events.py @@ -115,22 +115,24 @@ class TestKvEvents(CustomTestCase): if isinstance(event, BlockStored): # Validate BlockStored structure self.assertIsInstance(event.block_hashes, list) - self.assertEqual( - len(event.block_hashes), 1, "Should have one hash per block" + self.assertGreater( + len(event.block_hashes), + 0, + "Should have at least one block hash", ) self.assertIsInstance(event.token_ids, list) self.assertEqual( - event.block_size, len(event.token_ids), - "block_size should match token_ids length", + event.block_size * len(event.block_hashes), + "token_ids should contain one block_size chunk per hash", ) self.assertIsNone( event.lora_id, "lora_id should be None for basic test" ) - # Store this block for later validation - block_hash = event.block_hashes[0] - stored_blocks[block_hash] = event + # Store every block carried by this coalesced event. + for block_hash in event.block_hashes: + stored_blocks[block_hash] = event # If parent_block_hash is set, verify it was stored earlier if event.parent_block_hash is not None: @@ -151,6 +153,13 @@ class TestKvEvents(CustomTestCase): self.assertGreater( len(stored_blocks), 0, "Should have at least one BlockStored event" ) + self.assertTrue( + any( + isinstance(event, BlockStored) and len(event.block_hashes) > 1 + for event in events + ), + "Expected at least one coalesced BlockStored event", + ) # BlockRemoved events may not always occur in this short test, so just check if they do occur # that they reference previously stored blocks for removed_hash in removed_hashes: diff --git a/test/registered/unit/mem_cache/test_hiradix_cache_unit.py b/test/registered/unit/mem_cache/test_hiradix_cache_unit.py index 21c6e0851..2c0b0a205 100644 --- a/test/registered/unit/mem_cache/test_hiradix_cache_unit.py +++ b/test/registered/unit/mem_cache/test_hiradix_cache_unit.py @@ -111,14 +111,12 @@ class TestHiRadixCacheKVEvents(CustomTestCase): cache.writing_check(write_back=True) - # Both split fragments must be published, with intact parentage. + # Both split fragments must be published as one parent-linked batch. stored_cpu = self._stored_cpu_events(cache) - self.assertEqual( - [list(e.token_ids) for e in stored_cpu], - [[1, 2], [3, 4]], - ) + self.assertEqual(len(stored_cpu), 1) + self.assertEqual(list(stored_cpu[0].token_ids), [1, 2, 3, 4]) self.assertIsNone(stored_cpu[0].parent_block_hash) - self.assertEqual(stored_cpu[1].parent_block_hash, stored_cpu[0].block_hashes[0]) + self.assertEqual(len(stored_cpu[0].block_hashes), 2) if __name__ == "__main__": diff --git a/test/registered/unit/mem_cache/test_mamba_unittest.py b/test/registered/unit/mem_cache/test_mamba_unittest.py index d763e548a..c69d2e2d2 100755 --- a/test/registered/unit/mem_cache/test_mamba_unittest.py +++ b/test/registered/unit/mem_cache/test_mamba_unittest.py @@ -464,9 +464,11 @@ class TestMamba(unittest.TestCase): ) events = tree.take_events() stored_events = [e for e in events if isinstance(e, BlockStored)] - self.assertEqual(len(stored_events), 3) - self.assertEqual([e.token_ids[0] for e in stored_events], [1, 2, 3]) - stored_hashes.extend(e.block_hashes[0] for e in stored_events) + self.assertEqual(len(stored_events), 1) + self.assertEqual(list(stored_events[0].token_ids), [1, 2, 3]) + stored_hashes.extend( + block_hash for event in stored_events for block_hash in event.block_hashes + ) req2 = make_dummy_req() key2 = RadixKey(array("q", [1, 2, 3, 4, 5])) @@ -479,9 +481,11 @@ class TestMamba(unittest.TestCase): ) events = tree.take_events() stored_events = [e for e in events if isinstance(e, BlockStored)] - self.assertEqual(len(stored_events), 2) - self.assertEqual([e.token_ids[0] for e in stored_events], [4, 5]) - stored_hashes.extend(e.block_hashes[0] for e in stored_events) + self.assertEqual(len(stored_events), 1) + self.assertEqual(list(stored_events[0].token_ids), [4, 5]) + stored_hashes.extend( + block_hash for event in stored_events for block_hash in event.block_hashes + ) # Evicting an internal mamba state creates a tombstone but does not # remove full-attention KV blocks, so it must not emit BlockRemoved. @@ -517,8 +521,8 @@ class TestMamba(unittest.TestCase): first_insert_events = [ e for e in tree.take_events() if isinstance(e, BlockStored) ] - self.assertEqual(len(first_insert_events), 4) - split_parent_hash = first_insert_events[1].block_hashes[0] + self.assertEqual(len(first_insert_events), 1) + split_parent_hash = first_insert_events[0].block_hashes[1] req2 = make_dummy_req() key2 = RadixKey(array("q", [1, 2, 5, 6])) @@ -532,8 +536,8 @@ class TestMamba(unittest.TestCase): second_insert_events = [ e for e in tree.take_events() if isinstance(e, BlockStored) ] - self.assertEqual(len(second_insert_events), 2) - self.assertEqual(list(second_insert_events[0].token_ids), [5]) + self.assertEqual(len(second_insert_events), 1) + self.assertEqual(list(second_insert_events[0].token_ids), [5, 6]) self.assertEqual(second_insert_events[0].parent_block_hash, split_parent_hash) def _setup_tree_and_allocator(self, enable_kv_cache_events=False): diff --git a/test/registered/unit/mem_cache/test_radix_cache_unit.py b/test/registered/unit/mem_cache/test_radix_cache_unit.py index d9a2808cf..3a4656f4d 100644 --- a/test/registered/unit/mem_cache/test_radix_cache_unit.py +++ b/test/registered/unit/mem_cache/test_radix_cache_unit.py @@ -30,7 +30,14 @@ from array import array import torch -from sglang.srt.disaggregation.kv_events import BlockRemoved, BlockStored +from sglang.srt.disaggregation.kv_events import ( + AllBlocksCleared, + BlockRemoved, + BlockStored, + BlockStoredMetadata, + BlockStoredWithMetadata, + StorageMedium, +) from sglang.srt.mem_cache.allocator.token import TokenToKVPoolAllocator from sglang.srt.mem_cache.base_prefix_cache import ( EvictParams, @@ -38,6 +45,7 @@ from sglang.srt.mem_cache.base_prefix_cache import ( InsertParams, MatchPrefixParams, ) +from sglang.srt.mem_cache.events import KVCacheEventMixin from sglang.srt.mem_cache.mamba_radix_cache import TreeNode as MambaTreeNode from sglang.srt.mem_cache.radix_cache import RadixCache, RadixKey, TreeNode from sglang.srt.utils import get_device @@ -46,6 +54,100 @@ from sglang.srt.utils import get_device DEFAULT_PAGE_SIZE = 4 +class _KVCacheEventQueue(KVCacheEventMixin): + def __init__(self): + self.enable_kv_cache_events = True + self.kv_event_queue = [] + + +class TestKVCacheEventQueue(unittest.TestCase): + @staticmethod + def _store( + block_hash: int, + parent_block_hash: int | None, + *, + block_size: int = 2, + medium: StorageMedium = StorageMedium.GPU, + lora_id: int | None = None, + cache_salt: str | None = None, + ) -> BlockStored: + event_args = dict( + block_hashes=[block_hash], + parent_block_hash=parent_block_hash, + token_ids=[block_hash, block_hash + 1][:block_size], + block_size=block_size, + lora_id=lora_id, + medium=medium, + ) + if cache_salt is None: + return BlockStored(**event_args) + return BlockStoredWithMetadata( + **event_args, + metadata=BlockStoredMetadata(cache_salt=cache_salt), + ) + + def test_enqueue_coalesces_compatible_stores(self): + queue = _KVCacheEventQueue() + queue._enqueue_kv_event(self._store(1, None)) + queue._enqueue_kv_event(self._store(2, 1)) + + events = queue.take_events() + self.assertEqual(len(events), 1) + self.assertEqual(events[0].block_hashes, [1, 2]) + self.assertEqual(events[0].parent_block_hash, None) + self.assertEqual(events[0].token_ids, [1, 2, 2, 3]) + + def test_enqueue_coalesces_compatible_removes(self): + queue = _KVCacheEventQueue() + queue._enqueue_kv_event( + BlockRemoved(block_hashes=[1], medium=StorageMedium.GPU) + ) + queue._enqueue_kv_event( + BlockRemoved(block_hashes=[2, 3], medium=StorageMedium.GPU) + ) + + events = queue.take_events() + self.assertEqual(len(events), 1) + self.assertIsInstance(events[0], BlockRemoved) + self.assertEqual(events[0].block_hashes, [1, 2, 3]) + + def test_enqueue_preserves_fusion_boundaries(self): + incompatible_stores = [ + self._store(2, 1, medium=StorageMedium.CPU), + self._store(3, 1, lora_id=1), + self._store(4, 1, block_size=1), + self._store(5, None), + ] + for incoming in incompatible_stores: + queue = _KVCacheEventQueue() + queue._enqueue_kv_event(self._store(1, None)) + queue._enqueue_kv_event(incoming) + self.assertEqual(len(queue.take_events()), 2) + + queue = _KVCacheEventQueue() + queue._enqueue_kv_event(self._store(1, None)) + queue._enqueue_kv_event( + BlockRemoved(block_hashes=[1], medium=StorageMedium.GPU) + ) + queue._enqueue_kv_event(AllBlocksCleared()) + queue._enqueue_kv_event(self._store(2, None)) + self.assertEqual(len(queue.take_events()), 4) + + queue = _KVCacheEventQueue() + queue._enqueue_kv_event( + BlockRemoved(block_hashes=[1], medium=StorageMedium.GPU) + ) + queue._enqueue_kv_event( + BlockRemoved(block_hashes=[2], medium=StorageMedium.CPU) + ) + self.assertEqual(len(queue.take_events()), 2) + + queue = _KVCacheEventQueue() + queue._enqueue_kv_event(self._store(1, None, cache_salt="tenant-a")) + queue._enqueue_kv_event(self._store(2, 1, cache_salt="tenant-b")) + self.assertEqual(len(queue.take_events()), 2) + + class TestRadixKey(unittest.TestCase): """Test cases for RadixKey class.""" @@ -484,7 +586,11 @@ class TestRadixCache(unittest.TestCase): ] self.assertGreater(len(block_stored_events), 0) for event in block_stored_events: - self.assertLessEqual(len(event.token_ids), page_size) + self.assertLessEqual(event.block_size, page_size) + self.assertEqual( + len(event.token_ids), + event.block_size * len(event.block_hashes), + ) else: self.assertEqual(len(events), 0) @@ -524,7 +630,10 @@ class TestRadixCache(unittest.TestCase): self.assertIn("BlockStored", event_types) stored_hashes = [ - event.block_hashes[0] for event in events if isinstance(event, BlockStored) + block_hash + for event in events + if isinstance(event, BlockStored) + for block_hash in event.block_hashes ] self.assertEqual(len(stored_hashes), 2) @@ -640,26 +749,21 @@ class TestRadixCache(unittest.TestCase): stored = [event for event in events if isinstance(event, BlockStored)] removed = [event for event in events if isinstance(event, BlockRemoved)] - self.assertEqual(len(stored), 2) - self.assertTrue( - all(event.metadata.cache_salt == "tenant-a" for event in stored) - ) - self.assertEqual(stored[1].parent_block_hash, stored[0].block_hashes[0]) - self.assertEqual( - removed[0].block_hashes, - [event.block_hashes[0] for event in stored], - ) + self.assertEqual(len(stored), 1) + self.assertEqual(stored[0].metadata.cache_salt, "tenant-a") + self.assertEqual(stored[0].parent_block_hash, None) + self.assertEqual(len(stored[0].block_hashes), 2) + self.assertEqual(removed[0].block_hashes, stored[0].block_hashes) unsalted = RadixCache.create_simulated(page_size=2, enable_kv_cache_events=True) unsalted.insert(InsertParams(key=RadixKey(tokens), value=None)) unsalted_hashes = [ - event.block_hashes[0] + block_hash for event in unsalted.take_events() if isinstance(event, BlockStored) + for block_hash in event.block_hashes ] - self.assertNotEqual( - unsalted_hashes, [event.block_hashes[0] for event in stored] - ) + self.assertNotEqual(unsalted_hashes, stored[0].block_hashes) def test_cache_salt_event_hashes_are_preserved_across_node_split(self): cache = RadixCache.create_simulated(page_size=2, enable_kv_cache_events=True) @@ -882,12 +986,12 @@ class TestRadixCache(unittest.TestCase): events = cache.take_events() block_stored_events = [e for e in events if isinstance(e, BlockStored)] - # Should have 2 blocks (2 pages of size 4) - self.assertEqual(len(block_stored_events), 2) + # The two pages should be represented by one parent-linked store event. + self.assertEqual(len(block_stored_events), 1) + self.assertEqual(len(block_stored_events[0].block_hashes), 2) # Extract block hashes - block_hash_1 = block_stored_events[0].block_hashes[0] - block_hash_2 = block_stored_events[1].block_hashes[0] + block_hash_1, block_hash_2 = block_stored_events[0].block_hashes # The two blocks should have DIFFERENT hashes despite same content # because they are at different positions (sequence-aware hashing) @@ -897,12 +1001,9 @@ class TestRadixCache(unittest.TestCase): "Repeating token patterns should get different sequence-aware hashes", ) - # First block should have no parent + # The coalesced event keeps the original root parent and ordered hashes. self.assertIsNone(block_stored_events[0].parent_block_hash) - # Second block's parent should be the first block's hash - self.assertEqual(block_stored_events[1].parent_block_hash, block_hash_1) - def test_hash_value_split(self): """Test that hash_value is split correctly when nodes are split.""" cache = RadixCache.create_simulated( diff --git a/test/registered/unit/mem_cache/test_swa_unittest.py b/test/registered/unit/mem_cache/test_swa_unittest.py index 0b451af5c..ac1ccd04e 100644 --- a/test/registered/unit/mem_cache/test_swa_unittest.py +++ b/test/registered/unit/mem_cache/test_swa_unittest.py @@ -150,18 +150,20 @@ class TestSWA(unittest.TestCase): first_insert_events = [ e for e in tree.take_events() if isinstance(e, BlockStored) ] - self.assertEqual(len(first_insert_events), 4) - self.assertEqual([e.token_ids[0] for e in first_insert_events], [1, 2, 3, 4]) + self.assertEqual(len(first_insert_events), 1) + self.assertEqual(list(first_insert_events[0].token_ids), [1, 2, 3, 4]) _insert(tree, allocator, [1, 2, 3, 4, 5, 6]) second_insert_events = [ e for e in tree.take_events() if isinstance(e, BlockStored) ] - self.assertEqual(len(second_insert_events), 2) - self.assertEqual([e.token_ids[0] for e in second_insert_events], [5, 6]) + self.assertEqual(len(second_insert_events), 1) + self.assertEqual(list(second_insert_events[0].token_ids), [5, 6]) stored_hashes = [ - e.block_hashes[0] for e in first_insert_events + second_insert_events + block_hash + for event in first_insert_events + second_insert_events + for block_hash in event.block_hashes ] # Evicting only SWA tokens tombstones nodes but keeps full KV blocks. @@ -189,15 +191,15 @@ class TestSWA(unittest.TestCase): first_insert_events = [ e for e in tree.take_events() if isinstance(e, BlockStored) ] - self.assertEqual(len(first_insert_events), 4) - split_parent_hash = first_insert_events[1].block_hashes[0] + self.assertEqual(len(first_insert_events), 1) + split_parent_hash = first_insert_events[0].block_hashes[1] _insert(tree, allocator, [1, 2, 5, 6]) second_insert_events = [ e for e in tree.take_events() if isinstance(e, BlockStored) ] - self.assertEqual(len(second_insert_events), 2) - self.assertEqual(list(second_insert_events[0].token_ids), [5]) + self.assertEqual(len(second_insert_events), 1) + self.assertEqual(list(second_insert_events[0].token_ids), [5, 6]) self.assertEqual(second_insert_events[0].parent_block_hash, split_parent_hash) def test_swa_memory_pool_paged_free_clears_full_page_mapping(self): 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 2e17223ed..f8bb9866d 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 @@ -559,9 +559,9 @@ class TestUnifiedRadixCacheKVEvents(CustomTestCase): seq = [1, 2, 3, 4] self._insert(cache, allocator, seq) stored = self._stored_events(cache, 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] + self.assertEqual(len(stored), 1) + self.assertEqual(list(stored[0].token_ids), [1, 2, 3, 4]) + stored_hashes = self._event_hashes(stored) result = cache.evict(EvictParams(num_tokens=len(seq))) self.assertGreaterEqual(result.num_tokens_evicted, len(seq)) @@ -575,7 +575,7 @@ class TestUnifiedRadixCacheKVEvents(CustomTestCase): self._insert(cache, allocator, [1, 2, 3, 4]) first_insert = self._stored_events(cache, StorageMedium.GPU) - self.assertEqual(len(first_insert), 2) + self.assertEqual(len(first_insert), 1) split_parent_hash = first_insert[0].block_hashes[0] self._insert(cache, allocator, [1, 2, 5, 6]) @@ -599,13 +599,13 @@ class TestUnifiedRadixCacheKVEvents(CustomTestCase): seq = [1, 2, 3, 4] self._insert(cache, allocator, seq) stored_gpu = self._stored_events(cache, StorageMedium.GPU) - self.assertEqual(len(stored_gpu), 2) - stored_hashes = [e.block_hashes[0] for e in stored_gpu] + self.assertEqual(len(stored_gpu), 1) + stored_hashes = self._event_hashes(stored_gpu) node = self._leaf_for(cache, seq) self._backup_node(cache, node) stored_cpu = self._stored_events(cache, StorageMedium.CPU) - self.assertCountEqual([e.block_hashes[0] for e in stored_cpu], stored_hashes) + self.assertCountEqual(self._event_hashes(stored_cpu), stored_hashes) cache.evict(EvictParams(num_tokens=len(seq))) removed_gpu = self._removed_events(cache, StorageMedium.GPU) @@ -613,7 +613,7 @@ class TestUnifiedRadixCacheKVEvents(CustomTestCase): self._load_back_node(cache, node) restored_gpu = self._stored_events(cache, StorageMedium.GPU) - self.assertCountEqual([e.block_hashes[0] for e in restored_gpu], stored_hashes) + self.assertCountEqual(self._event_hashes(restored_gpu), stored_hashes) cache.evict(EvictParams(num_tokens=len(seq))) self._removed_events(cache, StorageMedium.GPU) @@ -648,14 +648,12 @@ class TestUnifiedRadixCacheKVEvents(CustomTestCase): [[1, 2], [3, 4]], ) - # Both split fragments must be published, with intact parentage. + # Both split fragments must be published as one parent-linked batch. stored_cpu = self._stored_events(cache, StorageMedium.CPU) - self.assertEqual( - [list(e.token_ids) for e in stored_cpu], - [[1, 2], [3, 4]], - ) + self.assertEqual(len(stored_cpu), 1) + self.assertEqual(list(stored_cpu[0].token_ids), [1, 2, 3, 4]) self.assertIsNone(stored_cpu[0].parent_block_hash) - self.assertEqual(stored_cpu[1].parent_block_hash, stored_cpu[0].block_hashes[0]) + self.assertEqual(len(stored_cpu[0].block_hashes), 2) def test_hicache_reinsert_evicted_node_emits_gpu_store(self): cache, allocator, _ = build_fixture(self.cfg, enable_kv_cache_events=True) @@ -665,8 +663,8 @@ class TestUnifiedRadixCacheKVEvents(CustomTestCase): seq = [1, 2, 3, 4] self._insert(cache, allocator, seq) stored_gpu = self._stored_events(cache, StorageMedium.GPU) - self.assertEqual(len(stored_gpu), 2) - stored_hashes = [e.block_hashes[0] for e in stored_gpu] + self.assertEqual(len(stored_gpu), 1) + stored_hashes = self._event_hashes(stored_gpu) node = self._leaf_for(cache, seq) self._backup_node(cache, node) @@ -680,7 +678,7 @@ class TestUnifiedRadixCacheKVEvents(CustomTestCase): self._insert(cache, allocator, seq) restored_gpu = self._stored_events(cache, StorageMedium.GPU) self.assertFalse(node.evicted) - self.assertCountEqual([e.block_hashes[0] for e in restored_gpu], stored_hashes) + self.assertCountEqual(self._event_hashes(restored_gpu), stored_hashes) class UnifiedRadixCacheSuite: