diff --git a/python/sglang/srt/arg_groups/kv_cache_hook.py b/python/sglang/srt/arg_groups/kv_cache_hook.py index a5780bf3a9e1..85b772b3feae 100644 --- a/python/sglang/srt/arg_groups/kv_cache_hook.py +++ b/python/sglang/srt/arg_groups/kv_cache_hook.py @@ -530,18 +530,10 @@ def handle_unified_memory_pool(server_args: Any) -> None: ) if cfg.disaggregation_decode_retraction_backup == "host_pool": model_config = model_config_of(server_args) - assert not cfg.disaggregation_decode_enable_radix_cache, ( - "--enable-unified-memory host-pool decode retraction does not " - "support decode radix-cache H2D/D2H transfers yet." - ) assert mambaish_config(model_config) is None, ( "--enable-unified-memory host-pool decode retraction does not " "support hybrid-Mamba models." ) - assert not model_config.is_hybrid_swa, ( - "--enable-unified-memory host-pool decode retraction does not " - "support hybrid-SWA H2D/D2H transfers yet." - ) assert cfg.speculative_algorithm in (None, "DSPARK"), ( "--enable-unified-memory only supports --speculative-algorithm " "DSPARK (chain draft); other speculative algorithms are not yet " @@ -578,6 +570,21 @@ def handle_unified_memory_pool(server_args: Any) -> None: "the LMCache offload path indexes the device buffers with the ids it " "is handed, and under the unified pool those are VIRTUAL." ) + assert not cfg.enable_unified_cache_external_linker, ( + "--enable-unified-memory does not support " + "--enable-unified-cache-external-linker: direct L3 transfers do not " + "preserve unified page-envelope indices and compaction lifetimes. " + "Use --enable-hierarchical-cache for supported L2/L3 transfers." + ) + if cfg.enable_hierarchical_cache: + assert cfg.pp_size == 1, ( + "--enable-unified-memory with hierarchical cache does not support " + "pipeline parallelism (--pp-size > 1)." + ) + assert not envs.SGLANG_DISABLE_LAZY_COMPACTION.get(), ( + "--enable-unified-memory with hierarchical cache requires lazy " + "compaction so pending H2D physical reservations remain stable." + ) if cfg.dcp_size > 1: _validate_unified_memory_dcp(server_args) # Prefill cuda-graph capture IS wired for the unified pool: the captured diff --git a/python/sglang/srt/disaggregation/decode.py b/python/sglang/srt/disaggregation/decode.py index 22c1e6e1bc70..f3956f561a5e 100644 --- a/python/sglang/srt/disaggregation/decode.py +++ b/python/sglang/srt/disaggregation/decode.py @@ -787,12 +787,19 @@ def _match_prefix_and_lock(self, req: Req) -> DecodePrefixMatch: Match a request against the decode-side radix cache, lock the matched node to prevent eviction, and return the matched prefix information. """ + max_prefix_len = None + if self._uses_swa_tail_prealloc(): + fill_len = self._pre_alloc_fill_len(req) + max_prefix_len = fill_len - self._swa_tail_len(fill_len) + # Match and lock only reusable FULL KV. The entire SWA tail must be + # freshly allocated, including when the prefix comes from L2/L3. result = match_prefix_for_req( self.tree_cache, req, req.origin_input_ids, cow_mamba=self.tree_cache.supports_mamba(), include_req=True, + max_prefix_len=max_prefix_len, ) # Keep aggregated scheduling semantics while preserving the SWA lock # boundary needed for the matching dec_lock_ref; the full receipt @@ -800,7 +807,9 @@ def _match_prefix_and_lock(self, req: Req) -> DecodePrefixMatch: req.lock_receipt = self.tree_cache.inc_lock_ref( result.last_device_node ).to_dec_params() - return self._build_decode_prefix_match(req, result) + return self._build_decode_prefix_match( + req, result, max_prefix_len=max_prefix_len + ) def _resolve_prefill_dp_rank(self, req: Req) -> Optional[int]: prefill_info = self.kv_manager.prefill_info_table.get(_bootstrap_addr(req)) @@ -1422,20 +1431,6 @@ def pop_preallocated( fill_len = self._pre_alloc_fill_len(decode_req.req) - # Cap full-attention prefix reuse at the sliding-window start so - # the SWA window lands entirely in the fresh delta, keeping - # alloc_extend_swa_tail's tail->full mapping in range. Costs reuse - # of only the last ~window_size full-attention tokens. - if uses_swa_tail_prealloc and prefix_len > 0: - swa_prefix_cap = fill_len - self._swa_tail_len(fill_len) - if prefix_len > swa_prefix_cap: - prefix_len = swa_prefix_cap - prefix_indices = prefix_indices[:prefix_len] - # Cap the prefill-committed prefix too: tokens past the - # cap are not device-resident, so prefill must transfer - # them. - total_prefix_len = prefix_len - # Decode transfers the SWA tail fresh, so retain only the # full-attention prefix lock needed for reuse. if ( @@ -2317,9 +2312,9 @@ def alloc_for_decode_prealloc( # the live window tail. kv_loc = allocator.alloc_extend_swa_tail( prefix_lens=torch.tensor( - [prefix_len], dtype=torch.int64, device=device + [total_prefix_len], dtype=torch.int64, device=device ), - prefix_lens_cpu=torch.tensor([prefix_len], dtype=torch.int64), + prefix_lens_cpu=torch.tensor([total_prefix_len], dtype=torch.int64), seq_lens=torch.tensor([fill_len], dtype=torch.int64, device=device), seq_lens_cpu=torch.tensor([fill_len], dtype=torch.int64), last_loc=last_loc, diff --git a/python/sglang/srt/disaggregation/decode_hicache_mixin.py b/python/sglang/srt/disaggregation/decode_hicache_mixin.py index 7d75dfa79f86..4127f685450d 100644 --- a/python/sglang/srt/disaggregation/decode_hicache_mixin.py +++ b/python/sglang/srt/disaggregation/decode_hicache_mixin.py @@ -61,7 +61,9 @@ class HiCacheRestoreResult(Enum): class DecodeHiCachePreallocMixin: """HiCache hooks for ``DecodePreallocQueue``: issue prefetch + reserve tokens.""" - def _build_decode_prefix_match(self, req: Req, result: Any) -> DecodePrefixMatch: + def _build_decode_prefix_match( + self, req: Req, result: Any, *, max_prefix_len: Optional[int] = None + ) -> DecodePrefixMatch: """Convert a ``match_prefix_for_req`` result into ``DecodePrefixMatch``. Performs the optional L3 storage hit length query when decode-side @@ -79,7 +81,7 @@ def _build_decode_prefix_match(self, req: Req, result: Any) -> DecodePrefixMatch last_host_node ): matched_len = l1_prefix_len + l2_host_hit_length - suffix_tokens = req.origin_input_ids[matched_len:] + suffix_tokens = req.origin_input_ids[matched_len:max_prefix_len] last_hash = self.tree_cache.get_last_hash_value(last_host_node) prefix_keys = ( self.tree_cache.get_prefix_hash_values(last_host_node) @@ -221,6 +223,7 @@ def _try_hicache_queue_load_back(self, dr: DecodeRequest) -> bool: dr.req.origin_input_ids, cow_mamba=False, include_req=True, + max_prefix_len=pm.decode_prefix_len, ) new_indices, restored_node = self.tree_cache.init_load_back( InitLoadBackParams( diff --git a/python/sglang/srt/managers/cache_controller.py b/python/sglang/srt/managers/cache_controller.py index 8261cd60fdd2..a05c7623dd5e 100644 --- a/python/sglang/srt/managers/cache_controller.py +++ b/python/sglang/srt/managers/cache_controller.py @@ -43,7 +43,7 @@ from sglang.srt.mem_cache.l2_transfer import L2Transfer, L2TransferEngine from sglang.srt.mem_cache.memory_pool import MLATokenToKVPool from sglang.srt.mem_cache.utils import get_storage_hash_str -from sglang.srt.runtime_context import get_parallel +from sglang.srt.runtime_context import get_memory, get_parallel from sglang.srt.utils import get_device_module logger = logging.getLogger(__name__) @@ -842,6 +842,10 @@ def start_writing(self) -> None: self._l2_transfers(host_indices, device_indices, pool_transfers) ) + self.mem_pool_device_allocator.set_hicache_transfer_done_event( + (id(self), "write"), completion.finish_event + ) + self.ack_write_queue.append( HiCacheAck( start_event=completion.start_event, @@ -986,6 +990,10 @@ def start_loading(self) -> int: transfer_layer_id_max=self.transfer_layer_id_max, ) + self.mem_pool_device_allocator.set_hicache_transfer_done_event( + (id(self), "load"), completion.finish_event + ) + self.ack_load_queue.append( HiCacheAck( start_event=completion.start_event, @@ -1108,13 +1116,22 @@ def _page_transfer(self, operation: PrefetchOperation) -> int: # Get one batch token, and update the completed_tokens if succeed extra_info = HiCacheStorageExtraInfo(prefix_keys=prefix_keys) - hit_pages = self._page_transfer_kv_batch( - operation, - batch_hashes, - batch_host_indices, - extra_info, - kv_derived_transfers, - ) + try: + hit_pages = self._page_transfer_kv_batch( + operation, + batch_hashes, + batch_host_indices, + extra_info, + kv_derived_transfers, + ) + except Exception: + if not get_memory().enable_unified_memory: + raise + logger.exception( + "HiCache prefetch transfer failed for request %s", + operation.request_id, + ) + hit_pages = 0 # Check termination if hit_pages != len(batch_hashes): all_success = False @@ -1175,10 +1192,13 @@ def prefetch_io_aux_func(self): while not self.storage_stop_event.is_set(): try: operation = self.prefetch_buffer.get(block=True, timeout=1) - if operation is None: - continue + except Empty: + continue + if operation is None: + continue + try: self._page_transfer(operation) - + finally: self.prefetch_sync_queue.put( PrefetchAck( rid=operation.request_id, @@ -1186,8 +1206,6 @@ def prefetch_io_aux_func(self): operation=operation, ) ) - except Empty: - continue def prefetch_rate_limited(self) -> bool: """ @@ -1209,6 +1227,24 @@ def prefetch_rate_limited(self) -> bool: # todo: more sophisticated rate limiting based on storage backend performance return False + def alloc_prefetch_host_buffers( + self, operation: StorageOperation, need_size: int + ) -> Optional[torch.Tensor]: + """Allocate the host bounce for a storage hit.""" + return self.mem_pool_host.alloc(need_size) + + def can_fit_prefetch_host_buffers( + self, operation: StorageOperation, need_size: int + ) -> bool: + """Whether a prefetch bounce can fit when its host pools are empty.""" + return need_size <= self.mem_pool_host.size + + def free_prefetch_host_buffers( + self, operation: StorageOperation, host_indices: torch.Tensor + ) -> None: + """Roll back a hit-sized host bounce before transfer ownership moves.""" + self.mem_pool_host.free(host_indices) + def _storage_hit_query(self, operation) -> tuple[list[str], int]: last_hash = operation.last_hash tokens_to_fetch = operation.token_ids @@ -1243,10 +1279,22 @@ def prefetch_thread_func(self): operation = self.prefetch_queue.get(block=True, timeout=1) if operation is None: continue - if operation.is_terminated(): + try: + if operation.is_terminated(): + hash_value, storage_hit_count = [], 0 + else: + hash_value, storage_hit_count = self._storage_hit_query( + operation + ) + except Exception: + if not get_memory().enable_unified_memory: + raise + logger.exception( + "HiCache storage query failed for request %s", + operation.request_id, + ) hash_value, storage_hit_count = [], 0 - else: - hash_value, storage_hit_count = self._storage_hit_query(operation) + storage_hit_count_tensor = torch.tensor( storage_hit_count, dtype=torch.int ) diff --git a/python/sglang/srt/managers/schedule_policy.py b/python/sglang/srt/managers/schedule_policy.py index b508eacb13ff..e7d10082d090 100644 --- a/python/sglang/srt/managers/schedule_policy.py +++ b/python/sglang/srt/managers/schedule_policy.py @@ -33,7 +33,7 @@ import os import random from collections import Counter -from contextlib import contextmanager +from contextlib import contextmanager, nullcontext from dataclasses import dataclass from enum import Enum, auto from functools import lru_cache @@ -68,6 +68,11 @@ zero_match_result, ) from sglang.srt.mem_cache.radix_cache import RadixCache, RadixKey, TreeNode +from sglang.srt.mem_cache.unified_cache.components import ( + CacheTransferPhase, + ComponentType, +) +from sglang.srt.mem_cache.unified_radix_cache import UnifiedRadixCache if TYPE_CHECKING: from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator @@ -161,6 +166,7 @@ def match_prefix_for_req( *, cow_mamba: bool = False, include_req: bool = False, + max_prefix_len: Optional[int] = None, ): if token_ids is None: token_ids = req.origin_input_ids + req.output_ids @@ -171,6 +177,10 @@ def match_prefix_for_req( # this request's SWA ring. No-op for other layouts. reprefill_tail = tree_cache.swa_reprefill_tail_tokens() key_limit = max(0, len(token_ids) - reprefill_tail) if reprefill_tail else None + if max_prefix_len is not None: + key_limit = ( + max_prefix_len if key_limit is None else min(key_limit, max_prefix_len) + ) match_result = tree_cache.match_prefix( MatchPrefixParams( @@ -1188,14 +1198,23 @@ def add_chunked_req(self, req: Req): return req if truncated else None @contextmanager - def _lock_node(self, last_node: TreeNode): - # Replay the acquire's receipt (SWA boundary uuid, mamba flag) so the - # release takes back exactly what this temporary lock took. - dec_lock_params = self.tree_cache.inc_lock_ref(last_node).to_dec_params() + def _lock_node(self, last_node: TreeNode, *, lock_host: bool = False): + host_lock_params = ( + self.tree_cache.inc_host_lock_ref(last_node).to_dec_params() + if lock_host + else None + ) try: - yield None + # Replay the acquire's receipt (SWA boundary uuid, mamba flag) so the + # release takes back exactly what this temporary lock took. + dec_lock_params = self.tree_cache.inc_lock_ref(last_node).to_dec_params() + try: + yield None + finally: + self.tree_cache.dec_lock_ref(last_node, dec_lock_params) finally: - self.tree_cache.dec_lock_ref(last_node, dec_lock_params) + if host_lock_params is not None: + self.tree_cache.dec_host_lock_ref(last_node, host_lock_params) def add_one_req_ignore_eos(self, req: Req): cand_extend_input_len = len(req.full_untruncated_fill_ids) - len( @@ -1398,14 +1417,55 @@ def add_one_req( return AddReqResult.OTHER if req.needs_host_load_back(): - promised_host_hit = req.host_hit_length - loaded = self.tree_cache.init_load_back( - InitLoadBackParams( - best_match_node=req.best_match_node, - host_hit_length=req.host_hit_length, - req=req, + load_max_new = min(max_new, admission.max_new_tokens) + # Reclaim can write back device victims and evict host leaves. + # Pin the selected host/aux match until load-back owns its locks. + with ( + self._lock_node(req.best_match_node, lock_host=True) + if isinstance(self.tree_cache, UnifiedRadixCache) + else nullcontext() + ): + full_load_tokens = req.host_hit_length + if ( + isinstance(self.tree_cache, UnifiedRadixCache) + and self.tree_cache.buffer_pipeline is None + and not ( + self.tree_cache.linker is not None + and self.tree_cache.linker.has_hit(req.rid) + ) + ): + # Host hits can include resident FULL behind host-only aux. + # Reuse the FULL transfer spec to count only new slots. + full_transfer = ( + self.tree_cache.tree_core.build_hicache_transfers( + ComponentType.FULL, + req.best_match_node, + CacheTransferPhase.LOAD_BACK, + )[0] + ) + full_load_tokens = len(full_transfer.host_indices) + if not self.memory_budget.prepare_load_back( + full_tokens=( + full_load_tokens + + admission.extend_len + + load_max_new + + self.page_size + + mamba_gap_reserve + ), + extend_input_len=admission.extend_len, + max_new_tokens=load_max_new, + swa_host_hit_length=req.swa_host_hit_length, + chunk_limit=self.rem_chunk_tokens, + ): + return AddReqResult.NO_TOKEN + promised_host_hit = req.host_hit_length + loaded = self.tree_cache.init_load_back( + InitLoadBackParams( + best_match_node=req.best_match_node, + host_hit_length=req.host_hit_length, + req=req, + ) ) - ) if loaded is None: return AddReqResult.OTHER new_indices, req.last_node = loaded diff --git a/python/sglang/srt/mem_cache/allocator/base.py b/python/sglang/srt/mem_cache/allocator/base.py index 898973ee8961..ab3831bb4359 100644 --- a/python/sglang/srt/mem_cache/allocator/base.py +++ b/python/sglang/srt/mem_cache/allocator/base.py @@ -16,7 +16,7 @@ from __future__ import annotations import abc -from typing import TYPE_CHECKING, Protocol +from typing import TYPE_CHECKING, Hashable, Protocol import torch @@ -228,6 +228,10 @@ def load_cpu_copy( ): raise NotImplementedError() + def set_hicache_transfer_done_event(self, transfer_key: Hashable, event) -> None: + """Record an asynchronous HiCache transfer completion event if needed.""" + return + def alloc_extend(self, *args, **kwargs): raise NotImplementedError("alloc_extend is only for paged allocator") diff --git a/python/sglang/srt/mem_cache/allocator/unified_hybrid_swa.py b/python/sglang/srt/mem_cache/allocator/unified_hybrid_swa.py index 069ab94c71be..2133496c1f12 100644 --- a/python/sglang/srt/mem_cache/allocator/unified_hybrid_swa.py +++ b/python/sglang/srt/mem_cache/allocator/unified_hybrid_swa.py @@ -19,7 +19,7 @@ import logging import math from abc import abstractmethod -from typing import Callable, List, Optional, Sequence, Tuple +from typing import Callable, Hashable, List, Optional, Sequence, Tuple import torch from torch.profiler import record_function @@ -502,11 +502,11 @@ def alloc_extend_swa_tail( The static composite allocates the two sides independently and records a full->swa index mapping. That is not representable here: the two sides SHARE one virtual id space (a virtual page names a full-physical - page and, if bound, a swa-physical one), which is why - `set_full_to_swa_mapping` is a no-op on this allocator and - `translate_loc_from_full_to_swa` derives the swa id from the virtual id - instead of a table. Running the static body would call `alloc_extend` - on the swa sub-allocator, which asserts it is not the id owner. + page and, if bound, a swa-physical one). Load-back installs that binding + through `set_full_to_swa_mapping`, and `translate_loc_from_full_to_swa` + resolves the SWA kernel-facing id from the shared virtual id. Running the + static body would call `alloc_extend` on the swa sub-allocator, which + asserts it is not the id owner. The tail is expressed by binding swa for the TAIL's virtual pages only. A new page left unbound has no swa-physical page, which reads as the @@ -678,8 +678,12 @@ def free_full_segment(self, free_index: torch.Tensor, *, start_pos: int) -> None def set_full_to_swa_mapping( self, full_indices: torch.Tensor, swa_indices: torch.Tensor ) -> None: - """Binding load-back rows already updates the shared SWA v2p mapping.""" - return + if full_indices.numel() == 0: + return + assert full_indices.numel() == swa_indices.numel() + full_pages = full_indices.to(torch.int64) // self.page_size + swa_pages = swa_indices.to(torch.int64) // self.page_size + self.swa_attn_allocator.bind(full_pages, swa_pages) def clear_full_to_swa_mapping(self, full_indices: torch.Tensor) -> None: # Paired with set_full_to_swa_mapping: shared mode has no mapping tensor. @@ -797,6 +801,10 @@ def _flush_targets(self): ... @abstractmethod def _ask_float_for_room(self, need_tokens: int) -> None: ... + def set_hicache_transfer_done_event(self, transfer_key: Hashable, event) -> None: + self.full_attn_allocator.set_hicache_transfer_done_event(transfer_key, event) + self.swa_attn_allocator.set_hicache_transfer_done_event(transfer_key, event) + class UnifiedSWATokenToKVPoolAllocator(UnifiedSWAAllocatorBase): """Two-ended FULL/SWA allocator with asymmetric shared-byte reservations.""" @@ -1204,7 +1212,10 @@ def evict_to_free_tokens( return full_reclaim, swa_reclaim = reclaim_plan if full_reclaim or swa_reclaim: - tree_cache.evict_for_alloc( + # The shared-byte plan returns cumulative eviction quotas. + # Per-component capacity targets can count the same shared bytes + # independently and stop before the joint allocation fits. + tree_cache.evict( EvictParams(num_tokens=full_reclaim, swa_num_tokens=swa_reclaim) ) # A zero-reclaim plan can still depend on compaction before allocation. diff --git a/python/sglang/srt/mem_cache/allocator/unified_mamba.py b/python/sglang/srt/mem_cache/allocator/unified_mamba.py index 71c25cabc668..37eb82877215 100644 --- a/python/sglang/srt/mem_cache/allocator/unified_mamba.py +++ b/python/sglang/srt/mem_cache/allocator/unified_mamba.py @@ -17,7 +17,7 @@ from __future__ import annotations import logging -from typing import Callable, List, Optional, Sequence +from typing import Callable, Hashable, List, Optional, Sequence import torch from torch.profiler import record_function @@ -443,3 +443,7 @@ def flush_opportunistic(self) -> int: ): return 0 return fa.flush_opportunistic() + ma.flush_opportunistic() + + def set_hicache_transfer_done_event(self, transfer_key: Hashable, event) -> None: + self.full_attn_allocator.set_hicache_transfer_done_event(transfer_key, event) + self.mamba_allocator.set_hicache_transfer_done_event(transfer_key, event) diff --git a/python/sglang/srt/mem_cache/allocator/unified_sub_pool.py b/python/sglang/srt/mem_cache/allocator/unified_sub_pool.py index f24ee2ad0f93..45c6404b17ae 100644 --- a/python/sglang/srt/mem_cache/allocator/unified_sub_pool.py +++ b/python/sglang/srt/mem_cache/allocator/unified_sub_pool.py @@ -29,6 +29,7 @@ Callable, Dict, Generic, + Hashable, List, Optional, Sequence, @@ -419,6 +420,12 @@ def __init__( # write-race check. Single slot: at most ONE forward in flight per call # site; `_flush` materializes the write-set lazily, avoiding a sync here. self._inflight_forward: Optional[Tuple[torch.cuda.Event, torch.Tensor]] = None + self._hicache_transfer_done_events: Dict[Hashable, torch.cuda.Event] = {} + # HiCache reserves SWA physical pages during load-back, then binds them + # into the tree before the H2D copy is submitted. Compaction must not + # move those pages in that interval because the queued PoolTransfer still + # carries the original physical ids. + self._pending_hicache_load_pages = 0 # Per-call move cap on NON-urgent `_flush`: bounds work per `on_idle()` so # a large backlog doesn't block ZMQ IPC. Urgent retries are uncapped. @@ -521,6 +528,8 @@ def clear(self) -> None: self.live_page_count = 0 self._inflight_forward = None self._latest_forward_done_event = None + self._hicache_transfer_done_events.clear() + self._pending_hicache_load_pages = 0 def clear_inverse_history(self) -> None: self._inverse_history.clear() @@ -726,6 +735,9 @@ def _peer_drainable_hole_bytes(self) -> int: def moves_blocked(self) -> bool: """Whether any installed gate currently forbids relocating pages.""" + if self._pending_hicache_load_pages > 0: + # Physical H2D reservations must stay put even before queueing. + return True for gate in (self.disagg_move_gate, self.host_transfer_move_gate): if gate is not None and not gate(): return True @@ -926,6 +938,8 @@ def _alloc_bind_fast_or_slow( with record_function("MultiEndedAlloc._alloc_bind_fast_or_slow"): if N == 0: return torch.empty(0, dtype=torch.int64, device=self.device) + if self.lazy_compaction and self._free_phys_pages.numel() > 0: + self._wait_hicache_transfers() # FAST PATH: eager, or lazy with no current holes. if not self.lazy_compaction or self._free_phys_pages.numel() == 0: @@ -1082,19 +1096,15 @@ def alloc(self, need_size: int) -> Optional[torch.Tensor]: if not _relieve_for_alloc(self, need_size): return None num_pages = need_size // self.page_size + if num_pages > len(self.free_virtual_ids): + return None v_pages = self.free_virtual_ids[:num_pages] self.free_virtual_ids = self.free_virtual_ids[num_pages:] phys_pages = self._alloc_bind_fast_or_slow(v_pages, num_pages) if phys_pages is None: self.free_virtual_ids = torch.cat([v_pages, self.free_virtual_ids]) return None - if self.page_size == 1: - return v_pages # v_pages already IS the token id list - # Expand page ids to token ids: (P, 1) * S + (S,) -> (P, S) -> (P*S,). - return ( - v_pages[:, None] * self.page_size - + torch.arange(self.page_size, device=self.device) - ).reshape(-1) + return self._expand_pages_to_tokens(v_pages) def alloc_with_virtual(self, virtual_pages: torch.Tensor) -> None: """Take physical PAGES for caller-supplied virtual PAGE ids (not token @@ -1264,6 +1274,7 @@ def free( self._free_lazy(free_index, pages=_pages) return # --- EAGER path --- + self._wait_hicache_transfers() # Near-no-op in normal mode (sampling's CPU sync already drained # forward_stream); in overlap mode it does the serializing. if self.forward_stream is not None: @@ -1713,6 +1724,7 @@ def _flush(self, *, urgent: bool) -> int: if self.moves_blocked(): # Holes stay in the free list; the next flush picks them up. return 0 + self._wait_hicache_transfers() self._stats_n_flush_calls += 1 with record_function("MultiEndedAlloc._flush"): self._drain_pending_reuse(urgent=urgent) @@ -1965,6 +1977,92 @@ def _set_capacity( self.size = self.max_slots self.clear() + _pending_hicache_load_pages: _CapacityField[int] = _CapacityField() + + def set_hicache_transfer_done_event(self, transfer_key: Hashable, event) -> None: + self._hicache_transfer_done_events[transfer_key] = event + if transfer_key == "load" or ( + isinstance(transfer_key, tuple) and transfer_key[-1] == "load" + ): + # The H2D copy is now ordered before this event. Future compaction + # may proceed after _wait_hicache_transfers() waits on it. + self._pending_hicache_load_pages = 0 + + def _wait_hicache_transfers(self) -> None: + if not self._hicache_transfer_done_events: + return + current_stream = torch.cuda.current_stream() + events = tuple(self._hicache_transfer_done_events.values()) + self._hicache_transfer_done_events.clear() + for event in events: + current_stream.wait_event(event) + + def _expand_pages_to_tokens(self, pages: torch.Tensor) -> torch.Tensor: + if self.page_size == 1: + return pages + return ( + pages[:, None] * self.page_size + + torch.arange(self.page_size, device=self.device) + ).reshape(-1) + + def alloc_physical(self, need_size: int) -> Optional[torch.Tensor]: + """Reserve physical token slots without assigning virtual page ids.""" + if need_size <= 0: + return torch.empty(0, dtype=torch.int64, device=self.device) + assert need_size % self.page_size == 0, ( + f"MultiEndedAllocator({self.sub_pool_name!r}).alloc_physical: need_size=" + f"{need_size} must be a multiple of page_size={self.page_size}" + ) + if need_size > self.available_size(): + if not _relieve_for_alloc(self, need_size): + return None + if self.lazy_compaction and self._free_phys_pages.numel() > 0: + self._wait_hicache_transfers() + physical_pages = self.take_physical_pages(need_size // self.page_size) + if physical_pages is None: + return None + if self.lazy_compaction: + self._pending_hicache_load_pages += int(physical_pages.shape[0]) + return self._expand_pages_to_tokens(physical_pages) + + def cancel_physical_reservation(self, free_index: torch.Tensor) -> None: + """Roll back a HiCache physical allocation before its H2D is submitted.""" + if free_index is None or free_index.numel() == 0: + return + if self.lazy_compaction: + num_pages = free_index.numel() // self.page_size + assert num_pages <= self._pending_hicache_load_pages, ( + f"MultiEndedAllocator({self.sub_pool_name!r}) released {num_pages} " + f"HiCache pages with only {self._pending_hicache_load_pages} pending" + ) + self._pending_hicache_load_pages -= num_pages + self.free_physical(free_index) + + def free_physical(self, free_index: torch.Tensor) -> None: + """Release physical token slots reserved by :meth:`alloc_physical`.""" + if free_index is None or free_index.numel() == 0: + return + assert free_index.numel() % self.page_size == 0, ( + f"MultiEndedAllocator({self.sub_pool_name!r}).free_physical: " + f"{free_index.numel()} tokens must contain whole pages of " + f"size {self.page_size}" + ) + physical_pages = ( + free_index.detach().to(torch.int64)[:: self.page_size] // self.page_size + ) + self._wait_hicache_transfers() + if self.forward_stream is not None and not self.lazy_compaction: + torch.cuda.current_stream().wait_stream(self.forward_stream) + virtual_pages = self.physical_to_virtual[physical_pages] + bound_mask = virtual_pages >= 0 + self.virtual_to_physical.index_fill_(0, virtual_pages[bound_mask], -1) + self.physical_to_virtual.index_fill_(0, physical_pages, -1) + if self.lazy_compaction: + self._free_phys_pages = torch.cat([self._free_phys_pages, physical_pages]) + self.live_page_count -= int(physical_pages.shape[0]) + return + self._compact_pending(physical_pages) + def _chain_byte_accounting_violations( chain: List[MultiEndedAllocator], @@ -2258,6 +2356,23 @@ def free( self._holes_dirty = True self._park_if_empty() + def free_physical(self, free_index: torch.Tensor) -> None: + """Release physical reservations using the float's hole/span bookkeeping.""" + if free_index is None or free_index.numel() == 0: + return + assert free_index.numel() % self.page_size == 0 + physical_pages = ( + free_index.detach().to(torch.int64)[:: self.page_size] // self.page_size + ) + self._wait_hicache_transfers() + virtual_pages = self.physical_to_virtual[physical_pages] + bound_pages = virtual_pages[virtual_pages >= 0] + self.virtual_to_physical.index_fill_(0, bound_pages, -1) + self.physical_to_virtual.index_fill_(0, physical_pages, -1) + self._free_phys_pages = torch.cat([self._free_phys_pages, physical_pages]) + self._holes_dirty = True + self._park_if_empty() + def _park_if_empty(self) -> bool: """Reset the span and go frontier-transparent once no live page remains. Sync-free: `_live_pages()` is span minus hole COUNT, both diff --git a/python/sglang/srt/mem_cache/buffer_mode/pipeline.py b/python/sglang/srt/mem_cache/buffer_mode/pipeline.py index 71f1de33c02d..f8870919a44a 100644 --- a/python/sglang/srt/mem_cache/buffer_mode/pipeline.py +++ b/python/sglang/srt/mem_cache/buffer_mode/pipeline.py @@ -28,6 +28,7 @@ import logging from array import array from collections import deque +from dataclasses import replace from typing import TYPE_CHECKING, Optional import msgspec @@ -49,6 +50,7 @@ PoolTransfer, SidecarPoolSpec, ) +from sglang.srt.mem_cache.pool_host.base import HostKVCache from sglang.srt.mem_cache.radix_cache import RadixKey from sglang.srt.mem_cache.storage_prefetch import StagedPrefetchPlan from sglang.srt.mem_cache.unified_cache.cache_action import RebuildFullToSWAMapping @@ -71,6 +73,13 @@ logger = logging.getLogger(__name__) +def _minimum_full_transfer_tokens( + host_pool: HostKVCache, prefetch_threshold: int +) -> int: + page_size = host_pool.page_size + return max(page_size, -(-prefetch_threshold // page_size) * page_size) + + class _UnifiedBackupIntent(msgspec.Struct): """Buffer-mode backup intent, unpinned while queued. @@ -95,6 +104,7 @@ class _UnifiedBufferBackupEntry(msgspec.Struct): host_indices: torch.Tensor aux_xfers: list[PoolTransfer] lock_params: DecLockRefParams + occupied_units: int class _StagedPrefetch(msgspec.Struct): @@ -159,6 +169,7 @@ def validate_buffer_only_stack( sidecar_pool_specs: list[SidecarPoolSpec], host_pool_group: HostPoolGroup, swa_component: Optional[SWAComponent], + storage_prefetch_threshold: int = 256, ) -> None: """Post-assembly buffer-mode fences. @@ -201,12 +212,40 @@ def validate_buffer_only_stack( # window), so every window-carrying intent would be dropped as # oversize and SWA storage coverage would silently be zero. window_tokens = swa.full_window_pages * swa._swa_kv_pool_host.page_size - if swa._swa_kv_pool_host.size < 2 * window_tokens: + shared_domain = ( + swa._swa_kv_pool_host.shared_allocation_domain + if isinstance(swa._swa_kv_pool_host, HostKVCache) + else None + ) + if shared_domain is not None: + full_host_pool = ( + swa.cache.cache_controller.mem_pool_host.anchor_entry.host_pool + ) + min_full_tokens = _minimum_full_transfer_tokens( + full_host_pool, storage_prefetch_threshold + ) + one_transfer_bytes = ( + window_tokens * swa._swa_kv_pool_host.size_per_token + + min_full_tokens * full_host_pool.size_per_token + ) + one_transfer = [ + (full_host_pool.pool_label, min_full_tokens), + (swa._swa_kv_pool_host.pool_label, window_tokens), + ] + enough_capacity = shared_domain.can_fit_many_then( + one_transfer, one_transfer, empty=True + ) + capacity = f"{shared_domain.capacity_bytes} shared bytes" + requirement = f"{2 * one_transfer_bytes} shared bytes" + else: + enough_capacity = swa._swa_kv_pool_host.size >= 2 * window_tokens + capacity = f"{swa._swa_kv_pool_host.size} SWA tokens" + requirement = f"{2 * window_tokens} SWA tokens" + if not enough_capacity: raise ValueError( - "--hicache-host-memory-mode buffer_only requires an SWA " - f"host pool of at least two trailing windows " - f"({2 * window_tokens} tokens; got " - f"{swa._swa_kv_pool_host.size}): one staging a write " + "--hicache-host-memory-mode buffer_only requires a host arena " + "large enough for two minimum Full/SWA transfers " + f"({requirement}; got {capacity}): one staging a write " "while one stays reserved for prefetch window allocs." ) @@ -300,8 +339,89 @@ def reset(self) -> None: self.anchor_locked_tokens_ = 0 self._anchor_lock_cap_skips = 0 + def _shared_host_domain(self): + cc = self._cache.cache_controller + anchor = cc.mem_pool_host.anchor_entry.host_pool + if not isinstance(anchor, HostKVCache): + return None + return anchor.shared_allocation_domain + + def _transfer_tokens(self, transfer: PoolTransfer) -> int: + if transfer.host_indices is not None: + return len(transfer.host_indices) + if transfer.keys is None: + return 0 + entry = self._cache.cache_controller.mem_pool_host.entry_map.get(transfer.name) + return 0 if entry is None else len(transfer.keys) * entry.host_pool.page_size + + def host_allocation_units( + self, + host_indices: Optional[torch.Tensor], + aux_xfers: Optional[list[PoolTransfer]], + ) -> int: + """Host usage expressed in anchor-token units for scheduler accounting.""" + if host_indices is None: + return 0 + return self._host_request_units(len(host_indices), aux_xfers) + + def _shared_host_requests( + self, kv_tokens: int, aux_xfers: Optional[list[PoolTransfer]] + ) -> Optional[list[tuple[str, int]]]: + if self._shared_host_domain() is None: + return None + return [ + (pool.pool_label, tokens) + for pool, tokens in self._host_staging_sizes(kv_tokens, aux_xfers) + ] + + def _host_staging_sizes(self, kv_tokens, aux_xfers): + cc = self._cache.cache_controller + yield cc.mem_pool_host.anchor_entry.host_pool, kv_tokens + for transfer in aux_xfers or (): + if transfer.indices_from_pool is not None: + continue + entry = cc.mem_pool_host.entry_map.get(transfer.name) + if entry is not None: + yield entry.host_pool, self._transfer_tokens(transfer) + + def _host_request_units( + self, kv_tokens: int, aux_xfers: Optional[list[PoolTransfer]] + ) -> int: + """Host staging in anchor-token accounting units.""" + cc = self._cache.cache_controller + anchor = cc.mem_pool_host.anchor_entry.host_pool + if self._shared_host_domain() is None: + return kv_tokens + num_bytes = sum( + tokens * pool.size_per_token + for pool, tokens in self._host_staging_sizes(kv_tokens, aux_xfers) + ) + return (num_bytes + anchor.size_per_token - 1) // anchor.size_per_token + + def _shared_backup_fits(self, requests, *, empty: bool = False) -> bool: + return self._shared_host_domain().can_fit_many_then( + requests, self._shared_load_reserve_requests(), empty=empty + ) + + def _shared_load_reserve_requests(self) -> list[tuple[str, int]]: + cc = self._cache.cache_controller + anchor = cc.mem_pool_host.anchor_entry.host_pool + swa_entry = cc.mem_pool_host.entry_map.get(PoolName.SWA) + min_full_tokens = _minimum_full_transfer_tokens( + anchor, self._cache.prefetch_threshold + ) + requests = [(anchor.pool_label, min_full_tokens)] + if swa_entry is not None and self._swa_window_pages: + requests.append( + ( + swa_entry.host_pool.pool_label, + self._swa_window_pages * swa_entry.host_pool.page_size, + ) + ) + return requests + def is_idle(self) -> bool: - """No pending or in-flight buffer-mode transfer work.""" + """No queued, staged, or in-flight buffer-mode work or anchor pins.""" return not ( self.pending_hit_allocs or self.staged_prefetches @@ -310,6 +430,7 @@ def is_idle(self) -> bool: or self.inflight_backup_node_ids or self.ongoing_write_through or self.ongoing_backup + or self.anchor_locks ) def swa_transient_size(self) -> int: @@ -446,11 +567,18 @@ def _backup_oversize( capacity (total for KV, total minus the loads-priority margin for aux pools — matching ``_aux_budget_blocked``'s admission ceiling): such an intent could never stage and would wedge the FIFO head.""" + if aux_xfers is None: + aux_xfers = self._build_aux_staging_transfers(node_id, hash_values) cc = self._cache.cache_controller + shared_requests = self._shared_host_requests(intent_tokens, aux_xfers) + if shared_requests is not None: + pool_tokens = cc.mem_pool_host.size + max_write_units = pool_tokens - pool_tokens // 10 + if self._host_request_units(intent_tokens, aux_xfers) > max_write_units: + return True + return not self._shared_backup_fits(shared_requests, empty=True) if intent_tokens > cc.mem_pool_host.size: return True - if aux_xfers is None: - aux_xfers = self._build_aux_staging_transfers(node_id, hash_values) for t in aux_xfers or (): entry = cc.mem_pool_host.entry_map.get(t.name) if entry is not None and ( @@ -540,7 +668,7 @@ def flush_pending_writes(self) -> None: self._log_backup_dropped(intent_tokens) continue if self.write_staged_tokens_ >= live_cap: - # Yield to live fetch demand; retry next round. + # Wait for current copies to free staging before rebuilding the head. break device_value, comp_xfers = self._cache.tree_core.build_backup_spec( snapshot.node_id @@ -560,6 +688,17 @@ def flush_pending_writes(self) -> None: self.write_backlog_tokens_ -= intent_tokens self._log_backup_dropped(intent_tokens) continue + shared_domain = self._shared_host_domain() + staging_at_limit = ( + self.write_staged_tokens_ + + self._host_request_units(intent_tokens, sizing_xfers) + > live_cap + if shared_domain is not None + else self.write_staged_tokens_ >= live_cap + ) + if staging_at_limit: + # Yield to live fetch demand; retry next round. + break if self._aux_budget_blocked(intent, sizing_xfers): # An aux pool lacks staging headroom: yield at the gate # instead of failing the alloc inside cc.write; acks free @@ -609,13 +748,15 @@ def _stage_backup_intent( # NOTE: no commit_backup — the node must never appear # host-resident in buffer mode; staging slots live in the entry. lock_params = cache.inc_lock_ref(snapshot.node_id).to_dec_params() + occupied_units = self.host_allocation_units(host_indices, aux_xfers) self.ongoing_write_through[snapshot.node_id] = _UnifiedBufferBackupEntry( intent=intent, host_indices=host_indices, aux_xfers=aux_xfers, lock_params=lock_params, + occupied_units=occupied_units, ) - self.write_staged_tokens_ += len(host_indices) + self.write_staged_tokens_ += occupied_units self.write_backlog_tokens_ -= len(snapshot.hash_values) * cache.page_size return True @@ -636,9 +777,14 @@ def _aux_budget_blocked( aux = self._build_aux_staging_transfers( snapshot.node_id, snapshot.hash_values ) + cc = self._cache.cache_controller + shared_requests = self._shared_host_requests( + len(snapshot.hash_values) * self._cache.page_size, aux + ) + if shared_requests is not None: + return not self._shared_backup_fits(shared_requests) if not aux: return False - cc = self._cache.cache_controller for t in aux: entry = cc.mem_pool_host.entry_map.get(t.name) if entry is None: @@ -740,7 +886,7 @@ def finish_storage_write_ack(self, operation_id: int) -> None: snapshot = intent.snapshot self._cache.storage_existence_cache.add(PoolName.KV, snapshot.hash_values) self._free_staging_now(entry.host_indices, entry.aux_xfers) - self.write_staged_tokens_ -= len(entry.host_indices) + self.write_staged_tokens_ -= entry.occupied_units self.inflight_backup_node_ids.discard(snapshot.node_id) _untrack_content_refs(self.inflight_backup_hashes, snapshot.hash_values) @@ -968,12 +1114,6 @@ def _refetch_staged(self, f: _StagedPrefetch) -> None: f.request.rid, f.matched_len + f.num_tokens ) - @staticmethod - def _occupied_span(host_indices) -> int: - """Occupancy units a buffer-mode prefetch holds: granted at - hit-alloc, sized to the allocation (0 while still querying).""" - return len(host_indices) if host_indices is not None else 0 - def stage_completed_prefetch( self, request: CacheRequestHandle, @@ -1001,6 +1141,9 @@ def stage_completed_prefetch( for transfer in operation.pool_transfers or () if transfer.indices_from_pool is not None ) + occupied_tokens = operation.buffer_host_occupied_units + if occupied_tokens is None: + occupied_tokens = self.host_allocation_units(host_indices, aux_xfers) has_aux = any( t.host_indices is not None and t.host_indices.numel() > 0 for t in aux_xfers @@ -1012,7 +1155,7 @@ def stage_completed_prefetch( cc.append_host_mem_release( host_indices[:num_tokens], extra_pools=aux_xfers or None ) - cc.prefetch_tokens_occupied -= self._occupied_span(host_indices) + cc.prefetch_tokens_occupied -= occupied_tokens cache.prefetch_loaded_tokens_by_reqid[request] = 0 cache.prefetch_loaded_storage_start_by_reqid.pop(request, None) return True @@ -1024,8 +1167,6 @@ def stage_completed_prefetch( # itself is the evidence, so feeding is sound even if this staged # prefetch is later dropped unconsumed. cache.storage_existence_cache.add(PoolName.KV, list(staged_hashes)) - occupied_tokens = self._occupied_span(host_indices) - self.staged_prefetches[request] = _StagedPrefetch( request=request, key_tokens=array( @@ -1150,6 +1291,52 @@ def _defer_for_capacity(pool: str) -> None: ), 0, ) + swa_entry = cc.mem_pool_host.entry_map.get(PoolName.SWA) + binds_swa_to_full = ( + swa_entry is not None + and swa_entry.device_indices_from_anchor_fn is not None + ) + repair_ranges = [] + if binds_swa_to_full and staged_swa: + window_start = span_end - staged_swa + repair_end = min(splice_base, span_end) + if window_start < repair_end: + repair_ranges = cache.tree_core.swa_tombstone_ranges( + key, window_start, repair_end + ) + # Only missing SWA rows belong to this load. Existing bindings may + # be in use by another request and must survive allocation rollback. + anchor_parts = [] + host_parts = [] + for repair_start, repair_end_ in repair_ranges: + anchor_parts.append(req.prefix_indices[repair_start:repair_end_]) + host_parts.append( + slice(repair_start - window_start, repair_end_ - window_start) + ) + tail_start = max(splice_base, window_start) + if tail_start < span_end: + anchor_parts.append( + slice(tail_start - splice_base, span_end - splice_base) + ) + host_parts.append( + slice(tail_start - window_start, span_end - window_start) + ) + for i, transfer in enumerate(load_xfers): + if transfer.name != PoolName.SWA: + continue + # Keep the original complete host bounce for ack/drop release. + load_xfers[i] = replace( + transfer, + host_indices=( + torch.cat([transfer.host_indices[part] for part in host_parts]) + if host_parts + else transfer.host_indices[:0] + ), + anchor_index_parts=anchor_parts, + ) + if not anchor_parts: + load_xfers = [t for t in load_xfers if t.name != PoolName.SWA] + device_indices = cc.load( host_indices=f.host_indices[trim_tokens:], node_id=load_back_id, @@ -1177,7 +1364,18 @@ def _defer_for_capacity(pool: str) -> None: None, ) aux_device_releases: list[tuple[PoolName, torch.Tensor]] = [] - if swa_dev is not None: + if swa_dev is not None and binds_swa_to_full: + # Binding already updated the allocator's virtual-to-physical table. + # The tree owns virtual FULL rows, not the H2D kernel-facing IDs. + for repair_start, repair_end_ in repair_ranges: + for action in cache.tree_core.attach_swa_window( + key, + repair_start, + repair_end_, + req.prefix_indices[repair_start:repair_end_], + ): + cache._apply_cache_action(action) + elif swa_dev is not None: # Register the window's FULL->SWA translation now (attention reads # through it). Keep SWA slots another request may still hold; their # redundant H2D destinations are reclaimed at the transfer ack. @@ -1185,7 +1383,7 @@ def _defer_for_capacity(pool: str) -> None: -len(swa_dev) : ] allocator = cache.token_to_kv_pool_allocator - old_swa = allocator.full_to_swa_index_mapping[full_window.to(torch.int64)] + old_swa = allocator.translate_swa_indices_for_transfer(full_window) missing = old_swa <= 0 window_start = span_end - len(swa_dev) repair_end = min(splice_base, span_end) @@ -1287,8 +1485,19 @@ def try_finish_load_back(self, ack_id: int) -> bool: self._free_staging_now(f.host_indices, f.aux_xfers) for pool_name, device_indices in f.aux_device_releases: entry = cc.mem_pool_host.entry_map[pool_name] - free_fn = entry.device_free_fn or entry.device_pool.free - free_fn(device_indices) + from sglang.srt.mem_cache.allocator.unified_hybrid_swa import ( + UnifiedSWAAllocatorBase, + ) + + allocator = cache.token_to_kv_pool_allocator + if pool_name == PoolName.SWA and isinstance( + allocator, UnifiedSWAAllocatorBase + ): + # H2D has completed; these redundant slots are no longer pending. + allocator.swa_attn_allocator.free_physical(device_indices) + else: + free_fn = entry.device_free_fn or entry.device_pool.free + free_fn(device_indices) cc.prefetch_tokens_occupied -= f.occupied_tokens logger.info( diff --git a/python/sglang/srt/mem_cache/hybrid_cache/hybrid_cache_controller.py b/python/sglang/srt/mem_cache/hybrid_cache/hybrid_cache_controller.py index bed036f0de45..ec4d7390dede 100644 --- a/python/sglang/srt/mem_cache/hybrid_cache/hybrid_cache_controller.py +++ b/python/sglang/srt/mem_cache/hybrid_cache/hybrid_cache_controller.py @@ -39,7 +39,9 @@ from sglang.srt.mem_cache.pool_host import HostPoolGroup, PoolEntry from sglang.srt.mem_cache.pool_host.base import uses_shared_host_layout from sglang.srt.mem_cache.pool_host.mha import MHATokenToKVPoolHost +from sglang.srt.mem_cache.pool_host.unified import UnifiedPageEnvelopeHostPool from sglang.srt.mem_cache.radix_cache import RadixKey +from sglang.srt.runtime_context import get_memory if TYPE_CHECKING: from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator @@ -163,6 +165,7 @@ def __init__( pool_transfers=pool_transfers, ) self.pool_transfers_done = not bool(pool_transfers) + self.buffer_host_occupied_units: Optional[int] = None # The Python transfer worker leaves the unfinished tail to the ACK drain; # a controller that releases it itself must set this False. self.ack_releases_incomplete_host_indices = True @@ -277,6 +280,11 @@ def _stop_storage_threads(self): self._stop_pp_prefetch_thread() super()._stop_storage_threads() + @staticmethod + def supports_page_envelope_host(storage_backend: str | None) -> bool: + """Only opt in backends whose FULL object is the complete shared page.""" + return storage_backend is None or storage_backend == "mori" + def attach_storage_backend( self, storage_backend: str, @@ -285,6 +293,14 @@ def attach_storage_backend( storage_backend_extra_config: Optional[dict] = None, host_pools: Optional[list[PoolEntry]] = None, ): + if isinstance( + self.storage_host_pool, UnifiedPageEnvelopeHostPool + ) and not self.supports_page_envelope_host(storage_backend): + # A runtime backend switch cannot replace the existing host arena. + raise ValueError( + f"Storage backend {storage_backend!r} requires separate host pools. " + "Restart with this backend selected to create compatible host pools." + ) enable_pp_ticket = ( self.host_memory_mode == "buffer_only" and storage_backend == "mooncake" @@ -572,6 +588,9 @@ def _alloc_shared_host_requests_with_reclaim( for entry, (_, need_size) in zip(entries, requests, strict=True) ] + # Scheduler drains commit the same logical host allocations and frees + # across attention ranks. Layout compaction waits for local I/O before + # allocating; transfer readiness does not change the capacity decision. allocated = domain.alloc_many(domain_requests) if allocated is not None: return allocated @@ -615,7 +634,8 @@ def _alloc_shared_host_requests_with_reclaim( // pool.page_size * pool.page_size ) - if entry.host_evict_fn(tokens) <= 0: + evicted = entry.host_evict_fn(tokens) + if evicted <= 0: continue made_progress = True allocated = domain.alloc_many(domain_requests) @@ -665,6 +685,134 @@ def allocate_shared_host_transfers( return None return host_indices, resolved + def _shared_prefetch_requests( + self, operation: PrefetchOperation, need_size: int + ) -> tuple[list[tuple[PoolName, int]], list[PoolTransfer]]: + anchor = self.mem_pool_host.anchor_entry + pool_transfers = operation.pool_transfers or [] + requests = [(anchor.name, need_size)] + sidecar_hit_pages = ( + operation.sidecar_hit_pages + if operation.sidecar_hash_values is not None + else need_size // self.page_size + ) + independent_transfers = [] + for transfer in pool_transfers: + if transfer.indices_from_pool is not None: + continue + if transfer.host_indices is not None: + raise AssertionError( + "Shared prefetch pools must be allocated atomically." + ) + entry = self.mem_pool_host.entry_map.get(transfer.name) + if entry is None: + continue + num_pages = len(transfer.keys or []) + if transfer.hit_policy == PoolHitPolicy.TRAILING_PAGES: + num_pages = min(num_pages or 1, sidecar_hit_pages) + requests.append((transfer.name, num_pages * entry.host_pool.page_size)) + independent_transfers.append(transfer) + return requests, independent_transfers + + def can_fit_prefetch_host_buffers( + self, operation: StorageOperation, need_size: int, *, empty: bool = True + ) -> bool: + anchor = self.mem_pool_host.anchor_entry + if not uses_shared_host_layout(anchor.host_pool): + return super().can_fit_prefetch_host_buffers(operation, need_size) + domain = anchor.host_pool.shared_allocation_domain + requests, _ = self._shared_prefetch_requests(operation, need_size) + domain_requests = [ + (self.mem_pool_host.entry_map[name].host_pool.pool_label, size) + for name, size in requests + ] + return domain.can_fit_many_then(domain_requests, (), empty=empty) + + def prefetch_rate_limited(self) -> bool: + if self.host_memory_mode == "buffer_only" and uses_shared_host_layout( + self.mem_pool_host.anchor_entry.host_pool + ): + # Shared-arena occupancy includes the Full and SWA byte footprint. + return self.prefetch_tokens_occupied >= self.prefetch_capacity_limit + return super().prefetch_rate_limited() + + def alloc_prefetch_host_buffers( + self, operation: StorageOperation, need_size: int + ) -> Optional[torch.Tensor]: + """Atomically allocate every prefetch slice from a shared arena.""" + anchor = self.mem_pool_host.anchor_entry + if not uses_shared_host_layout(anchor.host_pool): + return super().alloc_prefetch_host_buffers(operation, need_size) + + requests, independent_transfers = self._shared_prefetch_requests( + operation, need_size + ) + + allocated = self._alloc_shared_host_requests_with_reclaim(requests) + if allocated is None: + return None + self._sync_trailing_keys( + operation.pool_transfers or [], + operation.sidecar_hash_values or operation.hash_value, + ( + operation.sidecar_hit_pages + if operation.sidecar_hash_values is not None + else need_size // self.page_size + ), + ) + for transfer, indices in zip(independent_transfers, allocated[1:], strict=True): + transfer.host_indices = indices + return allocated[0] + + def allocate_storage_hit( + self, + operation: StorageOperation, + hit_tokens: int, + *, + allow_partial: bool, + min_tokens: int, + evict_host: Callable[[int], int], + ) -> tuple[Optional[torch.Tensor], int]: + host_indices = self.alloc_prefetch_host_buffers(operation, hit_tokens) + shared = uses_shared_host_layout(self.mem_pool_host.anchor_entry.host_pool) + if host_indices is None and not shared: + evict_host(hit_tokens) + host_indices = self.alloc_prefetch_host_buffers(operation, hit_tokens) + if host_indices is not None or not allow_partial: + return host_indices, hit_tokens + + if shared: + low, high = 0, hit_tokens // self.page_size + while low < high: + mid = (low + high + 1) // 2 + if self.can_fit_prefetch_host_buffers( + operation, mid * self.page_size, empty=False + ): + low = mid + else: + high = mid - 1 + alloc_len = low * self.page_size + else: + available = self.mem_pool_host.available_size() + alloc_len = min(hit_tokens, available - available % self.page_size) + if alloc_len >= min_tokens: + host_indices = self.alloc_prefetch_host_buffers(operation, alloc_len) + return host_indices, alloc_len + + def free_prefetch_host_buffers( + self, operation: StorageOperation, host_indices: torch.Tensor + ) -> None: + """Roll back an unsubmitted prefetch's anchor and independent pools.""" + pool_transfers = operation.pool_transfers or [] + self.mem_pool_host.free(host_indices) + for transfer in pool_transfers: + if transfer.indices_from_pool is not None or transfer.host_indices is None: + continue + entry = self.mem_pool_host.entry_map.get(transfer.name) + if entry is not None: + entry.host_pool.free(transfer.host_indices) + transfer.host_indices = None + def _move_op_indices( self, op: CacheOperation ) -> tuple[torch.Tensor, torch.Tensor, Optional[list[PoolTransfer]]]: @@ -1308,16 +1456,26 @@ def _page_transfer_sidecar( ) self._sync_trailing_keys(transfers_nonkv, sidecar_hashes, sidecar_hit_pages) self._resolve_sidecar_nonkv_derived_pool_transfers(operation) - results = {} - for extra_info, transfers in _trailing_chain_groups( - operation.prefix_keys, - sidecar_hashes, - transfers_nonkv, - ): - results.update( - self.storage_backend.batch_get_v2(transfers, extra_info=extra_info) + try: + results = {} + for extra_info, transfers in _trailing_chain_groups( + operation.prefix_keys, + sidecar_hashes, + transfers_nonkv, + ): + results.update( + self.storage_backend.batch_get_v2( + transfers, extra_info=extra_info + ) + ) + pool_hits = count_pool_hits(results) + except Exception: + if not get_memory().enable_unified_memory: + raise + logger.exception( + "HiCache sidecar prefetch failed for request %s", + operation.request_id, ) - pool_hits = count_pool_hits(results) # Emit PrefetchAck to prefetch_sync_queue, even the operation has been canceled by the # scheduler thread. The prefetch sync thread expects the same number of PrefetchAck objects # to perform all_reduce. @@ -1521,7 +1679,7 @@ def _resolve_device_transfers( anchor_transfers = [] def rollback_allocated() -> None: - for prev_pool, prev_free_fn, prev_indices in newly_allocated: + for prev_pool, prev_free_fn, prev_indices in reversed(newly_allocated): prev_free_fn(prev_indices) prev_pool.device_indices = None diff --git a/python/sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py b/python/sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py index 8baee144d0ad..8d3247e09160 100644 --- a/python/sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py +++ b/python/sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py @@ -273,6 +273,13 @@ def build_kv_only_group( def _swa_allocation_callbacks(allocator, bind=None, free_bound=None) -> dict: """Keep allocation and rollback in the same ID space for every SWA stack.""" + from sglang.srt.mem_cache.allocator.unified_sub_pool import MultiEndedAllocator + + if isinstance(allocator, MultiEndedAllocator): + return dict( + device_alloc_fn=allocator.alloc_physical, + device_free_fn=allocator.cancel_physical_reservation, + ) if bind is not None: assert free_bound is not None return dict( @@ -300,7 +307,9 @@ def _uses_unified_page_envelope_host( for pool in (full_kv_pool, swa_kv_pool) ) and {full_kv_pool.grow_direction, swa_kv_pool.grow_direction} == {"up", "down"} - and get_memory().hicache_host_memory_mode != "buffer_only" + and HybridCacheController.supports_page_envelope_host( + get_memory().hicache_storage_backend + ) ) diff --git a/python/sglang/srt/mem_cache/kv_cache_builder.py b/python/sglang/srt/mem_cache/kv_cache_builder.py index 1442a6ab2432..016d9c4cfa21 100644 --- a/python/sglang/srt/mem_cache/kv_cache_builder.py +++ b/python/sglang/srt/mem_cache/kv_cache_builder.py @@ -176,19 +176,6 @@ def resolve_decode_retraction_backup(*, tp_worker: BaseTpWorker) -> str: backend = disagg.disaggregation_decode_retraction_backup unified_hybrid_swa = memory.enable_unified_memory and tp_worker.is_hybrid_swa - unified_decode_radix = ( - memory.enable_unified_memory and disagg.disaggregation_decode_enable_radix_cache - ) - if backend == "host_pool" and unified_decode_radix: - raise ValueError( - "Unified-memory H2D/D2H does not support host-pool decode " - "retraction with decode radix cache yet." - ) - if backend == "host_pool" and unified_hybrid_swa: - raise ValueError( - "Unified-memory hybrid-SWA H2D/D2H does not support host-pool " - "decode retraction yet." - ) if backend is None: kv_cache = tp_worker.get_memory_pool()[1].get_kvcache() full_tokens_per_layer = ( diff --git a/python/sglang/srt/mem_cache/kv_cache_configurator.py b/python/sglang/srt/mem_cache/kv_cache_configurator.py index 405715bdb410..14789c2fd1e6 100644 --- a/python/sglang/srt/mem_cache/kv_cache_configurator.py +++ b/python/sglang/srt/mem_cache/kv_cache_configurator.py @@ -2174,6 +2174,7 @@ def _build_token_to_kv_pool_allocator( swa_allocator = token_to_kv_pool_allocator.logical_attn_allocator else: swa_allocator = token_to_kv_pool_allocator + assert isinstance(swa_allocator, SWATokenToKVPoolAllocator) uses_unified_virtual_ids = isinstance( swa_allocator, UnifiedSWAAllocatorBase ) @@ -2194,7 +2195,6 @@ def _build_token_to_kv_pool_allocator( identity_mapping[-1] = -1 token_to_kv_pool.register_mapping(identity_mapping) elif not uses_unified_virtual_ids: - assert isinstance(swa_allocator, SWATokenToKVPoolAllocator) token_to_kv_pool.register_mapping( swa_allocator.full_to_swa_index_mapping ) diff --git a/python/sglang/srt/mem_cache/memory_pool_host.py b/python/sglang/srt/mem_cache/memory_pool_host.py index 4c5963b74826..3dfc66867111 100644 --- a/python/sglang/srt/mem_cache/memory_pool_host.py +++ b/python/sglang/srt/mem_cache/memory_pool_host.py @@ -64,6 +64,7 @@ class LogicalHostPool: """ shared_allocation_domain = None + stores_page_envelope = False storage_format_tag = None def __init__(self, size: int, page_size: int, layout: str = "layer_first"): diff --git a/python/sglang/srt/mem_cache/pool_host/base.py b/python/sglang/srt/mem_cache/pool_host/base.py index 8469bd28d6ae..56588886ecd1 100644 --- a/python/sglang/srt/mem_cache/pool_host/base.py +++ b/python/sglang/srt/mem_cache/pool_host/base.py @@ -164,6 +164,7 @@ class HostKVCache(abc.ABC): dcp_size = 1 dcp_rank = 0 shared_allocation_domain = None + stores_page_envelope = False # Names this pool's page byte format in storage keys when it has one of its # own, so pages persisted in another format miss instead of loading. storage_format_tag: Optional[str] = None diff --git a/python/sglang/srt/mem_cache/pool_host/group.py b/python/sglang/srt/mem_cache/pool_host/group.py index 4abe9b3e0bc8..68d7a24639c8 100644 --- a/python/sglang/srt/mem_cache/pool_host/group.py +++ b/python/sglang/srt/mem_cache/pool_host/group.py @@ -291,6 +291,10 @@ def dsa_kv_cache_store_fp8(self): def size_per_token(self): return self.anchor_entry.host_pool.size_per_token + @property + def stores_page_envelope(self) -> bool: + return self.anchor_entry.host_pool.stores_page_envelope + def clear(self) -> None: for entry in self.entries: entry.host_pool.clear() diff --git a/python/sglang/srt/mem_cache/pool_host/unified.py b/python/sglang/srt/mem_cache/pool_host/unified.py index 48ca720a20ff..4ffd5e8843e5 100644 --- a/python/sglang/srt/mem_cache/pool_host/unified.py +++ b/python/sglang/srt/mem_cache/pool_host/unified.py @@ -303,15 +303,30 @@ def _can_reserve_without_compaction(self, page_counts: dict[str, int]) -> bool: extension_bytes += extension_pages * side.page_bytes return extension_bytes <= self._current_gap_bytes() - def _can_fit_packed(self, page_counts: dict[str, int]) -> bool: + def _can_fit_state( + self, + live_page_counts: dict[str, int], + free_logical_page_counts: dict[str, int], + request_page_counts: dict[str, int], + ) -> bool: used_bytes = 0 - for name, page_count in page_counts.items(): - side = self.sides[name] - if page_count > self._logical_free_page_count(side): + for name, page_count in request_page_counts.items(): + if page_count > free_logical_page_counts[name]: return False - used_bytes += (side.live_page_count + page_count) * side.page_bytes + side = self.sides[name] + used_bytes += (live_page_counts[name] + page_count) * side.page_bytes return used_bytes <= self._allocatable_bytes + def _can_fit_packed(self, page_counts: dict[str, int]) -> bool: + return self._can_fit_state( + {name: side.live_page_count for name, side in self.sides.items()}, + { + name: self._logical_free_page_count(side) + for name, side in self.sides.items() + }, + page_counts, + ) + def _begin_compaction(self) -> list: if self._layout_lease_state.depth: raise RuntimeError( @@ -543,6 +558,48 @@ def alloc_many( self.sides[name].free_logical_extents = extents return results + def can_fit_many_then( + self, + requests: Sequence[tuple[str, int]], + following_requests: Sequence[tuple[str, int]], + *, + empty: bool = False, + ) -> bool: + """Whether both allocation groups fit in order without mutating state.""" + with self.lock: + _, first_page_counts = self._request_page_counts(requests) + _, following_page_counts = self._request_page_counts(following_requests) + if empty: + live_page_counts = {name: 0 for name in self.sides} + free_logical_page_counts = { + name: side.page_num for name, side in self.sides.items() + } + else: + live_page_counts = { + name: side.live_page_count for name, side in self.sides.items() + } + free_logical_page_counts = { + name: self._logical_free_page_count(side) + for name, side in self.sides.items() + } + if not self._can_fit_state( + live_page_counts, free_logical_page_counts, first_page_counts + ): + return False + live_after_first = { + name: live_page_counts[name] + first_page_counts[name] + for name in self.sides + } + free_after_first = { + name: free_logical_page_counts[name] - first_page_counts[name] + for name in self.sides + } + return self._can_fit_state( + live_after_first, + free_after_first, + following_page_counts, + ) + def free(self, name: str, indices: torch.Tensor) -> int: with self.lock: side = self.sides[name] @@ -610,6 +667,7 @@ def translate_index(self, name: str, index: int) -> int: class UnifiedPageEnvelopeHostPool(HostKVCache): """Host mirror that transfers complete unified-memory page envelopes.""" + stores_page_envelope = True storage_format_tag = "unified-token-major" def __init__( diff --git a/python/sglang/srt/mem_cache/prefill_budget.py b/python/sglang/srt/mem_cache/prefill_budget.py index a1926abe0ebd..581c308ac343 100644 --- a/python/sglang/srt/mem_cache/prefill_budget.py +++ b/python/sglang/srt/mem_cache/prefill_budget.py @@ -15,7 +15,8 @@ The scheduler supplies token demand and its chunk/decode limits. These objects account for admitted but not yet allocated work and query live cache capacity: locking a prefix or preempting a request must affect the next admission check. -They neither select requests nor mutate the prefix cache or allocator. +Selection checks do not mutate cache state. Shared-pool load preparation +realizes the selected reservation before a host transfer pins device rows. """ from typing import Optional @@ -72,6 +73,18 @@ def __init__(self, allocator, tree_cache, *, num_mixed_decode_tokens: int = 0): def ceil_paged_tokens(self, tokens: int) -> int: return -(-tokens // self.page_size) * self.page_size + def prepare_load_back( + self, + *, + full_tokens: int, + extend_input_len: int, + max_new_tokens: int, + swa_host_hit_length: int, + chunk_limit: int | None, + ) -> bool: + """Prepare pools whose admission depends on movable shared space.""" + return True + def _available_and_evictable(self): evictable = ( self.tree_cache.full_evictable_size() @@ -322,6 +335,36 @@ def _fits(self, full_tokens, swa_tokens, *, empty_pool=False): require_token_slack=empty_pool, ) + def prepare_load_back( + self, + *, + full_tokens: int, + extend_input_len: int, + max_new_tokens: int, + swa_host_hit_length: int, + chunk_limit: int | None, + ) -> bool: + # H2D pins both bands. Realize the selected joint budget while their + # holes can still be compacted, including pending prefill/decode demand. + full_tokens = int(self.ceil_paged_tokens(full_tokens + self.total_offset)) + swa_tokens = int( + self.ceil_paged_tokens( + self.swa_tokens( + extend_input_len, + max_new_tokens, + chunk_limit=chunk_limit, + swa_host_hit_length=swa_host_hit_length, + ) + + self.swa_offset + ) + ) + return ( + self.allocator.reclaim_for_prealloc( + self.tree_cache, full_tokens, swa_tokens + ) + is None + ) + def _joint_chunk_cap(self, *, max_chunk_tokens, chunk_limit, swa_host_hit_length=0): lo, hi = 0, max(0, max_chunk_tokens) // self.page_size while lo < hi: diff --git a/python/sglang/srt/mem_cache/rust_tree_core/adapter.py b/python/sglang/srt/mem_cache/rust_tree_core/adapter.py index cd45c5a84db8..ae458af945d4 100644 --- a/python/sglang/srt/mem_cache/rust_tree_core/adapter.py +++ b/python/sglang/srt/mem_cache/rust_tree_core/adapter.py @@ -5,7 +5,7 @@ import sys from array import array -from typing import TYPE_CHECKING, Optional, Sequence +from typing import TYPE_CHECKING, Callable, Optional, Sequence import torch @@ -371,6 +371,19 @@ def __init__(self, params: CacheInitParams): ) self._page_size = params.page_size + self._swa_backup_index_mapper: Optional[ + Callable[[torch.Tensor], torch.Tensor] + ] = None + allocator = params.token_to_kv_pool_allocator + if allocator is not None and ComponentType.SWA in self.tree_components: + from sglang.srt.mem_cache.allocator.unified_hybrid_swa import ( + UnifiedSWAAllocatorBase, + ) + + if isinstance(allocator, UnifiedSWAAllocatorBase): + self._swa_backup_index_mapper = ( + allocator.translate_swa_indices_for_transfer + ) self.is_eagle = ( params.is_eagle and ComponentType.MAMBA not in self.tree_components ) @@ -555,6 +568,11 @@ def evict_device_next_node( result = EvictDeviceNextNodeResult( node_id=binding_result.node_id, made_progress=binding_result.made_progress, + backup_kv=( + _cache_action_from_tagged(binding_result.backup_kv) + if binding_result.backup_kv is not None + else None + ), unbacked_tokens=binding_result.unbacked_tokens, mamba_backup_node_id=binding_result.mamba_backup_node_id, swa_backup_node_id=binding_result.swa_backup_node_id, @@ -770,6 +788,9 @@ def evict_excess_path_states( def set_hicache_enabled(self) -> None: self._binding.set_hicache_enabled() + def enable_swa_write_back_eviction_barrier(self) -> None: + self._binding.enable_swa_write_back_eviction_barrier() + def set_host_memory_buffer_only(self) -> None: self._binding.set_host_memory_buffer_only() @@ -868,11 +889,7 @@ def build_backup_spec( def _refresh_swa_backup_indices(self, transfers: Sequence[PoolTransfer]) -> None: if not transfers: return - from sglang.srt.mem_cache.allocator.unified_hybrid_swa import ( - UnifiedSWAAllocatorBase, - ) - - if not isinstance(self._allocator, UnifiedSWAAllocatorBase): + if self._swa_backup_index_mapper is None: return for transfer in transfers: if transfer.name != PoolName.SWA or not transfer.nodes_to_load: @@ -883,7 +900,7 @@ def _refresh_swa_backup_indices(self, transfers: Sequence[PoolTransfer]) -> None self.get_component_device_value(node_id, ComponentType.FULL) for node_id in transfer.nodes_to_load ] - transfer.device_indices = self._allocator.translate_loc_from_full_to_swa( + transfer.device_indices = self._swa_backup_index_mapper( torch.cat(full_values) ).to(torch.int64) @@ -972,6 +989,9 @@ def prefetch_anchor_info( ) -> tuple[Optional[str], Optional[str]]: return self._binding.prefetch_anchor_info(node_id) + def is_write_through_compatible(self) -> bool: + return self._binding.is_write_through_compatible() + def is_backuped(self, node_id: NodeId) -> bool: return self._binding.node_backuped(node_id) diff --git a/python/sglang/srt/mem_cache/storage/umbp/umbp_store.py b/python/sglang/srt/mem_cache/storage/umbp/umbp_store.py index 019fad958577..eadeb91afb55 100644 --- a/python/sglang/srt/mem_cache/storage/umbp/umbp_store.py +++ b/python/sglang/srt/mem_cache/storage/umbp/umbp_store.py @@ -917,6 +917,7 @@ def _resolve_node_tags() -> List[str]: # Initialize before the optional constructor-time pool registration. self.registered_pools: dict = {} self._kv_anchor_is_logical = False + self._kv_anchor_is_page_envelope = False self._registered_regions: set = set() self.client = UMBPClient(cfg) @@ -1001,6 +1002,7 @@ def register_mem_pool_host(self, mem_pool_host: HostKVCache): # A logical anchor owns indices; side pools carry the data. self._kv_anchor_is_logical = self.mem_pool_host.kv_buffer is None + self._kv_anchor_is_page_envelope = self.mem_pool_host.stores_page_envelope self._zero_copy_registered = False # Side-pool registration needs the mode even for a logical anchor. @@ -1163,68 +1165,61 @@ def register_mem_host_pool_v2(self, host_pool: HostKVCache, host_pool_name): # ------------------------------------------------------------------ # Key suffix generation — mirrors MooncakeStore # ------------------------------------------------------------------ - def _get_mha_buffer_meta(self, keys, indices): - ptr_list, element_size_list = self.mem_pool_host.get_page_buffer_meta(indices) - key_list = [] - for key_ in keys: - key_list.append(f"{key_}_{self.mha_suffix}_k") - key_list.append(f"{key_}_{self.mha_suffix}_v") - assert len(key_list) == len(ptr_list) - return key_list, ptr_list, element_size_list - - def _get_mha_split_heads_buffer_meta(self, keys, indices): - ptr_list, element_size_list = ( - self.mem_pool_host.get_split_heads_page_buffer_meta( - indices, self.split_factor - ) + def _anchor_key_suffixes(self) -> tuple[str, ...]: + if self._kv_anchor_is_page_envelope: + return (f"_{self.mha_suffix}_kv",) + if self.is_mla_backend: + return (f"_{self.mla_suffix}_k",) + ranks = ( + self.mha_suffix + if self.storage_config and self.storage_config.should_split_heads + else (self.mha_suffix,) + ) + return tuple( + f"_{rank}_{component}" for rank in ranks for component in ("k", "v") ) - key_list = [] - for key_ in keys: - for suffix in self.mha_suffix: - key_list.append(f"{key_}_{suffix}_k") - key_list.append(f"{key_}_{suffix}_v") - assert len(key_list) == len(ptr_list) - return key_list, ptr_list, element_size_list - def _get_mla_buffer_meta(self, keys, indices): - ptr_list, element_size_list = self.mem_pool_host.get_page_buffer_meta(indices) - key_list = [] - for key_ in keys: - key_list.append(f"{key_}_{self.mla_suffix}_k") - assert len(key_list) == len(ptr_list) - return key_list, ptr_list, element_size_list + def _anchor_keys(self, keys) -> tuple[list[str], int]: + suffixes = self._anchor_key_suffixes() + return [f"{key}{suffix}" for key in keys for suffix in suffixes], len(suffixes) def _batch_preprocess(self, keys, host_indices): assert len(keys) > 0 assert len(keys) == len(host_indices) // self.mem_pool_host.page_size - if self.is_mla_backend: - return self._get_mla_buffer_meta(keys, host_indices) + key_list, _ = self._anchor_keys(keys) + if ( + not self._kv_anchor_is_page_envelope + and not self.is_mla_backend + and self.storage_config + and self.storage_config.should_split_heads + ): + ptr_list, element_size_list = ( + self.mem_pool_host.get_split_heads_page_buffer_meta( + host_indices, self.split_factor + ) + ) else: - if self.storage_config and self.storage_config.should_split_heads: - return self._get_mha_split_heads_buffer_meta(keys, host_indices) - else: - return self._get_mha_buffer_meta(keys, host_indices) + ptr_list, element_size_list = self.mem_pool_host.get_page_buffer_meta( + host_indices + ) + assert len(key_list) == len(ptr_list) + return key_list, ptr_list, element_size_list def _batch_postprocess(self, results: List[bool], is_set_operate=False): """Convert per-key-component results to per-page results. - For MHA: each page has K+V → group pairs. - For MLA: each page has K only. + Unified page envelopes and MLA have one object per page. Ordinary MHA + has K+V pairs, or one pair per split rank. """ - if self.is_mla_backend: + group_size = len(self._anchor_key_suffixes()) + if group_size == 1: return list(results) - else: - if self.storage_config and self.storage_config.should_split_heads: - group_size = self.split_factor * 2 - groups = [ - results[i : i + group_size] - for i in range(0, len(results), group_size) - ] - return [all(g) for g in groups] - else: - # Group K/V pairs - kv_pairs = zip(results[::2], results[1::2]) - return [k and v for k, v in kv_pairs] + result_count = len(results) + if not (self.storage_config and self.storage_config.should_split_heads): + result_count -= result_count % group_size + return [ + all(results[i : i + group_size]) for i in range(0, result_count, group_size) + ] # ------------------------------------------------------------------ # Zero-copy v1 interface @@ -1282,19 +1277,8 @@ def _compute_expanded_depths( depths_per_page = [prefix_len + i for i in range(len(keys))] # Expand to match the key_strs layout produced by _batch_preprocess. - expanded = [] - for d in depths_per_page: - if self.is_mla_backend: - expanded.append(d) # MLA: 1 key per page - elif self.storage_config and self.storage_config.should_split_heads: - # split heads: 2 keys per split rank, split_factor ranks per page - for _ in range(self.split_factor): - expanded.append(d) - expanded.append(d) - else: - expanded.append(d) # K - expanded.append(d) # V - return expanded + objects_per_page = len(self._anchor_key_suffixes()) + return [depth for depth in depths_per_page for _ in range(objects_per_page)] def batch_set_v1( self, @@ -1352,23 +1336,7 @@ def batch_exists( self, keys: List[str], extra_info: Optional[HiCacheStorageExtraInfo] = None ) -> int: """Return count of consecutive existing keys from start.""" - if self.is_mla_backend: - query_keys = [f"{key}_{self.mla_suffix}_k" for key in keys] - key_multiplier = 1 - else: - query_keys = [] - if self.storage_config and self.storage_config.should_split_heads: - for key in keys: - for suffix in self.mha_suffix: - query_keys.append(f"{key}_{suffix}_k") - query_keys.append(f"{key}_{suffix}_v") - key_multiplier = 2 * self.split_factor - else: - for key in keys: - query_keys.append(f"{key}_{self.mha_suffix}_k") - query_keys.append(f"{key}_{self.mha_suffix}_v") - key_multiplier = 2 - + query_keys, key_multiplier = self._anchor_keys(keys) hit_count = self.client.batch_exists_consecutive(query_keys) return hit_count // key_multiplier diff --git a/python/sglang/srt/mem_cache/unified_cache/components/swa.py b/python/sglang/srt/mem_cache/unified_cache/components/swa.py index 6198a0a9ab45..d6cfe785e928 100644 --- a/python/sglang/srt/mem_cache/unified_cache/components/swa.py +++ b/python/sglang/srt/mem_cache/unified_cache/components/swa.py @@ -221,9 +221,14 @@ def validate_session_state( ) def _translate_full_to_swa(self, full_indices: torch.Tensor) -> torch.Tensor: - return self.cache.token_to_kv_pool_allocator.translate_loc_from_full_to_swa( - full_indices + swa_indices = ( + self.cache.token_to_kv_pool_allocator.translate_swa_indices_for_transfer( + full_indices + ) ) + # Tree component values use int64 indices; normalize transfer ids at the + # tree boundary. + return swa_indices.to(torch.int64) def _unified_allocator(self): """The unified SWA composite, or None when running on the static pool.""" diff --git a/python/sglang/srt/mem_cache/unified_cache/storage_attachment.py b/python/sglang/srt/mem_cache/unified_cache/storage_attachment.py index ee264d325396..7f0c0899f1c3 100644 --- a/python/sglang/srt/mem_cache/unified_cache/storage_attachment.py +++ b/python/sglang/srt/mem_cache/unified_cache/storage_attachment.py @@ -14,9 +14,11 @@ import logging from typing import TYPE_CHECKING, Optional +from sglang.srt.mem_cache.buffer_mode.pipeline import validate_buffer_only_stack from sglang.srt.mem_cache.hybrid_cache.hybrid_cache_controller import ( HybridCacheController, ) +from sglang.srt.mem_cache.unified_cache.component_type import ComponentType from sglang.srt.observability.metrics_collector import ( STAT_LOGGER_ROLE_STORAGE, StorageMetricsCollector, @@ -73,6 +75,23 @@ def attach( "launch with --enable-hierarchical-cache to attach a backend.", ) + # Write-back can retain auxiliary host rows without FULL, or backed + # descendants below an unbacked FULL ancestor. Those states are invalid + # for write-through insertion and eviction. Check before changing any + # policy, including on the same-backend update and re-attach paths. + if ( + cache.is_write_back + and cache.host_memory_mode != "buffer_only" + and hicache_write_policy in ("write_through", "write_through_selective") + and not cache.tree_core.is_write_through_compatible() + ): + return ( + False, + "Cannot switch HiCache from write_back to write_through while " + "the cache has host data without its FULL prefix. " + "Flush the cache before changing the write policy.", + ) + if cache.enable_storage: current_backend = controller.storage_backend_type if current_backend != storage_backend: @@ -90,10 +109,6 @@ def attach( "policies updated.", ) - # Apply policies before the controller attach, so the storage threads - # observe the new values as soon as they start. - self._apply_policies(hicache_storage_prefetch_policy, hicache_write_policy) - logger.info(f"Attaching HiCache storage backend: {storage_backend}") try: ( @@ -113,8 +128,23 @@ def attach( f"'{storage_backend_extra_config_json}': {e}", ) + original_policies = ( + cache.prefetch_stop_policy, + controller.write_policy, + cache.write_through_threshold, + cache.is_write_back, + ) try: prefetch_threshold = self.resolve_prefetch_threshold(prefetch_threshold) + if cache.host_memory_mode == "buffer_only": + validate_buffer_only_stack( + sidecar_pool_specs=cache.sidecar_pool_specs, + host_pool_group=cache.host_pool_group, + swa_component=cache.components.get(ComponentType.SWA), + storage_prefetch_threshold=prefetch_threshold, + ) + # New workers must see the requested policy from their first operation. + self._apply_policies(hicache_storage_prefetch_policy, hicache_write_policy) controller.attach_storage_backend( storage_backend=storage_backend, prefetch_threshold=prefetch_threshold, @@ -123,6 +153,12 @@ def attach( host_pools=controller.mem_pool_host.entries, ) except Exception as e: + ( + cache.prefetch_stop_policy, + controller.write_policy, + cache.write_through_threshold, + cache.is_write_back, + ) = original_policies logger.exception( f"Failed to attach storage backend '{storage_backend}': {e}" ) @@ -395,15 +431,19 @@ def _release_pending_storage_ops(self) -> None: cache.buffer_pipeline.release_anchor_lock(handle) controller.append_host_mem_release( host_indices=info.host_indices[:completed_tokens], - extra_pools=[ - x for xfers in info.comp_xfers.values() for x in xfers - ], + extra_pools=( + [x for xfers in info.comp_xfers.values() for x in xfers] + if info.operation.pool_transfers_done + else None + ), ) controller.prefetch_tokens_occupied = max( 0, controller.prefetch_tokens_occupied - cache._prefetch_occupied_span( - info.prefetch_key, info.host_indices + info.prefetch_key, + info.host_indices, + operation=info.operation, ), ) except Exception: diff --git a/python/sglang/srt/mem_cache/unified_cache/unified_tree_core.py b/python/sglang/srt/mem_cache/unified_cache/unified_tree_core.py index d4f088a64f7a..b8fdb70c65f9 100644 --- a/python/sglang/srt/mem_cache/unified_cache/unified_tree_core.py +++ b/python/sglang/srt/mem_cache/unified_cache/unified_tree_core.py @@ -544,6 +544,22 @@ def node_by_id(self, node_id: NodeId) -> UnifiedTreeNode: """ return self._node_arena[node_id] + def is_write_through_compatible(self) -> bool: + for node in self._node_arena.values(): + if node is self.root_node: + continue + full_host = node.component_data[BASE_COMPONENT_TYPE].host_value + if full_host is None: + if any( + node.component_data[ct].host_value is not None + for ct in self.component_types + if ct != BASE_COMPONENT_TYPE + ): + return False + elif node.parent is not self.root_node and not node.parent.backuped: + return False + return True + def is_backuped(self, node_id: NodeId) -> bool: """Whether the node's KV is already backed up to host.""" return self._node_arena[node_id].backuped @@ -1756,14 +1772,20 @@ def drop_subtree_no_host(self, node_id: NodeId) -> DropSubtreeNoHostResult: # A failed backup never issues the D->H copy, so the subtree root has # no host state and no in-flight DMA reading its device slots. assert not node.backuped and node.write_through_pending_id is None - if any(cd.host_lock_ref > 0 for cd in node.component_data): + if node.load_back_pending_id is not None or any( + cd.lock_ref > 0 or cd.host_lock_ref > 0 for cd in node.component_data + ): return result descendants: list[UnifiedTreeNode] = [] stack = list(node.children.values()) while stack: cur = stack.pop() - if any( - cd.lock_ref > 0 or cd.host_lock_ref > 0 for cd in cur.component_data + if ( + cur.write_through_pending_id is not None + or cur.load_back_pending_id is not None + or any( + cd.lock_ref > 0 or cd.host_lock_ref > 0 for cd in cur.component_data + ) ): return result descendants.append(cur) @@ -2275,12 +2297,17 @@ def insert_host( child_key = key.child_key(self.page_size) matched_length = 0 + host_prefix_len = 0 + refilled_host_node = None cache_actions: list[CacheAction | ComponentAction] = [] while len(key) > 0 and child_key in node.children: node = node.children[child_key] self._touch_node(node) prefix_len = node.key.match(key, page_size=self.page_size) + matched_host_value = host_value[:prefix_len] + matched_hash_value = hash_value[: prefix_len // self.page_size] + key = key[prefix_len:] host_value = host_value[prefix_len:] hash_value = hash_value[prefix_len // self.page_size :] @@ -2291,18 +2318,37 @@ def insert_host( if action is not None: cache_actions.append(action) + if not self.is_write_back: + full_data = node.component_data[BASE_COMPONENT_TYPE] + if full_data.host_value is None: + full_data.host_value = matched_host_value.clone() + if node.hash_value is None: + node.hash_value = list(matched_hash_value) + self.kv_events.record_store(node, medium=StorageMedium.CPU) + self._update_evictable_leaf_sets(node) + if node.parent is not None: + self._update_evictable_leaf_sets(node.parent) + self._update_duplicate_tracking(node) + refilled_host_node = node + else: + assert refilled_host_node is None, ( + "write-through host-prefix invariant broken: encountered " + f"backed node {node.id} below a refilled host tombstone" + ) + host_prefix_len += prefix_len + if len(key): child_key = key.child_key(self.page_size) result = InsertResult( - prefix_len=matched_length, + prefix_len=(matched_length if self.is_write_back else host_prefix_len), total_len=total_len, cache_actions=cache_actions, ) if len(key) == 0: - if ( - node is not self.root_node - and node.component_data[BASE_COMPONENT_TYPE].host_value is not None + if node is not self.root_node and ( + refilled_host_node is not None + or node.component_data[BASE_COMPONENT_TYPE].host_value is not None ): result.inserted_host_node = node.id return result diff --git a/python/sglang/srt/mem_cache/unified_cache/unified_tree_core_interface.py b/python/sglang/srt/mem_cache/unified_cache/unified_tree_core_interface.py index 4698fc105e9e..32e2f1a57cd8 100644 --- a/python/sglang/srt/mem_cache/unified_cache/unified_tree_core_interface.py +++ b/python/sglang/srt/mem_cache/unified_cache/unified_tree_core_interface.py @@ -52,6 +52,7 @@ class EvictDeviceNextNodeResult(BaseEvictionResult): """ node_id: Optional[NodeId] = None + backup_kv: Optional[BackupKV] = None made_progress: bool = False unbacked_tokens: int = 0 mamba_backup_node_id: Optional[NodeId] = None @@ -183,6 +184,14 @@ def node_by_id(self, node_id: NodeId) -> UnifiedTreeNode: """ ... + @abstractmethod + def is_write_through_compatible(self) -> bool: + """Whether current host ownership satisfies write-through invariants. + + Read-only; the caller must have drained in-flight cache operations. + """ + ... + @abstractmethod def is_backuped(self, node_id: NodeId) -> bool: """Whether the node's KV is already backed up to host.""" @@ -510,6 +519,13 @@ def set_hicache_enabled(self) -> None: """Mark the host tier (HiCache) as wired.""" ... + def enable_swa_write_back_eviction_barrier(self) -> None: + """Enable a backend-managed barrier when needed. + + The Python core demotes through SWAComponent directly. Native cores + may return a backup action to the cache executor before eviction. + """ + @abstractmethod def set_host_memory_buffer_only(self) -> None: """Mark the host tier as buffer-only: one node staged per backup intent.""" diff --git a/python/sglang/srt/mem_cache/unified_radix_cache.py b/python/sglang/srt/mem_cache/unified_radix_cache.py index 7cc6549d3b9d..67f1d48698c2 100644 --- a/python/sglang/srt/mem_cache/unified_radix_cache.py +++ b/python/sglang/srt/mem_cache/unified_radix_cache.py @@ -287,6 +287,7 @@ def __init__( "declined_rate_limited": 0, "declined_anchor_lost": 0, "declined_device_covered": 0, + "declined_host_oversize": 0, "revoked_insufficient": 0, "revoked_full_miss": 0, "l3_demand_requests": 0, @@ -470,6 +471,7 @@ def init_hicache(self, server_args: ServerArgs, params: CacheInitParams) -> None sidecar_pool_specs=self.sidecar_pool_specs, host_pool_group=self.host_pool_group, swa_component=swa, + storage_prefetch_threshold=storage_prefetch_threshold, ) self.buffer_pipeline = BufferModePipeline( cache=self, @@ -485,11 +487,11 @@ def init_hicache(self, server_args: ServerArgs, params: CacheInitParams) -> None # content would drop-newest and punch storage holes. write_backlog_cap=2 * self.token_to_kv_pool_allocator.size_full, ) + # State initialization + if self.buffer_pipeline is not None: self.cache_controller.host_write_staged_tokens_fn = lambda: ( self.buffer_pipeline.write_staged_tokens_ ) - - # State initialization self.write_through_threshold = ( 1 if get_memory().hicache_write_policy == "write_through" else 2 ) @@ -497,6 +499,13 @@ def init_hicache(self, server_args: ServerArgs, params: CacheInitParams) -> None self.cache_controller is not None and self.cache_controller.write_policy == "write_back" ) + # Preserve the SWA host window before device eviction makes it unrecoverable. + if ( + get_memory().enable_unified_memory + and self.host_memory_mode == "cache" + and self.tree_core.has_swa_host_pool + ): + self.tree_core.enable_swa_write_back_eviction_barrier() # Pre-seed the logical dropped-tokens series. if self.metrics_collector is not None and self.cache_controller is not None: reasons = ["host_pressure"] @@ -749,45 +758,74 @@ def _accumulate_tracker( def _evict_device_next_node( self, component_type: ComponentType, tracker: dict[ComponentType, int] ) -> tuple[Optional[NodeId], bool]: - """Advance the eviction walk one node, consuming its step result.""" - result = self.tree_core.evict_device_next_node(component_type, tracker) - if result.mamba_backup_node_id is not None: - assert component_type == ComponentType.MAMBA and result.node_id is None - assert ( - not result.device_frees and not result.host_frees and not result.tracker - ) - # Reserve a host state slot and wait for the backup acknowledgment - # before freeing device state. If allocation fails, eviction still - # proceeds to make room on the device. - node_id = result.mamba_backup_node_id - mamba_host_pool = self.host_pool_group.get_pool(PoolName.MAMBA) - if mamba_host_pool is not None and mamba_host_pool.available_size() < 1: - self.evict_host(1, ComponentType.MAMBA) - self.backup_node_for_write_back(node_id) - result = self.tree_core.finish_mamba_state_eviction(node_id) - elif result.swa_backup_node_id is not None: - assert component_type == ComponentType.SWA and result.node_id is None - assert ( - not result.device_frees and not result.host_frees and not result.tracker - ) - # The backup can cover several unbacked SWA segments. Reserve host - # space for the whole window before copying it, then resume eviction - # even if host allocation fails. - node_id = result.swa_backup_node_id - needed = result.swa_backup_num_tokens - swa_host_pool = self.host_pool_group.get_pool(PoolName.SWA) - if swa_host_pool is not None and swa_host_pool.available_size() < needed: - self.evict_host(needed, ComponentType.SWA) - self.backup_node_for_write_back(node_id) - result = self.tree_core.finish_swa_state_eviction(node_id) - self._free_values(result.device_frees, result.host_frees) - if self._tracks_write_through_unbacked_evictions(): - self._record_dropped_tokens( - result.unbacked_tokens, - reason="write_through_unbacked_eviction", - ) - self._accumulate_tracker(tracker, result.tracker) - return result.node_id, result.made_progress + """Advance the walk, completing pending host-backup barriers.""" + while True: + result = self.tree_core.evict_device_next_node(component_type, tracker) + if result.mamba_backup_node_id is not None: + assert component_type == ComponentType.MAMBA and result.node_id is None + assert ( + not result.device_frees + and not result.host_frees + and not result.tracker + ) + # Reserve a host state slot and wait for the backup acknowledgment + # before freeing device state. If allocation fails, eviction still + # proceeds to make room on the device. + node_id = result.mamba_backup_node_id + mamba_host_pool = self.host_pool_group.get_pool(PoolName.MAMBA) + if mamba_host_pool is not None and mamba_host_pool.available_size() < 1: + self.evict_host(1, ComponentType.MAMBA) + self.backup_node_for_write_back(node_id) + result = self.tree_core.finish_mamba_state_eviction(node_id) + elif result.swa_backup_node_id is not None: + assert component_type == ComponentType.SWA and result.node_id is None + assert ( + not result.device_frees + and not result.host_frees + and not result.tracker + ) + # The backup can cover several unbacked SWA segments. Reserve host + # space for the whole window before copying it, then resume eviction + # even if host allocation fails. + node_id = result.swa_backup_node_id + needed = result.swa_backup_num_tokens + swa_host_pool = self.host_pool_group.get_pool(PoolName.SWA) + if ( + swa_host_pool is not None + and swa_host_pool.available_size() < needed + ): + self.evict_host(needed, ComponentType.SWA) + self.backup_node_for_write_back(node_id) + result = self.tree_core.finish_swa_state_eviction(node_id) + self._free_values(result.device_frees, result.host_frees) + if self._tracks_write_through_unbacked_evictions(): + self._record_dropped_tokens( + result.unbacked_tokens, + reason="write_through_unbacked_eviction", + ) + self._accumulate_tracker(tracker, result.tracker) + if result.backup_kv is None: + return result.node_id, result.made_progress + + assert result.node_id is None + assert self.buffer_pipeline is None, ( + "SWA write-back eviction barriers are cache-mode only" + ) + written = self._execute_and_commit_kv_backup( + result.backup_kv, write_back=True + ) + if written <= 0: + node_id = result.backup_kv.node_ids[0] + logger.warning( + "write_back: auxiliary backup failed under host pressure " + "(component=%s, node=%d); dropping only the component", + component_type.name, + node_id, + ) + # Match the Python SWA demotion: preserve FULL and descendants; + # the resumed native walk tombstones this component after one try. + continue + self.writing_check(write_back=True) def _evict_device_leaf( self, node_id: NodeId, tracker: dict[ComponentType, int] @@ -1366,14 +1404,17 @@ def _retraction_device_transfers( component_transfers: dict[ComponentType, list[PoolTransfer]] = {} if self.supports_swa(): - kv_cache = self.token_to_kv_pool_allocator.get_kvcache() assert self.sliding_window_size is not None window_start = max(0, num_tokens - self.sliding_window_size) window_start = window_start // self.page_size * self.page_size window_indices = self.req_to_token_pool.req_to_token[ req.kv.req_pool_idx, window_start:num_tokens ].to(torch.int64) - swa_indices = kv_cache.translate_loc_from_full_to_swa(window_indices) + swa_indices = ( + self.token_to_kv_pool_allocator.translate_swa_indices_for_transfer( + window_indices + ) + ) assert bool((swa_indices > 0).all()), ( f"unmapped SWA window positions for request {req.rid}" ) @@ -1392,13 +1433,12 @@ def _retraction_device_transfers( for transfers in component_transfers.values() for transfer in transfers ] - extra_transfers.extend( - self._build_sidecar_transfers( - CacheTransferPhase.BACKUP_HOST, - kv_transfer, - component_transfers, - ) + sidecar_transfers = self._build_sidecar_transfers( + CacheTransferPhase.BACKUP_HOST, + kv_transfer, + component_transfers, ) + extra_transfers.extend(sidecar_transfers) return full_indices, extra_transfers def _reclaim_retraction_host(self, num_tokens: int) -> int: @@ -2397,7 +2437,9 @@ def _check_hybrid_prefetch_result( self.buffer_pipeline.release_anchor_lock(request) del self.ongoing_prefetch[request] self.cache_controller.prefetch_tokens_occupied -= ( - self._prefetch_occupied_span(prefetch_key, host_indices) + self._prefetch_occupied_span( + prefetch_key, host_indices, operation=operation + ) ) self.prefetch_loaded_tokens_by_reqid[request] = 0 self.prefetch_loaded_storage_start_by_reqid.pop(request, None) @@ -2584,7 +2626,7 @@ def release_aborted_request(self, request: CacheRequestHandle) -> None: # Buffer mode granted occupancy at hit-alloc, sized to the bounce; # cache mode reserved the requested span at enqueue. self.cache_controller.prefetch_tokens_occupied -= self._prefetch_occupied_span( - prefetch_key, host_indices + prefetch_key, host_indices, operation=operation ) def _invalidate_absent_from_hit_query(self, operation) -> None: @@ -2621,11 +2663,18 @@ def _account_prefetch_outcome(self, operation, revoked: bool) -> None: def prefetch_outcome_stats_snapshot(self) -> dict: return self._prefetch_outcome_stats.copy() - def _prefetch_occupied_span(self, prefetch_key, host_indices) -> int: + def _prefetch_occupied_span( + self, prefetch_key, host_indices, *, operation=None + ) -> int: """Occupancy units held by a prefetch: cache mode reserves the requested span at enqueue; buffer mode grants at hit-alloc, sized to the allocation (0 while still querying / parked).""" if self.host_memory_mode == "buffer_only": + if ( + operation is not None + and operation.buffer_host_occupied_units is not None + ): + return operation.buffer_host_occupied_units return len(host_indices) if host_indices is not None else 0 return len(prefetch_key) @@ -2742,7 +2791,9 @@ def revoke_pending_prefetch(self, request: CacheRequestHandle) -> None: cc.prefetch_tokens_occupied = max( 0, cc.prefetch_tokens_occupied - - self._prefetch_occupied_span(prefetch_key, _host_indices), + - self._prefetch_occupied_span( + prefetch_key, _host_indices, operation=operation + ), ) def _drain_storage_control_queues_impl( @@ -2827,20 +2878,24 @@ def _try_alloc_storage_hit(operation) -> bool: else: aux_hit_tokens = hit_tokens alloc_len = hit_tokens - host_indices = cc.mem_pool_host.alloc(alloc_len) - if host_indices is None: - self.evict_host(alloc_len) - host_indices = cc.mem_pool_host.alloc(alloc_len) - if host_indices is None and not buffer_mode: - # Memory-pressure fallback: a shorter page-aligned prefix. - # (Cache mode only — buffer mode parks for the full hit.) - available_size = cc.mem_pool_host.available_size() - alloc_len = min( - hit_tokens, - available_size - (available_size % self.page_size), + if buffer_mode and not cc.can_fit_prefetch_host_buffers( + operation, alloc_len + ): + self._prefetch_outcome_stats["declined_host_oversize"] += 1 + logger.warning( + "HiCache buffer prefetch declined req=%s: the Full/sidecar " + "bounce cannot fit in an empty host arena", + request, ) - if alloc_len >= self.prefetch_threshold: - host_indices = cc.mem_pool_host.alloc(alloc_len) + self.revoke_pending_prefetch(request) + return True + host_indices, alloc_len = cc.allocate_storage_hit( + operation, + hit_tokens, + allow_partial=not buffer_mode, + min_tokens=self.prefetch_threshold, + evict_host=self.evict_host, + ) if host_indices is None: if buffer_mode: # Parked ops hold no pin: release and re-take at the next @@ -2890,7 +2945,12 @@ def _try_alloc_storage_hit(operation) -> bool: operation.host_indices = host_indices self.ongoing_prefetch[request] = info._replace(host_indices=host_indices) if buffer_mode: - cc.prefetch_tokens_occupied += alloc_len + operation.buffer_host_occupied_units = ( + self.buffer_pipeline.host_allocation_units( + host_indices, operation.pool_transfers + ) + ) + cc.prefetch_tokens_occupied += operation.buffer_host_occupied_units cc.prefetch_buffer.put(operation) return True diff --git a/rust/sglang-radix-tree/src/components/swa.rs b/rust/sglang-radix-tree/src/components/swa.rs index bfa547549fd9..3bde51e79de2 100644 --- a/rust/sglang-radix-tree/src/components/swa.rs +++ b/rust/sglang-radix-tree/src/components/swa.rs @@ -802,6 +802,20 @@ impl TreeComponent for SwaComponent { break 'step None; } } + if tree_core.is_write_back + && tree_core.swa_write_back_eviction_barrier_enabled + && tree_core.component_state(SWA).evict_device_last_backup != Some(x) + && !tree_core.arena.node(x).backuped() + { + // A later Full backup cannot recover SWA data after this + // internal node is tombstoned. Pause on the same cursor so + // the Controller can preserve the dirty path first. + tree_core.component_state_mut(SWA).evict_device_backup_node = Some(x); + tree_core.component_state_mut(SWA).evict_device_last_backup = Some(x); + cursor = Some(x); + break 'step None; + } + // Internal nodes are tombstoned inline (no IO). tree_core.evict_component_and_detach_lru_( x, ct, diff --git a/rust/sglang-radix-tree/src/python_bindings.rs b/rust/sglang-radix-tree/src/python_bindings.rs index fa6c04c07cff..f580134ff398 100644 --- a/rust/sglang-radix-tree/src/python_bindings.rs +++ b/rust/sglang-radix-tree/src/python_bindings.rs @@ -847,6 +847,7 @@ fn tracker_to_py(tracker: HashMap) -> HashMap { #[pyclass(get_all)] pub struct EvictDeviceNextNodeResultBinding { node_id: Option, + backup_kv: Option>, mamba_backup_node_id: Option, swa_backup_node_id: Option, swa_backup_num_tokens: usize, @@ -1430,11 +1431,16 @@ impl TreeCoreBinding { let (node_id, result) = py.allow_threads(move || self.core().evict_device_next_node(ct, &baseline)); let made_progress = node_id.is_some() + || result.backup_kv.is_some() || result.mamba_backup_node_id.is_some() || result.swa_backup_node_id.is_some() || !result.tracker.is_empty(); Ok(EvictDeviceNextNodeResultBinding { node_id, + backup_kv: result + .backup_kv + .map(|backup| cache_action_to_py(py, CacheAction::BackupKV(backup))) + .transpose()?, mamba_backup_node_id: result.mamba_backup_node_id, swa_backup_node_id: result.swa_backup_node_id, swa_backup_num_tokens: result.swa_backup_num_tokens, @@ -1473,6 +1479,7 @@ impl TreeCoreBinding { ) -> PyResult { Ok(EvictDeviceNextNodeResultBinding { node_id: None, + backup_kv: None, mamba_backup_node_id: None, swa_backup_node_id: None, swa_backup_num_tokens: 0, @@ -1642,6 +1649,11 @@ impl TreeCoreBinding { py.allow_threads(|| self.core().enable_hicache) } + /// Preserve dirty SWA data before cache-mode write-back eviction. + fn enable_swa_write_back_eviction_barrier(&self, py: Python<'_>) { + py.allow_threads(|| self.core().enable_swa_write_back_eviction_barrier()); + } + /// Mark the SWA host pool as wired (HiCache). fn set_has_swa_host_pool(&self, py: Python<'_>) { py.allow_threads(|| self.core().set_has_swa_host_pool()); @@ -2039,6 +2051,10 @@ impl TreeCoreBinding { Ok(()) } + fn is_write_through_compatible(&self, py: Python<'_>) -> bool { + py.allow_threads(|| self.core().is_write_through_compatible()) + } + /// Set the write-back (vs write-through) policy; decided at HiCache init. fn set_is_write_back(&self, py: Python<'_>, is_write_back: bool) { py.allow_threads(|| self.core().is_write_back = is_write_back); @@ -3051,6 +3067,11 @@ macro_rules! tree_core_binding { catch_native_panic(|| Ok(self.inner.enable_hicache(py))) } + /// Preserve dirty SWA data before cache-mode write-back eviction. + fn enable_swa_write_back_eviction_barrier(&self, py: Python<'_>) { + self.inner.enable_swa_write_back_eviction_barrier(py) + } + /// Mark the SWA host pool as wired (HiCache). fn set_has_swa_host_pool(&self, py: Python<'_>) -> PyResult<()> { catch_native_panic(|| { @@ -3329,6 +3350,10 @@ macro_rules! tree_core_binding { catch_native_panic(|| self.inner.dec_host_lock_ref(py, node_id, params)) } + fn is_write_through_compatible(&self, py: Python<'_>) -> bool { + self.inner.is_write_through_compatible(py) + } + /// Set the write-back (vs write-through) policy; decided at HiCache init. fn set_is_write_back(&self, py: Python<'_>, is_write_back: bool) -> PyResult<()> { catch_native_panic(|| { diff --git a/rust/sglang-radix-tree/src/tests/unified_tree_core.rs b/rust/sglang-radix-tree/src/tests/unified_tree_core.rs index 24f6a264055b..0cd2e63c6083 100644 --- a/rust/sglang-radix-tree/src/tests/unified_tree_core.rs +++ b/rust/sglang-radix-tree/src/tests/unified_tree_core.rs @@ -4103,23 +4103,22 @@ fn insert_host_allows_a_suffix_under_an_unbacked_write_back_parent() { fn insert_host_drops_a_suffix_under_an_unbacked_write_through_parent() { let mut tc = core(); tc.insert(&insert_params(&vec![1, 2], &[10, 11])); - let root = tc.arena.root(); + let anchor = tc + .match_prefix(&match_params(&vec![1, 2])) + .best_match_node_id; let nodes_before = tc.arena.len(); let result = tc .insert_host( - tc.arena.node(root).id, + anchor, /* extra_key = */ None, - vec![1, 2, 3, 4], - Tensor::from_slice(&[100i64, 101, 102, 103]), - vec!["h0", "h1", "h2", "h3"] - .into_iter() - .map(String::from) - .collect(), + vec![3, 4], + Tensor::from_slice(&[102i64, 103]), + vec!["h2", "h3"].into_iter().map(String::from).collect(), ) .expect("live test node"); - assert_eq!(result.prefix_len, 2); - assert_eq!(result.total_len, 4); + assert_eq!(result.prefix_len, 0); + assert_eq!(result.total_len, 2); assert_eq!(result.inserted_host_node, None); assert!(result.host_insert_dropped); assert!(result.cache_actions.is_empty()); @@ -4127,7 +4126,7 @@ fn insert_host_drops_a_suffix_under_an_unbacked_write_through_parent() { } #[test] -fn insert_host_drop_preserves_split_actions_and_lengths() { +fn insert_host_refill_preserves_split_actions_and_lengths() { let mut tc = core(); tc.insert(&insert_params(&vec![1, 2, 3], &[10, 11, 12])); let leaf = tc @@ -4146,10 +4145,10 @@ fn insert_host_drop_preserves_split_actions_and_lengths() { ) .expect("live test node"); - assert_eq!(result.prefix_len, 1); + assert_eq!(result.prefix_len, 0); assert_eq!(result.total_len, 2); - assert_eq!(result.inserted_host_node, None); - assert!(result.host_insert_dropped); + assert!(result.inserted_host_node.is_some()); + assert!(!result.host_insert_dropped); assert!(matches!( result.cache_actions.as_slice(), [CacheAction::ReplaceWriteThroughOnNodeSplit { ack_id, .. }] if *ack_id == leaf @@ -4252,7 +4251,7 @@ fn insert_host_full_match_reports_only_a_backuped_node() { .match_prefix(&match_params(&vec![1, 2])) .best_match_node_id; let root = tc.arena.root(); - // The device-only match reports no host node. + // The device-only match consumes the supplied host slots. let result = tc .insert_host( tc.arena.node(root).id, @@ -4262,14 +4261,15 @@ fn insert_host_full_match_reports_only_a_backuped_node() { vec!["h0".to_string(), "h1".to_string()], ) .expect("live test node"); - assert_eq!(result.prefix_len, 2); - assert_eq!(result.inserted_host_node, None); + assert_eq!(result.prefix_len, 0); + assert_eq!(result.inserted_host_node, Some(leaf)); assert!(!result.host_insert_dropped); - // Once backuped, the same insert reports the node. - tc.arena.set_host_value( - tc.arena.resolve(leaf).expect("live test node"), - FULL, - Tensor::from_slice(&[20i64, 21]), + // The refill owns these rows; a repeat reports only the duplicate prefix. + assert!( + tc.arena + .node(tc.arena.resolve(leaf).expect("live test node")) + .host_value(FULL) + .equal(&Tensor::from_slice(&[100i64, 101])) ); let result = tc .insert_host( diff --git a/rust/sglang-radix-tree/src/unified_tree_core.rs b/rust/sglang-radix-tree/src/unified_tree_core.rs index 609e57425e8e..303202fbbb04 100644 --- a/rust/sglang-radix-tree/src/unified_tree_core.rs +++ b/rust/sglang-radix-tree/src/unified_tree_core.rs @@ -186,7 +186,9 @@ pub struct InsertParams<'k, K: ChildKeyType> { pub struct InsertResult { /// The incoming pages use another chain rotation and were not adopted. pub rotation_tail_declined: bool, - /// Tokens of the insert key that overlapped existing nodes. + /// Incoming value rows the caller may release as duplicates. For a + /// write-through insert_host, structurally matched nodes that lacked Full + /// host state are refilled and therefore excluded from this count. pub prefix_len: usize, /// The inserted key's full (page-aligned) length. pub total_len: usize, @@ -411,7 +413,7 @@ pub struct PoolTransferResult { } /// A device->host backup work item for the cache to execute. -#[derive(Default)] +#[derive(Default, Debug)] pub struct BackupKV { /// Backup these nodes device->host in order, stopping at the first failure; the /// caller orders them parent-before-child for write-through and child-first for @@ -506,6 +508,11 @@ pub struct ComponentState { /// leaf may be freed: the leaf's parent for Full, the LRU predecessor for /// SWA and Mamba. pub(crate) evict_device_cursor: Option, + /// Internal node whose component value must be backed up before the walk + /// can tombstone it. The Controller consumes this request between steps. + pub(crate) evict_device_backup_node: Option, + /// A resumed victim is tombstoned after its best-effort backup attempt. + pub(crate) evict_device_last_backup: Option, /// Internal component victim waiting for the controller's host backup attempt. /// A generation-checked handle survives host eviction during that I/O. pub(crate) evict_device_pending_node: Option, @@ -579,6 +586,7 @@ pub struct EvictionStepResult { pub tracker: HashMap, pub device_frees: HashMap>, pub host_frees: HashMap>, + pub backup_kv: Option, /// Full device tokens freed without a host copy during this device step. pub unbacked_tokens: usize, /// Back up this internal Mamba state before resuming its device tombstone. @@ -632,6 +640,8 @@ pub struct UnifiedTreeCore { pub(crate) enable_external_cache_linker: bool, /// Whether the cache wired a host SWA pool (HiCache). pub(crate) has_swa_host_pool: bool, + /// Whether dirty internal SWA nodes must be backed up before eviction. + pub(crate) swa_write_back_eviction_barrier_enabled: bool, /// Whether tree mutations emit BlockStored/BlockRemoved events. pub(crate) enable_kv_cache_events: bool, /// Queued placement events, drained by take_events. @@ -744,6 +754,8 @@ impl UnifiedTreeCore { state.is_evict_device_ongoing = true; state.evict_device_request_cnt = request_cnt; state.evict_device_cursor = None; + state.evict_device_backup_node = None; + state.evict_device_last_backup = None; state.evict_device_pending_node = None; state.evict_device_pending_num_tokens = 0; } @@ -758,6 +770,8 @@ impl UnifiedTreeCore { ); state.is_evict_device_ongoing = false; state.evict_device_cursor = None; + state.evict_device_backup_node = None; + state.evict_device_last_backup = None; state.evict_device_pending_node = None; state.evict_device_pending_num_tokens = 0; } @@ -826,6 +840,7 @@ impl UnifiedTreeCore { enable_storage: false, enable_external_cache_linker: false, has_swa_host_pool: params.has_swa_host_pool, + swa_write_back_eviction_barrier_enabled: false, enable_kv_cache_events: params.enable_kv_cache_events, kv_event_queue: Vec::new(), namespaced_event_hashes: HashMap::new(), @@ -2382,6 +2397,15 @@ impl UnifiedTreeCore { &mut result.device_frees, &mut result.host_frees, ); + let backup_node = self + .component_state_mut(component_type) + .evict_device_backup_node + .take(); + if let Some(backup_node) = backup_node { + assert!(node_id.is_none()); + result.backup_kv = + Some(self.build_backup_kv_action_(self.arena.node(backup_node), true)); + } result.unbacked_tokens = self.tracked_unbacked_tokens.take().unwrap(); let state = self.component_state(component_type); if component_type == MAMBA { @@ -2548,7 +2572,7 @@ impl UnifiedTreeCore { // A failed backup never issues the D->H copy, so the subtree root has // no host state and no in-flight DMA reading its device slots. assert!(!node.backuped() && node.write_through_pending_id.is_none()); - if node.is_host_locked() { + if node.is_host_locked() || node.is_load_back_pending() { return Ok((false, result)); } } @@ -2562,7 +2586,11 @@ impl UnifiedTreeCore { .collect(); while let Some(cur_id) = stack.pop() { let cur = self.arena.node(cur_id); - if cur.is_device_locked() || cur.is_host_locked() { + if cur.is_device_locked() + || cur.is_host_locked() + || cur.write_through_pending_id.is_some() + || cur.is_load_back_pending() + { return Ok((false, result)); } descendants.push(cur_id); @@ -3133,6 +3161,11 @@ impl UnifiedTreeCore { self.enable_hicache = true; } + /// Preserve dirty internal SWA nodes before cache-mode write-back eviction. + pub fn enable_swa_write_back_eviction_barrier(&mut self) { + self.swa_write_back_eviction_barrier_enabled = true; + } + /// Mark the host tier as buffer-only; wired after the host pools are built. pub fn set_host_memory_buffer_only(&mut self) { self.is_host_memory_buffer_only = true; @@ -3438,8 +3471,14 @@ impl UnifiedTreeCore { }); } - // Walk cursor: atoms of `key` already matched (also the running prefix length). + // Walk cursor: atoms of `key` already matched structurally. let mut matched_length = 0; + // For write-through, only a leading already-host-resident prefix is a + // duplicate that the caller may free. Structurally matching device-only + // nodes consume fresh host indices and must not contribute to this count. + let mut host_prefix_len = 0; + let mut refilled_unbacked_node = false; + let mut inserted_host_node = None; let mut cache_actions: Vec = Vec::new(); while matched_length < total_len { let Some(child_id) = self.arena.child_on_page_in_namespace( @@ -3451,6 +3490,7 @@ impl UnifiedTreeCore { }; node_id = child_id; self.touch_node_(node_id); + let matched_start = matched_length; let node = self.arena.node(node_id); let prefix_len = key.match_len(matched_length, &node.key, self.page_size); let node_key_len = node.key.atom_len(); @@ -3463,13 +3503,50 @@ impl UnifiedTreeCore { cache_actions.push(action); } } + + if !self.is_write_back { + if self.arena.node(node_id).has_host_value(FULL) { + assert!( + !refilled_unbacked_node, + "insert_host: write-through path has a backed node below an unbacked node" + ); + host_prefix_len = matched_length; + continue; + } + + refilled_unbacked_node = true; + self.arena.set_host_value( + node_id, + FULL, + host_value + .narrow(0, matched_start as i64, prefix_len as i64) + .copy(), + ); + if self.arena.node(node_id).hash_value.is_none() { + let first_page = matched_start / self.page_size; + let last_page = matched_length / self.page_size; + self.arena.node_mut(node_id).hash_value = + Some(hash_value[first_page..last_page].to_vec()); + } + self.update_evictable_leaf_sets_(node_id); + if let Some(parent_id) = self.arena.node(node_id).try_parent() { + self.update_evictable_leaf_sets_(parent_id); + } + self.update_full_coexisting_host_tracking_(node_id); + self.record_store_event_(node_id, StorageMedium::Cpu, /* session_id = */ None); + inserted_host_node = Some(self.arena.node(node_id).id); + } } let mut result = InsertResult { - prefix_len: matched_length, + prefix_len: if self.is_write_back { + matched_length + } else { + host_prefix_len + }, total_len, last_device_node_id: None, - inserted_host_node: None, + inserted_host_node, host_insert_dropped: false, rotation_tail_declined: false, mamba_exist: false, @@ -3479,7 +3556,7 @@ impl UnifiedTreeCore { }; if matched_length == total_len { let node = self.arena.node(node_id); - if !node.is_root() && node.has_host_value(FULL) { + if result.inserted_host_node.is_none() && !node.is_root() && node.has_host_value(FULL) { result.inserted_host_node = Some(self.arena.node(node_id).id); } return Ok(result); @@ -4547,6 +4624,31 @@ impl UnifiedTreeCore { && self.arena.has_host_value(node_idx, component_type)) } + /// Read-only check of the host invariants required by write-through. + /// The administrative caller must first drain in-flight cache operations. + pub fn is_write_through_compatible(&self) -> bool { + for node_id in self.collect_all_nodes_() { + let node = self.arena.node(node_id); + if node.is_root() { + continue; + } + if !node.has_host_value(FULL) { + if self.components.iter().any(|component| { + let ct = component.component_type(); + ct != FULL && node.has_host_value(ct) + }) { + return false; + } + } else { + let parent = self.arena.node(node.parent()); + if !parent.is_root() && !parent.has_host_value(FULL) { + return false; + } + } + } + true + } + /// Verify tree-structure, leaf-set, LRU, size, and ongoing-op invariants; raise /// AssertionError on any violation. ongoing_* args are (id, node_id) pairs. pub fn sanity_check( diff --git a/test/registered/unit/disaggregation/test_unified_memory_move_gate.py b/test/registered/unit/disaggregation/test_unified_memory_move_gate.py index 201d8eee2087..5a23bc1386f1 100644 --- a/test/registered/unit/disaggregation/test_unified_memory_move_gate.py +++ b/test/registered/unit/disaggregation/test_unified_memory_move_gate.py @@ -150,6 +150,7 @@ class _Peer: def __init__(self, gate, host_gate=None): self.lazy_compaction = True self._free_phys_pages = [0, 1, 2, 3] # only len() is read + self._pending_hicache_load_pages = 0 self.entry_bytes_per_page = 512 self.disagg_move_gate = gate self.host_transfer_move_gate = host_gate diff --git a/test/registered/unit/mem_cache/test_buffer_mode_sidecar.py b/test/registered/unit/mem_cache/test_buffer_mode_sidecar.py index acd11ff09c0f..eb26245264be 100644 --- a/test/registered/unit/mem_cache/test_buffer_mode_sidecar.py +++ b/test/registered/unit/mem_cache/test_buffer_mode_sidecar.py @@ -37,7 +37,11 @@ class TestBufferModeSidecar(unittest.TestCase): def _swa_component(): return SimpleNamespace( full_window_pages=2, - _swa_kv_pool_host=SimpleNamespace(page_size=2, size=8), + _swa_kv_pool_host=SimpleNamespace( + page_size=2, + size=8, + shared_allocation_domain=None, + ), ) @staticmethod @@ -137,6 +141,7 @@ def test_write_stages_and_persists_dsv4_full_and_swa_sidecars(self): ] controller = MagicMock() + controller.mem_pool_host.anchor_entry.host_pool.shared_allocation_domain = None controller.mem_pool_host.entry_map = { PoolName.SWA: SimpleNamespace( host_pool=SimpleNamespace(page_size=page_size) @@ -234,12 +239,13 @@ def test_completed_prefetch_keeps_dsv4_full_and_swa_sidecars_for_h2d(self): ) for spec in self._dsv4_specs() ] + host_indices = torch.arange(4, dtype=torch.int64) operation = SimpleNamespace( id=23, pool_transfers=sidecars, storage_start=0, + buffer_host_occupied_units=len(host_indices), ) - host_indices = torch.arange(4, dtype=torch.int64) req_id = CacheRequestHandle("sidecar-prefetch", 0) cache = MagicMock() diff --git a/test/registered/unit/mem_cache/test_decode_radix_lock_ref.py b/test/registered/unit/mem_cache/test_decode_radix_lock_ref.py index a3b5792b02eb..832a59630e27 100644 --- a/test/registered/unit/mem_cache/test_decode_radix_lock_ref.py +++ b/test/registered/unit/mem_cache/test_decode_radix_lock_ref.py @@ -545,9 +545,8 @@ def test_pop_preallocated_rechecks_budget_after_lock(self): scheduler.output_streamer = MagicMock() queue.scheduler = scheduler - # The 4-token match is locked, then capped to zero because the whole - # 8-token request is inside the SWA window. Admission rejection must - # still release the original matched-node lock. + # Admission rejection must release the matched-node lock, including + # when the SWA lock was already released for fresh tail allocation. queue._allocatable_token_budgets = MagicMock(return_value=3) preallocated, failed = queue.pop_preallocated() diff --git a/test/registered/unit/mem_cache/test_hicache_staged_write_back_dispatch.py b/test/registered/unit/mem_cache/test_hicache_staged_write_back_dispatch.py index 01fb49a4fa18..2f9ecbb15654 100644 --- a/test/registered/unit/mem_cache/test_hicache_staged_write_back_dispatch.py +++ b/test/registered/unit/mem_cache/test_hicache_staged_write_back_dispatch.py @@ -225,6 +225,7 @@ def test_hybrid_load_forwards_merged_pool_transfers(self): op.pool_transfers, ) controller.mem_pool_host = _host_group_stub([], can_use_write_back_jit=False) + controller.mem_pool_device_allocator = mock.Mock() controller._l2_transfers.side_effect = lambda *args: ( HybridCacheController._l2_transfers(controller, *args) ) @@ -1139,6 +1140,7 @@ def test_write_back_jit_hybrid_write_keeps_extra_host_indices_on_cpu(self): captured, can_use_write_back_jit=True ) controller.mem_pool_device = None + controller.mem_pool_device_allocator = mock.Mock() controller.ack_write_queue = [] controller.move_hybrid_indices = mock.Mock( side_effect=AssertionError( @@ -1173,6 +1175,7 @@ def test_hybrid_write_moves_indices_without_write_back_jit(self): captured, can_use_write_back_jit=False ) controller.mem_pool_device = None + controller.mem_pool_device_allocator = mock.Mock() controller.ack_write_queue = [] controller.move_hybrid_indices = mock.Mock( return_value=(op.host_indices, op.device_indices, op.pool_transfers) @@ -1230,6 +1233,7 @@ def backup_from_device_all_layer( ] ) controller.mem_pool_device = None + controller.mem_pool_device_allocator = mock.Mock() controller.ack_write_queue = [] with mock.patch.object(transfer_module, "device_module", _FakeDeviceModule): controller.l2_transfer_engine = L2TransferEngine("kernel") @@ -1296,6 +1300,7 @@ def backup_from_device_all_layer( controller.io_backend = "kernel" controller.mem_pool_host = FakeHostPool.__new__(FakeHostPool) controller.mem_pool_device = None + controller.mem_pool_device_allocator = mock.Mock() controller.device = "cuda" controller.ack_write_queue = [] controller.move_indices = mock.Mock( @@ -1332,6 +1337,7 @@ def backup_from_device_all_layer( controller.io_backend = "kernel" controller.mem_pool_host = FakeHostPool.__new__(FakeHostPool) controller.mem_pool_device = None + controller.mem_pool_device_allocator = mock.Mock() controller.device = "cuda" controller.ack_write_queue = [] controller.move_indices = mock.Mock( diff --git a/test/registered/unit/mem_cache/test_hybrid_pool_assembler.py b/test/registered/unit/mem_cache/test_hybrid_pool_assembler.py index 7e970d0d9d96..299b5f07ba30 100644 --- a/test/registered/unit/mem_cache/test_hybrid_pool_assembler.py +++ b/test/registered/unit/mem_cache/test_hybrid_pool_assembler.py @@ -1,18 +1,23 @@ """Unit tests for hybrid HiCache pool assembly.""" import unittest +from queue import Queue from types import SimpleNamespace from unittest.mock import MagicMock, patch import msgspec import torch -from sglang.srt.mem_cache.base_prefix_cache import EvictParams -from sglang.srt.mem_cache.hicache_storage import PoolName +from sglang.srt.mem_cache.base_prefix_cache import CacheRequestHandle, EvictParams +from sglang.srt.mem_cache.hicache_storage import PoolHitPolicy, PoolName, PoolTransfer from sglang.srt.mem_cache.hybrid_cache import hybrid_pool_assembler from sglang.srt.mem_cache.hybrid_cache.host_pool_config import ( prepare_host_pool_config, ) +from sglang.srt.mem_cache.hybrid_cache.hybrid_cache_controller import ( + HybridCacheController, + PrefetchOperation, +) from sglang.srt.mem_cache.hybrid_cache.hybrid_pool_assembler import ( StackBuildResult, _check_declared_pools_present, @@ -29,6 +34,7 @@ build_hybrid_swa_group, ) from sglang.srt.mem_cache.memory_pool import DSATokenToKVPool, HybridLinearKVPool +from sglang.srt.mem_cache.pool_host import HostPoolGroup, PoolEntry from sglang.srt.mem_cache.pool_host import dsa as pool_host_dsa from sglang.srt.mem_cache.pool_host.host_pool_decl import ( HostPoolDecl, @@ -1003,6 +1009,113 @@ def _build_unified_host_pair(bundle): class TestUnifiedPageEnvelopeHostPool(CustomTestCase): + def test_sidecar_read_error_is_only_a_cache_miss_with_unified_memory(self): + from sglang.srt.runtime_context import publish, reset_context + from sglang.srt.server_args import ServerArgs + + seen_prefixes = [] + + class FailingStorage: + def batch_get_v2(self, transfers, *, extra_info=None): + seen_prefixes.append(extra_info.prefix_keys) + raise ValueError("sidecar read failed") + + self.addCleanup(reset_context) + for unified in (False, True): + with self.subTest(unified=unified): + reset_context() + publish( + ServerArgs(model_path="dummy", enable_unified_memory=unified), + role="tokenizer", + ) + cc = HybridCacheController.__new__(HybridCacheController) + cc.storage_backend = FailingStorage() + cc.prefetch_sync_queue = Queue() + operation = PrefetchOperation( + CacheRequestHandle("r", 0), + [1], + pool_transfers=[PoolTransfer(name=PoolName.SWA)], + ) + operation.hash_value = ["h0"] + operation.prefix_keys = ["prefix"] + if unified: + with self.assertLogs(level="ERROR"): + cc._page_transfer_sidecar(operation, kv_completed_pages=1) + ack = cc.prefetch_sync_queue.get_nowait() + self.assertIs(ack.operation, operation) + self.assertEqual(ack.pool_hits, {}) + else: + with self.assertRaisesRegex(ValueError, "sidecar read failed"): + cc._page_transfer_sidecar(operation, kv_completed_pages=1) + self.assertTrue(cc.prefetch_sync_queue.empty()) + self.assertEqual(seen_prefixes, [["prefix", "h0"], ["prefix", "h0"]]) + + def test_shorter_prefetch_reserves_full_and_swa_without_mutating_probe_keys(self): + page_size = 4 + full_pool, swa_pool = _build_unified_host_pair( + _build_unified_swa_pool(page_size) + ) + self.addCleanup(full_pool.destroy) + self.addCleanup(swa_pool.destroy) + cc = HybridCacheController.__new__(HybridCacheController) + cc.page_size = page_size + cc.host_memory_mode = "cache" + cc.attn_cp_group = cc.attn_tp_group = cc.tp_group = None + cc.mem_pool_host = HostPoolGroup( + [ + PoolEntry(PoolName.KV, full_pool, None, None), + PoolEntry(PoolName.SWA, swa_pool, None, None), + ] + ) + hit_tokens = full_pool.available_size() + hashes = [str(i) for i in range(hit_tokens // page_size)] + transfer = PoolTransfer( + name=PoolName.SWA, + keys=hashes[-2:], + hit_policy=PoolHitPolicy.TRAILING_PAGES, + ) + operation = PrefetchOperation( + CacheRequestHandle("r", 0), + list(range(hit_tokens)), + pool_transfers=[transfer], + ) + operation.hash_value = hashes + self.assertIsNone(cc.alloc_prefetch_host_buffers(operation, hit_tokens)) + fitting = [ + n + for n in range(page_size, hit_tokens + 1, page_size) + if cc.can_fit_prefetch_host_buffers(operation, n, empty=False) + ] + self.assertTrue(fitting) + self.assertEqual(transfer.keys, hashes[-2:]) + length = max(fitting) + self.assertLess(length, hit_tokens) + host_indices = cc.alloc_prefetch_host_buffers(operation, length) + self.assertEqual(host_indices.numel(), length) + self.assertEqual(transfer.host_indices.numel(), 2 * page_size) + self.assertEqual( + transfer.keys, hashes[length // page_size - 2 : length // page_size] + ) + cc.free_prefetch_host_buffers(operation, host_indices) + + # A hit-time rematch can trim some or all FULL pages while the entire + # trailing SWA window still needs staging from the same shared arena. + for full_tokens in (page_size, 0): + with self.subTest(full_tokens=full_tokens): + transfer.keys = hashes[-2:] + operation.hash_value = list(hashes) + operation.storage_hit_count = hit_tokens + operation.sidecar_hash_values = None + cc.trim_prefetch_full_head(operation, hit_tokens - full_tokens) + self.assertTrue( + cc.can_fit_prefetch_host_buffers(operation, full_tokens) + ) + host_indices = cc.alloc_prefetch_host_buffers(operation, full_tokens) + self.assertEqual(host_indices.numel(), full_tokens) + self.assertEqual(transfer.host_indices.numel(), 2 * page_size) + self.assertEqual(transfer.keys, hashes[-2:]) + cc.free_prefetch_host_buffers(operation, host_indices) + def test_shared_arena_can_reuse_bytes_across_sides(self): page_size = 4 full_pool, swa_pool = _build_unified_host_pair( diff --git a/test/registered/unit/mem_cache/test_rust_tree_core_integration.py b/test/registered/unit/mem_cache/test_rust_tree_core_integration.py index d28df118ea11..4c252cc426e3 100644 --- a/test/registered/unit/mem_cache/test_rust_tree_core_integration.py +++ b/test/registered/unit/mem_cache/test_rust_tree_core_integration.py @@ -2117,7 +2117,9 @@ def reclaim(count, component): @pytest.mark.parametrize("guard", ["none", "host_lock", "pending_dma"]) def test_swa_host_pressure_retains_device_resident_backups(backend, guard): core, allocator = _swa_transfer_core(backend) - allocator.translate_loc_from_full_to_swa.side_effect = lambda values: values + 100 + allocator.translate_swa_indices_for_transfer.side_effect = lambda values: ( + values + 100 + ) core.is_write_back = True nodes = [] for tokens, values in (([1, 2], [10, 11]), ([1, 2, 3], [20, 21, 12])): @@ -2127,7 +2129,7 @@ def test_swa_host_pressure_retains_device_resident_backups(backend, guard): core.set_component_device_value( action.node_id, ComponentType.SWA, - allocator.translate_loc_from_full_to_swa(action.source_value), + allocator.translate_swa_indices_for_transfer(action.source_value), ) else: assert isinstance(action, (FreeDeviceKV, FreeDeviceKVFullOnly)) @@ -2277,7 +2279,7 @@ def test_swa_backup_resolves_relocated_full_virtual_ids(backend, unified, entryp core, allocator = _swa_transfer_core(backend, unified=unified) mapping = torch.arange(64) + 100 - allocator.translate_loc_from_full_to_swa.side_effect = lambda indices: mapping[ + allocator.translate_swa_indices_for_transfer.side_effect = lambda indices: mapping[ indices ] values = [11, 7, 20, 4, 9, 13] @@ -2302,10 +2304,10 @@ def test_swa_backup_resolves_relocated_full_virtual_ids(backend, unified, entryp ] }, ) - # Relocation changes kernel-facing SWA addresses while tree-owned Full + # Relocation changes physical SWA transfer addresses while tree-owned Full # virtual IDs and cached SWA physical snapshots stay unchanged. mapping[torch.tensor([11, 7, 13])] = torch.tensor([511, 407, 613]) - allocator.translate_loc_from_full_to_swa.reset_mock() + allocator.translate_swa_indices_for_transfer.reset_mock() if entrypoint == "backup_spec": full, auxiliary = core.build_backup_spec(nodes[-1]) assert full.tolist() == [13] @@ -2320,14 +2322,16 @@ def test_swa_backup_resolves_relocated_full_virtual_ids(backend, unified, entryp [511, 407, 613] if unified else [111, 107, 113] ) if unified: - allocator.translate_loc_from_full_to_swa.assert_called_once() - assert allocator.translate_loc_from_full_to_swa.call_args.args[0].tolist() == [ + allocator.translate_swa_indices_for_transfer.assert_called_once() + assert allocator.translate_swa_indices_for_transfer.call_args.args[ + 0 + ].tolist() == [ 11, 7, 13, ] else: - allocator.translate_loc_from_full_to_swa.assert_not_called() + allocator.translate_swa_indices_for_transfer.assert_not_called() assert core.get_component_device_value(nodes[0], ComponentType.SWA).tolist() == [ 111, 107, @@ -4516,6 +4520,77 @@ def test_write_through_threshold_assignment_reaches_the_core(): assert any(isinstance(action, BackupKV) for action in result.cache_actions) +@pytest.mark.parametrize("backup_nodes", ["leaf", "parent_and_leaf", "parent"]) +def test_unified_swa_backup_when_full_already_has_a_host_copy(backup_nodes): + from sglang.srt.mem_cache.unified_memory_pool import init_unified_swa_pools + + bundle = init_unified_swa_pools( + device="cpu", + kv_cache_dtype=torch.float16, + head_num=1, + head_dim=8, + v_head_dim=8, + swa_head_num=1, + swa_head_dim=8, + swa_v_head_dim=8, + page_size=1, + start_layer=0, + end_layer=2, + swa_attention_layer_ids=[1], + full_attention_layer_ids=[0], + total_bytes=4096, + enable_memory_saver=False, + need_sort=False, + lazy_compaction=False, + ) + allocator = bundle.token_to_kv_pool_allocator + padding = allocator.alloc(4) + values = allocator.alloc(4) + core = _swa_tree_core(window=4, token_to_kv_pool_allocator=allocator) + core.set_hicache_enabled() + core.has_swa_host_pool = True + sources = [] + if backup_nodes != "leaf": + parent = _insert(core, [1, 2], values[:2].tolist()).last_device_node + sources.append((parent, values[:2])) + node = _insert(core, [1, 2, 3, 4], values.tolist()).last_device_node + node_values = values if backup_nodes == "leaf" else values[2:] + sources.append((node, node_values)) + for source_id, source_values in sources: + core.set_component_device_value( + source_id, + ComponentType.SWA, + allocator.translate_swa_indices_for_transfer(source_values), + ) + backed_up = {} + if backup_nodes == "parent": + backed_up[ComponentType.SWA] = [ + PoolTransfer( + name=PoolName.SWA, + host_indices=torch.tensor([200, 201]), + device_indices=allocator.translate_swa_indices_for_transfer( + node_values + ), + nodes_to_load=[node], + ) + ] + sources.pop() + core.commit_backup(node, torch.arange(100, 100 + node_values.numel()), backed_up) + allocator.free(padding) + + full_indices, transfers = core.build_backup_spec(node) + + assert full_indices.numel() == 0 + (swa_transfer,) = transfers[ComponentType.SWA] + assert swa_transfer.nodes_to_load == [source_id for source_id, _ in sources] + assert torch.equal( + swa_transfer.device_indices, + allocator.translate_swa_indices_for_transfer( + torch.cat([source_values for _, source_values in sources]) + ), + ) + + def test_swa_prefetch_commit_end_to_end(): from sglang.srt.mem_cache.unified_cache.components import CacheTransferPhase @@ -4847,7 +4922,7 @@ def test_swa_rebuild_applies_through_the_python_allocator(): # The cache executed SWARebuild through the allocator: the node holds the # full slice's SWA translation. stored = cache.tree_core.get_component_device_value(node, ComponentType.SWA) - expected = allocator.translate_loc_from_full_to_swa(full) + expected = allocator.translate_swa_indices_for_transfer(full) assert stored is not None assert stored.tolist() == expected.tolist() assert (allocator.full_to_swa_index_mapping[full.to(torch.int64)] > 0).all() @@ -4883,7 +4958,9 @@ def test_recover_with_locked_full_applies_through_the_python_allocator(): # The locked full keeps its slots, remapped onto the incoming full's SWA # translation; the incoming full is freed back to the allocator. stored = cache.tree_core.get_component_device_value(node, ComponentType.SWA) - assert stored.tolist() == allocator.translate_loc_from_full_to_swa(kept).tolist() + assert ( + stored.tolist() == allocator.translate_swa_indices_for_transfer(kept).tolist() + ) assert (allocator.full_to_swa_index_mapping[incoming.to(torch.int64)] == 0).all() assert ( allocator.full_attn_allocator.available_size() == before_free + incoming.numel() diff --git a/test/registered/unit/mem_cache/test_storage_prefetch_lifecycle.py b/test/registered/unit/mem_cache/test_storage_prefetch_lifecycle.py index a3a3b7728ae9..a1dfab160871 100644 --- a/test/registered/unit/mem_cache/test_storage_prefetch_lifecycle.py +++ b/test/registered/unit/mem_cache/test_storage_prefetch_lifecycle.py @@ -83,8 +83,12 @@ def _staged_fixture(full_match=2): cc.storage_backend.batch_exists.return_value = 0 cc.mem_pool_host = SimpleNamespace( free=Mock(), + anchor_entry=SimpleNamespace(host_pool=SimpleNamespace()), entry_map={ - PoolName.SWA: SimpleNamespace(host_pool=SimpleNamespace(free=Mock())) + PoolName.SWA: SimpleNamespace( + host_pool=SimpleNamespace(free=Mock()), + device_indices_from_anchor_fn=None, + ) }, ) cache.cache_controller = cc @@ -281,6 +285,7 @@ def test_trim_and_stage_preserve_raw_token_boundaries(self): swa = PoolTransfer(name=PoolName.SWA, host_indices=torch.arange(4)) operation.pool_transfers = [swa] operation.host_indices = torch.arange(hit_tokens) + operation.buffer_host_occupied_units = hit_tokens cache.ongoing_prefetch[req.cache_request_handle] = info._replace( host_indices=operation.host_indices, comp_xfers={"swa": [swa]} ) diff --git a/test/registered/unit/mem_cache/test_swa_locked_full_recover_unified.py b/test/registered/unit/mem_cache/test_swa_locked_full_recover_unified.py index 63ac884f9679..cbdc9eef6007 100644 --- a/test/registered/unit/mem_cache/test_swa_locked_full_recover_unified.py +++ b/test/registered/unit/mem_cache/test_swa_locked_full_recover_unified.py @@ -127,6 +127,9 @@ def set_full_to_swa_mapping(self, full, swa): self.mapping_calls.append((full, swa)) self.full_to_swa_index_mapping[full.to(torch.int64)] = swa.to(torch.int64) + def translate_swa_indices_for_transfer(self, full): + return self.translate_loc_from_full_to_swa(full) + def clear_full_to_swa_mapping(self, full): self.clear_calls.append(full) self.full_to_swa_index_mapping[full.to(torch.int64)] = 0 diff --git a/test/registered/unit/mem_cache/test_umbp_store.py b/test/registered/unit/mem_cache/test_umbp_store.py index 271b78a59fbf..2faac54b2055 100755 --- a/test/registered/unit/mem_cache/test_umbp_store.py +++ b/test/registered/unit/mem_cache/test_umbp_store.py @@ -36,6 +36,8 @@ class MockStorageConfig: class MockHostKVCache: """Mock HostKVCache that simulates page_first layout with real buffers.""" + stores_page_envelope = False + def __init__(self, num_pages=4, page_size=1, element_size=1024): self.layout = "page_first" self.page_size = page_size @@ -91,6 +93,7 @@ class MockLogicalHostPool: layout = "page_first" page_size = 1 kv_buffer = None + stores_page_envelope = False class MockHybridSidePool: diff --git a/test/registered/unit/mem_cache/test_unified_free_no_host_sync.py b/test/registered/unit/mem_cache/test_unified_free_no_host_sync.py index 05063d710797..736b8502e36d 100644 --- a/test/registered/unit/mem_cache/test_unified_free_no_host_sync.py +++ b/test/registered/unit/mem_cache/test_unified_free_no_host_sync.py @@ -63,8 +63,10 @@ def _paged_allocator(lazy: bool): _TOMBSTONE_METHODS = [ (mea.MultiEndedAllocator, "_free_lazy"), (mea.MultiEndedAllocator, "free"), + (mea.MultiEndedAllocator, "free_physical"), (mea.MultiEndedAllocator, "_commit_move_batch"), (mea.FloatMultiEndedAllocator, "free"), + (mea.FloatMultiEndedAllocator, "free_physical"), (mea.FloatMultiEndedAllocator, "make_room"), (mea.FloatMultiEndedAllocator, "_relocate_to_positions"), ] diff --git a/test/registered/unit/mem_cache/test_unified_radix_cache_unittest.py b/test/registered/unit/mem_cache/test_unified_radix_cache_unittest.py index eee309cf769f..6893cb49cf81 100644 --- a/test/registered/unit/mem_cache/test_unified_radix_cache_unittest.py +++ b/test/registered/unit/mem_cache/test_unified_radix_cache_unittest.py @@ -757,6 +757,50 @@ def _build_internal_chain(self, component_type, enable_session_radix_cache): self.assertIsNotNone(_device_value(cache, node_id, ComponentType.FULL)) return cache, first, second, leaf + def test_swa_write_back_host_full_preserves_full_and_continues(self): + ct = ComponentType.SWA + for session in _session_radix_cache_test_values(): + for pinned in (False, True): + with self.subTest(session=session, pinned=pinned): + # This exercises the Rust backup barrier, independently of + # the shared suite's default backend. + with mock.patch(f"{__name__}._TREE_CORE_TEST_BACKEND", "rust"): + cache, first, second, leaf = self._build_internal_chain( + ct, session + ) + cache.tree_core.is_write_back = True + cache.tree_core.has_swa_host_pool = True + cache.tree_core.enable_swa_write_back_eviction_barrier() + if session: + # Session caches use main's Python fallback, whose + # component backup is gated by the HiCache attachment. + cache.tree_core.enable_hicache = True + cache.host_pool_group = mock.Mock() + cache.host_pool_group.get_pool.return_value = None + if pinned: + receipt = cache.inc_host_lock_ref(first).to_dec_params() + tracker = {ComponentType.FULL: 0, ct: 0} + # Only host allocation fails; tree walking and freeing are real. + with ( + mock.patch.object(cache, "cache_controller", mock.Mock()), + mock.patch.object( + cache, "_execute_and_commit_kv_backup", return_value=0 + ) as backup, + ): + cache._evict_components({ComponentType.FULL: 0, ct: 2}, tracker) + backup.assert_called_once() + self.assertEqual(tracker[ct], 2) + self.assertEqual(tracker[ComponentType.FULL], 0) + for node_id in (first, second, leaf): + self.assertIsNotNone( + _device_value(cache, node_id, ComponentType.FULL) + ) + if pinned: + # A host pin does not pin a device-only SWA value. + self.assertIsNone(_device_value(cache, first, ct)) + cache.dec_host_lock_ref(first, receipt) + cache.sanity_check() + def _evict_for_alloc_after_first_drain(self, cache, component_type): capacity = {"available": 0} auxiliary_drains = {"count": 0} @@ -1092,6 +1136,7 @@ def test_sanity_check_reads_buffer_backup_node_id_from_snapshot(self): host_indices=torch.empty(0, dtype=torch.int64), aux_xfers=[], lock_params=lock_params, + occupied_units=0, ) } cache.buffer_pipeline = pipeline @@ -9168,6 +9213,39 @@ def _component_with_cache(component_type, cache): class TestUnifiedRadixCacheActionRouting(CustomTestCase): """CacheAction routing: each type forwards to the right Controller API.""" + def test_retraction_uses_engine_transfer_index_domains(self): + cache = object.__new__(UnifiedRadixCache) + cache.is_swa_enabled = True + cache._sliding_window_size = 3 + cache.tree_core = mock.Mock(page_size=2) + cache.page_size = 2 + cache.sidecar_pool_specs = () + cache.req_to_token_pool = mock.Mock() + cache.req_to_token_pool.req_to_token = torch.tensor( + [[1, 2, 3, 4, 5, 0]], dtype=torch.int64 + ) + cache.token_to_kv_pool_allocator = mock.Mock() + cache.token_to_kv_pool_allocator.translate_swa_indices_for_transfer.side_effect = ( + lambda indices: indices + 100 + ) + req = mock.Mock(rid="req", seqlen=6) + req.kv.req_pool_idx = 0 + + full_indices, transfers = cache._retraction_device_transfers(req) + + self.assertTrue(torch.equal(full_indices, torch.tensor([1, 2, 3, 4, 5, 6]))) + self.assertEqual(len(transfers), 1) + self.assertEqual(transfers[0].name, PoolName.SWA) + self.assertTrue( + torch.equal( + transfers[0].device_indices, + torch.tensor([103, 104, 105, 106]), + ) + ) + cache.token_to_kv_pool_allocator.translate_swa_indices_for_transfer.assert_called_once() + cache.token_to_kv_pool_allocator.translate_kv_indices_for_transfer.assert_not_called() + cache.token_to_kv_pool_allocator.translate_loc_from_full_to_swa.assert_not_called() + def test_backup_publish_node_ids_collects_component_nodes_once(self): comp_xfers = { ComponentType.SWA: [PoolTransfer(name=PoolName.SWA, nodes_to_load=[3, 4])], @@ -9252,12 +9330,13 @@ def test_apply_component_action_swa_rebuild(self): cache = mock.MagicMock() alloc = cache.token_to_kv_pool_allocator source_value = torch.tensor([3, 4], dtype=torch.int64) - swa_value = alloc.translate_loc_from_full_to_swa.return_value + swa_value = torch.tensor([7, 8], dtype=torch.int64) + alloc.translate_swa_indices_for_transfer.return_value = swa_value _component_with_cache(ComponentType.SWA, cache).apply_component_action( SWARebuild(node_id=5, source_value=source_value), ) # translate the source full to SWA and store it on the node (no free) - alloc.translate_loc_from_full_to_swa.assert_called_once_with(source_value) + alloc.translate_swa_indices_for_transfer.assert_called_once_with(source_value) alloc.free.assert_not_called() alloc.free_full.assert_not_called() cache.tree_core.set_component_device_value.assert_called_once_with( @@ -9269,7 +9348,8 @@ def test_apply_component_action_swa_recover_on_full_locked(self): alloc = cache.token_to_kv_pool_allocator kept_full = torch.tensor([1, 2], dtype=torch.int64) incoming_full = torch.tensor([3, 4], dtype=torch.int64) - swa_value = alloc.translate_loc_from_full_to_swa.return_value + swa_value = torch.tensor([7, 8], dtype=torch.int64) + alloc.translate_swa_indices_for_transfer.return_value = swa_value _component_with_cache(ComponentType.SWA, cache).apply_component_action( RecoverSWAWithLockedFull( node_id=5, @@ -9278,7 +9358,7 @@ def test_apply_component_action_swa_recover_on_full_locked(self): ), ) # keep the locked full, remap it onto the incoming full's SWA translation - alloc.translate_loc_from_full_to_swa.assert_called_once_with(incoming_full) + alloc.translate_swa_indices_for_transfer.assert_called_once_with(incoming_full) alloc.set_full_to_swa_mapping.assert_called_once_with(kept_full, swa_value) # the incoming full's stale mapping is cleared, then its slot freed (full-only) alloc.clear_full_to_swa_mapping.assert_called_once_with(incoming_full) @@ -10703,15 +10783,21 @@ def test_storage_cleanup_releases_buffer_prefetch_anchor(self): anchor_node_id=99, prefetch_key=RadixKey(array("q", range(16))), host_indices=host_indices, - operation=mock.Mock(), + operation=mock.Mock( + buffer_host_occupied_units=4, pool_transfers_done=True + ), anchor_lock_params=None, comp_xfers={}, ) } cache.ongoing_backup = {} cache.host_memory_mode = "buffer_only" - cache._prefetch_occupied_span.side_effect = lambda key, indices: ( - UnifiedRadixCache._prefetch_occupied_span(cache, key, indices) + cache._prefetch_occupied_span.side_effect = ( + lambda key, indices, *, operation=None: ( + UnifiedRadixCache._prefetch_occupied_span( + cache, key, indices, operation=operation + ) + ) ) controller = cache.cache_controller controller.terminate_prefetch.return_value = (4, None)