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
@@ -11,6 +11,7 @@ from typing import TYPE_CHECKING, Any, Optional
import torch
from sglang.srt.disaggregation.kv_events import StorageMedium
from sglang.srt.mem_cache.base_prefix_cache import (
BasePrefixCache,
DecLockRefParams,
@@ -24,6 +25,7 @@ from sglang.srt.mem_cache.base_prefix_cache import (
MatchPrefixParams,
MatchResult,
)
from sglang.srt.mem_cache.events import KVCacheEventMixin
from sglang.srt.mem_cache.hicache_storage import (
PoolName,
PoolTransfer,
@@ -220,7 +222,7 @@ COMPONENT_REGISTRY: dict[ComponentType, type[TreeComponent]] = {
logger = logging.getLogger(__name__)
class UnifiedRadixCache(BasePrefixCache):
class UnifiedRadixCache(KVCacheEventMixin, BasePrefixCache):
def __init__(
self,
params: CacheInitParams,
@@ -230,6 +232,8 @@ class UnifiedRadixCache(BasePrefixCache):
self.page_size = params.page_size
self.disable = params.disable
self.is_eagle = params.is_eagle
self.enable_kv_cache_events = params.enable_kv_cache_events
self.kv_event_queue = []
if self.token_to_kv_pool_allocator:
self.device = self.token_to_kv_pool_allocator.device
@@ -328,6 +332,7 @@ class UnifiedRadixCache(BasePrefixCache):
last_host_node=self.root_node,
best_match_node=self.root_node,
)
self._record_all_cleared_event()
def init_hicache(self, server_args: ServerArgs, params: CacheInitParams) -> None:
"""Initialize HiCache infrastructure."""
@@ -869,6 +874,7 @@ class UnifiedRadixCache(BasePrefixCache):
self._update_evictable_leaf_sets(new_node)
self._update_evictable_leaf_sets(parent)
self._record_store_event(new_node)
return new_node
def _unevict_node_on_insert(
@@ -885,6 +891,7 @@ class UnifiedRadixCache(BasePrefixCache):
self._update_evictable_leaf_sets(node)
if node.parent is not None:
self._update_evictable_leaf_sets(node.parent)
self._record_store_event(node, medium=StorageMedium.GPU)
def _insert_helper(
self,
@@ -1241,6 +1248,7 @@ class UnifiedRadixCache(BasePrefixCache):
node, trigger, target=EvictLayer.DEVICE, tracker=tracker
)
self._cascade_evict(node, trigger, tracker)
self._record_remove_event(node, medium=StorageMedium.GPU)
# after device eviction, insert aux components into host LRU.
self._for_each_component_lru(
@@ -1271,6 +1279,7 @@ class UnifiedRadixCache(BasePrefixCache):
return
else:
# Write-through: node has no backup, delete entirely.
self._record_remove_event(node, medium=StorageMedium.GPU)
for comp in self._components_tuple:
self._evict_component_and_detach_lru(
node, comp, target=EvictLayer.ALL, tracker=tracker
@@ -1291,6 +1300,7 @@ class UnifiedRadixCache(BasePrefixCache):
All freed tokens are accumulated into *tracker*."""
assert self._is_host_leaf(node), f"node {node.id} is not an H-leaf"
self._record_remove_event(node, medium=StorageMedium.CPU)
for comp in self._components_tuple:
_, hf = self._evict_component_and_detach_lru(
node, comp, target=EvictLayer.ALL, tracker=None
@@ -1437,6 +1447,8 @@ class UnifiedRadixCache(BasePrefixCache):
CacheTransferPhase.LOAD_BACK,
[kv_xfer],
)
for node in kv_xfer.nodes_to_load or ():
self._record_store_event(node, medium=StorageMedium.GPU)
for ct, xfers in comp_xfers.items():
self.components[ct].commit_hicache_transfer(
best_match_node,
@@ -2031,6 +2043,7 @@ class UnifiedRadixCache(BasePrefixCache):
entry = self.ongoing_write_through.pop(ack_id, None)
if entry is not None:
node, params = entry
self._record_store_event(node, medium=StorageMedium.CPU)
if params is not None:
self.dec_lock_ref(node, params)
if self.enable_storage:
@@ -2062,6 +2075,7 @@ class UnifiedRadixCache(BasePrefixCache):
finish_event.synchronize()
for ack_id in ack_list:
node, params = self.ongoing_write_through.pop(ack_id)
self._record_store_event(node, medium=StorageMedium.CPU)
self.dec_lock_ref(node, params)
if self.enable_storage:
self.write_backup_storage(node)
@@ -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