From aa718f734356f3c9d1ad736337a0130a24244f9e Mon Sep 17 00:00:00 2001 From: cctry Date: Tue, 25 Aug 2026 16:31:57 -0700 Subject: [PATCH] Refactor HiCache host pool management (#36232) --- .../sglang/srt/managers/cache_controller.py | 175 +---------------- .../hybrid_cache/hybrid_cache_controller.py | 132 ++++++------- .../hybrid_cache/hybrid_pool_assembler.py | 22 +-- .../sglang/srt/mem_cache/kv_cache_builder.py | 57 +----- .../sglang/srt/mem_cache/memory_pool_host.py | 134 +------------ .../srt/mem_cache/pool_host/__init__.py | 3 + .../sglang/srt/mem_cache/pool_host/group.py | 181 ++++++++++++++++++ .../components/full_component.py | 2 +- .../components/mamba_component.py | 11 +- .../unified_cache/components/swa_component.py | 11 +- .../srt/mem_cache/unified_radix_cache.py | 38 ++-- .../test_decode_retraction_backup.py | 1 - ...test_hicache_staged_write_back_dispatch.py | 30 ++- .../unit/mem_cache/test_mem_pool_host.py | 52 ++++- .../test_unified_radix_cache_unittest.py | 2 +- ...test_supplied_instance_exposure_ratchet.py | 1 - 16 files changed, 350 insertions(+), 502 deletions(-) create mode 100644 python/sglang/srt/mem_cache/pool_host/group.py diff --git a/python/sglang/srt/managers/cache_controller.py b/python/sglang/srt/managers/cache_controller.py index d3f2a6128..41eb16383 100644 --- a/python/sglang/srt/managers/cache_controller.py +++ b/python/sglang/srt/managers/cache_controller.py @@ -318,6 +318,7 @@ class HiCacheController: mem_pool_device = mem_pool_device.full_kv_pool self.mem_pool_device = mem_pool_device self.mem_pool_host = mem_pool_host + self.storage_host_pool = mem_pool_host self.write_policy = write_policy self.page_size = page_size self.io_backend = io_backend @@ -329,15 +330,6 @@ class HiCacheController: # limiter subtracts write staging from actual pool usage. self.host_write_staged_tokens_fn: Optional[Callable[[], int]] = None - # Draft KV pool support (best-effort piggyback on target L2/L3 ops). - self.has_draft = False - self.mem_pool_device_draft = None - self.mem_pool_host_draft = None - self.draft_page_get_func = None - self.draft_page_set_func = None - self.has_mtp_draft = False - self.mtp_draft_device_pools = () - # Default storage page IO functions (may be overridden by attach). self.page_get_func = self._generic_page_get self.page_set_func = self._generic_page_set @@ -570,9 +562,9 @@ class HiCacheController: try: self.storage_backend = StorageBackendFactory.create_backend( - storage_backend, self.storage_config, self.mem_pool_host + storage_backend, self.storage_config, self.storage_host_pool ) - self.storage_backend.register_mem_pool_host(self.mem_pool_host) + self.storage_backend.register_mem_pool_host(self.storage_host_pool) self.enable_storage = True # todo: threshold policy for prefetching @@ -609,8 +601,6 @@ class HiCacheController: self.page_get_func = self._page_get_zero_copy self.page_set_func = self._page_set_zero_copy - self._maybe_register_draft_with_storage() - # Ensure stop_event is clear before starting threads. self.storage_stop_event.clear() self._start_storage_threads() @@ -638,8 +628,6 @@ class HiCacheController: self.enable_storage = False self.page_get_func = self._generic_page_get self.page_set_func = self._generic_page_set - self.draft_page_get_func = None - self.draft_page_set_func = None raise def detach_storage_backend(self): @@ -685,8 +673,6 @@ class HiCacheController: self.enable_storage = False self.page_get_func = self._generic_page_get self.page_set_func = self._generic_page_set - self.draft_page_get_func = None - self.draft_page_set_func = None # Now it's safe to clear the stop event for future re-attach. self.storage_stop_event.clear() @@ -838,12 +824,7 @@ class HiCacheController: ) def _transfer_num_bytes(self, op: CacheOperation) -> int: - """Total bytes moved by a merged transfer op (draft piggyback included).""" - num_tokens = len(op.device_indices) - num_bytes = num_tokens * self.mem_pool_host.size_per_token - if self.has_draft: - num_bytes += num_tokens * self.mem_pool_host_draft.size_per_token - return num_bytes + return len(op.device_indices) * self.mem_pool_host.size_per_token def _num_tokens_by_pool(self, op: CacheOperation) -> dict[str, int]: return {PoolName.KV.value: len(op.device_indices)} @@ -920,15 +901,6 @@ class HiCacheController: device_indices=device_indices, ) ] - if self.has_draft and host_indices.numel() > 0: - transfers.append( - L2Transfer( - host_pool=self.mem_pool_host_draft, - device_pool=self.mem_pool_device_draft, - host_indices=host_indices, - device_indices=device_indices, - ) - ) return transfers def _l2_load_transfers( @@ -981,63 +953,6 @@ class HiCacheController: self.mem_pool_host.free(host_indices) return len(host_indices) - def set_draft_kv_pool(self, draft_device_pool, draft_host_pool) -> None: - """Register draft KV pools so L2/L3 ops piggyback draft transfers.""" - self.has_draft = True - self.mem_pool_device_draft = draft_device_pool - self.mem_pool_host_draft = draft_host_pool - logger.info( - "HiCache draft KV registered: %s (host %d slots)", - type(draft_device_pool).__name__, - draft_host_pool.size, - ) - - # If storage is already attached, wire up the draft I/O path now. - # Otherwise this will be deferred until attach_storage_backend(). - self._maybe_register_draft_with_storage() - - def set_mtp_draft_pools(self, device_pools) -> None: - """Register MTP device pools used for L2 load-back.""" - self.mtp_draft_device_pools = tuple(device_pools) - self.has_mtp_draft = bool(self.mtp_draft_device_pools) - - def _maybe_register_draft_with_storage(self) -> None: - """Pick the draft L3 IO implementation.""" - self.draft_page_get_func = None - self.draft_page_set_func = None - if not self.has_draft or not self.enable_storage: - return - - backend = self.storage_backend_type - - # Multi-pool zero-copy backends. - if backend == "mooncake": - if self.storage_config.should_split_heads: - logger.warning( - "HiCache draft L3 disabled: should_split_heads not yet " - "supported on the mooncake v2 path." - ) - return - self.storage_backend.register_mem_host_pool_v2( - self.mem_pool_host_draft, PoolName.DRAFT - ) - self.draft_page_get_func = self._draft_page_get_v2 - self.draft_page_set_func = self._draft_page_set_v2 - return - - # TODO: support "hf3fs", "eic", "nixl", "simm" - if backend in {"hf3fs", "eic", "nixl", "simm"}: - logger.warning( - "HiCache draft L3 disabled: backend %s does not yet support " - "draft pool registration.", - backend, - ) - return - - # Generic backends. - self.draft_page_get_func = self._draft_page_get_generic - self.draft_page_set_func = self._draft_page_set_generic - def prefetch( self, request_id: str, @@ -1094,7 +1009,7 @@ class HiCacheController: self, operation, hash_values, host_indices, extra_info=None ) -> int: dummy_page_dst = [ - self.mem_pool_host.get_dummy_flat_data_page() for _ in hash_values + self.storage_host_pool.get_dummy_flat_data_page() for _ in hash_values ] page_data = self.storage_backend.batch_get(hash_values, dummy_page_dst) if page_data is None: @@ -1108,7 +1023,7 @@ class HiCacheController: break if operation.is_terminated(): break - self.mem_pool_host.set_from_flat_data_page( + self.storage_host_pool.set_from_flat_data_page( host_indices[i * self.page_size], page_data[i], ) @@ -1137,12 +1052,6 @@ class HiCacheController: i * self.page_size : (i + len(batch_hashes)) * self.page_size ] - # Best-effort draft L3 read before publishing target completion. - # Otherwise wait_complete can race and load back target KV before - # draft KV reaches host memory. - if self.has_draft: - self._draft_page_get(batch_hashes, batch_host_indices) - # Get one batch token, and update the completed_tokens if succeed extra_info = HiCacheStorageExtraInfo(prefix_keys=prefix_keys) @@ -1325,7 +1234,7 @@ class HiCacheController: # todo: deprecate def _generic_page_set(self, hash_values, host_indices, extra_info=None) -> bool: data = [ - self.mem_pool_host.get_data_page(host_indices[i * self.page_size]) + self.storage_host_pool.get_data_page(host_indices[i * self.page_size]) for i in range(len(hash_values)) ] return self.storage_backend.batch_set(hash_values, data) @@ -1335,72 +1244,6 @@ class HiCacheController: self.storage_backend.batch_set_v1(hash_values, host_indices, extra_info) ) - def _draft_page_set(self, hash_values, host_indices) -> None: - """Best-effort write draft KV pages to L3 alongside the target backup.""" - if self.draft_page_set_func is None: - return - try: - self.draft_page_set_func(hash_values, host_indices) - except Exception: - logger.debug( - "Draft L3 write failed (best-effort), skipping.", exc_info=True - ) - - def _draft_page_get(self, hash_values, host_indices) -> None: - """Best-effort read draft KV pages from L3 (mirrors `_draft_page_set`).""" - if self.draft_page_get_func is None: - return - try: - self.draft_page_get_func(hash_values, host_indices) - except Exception: - logger.debug("Draft L3 read failed (best-effort), skipping.", exc_info=True) - - def _draft_page_set_v2(self, hash_values, host_indices) -> None: - self.storage_backend.batch_set_v2( - [ - PoolTransfer( - name=PoolName.DRAFT, - host_indices=host_indices, - keys=list(hash_values), - ) - ] - ) - - def _draft_page_get_v2(self, hash_values, host_indices) -> None: - self.storage_backend.batch_get_v2( - [ - PoolTransfer( - name=PoolName.DRAFT, - host_indices=host_indices, - keys=list(hash_values), - ) - ] - ) - - def _draft_page_set_generic(self, hash_values, host_indices) -> None: - # `{hash}.draft` mirrors HiCacheStorage._get_component_key's - # `{key}.{pool_name}` convention so target/draft pages never collide. - draft_keys = [f"{h}.{PoolName.DRAFT}" for h in hash_values] - draft_data = [ - self.mem_pool_host_draft.get_data_page(host_indices[i * self.page_size]) - for i in range(len(draft_keys)) - ] - self.storage_backend.batch_set(draft_keys, draft_data) - - def _draft_page_get_generic(self, hash_values, host_indices) -> None: - draft_keys = [f"{h}.{PoolName.DRAFT}" for h in hash_values] - draft_dummy = [ - self.mem_pool_host_draft.get_dummy_flat_data_page() for _ in draft_keys - ] - draft_pages = self.storage_backend.batch_get(draft_keys, draft_dummy) - if draft_pages is None: - return - for i, p in enumerate(draft_pages): - if p is not None: - self.mem_pool_host_draft.set_from_flat_data_page( - host_indices[i * self.page_size], p - ) - # Backup batch by batch def _page_backup(self, operation): # Backup batch by batch @@ -1420,10 +1263,6 @@ class HiCacheController: ) break - # Best-effort draft L3 write alongside target. - if self.has_draft: - self._draft_page_set(batch_hashes, batch_host_indices) - if prefix_keys and len(prefix_keys) > 0: prefix_keys += batch_hashes operation.completed_tokens += self.page_size * len(batch_hashes) 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 94b22e349..eb3a2b6f5 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 @@ -33,7 +33,7 @@ from sglang.srt.mem_cache.hicache_storage import ( count_pool_hits, ) from sglang.srt.mem_cache.l2_transfer import L2Transfer -from sglang.srt.mem_cache.memory_pool_host import HostPoolGroup, PoolEntry +from sglang.srt.mem_cache.pool_host import HostPoolGroup, PoolEntry from sglang.srt.mem_cache.pool_host.mha import MHATokenToKVPoolHost if TYPE_CHECKING: @@ -136,6 +136,7 @@ class HybridCacheController(BaseHiCacheController): self.layer_num = transfer_layer_num self.layer_done_counter = LayerDoneCounter(self.layer_num) + self.storage_host_pool = mem_pool_host.anchor_entry.host_pool if startup_storage_backend is not None: self.attach_storage_backend( storage_backend=startup_storage_backend, @@ -315,11 +316,10 @@ class HybridCacheController(BaseHiCacheController): host_indices = self.mem_pool_host.alloc(len(device_indices)) if host_indices is None: return None - pool_transfers = self._resolve_pool_transfers_allocation( + pool_transfers = self.mem_pool_host.resolve_host_transfers( extra_pools, - alloc_host=True, - kv_device_indices=device_indices, - kv_host_indices=host_indices, + primary_device_indices=device_indices, + primary_host_indices=host_indices, ) if pool_transfers is None and extra_pools: self.mem_pool_host.free(host_indices) @@ -416,15 +416,6 @@ class HybridCacheController(BaseHiCacheController): layer_mapper=entry.layer_mapper, ) ) - if self.has_draft and host_indices.numel() > 0: - transfers.append( - L2Transfer( - host_pool=self.mem_pool_host_draft, - device_pool=self.mem_pool_device_draft, - host_indices=host_indices, - device_indices=device_indices, - ) - ) return transfers def _l2_load_transfers( @@ -434,36 +425,40 @@ class HybridCacheController(BaseHiCacheController): pool_transfers: Optional[list[PoolTransfer]] = None, ) -> list[L2Transfer]: transfers = self._l2_transfers(host_indices, device_indices, pool_transfers) - if getattr(self, "has_mtp_draft", False): - target_transfers = list(transfers) - for depth, draft_device_pool in enumerate(self.mtp_draft_device_pools): - for transfer in target_transfers: - if transfer.layer_mapper is None: - continue - draft_host_layer = transfer.layer_mapper(self.layer_num + depth) - if draft_host_layer is None: - continue + transfers_by_entry = { + (id(t.host_pool), id(t.device_pool)): t for t in transfers + } + for entry in self.mem_pool_host.entry_map.values(): + target_transfer = transfers_by_entry.get( + (id(entry.host_pool), id(entry.device_pool)) + ) + if target_transfer is None or target_transfer.layer_mapper is None: + continue + for depth, draft_device_pool in enumerate(entry.packed_draft_device_pools): + draft_host_layer = target_transfer.layer_mapper(self.layer_num + depth) + if draft_host_layer is None: + continue - def draft_layer_mapper( - layer_id: int, - *, - expected_layer_id: int = depth, - host_layer_id: int = draft_host_layer, - ) -> Optional[int]: - if layer_id == expected_layer_id: - return host_layer_id - return None + def draft_layer_mapper( + layer_id: int, + *, + expected_layer_id: int = depth, + host_layer_id: int = draft_host_layer, + ) -> Optional[int]: + if layer_id == expected_layer_id: + return host_layer_id + return None - transfers.append( - L2Transfer( - host_pool=transfer.host_pool, - device_pool=draft_device_pool, - host_indices=transfer.host_indices, - device_indices=transfer.device_indices, - layer_mapper=draft_layer_mapper, - is_draft=True, - ) + transfers.append( + L2Transfer( + host_pool=target_transfer.host_pool, + device_pool=draft_device_pool, + host_indices=target_transfer.host_indices, + device_indices=target_transfer.device_indices, + layer_mapper=draft_layer_mapper, + is_draft=True, ) + ) return transfers def _num_tokens_by_pool(self, op: CacheOperation) -> dict[str, int]: @@ -479,13 +474,13 @@ class HybridCacheController(BaseHiCacheController): return counts def _transfer_num_bytes(self, op: CacheOperation) -> int: - """Total bytes moved by a merged transfer op across all pools, - including draft piggyback and sidecar transfers riding another - pool's indices (both excluded from the per-pool token counts).""" + """Total bytes moved by a merged transfer op across all pools. + + Sidecar transfers riding another pool's indices are included here but + excluded from the per-pool token counts. + """ kv_tokens = len(op.device_indices) num_bytes = kv_tokens * self.mem_pool_host.anchor_entry.host_pool.size_per_token - if self.has_draft: - num_bytes += kv_tokens * self.mem_pool_host_draft.size_per_token # Slot counts of the pools sidecars can ride on. source_len = {self.mem_pool_host.anchor_entry.name: kv_tokens} for t in op.pool_transfers or []: @@ -523,9 +518,8 @@ class HybridCacheController(BaseHiCacheController): if device_indices is None: return None - pool_transfers = self._resolve_pool_transfers_allocation( + pool_transfers = self._resolve_device_transfers( extra_pools, - alloc_host=False, kv_device_indices=device_indices, kv_host_indices=host_indices, ) @@ -833,27 +827,22 @@ class HybridCacheController(BaseHiCacheController): ) transfer.host_indices = transfer.host_indices[:needed] - def _resolve_pool_transfers_allocation( + def _resolve_device_transfers( self, extra_pools: Optional[list[PoolTransfer]], - alloc_host: bool, kv_device_indices: Optional[torch.Tensor] = None, kv_host_indices: Optional[torch.Tensor] = None, ) -> Optional[list[PoolTransfer]]: - """Auto-alloc host or device indices for PoolTransfers where they are None.""" + """Allocate unresolved side-pool device indices atomically.""" if not extra_pools: return None - # (pool, free_fn, indices) for atomic rollback on failure. newly_allocated: list[tuple[PoolTransfer, Callable, torch.Tensor]] = [] derived_transfers: list[PoolTransfer] = [] def rollback_allocated() -> None: for prev_pool, prev_free_fn, prev_indices in newly_allocated: prev_free_fn(prev_indices) - if alloc_host: - prev_pool.host_indices = None - else: - prev_pool.device_indices = None + prev_pool.device_indices = None for pool in extra_pools: if pool.indices_from_pool is not None: @@ -862,23 +851,15 @@ class HybridCacheController(BaseHiCacheController): entry = self.mem_pool_host.entry_map.get(pool.name) if entry is None: continue - if alloc_host: - if pool.host_indices is not None or pool.device_indices is None: - continue - alloc_fn = entry.host_pool.alloc - free_fn = entry.host_pool.free - evict_fn = entry.host_evict_fn - size = len(pool.device_indices) - else: - if pool.device_indices is not None or pool.host_indices is None: - continue - # device_alloc_fn / device_free_fn override entry.device_pool's - # methods for pools whose device_pool is a raw KV pool (layout) - # rather than an allocator (e.g. SWA). - alloc_fn = entry.device_alloc_fn or entry.device_pool.alloc - free_fn = entry.device_free_fn or entry.device_pool.free - evict_fn = entry.device_evict_fn - size = len(pool.host_indices) + if pool.device_indices is not None or pool.host_indices is None: + continue + # device_alloc_fn / device_free_fn override entry.device_pool's + # methods for pools whose device_pool is a raw KV pool (layout) + # rather than an allocator (e.g. SWA). + alloc_fn = entry.device_alloc_fn or entry.device_pool.alloc + free_fn = entry.device_free_fn or entry.device_pool.free + evict_fn = entry.device_evict_fn + size = len(pool.host_indices) indices = alloc_fn(size) if indices is None and evict_fn: evict_fn(size) @@ -887,10 +868,7 @@ class HybridCacheController(BaseHiCacheController): # Atomic rollback: free everything we successfully allocated. rollback_allocated() return None - if alloc_host: - pool.host_indices = indices - else: - pool.device_indices = indices + pool.device_indices = indices newly_allocated.append((pool, free_fn, indices)) # Assign indices to deferred pools from their source. 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 22b1c3307..ad52d2bef 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 @@ -15,10 +15,9 @@ from sglang.srt.mem_cache.hybrid_cache.hybrid_cache_controller import ( from sglang.srt.mem_cache.memory_pool_host import ( DeepSeekV4PagedHostPool, DeepSeekV4StateHostPool, - HostPoolGroup, LogicalHostPool, - PoolEntry, ) +from sglang.srt.mem_cache.pool_host import HostPoolGroup, PoolEntry from sglang.srt.mem_cache.pool_host.common import get_allocator_type from sglang.srt.mem_cache.pool_host.dsa import DSAIndexerPoolHost from sglang.srt.mem_cache.pool_host.mamba import MambaPoolHost @@ -141,6 +140,7 @@ def build_pool_entry( device_evict_fn: Optional[Callable[[int], Any]] = None, device_alloc_fn: Optional[Callable[[int], Any]] = None, device_free_fn: Optional[Callable[[Any], Any]] = None, + packed_draft_device_pools: tuple[Any, ...] = (), ) -> PoolEntry: return PoolEntry( name=name, @@ -152,6 +152,7 @@ def build_pool_entry( device_evict_fn=device_evict_fn, device_alloc_fn=device_alloc_fn, device_free_fn=device_free_fn, + packed_draft_device_pools=packed_draft_device_pools, ) @@ -193,6 +194,7 @@ def build_kv_only_group( layer_mapping=full_layer_mapping, transfer_layer_num=transfer_layer_num + len(mtp_draft_device_pools), is_anchor=True, + packed_draft_device_pools=mtp_draft_device_pools, ) ] ) @@ -264,6 +266,7 @@ def build_hybrid_swa_group( device_free_fn=( swa_attn_allocator.free if swa_attn_allocator is not None else None ), + packed_draft_device_pools=mtp_swa_device_pools, ), ] ) @@ -313,9 +316,6 @@ def build_kv_only_stack( enable_storage_metrics=enable_storage_metrics, host_memory_mode=server_args.hicache_host_memory_mode, ) - if params.mtp_draft_device_pools: - cache_controller.set_mtp_draft_pools(params.mtp_draft_device_pools) - return host_pool_group, cache_controller @@ -384,8 +384,6 @@ def build_hybrid_swa_stack( enable_storage_metrics=enable_storage_metrics, host_memory_mode=server_args.hicache_host_memory_mode, ) - if mtp_swa_device_pools: - cache_controller.set_mtp_draft_pools(mtp_swa_device_pools) return host_pool_group, cache_controller @@ -538,6 +536,7 @@ def build_deepseek_v4_hicache_stack( device_evict_fn=device_swa_evict_fn, device_alloc_fn=swa_attn_allocator.alloc, device_free_fn=swa_attn_allocator.free, + packed_draft_device_pools=tuple(mtp_swa_device_buffers), ) ) @@ -672,8 +671,6 @@ def build_deepseek_v4_hicache_stack( enable_storage_metrics=enable_storage_metrics, host_memory_mode=server_args.hicache_host_memory_mode, ) - if mtp_swa_device_buffers: - cache_controller.set_mtp_draft_pools(mtp_swa_device_buffers) return host_pool_group, cache_controller @@ -735,6 +732,7 @@ def build_hybrid_mamba_stack( layer_mapping=full_layer_mapping, transfer_layer_num=transfer_layer_num + len(mtp_draft_device_pools), is_anchor=True, + packed_draft_device_pools=mtp_draft_device_pools, ), build_pool_entry( name=PoolName.MAMBA, @@ -768,8 +766,6 @@ def build_hybrid_mamba_stack( enable_storage_metrics=enable_storage_metrics, host_memory_mode=server_args.hicache_host_memory_mode, ) - if mtp_draft_device_pools: - cache_controller.set_mtp_draft_pools(mtp_draft_device_pools) return host_pool_group, cache_controller @@ -933,6 +929,7 @@ def build_anchor_sidecar_stack( layer_mapping=full_layer_mapping, transfer_layer_num=transfer_layer_num + len(mtp_draft_device_pools), is_anchor=True, + packed_draft_device_pools=mtp_draft_device_pools, ), build_pool_entry( name=sidecar_pool_name, @@ -940,6 +937,7 @@ def build_anchor_sidecar_stack( device_pool=kv_pool, layer_mapping=full_layer_mapping, transfer_layer_num=transfer_layer_num + len(mtp_draft_device_pools), + packed_draft_device_pools=mtp_draft_device_pools, ), ] host_pool_group = HostPoolGroup(entries) @@ -962,8 +960,6 @@ def build_anchor_sidecar_stack( enable_storage_metrics=enable_storage_metrics, host_memory_mode=server_args.hicache_host_memory_mode, ) - if mtp_draft_device_pools: - cache_controller.set_mtp_draft_pools(mtp_draft_device_pools) return host_pool_group, cache_controller diff --git a/python/sglang/srt/mem_cache/kv_cache_builder.py b/python/sglang/srt/mem_cache/kv_cache_builder.py index 7f36dab62..08ee50f9f 100644 --- a/python/sglang/srt/mem_cache/kv_cache_builder.py +++ b/python/sglang/srt/mem_cache/kv_cache_builder.py @@ -66,7 +66,6 @@ def maybe_register_hicache_draft( tree_cache, draft_plan: HiCacheDraftPlan, server_args: ServerArgs, - page_size: int, ) -> None: from sglang.srt.speculative.base_spec_worker import HiCacheDraftMode @@ -76,13 +75,7 @@ def maybe_register_hicache_draft( from sglang.srt.mem_cache.unified_radix_cache import UnifiedRadixCache if not isinstance(tree_cache, UnifiedRadixCache): - _register_legacy_hicache_draft( - tree_cache=tree_cache, - draft_pool=draft_plan.device_pools[0], - server_args=server_args, - page_size=page_size, - ) - return + raise NotImplementedError("HiCache draft pools require UnifiedRadixCache.") from sglang.srt.mem_cache.hybrid_cache.hybrid_pool_assembler import ( build_hicache_draft_sidecars, @@ -93,51 +86,8 @@ def maybe_register_hicache_draft( tree_cache=tree_cache, server_args=server_args, ) - tree_cache.register_hicache_draft_pools(specs, entries) - - -def _register_legacy_hicache_draft( - *, - tree_cache, - draft_pool, - server_args: ServerArgs, - page_size: int, -) -> None: - from sglang.srt.mem_cache.memory_pool import ( - MHATokenToKVPool, - MLATokenToKVPool, - ) - from sglang.srt.mem_cache.pool_host.mha import get_mha_host_pool_cls - from sglang.srt.mem_cache.pool_host.mla import MLATokenToKVPoolHost - - pool = draft_pool - if pool.layer_num == 0: - return - - # Create host pool for draft with the same slot count as the target host pool, - # so that host indices stay 1-to-1 between target and draft KV caches. - primary_host_pool = tree_cache.cache_controller.mem_pool_host - host_pool_kwargs = dict( - host_to_device_ratio=primary_host_pool.logical_size / pool.size, - host_size=0, - page_size=page_size, - layout=get_memory().hicache_mem_layout, - allocator_type=server_args.hicache_storage_backend, - pool_label="draft", - ) - if isinstance(pool, MHATokenToKVPool): - draft_host_pool = get_mha_host_pool_cls(pool)(pool, **host_pool_kwargs) - elif isinstance(pool, MLATokenToKVPool): - draft_host_pool = MLATokenToKVPoolHost(pool, **host_pool_kwargs) - else: - logger.warning( - "Draft pool type %s is not supported by the legacy HiCache path; " - "skipping draft KV registration.", - type(pool).__name__, - ) - return - - tree_cache.cache_controller.set_draft_kv_pool(pool, draft_host_pool) + for spec, entry in zip(specs, entries, strict=True): + tree_cache.register_sidecar_pool(spec, entry) # Host slots a backup-only retraction pool gets, as a fraction of the device @@ -369,7 +319,6 @@ def build_kv_cache( tree_cache=tree_cache, draft_plan=hicache_draft_plan, server_args=server_args, - page_size=page_size, ) if retraction_backup == "host_pool": diff --git a/python/sglang/srt/mem_cache/memory_pool_host.py b/python/sglang/srt/mem_cache/memory_pool_host.py index 9aea9427b..2ffb5b4b3 100644 --- a/python/sglang/srt/mem_cache/memory_pool_host.py +++ b/python/sglang/srt/mem_cache/memory_pool_host.py @@ -2,11 +2,7 @@ from __future__ import annotations import logging import threading -from dataclasses import dataclass -from typing import TYPE_CHECKING, Any, Callable, Optional - -if TYPE_CHECKING: - from sglang.srt.mem_cache.hicache_storage import PoolName +from typing import Optional import torch @@ -955,131 +951,3 @@ class DeepSeekV4StateHostPool(HostKVCache): self.kv_buffer.data_ptr() % page_size_bytes == 0 and page_bytes % page_size_bytes == 0 ) - - -@dataclass -class PoolEntry: - name: PoolName - host_pool: Any - device_pool: Any - layer_mapper: Callable[[int], Optional[int]] - is_primary_index_anchor: bool = False - # Optional eviction callbacks for auto-alloc in HybridCacheController. - # host_evict_fn(n): evict n slots from the host pool (used by write()). - # device_evict_fn(n): evict n slots from the device pool (used by load()). - host_evict_fn: Optional[Callable] = None - device_evict_fn: Optional[Callable] = None - # Optional alloc/free overrides for the device side, used by - # _resolve_pool_transfers_allocation. Set when entry.device_pool is the - # raw KV/state pool (layout) rather than an allocator (e.g. SWA/Mamba, - # where alloc lives on a separate allocator object). - # When None, fall back to entry.device_pool.alloc/free. - device_alloc_fn: Optional[Callable] = None - device_free_fn: Optional[Callable] = None - - -class HostPoolGroup: - def __init__(self, entries: list[PoolEntry]): - if not entries: - raise ValueError("HostPoolGroup requires at least one pool entry.") - self.entries = entries - self.entry_map = {entry.name: entry for entry in entries} - self.anchor_entry = next( - (entry for entry in entries if entry.is_primary_index_anchor), - entries[0], - ) - - self.layout = self.anchor_entry.host_pool.layout - self.page_size = self.anchor_entry.host_pool.page_size - self.device = self.anchor_entry.host_pool.device - self.size = self.anchor_entry.host_pool.size - self.logical_size = self.anchor_entry.host_pool.logical_size - child_write_back_jit = [ - getattr(entry.host_pool, "can_use_write_back_jit", False) - for entry in entries - ] - self.can_use_write_back_jit = all(child_write_back_jit) - self.supports_per_pool_backup_indices = any(child_write_back_jit) - - def add_entry(self, entry: PoolEntry) -> None: - if entry.name in self.entry_map: - raise ValueError(f"Host pool {entry.name} is already registered.") - self.entries.append(entry) - self.entry_map[entry.name] = entry - self.can_use_write_back_jit = ( - self.can_use_write_back_jit and entry.host_pool.can_use_write_back_jit - ) - self.supports_per_pool_backup_indices = ( - self.supports_per_pool_backup_indices - or entry.host_pool.can_use_write_back_jit - ) - - @property - def kv_buffer(self): - return self.anchor_entry.host_pool.kv_buffer - - @property - def size_per_token(self): - return self.anchor_entry.host_pool.size_per_token - - @property - def allocator(self): - return self.anchor_entry.host_pool.allocator - - @property - def dtype(self): - return self.anchor_entry.host_pool.dtype - - @property - def start_layer(self): - return self.anchor_entry.host_pool.start_layer - - @property - def end_layer(self): - return self.anchor_entry.host_pool.end_layer - - def get_ksize_per_token(self): - return self.anchor_entry.host_pool.get_ksize_per_token() - - def get_size_per_token(self): - return self.anchor_entry.host_pool.get_size_per_token() - - def get_pool(self, name: PoolName): - return self.entry_map[name].host_pool - - def get_page_buffer_meta(self, indices): - return self.anchor_entry.host_pool.get_page_buffer_meta(indices) - - def get_split_heads_page_buffer_meta(self, indices, split_factor: int): - return self.anchor_entry.host_pool.get_split_heads_page_buffer_meta( - indices, split_factor - ) - - def is_stride_page_aligned(self, page_size_bytes: int = 4096) -> bool: - return self.anchor_entry.host_pool.is_stride_page_aligned(page_size_bytes) - - def clear(self) -> None: - for entry in self.entries: - entry.host_pool.clear() - - def destroy(self) -> None: - for entry in self.entries: - entry.host_pool.destroy() - - def available_size(self): - return self.anchor_entry.host_pool.available_size() - - def alloc(self, need_size: int) -> Optional[torch.Tensor]: - return self.anchor_entry.host_pool.alloc(need_size) - - def free(self, indices: torch.Tensor) -> int: - return self.anchor_entry.host_pool.free(indices) - - def get_data_page(self, index, flat: bool = True): - return self.anchor_entry.host_pool.get_data_page(index, flat) - - def get_dummy_flat_data_page(self): - return self.anchor_entry.host_pool.get_dummy_flat_data_page() - - def set_from_flat_data_page(self, index: int, data_page) -> None: - return self.anchor_entry.host_pool.set_from_flat_data_page(index, data_page) diff --git a/python/sglang/srt/mem_cache/pool_host/__init__.py b/python/sglang/srt/mem_cache/pool_host/__init__.py index 6855269eb..2ae9ccc99 100644 --- a/python/sglang/srt/mem_cache/pool_host/__init__.py +++ b/python/sglang/srt/mem_cache/pool_host/__init__.py @@ -1,7 +1,10 @@ from sglang.srt.mem_cache.pool_host.base import HostKVCache from sglang.srt.mem_cache.pool_host.common import HostTensorAllocator +from sglang.srt.mem_cache.pool_host.group import HostPoolGroup, PoolEntry __all__ = [ "HostKVCache", + "HostPoolGroup", "HostTensorAllocator", + "PoolEntry", ] diff --git a/python/sglang/srt/mem_cache/pool_host/group.py b/python/sglang/srt/mem_cache/pool_host/group.py new file mode 100644 index 000000000..498b60696 --- /dev/null +++ b/python/sglang/srt/mem_cache/pool_host/group.py @@ -0,0 +1,181 @@ +from __future__ import annotations + +from collections.abc import Callable +from dataclasses import dataclass +from typing import Any + +import torch + +from sglang.srt.mem_cache.hicache_storage import PoolName, PoolTransfer + + +@dataclass +class PoolEntry: + name: PoolName + host_pool: Any + device_pool: Any + layer_mapper: Callable[[int], int | None] + is_primary_index_anchor: bool = False + host_evict_fn: Callable[[int], Any] | None = None + device_evict_fn: Callable[[int], Any] | None = None + device_alloc_fn: Callable[[int], Any] | None = None + device_free_fn: Callable[[Any], Any] | None = None + packed_draft_device_pools: tuple[Any, ...] = () + + +class HostPoolGroup: + """Allocation facade for an anchor host pool and its side pools.""" + + def __init__(self, entries: list[PoolEntry]): + if not entries: + raise ValueError("HostPoolGroup requires at least one pool entry.") + if len({entry.name for entry in entries}) != len(entries): + raise ValueError("HostPoolGroup pool names must be unique.") + + anchors = [entry for entry in entries if entry.is_primary_index_anchor] + if len(anchors) > 1: + raise ValueError("HostPoolGroup requires at most one anchor pool.") + + self.entries = list(entries) + self.entry_map = {entry.name: entry for entry in entries} + self.anchor_entry = anchors[0] if anchors else entries[0] + + self.layout = self.anchor_entry.host_pool.layout + self.page_size = self.anchor_entry.host_pool.page_size + self.device = self.anchor_entry.host_pool.device + self.size = self.anchor_entry.host_pool.size + self.logical_size = self.anchor_entry.host_pool.logical_size + self._refresh_transfer_capabilities() + + def _refresh_transfer_capabilities(self) -> None: + child_write_back_jit = [ + entry.host_pool.can_use_write_back_jit for entry in self.entries + ] + self.can_use_write_back_jit = all(child_write_back_jit) + self.supports_per_pool_backup_indices = any(child_write_back_jit) + + def add_entry(self, entry: PoolEntry) -> None: + if entry.name in self.entry_map: + raise ValueError(f"Host pool {entry.name} is already registered.") + if entry.is_primary_index_anchor: + raise ValueError("Cannot replace the anchor of an existing HostPoolGroup.") + self.entries.append(entry) + self.entry_map[entry.name] = entry + self._refresh_transfer_capabilities() + + def get_entry(self, name: PoolName | None = None) -> PoolEntry: + return self.anchor_entry if name is None else self.entry_map[name] + + def get_pool(self, name: PoolName): + return self.get_entry(name).host_pool + + def alloc( + self, + need_size: int, + *, + pool: PoolName | None = None, + reclaim: Callable[[int], Any] | None = None, + ) -> torch.Tensor | None: + """Allocate from one pool, optionally reclaiming once before retrying.""" + host_pool = self.get_entry(pool).host_pool + indices = host_pool.alloc(need_size) + if indices is None and reclaim is not None: + reclaim(need_size) + indices = host_pool.alloc(need_size) + return indices + + def free(self, indices: torch.Tensor, *, pool: PoolName | None = None) -> int: + return self.get_entry(pool).host_pool.free(indices) + + def resolve_host_transfers( + self, + transfers: list[PoolTransfer] | None, + *, + primary_device_indices: torch.Tensor | None = None, + primary_host_indices: torch.Tensor | None = None, + ) -> list[PoolTransfer] | None: + """Allocate unresolved side-pool host indices atomically. + + On failure, every allocation made by this call is released and the + corresponding transfer is restored to its unresolved state. + """ + if not transfers: + return None + + allocated: list[tuple[PoolTransfer, torch.Tensor]] = [] + derived_transfers: list[PoolTransfer] = [] + + def rollback() -> None: + for transfer, indices in allocated: + self.free(indices, pool=transfer.name) + transfer.host_indices = None + + for transfer in transfers: + if transfer.indices_from_pool is not None: + derived_transfers.append(transfer) + continue + if transfer.host_indices is not None or transfer.device_indices is None: + continue + entry = self.entry_map.get(transfer.name) + if entry is None: + continue + indices = self.alloc( + len(transfer.device_indices), + pool=transfer.name, + reclaim=entry.host_evict_fn, + ) + if indices is None: + rollback() + return None + transfer.host_indices = indices + allocated.append((transfer, indices)) + + for transfer in derived_transfers: + if transfer.indices_from_pool == self.anchor_entry.name: + transfer.host_indices = primary_host_indices + transfer.device_indices = primary_device_indices + continue + + source = next( + ( + candidate + for candidate in transfers + if candidate.indices_from_pool is None + and candidate.name == transfer.indices_from_pool + ), + None, + ) + if source is None: + rollback() + return None + transfer.host_indices = source.host_indices + transfer.device_indices = source.device_indices + return transfers + + def release_transfers(self, transfers: list[PoolTransfer] | None) -> int: + """Release independently allocated side-pool indices. + + Derived transfers share another pool's indices and are deliberately + skipped so each allocation is released exactly once. + """ + released = 0 + for transfer in transfers or []: + if transfer.indices_from_pool is not None or transfer.host_indices is None: + continue + released += self.free(transfer.host_indices, pool=transfer.name) + return released + + @property + def size_per_token(self): + return self.anchor_entry.host_pool.size_per_token + + def clear(self) -> None: + for entry in self.entries: + entry.host_pool.clear() + + def destroy(self) -> None: + for entry in self.entries: + entry.host_pool.destroy() + + def available_size(self, pool: PoolName | None = None): + return self.get_entry(pool).host_pool.available_size() diff --git a/python/sglang/srt/mem_cache/unified_cache/components/full_component.py b/python/sglang/srt/mem_cache/unified_cache/components/full_component.py index 3ec838c38..b927e0b0f 100644 --- a/python/sglang/srt/mem_cache/unified_cache/components/full_component.py +++ b/python/sglang/srt/mem_cache/unified_cache/components/full_component.py @@ -428,7 +428,7 @@ class FullComponent(TreeComponent): if self._full_kv_pool_host is None: return for host_value in host_values: - self._full_kv_pool_host.free(host_value) + self.cache.host_pool_group.free(host_value, pool=PoolName.KV) def apply_component_action(self, action: ComponentAction) -> None: if isinstance(action, FreeComponentDeviceSlot): diff --git a/python/sglang/srt/mem_cache/unified_cache/components/mamba_component.py b/python/sglang/srt/mem_cache/unified_cache/components/mamba_component.py index fe5347ac3..3fbd7ad2b 100644 --- a/python/sglang/srt/mem_cache/unified_cache/components/mamba_component.py +++ b/python/sglang/srt/mem_cache/unified_cache/components/mamba_component.py @@ -680,10 +680,11 @@ class MambaComponent(TreeComponent): *, prefetch_tokens: int = 0, ) -> PreparePrefetchResult: - host_indices = self._mamba_pool_host.alloc(1) - if host_indices is None: - self.cache.evict_host(1, ComponentType.MAMBA) - host_indices = self._mamba_pool_host.alloc(1) + host_indices = self.cache.host_pool_group.alloc( + 1, + pool=PoolName.MAMBA, + reclaim=lambda size: self.cache.evict_host(size, ComponentType.MAMBA), + ) if host_indices is None: return PreparePrefetchResult(alloc_failed=True) return PreparePrefetchResult(host_indices=host_indices) @@ -897,7 +898,7 @@ class MambaComponent(TreeComponent): if self._mamba_pool_host is None: return for host_value in host_values: - self._mamba_pool_host.free(host_value) + self.cache.host_pool_group.free(host_value, pool=PoolName.MAMBA) def apply_component_action(self, action: ComponentAction) -> None: if isinstance(action, MambaEvictExcessPathStates): diff --git a/python/sglang/srt/mem_cache/unified_cache/components/swa_component.py b/python/sglang/srt/mem_cache/unified_cache/components/swa_component.py index 74bb8de90..a252541d2 100644 --- a/python/sglang/srt/mem_cache/unified_cache/components/swa_component.py +++ b/python/sglang/srt/mem_cache/unified_cache/components/swa_component.py @@ -794,10 +794,11 @@ class SWAComponent(TreeComponent): # device-guaranteed, require a full window. return PreparePrefetchResult() num_tokens = num_pages * self.cache.page_size - host_indices = self._swa_kv_pool_host.alloc(num_tokens) - if host_indices is None: - self.cache.evict_host(num_tokens, ComponentType.SWA) - host_indices = self._swa_kv_pool_host.alloc(num_tokens) + host_indices = self.cache.host_pool_group.alloc( + num_tokens, + pool=PoolName.SWA, + reclaim=lambda size: self.cache.evict_host(size, ComponentType.SWA), + ) if host_indices is None: return PreparePrefetchResult(alloc_failed=True) return PreparePrefetchResult(host_indices=host_indices) @@ -1121,7 +1122,7 @@ class SWAComponent(TreeComponent): if self._swa_kv_pool_host is None: return for host_value in host_values: - self._swa_kv_pool_host.free(host_value) + self.cache.host_pool_group.free(host_value, pool=PoolName.SWA) def apply_component_action(self, action: ComponentAction) -> None: alloc = self.cache.token_to_kv_pool_allocator diff --git a/python/sglang/srt/mem_cache/unified_radix_cache.py b/python/sglang/srt/mem_cache/unified_radix_cache.py index 091d22d16..97f64df1a 100644 --- a/python/sglang/srt/mem_cache/unified_radix_cache.py +++ b/python/sglang/srt/mem_cache/unified_radix_cache.py @@ -91,7 +91,7 @@ if TYPE_CHECKING: from sglang.srt.mem_cache.hybrid_cache.hybrid_cache_controller import ( PrefetchOperation, ) - from sglang.srt.mem_cache.memory_pool_host import PoolEntry + from sglang.srt.mem_cache.pool_host import PoolEntry from sglang.srt.server_args import ServerArgs from sglang.srt.utils.rank_consensus_checker import rank_consensus @@ -475,17 +475,14 @@ class UnifiedRadixCache(BasePrefixCache): extra_metric_labels=self.extra_metric_labels, ) - def register_sidecar_pool(self, spec: SidecarPoolSpec) -> None: - self.sidecar_pool_specs.append(spec) - - def register_hicache_draft_pools( - self, specs: list[SidecarPoolSpec], entries: list[PoolEntry] + def register_sidecar_pool( + self, spec: SidecarPoolSpec, entry: Optional[PoolEntry] = None ) -> None: - if self.cache_controller is None: - raise RuntimeError("HiCache controller is not attached.") - for spec, entry in zip(specs, entries, strict=True): + if entry is not None: + if self.cache_controller is None: + raise RuntimeError("HiCache controller is not attached.") self.cache_controller.register_host_pool_entry(entry) - self.register_sidecar_pool(spec) + self.sidecar_pool_specs.append(spec) def release_host_resources(self) -> None: if self.host_pool_group is not None: @@ -1137,11 +1134,10 @@ class UnifiedRadixCache(BasePrefixCache): if host_indices is None: return None - resolved = self.cache_controller._resolve_pool_transfers_allocation( + resolved = self.host_pool_group.resolve_host_transfers( extra_transfers or None, - alloc_host=True, - kv_device_indices=device_indices, - kv_host_indices=host_indices, + primary_device_indices=device_indices, + primary_host_indices=host_indices, ) if resolved is None and extra_transfers: self.host_pool_group.free(host_indices) @@ -1195,9 +1191,8 @@ class UnifiedRadixCache(BasePrefixCache): ) for name, saved in saved_by_name.items() ] - resolved = self.cache_controller._resolve_pool_transfers_allocation( + resolved = self.cache_controller._resolve_device_transfers( restored_transfers or None, - alloc_host=False, kv_device_indices=device_indices, kv_host_indices=backup.host_indices, ) @@ -1223,10 +1218,7 @@ class UnifiedRadixCache(BasePrefixCache): def retraction_discard(self, backup: RetractionBackup) -> None: self.host_pool_group.free(backup.host_indices) - for transfer in backup.pool_transfers or []: - if transfer.indices_from_pool is None: - assert transfer.host_indices is not None - self.host_pool_group.get_pool(transfer.name).free(transfer.host_indices) + self.host_pool_group.release_transfers(backup.pool_transfers) # ---- HiCache: Backup / LoadBack ---- @@ -2259,9 +2251,9 @@ class UnifiedRadixCache(BasePrefixCache): host_indices_list.append(host_indices) released_tokens += len(host_indices) if host_indices_list: - entry = cc.mem_pool_host.entry_map.get(pool_name) - if entry is not None: - entry.host_pool.free(torch.cat(host_indices_list, dim=0)) + cc.mem_pool_host.free( + torch.cat(host_indices_list, dim=0), pool=pool_name + ) drained[pool_name] = (len(host_indices_list), released_tokens) return drained diff --git a/test/registered/unit/mem_cache/test_decode_retraction_backup.py b/test/registered/unit/mem_cache/test_decode_retraction_backup.py index 43b68d45b..f20ea8443 100644 --- a/test/registered/unit/mem_cache/test_decode_retraction_backup.py +++ b/test/registered/unit/mem_cache/test_decode_retraction_backup.py @@ -117,7 +117,6 @@ class TestDecodeRetractionBackup(unittest.TestCase): device_pools=(draft_pool,), ), server_args=server_args, - page_size=1, ) self.assertIn(PoolName.DRAFT, cache.host_pool_group.entry_map) cache.validate_retraction_host_capacity() 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 15f77b64e..f5711f2f6 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 @@ -21,10 +21,9 @@ from sglang.srt.mem_cache.l2_transfer import L2Transfer, L2TransferEngine from sglang.srt.mem_cache.memory_pool_host import ( DeepSeekV4PagedHostPool, DeepSeekV4StateHostPool, - HostPoolGroup, LogicalHostPool, - PoolEntry, ) +from sglang.srt.mem_cache.pool_host import HostPoolGroup, PoolEntry from sglang.srt.mem_cache.pool_host.dsa import DSAIndexerPoolHost from sglang.srt.mem_cache.pool_host.mamba import MambaPoolHost from sglang.srt.mem_cache.pool_host.mha import MHATokenToKVPoolHost @@ -222,8 +221,6 @@ class TestHiCacheStagedWriteBackDispatch(CustomTestCase): op.pool_transfers, ) controller.mem_pool_host = _host_group_stub([], can_use_write_back_jit=False) - controller.has_draft = False - controller.has_mtp_draft = False controller._l2_transfers.side_effect = lambda *args: ( HybridCacheController._l2_transfers(controller, *args) ) @@ -287,20 +284,19 @@ class TestHiCacheStagedWriteBackDispatch(CustomTestCase): def test_packed_draft_load_is_flattened_into_l2_transfers(self): host_pool = mock.Mock() controller = HybridCacheController.__new__(HybridCacheController) + entry = PoolEntry( + name=PoolName.KV, + host_pool=host_pool, + device_pool=mock.sentinel.target_device_pool, + layer_mapper={0: 0, 1: 1, 2: 2}.get, + is_primary_index_anchor=True, + packed_draft_device_pools=(mock.sentinel.draft_device_pool,), + ) controller.mem_pool_host = SimpleNamespace( - anchor_entry=PoolEntry( - name=PoolName.KV, - host_pool=host_pool, - device_pool=mock.sentinel.target_device_pool, - layer_mapper={0: 0, 1: 1, 2: 2}.get, - is_primary_index_anchor=True, - ), - entry_map={}, + anchor_entry=entry, + entry_map={entry.name: entry}, ) controller.layer_num = 2 - controller.has_mtp_draft = True - controller.mtp_draft_device_pools = (mock.sentinel.draft_device_pool,) - controller.has_draft = False self.assertEqual( len(controller._l2_transfers(_indices(0, 2), _indices(2, 4))), 1 @@ -937,7 +933,6 @@ class TestHiCacheStagedWriteBackDispatch(CustomTestCase): captured, can_use_write_back_jit=True ) controller.mem_pool_device = None - controller.has_draft = False controller.ack_write_queue = [] controller.move_hybrid_indices = mock.Mock( side_effect=AssertionError( @@ -972,7 +967,6 @@ class TestHiCacheStagedWriteBackDispatch(CustomTestCase): captured, can_use_write_back_jit=False ) controller.mem_pool_device = None - controller.has_draft = False controller.ack_write_queue = [] controller.move_hybrid_indices = mock.Mock( return_value=(op.host_indices, op.device_indices, op.pool_transfers) @@ -1007,7 +1001,6 @@ class TestHiCacheStagedWriteBackDispatch(CustomTestCase): controller.io_backend = "kernel" controller.mem_pool_host = FakeHostPool() controller.mem_pool_device = None - controller.has_draft = False controller.device = "cuda" controller.ack_write_queue = [] controller.move_indices = mock.Mock( @@ -1044,7 +1037,6 @@ class TestHiCacheStagedWriteBackDispatch(CustomTestCase): controller.io_backend = "kernel" controller.mem_pool_host = FakeHostPool() controller.mem_pool_device = None - controller.has_draft = False controller.device = "cuda" controller.ack_write_queue = [] controller.move_indices = mock.Mock( diff --git a/test/registered/unit/mem_cache/test_mem_pool_host.py b/test/registered/unit/mem_cache/test_mem_pool_host.py index be68761cc..2352e0127 100644 --- a/test/registered/unit/mem_cache/test_mem_pool_host.py +++ b/test/registered/unit/mem_cache/test_mem_pool_host.py @@ -6,12 +6,13 @@ import unittest.mock import torch +from sglang.srt.mem_cache.hicache_storage import PoolName, PoolTransfer from sglang.srt.mem_cache.memory_pool import MHATokenToKVPool from sglang.srt.mem_cache.memory_pool_host import ( DeepSeekV4PagedHostPool, LogicalHostPool, ) -from sglang.srt.mem_cache.pool_host import base +from sglang.srt.mem_cache.pool_host import HostPoolGroup, PoolEntry, base from sglang.srt.mem_cache.pool_host.mamba import MambaPoolHost from sglang.srt.mem_cache.pool_host.mha import MHATokenToKVPoolHost from sglang.srt.runtime_context import get_context @@ -238,5 +239,54 @@ class TestHostMemoryBudget(CustomTestCase): self.assertEqual(base.ranks_per_host(), 8) +class TestHostPoolGroup(CustomTestCase): + @staticmethod + def _group(**sizes): + return HostPoolGroup( + [ + PoolEntry( + name=PoolName(name), + host_pool=LogicalHostPool(size=size, page_size=1), + device_pool=None, + layer_mapper=lambda layer_id: layer_id, + is_primary_index_anchor=name == PoolName.KV.value, + ) + for name, size in sizes.items() + ] + ) + + def test_resolve_and_release_multi_pool_allocation(self): + group = self._group(kv=4, swa=2) + primary = group.alloc(2) + transfers = [ + PoolTransfer(name=PoolName.SWA, device_indices=torch.arange(2)), + PoolTransfer(name=PoolName.INDEXER, indices_from_pool=PoolName.SWA), + ] + + self.assertIsNotNone( + group.resolve_host_transfers( + transfers, + primary_device_indices=torch.arange(2), + primary_host_indices=primary, + ) + ) + self.assertIs(transfers[1].host_indices, transfers[0].host_indices) + group.free(primary) + group.release_transfers(transfers) + self.assertEqual(group.available_size(), 4) + self.assertEqual(group.available_size(PoolName.SWA), 2) + + def test_resolve_rolls_back_partial_allocation(self): + group = self._group(kv=4, swa=2, mamba=1) + transfers = [ + PoolTransfer(name=PoolName.SWA, device_indices=torch.arange(2)), + PoolTransfer(name=PoolName.MAMBA, device_indices=torch.arange(2)), + ] + + self.assertIsNone(group.resolve_host_transfers(transfers)) + self.assertIsNone(transfers[0].host_indices) + self.assertEqual(group.available_size(PoolName.SWA), 2) + + if __name__ == "__main__": unittest.main() 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 67a8a94a0..90a14b8c6 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 @@ -5518,7 +5518,7 @@ class UnifiedRadixCacheSuite: self.assertEqual(xfer.nodes_to_load, [n.id for n in loaded_nodes]) # Allocate SWA device slots from the inner allocator (mirrors how - # _resolve_pool_transfers_allocation routes via device_alloc_fn -> + # _resolve_device_transfers routes via device_alloc_fn -> # swa_attn_allocator.alloc on the load-back path). n_swa = int(xfer.host_indices.numel()) new_swa = allocator.swa_attn_allocator.alloc(n_swa) diff --git a/test/registered/unit/test_supplied_instance_exposure_ratchet.py b/test/registered/unit/test_supplied_instance_exposure_ratchet.py index 3b4f6b31c..333855ed4 100644 --- a/test/registered/unit/test_supplied_instance_exposure_ratchet.py +++ b/test/registered/unit/test_supplied_instance_exposure_ratchet.py @@ -219,7 +219,6 @@ _OVERRIDDEN_AND_READ = { ("weight_cache/daemon.py", "model_path"), ("configs/model_config.py", "dtype"), ("configs/model_config.py", "model_path"), - ("mem_cache/kv_cache_builder.py", "hicache_storage_backend"), ("mem_cache/pool_host/common.py", "hicache_storage_backend"), ("mem_cache/pool_host/common.py", "hicache_storage_backend_extra_config"), ("mem_cache/unified_radix_cache.py", "hicache_storage_backend"),