fix: malformed KV events for NVIDIA Dynamo (#13488)
Signed-off-by: PeaBrane <yanrpei@gmail.com> Co-authored-by: ishandhanani <82981111+ishandhanani@users.noreply.github.com> Co-authored-by: Zhiqiang Xie <xiezhq@stanford.edu>
This commit is contained in:
@@ -23,6 +23,7 @@ The radix tree data structure for managing the KV cache.
|
||||
"""
|
||||
|
||||
import heapq
|
||||
import logging
|
||||
import sys
|
||||
import time
|
||||
from collections import defaultdict
|
||||
@@ -31,6 +32,8 @@ from typing import TYPE_CHECKING, Any, Iterator, List, Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
from sglang.srt.disaggregation.kv_events import (
|
||||
AllBlocksCleared,
|
||||
BlockRemoved,
|
||||
@@ -46,6 +49,7 @@ from sglang.srt.mem_cache.evict_policy import (
|
||||
MRUStrategy,
|
||||
PriorityStrategy,
|
||||
)
|
||||
from sglang.srt.mem_cache.hicache_storage import get_hash_str, hash_str_to_int64
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.managers.schedule_batch import Req
|
||||
@@ -185,6 +189,66 @@ def get_child_key(key: RadixKey, page_size: int = 1):
|
||||
return (key.extra_key, plain_key)
|
||||
|
||||
|
||||
def compute_node_hash_values(node: "TreeNode", page_size: int) -> List[str]:
|
||||
"""Compute SHA256-based hash values for position-aware identification.
|
||||
|
||||
Args:
|
||||
node: The TreeNode to compute hash values for
|
||||
page_size: The page size for chunking tokens
|
||||
|
||||
Returns:
|
||||
List of SHA256 hex strings, one per page
|
||||
"""
|
||||
hash_values = []
|
||||
|
||||
# Get parent's last hash value if parent exists
|
||||
parent_hash = None
|
||||
if node.parent is not None and node.parent.hash_value is not None:
|
||||
# Check if parent is root by checking if it has empty key
|
||||
if len(node.parent.key) > 0 and len(node.parent.hash_value) > 0:
|
||||
parent_hash = node.parent.hash_value[-1]
|
||||
|
||||
# Iterate through node's pages
|
||||
for start in range(0, len(node.key), page_size):
|
||||
page_tokens = node.key.token_ids[start : start + page_size]
|
||||
if not page_tokens:
|
||||
continue
|
||||
|
||||
# Use SHA256-based chaining via get_hash_str
|
||||
hash_val = get_hash_str(page_tokens, prior_hash=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):
|
||||
self.disable = params.disable
|
||||
@@ -258,6 +322,7 @@ class RadixCache(BasePrefixCache):
|
||||
self.root_node.value = []
|
||||
self.root_node.host_value = []
|
||||
self.root_node.lock_ref = 1
|
||||
self.root_node.hash_value = []
|
||||
self.evictable_size_ = 0
|
||||
self.protected_size_ = 0
|
||||
self._record_all_cleared_event()
|
||||
@@ -596,8 +661,10 @@ class RadixCache(BasePrefixCache):
|
||||
child.value = child.value[split_len:]
|
||||
new_node.parent.children[self.get_child_key_fn(key)] = new_node
|
||||
|
||||
self._record_store_event(new_node)
|
||||
self._record_store_event(child)
|
||||
# Split hash_value if it was already computed, otherwise leave as None
|
||||
new_node.hash_value, child.hash_value = split_node_hash_value(
|
||||
child.hash_value, split_len, self.page_size
|
||||
)
|
||||
|
||||
return new_node
|
||||
|
||||
@@ -640,6 +707,7 @@ class RadixCache(BasePrefixCache):
|
||||
new_node.value = value
|
||||
node.children[child_key] = new_node
|
||||
self.evictable_size_ += len(key)
|
||||
# Hash will be computed lazily during event emission
|
||||
self._record_store_event(new_node)
|
||||
return total_prefix_length
|
||||
|
||||
@@ -697,22 +765,26 @@ class RadixCache(BasePrefixCache):
|
||||
def _record_store_event(self, node: TreeNode):
|
||||
# One BlockStored per ``page_size`` chunk.
|
||||
if self.enable_kv_cache_events:
|
||||
# First chunk links to the last page of the parent node (if any).
|
||||
if node.parent is None or node != self.root_node:
|
||||
parent_block_hash = None
|
||||
else:
|
||||
last_page_start = (
|
||||
(len(node.parent.key) - 1) // self.page_size
|
||||
) * self.page_size
|
||||
parent_parent_tokens = node.parent.key.token_ids[last_page_start:]
|
||||
parent_block_hash = hash(tuple(parent_parent_tokens))
|
||||
# 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
|
||||
for start in range(0, len(node.key), self.page_size):
|
||||
page_tokens = node.key.token_ids[start : start + self.page_size]
|
||||
if not page_tokens:
|
||||
continue
|
||||
|
||||
block_hash = hash(tuple(page_tokens))
|
||||
block_hash = hash_str_to_int64(node.hash_value[page_index])
|
||||
|
||||
self.kv_event_queue.append(
|
||||
BlockStored(
|
||||
@@ -724,19 +796,28 @@ class RadixCache(BasePrefixCache):
|
||||
)
|
||||
)
|
||||
|
||||
# Chain next chunk to this one.
|
||||
parent_block_hash = block_hash
|
||||
page_index += 1
|
||||
|
||||
def _record_remove_event(self, node: TreeNode):
|
||||
# One BlockRemoved per chunk.
|
||||
if self.enable_kv_cache_events:
|
||||
# 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
|
||||
for start in range(0, len(node.key), self.page_size):
|
||||
page_tokens = node.key.token_ids[start : start + self.page_size]
|
||||
if not page_tokens:
|
||||
continue
|
||||
block_hash = hash(tuple(page_tokens))
|
||||
|
||||
block_hash = hash_str_to_int64(node.hash_value[page_index])
|
||||
|
||||
self.kv_event_queue.append(BlockRemoved(block_hashes=[block_hash]))
|
||||
|
||||
page_index += 1
|
||||
|
||||
def _record_all_cleared_event(self):
|
||||
if self.enable_kv_cache_events:
|
||||
self.kv_event_queue.append(AllBlocksCleared())
|
||||
|
||||
Reference in New Issue
Block a user