[Rust TreeCore] Support external cache linker (#37306)

This commit is contained in:
Jialin Ouyang
2026-09-10 19:22:24 +08:00
committed by GitHub
parent 334e94d8ac
commit 908226fea2
18 changed files with 1549 additions and 90 deletions
@@ -679,15 +679,11 @@ class RustUnifiedTreeCore(UnifiedTreeCoreInterface):
@property
def enable_external_cache_linker(self) -> bool:
return False
return self._binding.enable_external_cache_linker()
@enable_external_cache_linker.setter
def enable_external_cache_linker(self, value: bool) -> None:
# TODO(Jialin): Port external cache linker support from #37091 and #37151.
if value:
raise ValueError(
"External cache linker is not supported by the Rust TreeCore"
)
self._binding.set_enable_external_cache_linker(value)
def insert_host(
self,
@@ -913,6 +909,27 @@ class RustUnifiedTreeCore(UnifiedTreeCoreInterface):
def finish_load_back(self, anchor_node_id: NodeId) -> None:
self._binding.finish_load_back(anchor_node_id)
def build_external_linker_offload_transfers(
self, node_id: NodeId
) -> Optional[list[PoolTransfer]]:
transfers = self._binding.build_external_linker_offload_transfers(node_id)
if transfers is None:
return None
return [_transfer_from_binding(transfer) for transfer in transfers]
def mark_external_cache_stored_path(
self, from_node_id: NodeId, until_node_id: NodeId
) -> None:
self._binding.mark_external_cache_stored_path(from_node_id, until_node_id)
def mark_external_linker_offload_pending(self, node_id: NodeId) -> None:
self._binding.mark_external_linker_offload_pending(node_id)
def finish_external_linker_offload(
self, node_ids: Sequence[NodeId], ack_id: NodeId, success: bool
) -> None:
self._binding.finish_external_linker_offload(list(node_ids), ack_id, success)
@property
def write_back_duplicate_reclaim_digest(self) -> int:
return self._binding.write_back_duplicate_reclaim_digest()
@@ -36,6 +36,7 @@ from sglang.srt.mem_cache.hicache_storage import (
)
from sglang.srt.mem_cache.radix_cache import RadixKey
from sglang.srt.mem_cache.unified_cache.components import (
ComponentType,
ExternalLinkerLoadPhase,
LinkerTransferPhase,
TreeComponent,
@@ -48,6 +49,14 @@ if TYPE_CHECKING:
from sglang.srt.mem_cache.unified_radix_cache import UnifiedRadixCache
_EXTERNAL_LINKER_SUPPORTED_COMPONENTS = frozenset(
{
ComponentType.FULL,
ComponentType.SWA,
}
)
class UnifiedCacheLinker(ABC):
"""External KV store reached directly from the device pools."""
@@ -139,6 +148,16 @@ class UnifiedCacheLinkerWrapper:
cache: UnifiedRadixCache,
cache_linker: UnifiedCacheLinker,
):
unsupported = set(cache.tree_components) - _EXTERNAL_LINKER_SUPPORTED_COMPONENTS
if unsupported:
names = ", ".join(
component.name for component in sorted(unsupported, key=int)
)
raise ValueError(
"External cache linker supports only Full and SWA tree "
f"components; unsupported: {names}"
)
self.cache = cache
self.cache_linker = cache_linker
# rid -> what match found, consumed by the next init_load_back.
@@ -162,6 +181,7 @@ class UnifiedCacheLinkerWrapper:
def match(self, key: RadixKey, req: Req, result: MatchResult) -> MatchResult:
cache = self.cache
key, _ = key.maybe_to_bigram_view(cache.tree_core.is_eagle)
page = cache.page_size
device_hit_len = int(result.device_indices.numel())
if device_hit_len >= len(key):
@@ -343,10 +363,9 @@ class UnifiedCacheLinkerWrapper:
self._queue_load(req.rid, insert_result.last_device_node, load_transfers)
node = cache.resolve_node_handle(insert_result.last_device_node)
while node.id != req.last_node:
node.external_cache_stored = True
node = node.parent
cache.tree_core.mark_external_cache_stored_path(
insert_result.last_device_node, req.last_node
)
return canonical_tail, insert_result.last_device_node
def _queue_load(
@@ -461,20 +480,14 @@ class UnifiedCacheLinkerWrapper:
def offload_nodes(self, node_ids: Sequence[NodeId]) -> None:
"""Persist a write-through chain, skipping nodes already in the store."""
for node_id in node_ids:
if not self.cache.resolve_node_handle(node_id).external_cache_stored:
self._offload_node(node_id)
def _offload_node(self, node_id: NodeId) -> None:
cache = self.cache
node = cache.resolve_node_handle(node_id)
transfers = []
for component in cache._components_tuple:
transfer = component.build_external_linker_transfer(
LinkerTransferPhase.OFFLOAD, node, None
transfers = self.cache.tree_core.build_external_linker_offload_transfers(
node_id
)
if transfer is not None:
transfers.append(transfer)
if transfers is not None:
self._offload_node(node_id, transfers)
def _offload_node(self, node_id: NodeId, transfers: list[PoolTransfer]) -> None:
cache = self.cache
lock_params = cache.inc_lock_ref(node_id).to_dec_params()
try:
queued = self.cache_linker.offload(transfers)
@@ -485,8 +498,7 @@ class UnifiedCacheLinkerWrapper:
cache.dec_lock_ref(node_id, lock_params)
return
cache.tree_core.mark_write_through_pending([node_id], ack_id=node_id)
node.external_cache_stored = True
cache.tree_core.mark_external_linker_offload_pending(node_id)
self.pending_offloads.append(_PendingOffload(node_id, lock_params, [node_id]))
def replace_pending_offload_node(
@@ -528,11 +540,9 @@ class UnifiedCacheLinkerWrapper:
assert len(successes) <= len(self.pending_offloads)
for success in successes:
pending = self.pending_offloads.pop(0)
for node_id in pending.publish_node_ids:
node = self.cache.resolve_node_handle(node_id)
if node.write_through_pending_id == pending.lock_node_id:
node.write_through_pending_id = None
node.external_cache_stored = success
self.cache.tree_core.finish_external_linker_offload(
pending.publish_node_ids, pending.lock_node_id, success
)
self.cache.dec_lock_ref(pending.lock_node_id, pending.lock_params)
def start_layer_wise_loading(self) -> int:
@@ -550,11 +560,9 @@ class UnifiedCacheLinkerWrapper:
self.cache.dec_lock_ref(node_id, lock_params)
self.pending_loads.clear()
for pending in self.pending_offloads:
for node_id in pending.publish_node_ids:
node = self.cache.resolve_node_handle(node_id)
if node.write_through_pending_id == pending.lock_node_id:
node.write_through_pending_id = None
node.external_cache_stored = False
self.cache.tree_core.finish_external_linker_offload(
pending.publish_node_ids, pending.lock_node_id, False
)
self.cache.dec_lock_ref(pending.lock_node_id, pending.lock_params)
self.pending_offloads.clear()
@@ -60,6 +60,7 @@ from sglang.srt.mem_cache.unified_cache.components import (
ComponentData,
ComponentType,
EvictLayer,
LinkerTransferPhase,
LRURefreshPhase,
TreeComponent,
get_and_increase_time_counter,
@@ -954,7 +955,7 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface):
if self.enable_external_cache_linker:
return (
not node.external_cache_stored
self._needs_external_linker_offload(node)
and node.hit_count >= self.write_through_threshold
)
@@ -964,6 +965,11 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface):
and node.hit_count >= self.write_through_threshold
)
@staticmethod
def _needs_external_linker_offload(node: UnifiedTreeNode) -> bool:
"""Whether neither a confirmed nor an in-flight external copy exists."""
return not node.external_cache_stored and node.write_through_pending_id is None
def begin_insert(self, params: InsertParams) -> InsertStepResult:
"""Start the insert, running to its first barrier or completion."""
# Insert walks are single-flight; a live walk means re-entrancy.
@@ -2140,7 +2146,12 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface):
while (
ancestor is not None
and ancestor is not self.root_node
and not (ancestor.backuped or ancestor.external_cache_stored)
and not ancestor.backuped
and not ancestor.external_cache_stored
and (
not self.enable_external_cache_linker
or ancestor.write_through_pending_id is None
)
):
chain.append(ancestor)
ancestor = ancestor.parent
@@ -2279,6 +2290,69 @@ class UnifiedTreeCore(UnifiedTreeCoreInterface):
node = node.parent
return depth
def build_external_linker_offload_transfers(
self, node_id: NodeId
) -> Optional[list[PoolTransfer]]:
"""Build transfers for a node with no stored or pending external copy."""
node = self.node_by_id(node_id)
if not self._needs_external_linker_offload(node):
return None
transfers = []
for component in self.components:
transfer = component.build_external_linker_transfer(
LinkerTransferPhase.OFFLOAD, node, None
)
if transfer is not None:
transfers.append(transfer)
return transfers
def mark_external_cache_stored_path(
self, from_node_id: NodeId, until_node_id: NodeId
) -> None:
"""Mark an externally restored path, excluding its existing anchor."""
until_node = self.node_by_id(until_node_id)
node = self.node_by_id(from_node_id)
path = []
while node is not until_node:
if node.parent is None:
raise RuntimeError(
f"node {until_node_id} is not an ancestor of node {from_node_id}"
)
path.append(node)
node = node.parent
for node in path:
node.external_cache_stored = True
def mark_external_linker_offload_pending(self, node_id: NodeId) -> None:
"""Publish an accepted external offload as pending."""
node = self.node_by_id(node_id)
if not self._needs_external_linker_offload(node):
raise AssertionError(
f"invalid external offload state for node {node_id}: "
f"stored={node.external_cache_stored}, "
f"pending={node.write_through_pending_id}"
)
node.write_through_pending_id = node_id
def finish_external_linker_offload(
self, node_ids: Sequence[NodeId], ack_id: NodeId, success: bool
) -> None:
"""Finalize external-store state for an offload and its split fragments."""
nodes = [self.node_by_id(node_id) for node_id in node_ids]
for node_id, node in zip(node_ids, nodes):
if node.write_through_pending_id != ack_id:
raise AssertionError(
f"invalid external offload state for node {node_id}: "
f"expected pending={ack_id}; got "
f"stored={node.external_cache_stored}, "
f"pending={node.write_through_pending_id}"
)
for node in nodes:
node.write_through_pending_id = None
node.external_cache_stored |= success
def finish_write_through(self, node_ids: list[NodeId], ack_id: int) -> None:
"""Clear the write-through-pending mark (when it matches ack_id) and record the
host store event for each acked node."""
@@ -148,6 +148,7 @@ class UnifiedTreeCoreInterface(ABC):
device: torch.device
enable_hicache: bool
enable_storage: bool
enable_external_cache_linker: bool
write_through_threshold: int
is_write_back: bool
has_swa_host_pool: bool
@@ -533,6 +534,41 @@ class UnifiedTreeCoreInterface(ABC):
"""Clear the in-flight H->D marks on the anchor's root path at ack time."""
...
# ==== External Cache Linker ====
@abstractmethod
def build_external_linker_offload_transfers(
self, node_id: NodeId
) -> Optional[list[PoolTransfer]]:
"""Build direct device-to-external-store transfers for an eligible node.
Return None when the node is stored externally or has an offload pending.
"""
...
@abstractmethod
def mark_external_cache_stored_path(
self, from_node_id: NodeId, until_node_id: NodeId
) -> None:
"""Mark the path from ``from_node_id`` to, but excluding, ``until_node_id``."""
...
@abstractmethod
def mark_external_linker_offload_pending(self, node_id: NodeId) -> None:
"""Publish an accepted external offload as pending."""
...
@abstractmethod
def finish_external_linker_offload(
self, node_ids: Sequence[NodeId], ack_id: NodeId, success: bool
) -> None:
"""Finish one external offload for every current fragment of its node.
A successful write confirms external storage. A failed redundant write
preserves storage already confirmed independently by a concurrent load.
"""
...
# Order-sensitive digest of write_back duplicate-reclaim victim ids,
# cross-checked across TP ranks; cores that never reclaim keep 0.
write_back_duplicate_reclaim_digest: int = 0