feat(agent sessions): attribute stored KV cache blocks to sessions (#37482)
Signed-off-by: Ishan Dhanani <ishandhanani@gmail.com>
This commit is contained in:
@@ -12,9 +12,9 @@ import unittest
|
||||
import msgspec
|
||||
|
||||
from sglang.srt.disaggregation.kv_events import (
|
||||
AllBlocksCleared,
|
||||
BlockRemoved,
|
||||
BlockStored,
|
||||
BlockStoredMetadata,
|
||||
BlockStoredWithMetadata,
|
||||
KVEventBatch,
|
||||
StorageMedium,
|
||||
ZmqEventPublisher,
|
||||
@@ -186,42 +186,67 @@ class TestSelectKvPublisherDpRank(CustomTestCase):
|
||||
|
||||
|
||||
class TestBlockStoredWireFormat(CustomTestCase):
|
||||
def _event(self, metadata=None):
|
||||
event_type = BlockStored if metadata is None else BlockStoredWithMetadata
|
||||
kwargs = dict(
|
||||
def _event(self, **extra):
|
||||
return BlockStored(
|
||||
block_hashes=[123],
|
||||
parent_block_hash=None,
|
||||
token_ids=[1, 2],
|
||||
block_size=2,
|
||||
lora_id=None,
|
||||
medium=StorageMedium.GPU,
|
||||
**extra,
|
||||
)
|
||||
if metadata is not None:
|
||||
kwargs["metadata"] = metadata
|
||||
return event_type(**kwargs)
|
||||
|
||||
def test_unsalted_event_keeps_legacy_array_shape(self):
|
||||
def test_event_is_a_tagged_map(self):
|
||||
decoded = msgspec.msgpack.decode(msgspec.msgpack.encode(self._event()))
|
||||
self.assertEqual(len(decoded), 7)
|
||||
self.assertIsInstance(decoded, dict)
|
||||
self.assertEqual(decoded["type"], "BlockStored")
|
||||
self.assertEqual(
|
||||
set(decoded),
|
||||
{
|
||||
"type",
|
||||
"block_hashes",
|
||||
"parent_block_hash",
|
||||
"token_ids",
|
||||
"block_size",
|
||||
"lora_id",
|
||||
"medium",
|
||||
},
|
||||
)
|
||||
|
||||
def test_salted_event_appends_typed_metadata(self):
|
||||
event = self._event(BlockStoredMetadata(cache_salt="tenant-a"))
|
||||
encoded = msgspec.msgpack.encode(event)
|
||||
decoded = msgspec.msgpack.decode(encoded)
|
||||
round_tripped = msgspec.msgpack.decode(encoded, type=BlockStoredWithMetadata)
|
||||
self.assertEqual(len(decoded), 8)
|
||||
self.assertEqual(decoded[7], {"cache_salt": "tenant-a"})
|
||||
self.assertEqual(round_tripped.metadata.cache_salt, "tenant-a")
|
||||
def test_salt_and_session_are_named_fields(self):
|
||||
event = self._event(cache_salt="tenant-a", session_id="session-a")
|
||||
decoded = msgspec.msgpack.decode(msgspec.msgpack.encode(event))
|
||||
self.assertEqual(decoded["cache_salt"], "tenant-a")
|
||||
self.assertEqual(decoded["session_id"], "session-a")
|
||||
|
||||
def test_salted_event_remains_compatible_with_typed_batch_consumers(self):
|
||||
def test_one_decoder_reads_a_mixed_batch(self):
|
||||
batch = KVEventBatch(
|
||||
ts=1.0,
|
||||
events=[self._event(BlockStoredMetadata(cache_salt="tenant-a"))],
|
||||
events=[
|
||||
self._event(),
|
||||
self._event(cache_salt="tenant-a"),
|
||||
self._event(session_id="session-a"),
|
||||
BlockRemoved(block_hashes=[123], medium=StorageMedium.GPU),
|
||||
AllBlocksCleared(),
|
||||
],
|
||||
)
|
||||
round_tripped = msgspec.msgpack.decode(
|
||||
msgspec.msgpack.encode(batch), type=KVEventBatch
|
||||
)
|
||||
self.assertEqual(round_tripped.events[0].block_hashes, [123])
|
||||
stored = round_tripped.events[:3]
|
||||
self.assertEqual([e.cache_salt for e in stored], [None, "tenant-a", None])
|
||||
self.assertEqual([e.session_id for e in stored], [None, None, "session-a"])
|
||||
self.assertIsInstance(round_tripped.events[3], BlockRemoved)
|
||||
self.assertIsInstance(round_tripped.events[4], AllBlocksCleared)
|
||||
|
||||
def test_batch_stays_a_positional_array_of_maps(self):
|
||||
batch = KVEventBatch(ts=1.0, events=[self._event()], attn_dp_rank=0)
|
||||
decoded = msgspec.msgpack.decode(msgspec.msgpack.encode(batch))
|
||||
self.assertEqual(decoded[0], 1.0)
|
||||
self.assertEqual(decoded[2], 0)
|
||||
self.assertIsInstance(decoded[1][0], dict)
|
||||
self.assertEqual(len(decoded), 3)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
@@ -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