Support KV events for UnifiedRadixCache (#26387)

Co-authored-by: Zhangheng <hzh0425@apache.org>
This commit is contained in:
weireweire
2026-05-27 22:10:40 +08:00
committed by GitHub
co-authored by Zhangheng
parent 5f8911183b
commit 034dd39189
2 changed files with 184 additions and 2 deletions
@@ -9,6 +9,11 @@ from unittest import mock
import torch
from sglang.srt.configs.mamba_utils import Mamba2CacheParams, Mamba2StateShape
from sglang.srt.disaggregation.kv_events import (
BlockRemoved,
BlockStored,
StorageMedium,
)
from sglang.srt.environ import envs
from sglang.srt.layers.attention.fla.chunk_delta_h import CHUNK_SIZE as FLA_CHUNK_SIZE
from sglang.srt.managers.schedule_batch import Req
@@ -113,7 +118,7 @@ class CacheConfig:
return "_".join(parts)
def build_fixture(cfg: CacheConfig):
def build_fixture(cfg: CacheConfig, *, enable_kv_cache_events: bool = False):
"""Create (tree, allocator, req_to_token_pool) from a CacheConfig."""
server_args = ServerArgs(model_path="dummy", page_size=cfg.page_size)
# MambaRadixCache reads mamba_cache_chunk_size, whose property otherwise
@@ -227,6 +232,7 @@ def build_fixture(cfg: CacheConfig):
sliding_window_size=cfg.sliding_window_size,
tree_components=cfg.components,
enable_mamba_extra_buffer=cfg.enable_mamba_extra_buffer,
enable_kv_cache_events=enable_kv_cache_events,
)
tree = UnifiedRadixCache(params=cache_init_params)
tree.cache_init_params = cache_init_params
@@ -234,6 +240,168 @@ def build_fixture(cfg: CacheConfig):
return tree, allocator, req_to_token_pool
class TestUnifiedRadixCacheKVEvents(CustomTestCase):
cfg = CacheConfig(page_size=2, kv_size=64, max_context_len=64)
def _insert(self, tree, allocator, tokens):
key = RadixKey(array("q", tokens))
value = allocator.alloc(len(tokens))
self.assertIsNotNone(value)
return tree.insert(InsertParams(key=key, value=value[: len(key)]))
def _stored_events(self, tree, medium=None):
events = [e for e in tree.take_events() if isinstance(e, BlockStored)]
if medium is not None:
events = [e for e in events if e.medium == medium]
return events
def _removed_events(self, tree, medium=None):
events = [e for e in tree.take_events() if isinstance(e, BlockRemoved)]
if medium is not None:
events = [e for e in events if e.medium == medium]
return events
def _leaf_for(self, tree, tokens):
match = tree.match_prefix(MatchPrefixParams(key=RadixKey(array("q", tokens))))
self.assertIsNot(match.last_device_node, tree.root_node)
return match.last_device_node
def _init_hicache(self, tree, *, write_policy: str = "write_through"):
import sglang.srt.mem_cache.hybrid_cache.hybrid_pool_assembler as assembler
orig_kv_host_pool = assembler.MHATokenToKVPoolHost
def kv_host_pool_wrapper(*args, **kwargs):
kwargs["pin_memory"] = False
return orig_kv_host_pool(*args, **kwargs)
patcher = mock.patch.object(
assembler,
"MHATokenToKVPoolHost",
side_effect=kv_host_pool_wrapper,
)
patcher.start()
self.addCleanup(patcher.stop)
server_args = ServerArgs(
model_path="dummy",
page_size=self.cfg.page_size,
hicache_io_backend="direct",
hicache_write_policy=write_policy,
)
set_global_server_args_for_scheduler(server_args)
tree.init_hicache(server_args, tree.cache_init_params)
tree.write_through_threshold = 1 << 30
tree.load_back_threshold = 0
def _backup_node(self, tree, node):
backed_up = tree.write_backup(node, write_back=True)
self.assertGreater(backed_up, 0)
tree.writing_check(write_back=True)
def _load_back_node(self, tree, node):
loaded = tree.load_back(node)
self.assertTrue(loaded)
producer_id = tree.ready_to_load_host_cache()
self.assertNotEqual(producer_id, -1)
for _, finish_event, _ in list(tree.cache_controller.ack_load_queue):
finish_event.synchronize()
tree.loading_check()
def test_kv_events_store_and_remove_full_blocks(self):
tree, allocator, _ = build_fixture(self.cfg, enable_kv_cache_events=True)
tree.take_events() # Clear the reset event.
seq = [1, 2, 3, 4]
self._insert(tree, allocator, seq)
stored = self._stored_events(tree, 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]
result = tree.evict(EvictParams(num_tokens=len(seq)))
self.assertGreaterEqual(result.num_tokens_evicted, len(seq))
removed = self._removed_events(tree, StorageMedium.GPU)
self.assertCountEqual([e.block_hashes[0] for e in removed], stored_hashes)
def test_kv_events_split_preserves_block_hash_parentage(self):
tree, allocator, _ = build_fixture(self.cfg, enable_kv_cache_events=True)
tree.take_events() # Clear the reset event.
self._insert(tree, allocator, [1, 2, 3, 4])
first_insert = self._stored_events(tree, StorageMedium.GPU)
self.assertEqual(len(first_insert), 2)
split_parent_hash = first_insert[0].block_hashes[0]
self._insert(tree, allocator, [1, 2, 5, 6])
second_insert = self._stored_events(tree, StorageMedium.GPU)
self.assertEqual(len(second_insert), 1)
self.assertEqual(list(second_insert[0].token_ids), [5, 6])
self.assertEqual(second_insert[0].parent_block_hash, split_parent_hash)
split_parent = next(iter(tree.root_node.children.values()))
split_child = split_parent.children.get((3, 4))
self.assertIsNotNone(split_child)
self.assertEqual(len(split_parent.hash_value), 1)
self.assertIsNotNone(split_child.hash_value)
self.assertEqual(len(split_child.hash_value), 1)
def test_hicache_kv_events_track_gpu_cpu_transitions(self):
tree, allocator, _ = build_fixture(self.cfg, enable_kv_cache_events=True)
self._init_hicache(tree)
tree.take_events() # Clear reset / init events.
seq = [1, 2, 3, 4]
self._insert(tree, allocator, seq)
stored_gpu = self._stored_events(tree, StorageMedium.GPU)
self.assertEqual(len(stored_gpu), 2)
stored_hashes = [e.block_hashes[0] for e in stored_gpu]
node = self._leaf_for(tree, seq)
self._backup_node(tree, node)
stored_cpu = self._stored_events(tree, StorageMedium.CPU)
self.assertCountEqual([e.block_hashes[0] for e in stored_cpu], stored_hashes)
tree.evict(EvictParams(num_tokens=len(seq)))
removed_gpu = self._removed_events(tree, StorageMedium.GPU)
self.assertCountEqual([e.block_hashes[0] for e in removed_gpu], stored_hashes)
self._load_back_node(tree, node)
restored_gpu = self._stored_events(tree, StorageMedium.GPU)
self.assertCountEqual([e.block_hashes[0] for e in restored_gpu], stored_hashes)
tree.evict(EvictParams(num_tokens=len(seq)))
self._removed_events(tree, StorageMedium.GPU)
tree.evict_host(len(seq))
removed_cpu = self._removed_events(tree, StorageMedium.CPU)
self.assertCountEqual([e.block_hashes[0] for e in removed_cpu], stored_hashes)
def test_hicache_reinsert_evicted_node_emits_gpu_store(self):
tree, allocator, _ = build_fixture(self.cfg, enable_kv_cache_events=True)
self._init_hicache(tree)
tree.take_events() # Clear reset / init events.
seq = [1, 2, 3, 4]
self._insert(tree, allocator, seq)
stored_gpu = self._stored_events(tree, StorageMedium.GPU)
self.assertEqual(len(stored_gpu), 2)
stored_hashes = [e.block_hashes[0] for e in stored_gpu]
node = self._leaf_for(tree, seq)
self._backup_node(tree, node)
self._stored_events(tree, StorageMedium.CPU)
tree.evict(EvictParams(num_tokens=len(seq)))
self._removed_events(tree, StorageMedium.GPU)
self.assertTrue(node.evicted)
self.assertTrue(node.backuped)
self._insert(tree, allocator, seq)
restored_gpu = self._stored_events(tree, StorageMedium.GPU)
self.assertFalse(node.evicted)
self.assertCountEqual([e.block_hashes[0] for e in restored_gpu], stored_hashes)
class UnifiedRadixCacheSuite:
cfg: CacheConfig