feat: Add KV events for Mamba radix cache (#23678)

Signed-off-by: zhongdaor-nv <220807034+zhongdaor-nv@users.noreply.github.com>
Co-authored-by: zhongdaor-nv <220807034+zhongdaor-nv@users.noreply.github.com>
This commit is contained in:
zhongdaor-nv
2026-05-08 11:53:36 -07:00
committed by GitHub
co-authored by zhongdaor-nv
parent ca7a8cc61d
commit 2cf1a4ab38
8 changed files with 344 additions and 158 deletions
+127
View File
@@ -0,0 +1,127 @@
# Copyright 2025 SGLang Team
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# ==============================================================================
"""KV cache placement event emission mixin.
The mixin produces the ``BlockStored`` / ``BlockRemoved`` / ``AllBlocksCleared``
events consumed by KV-aware routers (e.g. dynamo).
"""
from typing import Any
from sglang.srt.disaggregation.kv_events import (
AllBlocksCleared,
BlockRemoved,
BlockStored,
StorageMedium,
)
from sglang.srt.mem_cache.utils import (
compute_node_hash_values,
hash_str_to_int64,
)
class KVCacheEventMixin:
def _record_store_event(self, node: Any, medium=None):
# One BlockStored per ``page_size`` chunk.
# ``medium`` defaults to StorageMedium.GPU but callers may override
# for lower-tier insertions (e.g. StorageMedium.CPU for host/L2 cache).
if self.enable_kv_cache_events:
if medium is None:
medium = StorageMedium.GPU
# Compute hash_value lazily if not already set
if node.hash_value is None:
node.hash_value = compute_node_hash_values(node, self.page_size)
# Get parent's last hash value for first page
parent_block_hash = None
if node.parent is not None and node.parent != self.root_node:
if (
node.parent.hash_value is not None
and len(node.parent.hash_value) > 0
):
parent_block_hash = hash_str_to_int64(node.parent.hash_value[-1])
page_index = 0
logical_len = len(node.key)
is_bigram = node.key.is_bigram
raw = node.key.token_ids
for start in range(0, logical_len, self.page_size):
end = min(start + self.page_size, logical_len)
if end <= start:
continue
# Preserve historical event payload: bigram pages expose tuples.
if is_bigram:
page_tokens = [(raw[j], raw[j + 1]) for j in range(start, end)]
else:
page_tokens = raw[start:end]
block_hash = hash_str_to_int64(node.hash_value[page_index])
self.kv_event_queue.append(
BlockStored(
block_hashes=[block_hash],
parent_block_hash=parent_block_hash,
token_ids=page_tokens,
block_size=len(page_tokens),
lora_id=None,
medium=medium,
)
)
parent_block_hash = block_hash
page_index += 1
def _record_remove_event(self, node: Any, medium=None):
# One BlockRemoved per chunk.
# ``medium`` defaults to StorageMedium.GPU but callers may override for
# lower-tier removals (e.g. StorageMedium.CPU when evicting from host).
if self.enable_kv_cache_events:
if medium is None:
medium = StorageMedium.GPU
# Compute hash_value lazily if not already set (must match what was stored)
if node.hash_value is None:
node.hash_value = compute_node_hash_values(node, self.page_size)
page_index = 0
logical_len = len(node.key)
for start in range(0, logical_len, self.page_size):
end = min(start + self.page_size, logical_len)
if end <= start:
continue
block_hash = hash_str_to_int64(node.hash_value[page_index])
self.kv_event_queue.append(
BlockRemoved(block_hashes=[block_hash], medium=medium)
)
page_index += 1
def _record_all_cleared_event(self):
if self.enable_kv_cache_events:
self.kv_event_queue.append(AllBlocksCleared())
def take_events(self):
"""Atomically takes all events and clears the queue.
Returns:
A list of KV cache events.
"""
if not self.enable_kv_cache_events:
return []
events = self.kv_event_queue
self.kv_event_queue = []
return events
@@ -12,6 +12,7 @@ from typing import TYPE_CHECKING, Dict, List, Optional, Tuple
import torch import torch
from sglang.srt.disaggregation.kv_events import StorageMedium
from sglang.srt.mem_cache.base_prefix_cache import ( from sglang.srt.mem_cache.base_prefix_cache import (
DecLockRefParams, DecLockRefParams,
DecLockRefResult, DecLockRefResult,
@@ -38,9 +39,8 @@ from sglang.srt.mem_cache.mamba_radix_cache import (
from sglang.srt.mem_cache.memory_pool import HybridLinearKVPool, HybridReqToTokenPool from sglang.srt.mem_cache.memory_pool import HybridLinearKVPool, HybridReqToTokenPool
from sglang.srt.mem_cache.radix_cache import ( from sglang.srt.mem_cache.radix_cache import (
RadixKey, RadixKey,
compute_node_hash_values,
split_node_hash_value,
) )
from sglang.srt.mem_cache.utils import compute_node_hash_values, split_node_hash_value
from sglang.srt.observability.metrics_collector import StorageMetricsCollector from sglang.srt.observability.metrics_collector import StorageMetricsCollector
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -319,6 +319,7 @@ class HiMambaRadixCache(MambaRadixCache):
n.value = full_device_indices[offset : offset + n_len].clone() n.value = full_device_indices[offset : offset + n_len].clone()
offset += n_len offset += n_len
self._record_store_event(n, medium=StorageMedium.GPU)
self.full_lru_list.insert_mru(n) self.full_lru_list.insert_mru(n)
self.full_evictable_size_ += n_len self.full_evictable_size_ += n_len
self._update_leaf_status(n) self._update_leaf_status(n)
@@ -379,6 +380,9 @@ class HiMambaRadixCache(MambaRadixCache):
finish_event.synchronize() finish_event.synchronize()
for ack_id in ack_list: for ack_id in ack_list:
backuped_node = self.ongoing_write_through.pop(ack_id) backuped_node = self.ongoing_write_through.pop(ack_id)
self._record_store_event(
backuped_node, medium=StorageMedium.CPU
)
if self.enable_storage: if self.enable_storage:
self.write_backup_storage(backuped_node) self.write_backup_storage(backuped_node)
self.cache_controller.ack_write_queue.clear() self.cache_controller.ack_write_queue.clear()
@@ -408,6 +412,7 @@ class HiMambaRadixCache(MambaRadixCache):
finish_event.synchronize() finish_event.synchronize()
for ack_id in ack_list: for ack_id in ack_list:
backuped_node = self.ongoing_write_through.pop(ack_id) backuped_node = self.ongoing_write_through.pop(ack_id)
self._record_store_event(backuped_node, medium=StorageMedium.CPU)
self.dec_lock_ref(backuped_node) self.dec_lock_ref(backuped_node)
if self.enable_storage: if self.enable_storage:
self.write_backup_storage(backuped_node) self.write_backup_storage(backuped_node)
@@ -514,6 +519,7 @@ class HiMambaRadixCache(MambaRadixCache):
num_full = len(node.value) num_full = len(node.value)
self._record_remove_event(node, medium=StorageMedium.GPU)
self.cache_controller.evict_device(node.value) self.cache_controller.evict_device(node.value)
self.full_evictable_size_ -= num_full self.full_evictable_size_ -= num_full
if self.full_lru_list.in_list(node): if self.full_lru_list.in_list(node):
@@ -534,6 +540,7 @@ class HiMambaRadixCache(MambaRadixCache):
full_num_evicted = len(node.value) full_num_evicted = len(node.value)
self._record_remove_event(node, medium=StorageMedium.GPU)
self.cache_controller.evict_device(node.value) self.cache_controller.evict_device(node.value)
self.full_evictable_size_ -= full_num_evicted self.full_evictable_size_ -= full_num_evicted
if self.full_lru_list.in_list(node): if self.full_lru_list.in_list(node):
@@ -576,6 +583,7 @@ class HiMambaRadixCache(MambaRadixCache):
node.host_mamba_ref_counter == 0 node.host_mamba_ref_counter == 0
), f"host mamba in use, {node.id=} {node.host_mamba_ref_counter=}" ), f"host mamba in use, {node.id=} {node.host_mamba_ref_counter=}"
self._record_remove_event(node, medium=StorageMedium.CPU)
full_num_evicted = self.cache_controller.evict_host(node.host_value) full_num_evicted = self.cache_controller.evict_host(node.host_value)
node.host_value = None node.host_value = None
@@ -612,6 +620,7 @@ class HiMambaRadixCache(MambaRadixCache):
and node.host_ref_counter == 0 and node.host_ref_counter == 0
and node.host_mamba_ref_counter == 0 and node.host_mamba_ref_counter == 0
): ):
self._record_remove_event(node, medium=StorageMedium.CPU)
self.cache_controller.evict_host(node.host_value) self.cache_controller.evict_host(node.host_value)
node.host_value = None node.host_value = None
@@ -641,6 +650,7 @@ class HiMambaRadixCache(MambaRadixCache):
parent = node.parent parent = node.parent
if not parent.evicted: if not parent.evicted:
self._record_remove_event(parent, medium=StorageMedium.GPU)
full_num_evicted += len(parent.value) full_num_evicted += len(parent.value)
self.full_evictable_size_ -= len(parent.value) self.full_evictable_size_ -= len(parent.value)
self.cache_controller.evict_device(parent.value) self.cache_controller.evict_device(parent.value)
@@ -808,6 +818,7 @@ class HiMambaRadixCache(MambaRadixCache):
node.value = fresh_value.clone() node.value = fresh_value.clone()
self.full_lru_list.insert_mru(node) self.full_lru_list.insert_mru(node)
self.full_evictable_size_ += n self.full_evictable_size_ += n
self._record_store_event(node, medium=StorageMedium.GPU)
self._update_leaf_status(node) self._update_leaf_status(node)
if node.parent is not None: if node.parent is not None:
@@ -903,8 +914,9 @@ class HiMambaRadixCache(MambaRadixCache):
parent.children[child_key] = new_node parent.children[child_key] = new_node
self.full_evictable_size_ += len(value) self.full_evictable_size_ += len(value)
self.mamba_evictable_size_ += len(mamba_value) self.mamba_evictable_size_ += len(mamba_value)
if self.enable_storage: if self.enable_storage or self.enable_kv_cache_events:
new_node.hash_value = compute_node_hash_values(new_node, self.page_size) new_node.hash_value = compute_node_hash_values(new_node, self.page_size)
self._record_store_event(new_node, medium=StorageMedium.GPU)
self._update_full_device_leaf_status(new_node) self._update_full_device_leaf_status(new_node)
self._update_full_device_leaf_status(parent) self._update_full_device_leaf_status(parent)
return new_node return new_node
@@ -1067,10 +1079,6 @@ class HiMambaRadixCache(MambaRadixCache):
new_node.host_value = child.host_value[:split_len].clone() new_node.host_value = child.host_value[:split_len].clone()
child.host_value = child.host_value[split_len:].clone() child.host_value = child.host_value[split_len:].clone()
new_node.hash_value, child.hash_value = split_node_hash_value(
child.hash_value, split_len, self.page_size
)
self._update_leaf_status(new_node) self._update_leaf_status(new_node)
self._update_leaf_status(child) self._update_leaf_status(child)
@@ -1879,6 +1887,7 @@ class HiMambaRadixCache(MambaRadixCache):
leaf_node = new_node leaf_node = new_node
self._update_full_host_leaf_status(new_node) self._update_full_host_leaf_status(new_node)
self._update_full_host_leaf_status(node) self._update_full_host_leaf_status(node)
self._record_store_event(new_node, medium=StorageMedium.CPU)
# Attach mamba state to the new leaf # Attach mamba state to the new leaf
if leaf_node is not None and mamba_host_value is not None and mamba_loaded: if leaf_node is not None and mamba_host_value is not None and mamba_loaded:
@@ -50,6 +50,8 @@ from sglang.srt.mem_cache.radix_cache import (
RadixCache, RadixCache,
RadixKey, RadixKey,
TreeNode, TreeNode,
)
from sglang.srt.mem_cache.utils import (
compute_node_hash_values, compute_node_hash_values,
split_node_hash_value, split_node_hash_value,
) )
@@ -45,8 +45,10 @@ from sglang.srt.mem_cache.base_prefix_cache import (
MatchPrefixParams, MatchPrefixParams,
MatchResult, MatchResult,
) )
from sglang.srt.mem_cache.events import KVCacheEventMixin
from sglang.srt.mem_cache.memory_pool import HybridReqToTokenPool from sglang.srt.mem_cache.memory_pool import HybridReqToTokenPool
from sglang.srt.mem_cache.radix_cache import RadixKey from sglang.srt.mem_cache.radix_cache import RadixKey
from sglang.srt.mem_cache.utils import split_node_hash_value
from sglang.srt.server_args import get_global_server_args from sglang.srt.server_args import get_global_server_args
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -415,7 +417,7 @@ class LRUList:
raise Exception(msg) raise Exception(msg)
class MambaRadixCache(BasePrefixCache): class MambaRadixCache(KVCacheEventMixin, BasePrefixCache):
def __init__(self, params: CacheInitParams): def __init__(self, params: CacheInitParams):
assert isinstance( assert isinstance(
params.token_to_kv_pool_allocator, TokenToKVPoolAllocator params.token_to_kv_pool_allocator, TokenToKVPoolAllocator
@@ -425,7 +427,9 @@ class MambaRadixCache(BasePrefixCache):
self.page_size = params.page_size self.page_size = params.page_size
self.disable = params.disable self.disable = params.disable
self.enable_kv_cache_events = params.enable_kv_cache_events
self.enable_mamba_extra_buffer = params.enable_mamba_extra_buffer self.enable_mamba_extra_buffer = params.enable_mamba_extra_buffer
self.kv_event_queue = []
if not self.enable_mamba_extra_buffer: if not self.enable_mamba_extra_buffer:
assert ( assert (
@@ -451,6 +455,7 @@ class MambaRadixCache(BasePrefixCache):
self.root_node = TreeNode() self.root_node = TreeNode()
self.root_node.key = RadixKey([], None) self.root_node.key = RadixKey([], None)
self.root_node.value = [] self.root_node.value = []
self.root_node.hash_value = []
self.root_node.full_lock_ref = 1 self.root_node.full_lock_ref = 1
self.root_node.mamba_lock_ref = 1 self.root_node.mamba_lock_ref = 1
self.full_evictable_size_ = 0 self.full_evictable_size_ = 0
@@ -460,6 +465,7 @@ class MambaRadixCache(BasePrefixCache):
# LRU lists are used to maintain the order of eviction of the nodes in the tree # LRU lists are used to maintain the order of eviction of the nodes in the tree
self.full_lru_list = LRUList(mamba=False) self.full_lru_list = LRUList(mamba=False)
self.mamba_lru_list = LRUList(mamba=True) self.mamba_lru_list = LRUList(mamba=True)
self._record_all_cleared_event()
def match_prefix(self, params: MatchPrefixParams) -> MatchResult: def match_prefix(self, params: MatchPrefixParams) -> MatchResult:
"""Find the matching prefix from the radix tree. """Find the matching prefix from the radix tree.
@@ -725,6 +731,7 @@ class MambaRadixCache(BasePrefixCache):
assert x.mamba_value is not None, f"leaf node mamba value is not None, {x.id=}" assert x.mamba_value is not None, f"leaf node mamba value is not None, {x.id=}"
# 1. a leaf node, free full tokens and mamba # 1. a leaf node, free full tokens and mamba
self._record_remove_event(x)
self.token_to_kv_pool_allocator.free(x.value) self.token_to_kv_pool_allocator.free(x.value)
full_num_evicted = len(x.value) full_num_evicted = len(x.value)
self.req_to_token_pool.mamba_pool.free(x.mamba_value) self.req_to_token_pool.mamba_pool.free(x.mamba_value)
@@ -1089,6 +1096,9 @@ class MambaRadixCache(BasePrefixCache):
child.key = child.key[split_len:] child.key = child.key[split_len:]
child.value = child.value[split_len:].clone() child.value = child.value[split_len:].clone()
new_node.parent.children[key.child_key(self.page_size)] = new_node new_node.parent.children[key.child_key(self.page_size)] = new_node
new_node.hash_value, child.hash_value = split_node_hash_value(
child.hash_value, split_len, self.page_size
)
# insert the new node and child into the lru lists, insert # insert the new node and child into the lru lists, insert
# parent first so that parent is after child in the lru list # parent first so that parent is after child in the lru list
@@ -1156,6 +1166,7 @@ class MambaRadixCache(BasePrefixCache):
node.children[child_key] = new_node node.children[child_key] = new_node
self.full_evictable_size_ += len(value) self.full_evictable_size_ += len(value)
self.mamba_evictable_size_ += len(mamba_value) self.mamba_evictable_size_ += len(mamba_value)
self._record_store_event(new_node)
elif node.mamba_value is None: # add for mamba tombstone elif node.mamba_value is None: # add for mamba tombstone
node.mamba_value = mamba_value node.mamba_value = mamba_value
self.full_lru_list.reset_node_mru(node) self.full_lru_list.reset_node_mru(node)
@@ -1185,6 +1196,7 @@ class MambaRadixCache(BasePrefixCache):
node.parent.mamba_lock_ref == 0 node.parent.mamba_lock_ref == 0
), f"tombstone mamba_lock_ref should always be 0, {node.parent.full_lock_ref=}, {node.parent.mamba_lock_ref=}, {node.parent.id=}" ), f"tombstone mamba_lock_ref should always be 0, {node.parent.full_lock_ref=}, {node.parent.mamba_lock_ref=}, {node.parent.id=}"
# delete tombstone node evicts full tokens # delete tombstone node evicts full tokens
self._record_remove_event(node.parent)
self.token_to_kv_pool_allocator.free(node.parent.value) self.token_to_kv_pool_allocator.free(node.parent.value)
full_num_evicted += len(node.parent.value) full_num_evicted += len(node.parent.value)
self.full_lru_list.remove_node(node.parent) self.full_lru_list.remove_node(node.parent)
+3 -149
View File
@@ -34,12 +34,6 @@ import torch
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
from sglang.srt.disaggregation.kv_events import (
AllBlocksCleared,
BlockRemoved,
BlockStored,
StorageMedium,
)
from sglang.srt.mem_cache.base_prefix_cache import ( from sglang.srt.mem_cache.base_prefix_cache import (
BasePrefixCache, BasePrefixCache,
DecLockRefParams, DecLockRefParams,
@@ -52,6 +46,7 @@ from sglang.srt.mem_cache.base_prefix_cache import (
MatchPrefixParams, MatchPrefixParams,
MatchResult, MatchResult,
) )
from sglang.srt.mem_cache.events import KVCacheEventMixin
from sglang.srt.mem_cache.evict_policy import ( from sglang.srt.mem_cache.evict_policy import (
EvictionStrategy, EvictionStrategy,
FIFOStrategy, FIFOStrategy,
@@ -62,7 +57,7 @@ from sglang.srt.mem_cache.evict_policy import (
PriorityStrategy, PriorityStrategy,
SLRUStrategy, SLRUStrategy,
) )
from sglang.srt.mem_cache.utils import hash_str_to_int64 from sglang.srt.mem_cache.utils import split_node_hash_value
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.managers.schedule_batch import Req from sglang.srt.managers.schedule_batch import Req
@@ -271,54 +266,7 @@ class TreeNode:
return self.last_access_time < other.last_access_time return self.last_access_time < other.last_access_time
def compute_node_hash_values(node: "TreeNode", page_size: int) -> List[str]: class RadixCache(KVCacheEventMixin, BasePrefixCache):
"""Compute SHA256-based hash values for position-aware identification."""
hash_values = []
parent_hash = None
if node.parent is not None and node.parent.hash_value is not None:
if len(node.parent.key) > 0 and len(node.parent.hash_value) > 0:
parent_hash = node.parent.hash_value[-1]
logical_len = len(node.key)
for start in range(0, logical_len, page_size):
end = min(start + page_size, logical_len)
if end <= start:
continue
hash_val = node.key.hash_page(start, end, parent_hash)
hash_values.append(hash_val)
parent_hash = hash_val
return hash_values
def split_node_hash_value(
child_hash_value: Optional[List[str]], split_len: int, page_size: int
) -> tuple[Optional[List[str]], Optional[List[str]]]:
"""Split hash_value between parent and child nodes during node splitting.
Args:
child_hash_value: The hash_value list from the child node being split
split_len: The length at which to split (in tokens)
page_size: The page size for calculating number of pages
Returns:
Tuple of (new_node_hash_value, updated_child_hash_value)
"""
if child_hash_value is None:
return None, None
if page_size == 1:
split_pages = split_len
else:
split_pages = split_len // page_size
new_node_hash = child_hash_value[:split_pages]
child_hash = child_hash_value[split_pages:]
return new_node_hash, child_hash
class RadixCache(BasePrefixCache):
def __init__(self, params: CacheInitParams): def __init__(self, params: CacheInitParams):
self.disable = params.disable self.disable = params.disable
self.req_to_token_pool = params.req_to_token_pool self.req_to_token_pool = params.req_to_token_pool
@@ -857,100 +805,6 @@ class RadixCache(BasePrefixCache):
stack.append(child) stack.append(child)
return total_size return total_size
def _record_store_event(self, node: TreeNode, medium=None):
# One BlockStored per ``page_size`` chunk.
# ``medium`` defaults to StorageMedium.GPU but callers may override
# for lower-tier insertions (e.g. StorageMedium.CPU for host/L2 cache).
if self.enable_kv_cache_events:
if medium is None:
medium = StorageMedium.GPU
# Compute hash_value lazily if not already set
if node.hash_value is None:
node.hash_value = compute_node_hash_values(node, self.page_size)
# Get parent's last hash value for first page
parent_block_hash = None
if node.parent is not None and node.parent != self.root_node:
if (
node.parent.hash_value is not None
and len(node.parent.hash_value) > 0
):
parent_block_hash = hash_str_to_int64(node.parent.hash_value[-1])
page_index = 0
logical_len = len(node.key)
is_bigram = node.key.is_bigram
raw = node.key.token_ids
for start in range(0, logical_len, self.page_size):
end = min(start + self.page_size, logical_len)
if end <= start:
continue
# Preserve historical event payload: bigram pages expose tuples.
if is_bigram:
page_tokens = [(raw[j], raw[j + 1]) for j in range(start, end)]
else:
page_tokens = raw[start:end]
block_hash = hash_str_to_int64(node.hash_value[page_index])
self.kv_event_queue.append(
BlockStored(
block_hashes=[block_hash],
parent_block_hash=parent_block_hash,
token_ids=page_tokens,
block_size=len(page_tokens),
lora_id=None,
medium=medium,
)
)
parent_block_hash = block_hash
page_index += 1
def _record_remove_event(self, node: TreeNode, medium=None):
# One BlockRemoved per chunk.
# ``medium`` defaults to StorageMedium.GPU but callers may override for
# lower-tier removals (e.g. StorageMedium.CPU when evicting from host).
if self.enable_kv_cache_events:
if medium is None:
medium = StorageMedium.GPU
# Compute hash_value lazily if not already set (must match what was stored)
if node.hash_value is None:
node.hash_value = compute_node_hash_values(node, self.page_size)
page_index = 0
logical_len = len(node.key)
for start in range(0, logical_len, self.page_size):
end = min(start + self.page_size, logical_len)
if end <= start:
continue
block_hash = hash_str_to_int64(node.hash_value[page_index])
self.kv_event_queue.append(
BlockRemoved(block_hashes=[block_hash], medium=medium)
)
page_index += 1
def _record_all_cleared_event(self):
if self.enable_kv_cache_events:
self.kv_event_queue.append(AllBlocksCleared())
def take_events(self):
"""Atomically takes all events and clears the queue.
Returns:
A list of KV cache events.
"""
if not self.enable_kv_cache_events:
return []
events = self.kv_event_queue
self.kv_event_queue = []
return events
if __name__ == "__main__": if __name__ == "__main__":
tree = RadixCache.create_simulated() tree = RadixCache.create_simulated()
+47
View File
@@ -389,3 +389,50 @@ def hash_str_to_int64(hash_str: str) -> int:
if uint64_val >= 2**63: if uint64_val >= 2**63:
return uint64_val - 2**64 return uint64_val - 2**64
return uint64_val return uint64_val
def compute_node_hash_values(node: Any, page_size: int) -> List[str]:
"""Compute SHA256-based hash values for position-aware KV block IDs."""
hash_values = []
parent_hash = None
if node.parent is not None and node.parent.hash_value is not None:
if len(node.parent.key) > 0 and len(node.parent.hash_value) > 0:
parent_hash = node.parent.hash_value[-1]
logical_len = len(node.key)
for start in range(0, logical_len, page_size):
end = min(start + page_size, logical_len)
if end <= start:
continue
hash_val = node.key.hash_page(start, end, parent_hash)
hash_values.append(hash_val)
parent_hash = hash_val
return hash_values
def split_node_hash_value(
child_hash_value: Optional[List[str]], split_len: int, page_size: int
) -> tuple[Optional[List[str]], Optional[List[str]]]:
"""Split hash_value between parent and child nodes during node splitting.
Args:
child_hash_value: The hash_value list from the child node being split
split_len: The length at which to split (in tokens)
page_size: The page size for calculating number of pages
Returns:
Tuple of (new_node_hash_value, updated_child_hash_value)
"""
if child_hash_value is None:
return None, None
if page_size == 1:
split_pages = split_len
else:
split_pages = split_len // page_size
new_node_hash = child_hash_value[:split_pages]
child_hash = child_hash_value[split_pages:]
return new_node_hash, child_hash
@@ -1,10 +1,14 @@
import shutil import shutil
import tempfile import tempfile
import time
import unittest import unittest
from types import SimpleNamespace from types import SimpleNamespace
import requests import requests
import zmq
from msgspec.msgpack import Decoder
from sglang.srt.disaggregation.kv_events import BlockStored, KVEventBatch
from sglang.srt.utils import kill_process_tree from sglang.srt.utils import kill_process_tree
from sglang.test.ci.ci_register import register_cuda_ci from sglang.test.ci.ci_register import register_cuda_ci
@@ -72,6 +76,8 @@ class TestQwen35WithHiCache(CustomTestCase):
"file", "file",
"--hicache-storage-prefetch-policy", "--hicache-storage-prefetch-policy",
"wait_complete", "wait_complete",
"--kv-events-config",
'{"publisher": "zmq", "topic": "kv-events"}',
], ],
) )
@@ -125,6 +131,42 @@ class TestQwen35WithHiCache(CustomTestCase):
f"first={first_metrics['score']}, second={second_metrics['score']}", f"first={first_metrics['score']}, second={second_metrics['score']}",
) )
def test_kv_events_smoke(self):
decoder = Decoder(type=KVEventBatch)
context = zmq.Context()
sub = context.socket(zmq.SUB)
sub.connect("tcp://localhost:5557")
sub.setsockopt_string(zmq.SUBSCRIBE, "kv-events")
try:
time.sleep(1.0)
res = requests.post(
f"{self.base_url}/generate",
json={
"text": "HiCache KV event compatibility check. " * 64,
"sampling_params": {"temperature": 0, "max_new_tokens": 1},
},
timeout=120,
)
res.raise_for_status()
events = []
deadline = time.time() + 10
while time.time() < deadline and not any(
isinstance(event, BlockStored) for event in events
):
if sub.poll(timeout=100):
_, _, payload = sub.recv_multipart()
events.extend(decoder.decode(payload).events)
self.assertTrue(
any(isinstance(event, BlockStored) for event in events),
"Expected at least one BlockStored event from Qwen3.5 HiCache server",
)
finally:
sub.close()
context.term()
if __name__ == "__main__": if __name__ == "__main__":
unittest.main() unittest.main()
@@ -3,6 +3,7 @@ import unittest
import torch import torch
from sglang.srt.configs.mamba_utils import Mamba2CacheParams, Mamba2StateShape from sglang.srt.configs.mamba_utils import Mamba2CacheParams, Mamba2StateShape
from sglang.srt.disaggregation.kv_events import BlockRemoved, BlockStored
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.managers.schedule_batch import Req from sglang.srt.managers.schedule_batch import Req
from sglang.srt.mem_cache.allocator import TokenToKVPoolAllocator from sglang.srt.mem_cache.allocator import TokenToKVPoolAllocator
@@ -305,7 +306,98 @@ class TestMamba(unittest.TestCase):
print(available_and_evictable_str(tree)) print(available_and_evictable_str(tree))
tree.sanity_check() tree.sanity_check()
def _setup_tree_and_allocator(self): def test_mamba_radix_cache_kv_events(self):
tree, allocator, _, make_dummy_req = self._setup_tree_and_allocator(
enable_kv_cache_events=True
)
tree.take_events() # Clear the reset event.
stored_hashes = []
req1 = make_dummy_req()
key1 = RadixKey([1, 2, 3])
tree.insert(
InsertParams(
key=key1,
value=allocator.alloc(3)[: len(key1)],
mamba_value=req1.mamba_pool_idx.unsqueeze(0),
)
)
events = tree.take_events()
stored_events = [e for e in events if isinstance(e, BlockStored)]
self.assertEqual(len(stored_events), 3)
self.assertEqual([e.token_ids[0] for e in stored_events], [1, 2, 3])
stored_hashes.extend(e.block_hashes[0] for e in stored_events)
req2 = make_dummy_req()
key2 = RadixKey([1, 2, 3, 4, 5])
tree.insert(
InsertParams(
key=key2,
value=allocator.alloc(5)[: len(key2)],
mamba_value=req2.mamba_pool_idx.unsqueeze(0),
)
)
events = tree.take_events()
stored_events = [e for e in events if isinstance(e, BlockStored)]
self.assertEqual(len(stored_events), 2)
self.assertEqual([e.token_ids[0] for e in stored_events], [4, 5])
stored_hashes.extend(e.block_hashes[0] for e in stored_events)
# Evicting an internal mamba state creates a tombstone but does not
# remove full-attention KV blocks, so it must not emit BlockRemoved.
result = tree.evict(EvictParams(num_tokens=0, mamba_num=1))
self.assertEqual(result.num_tokens_evicted, 0)
self.assertEqual(result.mamba_num_evicted, 1)
events = tree.take_events()
self.assertEqual([e for e in events if isinstance(e, BlockRemoved)], [])
result = tree.evict(EvictParams(num_tokens=1))
self.assertGreaterEqual(result.num_tokens_evicted, 1)
events = tree.take_events()
removed_hashes = [
e.block_hashes[0] for e in events if isinstance(e, BlockRemoved)
]
self.assertCountEqual(removed_hashes, stored_hashes)
def test_mamba_radix_cache_kv_events_split_hash(self):
tree, allocator, _, make_dummy_req = self._setup_tree_and_allocator(
enable_kv_cache_events=True
)
tree.take_events() # Clear the reset event.
req1 = make_dummy_req()
key1 = RadixKey([1, 2, 3, 4])
tree.insert(
InsertParams(
key=key1,
value=allocator.alloc(4)[: len(key1)],
mamba_value=req1.mamba_pool_idx.unsqueeze(0),
)
)
first_insert_events = [
e for e in tree.take_events() if isinstance(e, BlockStored)
]
self.assertEqual(len(first_insert_events), 4)
split_parent_hash = first_insert_events[1].block_hashes[0]
req2 = make_dummy_req()
key2 = RadixKey([1, 2, 5, 6])
tree.insert(
InsertParams(
key=key2,
value=allocator.alloc(4)[: len(key2)],
mamba_value=req2.mamba_pool_idx.unsqueeze(0),
)
)
second_insert_events = [
e for e in tree.take_events() if isinstance(e, BlockStored)
]
self.assertEqual(len(second_insert_events), 2)
self.assertEqual(second_insert_events[0].token_ids, [5])
self.assertEqual(second_insert_events[0].parent_block_hash, split_parent_hash)
def _setup_tree_and_allocator(self, enable_kv_cache_events=False):
"""Helper to create a MambaRadixCache with allocator for testing.""" """Helper to create a MambaRadixCache with allocator for testing."""
set_global_server_args_for_scheduler( set_global_server_args_for_scheduler(
ServerArgs(model_path="dummy", page_size=1) ServerArgs(model_path="dummy", page_size=1)
@@ -374,6 +466,7 @@ class TestMamba(unittest.TestCase):
token_to_kv_pool_allocator=allocator, token_to_kv_pool_allocator=allocator,
page_size=1, page_size=1,
disable=False, disable=False,
enable_kv_cache_events=enable_kv_cache_events,
) )
tree = MambaRadixCache(params=params) tree = MambaRadixCache(params=params)