[Unified Cache Linker][3/N]: Add backend-independent linker core (#37151)
Co-authored-by: 晟海 <huangtingwei.htw@antgroup.com>
This commit is contained in:
@@ -0,0 +1,571 @@
|
||||
"""Direct L3 support for :class:`UnifiedRadixCache`.
|
||||
|
||||
Links the cache's device pools straight to an external KV store, with no host
|
||||
tier in between. The transport contract and tree-side wrapper live here, while
|
||||
each backend owns its device-pool layout and physical I/O.
|
||||
|
||||
* :class:`UnifiedCacheLinker` -- the transport interface a backend implements.
|
||||
* :class:`UnifiedCacheLinkerWrapper` -- the tree-side flow that drives it. The
|
||||
cache owns one as a plain attribute, keeping the whole external-cache path out
|
||||
of the main tree file.
|
||||
|
||||
The tree only needs a handful of guarded hooks:
|
||||
|
||||
* ``match_prefix`` -> :meth:`UnifiedCacheLinkerWrapper.match`
|
||||
* ``init_load_back`` -> :meth:`UnifiedCacheLinkerWrapper.load_back`
|
||||
* ``BackupKV`` actions -> :meth:`UnifiedCacheLinkerWrapper.offload_nodes`
|
||||
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from collections.abc import Sequence
|
||||
from typing import TYPE_CHECKING, NamedTuple
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.mem_cache.base_prefix_cache import (
|
||||
DecLockRefParams,
|
||||
InsertParams,
|
||||
MatchResult,
|
||||
)
|
||||
from sglang.srt.mem_cache.hicache_storage import (
|
||||
PoolName,
|
||||
PoolTransfer,
|
||||
)
|
||||
from sglang.srt.mem_cache.radix_cache import RadixKey
|
||||
from sglang.srt.mem_cache.unified_cache.components import (
|
||||
ExternalLinkerLoadPhase,
|
||||
LinkerTransferPhase,
|
||||
TreeComponent,
|
||||
)
|
||||
from sglang.srt.mem_cache.utils import get_hash_str
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.managers.schedule_batch import Req
|
||||
from sglang.srt.mem_cache.unified_cache.unified_tree_core_interface import NodeId
|
||||
from sglang.srt.mem_cache.unified_radix_cache import UnifiedRadixCache
|
||||
|
||||
|
||||
class UnifiedCacheLinker(ABC):
|
||||
"""External KV store reached directly from the device pools."""
|
||||
|
||||
layer_done_counter: object
|
||||
|
||||
@abstractmethod
|
||||
def lookup(self, rid: str, transfers: list[PoolTransfer]) -> list[int]:
|
||||
"""Return every prefix length (in pages) that is fully restorable.
|
||||
|
||||
A length is included only when *all* pools satisfy their hit policy at
|
||||
that exact boundary (contiguous prefix pools, plus each trailing-window
|
||||
pool's window ending there). Trailing-window state (SWA / compress
|
||||
state) only exists at offloaded node boundaries, so the set is sparse
|
||||
and generally non-contiguous -- returning just the local maximum would
|
||||
let the tree pick a length that is invalid on another rank.
|
||||
|
||||
Local to this rank; the tree intersects the sets across ranks.
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
def load(self, rid: str, transfers: list[PoolTransfer]) -> bool:
|
||||
"""Queue a load into the given device indices.
|
||||
|
||||
The transfer is executed by the next ``start_layer_wise_loading`` call,
|
||||
not here.
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
def start_layer_wise_loading(self) -> int:
|
||||
"""Start queued loads and return the layer-counter consumer index."""
|
||||
|
||||
@abstractmethod
|
||||
def cancel_queued_load(self, rid: str) -> bool:
|
||||
"""Cancel a load that has not started yet."""
|
||||
|
||||
@abstractmethod
|
||||
def num_completed_loads(self) -> int:
|
||||
"""Return the number of completed load batches waiting to be consumed."""
|
||||
|
||||
@abstractmethod
|
||||
def pop_completed_load(self) -> list[str]:
|
||||
"""Consume the oldest completed load batch and return its request IDs."""
|
||||
|
||||
@abstractmethod
|
||||
def offload(self, transfers: list[PoolTransfer]) -> bool:
|
||||
"""Queue every transfer for atomic persistence."""
|
||||
|
||||
@abstractmethod
|
||||
def num_completed_offloads(self) -> int:
|
||||
"""Return the number of completed offloads waiting to be consumed."""
|
||||
|
||||
@abstractmethod
|
||||
def pop_completed_offload(self) -> bool:
|
||||
"""Consume the oldest completed offload and return its result."""
|
||||
|
||||
@abstractmethod
|
||||
def reset(self) -> None:
|
||||
"""Quiesce all transfers and reset backend state before returning."""
|
||||
|
||||
@abstractmethod
|
||||
def close(self) -> None:
|
||||
"""Quiesce all transfers and release backend resources."""
|
||||
|
||||
|
||||
class ExternalCacheHitMarker(NamedTuple):
|
||||
"""What ``match`` found in the external store, consumed by ``load_back``.
|
||||
|
||||
``prefix_key`` covers the device-cached prefix plus the restorable tail, so
|
||||
it is what gets inserted once the tail lands. ``tail_hashes`` are the
|
||||
per-page storage hashes of that tail alone, starting at ``device_hit_len``.
|
||||
"""
|
||||
|
||||
prefix_key: RadixKey
|
||||
tail_hashes: list[str]
|
||||
device_hit_len: int
|
||||
|
||||
|
||||
class _PendingOffload(NamedTuple):
|
||||
lock_node_id: NodeId
|
||||
lock_params: DecLockRefParams
|
||||
publish_node_ids: list[NodeId]
|
||||
|
||||
|
||||
class UnifiedCacheLinkerWrapper:
|
||||
"""Drives an external KV store on behalf of one :class:`UnifiedRadixCache`."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
cache: UnifiedRadixCache,
|
||||
cache_linker: UnifiedCacheLinker,
|
||||
):
|
||||
self.cache = cache
|
||||
self.cache_linker = cache_linker
|
||||
# rid -> what match found, consumed by the next init_load_back.
|
||||
self.hit_markers: dict[str, ExternalCacheHitMarker] = {}
|
||||
# Loads in flight, each pinning its inserted endpoint until DMA completes.
|
||||
self.pending_loads: dict[str, tuple[NodeId, DecLockRefParams]] = {}
|
||||
# Offloads in flight, each holding a lock on its node until it lands.
|
||||
self.pending_offloads: list[_PendingOffload] = []
|
||||
|
||||
cache.tree_core.enable_external_cache_linker = True
|
||||
cache.write_through_threshold = 1
|
||||
|
||||
@property
|
||||
def layer_done_counter(self) -> object:
|
||||
return self.cache_linker.layer_done_counter
|
||||
|
||||
def has_hit(self, rid: str) -> bool:
|
||||
return rid in self.hit_markers
|
||||
|
||||
# ---- match: probe the remote store and report host_hit_length ----
|
||||
|
||||
def match(self, key: RadixKey, req: Req, result: MatchResult) -> MatchResult:
|
||||
cache = self.cache
|
||||
page = cache.page_size
|
||||
device_hit_len = int(result.device_indices.numel())
|
||||
if device_hit_len >= len(key):
|
||||
return result
|
||||
|
||||
tail_hashes = self._tail_hashes(key, result, device_hit_len)
|
||||
if not tail_hashes:
|
||||
return result
|
||||
|
||||
lookup_transfers = []
|
||||
for component in cache._components_tuple:
|
||||
transfer = component.build_external_linker_transfer(
|
||||
LinkerTransferPhase.LOOKUP, None, tail_hashes
|
||||
)
|
||||
if transfer is None:
|
||||
return result
|
||||
lookup_transfers.append(transfer)
|
||||
by_pool = {transfer.name: transfer for transfer in lookup_transfers}
|
||||
|
||||
# Tail-relative: page 0 of `tail_hashes` is the first uncached page.
|
||||
hit_pages = self._sync_restorable_prefix(
|
||||
self.cache_linker.lookup(req.rid, lookup_transfers),
|
||||
num_pages=len(tail_hashes),
|
||||
device_hit_pages=0,
|
||||
)
|
||||
if hit_pages == 0:
|
||||
return result
|
||||
hit_tokens = hit_pages * page
|
||||
|
||||
swa_transfer = by_pool.get(PoolName.SWA)
|
||||
swa_host_hit_length = (
|
||||
min(len(swa_transfer.keys), hit_pages) * page
|
||||
if swa_transfer is not None
|
||||
else 0
|
||||
)
|
||||
# Mamba keeps a single state slot per node, so a hit is worth one slot.
|
||||
mamba_host_hit_length = 1 if PoolName.MAMBA in by_pool else 0
|
||||
|
||||
self.hit_markers[req.rid] = ExternalCacheHitMarker(
|
||||
prefix_key=key[: device_hit_len + hit_tokens],
|
||||
tail_hashes=list(tail_hashes[:hit_pages]),
|
||||
device_hit_len=device_hit_len,
|
||||
)
|
||||
return result._replace(
|
||||
last_host_node=result.best_match_node,
|
||||
host_hit_length=hit_tokens,
|
||||
swa_host_hit_length=max(result.swa_host_hit_length, swa_host_hit_length),
|
||||
mamba_host_hit_length=max(
|
||||
result.mamba_host_hit_length, mamba_host_hit_length
|
||||
),
|
||||
)
|
||||
|
||||
def _sync_restorable_prefix(
|
||||
self, restorable: list[int], *, num_pages: int, device_hit_pages: int
|
||||
) -> int:
|
||||
"""Intersect the per-rank sets of restorable prefix lengths and return the
|
||||
longest one, or 0 when the ranks share none beyond the device prefix.
|
||||
|
||||
A rank's set is sparse, so reducing per-rank maxima could land on a
|
||||
length that only some ranks can restore. On a 0/1 mask MIN is AND, which
|
||||
makes the reduction an intersection.
|
||||
"""
|
||||
mask = torch.zeros(num_pages + 1, dtype=torch.int)
|
||||
for pages in restorable:
|
||||
if device_hit_pages < pages <= num_pages:
|
||||
mask[pages] = 1
|
||||
self.cache._all_reduce_attn_groups(mask, torch.distributed.ReduceOp.MIN)
|
||||
common = mask.nonzero()
|
||||
if common.numel() == 0:
|
||||
return 0
|
||||
return int(common[-1].item())
|
||||
|
||||
def _tail_hashes(
|
||||
self, key: RadixKey, result: MatchResult, device_hit_len: int
|
||||
) -> list[str]:
|
||||
"""Per-page storage hashes for the device-uncached tail of the prefix."""
|
||||
last_hash = None
|
||||
if device_hit_len > 0:
|
||||
last_hash = self.cache.get_last_hash_value(result.last_device_node)
|
||||
if last_hash is None:
|
||||
# Without the anchor the tail would hash as if it started at the
|
||||
# sequence head, yielding keys that can never match.
|
||||
return []
|
||||
page = self.cache.page_size
|
||||
tail_len = (len(key) - device_hit_len) // page * page
|
||||
if tail_len == 0:
|
||||
return []
|
||||
return get_hash_str(
|
||||
key[device_hit_len : device_hit_len + tail_len],
|
||||
last_hash,
|
||||
page_size=page,
|
||||
)
|
||||
|
||||
# ---- init_load_back: remote -> device, then insert ----
|
||||
|
||||
def load_back(self, req: Req) -> tuple[torch.Tensor, NodeId]:
|
||||
cache = self.cache
|
||||
empty_indices = cache.tree_core.empty_match_result.device_indices
|
||||
hit = self.hit_markers.pop(req.rid, None)
|
||||
if hit is None:
|
||||
return empty_indices, req.last_node
|
||||
|
||||
device_hit_len = hit.device_hit_len
|
||||
tail_hashes = hit.tail_hashes
|
||||
prefix_len = device_hit_len + len(tail_hashes) * cache.page_size
|
||||
|
||||
# Build per-component linker transfers.
|
||||
component_transfers: list[tuple[TreeComponent, PoolTransfer]] = []
|
||||
for component in cache._components_tuple:
|
||||
transfer = component.build_external_linker_transfer(
|
||||
LinkerTransferPhase.LOAD, None, tail_hashes
|
||||
)
|
||||
if transfer is None:
|
||||
self._update_load(
|
||||
ExternalLinkerLoadPhase.ABORT,
|
||||
req,
|
||||
component_transfers,
|
||||
prefix_len,
|
||||
)
|
||||
return empty_indices, req.last_node
|
||||
component_transfers.append((component, transfer))
|
||||
|
||||
full_transfer = component_transfers[0][1]
|
||||
assert full_transfer.name == PoolName.KV
|
||||
self._update_load(
|
||||
ExternalLinkerLoadPhase.PREPARE,
|
||||
req,
|
||||
component_transfers,
|
||||
prefix_len,
|
||||
)
|
||||
|
||||
# Insert the newly loaded tail into the tree.
|
||||
prefix_indices = torch.cat(
|
||||
[req.prefix_indices.to(torch.int64), full_transfer.device_indices]
|
||||
)
|
||||
mamba_transfer = next(
|
||||
(
|
||||
transfer
|
||||
for _, transfer in component_transfers
|
||||
if transfer.name == PoolName.MAMBA
|
||||
),
|
||||
None,
|
||||
)
|
||||
insert_result = cache.insert(
|
||||
InsertParams(
|
||||
key=hit.prefix_key,
|
||||
value=prefix_indices,
|
||||
mamba_value=(
|
||||
mamba_transfer.device_indices[:1]
|
||||
if mamba_transfer is not None
|
||||
else None
|
||||
),
|
||||
prev_prefix_len=device_hit_len,
|
||||
swa_evicted_seqlen=(
|
||||
req.kv.swa_evicted_seqlen if req.kv is not None else 0
|
||||
),
|
||||
chunked=True,
|
||||
priority=getattr(req, "priority", 0) or 0,
|
||||
track_adopted_ranges=True,
|
||||
)
|
||||
)
|
||||
if mamba_transfer is not None and insert_result.mamba_exist:
|
||||
cache.req_to_token_pool.mamba_allocator.free(
|
||||
mamba_transfer.device_indices[:1]
|
||||
)
|
||||
|
||||
canonical_tail = cache.tree_core.collect_full_device_indices(
|
||||
insert_result.last_device_node, req.last_node
|
||||
)
|
||||
assert canonical_tail.numel() == len(tail_hashes) * cache.page_size
|
||||
load_transfers = self._update_load(
|
||||
ExternalLinkerLoadPhase.COMMIT,
|
||||
req,
|
||||
component_transfers,
|
||||
prefix_len,
|
||||
insert_result=insert_result,
|
||||
canonical_full=canonical_tail,
|
||||
)
|
||||
|
||||
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
|
||||
return canonical_tail, insert_result.last_device_node
|
||||
|
||||
def _queue_load(
|
||||
self, rid: str, node_id: NodeId, transfers: list[PoolTransfer]
|
||||
) -> None:
|
||||
if not transfers:
|
||||
return
|
||||
assert rid not in self.pending_loads
|
||||
lock_params = self.cache.inc_lock_ref(node_id).to_dec_params()
|
||||
try:
|
||||
queued = self.cache_linker.load(rid, transfers)
|
||||
except BaseException:
|
||||
self.cache.dec_lock_ref(node_id, lock_params)
|
||||
raise
|
||||
if not queued:
|
||||
self.cache.dec_lock_ref(node_id, lock_params)
|
||||
raise RuntimeError(f"Failed to queue the linker load for rid={rid!r}.")
|
||||
self.pending_loads[rid] = (node_id, lock_params)
|
||||
|
||||
def _update_load(
|
||||
self,
|
||||
phase: ExternalLinkerLoadPhase,
|
||||
req: Req,
|
||||
component_transfers: list[tuple[TreeComponent, PoolTransfer]],
|
||||
prefix_len: int,
|
||||
*,
|
||||
insert_result=None,
|
||||
canonical_full: torch.Tensor | None = None,
|
||||
) -> list[PoolTransfer]:
|
||||
if not component_transfers:
|
||||
return []
|
||||
full = component_transfers[0][1]
|
||||
result = []
|
||||
transfers = (
|
||||
reversed(component_transfers)
|
||||
if phase == ExternalLinkerLoadPhase.ABORT
|
||||
else component_transfers
|
||||
)
|
||||
for component, transfer in transfers:
|
||||
component_canonical = canonical_full
|
||||
if phase == ExternalLinkerLoadPhase.COMMIT:
|
||||
assert insert_result.adopted_ranges is not None
|
||||
coverage_start = prefix_len - len(transfer.device_indices)
|
||||
ranges = [
|
||||
(max(start, coverage_start), min(end, prefix_len))
|
||||
for start, end in insert_result.adopted_ranges.get(
|
||||
component.component_type, ()
|
||||
)
|
||||
if max(start, coverage_start) < min(end, prefix_len)
|
||||
]
|
||||
indices, keys = self._select_adopted_pages(
|
||||
transfer.device_indices,
|
||||
ranges,
|
||||
prefix_len,
|
||||
transfer.keys,
|
||||
)
|
||||
if not keys:
|
||||
continue
|
||||
transfer.device_indices = indices
|
||||
transfer.keys = keys
|
||||
component_canonical, _ = self._select_adopted_pages(
|
||||
canonical_full, ranges, prefix_len
|
||||
)
|
||||
transfer = component.update_external_linker_load(
|
||||
phase,
|
||||
req,
|
||||
full,
|
||||
transfer,
|
||||
prefix_len,
|
||||
insert_result=insert_result,
|
||||
canonical_full=component_canonical,
|
||||
)
|
||||
if transfer is not None:
|
||||
result.append(transfer)
|
||||
return result
|
||||
|
||||
def _select_adopted_pages(
|
||||
self,
|
||||
indices: torch.Tensor,
|
||||
ranges: Sequence[tuple[int, int]],
|
||||
prefix_len: int,
|
||||
keys: Sequence[str] | None = None,
|
||||
) -> tuple[torch.Tensor, list[str]]:
|
||||
page = self.cache.page_size
|
||||
coverage_start = prefix_len - len(indices)
|
||||
pages = indices.reshape(-1, page)
|
||||
if keys is not None:
|
||||
assert len(keys) == len(pages)
|
||||
|
||||
chunks = []
|
||||
selected_keys = []
|
||||
for start, end in ranges:
|
||||
start = max(start, coverage_start)
|
||||
end = min(end, prefix_len)
|
||||
if start >= end:
|
||||
continue
|
||||
assert (start - coverage_start) % page == 0
|
||||
assert (end - coverage_start) % page == 0
|
||||
first = (start - coverage_start) // page
|
||||
last = (end - coverage_start) // page
|
||||
chunks.append(pages[first:last].reshape(-1))
|
||||
if keys is not None:
|
||||
selected_keys.extend(keys[first:last])
|
||||
|
||||
if not chunks:
|
||||
return indices[:0], selected_keys
|
||||
selected = chunks[0] if len(chunks) == 1 else torch.cat(chunks)
|
||||
return selected, selected_keys
|
||||
|
||||
# ---- offload: device -> remote, driven by the write-through chain ----
|
||||
|
||||
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
|
||||
)
|
||||
if transfer is not None:
|
||||
transfers.append(transfer)
|
||||
|
||||
lock_params = cache.inc_lock_ref(node_id).to_dec_params()
|
||||
try:
|
||||
queued = self.cache_linker.offload(transfers)
|
||||
except BaseException:
|
||||
cache.dec_lock_ref(node_id, lock_params)
|
||||
raise
|
||||
if not queued:
|
||||
cache.dec_lock_ref(node_id, lock_params)
|
||||
return
|
||||
|
||||
cache.tree_core.mark_write_through_pending(node_id)
|
||||
node.external_cache_stored = True
|
||||
self.pending_offloads.append(_PendingOffload(node_id, lock_params, [node_id]))
|
||||
|
||||
def replace_pending_offload_node(
|
||||
self, ack_id: NodeId, old_node_id: NodeId, new_node_ids: list[NodeId]
|
||||
) -> None:
|
||||
for index, pending in enumerate(self.pending_offloads):
|
||||
if pending.lock_node_id != ack_id:
|
||||
continue
|
||||
publish_node_ids = []
|
||||
for node_id in pending.publish_node_ids:
|
||||
if node_id == old_node_id:
|
||||
publish_node_ids.extend(new_node_ids)
|
||||
else:
|
||||
publish_node_ids.append(node_id)
|
||||
self.pending_offloads[index] = pending._replace(
|
||||
publish_node_ids=publish_node_ids
|
||||
)
|
||||
return
|
||||
|
||||
def num_completed_offloads(self) -> int:
|
||||
return min(
|
||||
self.cache_linker.num_completed_offloads(), len(self.pending_offloads)
|
||||
)
|
||||
|
||||
def num_completed_loads(self) -> int:
|
||||
return self.cache_linker.num_completed_loads()
|
||||
|
||||
def drain_loads(self, finish_count: int) -> None:
|
||||
for _ in range(finish_count):
|
||||
for rid in self.cache_linker.pop_completed_load():
|
||||
node_id, lock_params = self.pending_loads.pop(rid)
|
||||
self.cache.dec_lock_ref(node_id, lock_params)
|
||||
|
||||
def take_completed_offloads(self, finish_count: int) -> list[bool]:
|
||||
assert finish_count <= len(self.pending_offloads)
|
||||
return [self.cache_linker.pop_completed_offload() for _ in range(finish_count)]
|
||||
|
||||
def commit_completed_offloads(self, successes: Sequence[bool]) -> None:
|
||||
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.dec_lock_ref(pending.lock_node_id, pending.lock_params)
|
||||
|
||||
def start_layer_wise_loading(self) -> int:
|
||||
return self.cache_linker.start_layer_wise_loading()
|
||||
|
||||
# ---- lifecycle ----
|
||||
|
||||
def reset(self) -> None:
|
||||
self.cache_linker.reset()
|
||||
self.hit_markers.clear()
|
||||
self._release_pending_locks()
|
||||
|
||||
def _release_pending_locks(self) -> None:
|
||||
for node_id, lock_params in self.pending_loads.values():
|
||||
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.dec_lock_ref(pending.lock_node_id, pending.lock_params)
|
||||
self.pending_offloads.clear()
|
||||
|
||||
def release_request(self, rid: str) -> None:
|
||||
self.hit_markers.pop(rid, None)
|
||||
# TODO: Roll back the published tree and component state atomically before
|
||||
# canceling; otherwise the tree may retain device slots that were never loaded.
|
||||
if self.cache_linker.cancel_queued_load(rid):
|
||||
node_id, lock_params = self.pending_loads.pop(rid)
|
||||
self.cache.dec_lock_ref(node_id, lock_params)
|
||||
|
||||
def close(self) -> None:
|
||||
self.cache_linker.close()
|
||||
self._release_pending_locks()
|
||||
@@ -72,6 +72,10 @@ from sglang.srt.mem_cache.unified_cache.session_ref_tracker import (
|
||||
)
|
||||
from sglang.srt.mem_cache.unified_cache.storage_attachment import StorageAttachment
|
||||
from sglang.srt.mem_cache.unified_cache.tree_core_registry import create_tree_core
|
||||
from sglang.srt.mem_cache.unified_cache.unified_cache_linker import (
|
||||
UnifiedCacheLinker,
|
||||
UnifiedCacheLinkerWrapper,
|
||||
)
|
||||
from sglang.srt.mem_cache.unified_cache.unified_tree_core import ( # noqa: F401
|
||||
NodeId,
|
||||
UnifiedLRUList,
|
||||
@@ -238,6 +242,7 @@ class UnifiedRadixCache(BasePrefixCache):
|
||||
self.host_pool_group = None # set by attach_hybrid_pool_to_unified_cache
|
||||
# Owns the storage backend lifecycle; built by init_hicache.
|
||||
self._storage_attachment: Optional[StorageAttachment] = None
|
||||
self.linker: Optional[UnifiedCacheLinkerWrapper] = None
|
||||
self.prefetch_stop_policy = "best_effort"
|
||||
self.prefetch_threshold = 256
|
||||
self.prefetch_timeout_base = 1.0
|
||||
@@ -342,7 +347,13 @@ class UnifiedRadixCache(BasePrefixCache):
|
||||
)
|
||||
self.work_list.append(send_work)
|
||||
|
||||
def init_cache_linker(self, cache_linker: UnifiedCacheLinker) -> None:
|
||||
"""Attach an external KV store directly to the device pools."""
|
||||
self.linker = UnifiedCacheLinkerWrapper(self, cache_linker)
|
||||
|
||||
def reset(self) -> None:
|
||||
if self.linker is not None:
|
||||
self.linker.reset()
|
||||
self._reset_full()
|
||||
|
||||
def _reset_full(self) -> None:
|
||||
@@ -497,6 +508,8 @@ class UnifiedRadixCache(BasePrefixCache):
|
||||
self.sidecar_pool_specs.append(spec)
|
||||
|
||||
def release_host_resources(self) -> None:
|
||||
if self.linker is not None:
|
||||
self.linker.close()
|
||||
if self.host_pool_group is not None:
|
||||
self.host_pool_group.destroy()
|
||||
|
||||
@@ -518,6 +531,8 @@ class UnifiedRadixCache(BasePrefixCache):
|
||||
result = component.finalize_match_result_in_cache(params, result)
|
||||
# Finalizers must not emit actions; the walk's were applied above.
|
||||
assert not result.cache_actions
|
||||
if self.linker is not None and params.req is not None:
|
||||
result = self.linker.match(params.key, params.req, result)
|
||||
return result
|
||||
|
||||
def is_chunk_cache(self) -> bool:
|
||||
@@ -1059,6 +1074,12 @@ class UnifiedRadixCache(BasePrefixCache):
|
||||
action.old_node_id,
|
||||
[action.new_node_id, action.new_child_node_id],
|
||||
)
|
||||
if self.linker is not None:
|
||||
self.linker.replace_pending_offload_node(
|
||||
action.ack_id,
|
||||
action.old_node_id,
|
||||
[action.new_node_id, action.new_child_node_id],
|
||||
)
|
||||
elif isinstance(action, FreeDeviceKV):
|
||||
# tree values are page-aligned copies of a kv row: page-exact segments
|
||||
for indices in action.indices:
|
||||
@@ -1067,7 +1088,10 @@ class UnifiedRadixCache(BasePrefixCache):
|
||||
for indices in action.indices:
|
||||
self.token_to_kv_pool_allocator.free_full(indices)
|
||||
elif isinstance(action, BackupKV):
|
||||
self._execute_and_commit_kv_backup(action)
|
||||
if self.linker is not None:
|
||||
self.linker.offload_nodes(action.node_ids)
|
||||
else:
|
||||
self._execute_and_commit_kv_backup(action)
|
||||
else:
|
||||
raise AssertionError(f"unhandled CacheAction: {type(action).__name__}")
|
||||
|
||||
@@ -2083,6 +2107,8 @@ class UnifiedRadixCache(BasePrefixCache):
|
||||
|
||||
@rank_consensus(same_params=True)
|
||||
def release_aborted_request(self, rid: str) -> None:
|
||||
if self.linker is not None:
|
||||
self.linker.release_request(rid)
|
||||
self.prefetch_loaded_tokens_by_reqid.pop(rid, None)
|
||||
self._storage_prefetch_missed_rids.discard(rid)
|
||||
if (
|
||||
@@ -2723,6 +2749,8 @@ class UnifiedRadixCache(BasePrefixCache):
|
||||
mem_quota = params.mem_quota
|
||||
req = params.req
|
||||
assert req is not None
|
||||
if self.linker is not None and self.linker.has_hit(req.rid):
|
||||
return self.linker.load_back(req)
|
||||
last_best_match_device_node_id = req.last_node
|
||||
|
||||
if (
|
||||
@@ -2757,6 +2785,27 @@ class UnifiedRadixCache(BasePrefixCache):
|
||||
|
||||
def check_hicache_events(self) -> None:
|
||||
"""Called per scheduler step to poll async HiCache events."""
|
||||
if self.linker is not None:
|
||||
finish_counts = torch.tensor(
|
||||
[
|
||||
self.linker.num_completed_loads(),
|
||||
self.linker.num_completed_offloads(),
|
||||
],
|
||||
dtype=torch.int,
|
||||
device="cpu",
|
||||
)
|
||||
self._all_reduce_attn_groups(finish_counts, torch.distributed.ReduceOp.MIN)
|
||||
load_count, offload_count = map(int, finish_counts.tolist())
|
||||
self.linker.drain_loads(load_count)
|
||||
local_successes = self.linker.take_completed_offloads(offload_count)
|
||||
if local_successes:
|
||||
successes = torch.tensor(local_successes, dtype=torch.int, device="cpu")
|
||||
self._all_reduce_attn_groups(successes, torch.distributed.ReduceOp.MIN)
|
||||
self.linker.commit_completed_offloads(
|
||||
[bool(success) for success in successes.tolist()]
|
||||
)
|
||||
return
|
||||
|
||||
# Reap the previous round's PP-sync sends before issuing new ones.
|
||||
self._drain_async_work()
|
||||
|
||||
@@ -2817,6 +2866,8 @@ class UnifiedRadixCache(BasePrefixCache):
|
||||
|
||||
def ready_to_load_host_cache(self) -> int:
|
||||
"""Notify the cache controller to start the KV cache loading."""
|
||||
if self.linker is not None:
|
||||
return self.linker.start_layer_wise_loading()
|
||||
if self.cache_controller is not None:
|
||||
return self.cache_controller.start_loading()
|
||||
return 0
|
||||
|
||||
@@ -57,6 +57,7 @@ class TestUnifiedPPSyncBatching(unittest.TestCase):
|
||||
cache.enable_storage_metrics = False
|
||||
cache.storage_metrics_collector = None
|
||||
cache.buffer_pipeline = None
|
||||
cache.linker = None
|
||||
cache._drain_async_work = MagicMock()
|
||||
cache._all_reduce = MagicMock()
|
||||
cache.writing_check = MagicMock()
|
||||
|
||||
@@ -0,0 +1,447 @@
|
||||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from sglang.srt.mem_cache.base_prefix_cache import InsertResult
|
||||
from sglang.srt.mem_cache.hicache_storage import PoolName, PoolTransfer
|
||||
from sglang.srt.mem_cache.unified_cache.cache_action import (
|
||||
ReplaceWriteThroughOnNodeSplit,
|
||||
)
|
||||
from sglang.srt.mem_cache.unified_cache.component_type import ComponentType
|
||||
from sglang.srt.mem_cache.unified_cache.components.full_component import FullComponent
|
||||
from sglang.srt.mem_cache.unified_cache.components.swa_component import SWAComponent
|
||||
from sglang.srt.mem_cache.unified_cache.components.tree_component import (
|
||||
ExternalLinkerLoadPhase,
|
||||
LinkerTransferPhase,
|
||||
)
|
||||
from sglang.srt.mem_cache.unified_cache.unified_cache_linker import (
|
||||
UnifiedCacheLinker,
|
||||
UnifiedCacheLinkerWrapper,
|
||||
)
|
||||
from sglang.srt.mem_cache.unified_radix_cache import UnifiedRadixCache
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
|
||||
register_cpu_ci(est_time=1, suite="base-a-test-cpu")
|
||||
|
||||
|
||||
class _FakeLinker(UnifiedCacheLinker):
|
||||
def __init__(self):
|
||||
self.layer_done_counter = object()
|
||||
self.restorable = []
|
||||
self.queued_loads = {}
|
||||
self.queued_offloads = []
|
||||
self.completed_loads = []
|
||||
self.completed_offloads = []
|
||||
self.reset_count = 0
|
||||
self.closed = False
|
||||
|
||||
def lookup(self, rid, transfers):
|
||||
return list(self.restorable)
|
||||
|
||||
def load(self, rid, transfers):
|
||||
self.queued_loads[rid] = list(transfers)
|
||||
return True
|
||||
|
||||
def start_layer_wise_loading(self):
|
||||
return 3
|
||||
|
||||
def cancel_queued_load(self, rid):
|
||||
if rid not in self.queued_loads:
|
||||
return False
|
||||
del self.queued_loads[rid]
|
||||
return True
|
||||
|
||||
def num_completed_loads(self):
|
||||
return len(self.completed_loads)
|
||||
|
||||
def pop_completed_load(self):
|
||||
return self.completed_loads.pop(0)
|
||||
|
||||
def offload(self, transfers):
|
||||
self.queued_offloads.append(list(transfers))
|
||||
return True
|
||||
|
||||
def num_completed_offloads(self):
|
||||
return len(self.completed_offloads)
|
||||
|
||||
def pop_completed_offload(self):
|
||||
return self.completed_offloads.pop(0)
|
||||
|
||||
def reset(self):
|
||||
self.reset_count += 1
|
||||
|
||||
def close(self):
|
||||
self.closed = True
|
||||
|
||||
|
||||
class _MappingRecorder:
|
||||
def __init__(self):
|
||||
self.mapping = []
|
||||
|
||||
def set_full_to_swa_mapping(self, full, swa):
|
||||
self.mapping.append((full.clone(), swa.clone()))
|
||||
|
||||
|
||||
def _cache_for_wrapper(**kwargs):
|
||||
defaults = {
|
||||
"tree_core": SimpleNamespace(enable_external_cache_linker=False),
|
||||
"write_through_threshold": 256,
|
||||
"pp_size": 1,
|
||||
"pp_group": None,
|
||||
}
|
||||
defaults.update(kwargs)
|
||||
return SimpleNamespace(**defaults)
|
||||
|
||||
|
||||
def test_cache_linker_attachment_is_backend_independent():
|
||||
cache = UnifiedRadixCache.__new__(UnifiedRadixCache)
|
||||
cache.tree_core = SimpleNamespace(
|
||||
enable_external_cache_linker=False,
|
||||
write_through_threshold=256,
|
||||
)
|
||||
cache.linker = None
|
||||
linker = _FakeLinker()
|
||||
|
||||
cache.init_cache_linker(linker)
|
||||
|
||||
assert cache.linker.cache_linker is linker
|
||||
assert cache.tree_core.enable_external_cache_linker
|
||||
assert cache.write_through_threshold == 1
|
||||
assert cache.linker.layer_done_counter is linker.layer_done_counter
|
||||
|
||||
|
||||
def test_restorable_prefix_intersects_sparse_rank_results():
|
||||
remote_mask = torch.tensor([0, 0, 1, 0, 0], dtype=torch.int)
|
||||
|
||||
def intersect_remote_mask(mask, op):
|
||||
assert op == torch.distributed.ReduceOp.MIN
|
||||
mask.copy_(torch.minimum(mask, remote_mask))
|
||||
|
||||
cache = _cache_for_wrapper(_all_reduce_attn_groups=intersect_remote_mask)
|
||||
wrapper = UnifiedCacheLinkerWrapper(cache, _FakeLinker())
|
||||
|
||||
hit_pages = wrapper._sync_restorable_prefix([2, 4], num_pages=4, device_hit_pages=0)
|
||||
|
||||
assert hit_pages == 2
|
||||
|
||||
|
||||
def test_async_offload_pins_node_until_completion():
|
||||
class _Component:
|
||||
def build_external_linker_transfer(self, phase, node, keys):
|
||||
assert phase == LinkerTransferPhase.OFFLOAD
|
||||
return PoolTransfer(name=PoolName.KV, keys=["page"])
|
||||
|
||||
linker = _FakeLinker()
|
||||
lock_params = object()
|
||||
locks = []
|
||||
unlocks = []
|
||||
|
||||
def inc_lock_ref(node):
|
||||
locks.append(node)
|
||||
return SimpleNamespace(to_dec_params=lambda: lock_params)
|
||||
|
||||
node_id = 7
|
||||
node = SimpleNamespace(
|
||||
id=node_id,
|
||||
external_cache_stored=False,
|
||||
write_through_pending_id=None,
|
||||
)
|
||||
cache = _cache_for_wrapper(
|
||||
tree_core=SimpleNamespace(
|
||||
enable_external_cache_linker=False,
|
||||
mark_write_through_pending=lambda value: setattr(
|
||||
node, "write_through_pending_id", value
|
||||
),
|
||||
),
|
||||
_components_tuple=(_Component(),),
|
||||
inc_lock_ref=inc_lock_ref,
|
||||
dec_lock_ref=lambda node, params: unlocks.append((node, params)),
|
||||
resolve_node_handle=lambda value: node if value == node_id else None,
|
||||
)
|
||||
wrapper = UnifiedCacheLinkerWrapper(cache, linker)
|
||||
|
||||
wrapper.offload_nodes([node_id])
|
||||
|
||||
assert locks == [node_id]
|
||||
assert node.external_cache_stored
|
||||
assert not unlocks
|
||||
|
||||
linker.completed_offloads.append(False)
|
||||
completed = wrapper.take_completed_offloads(finish_count=1)
|
||||
wrapper.commit_completed_offloads(completed)
|
||||
|
||||
assert not node.external_cache_stored
|
||||
assert unlocks == [(node_id, lock_params)]
|
||||
|
||||
|
||||
def test_async_load_pins_node_until_completion():
|
||||
linker = _FakeLinker()
|
||||
lock_params = object()
|
||||
locks = []
|
||||
unlocks = []
|
||||
|
||||
def inc_lock_ref(node):
|
||||
locks.append(node)
|
||||
return SimpleNamespace(to_dec_params=lambda: lock_params)
|
||||
|
||||
node_id = 7
|
||||
cache = _cache_for_wrapper(
|
||||
inc_lock_ref=inc_lock_ref,
|
||||
dec_lock_ref=lambda node, params: unlocks.append((node, params)),
|
||||
)
|
||||
wrapper = UnifiedCacheLinkerWrapper(cache, linker)
|
||||
|
||||
wrapper._queue_load("rid", node_id, [object()])
|
||||
|
||||
assert locks == [node_id]
|
||||
assert not unlocks
|
||||
|
||||
linker.completed_loads.append(["rid"])
|
||||
wrapper.drain_loads(finish_count=1)
|
||||
|
||||
assert unlocks == [(node_id, lock_params)]
|
||||
|
||||
|
||||
def test_release_request_cancels_queued_load():
|
||||
linker = _FakeLinker()
|
||||
lock_params = object()
|
||||
unlocks = []
|
||||
cache = _cache_for_wrapper(
|
||||
dec_lock_ref=lambda node, params: unlocks.append((node, params))
|
||||
)
|
||||
wrapper = UnifiedCacheLinkerWrapper(cache, linker)
|
||||
wrapper.hit_markers["rid"] = object()
|
||||
wrapper.pending_loads["rid"] = (7, lock_params)
|
||||
linker.queued_loads["rid"] = [object()]
|
||||
|
||||
wrapper.release_request("rid")
|
||||
|
||||
assert wrapper.hit_markers == {}
|
||||
assert wrapper.pending_loads == {}
|
||||
assert "rid" not in linker.queued_loads
|
||||
assert unlocks == [(7, lock_params)]
|
||||
|
||||
|
||||
def test_failed_offload_rolls_back_split_fragments():
|
||||
class _Component:
|
||||
def build_external_linker_transfer(self, phase, node, keys):
|
||||
return PoolTransfer(name=PoolName.KV, keys=["page"])
|
||||
|
||||
linker = _FakeLinker()
|
||||
lock_params = object()
|
||||
unlocks = []
|
||||
child = SimpleNamespace(
|
||||
id=7,
|
||||
external_cache_stored=False,
|
||||
write_through_pending_id=None,
|
||||
)
|
||||
parent = SimpleNamespace(
|
||||
id=8,
|
||||
external_cache_stored=False,
|
||||
write_through_pending_id=None,
|
||||
)
|
||||
nodes = {child.id: child, parent.id: parent}
|
||||
|
||||
def mark_pending(node_id):
|
||||
nodes[node_id].write_through_pending_id = node_id
|
||||
|
||||
cache = _cache_for_wrapper(
|
||||
tree_core=SimpleNamespace(
|
||||
enable_external_cache_linker=False,
|
||||
mark_write_through_pending=mark_pending,
|
||||
),
|
||||
_components_tuple=(_Component(),),
|
||||
inc_lock_ref=lambda node_id: SimpleNamespace(to_dec_params=lambda: lock_params),
|
||||
dec_lock_ref=lambda node_id, params: unlocks.append((node_id, params)),
|
||||
resolve_node_handle=nodes.__getitem__,
|
||||
)
|
||||
wrapper = UnifiedCacheLinkerWrapper(cache, linker)
|
||||
wrapper.offload_nodes([child.id])
|
||||
|
||||
parent.external_cache_stored = child.external_cache_stored
|
||||
parent.write_through_pending_id = child.write_through_pending_id
|
||||
wrapper.replace_pending_offload_node(child.id, child.id, [parent.id, child.id])
|
||||
linker.completed_offloads.append(False)
|
||||
wrapper.commit_completed_offloads(wrapper.take_completed_offloads(finish_count=1))
|
||||
|
||||
assert not parent.external_cache_stored
|
||||
assert not child.external_cache_stored
|
||||
assert parent.write_through_pending_id is None
|
||||
assert child.write_through_pending_id is None
|
||||
assert unlocks == [(child.id, lock_params)]
|
||||
|
||||
|
||||
def test_split_action_retargets_pending_external_offload():
|
||||
calls = []
|
||||
cache = UnifiedRadixCache.__new__(UnifiedRadixCache)
|
||||
cache.linker = SimpleNamespace(
|
||||
replace_pending_offload_node=lambda *args: calls.append(("linker", *args))
|
||||
)
|
||||
cache._replace_pending_write_through_node = lambda *args: calls.append(
|
||||
("hicache", *args)
|
||||
)
|
||||
action = ReplaceWriteThroughOnNodeSplit(
|
||||
ack_id=7,
|
||||
old_node_id=7,
|
||||
new_node_id=8,
|
||||
new_child_node_id=7,
|
||||
)
|
||||
|
||||
cache._apply_cache_action(action)
|
||||
|
||||
assert calls == [
|
||||
("hicache", 7, 7, [8, 7]),
|
||||
("linker", 7, 7, [8, 7]),
|
||||
]
|
||||
|
||||
|
||||
def test_reset_quiesces_backend_before_releasing_pending_locks():
|
||||
class _Component:
|
||||
def build_external_linker_transfer(self, phase, node, keys):
|
||||
return PoolTransfer(name=PoolName.KV, keys=["page"])
|
||||
|
||||
events = []
|
||||
|
||||
class _QuiescentFakeLinker(_FakeLinker):
|
||||
def reset(self):
|
||||
events.append("backend")
|
||||
super().reset()
|
||||
|
||||
linker = _QuiescentFakeLinker()
|
||||
node = SimpleNamespace(
|
||||
id=7,
|
||||
external_cache_stored=False,
|
||||
write_through_pending_id=None,
|
||||
)
|
||||
cache = _cache_for_wrapper(
|
||||
tree_core=SimpleNamespace(
|
||||
enable_external_cache_linker=False,
|
||||
mark_write_through_pending=lambda value: setattr(
|
||||
node, "write_through_pending_id", value
|
||||
),
|
||||
),
|
||||
_components_tuple=(_Component(),),
|
||||
inc_lock_ref=lambda node_id: SimpleNamespace(to_dec_params=object),
|
||||
dec_lock_ref=lambda node_id, params: events.append(("unlock", node_id)),
|
||||
resolve_node_handle=lambda node_id: node,
|
||||
)
|
||||
wrapper = UnifiedCacheLinkerWrapper(cache, linker)
|
||||
wrapper._queue_load("rid", node.id, [object()])
|
||||
wrapper.offload_nodes([node.id])
|
||||
|
||||
wrapper.reset()
|
||||
|
||||
assert events == ["backend", ("unlock", node.id), ("unlock", node.id)]
|
||||
assert wrapper.pending_loads == {}
|
||||
assert wrapper.pending_offloads == []
|
||||
assert not node.external_cache_stored
|
||||
assert node.write_through_pending_id is None
|
||||
|
||||
|
||||
def test_close_quiesces_backend_before_releasing_pending_loads():
|
||||
events = []
|
||||
|
||||
class _ClosingFakeLinker(_FakeLinker):
|
||||
def close(self):
|
||||
events.append("backend")
|
||||
super().close()
|
||||
|
||||
linker = _ClosingFakeLinker()
|
||||
cache = _cache_for_wrapper(
|
||||
dec_lock_ref=lambda node_id, params: events.append(("unlock", node_id))
|
||||
)
|
||||
wrapper = UnifiedCacheLinkerWrapper(cache, linker)
|
||||
wrapper.pending_loads["rid"] = (7, object())
|
||||
|
||||
wrapper.close()
|
||||
|
||||
assert events == ["backend", ("unlock", 7)]
|
||||
assert linker.closed
|
||||
assert wrapper.pending_loads == {}
|
||||
|
||||
|
||||
def test_check_hicache_events_commits_common_rank_results():
|
||||
committed = []
|
||||
cache = UnifiedRadixCache.__new__(UnifiedRadixCache)
|
||||
cache.linker = SimpleNamespace(
|
||||
num_completed_loads=lambda: 1,
|
||||
drain_loads=lambda count: committed.append(("load", count)),
|
||||
num_completed_offloads=lambda: 3,
|
||||
take_completed_offloads=lambda count: [True] * count,
|
||||
commit_completed_offloads=committed.append,
|
||||
)
|
||||
|
||||
reduce_calls = 0
|
||||
|
||||
def reduce_to_common_state(value, op):
|
||||
nonlocal reduce_calls
|
||||
assert op == torch.distributed.ReduceOp.MIN
|
||||
reduce_calls += 1
|
||||
if reduce_calls == 1:
|
||||
value.copy_(torch.tensor([1, 1]))
|
||||
else:
|
||||
value.fill_(0)
|
||||
|
||||
cache._all_reduce_attn_groups = reduce_to_common_state
|
||||
|
||||
cache.check_hicache_events()
|
||||
|
||||
assert committed == [("load", 1), [False]]
|
||||
|
||||
|
||||
def test_component_commit_keeps_only_adopted_pages():
|
||||
mapping = _MappingRecorder()
|
||||
cache = _cache_for_wrapper(
|
||||
page_size=2,
|
||||
token_to_kv_pool_allocator=SimpleNamespace(
|
||||
set_full_to_swa_mapping=mapping.set_full_to_swa_mapping
|
||||
),
|
||||
)
|
||||
wrapper = UnifiedCacheLinkerWrapper(cache, _FakeLinker())
|
||||
full_component = FullComponent.__new__(FullComponent)
|
||||
full_component.cache = cache
|
||||
full_component.component_type = ComponentType.FULL
|
||||
swa_component = SWAComponent.__new__(SWAComponent)
|
||||
swa_component.cache = cache
|
||||
swa_component.component_type = ComponentType.SWA
|
||||
full = PoolTransfer(
|
||||
name=PoolName.KV,
|
||||
keys=["a", "b", "c", "d"],
|
||||
device_indices=torch.tensor([100, 101, 102, 103, 104, 105, 106, 107]),
|
||||
)
|
||||
canonical_tail = torch.tensor([10, 11, 102, 103, 14, 15, 106, 107])
|
||||
swa = PoolTransfer(
|
||||
name=PoolName.SWA,
|
||||
keys=["a", "b", "c", "d"],
|
||||
device_indices=torch.tensor([200, 201, 202, 203, 204, 205, 206, 207]),
|
||||
)
|
||||
insert_result = InsertResult(
|
||||
prefix_len=0,
|
||||
adopted_ranges={
|
||||
ComponentType.FULL: [(2, 4), (6, 8)],
|
||||
ComponentType.SWA: [(2, 4), (6, 8)],
|
||||
},
|
||||
)
|
||||
|
||||
filtered = wrapper._update_load(
|
||||
ExternalLinkerLoadPhase.COMMIT,
|
||||
SimpleNamespace(),
|
||||
[(full_component, full), (swa_component, swa)],
|
||||
prefix_len=8,
|
||||
insert_result=insert_result,
|
||||
canonical_full=canonical_tail,
|
||||
)
|
||||
|
||||
assert filtered == [full, swa]
|
||||
assert full.keys == ["b", "d"]
|
||||
assert full.device_indices.tolist() == [102, 103, 106, 107]
|
||||
assert swa.keys == ["b", "d"]
|
||||
assert swa.device_indices.tolist() == [202, 203, 206, 207]
|
||||
mapped_full, mapped_swa = mapping.mapping[0]
|
||||
assert mapped_full.tolist() == [102, 103, 106, 107]
|
||||
assert mapped_swa.tolist() == [202, 203, 206, 207]
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(pytest.main([__file__, "-v"]))
|
||||
Reference in New Issue
Block a user