Refactor kv cache event mixin into a recorder (#35164)
This commit is contained in:
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user