feat(agent sessions): attribute stored KV cache blocks to sessions (#37482)
Signed-off-by: Ishan Dhanani <ishandhanani@gmail.com>
This commit is contained in:
@@ -612,6 +612,7 @@ class _GraftReq:
|
||||
self.swa_prefix_lock_released = False
|
||||
self.finished_reason = None
|
||||
self.session = None
|
||||
self.session_id = None
|
||||
|
||||
def get_fill_ids(self):
|
||||
return array("q", self.fill_ids)
|
||||
|
||||
@@ -34,8 +34,6 @@ from sglang.srt.disaggregation.kv_events import (
|
||||
AllBlocksCleared,
|
||||
BlockRemoved,
|
||||
BlockStored,
|
||||
BlockStoredMetadata,
|
||||
BlockStoredWithMetadata,
|
||||
StorageMedium,
|
||||
)
|
||||
from sglang.srt.managers.schedule_batch import ReqKvInfo
|
||||
@@ -66,20 +64,17 @@ class TestKVCacheEventQueue(unittest.TestCase):
|
||||
medium: StorageMedium = StorageMedium.GPU,
|
||||
lora_id: int | None = None,
|
||||
cache_salt: str | None = None,
|
||||
session_id: str | None = None,
|
||||
) -> BlockStored:
|
||||
event_args = dict(
|
||||
return BlockStored(
|
||||
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),
|
||||
cache_salt=cache_salt,
|
||||
session_id=session_id,
|
||||
)
|
||||
|
||||
def test_enqueue_coalesces_compatible_stores(self):
|
||||
@@ -133,6 +128,11 @@ class TestKVCacheEventQueue(unittest.TestCase):
|
||||
queue.enqueue(self._store(2, 1, cache_salt="tenant-b"))
|
||||
self.assertEqual(len(queue.take()), 2)
|
||||
|
||||
queue = KVCacheEventRecorder(enabled=True, page_size=DEFAULT_PAGE_SIZE)
|
||||
queue.enqueue(self._store(1, None, session_id="session-a"))
|
||||
queue.enqueue(self._store(2, 1, session_id="session-b"))
|
||||
self.assertEqual(len(queue.take()), 2)
|
||||
|
||||
|
||||
class TestRadixKey(unittest.TestCase):
|
||||
"""Test cases for RadixKey class."""
|
||||
@@ -781,7 +781,7 @@ class TestRadixCache(CustomTestCase):
|
||||
removed = [event for event in events if isinstance(event, BlockRemoved)]
|
||||
|
||||
self.assertEqual(len(stored), 1)
|
||||
self.assertEqual(stored[0].metadata.cache_salt, "tenant-a")
|
||||
self.assertEqual(stored[0].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)
|
||||
|
||||
@@ -20,8 +20,6 @@ from sglang.srt.disaggregation.kv_events import (
|
||||
AllBlocksCleared,
|
||||
BlockRemoved,
|
||||
BlockStored,
|
||||
BlockStoredMetadata,
|
||||
BlockStoredWithMetadata,
|
||||
StorageMedium,
|
||||
)
|
||||
from sglang.srt.environ import envs
|
||||
@@ -971,14 +969,14 @@ def test_salted_events_match_python_hash_and_metadata_contract():
|
||||
for value in mem_cache.get_hash_str(array("q", [1, 2, 7, 8]), seed, 2)
|
||||
]
|
||||
assert core.take_events() == [
|
||||
BlockStoredWithMetadata(
|
||||
BlockStored(
|
||||
block_hashes=hashes,
|
||||
parent_block_hash=None,
|
||||
token_ids=[1, 2, 7, 8],
|
||||
block_size=2,
|
||||
lora_id=None,
|
||||
medium=StorageMedium.GPU,
|
||||
metadata=BlockStoredMetadata(cache_salt="tenant-a"),
|
||||
cache_salt="tenant-a",
|
||||
)
|
||||
]
|
||||
|
||||
@@ -1012,14 +1010,14 @@ def test_salted_eagle_events_match_the_bigram_hash_contract():
|
||||
for value in mem_cache.get_hash_str(raw_tokens, seed, 2, is_bigram=True)
|
||||
]
|
||||
assert core.take_events() == [
|
||||
BlockStoredWithMetadata(
|
||||
BlockStored(
|
||||
block_hashes=hashes,
|
||||
parent_block_hash=None,
|
||||
token_ids=[(1, 2), (2, 3), (3, 4), (4, 5)],
|
||||
block_size=2,
|
||||
lora_id=None,
|
||||
medium=StorageMedium.GPU,
|
||||
metadata=BlockStoredMetadata(cache_salt="tenant-a"),
|
||||
cache_salt="tenant-a",
|
||||
)
|
||||
]
|
||||
|
||||
|
||||
@@ -22,7 +22,6 @@ from sglang.srt.configs.mamba_utils import Mamba2CacheParams, Mamba2StateShape
|
||||
from sglang.srt.disaggregation.kv_events import (
|
||||
BlockRemoved,
|
||||
BlockStored,
|
||||
BlockStoredWithMetadata,
|
||||
StorageMedium,
|
||||
)
|
||||
from sglang.srt.environ import envs
|
||||
@@ -1012,11 +1011,18 @@ class TestUnifiedRadixCacheKVEvents(CustomTestCase):
|
||||
*,
|
||||
extra_key=None,
|
||||
cache_salt=None,
|
||||
session_id=None,
|
||||
):
|
||||
key = RadixKey(array("q", tokens), extra_key=extra_key, cache_salt=cache_salt)
|
||||
value = allocator.alloc(len(tokens))
|
||||
self.assertIsNotNone(value)
|
||||
return cache.insert(InsertParams(key=key, value=value[: len(key)]))
|
||||
return cache.insert(
|
||||
InsertParams(
|
||||
key=key,
|
||||
value=value[: len(key)],
|
||||
session_id=session_id,
|
||||
)
|
||||
)
|
||||
|
||||
def _stored_events(self, cache, medium=None):
|
||||
events = [e for e in cache.take_events() if isinstance(e, BlockStored)]
|
||||
@@ -1096,8 +1102,7 @@ class TestUnifiedRadixCacheKVEvents(CustomTestCase):
|
||||
self._insert(cache, allocator, seq, cache_salt="tenant-a")
|
||||
stored = self._stored_events(cache, StorageMedium.GPU)
|
||||
self.assertEqual(len(stored), 1)
|
||||
self.assertIsInstance(stored[0], BlockStoredWithMetadata)
|
||||
self.assertEqual(stored[0].metadata.cache_salt, "tenant-a")
|
||||
self.assertEqual(stored[0].cache_salt, "tenant-a")
|
||||
salted_hashes = self._event_hashes(stored)
|
||||
|
||||
cache.evict(EvictParams(num_tokens=len(seq)))
|
||||
@@ -1115,6 +1120,73 @@ class TestUnifiedRadixCacheKVEvents(CustomTestCase):
|
||||
salted_hashes,
|
||||
)
|
||||
|
||||
def test_session_id_is_attributed_without_changing_block_hash(self):
|
||||
cache_a, allocator_a, _ = build_fixture(self.cfg, enable_kv_cache_events=True)
|
||||
cache_a.take_events()
|
||||
self._insert(cache_a, allocator_a, [1, 2, 3, 4], session_id="session-a")
|
||||
stored_a = self._stored_events(cache_a, StorageMedium.GPU)
|
||||
self.assertEqual(len(stored_a), 1)
|
||||
self.assertEqual(stored_a[0].session_id, "session-a")
|
||||
|
||||
cache_b, allocator_b, _ = build_fixture(self.cfg, enable_kv_cache_events=True)
|
||||
cache_b.take_events()
|
||||
self._insert(cache_b, allocator_b, [1, 2, 3, 4], session_id="session-b")
|
||||
stored_b = self._stored_events(cache_b, StorageMedium.GPU)
|
||||
self.assertEqual(stored_b[0].session_id, "session-b")
|
||||
self.assertEqual(self._event_hashes(stored_a), self._event_hashes(stored_b))
|
||||
|
||||
def test_shared_prefix_hit_is_quiet_and_divergent_tails_are_attributed(self):
|
||||
cache, allocator, _ = build_fixture(self.cfg, enable_kv_cache_events=True)
|
||||
cache.take_events()
|
||||
|
||||
shared_prefix = [1, 2, 3, 4]
|
||||
self._insert(cache, allocator, shared_prefix, session_id="session-a")
|
||||
initial = self._stored_events(cache, StorageMedium.GPU)
|
||||
self.assertEqual(len(initial), 1)
|
||||
shared_parent = initial[0].block_hashes[-1]
|
||||
|
||||
self._insert(cache, allocator, shared_prefix, session_id="session-b")
|
||||
self.assertEqual(self._stored_events(cache, StorageMedium.GPU), [])
|
||||
|
||||
self._insert(
|
||||
cache,
|
||||
allocator,
|
||||
shared_prefix + [5, 6],
|
||||
session_id="session-a",
|
||||
)
|
||||
session_a_tail = self._stored_events(cache, StorageMedium.GPU)
|
||||
self.assertEqual(len(session_a_tail), 1)
|
||||
self.assertEqual(session_a_tail[0].parent_block_hash, shared_parent)
|
||||
self.assertEqual(list(session_a_tail[0].token_ids), [5, 6])
|
||||
self.assertEqual(session_a_tail[0].session_id, "session-a")
|
||||
|
||||
self._insert(
|
||||
cache,
|
||||
allocator,
|
||||
shared_prefix + [7, 8],
|
||||
session_id="session-b",
|
||||
)
|
||||
session_b_tail = self._stored_events(cache, StorageMedium.GPU)
|
||||
self.assertEqual(len(session_b_tail), 1)
|
||||
self.assertEqual(session_b_tail[0].parent_block_hash, shared_parent)
|
||||
self.assertEqual(list(session_b_tail[0].token_ids), [7, 8])
|
||||
self.assertEqual(session_b_tail[0].session_id, "session-b")
|
||||
|
||||
def test_session_id_and_cache_salt_are_both_attributed(self):
|
||||
cache, allocator, _ = build_fixture(self.cfg, enable_kv_cache_events=True)
|
||||
cache.take_events()
|
||||
self._insert(
|
||||
cache,
|
||||
allocator,
|
||||
[1, 2, 3, 4],
|
||||
cache_salt="tenant-a",
|
||||
session_id="session-a",
|
||||
)
|
||||
stored = self._stored_events(cache, StorageMedium.GPU)
|
||||
self.assertEqual(len(stored), 1)
|
||||
self.assertEqual(stored[0].cache_salt, "tenant-a")
|
||||
self.assertEqual(stored[0].session_id, "session-a")
|
||||
|
||||
def test_cache_salt_event_parentage_survives_node_split(self):
|
||||
cache, allocator, _ = build_fixture(self.cfg, enable_kv_cache_events=True)
|
||||
cache.take_events()
|
||||
@@ -1127,8 +1199,7 @@ class TestUnifiedRadixCacheKVEvents(CustomTestCase):
|
||||
self._insert(cache, allocator, [1, 2, 5, 6], cache_salt="tenant-a")
|
||||
branch = self._stored_events(cache, StorageMedium.GPU)
|
||||
self.assertEqual(len(branch), 1)
|
||||
self.assertIsInstance(branch[0], BlockStoredWithMetadata)
|
||||
self.assertEqual(branch[0].metadata.cache_salt, "tenant-a")
|
||||
self.assertEqual(branch[0].cache_salt, "tenant-a")
|
||||
self.assertEqual(branch[0].parent_block_hash, original[0].block_hashes[0])
|
||||
self.assertEqual(list(branch[0].token_ids), [5, 6])
|
||||
|
||||
@@ -1296,10 +1367,11 @@ class TestUnifiedRadixCacheKVEvents(CustomTestCase):
|
||||
self.assertTrue(cache.tree_core.is_full_device_evicted(node))
|
||||
self.assertTrue(cache.tree_core.is_backuped(node))
|
||||
|
||||
self._insert(cache, allocator, seq)
|
||||
self._insert(cache, allocator, seq, session_id="session-a")
|
||||
restored_gpu = self._stored_events(cache, StorageMedium.GPU)
|
||||
self.assertFalse(cache.tree_core.is_full_device_evicted(node))
|
||||
self.assertCountEqual(self._event_hashes(restored_gpu), stored_hashes)
|
||||
self.assertEqual(restored_gpu[0].session_id, "session-a")
|
||||
|
||||
|
||||
class UnifiedRadixCacheSuite:
|
||||
|
||||
Reference in New Issue
Block a user