feat(kv-events): publish SWA radix cache events (#24718)
Signed-off-by: PeaBrane <yanrpei@gmail.com>
This commit is contained in:
@@ -41,9 +41,10 @@ from sglang.srt.mem_cache.base_prefix_cache import (
|
|||||||
MatchResult,
|
MatchResult,
|
||||||
)
|
)
|
||||||
from sglang.srt.mem_cache.cache_init_params import CacheInitParams
|
from sglang.srt.mem_cache.cache_init_params import CacheInitParams
|
||||||
|
from sglang.srt.mem_cache.events import KVCacheEventMixin
|
||||||
from sglang.srt.mem_cache.radix_cache import RadixKey
|
from sglang.srt.mem_cache.radix_cache import RadixKey
|
||||||
from sglang.srt.mem_cache.swa_memory_pool import SWATokenToKVPoolAllocator
|
from sglang.srt.mem_cache.swa_memory_pool import SWATokenToKVPoolAllocator
|
||||||
from sglang.srt.mem_cache.utils import convert_to_bigram_key
|
from sglang.srt.mem_cache.utils import convert_to_bigram_key, 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
|
||||||
@@ -77,6 +78,8 @@ class TreeNode:
|
|||||||
self.hit_count = 0
|
self.hit_count = 0
|
||||||
# store the host indices of KV cache
|
# store the host indices of KV cache
|
||||||
self.host_value = None
|
self.host_value = None
|
||||||
|
# store hash values of each page
|
||||||
|
self.hash_value: Optional[List[str]] = None
|
||||||
|
|
||||||
# for lru list, invariant:
|
# for lru list, invariant:
|
||||||
# 1. prev has greater last_access_time
|
# 1. prev has greater last_access_time
|
||||||
@@ -334,7 +337,7 @@ class LRUList:
|
|||||||
raise Exception(msg)
|
raise Exception(msg)
|
||||||
|
|
||||||
|
|
||||||
class SWARadixCache(BasePrefixCache):
|
class SWARadixCache(KVCacheEventMixin, BasePrefixCache):
|
||||||
def __init__(self, params: CacheInitParams):
|
def __init__(self, params: CacheInitParams):
|
||||||
assert isinstance(params.token_to_kv_pool_allocator, SWATokenToKVPoolAllocator)
|
assert isinstance(params.token_to_kv_pool_allocator, SWATokenToKVPoolAllocator)
|
||||||
self.req_to_token_pool = params.req_to_token_pool
|
self.req_to_token_pool = params.req_to_token_pool
|
||||||
@@ -342,6 +345,8 @@ class SWARadixCache(BasePrefixCache):
|
|||||||
self.page_size = params.page_size
|
self.page_size = params.page_size
|
||||||
self.disable = params.disable
|
self.disable = params.disable
|
||||||
self.is_eagle = params.is_eagle
|
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:
|
if self.token_to_kv_pool_allocator:
|
||||||
self.device = self.token_to_kv_pool_allocator.device
|
self.device = self.token_to_kv_pool_allocator.device
|
||||||
@@ -371,6 +376,7 @@ class SWARadixCache(BasePrefixCache):
|
|||||||
self.root_node = TreeNode()
|
self.root_node = TreeNode()
|
||||||
self.root_node.key = []
|
self.root_node.key = []
|
||||||
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.swa_lock_ref = 1
|
self.root_node.swa_lock_ref = 1
|
||||||
self.full_evictable_size_ = 0
|
self.full_evictable_size_ = 0
|
||||||
@@ -380,6 +386,7 @@ class SWARadixCache(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(is_swa_list=False)
|
self.full_lru_list = LRUList(is_swa_list=False)
|
||||||
self.swa_lru_list = LRUList(is_swa_list=True)
|
self.swa_lru_list = LRUList(is_swa_list=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.
|
||||||
@@ -573,6 +580,7 @@ class SWARadixCache(BasePrefixCache):
|
|||||||
assert x.full_lock_ref == 0, f"node is in use, {x.id=}"
|
assert x.full_lock_ref == 0, f"node is in use, {x.id=}"
|
||||||
|
|
||||||
# 1. free node kv indices, evict full and swa tokens
|
# 1. free node kv indices, evict full and swa tokens
|
||||||
|
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)
|
||||||
# Tombstoned leaves had their SWA freed earlier in `dec_swa_lock_only`
|
# Tombstoned leaves had their SWA freed earlier in `dec_swa_lock_only`
|
||||||
@@ -637,6 +645,7 @@ class SWARadixCache(BasePrefixCache):
|
|||||||
x.full_lock_ref == 0
|
x.full_lock_ref == 0
|
||||||
), f"leaf node with full lock must also have swa lock, {x.id=}"
|
), f"leaf node with full lock must also have swa lock, {x.id=}"
|
||||||
# 1. a leaf node, free full and swa tokens
|
# 1. a leaf node, free full and swa tokens
|
||||||
|
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)
|
||||||
swa_num_evicted += len(x.value)
|
swa_num_evicted += len(x.value)
|
||||||
@@ -995,6 +1004,11 @@ class SWARadixCache(BasePrefixCache):
|
|||||||
if child.swa_uuid is not None:
|
if child.swa_uuid is not None:
|
||||||
node.swa_uuid = child.swa_uuid
|
node.swa_uuid = child.swa_uuid
|
||||||
|
|
||||||
|
if node.hash_value is not None and child.hash_value is not None:
|
||||||
|
node.hash_value = list(node.hash_value) + list(child.hash_value)
|
||||||
|
else:
|
||||||
|
node.hash_value = None
|
||||||
|
|
||||||
self.full_lru_list.remove_node(child)
|
self.full_lru_list.remove_node(child)
|
||||||
if not child.swa_tombstone:
|
if not child.swa_tombstone:
|
||||||
self.swa_lru_list.remove_node(child)
|
self.swa_lru_list.remove_node(child)
|
||||||
@@ -1062,6 +1076,9 @@ class SWARadixCache(BasePrefixCache):
|
|||||||
assert len(child.key) > 0, f"child.key should not be empty"
|
assert len(child.key) > 0, f"child.key should not be empty"
|
||||||
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
|
||||||
@@ -1219,6 +1236,7 @@ class SWARadixCache(BasePrefixCache):
|
|||||||
if not swa_tombstone:
|
if not swa_tombstone:
|
||||||
self.swa_lru_list.insert_mru(new_node)
|
self.swa_lru_list.insert_mru(new_node)
|
||||||
self.swa_evictable_size_ += len(value)
|
self.swa_evictable_size_ += len(value)
|
||||||
|
self._record_store_event(new_node)
|
||||||
return new_node
|
return new_node
|
||||||
|
|
||||||
def _iteratively_delete_tombstone_leaf(
|
def _iteratively_delete_tombstone_leaf(
|
||||||
@@ -1236,6 +1254,7 @@ class SWARadixCache(BasePrefixCache):
|
|||||||
node.parent.swa_lock_ref == 0
|
node.parent.swa_lock_ref == 0
|
||||||
), f"tombstone swa_lock_ref should always be 0, {node.parent.full_lock_ref=}, {node.parent.swa_lock_ref=}, {node.parent.id=}"
|
), f"tombstone swa_lock_ref should always be 0, {node.parent.full_lock_ref=}, {node.parent.swa_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)
|
||||||
|
|||||||
@@ -2,6 +2,7 @@ import unittest
|
|||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.disaggregation.kv_events import BlockRemoved, BlockStored
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.mem_cache.base_prefix_cache import (
|
from sglang.srt.mem_cache.base_prefix_cache import (
|
||||||
DecLockRefParams,
|
DecLockRefParams,
|
||||||
@@ -41,6 +42,7 @@ def _build_swa_tree(
|
|||||||
kv_size: int = 64,
|
kv_size: int = 64,
|
||||||
kv_size_swa: int = 32,
|
kv_size_swa: int = 32,
|
||||||
sliding_window_size: int = 4,
|
sliding_window_size: int = 4,
|
||||||
|
enable_kv_cache_events: bool = False,
|
||||||
):
|
):
|
||||||
head_num = 8
|
head_num = 8
|
||||||
head_dim = 128
|
head_dim = 128
|
||||||
@@ -89,6 +91,7 @@ def _build_swa_tree(
|
|||||||
disable=False,
|
disable=False,
|
||||||
is_eagle=is_eagle,
|
is_eagle=is_eagle,
|
||||||
sliding_window_size=sliding_window_size,
|
sliding_window_size=sliding_window_size,
|
||||||
|
enable_kv_cache_events=enable_kv_cache_events,
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
return tree, allocator, req_to_token_pool
|
return tree, allocator, req_to_token_pool
|
||||||
@@ -133,6 +136,66 @@ class TestSWA(unittest.TestCase):
|
|||||||
def tearDownClass(cls):
|
def tearDownClass(cls):
|
||||||
pass
|
pass
|
||||||
|
|
||||||
|
def test_swa_radix_cache_kv_events(self):
|
||||||
|
tree, allocator, _ = _build_swa_tree(
|
||||||
|
is_eagle=False, enable_kv_cache_events=True
|
||||||
|
)
|
||||||
|
tree.take_events() # Clear the reset event.
|
||||||
|
|
||||||
|
_insert(tree, allocator, [1, 2, 3, 4])
|
||||||
|
first_insert_events = [
|
||||||
|
e for e in tree.take_events() if isinstance(e, BlockStored)
|
||||||
|
]
|
||||||
|
self.assertEqual(len(first_insert_events), 4)
|
||||||
|
self.assertEqual([e.token_ids[0] for e in first_insert_events], [1, 2, 3, 4])
|
||||||
|
|
||||||
|
_insert(tree, allocator, [1, 2, 3, 4, 5, 6])
|
||||||
|
second_insert_events = [
|
||||||
|
e for e in tree.take_events() if isinstance(e, BlockStored)
|
||||||
|
]
|
||||||
|
self.assertEqual(len(second_insert_events), 2)
|
||||||
|
self.assertEqual([e.token_ids[0] for e in second_insert_events], [5, 6])
|
||||||
|
|
||||||
|
stored_hashes = [
|
||||||
|
e.block_hashes[0] for e in first_insert_events + second_insert_events
|
||||||
|
]
|
||||||
|
|
||||||
|
# Evicting only SWA tokens tombstones nodes but keeps full KV blocks.
|
||||||
|
result = tree.evict(EvictParams(num_tokens=0, swa_num_tokens=1))
|
||||||
|
self.assertEqual(result.num_tokens_evicted, 0)
|
||||||
|
self.assertGreaterEqual(result.swa_num_tokens_evicted, 1)
|
||||||
|
self.assertEqual(
|
||||||
|
[e for e in tree.take_events() if isinstance(e, BlockRemoved)], []
|
||||||
|
)
|
||||||
|
|
||||||
|
result = tree.evict(EvictParams(num_tokens=1, swa_num_tokens=0))
|
||||||
|
self.assertGreaterEqual(result.num_tokens_evicted, 1)
|
||||||
|
removed_hashes = [
|
||||||
|
e.block_hashes[0] for e in tree.take_events() if isinstance(e, BlockRemoved)
|
||||||
|
]
|
||||||
|
self.assertCountEqual(removed_hashes, stored_hashes)
|
||||||
|
|
||||||
|
def test_swa_radix_cache_kv_events_split_hash(self):
|
||||||
|
tree, allocator, _ = _build_swa_tree(
|
||||||
|
is_eagle=False, enable_kv_cache_events=True
|
||||||
|
)
|
||||||
|
tree.take_events() # Clear the reset event.
|
||||||
|
|
||||||
|
_insert(tree, allocator, [1, 2, 3, 4])
|
||||||
|
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]
|
||||||
|
|
||||||
|
_insert(tree, allocator, [1, 2, 5, 6])
|
||||||
|
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 test_swa_memory_pool(self):
|
def test_swa_memory_pool(self):
|
||||||
size = 16
|
size = 16
|
||||||
size_swa = 16
|
size_swa = 16
|
||||||
|
|||||||
Reference in New Issue
Block a user