fix(hicache): emit KV events for L2 host cache insertions (#22894)
Signed-off-by: jthomson04 <jwillthomson19@gmail.com> Signed-off-by: Ishan Dhanani <ishandhanani@gmail.com> Co-authored-by: jthomson04 <jwillthomson19@gmail.com>
This commit is contained in:
co-authored by
jthomson04
parent
ac08ebed65
commit
3c007ee5d4
@@ -18,6 +18,7 @@ KV caching events
|
||||
"""
|
||||
|
||||
import atexit
|
||||
import enum
|
||||
import logging
|
||||
import queue
|
||||
import threading
|
||||
@@ -56,9 +57,13 @@ class KVCacheEvent(
|
||||
"""Base class for all KV cache-related events"""
|
||||
|
||||
|
||||
# Medium values for hicache storage tiers
|
||||
MEDIUM_GPU = "GPU"
|
||||
MEDIUM_CPU = "CPU_PINNED"
|
||||
class StorageMedium(str, enum.Enum):
|
||||
"""Storage tier for KV cache events."""
|
||||
|
||||
GPU = "GPU" # L1: device HBM
|
||||
CPU = "CPU_PINNED" # L2: host pinned memory
|
||||
DISK = "DISK" # L3: SSD / NVMe
|
||||
EXTERNAL = "EXTERNAL" # L4: shared / remote pool (e.g. Mooncake)
|
||||
|
||||
|
||||
class OffloadedState:
|
||||
|
||||
@@ -12,6 +12,7 @@ from typing import TYPE_CHECKING, Dict, List, Optional
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.disaggregation.kv_events import StorageMedium
|
||||
from sglang.srt.managers.cache_controller import HiCacheController, PrefetchOperation
|
||||
from sglang.srt.mem_cache.base_prefix_cache import (
|
||||
DecLockRefParams,
|
||||
@@ -677,6 +678,8 @@ class HiRadixCache(RadixCache):
|
||||
if not write_back:
|
||||
# no need to lock nodes if write back
|
||||
self.inc_lock_ref(node)
|
||||
# Note: store(CPU) event is deferred to writing_check() after the
|
||||
# async DMA transfer is confirmed complete.
|
||||
else:
|
||||
return 0
|
||||
|
||||
@@ -718,6 +721,10 @@ class HiRadixCache(RadixCache):
|
||||
finish_event.synchronize()
|
||||
for ack_id in ack_list:
|
||||
backuped_node = self.ongoing_write_through.pop(ack_id)
|
||||
# DMA confirmed -- block is now on host.
|
||||
self._record_store_event(
|
||||
backuped_node, medium=StorageMedium.CPU
|
||||
)
|
||||
if self.enable_storage:
|
||||
self.write_backup_storage(backuped_node)
|
||||
self.cache_controller.ack_write_queue.clear()
|
||||
@@ -748,6 +755,8 @@ class HiRadixCache(RadixCache):
|
||||
finish_event.synchronize()
|
||||
for ack_id in ack_list:
|
||||
backuped_node = self.ongoing_write_through.pop(ack_id)
|
||||
# DMA confirmed -- block is now on host.
|
||||
self._record_store_event(backuped_node, medium=StorageMedium.CPU)
|
||||
self.dec_lock_ref(backuped_node)
|
||||
if self.enable_storage:
|
||||
self.write_backup_storage(backuped_node)
|
||||
@@ -881,7 +890,10 @@ class HiRadixCache(RadixCache):
|
||||
return EvictResult(num_tokens_evicted=num_evicted)
|
||||
|
||||
def _evict_backuped(self, node: TreeNode):
|
||||
# GPU -> CPU demotion: no BlockRemoved since block is still reachable via load_back
|
||||
# GPU -> CPU demotion: block moves from device to host.
|
||||
# Emit remove(GPU) so downstream indexers stop scoring it as device-local.
|
||||
# The matching store(CPU) was emitted when write_backup() copied to host.
|
||||
self._record_remove_event(node, medium=StorageMedium.GPU)
|
||||
num_evicted = self.cache_controller.evict_device(node.value)
|
||||
assert num_evicted > 0
|
||||
self.evictable_size_ -= num_evicted
|
||||
@@ -922,8 +934,8 @@ class HiRadixCache(RadixCache):
|
||||
continue
|
||||
|
||||
# Block deleted entirely (GPU already evicted, now CPU freed) --
|
||||
# emit BlockRemoved so the router removes this block from its index.
|
||||
self._record_remove_event(x)
|
||||
# emit remove(CPU) so the router drops the host-tier entry.
|
||||
self._record_remove_event(x, medium=StorageMedium.CPU)
|
||||
num_evicted += self.cache_controller.evict_host(x.host_value)
|
||||
|
||||
key = self.get_child_key_fn(x.key)
|
||||
@@ -995,6 +1007,9 @@ class HiRadixCache(RadixCache):
|
||||
for node in nodes_to_load:
|
||||
node.value = device_indices[offset : offset + len(node.host_value)].clone()
|
||||
offset += len(node.host_value)
|
||||
# Block promoted from host to GPU -- emit store(GPU) so downstream
|
||||
# indexers see it as device-local again.
|
||||
self._record_store_event(node, medium=StorageMedium.GPU)
|
||||
self.evictable_size_ += len(device_indices)
|
||||
self.inc_lock_ref(last_hit_node)
|
||||
|
||||
@@ -1327,6 +1342,9 @@ class HiRadixCache(RadixCache):
|
||||
self._update_host_leaf_status(new_node)
|
||||
self._update_leaf_status(node)
|
||||
self._update_host_leaf_status(node)
|
||||
# Publish the newly materialized host suffix immediately so downstream
|
||||
# cache indexers can resolve descendants that extend this L2-only prefix.
|
||||
self._record_store_event(new_node, medium=StorageMedium.CPU)
|
||||
|
||||
return matched_length
|
||||
|
||||
|
||||
@@ -34,10 +34,10 @@ import torch
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
from sglang.srt.disaggregation.kv_events import (
|
||||
MEDIUM_GPU,
|
||||
AllBlocksCleared,
|
||||
BlockRemoved,
|
||||
BlockStored,
|
||||
StorageMedium,
|
||||
)
|
||||
from sglang.srt.mem_cache.base_prefix_cache import (
|
||||
BasePrefixCache,
|
||||
@@ -898,9 +898,14 @@ class RadixCache(BasePrefixCache):
|
||||
stack.append(child)
|
||||
return total_size
|
||||
|
||||
def _record_store_event(self, node: TreeNode):
|
||||
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)
|
||||
@@ -937,16 +942,21 @@ class RadixCache(BasePrefixCache):
|
||||
token_ids=page_tokens,
|
||||
block_size=len(page_tokens),
|
||||
lora_id=None,
|
||||
medium=MEDIUM_GPU,
|
||||
medium=medium,
|
||||
)
|
||||
)
|
||||
|
||||
parent_block_hash = block_hash
|
||||
page_index += 1
|
||||
|
||||
def _record_remove_event(self, node: TreeNode):
|
||||
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)
|
||||
@@ -961,7 +971,7 @@ class RadixCache(BasePrefixCache):
|
||||
block_hash = hash_str_to_int64(node.hash_value[page_index])
|
||||
|
||||
self.kv_event_queue.append(
|
||||
BlockRemoved(block_hashes=[block_hash], medium=MEDIUM_GPU)
|
||||
BlockRemoved(block_hashes=[block_hash], medium=medium)
|
||||
)
|
||||
|
||||
page_index += 1
|
||||
|
||||
Reference in New Issue
Block a user