diff --git a/.gitignore b/.gitignore index 8ac442295..6f94a4e54 100644 --- a/.gitignore +++ b/.gitignore @@ -275,3 +275,4 @@ sgl-kernel/csrc/**/*_musa/ *.glb *.ply *.npz +artifacts/ diff --git a/python/sglang/srt/disaggregation/kv_events.py b/python/sglang/srt/disaggregation/kv_events.py index 992798c79..0a5549bb1 100644 --- a/python/sglang/srt/disaggregation/kv_events.py +++ b/python/sglang/srt/disaggregation/kv_events.py @@ -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: diff --git a/python/sglang/srt/mem_cache/hiradix_cache.py b/python/sglang/srt/mem_cache/hiradix_cache.py index d342bbef6..e8e463c61 100644 --- a/python/sglang/srt/mem_cache/hiradix_cache.py +++ b/python/sglang/srt/mem_cache/hiradix_cache.py @@ -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 diff --git a/python/sglang/srt/mem_cache/radix_cache.py b/python/sglang/srt/mem_cache/radix_cache.py index f4d54bdc3..3e01bca91 100644 --- a/python/sglang/srt/mem_cache/radix_cache.py +++ b/python/sglang/srt/mem_cache/radix_cache.py @@ -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