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:
ishandhanani
2026-04-20 19:07:03 -07:00
committed by GitHub
co-authored by jthomson04
parent ac08ebed65
commit 3c007ee5d4
4 changed files with 45 additions and 11 deletions
+1
View File
@@ -275,3 +275,4 @@ sgl-kernel/csrc/**/*_musa/
*.glb
*.ply
*.npz
artifacts/
@@ -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:
+21 -3
View File
@@ -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
+15 -5
View File
@@ -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