[Unified Cache][1/N]: Support cache contract for external linker (#37091)

Co-authored-by: 晟海 <huangtingwei.htw@antgroup.com>
This commit is contained in:
Zhangheng
2026-08-30 22:24:13 +08:00
committed by GitHub
co-authored by 晟海
parent fe694986a2
commit c9eb475a88
9 changed files with 319 additions and 12 deletions
@@ -19,6 +19,7 @@ import torch
from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator
from sglang.srt.mem_cache.events import KVCacheEventRecorder
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
from sglang.srt.mem_cache.unified_cache.component_type import ComponentType
from sglang.srt.observability.metrics_collector import (
STAT_LOGGER_ROLE_RADIX_CACHE,
RadixCacheMetricsCollector,
@@ -34,9 +35,6 @@ if TYPE_CHECKING:
CacheAction,
ComponentAction,
)
from sglang.srt.mem_cache.unified_cache.components.tree_component import (
ComponentType,
)
@runtime_checkable
@@ -78,6 +76,7 @@ class InsertParams:
# General
chunked: bool = False
priority: int = 0
track_adopted_ranges: bool = False
@dataclasses.dataclass
@@ -90,11 +89,24 @@ class InsertResult:
mamba_exist: bool = False
inserted_host_node: Any = None
host_insert_dropped: bool = False
adopted_ranges: Optional[dict[ComponentType, list[tuple[int, int]]]] = None
# Controller-applied actions from the non-stepped channels (e.g. insert_host); the stepped insert emits via InsertStepResult.actions.
cache_actions: list[CacheAction | ComponentAction] = dataclasses.field(
default_factory=list
)
def record_adopted_range(
self, component_type: ComponentType, start: int, end: int
) -> None:
if self.adopted_ranges is None or start >= end:
return
ranges = self.adopted_ranges.setdefault(component_type, [])
if ranges and start <= ranges[-1][1]:
prev_start, prev_end = ranges[-1]
ranges[-1] = (min(prev_start, start), max(prev_end, end))
else:
ranges.append((start, end))
@dataclasses.dataclass
class EvictParams:
@@ -125,6 +125,13 @@ class PoolTransferResult:
kv_hit_pages: int
extra_pool_hit_pages: dict[str, int]
# Pools with TRAILING_PAGES (SWA, Mamba state) only hold a window that ends on an
# offloaded node boundary.
# Each rank owns its own shard and may hold a different set, so reducing a
# per-rank maximum would pick a length that is illegal on another rank; the
# caller intersects these sets instead.
restorable_prefix_pages: Optional[List[int]] = None
@classmethod
def empty(cls) -> PoolTransferResult:
return cls(0, {})
@@ -8,6 +8,8 @@ from sglang.srt.mem_cache.unified_cache.components.tree_component import (
ComponentData,
ComponentType,
EvictLayer,
ExternalLinkerLoadPhase,
LinkerTransferPhase,
LRURefreshPhase,
PrepareLoadBackResult,
PreparePrefetchResult,
@@ -20,6 +22,8 @@ __all__ = [
"BASE_COMPONENT_TYPE",
"ComponentData",
"ComponentType",
"ExternalLinkerLoadPhase",
"LinkerTransferPhase",
"EvictLayer",
"FullComponent",
"CacheTransferPhase",
@@ -7,6 +7,7 @@ import torch
from sglang.srt.mem_cache.base_prefix_cache import (
DecLockRefParams,
EvictParams,
IncLockRefResult,
InsertResult,
MatchPrefixParams,
@@ -22,10 +23,13 @@ from sglang.srt.mem_cache.unified_cache.components.tree_component import (
CacheTransferPhase,
ComponentType,
EvictLayer,
ExternalLinkerLoadPhase,
LinkerTransferPhase,
TreeComponent,
)
if TYPE_CHECKING:
from sglang.srt.managers.schedule_batch import Req
from sglang.srt.mem_cache.unified_cache.cache_action import (
CacheAction,
ComponentAction,
@@ -424,6 +428,68 @@ class FullComponent(TreeComponent):
self.tree_core._update_evictable_leaf_sets(node)
def _full_allocator(self):
"""The allocator that owns the full-attention pool alone."""
allocator = self.cache.token_to_kv_pool_allocator
return allocator.full_attn_allocator if self.cache.is_swa_enabled else allocator
def build_external_linker_transfer(
self,
phase: LinkerTransferPhase,
node: Optional[UnifiedTreeNode],
keys: Optional[Sequence[str]],
) -> Optional[PoolTransfer]:
if phase == LinkerTransferPhase.OFFLOAD:
if node is None or not node.hash_value:
return None
value = node.component_data[self.component_type].value
if value is None:
return None
return PoolTransfer(
name=PoolName.KV,
device_indices=value.to(torch.int64),
keys=list(node.hash_value),
)
if not keys:
return None
if phase == LinkerTransferPhase.LOOKUP:
return PoolTransfer(name=PoolName.KV, keys=list(keys))
if phase == LinkerTransferPhase.LOAD:
allocator = self._full_allocator()
num_tokens = len(keys) * self.cache.page_size
shortfall = max(0, num_tokens - allocator.available_size())
if shortfall:
self.cache.evict(EvictParams(num_tokens=shortfall))
slots = allocator.alloc(num_tokens)
if slots is None:
return None
return PoolTransfer(
name=PoolName.KV,
device_indices=slots.to(torch.int64),
keys=list(keys),
)
def update_external_linker_load(
self,
phase: ExternalLinkerLoadPhase,
req: Req,
full_transfer: PoolTransfer,
transfer: PoolTransfer,
prefix_len: int,
*,
insert_result: Optional[InsertResult] = None,
canonical_full: Optional[torch.Tensor] = None,
) -> Optional[PoolTransfer]:
if phase == ExternalLinkerLoadPhase.ABORT:
self._full_allocator().free(transfer.device_indices)
return None
return transfer
def free_host_values(self, host_values: list[torch.Tensor]) -> None:
if self._full_kv_pool_host is None:
return
@@ -29,6 +29,7 @@ from sglang.srt.mem_cache.unified_cache.components.tree_component import (
CacheTransferPhase,
ComponentType,
EvictLayer,
LinkerTransferPhase,
LRURefreshPhase,
PrepareLoadBackResult,
PreparePrefetchResult,
@@ -648,6 +649,16 @@ class MambaComponent(TreeComponent):
self._free_mamba_value(insert_params.mamba_value)
req.mamba_last_track_seqlen = None
def build_external_linker_transfer(
self,
phase: LinkerTransferPhase,
node: Optional[UnifiedTreeNode],
keys: Optional[Sequence[str]],
) -> Optional[PoolTransfer]:
raise AssertionError(
"MambaComponent does not support external linker mode, will support soon"
)
# ---- HiCache Hooks ----
def prepare_load_back(
@@ -6,6 +6,7 @@ import torch
from sglang.srt.mem_cache.base_prefix_cache import (
DecLockRefParams,
EvictParams,
IncLockRefResult,
InsertParams,
InsertResult,
@@ -32,6 +33,8 @@ from sglang.srt.mem_cache.unified_cache.components.tree_component import (
CacheTransferPhase,
ComponentType,
EvictLayer,
ExternalLinkerLoadPhase,
LinkerTransferPhase,
LRURefreshPhase,
PreparePrefetchResult,
TreeComponent,
@@ -268,6 +271,7 @@ class SWAComponent(TreeComponent):
total_prefix_len: int,
value_slice: torch.Tensor,
params: InsertParams,
result: InsertResult,
cache_actions: list[CacheAction | ComponentAction],
) -> int:
if params.prev_prefix_len >= total_prefix_len + prefix_len:
@@ -288,12 +292,22 @@ class SWAComponent(TreeComponent):
if swa_evicted_seqlen <= total_prefix_len:
# Branch 1: entire value_slice is within SWA window — recover
result.record_adopted_range(
self.component_type,
total_prefix_len,
total_prefix_len + prefix_len,
)
old_full = full_cd.value
if full_cd.lock_ref > 0:
cache_actions.append(
RecoverSWAWithLockedFull(node.id, old_full, value_slice)
)
return 0
result.record_adopted_range(
BASE_COMPONENT_TYPE,
total_prefix_len,
total_prefix_len + prefix_len,
)
full_cd.value = value_slice.clone()
cache_actions.append(FreeDeviceKVFullOnly([old_full]))
cache_actions.append(SWARebuild(node.id, value_slice))
@@ -301,6 +315,11 @@ class SWAComponent(TreeComponent):
elif swa_evicted_seqlen < total_prefix_len + prefix_len:
# Branch 2: value_slice[start_idx:] is within SWA window — partial recover
start_idx = swa_evicted_seqlen - total_prefix_len
result.record_adopted_range(
self.component_type,
swa_evicted_seqlen,
total_prefix_len + prefix_len,
)
is_locked = full_cd.lock_ref > 0
old_full = full_cd.value[start_idx:]
_, action = self.tree_core._split_node(node.key, node, start_idx)
@@ -312,6 +331,11 @@ class SWAComponent(TreeComponent):
RecoverSWAWithLockedFull(node.id, old_full, new_full)
)
return start_idx
result.record_adopted_range(
BASE_COMPONENT_TYPE,
swa_evicted_seqlen,
total_prefix_len + prefix_len,
)
node.component_data[BASE_COMPONENT_TYPE].value = new_full.clone()
cache_actions.append(FreeDeviceKVFullOnly([old_full]))
cache_actions.append(SWARebuild(node.id, new_full))
@@ -326,6 +350,7 @@ class SWAComponent(TreeComponent):
prefix_len: int,
total_prefix_len: int,
params: InsertParams,
result: InsertResult,
cache_actions: list[CacheAction | ComponentAction],
) -> None:
# _unevict_node_on_insert already wrote the request's fresh KV slice
@@ -351,6 +376,11 @@ class SWAComponent(TreeComponent):
cache_actions.append(action)
else:
return
result.record_adopted_range(
self.component_type,
max(total_prefix_len, swa_evicted_seqlen),
total_prefix_len + prefix_len,
)
cache_actions.append(
SWARebuild(
node.id,
@@ -370,10 +400,16 @@ class SWAComponent(TreeComponent):
return
node_start = result.prefix_len
node_end = node_start + len(node.key)
split_pos = params.swa_evicted_seqlen - node_start
if split_pos >= len(node.key):
# Entire leaf is outside the SWA window — left as a tombstone.
return
result.record_adopted_range(
self.component_type,
max(node_start, params.swa_evicted_seqlen),
node_end,
)
if split_pos > 0:
# Node straddles the boundary: split into an out-of-window parent
# (tombstone) and an in-window child; `node` becomes the child.
@@ -908,6 +944,99 @@ class SWAComponent(TreeComponent):
return None
def build_external_linker_transfer(
self,
phase: LinkerTransferPhase,
node: Optional[UnifiedTreeNode],
keys: Optional[Sequence[str]],
) -> Optional[PoolTransfer]:
page = self.cache.page_size
window_pages = (self.sliding_window_size + page - 1) // page
if phase == LinkerTransferPhase.OFFLOAD:
if node is None or not node.hash_value:
return None
value = node.component_data[self.component_type].value
if value is None or len(value) < page:
return None
num_pages = len(value) // page
return PoolTransfer(
name=PoolName.SWA,
device_indices=value[-num_pages * page :].to(torch.int64),
keys=node.hash_value[-num_pages:],
hit_policy=PoolHitPolicy.TRAILING_PAGES,
)
if not keys:
return None
# `keys` already start at the first device-uncached page, so the trailing
# window is simply their tail.
tail_keys = list(keys[max(0, len(keys) - window_pages) :])
if not tail_keys:
return None
transfer = PoolTransfer(
name=PoolName.SWA,
keys=tail_keys,
hit_policy=PoolHitPolicy.TRAILING_PAGES,
)
if phase == LinkerTransferPhase.LOAD:
num_tokens = len(tail_keys) * page
allocator = self.cache.token_to_kv_pool_allocator.swa_attn_allocator
shortfall = max(0, num_tokens - allocator.available_size())
if shortfall:
self.cache.evict(EvictParams(swa_num_tokens=shortfall))
transfer.device_indices = allocator.alloc(num_tokens)
if transfer.device_indices is None:
return None
transfer.device_indices = transfer.device_indices.to(torch.int64)
return transfer
def update_external_linker_load(
self,
phase: ExternalLinkerLoadPhase,
req: Req,
full_transfer: PoolTransfer,
transfer: PoolTransfer,
prefix_len: int,
*,
insert_result: Optional[InsertResult] = None,
canonical_full: Optional[torch.Tensor] = None,
) -> Optional[PoolTransfer]:
if phase == ExternalLinkerLoadPhase.ABORT:
self.cache.token_to_kv_pool_allocator.swa_attn_allocator.free(
transfer.device_indices
)
return None
allocator = self.cache.token_to_kv_pool_allocator
if phase == ExternalLinkerLoadPhase.PREPARE:
swa_len = len(transfer.device_indices)
allocator.set_full_to_swa_mapping(
full_transfer.device_indices[-swa_len:], transfer.device_indices
)
page = self.cache.page_size
window = ((self.sliding_window_size + page - 1) // page) * page
boundary = max(0, prefix_len - window)
if req.kv is None:
from sglang.srt.managers.schedule_batch import ReqKvInfo
req.kv = ReqKvInfo(
kv_allocated_len=prefix_len,
swa_evicted_seqlen=boundary,
)
else:
req.kv.swa_evicted_seqlen = max(req.kv.swa_evicted_seqlen, boundary)
return transfer
assert phase == ExternalLinkerLoadPhase.COMMIT
assert insert_result is not None and canonical_full is not None
assert len(canonical_full) == len(transfer.device_indices)
allocator.set_full_to_swa_mapping(canonical_full, transfer.device_indices)
return transfer
def commit_hicache_transfer(
self,
node: UnifiedTreeNode,
@@ -89,8 +89,19 @@ class CacheTransferPhase(str, Enum):
PREFETCH = "prefetch" # Storage→H
class LRURefreshPhase(str, Enum):
class LinkerTransferPhase(str, Enum):
LOOKUP = "lookup"
LOAD = "load"
OFFLOAD = "offload"
class ExternalLinkerLoadPhase(str, Enum):
PREPARE = "prepare"
COMMIT = "commit"
ABORT = "abort"
class LRURefreshPhase(str, Enum):
WALKDOWN = "walkdown" # touching a node while walking through the tree
MATCH_END = "match_end" # end of a successful prefix match
INSERT_END = "insert_end" # after a new/updated leaf is committed
@@ -382,6 +393,7 @@ class TreeComponent(ABC):
total_prefix_len: int,
value_slice: torch.Tensor,
params: InsertParams,
result: InsertResult,
cache_actions: list[CacheAction | ComponentAction],
) -> int:
"""Called per-node when an insert's key overlaps an existing node.
@@ -398,6 +410,7 @@ class TreeComponent(ABC):
prefix_len: int,
total_prefix_len: int,
params: InsertParams,
result: InsertResult,
cache_actions: list[CacheAction | ComponentAction],
) -> None:
"""Called after _unevict_node_on_insert restores the base (Full) value
@@ -703,3 +716,34 @@ class TreeComponent(ABC):
raise NotImplementedError(
f"{self.component_type} cannot apply {type(action).__name__}"
)
# ---- External Cache Linker Hooks ----
def build_external_linker_transfer(
self,
phase: LinkerTransferPhase,
node: Optional[UnifiedTreeNode],
keys: Optional[Sequence[str]],
) -> Optional[PoolTransfer]:
"""Build this component's direct device/storage transfer.
``node`` carries the device pages to persist on OFFLOAD and is None
otherwise. ``keys`` are the per-page hashes of the device-uncached tail
(page 0 is the first uncached page) on LOOKUP / LOAD, and None on
OFFLOAD, where the keys come from ``node.hash_value``.
"""
return None
def update_external_linker_load(
self,
phase: ExternalLinkerLoadPhase,
req: Req,
full_transfer: PoolTransfer,
transfer: PoolTransfer,
prefix_len: int,
*,
insert_result: Optional[InsertResult] = None,
canonical_full: Optional[torch.Tensor] = None,
) -> Optional[PoolTransfer]:
"""Prepare, commit, or abort this component's direct load."""
return transfer
@@ -122,6 +122,7 @@ class UnifiedTreeNode:
# Namespace-aware hashes used only for external KV events.
self.event_hash_value: Optional[list[str]] = None
self.hit_count = 0
self.external_cache_stored = False
self.priority = priority
self.lru_prev: list[UnifiedTreeNode | None] = [None] * (
_NUM_COMPONENT_TYPES * 2
@@ -371,10 +372,10 @@ class _InsertWalkState(msgspec.Struct):
value: torch.Tensor
params: InsertParams
priority: int
result: InsertResult
total_prefix_length: int = 0
is_new_leaf: bool = False
target_node: Optional[UnifiedTreeNode] = None
result: Optional[InsertResult] = None
# Emitted actions awaiting the next barrier flush (or the final step).
pending_actions: list[CacheAction | ComponentAction] = []
@@ -395,6 +396,7 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface):
self.is_eagle = params.is_eagle and ComponentType.MAMBA not in components
self.enable_hicache = False
self.enable_storage = False
self.enable_external_cache_linker = False
self.write_through_threshold = 256
self.is_write_back = False
self.has_swa_host_pool = False
@@ -854,6 +856,13 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface):
if self.is_write_back:
return False
node.hit_count += 1
if self.enable_external_cache_linker:
return (
not node.external_cache_stored
and node.hit_count >= self.write_through_threshold
)
return (
self.enable_hicache
and not node.backuped
@@ -895,6 +904,10 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface):
value=value,
params=params,
priority=priority,
result=InsertResult(
prefix_len=0,
adopted_ranges={} if params.track_adopted_ranges else None,
),
)
return self._advance_insert()
@@ -964,6 +977,11 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface):
if node.evicted:
self._unevict_node_on_insert(node, state.value[:prefix_len])
state.result.record_adopted_range(
BASE_COMPONENT_TYPE,
state.total_prefix_length,
state.total_prefix_length + prefix_len,
)
# FULL was restored from the request's fresh KV. Aux
# components (e.g. SWA) may still hold tombstones and need
# to rebuild their value from the same slice.
@@ -975,6 +993,7 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface):
prefix_len=prefix_len,
total_prefix_len=state.total_prefix_length,
params=state.params,
result=state.result,
cache_actions=step_actions,
)
else:
@@ -988,6 +1007,7 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface):
total_prefix_len=state.total_prefix_length,
value_slice=value_slice,
params=state.params,
result=state.result,
cache_actions=step_actions,
)
consumed_from = min(consumed_from, comp_consumed_from)
@@ -1012,6 +1032,11 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface):
# only a tombstone for this span (e.g. the whole leaf is outside the SWA
# window). Materialize it anyway so the Full KV stays cacheable.
if len(state.key):
state.result.record_adopted_range(
BASE_COMPONENT_TYPE,
state.total_prefix_length,
state.total_prefix_length + len(state.key),
)
state.target_node = self._add_new_node(
state.node, state.key, state.value, priority=state.priority
)
@@ -1023,10 +1048,8 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface):
# e.g. Mamba attaches mamba_value to the leaf node
# All hooks run before their emitted actions execute; an action failure
# fail-stops the process, so partial-commit state is never observed.
state.result = InsertResult(
prefix_len=state.total_prefix_length,
last_device_node=state.target_node.id,
)
state.result.prefix_len = state.total_prefix_length
state.result.last_device_node = state.target_node.id
for component in self.components:
component.commit_insert_component_data(
node=state.target_node,
@@ -1083,6 +1106,7 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface):
new_node.parent = child.parent
new_node.key = child.key[:split_len]
new_node.hit_count = child.hit_count
new_node.external_cache_stored = child.external_cache_stored
new_node.creation_time = child.creation_time
# Split fragments stay on the anchor's root path for the ack's walk.
new_node.load_back_pending_id = child.load_back_pending_id
@@ -1141,7 +1165,7 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface):
new_node.component_data[BASE_COMPONENT_TYPE].value = value.clone()
parent.children[key.child_key(self.page_size)] = new_node
self.component_evictable_size_[BASE_COMPONENT_TYPE] += len(value)
if self.enable_storage:
if self.enable_storage or self.enable_external_cache_linker:
new_node.hash_value = compute_node_hash_values(new_node, self.page_size)
self._update_evictable_leaf_sets(new_node)
@@ -1955,14 +1979,14 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface):
def _build_backup_kv_action(
self, node: UnifiedTreeNode, write_back: bool = False
) -> BackupKV:
"""Build the backup action for a node and its unbacked ancestors."""
"""Build the backup action for a node and its not-yet-persisted ancestors."""
chain = [node]
if not write_back:
ancestor = node.parent
while (
ancestor is not None
and ancestor is not self.root_node
and not ancestor.backuped
and not (ancestor.backuped or ancestor.external_cache_stored)
):
chain.append(ancestor)
ancestor = ancestor.parent