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
@@ -275,3 +275,4 @@ sgl-kernel/csrc/**/*_musa/
|
|||||||
*.glb
|
*.glb
|
||||||
*.ply
|
*.ply
|
||||||
*.npz
|
*.npz
|
||||||
|
artifacts/
|
||||||
|
|||||||
@@ -18,6 +18,7 @@ KV caching events
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
import atexit
|
import atexit
|
||||||
|
import enum
|
||||||
import logging
|
import logging
|
||||||
import queue
|
import queue
|
||||||
import threading
|
import threading
|
||||||
@@ -56,9 +57,13 @@ class KVCacheEvent(
|
|||||||
"""Base class for all KV cache-related events"""
|
"""Base class for all KV cache-related events"""
|
||||||
|
|
||||||
|
|
||||||
# Medium values for hicache storage tiers
|
class StorageMedium(str, enum.Enum):
|
||||||
MEDIUM_GPU = "GPU"
|
"""Storage tier for KV cache events."""
|
||||||
MEDIUM_CPU = "CPU_PINNED"
|
|
||||||
|
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:
|
class OffloadedState:
|
||||||
|
|||||||
@@ -12,6 +12,7 @@ from typing import TYPE_CHECKING, Dict, List, Optional
|
|||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.disaggregation.kv_events import StorageMedium
|
||||||
from sglang.srt.managers.cache_controller import HiCacheController, PrefetchOperation
|
from sglang.srt.managers.cache_controller import HiCacheController, PrefetchOperation
|
||||||
from sglang.srt.mem_cache.base_prefix_cache import (
|
from sglang.srt.mem_cache.base_prefix_cache import (
|
||||||
DecLockRefParams,
|
DecLockRefParams,
|
||||||
@@ -677,6 +678,8 @@ class HiRadixCache(RadixCache):
|
|||||||
if not write_back:
|
if not write_back:
|
||||||
# no need to lock nodes if write back
|
# no need to lock nodes if write back
|
||||||
self.inc_lock_ref(node)
|
self.inc_lock_ref(node)
|
||||||
|
# Note: store(CPU) event is deferred to writing_check() after the
|
||||||
|
# async DMA transfer is confirmed complete.
|
||||||
else:
|
else:
|
||||||
return 0
|
return 0
|
||||||
|
|
||||||
@@ -718,6 +721,10 @@ class HiRadixCache(RadixCache):
|
|||||||
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)
|
||||||
|
# DMA confirmed -- block is now on host.
|
||||||
|
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()
|
||||||
@@ -748,6 +755,8 @@ class HiRadixCache(RadixCache):
|
|||||||
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)
|
||||||
|
# DMA confirmed -- block is now on host.
|
||||||
|
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)
|
||||||
@@ -881,7 +890,10 @@ class HiRadixCache(RadixCache):
|
|||||||
return EvictResult(num_tokens_evicted=num_evicted)
|
return EvictResult(num_tokens_evicted=num_evicted)
|
||||||
|
|
||||||
def _evict_backuped(self, node: TreeNode):
|
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)
|
num_evicted = self.cache_controller.evict_device(node.value)
|
||||||
assert num_evicted > 0
|
assert num_evicted > 0
|
||||||
self.evictable_size_ -= num_evicted
|
self.evictable_size_ -= num_evicted
|
||||||
@@ -922,8 +934,8 @@ class HiRadixCache(RadixCache):
|
|||||||
continue
|
continue
|
||||||
|
|
||||||
# Block deleted entirely (GPU already evicted, now CPU freed) --
|
# Block deleted entirely (GPU already evicted, now CPU freed) --
|
||||||
# emit BlockRemoved so the router removes this block from its index.
|
# emit remove(CPU) so the router drops the host-tier entry.
|
||||||
self._record_remove_event(x)
|
self._record_remove_event(x, medium=StorageMedium.CPU)
|
||||||
num_evicted += self.cache_controller.evict_host(x.host_value)
|
num_evicted += self.cache_controller.evict_host(x.host_value)
|
||||||
|
|
||||||
key = self.get_child_key_fn(x.key)
|
key = self.get_child_key_fn(x.key)
|
||||||
@@ -995,6 +1007,9 @@ class HiRadixCache(RadixCache):
|
|||||||
for node in nodes_to_load:
|
for node in nodes_to_load:
|
||||||
node.value = device_indices[offset : offset + len(node.host_value)].clone()
|
node.value = device_indices[offset : offset + len(node.host_value)].clone()
|
||||||
offset += len(node.host_value)
|
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.evictable_size_ += len(device_indices)
|
||||||
self.inc_lock_ref(last_hit_node)
|
self.inc_lock_ref(last_hit_node)
|
||||||
|
|
||||||
@@ -1327,6 +1342,9 @@ class HiRadixCache(RadixCache):
|
|||||||
self._update_host_leaf_status(new_node)
|
self._update_host_leaf_status(new_node)
|
||||||
self._update_leaf_status(node)
|
self._update_leaf_status(node)
|
||||||
self._update_host_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
|
return matched_length
|
||||||
|
|
||||||
|
|||||||
@@ -34,10 +34,10 @@ import torch
|
|||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
from sglang.srt.disaggregation.kv_events import (
|
from sglang.srt.disaggregation.kv_events import (
|
||||||
MEDIUM_GPU,
|
|
||||||
AllBlocksCleared,
|
AllBlocksCleared,
|
||||||
BlockRemoved,
|
BlockRemoved,
|
||||||
BlockStored,
|
BlockStored,
|
||||||
|
StorageMedium,
|
||||||
)
|
)
|
||||||
from sglang.srt.mem_cache.base_prefix_cache import (
|
from sglang.srt.mem_cache.base_prefix_cache import (
|
||||||
BasePrefixCache,
|
BasePrefixCache,
|
||||||
@@ -898,9 +898,14 @@ class RadixCache(BasePrefixCache):
|
|||||||
stack.append(child)
|
stack.append(child)
|
||||||
return total_size
|
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.
|
# 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 self.enable_kv_cache_events:
|
||||||
|
if medium is None:
|
||||||
|
medium = StorageMedium.GPU
|
||||||
|
|
||||||
# Compute hash_value lazily if not already set
|
# Compute hash_value lazily if not already set
|
||||||
if node.hash_value is None:
|
if node.hash_value is None:
|
||||||
node.hash_value = compute_node_hash_values(node, self.page_size)
|
node.hash_value = compute_node_hash_values(node, self.page_size)
|
||||||
@@ -937,16 +942,21 @@ class RadixCache(BasePrefixCache):
|
|||||||
token_ids=page_tokens,
|
token_ids=page_tokens,
|
||||||
block_size=len(page_tokens),
|
block_size=len(page_tokens),
|
||||||
lora_id=None,
|
lora_id=None,
|
||||||
medium=MEDIUM_GPU,
|
medium=medium,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
parent_block_hash = block_hash
|
parent_block_hash = block_hash
|
||||||
page_index += 1
|
page_index += 1
|
||||||
|
|
||||||
def _record_remove_event(self, node: TreeNode):
|
def _record_remove_event(self, node: TreeNode, medium=None):
|
||||||
# One BlockRemoved per chunk.
|
# 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 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)
|
# Compute hash_value lazily if not already set (must match what was stored)
|
||||||
if node.hash_value is None:
|
if node.hash_value is None:
|
||||||
node.hash_value = compute_node_hash_values(node, self.page_size)
|
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])
|
block_hash = hash_str_to_int64(node.hash_value[page_index])
|
||||||
|
|
||||||
self.kv_event_queue.append(
|
self.kv_event_queue.append(
|
||||||
BlockRemoved(block_hashes=[block_hash], medium=MEDIUM_GPU)
|
BlockRemoved(block_hashes=[block_hash], medium=medium)
|
||||||
)
|
)
|
||||||
|
|
||||||
page_index += 1
|
page_index += 1
|
||||||
|
|||||||
Reference in New Issue
Block a user