[Rust TreeCore] Support external cache linker (#37306)
This commit is contained in:
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user