Refactor kv cache event mixin into a recorder (#35164)

This commit is contained in:
Ke Bao
2026-08-19 01:20:17 +08:00
committed by GitHub
parent 83d7d45330
commit 480033def0
12 changed files with 207 additions and 195 deletions
@@ -45,7 +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.events import KVCacheEventRecorder
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
@@ -54,12 +54,6 @@ 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(
@@ -87,26 +81,22 @@ class TestKVCacheEventQueue(unittest.TestCase):
)
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))
queue = KVCacheEventRecorder(enabled=True, page_size=DEFAULT_PAGE_SIZE)
queue.enqueue(self._store(1, None))
queue.enqueue(self._store(2, 1))
events = queue.take_events()
events = queue.take()
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)
)
queue = KVCacheEventRecorder(enabled=True, page_size=DEFAULT_PAGE_SIZE)
queue.enqueue(BlockRemoved(block_hashes=[1], medium=StorageMedium.GPU))
queue.enqueue(BlockRemoved(block_hashes=[2, 3], medium=StorageMedium.GPU))
events = queue.take_events()
events = queue.take()
self.assertEqual(len(events), 1)
self.assertIsInstance(events[0], BlockRemoved)
self.assertEqual(events[0].block_hashes, [1, 2, 3])
@@ -119,33 +109,27 @@ class TestKVCacheEventQueue(unittest.TestCase):
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 = KVCacheEventRecorder(enabled=True, page_size=DEFAULT_PAGE_SIZE)
queue.enqueue(self._store(1, None))
queue.enqueue(incoming)
self.assertEqual(len(queue.take()), 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 = KVCacheEventRecorder(enabled=True, page_size=DEFAULT_PAGE_SIZE)
queue.enqueue(self._store(1, None))
queue.enqueue(BlockRemoved(block_hashes=[1], medium=StorageMedium.GPU))
queue.enqueue(AllBlocksCleared())
queue.enqueue(self._store(2, None))
self.assertEqual(len(queue.take()), 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 = KVCacheEventRecorder(enabled=True, page_size=DEFAULT_PAGE_SIZE)
queue.enqueue(BlockRemoved(block_hashes=[1], medium=StorageMedium.GPU))
queue.enqueue(BlockRemoved(block_hashes=[2], medium=StorageMedium.CPU))
self.assertEqual(len(queue.take()), 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)
queue = KVCacheEventRecorder(enabled=True, page_size=DEFAULT_PAGE_SIZE)
queue.enqueue(self._store(1, None, cache_salt="tenant-a"))
queue.enqueue(self._store(2, 1, cache_salt="tenant-b"))
self.assertEqual(len(queue.take()), 2)
class TestRadixKey(unittest.TestCase):
@@ -410,7 +394,7 @@ class TestRadixCache(unittest.TestCase):
self.assertEqual(cache.page_size, page_size)
self.assertEqual(cache.disable, disable)
self.assertEqual(cache.enable_kv_cache_events, enable_events)
self.assertEqual(cache.kv_events.enabled, enable_events)
self.assertEqual(cache.device, torch.device("cpu"))
self.assertIsNotNone(cache.root_node)
self.assertEqual(len(cache.root_node.key), 0)