Refactor HiCache host pool management (#36232)

This commit is contained in:
cctry
2026-08-25 16:31:57 -07:00
committed by GitHub
parent 0c42a44cd7
commit aa718f7343
16 changed files with 350 additions and 502 deletions
+7 -168
View File
@@ -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)
@@ -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.
@@ -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
@@ -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":
+1 -133
View File
@@ -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)
@@ -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",
]
@@ -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()
@@ -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):
@@ -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):
@@ -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
@@ -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
@@ -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()
@@ -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(
@@ -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()
@@ -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)
@@ -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"),