Refactor HiCache host pool management (#36232)
This commit is contained in:
@@ -318,6 +318,7 @@ class HiCacheController:
|
|||||||
mem_pool_device = mem_pool_device.full_kv_pool
|
mem_pool_device = mem_pool_device.full_kv_pool
|
||||||
self.mem_pool_device = mem_pool_device
|
self.mem_pool_device = mem_pool_device
|
||||||
self.mem_pool_host = mem_pool_host
|
self.mem_pool_host = mem_pool_host
|
||||||
|
self.storage_host_pool = mem_pool_host
|
||||||
self.write_policy = write_policy
|
self.write_policy = write_policy
|
||||||
self.page_size = page_size
|
self.page_size = page_size
|
||||||
self.io_backend = io_backend
|
self.io_backend = io_backend
|
||||||
@@ -329,15 +330,6 @@ class HiCacheController:
|
|||||||
# limiter subtracts write staging from actual pool usage.
|
# limiter subtracts write staging from actual pool usage.
|
||||||
self.host_write_staged_tokens_fn: Optional[Callable[[], int]] = None
|
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).
|
# Default storage page IO functions (may be overridden by attach).
|
||||||
self.page_get_func = self._generic_page_get
|
self.page_get_func = self._generic_page_get
|
||||||
self.page_set_func = self._generic_page_set
|
self.page_set_func = self._generic_page_set
|
||||||
@@ -570,9 +562,9 @@ class HiCacheController:
|
|||||||
|
|
||||||
try:
|
try:
|
||||||
self.storage_backend = StorageBackendFactory.create_backend(
|
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
|
self.enable_storage = True
|
||||||
# todo: threshold policy for prefetching
|
# todo: threshold policy for prefetching
|
||||||
@@ -609,8 +601,6 @@ class HiCacheController:
|
|||||||
self.page_get_func = self._page_get_zero_copy
|
self.page_get_func = self._page_get_zero_copy
|
||||||
self.page_set_func = self._page_set_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.
|
# Ensure stop_event is clear before starting threads.
|
||||||
self.storage_stop_event.clear()
|
self.storage_stop_event.clear()
|
||||||
self._start_storage_threads()
|
self._start_storage_threads()
|
||||||
@@ -638,8 +628,6 @@ class HiCacheController:
|
|||||||
self.enable_storage = False
|
self.enable_storage = False
|
||||||
self.page_get_func = self._generic_page_get
|
self.page_get_func = self._generic_page_get
|
||||||
self.page_set_func = self._generic_page_set
|
self.page_set_func = self._generic_page_set
|
||||||
self.draft_page_get_func = None
|
|
||||||
self.draft_page_set_func = None
|
|
||||||
raise
|
raise
|
||||||
|
|
||||||
def detach_storage_backend(self):
|
def detach_storage_backend(self):
|
||||||
@@ -685,8 +673,6 @@ class HiCacheController:
|
|||||||
self.enable_storage = False
|
self.enable_storage = False
|
||||||
self.page_get_func = self._generic_page_get
|
self.page_get_func = self._generic_page_get
|
||||||
self.page_set_func = self._generic_page_set
|
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.
|
# Now it's safe to clear the stop event for future re-attach.
|
||||||
self.storage_stop_event.clear()
|
self.storage_stop_event.clear()
|
||||||
|
|
||||||
@@ -838,12 +824,7 @@ class HiCacheController:
|
|||||||
)
|
)
|
||||||
|
|
||||||
def _transfer_num_bytes(self, op: CacheOperation) -> int:
|
def _transfer_num_bytes(self, op: CacheOperation) -> int:
|
||||||
"""Total bytes moved by a merged transfer op (draft piggyback included)."""
|
return len(op.device_indices) * self.mem_pool_host.size_per_token
|
||||||
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
|
|
||||||
|
|
||||||
def _num_tokens_by_pool(self, op: CacheOperation) -> dict[str, int]:
|
def _num_tokens_by_pool(self, op: CacheOperation) -> dict[str, int]:
|
||||||
return {PoolName.KV.value: len(op.device_indices)}
|
return {PoolName.KV.value: len(op.device_indices)}
|
||||||
@@ -920,15 +901,6 @@ class HiCacheController:
|
|||||||
device_indices=device_indices,
|
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
|
return transfers
|
||||||
|
|
||||||
def _l2_load_transfers(
|
def _l2_load_transfers(
|
||||||
@@ -981,63 +953,6 @@ class HiCacheController:
|
|||||||
self.mem_pool_host.free(host_indices)
|
self.mem_pool_host.free(host_indices)
|
||||||
return len(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(
|
def prefetch(
|
||||||
self,
|
self,
|
||||||
request_id: str,
|
request_id: str,
|
||||||
@@ -1094,7 +1009,7 @@ class HiCacheController:
|
|||||||
self, operation, hash_values, host_indices, extra_info=None
|
self, operation, hash_values, host_indices, extra_info=None
|
||||||
) -> int:
|
) -> int:
|
||||||
dummy_page_dst = [
|
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)
|
page_data = self.storage_backend.batch_get(hash_values, dummy_page_dst)
|
||||||
if page_data is None:
|
if page_data is None:
|
||||||
@@ -1108,7 +1023,7 @@ class HiCacheController:
|
|||||||
break
|
break
|
||||||
if operation.is_terminated():
|
if operation.is_terminated():
|
||||||
break
|
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],
|
host_indices[i * self.page_size],
|
||||||
page_data[i],
|
page_data[i],
|
||||||
)
|
)
|
||||||
@@ -1137,12 +1052,6 @@ class HiCacheController:
|
|||||||
i * self.page_size : (i + len(batch_hashes)) * self.page_size
|
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
|
# Get one batch token, and update the completed_tokens if succeed
|
||||||
extra_info = HiCacheStorageExtraInfo(prefix_keys=prefix_keys)
|
extra_info = HiCacheStorageExtraInfo(prefix_keys=prefix_keys)
|
||||||
|
|
||||||
@@ -1325,7 +1234,7 @@ class HiCacheController:
|
|||||||
# todo: deprecate
|
# todo: deprecate
|
||||||
def _generic_page_set(self, hash_values, host_indices, extra_info=None) -> bool:
|
def _generic_page_set(self, hash_values, host_indices, extra_info=None) -> bool:
|
||||||
data = [
|
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))
|
for i in range(len(hash_values))
|
||||||
]
|
]
|
||||||
return self.storage_backend.batch_set(hash_values, data)
|
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)
|
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
|
# Backup batch by batch
|
||||||
def _page_backup(self, operation):
|
def _page_backup(self, operation):
|
||||||
# Backup batch by batch
|
# Backup batch by batch
|
||||||
@@ -1420,10 +1263,6 @@ class HiCacheController:
|
|||||||
)
|
)
|
||||||
break
|
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:
|
if prefix_keys and len(prefix_keys) > 0:
|
||||||
prefix_keys += batch_hashes
|
prefix_keys += batch_hashes
|
||||||
operation.completed_tokens += self.page_size * len(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,
|
count_pool_hits,
|
||||||
)
|
)
|
||||||
from sglang.srt.mem_cache.l2_transfer import L2Transfer
|
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
|
from sglang.srt.mem_cache.pool_host.mha import MHATokenToKVPoolHost
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
@@ -136,6 +136,7 @@ class HybridCacheController(BaseHiCacheController):
|
|||||||
self.layer_num = transfer_layer_num
|
self.layer_num = transfer_layer_num
|
||||||
self.layer_done_counter = LayerDoneCounter(self.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:
|
if startup_storage_backend is not None:
|
||||||
self.attach_storage_backend(
|
self.attach_storage_backend(
|
||||||
storage_backend=startup_storage_backend,
|
storage_backend=startup_storage_backend,
|
||||||
@@ -315,11 +316,10 @@ class HybridCacheController(BaseHiCacheController):
|
|||||||
host_indices = self.mem_pool_host.alloc(len(device_indices))
|
host_indices = self.mem_pool_host.alloc(len(device_indices))
|
||||||
if host_indices is None:
|
if host_indices is None:
|
||||||
return None
|
return None
|
||||||
pool_transfers = self._resolve_pool_transfers_allocation(
|
pool_transfers = self.mem_pool_host.resolve_host_transfers(
|
||||||
extra_pools,
|
extra_pools,
|
||||||
alloc_host=True,
|
primary_device_indices=device_indices,
|
||||||
kv_device_indices=device_indices,
|
primary_host_indices=host_indices,
|
||||||
kv_host_indices=host_indices,
|
|
||||||
)
|
)
|
||||||
if pool_transfers is None and extra_pools:
|
if pool_transfers is None and extra_pools:
|
||||||
self.mem_pool_host.free(host_indices)
|
self.mem_pool_host.free(host_indices)
|
||||||
@@ -416,15 +416,6 @@ class HybridCacheController(BaseHiCacheController):
|
|||||||
layer_mapper=entry.layer_mapper,
|
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
|
return transfers
|
||||||
|
|
||||||
def _l2_load_transfers(
|
def _l2_load_transfers(
|
||||||
@@ -434,13 +425,17 @@ class HybridCacheController(BaseHiCacheController):
|
|||||||
pool_transfers: Optional[list[PoolTransfer]] = None,
|
pool_transfers: Optional[list[PoolTransfer]] = None,
|
||||||
) -> list[L2Transfer]:
|
) -> list[L2Transfer]:
|
||||||
transfers = self._l2_transfers(host_indices, device_indices, pool_transfers)
|
transfers = self._l2_transfers(host_indices, device_indices, pool_transfers)
|
||||||
if getattr(self, "has_mtp_draft", False):
|
transfers_by_entry = {
|
||||||
target_transfers = list(transfers)
|
(id(t.host_pool), id(t.device_pool)): t for t in transfers
|
||||||
for depth, draft_device_pool in enumerate(self.mtp_draft_device_pools):
|
}
|
||||||
for transfer in target_transfers:
|
for entry in self.mem_pool_host.entry_map.values():
|
||||||
if transfer.layer_mapper is None:
|
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
|
continue
|
||||||
draft_host_layer = transfer.layer_mapper(self.layer_num + depth)
|
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:
|
if draft_host_layer is None:
|
||||||
continue
|
continue
|
||||||
|
|
||||||
@@ -456,10 +451,10 @@ class HybridCacheController(BaseHiCacheController):
|
|||||||
|
|
||||||
transfers.append(
|
transfers.append(
|
||||||
L2Transfer(
|
L2Transfer(
|
||||||
host_pool=transfer.host_pool,
|
host_pool=target_transfer.host_pool,
|
||||||
device_pool=draft_device_pool,
|
device_pool=draft_device_pool,
|
||||||
host_indices=transfer.host_indices,
|
host_indices=target_transfer.host_indices,
|
||||||
device_indices=transfer.device_indices,
|
device_indices=target_transfer.device_indices,
|
||||||
layer_mapper=draft_layer_mapper,
|
layer_mapper=draft_layer_mapper,
|
||||||
is_draft=True,
|
is_draft=True,
|
||||||
)
|
)
|
||||||
@@ -479,13 +474,13 @@ class HybridCacheController(BaseHiCacheController):
|
|||||||
return counts
|
return counts
|
||||||
|
|
||||||
def _transfer_num_bytes(self, op: CacheOperation) -> int:
|
def _transfer_num_bytes(self, op: CacheOperation) -> int:
|
||||||
"""Total bytes moved by a merged transfer op across all pools,
|
"""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)."""
|
Sidecar transfers riding another pool's indices are included here but
|
||||||
|
excluded from the per-pool token counts.
|
||||||
|
"""
|
||||||
kv_tokens = len(op.device_indices)
|
kv_tokens = len(op.device_indices)
|
||||||
num_bytes = kv_tokens * self.mem_pool_host.anchor_entry.host_pool.size_per_token
|
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.
|
# Slot counts of the pools sidecars can ride on.
|
||||||
source_len = {self.mem_pool_host.anchor_entry.name: kv_tokens}
|
source_len = {self.mem_pool_host.anchor_entry.name: kv_tokens}
|
||||||
for t in op.pool_transfers or []:
|
for t in op.pool_transfers or []:
|
||||||
@@ -523,9 +518,8 @@ class HybridCacheController(BaseHiCacheController):
|
|||||||
if device_indices is None:
|
if device_indices is None:
|
||||||
return None
|
return None
|
||||||
|
|
||||||
pool_transfers = self._resolve_pool_transfers_allocation(
|
pool_transfers = self._resolve_device_transfers(
|
||||||
extra_pools,
|
extra_pools,
|
||||||
alloc_host=False,
|
|
||||||
kv_device_indices=device_indices,
|
kv_device_indices=device_indices,
|
||||||
kv_host_indices=host_indices,
|
kv_host_indices=host_indices,
|
||||||
)
|
)
|
||||||
@@ -833,26 +827,21 @@ class HybridCacheController(BaseHiCacheController):
|
|||||||
)
|
)
|
||||||
transfer.host_indices = transfer.host_indices[:needed]
|
transfer.host_indices = transfer.host_indices[:needed]
|
||||||
|
|
||||||
def _resolve_pool_transfers_allocation(
|
def _resolve_device_transfers(
|
||||||
self,
|
self,
|
||||||
extra_pools: Optional[list[PoolTransfer]],
|
extra_pools: Optional[list[PoolTransfer]],
|
||||||
alloc_host: bool,
|
|
||||||
kv_device_indices: Optional[torch.Tensor] = None,
|
kv_device_indices: Optional[torch.Tensor] = None,
|
||||||
kv_host_indices: Optional[torch.Tensor] = None,
|
kv_host_indices: Optional[torch.Tensor] = None,
|
||||||
) -> Optional[list[PoolTransfer]]:
|
) -> 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:
|
if not extra_pools:
|
||||||
return None
|
return None
|
||||||
# (pool, free_fn, indices) for atomic rollback on failure.
|
|
||||||
newly_allocated: list[tuple[PoolTransfer, Callable, torch.Tensor]] = []
|
newly_allocated: list[tuple[PoolTransfer, Callable, torch.Tensor]] = []
|
||||||
derived_transfers: list[PoolTransfer] = []
|
derived_transfers: list[PoolTransfer] = []
|
||||||
|
|
||||||
def rollback_allocated() -> None:
|
def rollback_allocated() -> None:
|
||||||
for prev_pool, prev_free_fn, prev_indices in newly_allocated:
|
for prev_pool, prev_free_fn, prev_indices in newly_allocated:
|
||||||
prev_free_fn(prev_indices)
|
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:
|
for pool in extra_pools:
|
||||||
@@ -862,14 +851,6 @@ class HybridCacheController(BaseHiCacheController):
|
|||||||
entry = self.mem_pool_host.entry_map.get(pool.name)
|
entry = self.mem_pool_host.entry_map.get(pool.name)
|
||||||
if entry is None:
|
if entry is None:
|
||||||
continue
|
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:
|
if pool.device_indices is not None or pool.host_indices is None:
|
||||||
continue
|
continue
|
||||||
# device_alloc_fn / device_free_fn override entry.device_pool's
|
# device_alloc_fn / device_free_fn override entry.device_pool's
|
||||||
@@ -887,9 +868,6 @@ class HybridCacheController(BaseHiCacheController):
|
|||||||
# Atomic rollback: free everything we successfully allocated.
|
# Atomic rollback: free everything we successfully allocated.
|
||||||
rollback_allocated()
|
rollback_allocated()
|
||||||
return None
|
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))
|
newly_allocated.append((pool, free_fn, indices))
|
||||||
|
|
||||||
|
|||||||
@@ -15,10 +15,9 @@ from sglang.srt.mem_cache.hybrid_cache.hybrid_cache_controller import (
|
|||||||
from sglang.srt.mem_cache.memory_pool_host import (
|
from sglang.srt.mem_cache.memory_pool_host import (
|
||||||
DeepSeekV4PagedHostPool,
|
DeepSeekV4PagedHostPool,
|
||||||
DeepSeekV4StateHostPool,
|
DeepSeekV4StateHostPool,
|
||||||
HostPoolGroup,
|
|
||||||
LogicalHostPool,
|
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.common import get_allocator_type
|
||||||
from sglang.srt.mem_cache.pool_host.dsa import DSAIndexerPoolHost
|
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.mamba import MambaPoolHost
|
||||||
@@ -141,6 +140,7 @@ def build_pool_entry(
|
|||||||
device_evict_fn: Optional[Callable[[int], Any]] = None,
|
device_evict_fn: Optional[Callable[[int], Any]] = None,
|
||||||
device_alloc_fn: Optional[Callable[[int], Any]] = None,
|
device_alloc_fn: Optional[Callable[[int], Any]] = None,
|
||||||
device_free_fn: Optional[Callable[[Any], Any]] = None,
|
device_free_fn: Optional[Callable[[Any], Any]] = None,
|
||||||
|
packed_draft_device_pools: tuple[Any, ...] = (),
|
||||||
) -> PoolEntry:
|
) -> PoolEntry:
|
||||||
return PoolEntry(
|
return PoolEntry(
|
||||||
name=name,
|
name=name,
|
||||||
@@ -152,6 +152,7 @@ def build_pool_entry(
|
|||||||
device_evict_fn=device_evict_fn,
|
device_evict_fn=device_evict_fn,
|
||||||
device_alloc_fn=device_alloc_fn,
|
device_alloc_fn=device_alloc_fn,
|
||||||
device_free_fn=device_free_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,
|
layer_mapping=full_layer_mapping,
|
||||||
transfer_layer_num=transfer_layer_num + len(mtp_draft_device_pools),
|
transfer_layer_num=transfer_layer_num + len(mtp_draft_device_pools),
|
||||||
is_anchor=True,
|
is_anchor=True,
|
||||||
|
packed_draft_device_pools=mtp_draft_device_pools,
|
||||||
)
|
)
|
||||||
]
|
]
|
||||||
)
|
)
|
||||||
@@ -264,6 +266,7 @@ def build_hybrid_swa_group(
|
|||||||
device_free_fn=(
|
device_free_fn=(
|
||||||
swa_attn_allocator.free if swa_attn_allocator is not None else None
|
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,
|
enable_storage_metrics=enable_storage_metrics,
|
||||||
host_memory_mode=server_args.hicache_host_memory_mode,
|
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
|
return host_pool_group, cache_controller
|
||||||
|
|
||||||
|
|
||||||
@@ -384,8 +384,6 @@ def build_hybrid_swa_stack(
|
|||||||
enable_storage_metrics=enable_storage_metrics,
|
enable_storage_metrics=enable_storage_metrics,
|
||||||
host_memory_mode=server_args.hicache_host_memory_mode,
|
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
|
return host_pool_group, cache_controller
|
||||||
|
|
||||||
|
|
||||||
@@ -538,6 +536,7 @@ def build_deepseek_v4_hicache_stack(
|
|||||||
device_evict_fn=device_swa_evict_fn,
|
device_evict_fn=device_swa_evict_fn,
|
||||||
device_alloc_fn=swa_attn_allocator.alloc,
|
device_alloc_fn=swa_attn_allocator.alloc,
|
||||||
device_free_fn=swa_attn_allocator.free,
|
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,
|
enable_storage_metrics=enable_storage_metrics,
|
||||||
host_memory_mode=server_args.hicache_host_memory_mode,
|
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
|
return host_pool_group, cache_controller
|
||||||
|
|
||||||
|
|
||||||
@@ -735,6 +732,7 @@ def build_hybrid_mamba_stack(
|
|||||||
layer_mapping=full_layer_mapping,
|
layer_mapping=full_layer_mapping,
|
||||||
transfer_layer_num=transfer_layer_num + len(mtp_draft_device_pools),
|
transfer_layer_num=transfer_layer_num + len(mtp_draft_device_pools),
|
||||||
is_anchor=True,
|
is_anchor=True,
|
||||||
|
packed_draft_device_pools=mtp_draft_device_pools,
|
||||||
),
|
),
|
||||||
build_pool_entry(
|
build_pool_entry(
|
||||||
name=PoolName.MAMBA,
|
name=PoolName.MAMBA,
|
||||||
@@ -768,8 +766,6 @@ def build_hybrid_mamba_stack(
|
|||||||
enable_storage_metrics=enable_storage_metrics,
|
enable_storage_metrics=enable_storage_metrics,
|
||||||
host_memory_mode=server_args.hicache_host_memory_mode,
|
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
|
return host_pool_group, cache_controller
|
||||||
|
|
||||||
|
|
||||||
@@ -933,6 +929,7 @@ def build_anchor_sidecar_stack(
|
|||||||
layer_mapping=full_layer_mapping,
|
layer_mapping=full_layer_mapping,
|
||||||
transfer_layer_num=transfer_layer_num + len(mtp_draft_device_pools),
|
transfer_layer_num=transfer_layer_num + len(mtp_draft_device_pools),
|
||||||
is_anchor=True,
|
is_anchor=True,
|
||||||
|
packed_draft_device_pools=mtp_draft_device_pools,
|
||||||
),
|
),
|
||||||
build_pool_entry(
|
build_pool_entry(
|
||||||
name=sidecar_pool_name,
|
name=sidecar_pool_name,
|
||||||
@@ -940,6 +937,7 @@ def build_anchor_sidecar_stack(
|
|||||||
device_pool=kv_pool,
|
device_pool=kv_pool,
|
||||||
layer_mapping=full_layer_mapping,
|
layer_mapping=full_layer_mapping,
|
||||||
transfer_layer_num=transfer_layer_num + len(mtp_draft_device_pools),
|
transfer_layer_num=transfer_layer_num + len(mtp_draft_device_pools),
|
||||||
|
packed_draft_device_pools=mtp_draft_device_pools,
|
||||||
),
|
),
|
||||||
]
|
]
|
||||||
host_pool_group = HostPoolGroup(entries)
|
host_pool_group = HostPoolGroup(entries)
|
||||||
@@ -962,8 +960,6 @@ def build_anchor_sidecar_stack(
|
|||||||
enable_storage_metrics=enable_storage_metrics,
|
enable_storage_metrics=enable_storage_metrics,
|
||||||
host_memory_mode=server_args.hicache_host_memory_mode,
|
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
|
return host_pool_group, cache_controller
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -66,7 +66,6 @@ def maybe_register_hicache_draft(
|
|||||||
tree_cache,
|
tree_cache,
|
||||||
draft_plan: HiCacheDraftPlan,
|
draft_plan: HiCacheDraftPlan,
|
||||||
server_args: ServerArgs,
|
server_args: ServerArgs,
|
||||||
page_size: int,
|
|
||||||
) -> None:
|
) -> None:
|
||||||
from sglang.srt.speculative.base_spec_worker import HiCacheDraftMode
|
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
|
from sglang.srt.mem_cache.unified_radix_cache import UnifiedRadixCache
|
||||||
|
|
||||||
if not isinstance(tree_cache, UnifiedRadixCache):
|
if not isinstance(tree_cache, UnifiedRadixCache):
|
||||||
_register_legacy_hicache_draft(
|
raise NotImplementedError("HiCache draft pools require UnifiedRadixCache.")
|
||||||
tree_cache=tree_cache,
|
|
||||||
draft_pool=draft_plan.device_pools[0],
|
|
||||||
server_args=server_args,
|
|
||||||
page_size=page_size,
|
|
||||||
)
|
|
||||||
return
|
|
||||||
|
|
||||||
from sglang.srt.mem_cache.hybrid_cache.hybrid_pool_assembler import (
|
from sglang.srt.mem_cache.hybrid_cache.hybrid_pool_assembler import (
|
||||||
build_hicache_draft_sidecars,
|
build_hicache_draft_sidecars,
|
||||||
@@ -93,51 +86,8 @@ def maybe_register_hicache_draft(
|
|||||||
tree_cache=tree_cache,
|
tree_cache=tree_cache,
|
||||||
server_args=server_args,
|
server_args=server_args,
|
||||||
)
|
)
|
||||||
tree_cache.register_hicache_draft_pools(specs, entries)
|
for spec, entry in zip(specs, entries, strict=True):
|
||||||
|
tree_cache.register_sidecar_pool(spec, entry)
|
||||||
|
|
||||||
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)
|
|
||||||
|
|
||||||
|
|
||||||
# Host slots a backup-only retraction pool gets, as a fraction of the device
|
# 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,
|
tree_cache=tree_cache,
|
||||||
draft_plan=hicache_draft_plan,
|
draft_plan=hicache_draft_plan,
|
||||||
server_args=server_args,
|
server_args=server_args,
|
||||||
page_size=page_size,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
if retraction_backup == "host_pool":
|
if retraction_backup == "host_pool":
|
||||||
|
|||||||
@@ -2,11 +2,7 @@ from __future__ import annotations
|
|||||||
|
|
||||||
import logging
|
import logging
|
||||||
import threading
|
import threading
|
||||||
from dataclasses import dataclass
|
from typing import Optional
|
||||||
from typing import TYPE_CHECKING, Any, Callable, Optional
|
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
|
||||||
from sglang.srt.mem_cache.hicache_storage import PoolName
|
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
@@ -955,131 +951,3 @@ class DeepSeekV4StateHostPool(HostKVCache):
|
|||||||
self.kv_buffer.data_ptr() % page_size_bytes == 0
|
self.kv_buffer.data_ptr() % page_size_bytes == 0
|
||||||
and page_bytes % 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.base import HostKVCache
|
||||||
from sglang.srt.mem_cache.pool_host.common import HostTensorAllocator
|
from sglang.srt.mem_cache.pool_host.common import HostTensorAllocator
|
||||||
|
from sglang.srt.mem_cache.pool_host.group import HostPoolGroup, PoolEntry
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
"HostKVCache",
|
"HostKVCache",
|
||||||
|
"HostPoolGroup",
|
||||||
"HostTensorAllocator",
|
"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:
|
if self._full_kv_pool_host is None:
|
||||||
return
|
return
|
||||||
for host_value in host_values:
|
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:
|
def apply_component_action(self, action: ComponentAction) -> None:
|
||||||
if isinstance(action, FreeComponentDeviceSlot):
|
if isinstance(action, FreeComponentDeviceSlot):
|
||||||
|
|||||||
@@ -680,10 +680,11 @@ class MambaComponent(TreeComponent):
|
|||||||
*,
|
*,
|
||||||
prefetch_tokens: int = 0,
|
prefetch_tokens: int = 0,
|
||||||
) -> PreparePrefetchResult:
|
) -> PreparePrefetchResult:
|
||||||
host_indices = self._mamba_pool_host.alloc(1)
|
host_indices = self.cache.host_pool_group.alloc(
|
||||||
if host_indices is None:
|
1,
|
||||||
self.cache.evict_host(1, ComponentType.MAMBA)
|
pool=PoolName.MAMBA,
|
||||||
host_indices = self._mamba_pool_host.alloc(1)
|
reclaim=lambda size: self.cache.evict_host(size, ComponentType.MAMBA),
|
||||||
|
)
|
||||||
if host_indices is None:
|
if host_indices is None:
|
||||||
return PreparePrefetchResult(alloc_failed=True)
|
return PreparePrefetchResult(alloc_failed=True)
|
||||||
return PreparePrefetchResult(host_indices=host_indices)
|
return PreparePrefetchResult(host_indices=host_indices)
|
||||||
@@ -897,7 +898,7 @@ class MambaComponent(TreeComponent):
|
|||||||
if self._mamba_pool_host is None:
|
if self._mamba_pool_host is None:
|
||||||
return
|
return
|
||||||
for host_value in host_values:
|
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:
|
def apply_component_action(self, action: ComponentAction) -> None:
|
||||||
if isinstance(action, MambaEvictExcessPathStates):
|
if isinstance(action, MambaEvictExcessPathStates):
|
||||||
|
|||||||
@@ -794,10 +794,11 @@ class SWAComponent(TreeComponent):
|
|||||||
# device-guaranteed, require a full window.
|
# device-guaranteed, require a full window.
|
||||||
return PreparePrefetchResult()
|
return PreparePrefetchResult()
|
||||||
num_tokens = num_pages * self.cache.page_size
|
num_tokens = num_pages * self.cache.page_size
|
||||||
host_indices = self._swa_kv_pool_host.alloc(num_tokens)
|
host_indices = self.cache.host_pool_group.alloc(
|
||||||
if host_indices is None:
|
num_tokens,
|
||||||
self.cache.evict_host(num_tokens, ComponentType.SWA)
|
pool=PoolName.SWA,
|
||||||
host_indices = self._swa_kv_pool_host.alloc(num_tokens)
|
reclaim=lambda size: self.cache.evict_host(size, ComponentType.SWA),
|
||||||
|
)
|
||||||
if host_indices is None:
|
if host_indices is None:
|
||||||
return PreparePrefetchResult(alloc_failed=True)
|
return PreparePrefetchResult(alloc_failed=True)
|
||||||
return PreparePrefetchResult(host_indices=host_indices)
|
return PreparePrefetchResult(host_indices=host_indices)
|
||||||
@@ -1121,7 +1122,7 @@ class SWAComponent(TreeComponent):
|
|||||||
if self._swa_kv_pool_host is None:
|
if self._swa_kv_pool_host is None:
|
||||||
return
|
return
|
||||||
for host_value in host_values:
|
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:
|
def apply_component_action(self, action: ComponentAction) -> None:
|
||||||
alloc = self.cache.token_to_kv_pool_allocator
|
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 (
|
from sglang.srt.mem_cache.hybrid_cache.hybrid_cache_controller import (
|
||||||
PrefetchOperation,
|
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.server_args import ServerArgs
|
||||||
|
|
||||||
from sglang.srt.utils.rank_consensus_checker import rank_consensus
|
from sglang.srt.utils.rank_consensus_checker import rank_consensus
|
||||||
@@ -475,17 +475,14 @@ class UnifiedRadixCache(BasePrefixCache):
|
|||||||
extra_metric_labels=self.extra_metric_labels,
|
extra_metric_labels=self.extra_metric_labels,
|
||||||
)
|
)
|
||||||
|
|
||||||
def register_sidecar_pool(self, spec: SidecarPoolSpec) -> None:
|
def register_sidecar_pool(
|
||||||
self.sidecar_pool_specs.append(spec)
|
self, spec: SidecarPoolSpec, entry: Optional[PoolEntry] = None
|
||||||
|
|
||||||
def register_hicache_draft_pools(
|
|
||||||
self, specs: list[SidecarPoolSpec], entries: list[PoolEntry]
|
|
||||||
) -> None:
|
) -> None:
|
||||||
|
if entry is not None:
|
||||||
if self.cache_controller is None:
|
if self.cache_controller is None:
|
||||||
raise RuntimeError("HiCache controller is not attached.")
|
raise RuntimeError("HiCache controller is not attached.")
|
||||||
for spec, entry in zip(specs, entries, strict=True):
|
|
||||||
self.cache_controller.register_host_pool_entry(entry)
|
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:
|
def release_host_resources(self) -> None:
|
||||||
if self.host_pool_group is not None:
|
if self.host_pool_group is not None:
|
||||||
@@ -1137,11 +1134,10 @@ class UnifiedRadixCache(BasePrefixCache):
|
|||||||
if host_indices is None:
|
if host_indices is None:
|
||||||
return None
|
return None
|
||||||
|
|
||||||
resolved = self.cache_controller._resolve_pool_transfers_allocation(
|
resolved = self.host_pool_group.resolve_host_transfers(
|
||||||
extra_transfers or None,
|
extra_transfers or None,
|
||||||
alloc_host=True,
|
primary_device_indices=device_indices,
|
||||||
kv_device_indices=device_indices,
|
primary_host_indices=host_indices,
|
||||||
kv_host_indices=host_indices,
|
|
||||||
)
|
)
|
||||||
if resolved is None and extra_transfers:
|
if resolved is None and extra_transfers:
|
||||||
self.host_pool_group.free(host_indices)
|
self.host_pool_group.free(host_indices)
|
||||||
@@ -1195,9 +1191,8 @@ class UnifiedRadixCache(BasePrefixCache):
|
|||||||
)
|
)
|
||||||
for name, saved in saved_by_name.items()
|
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,
|
restored_transfers or None,
|
||||||
alloc_host=False,
|
|
||||||
kv_device_indices=device_indices,
|
kv_device_indices=device_indices,
|
||||||
kv_host_indices=backup.host_indices,
|
kv_host_indices=backup.host_indices,
|
||||||
)
|
)
|
||||||
@@ -1223,10 +1218,7 @@ class UnifiedRadixCache(BasePrefixCache):
|
|||||||
|
|
||||||
def retraction_discard(self, backup: RetractionBackup) -> None:
|
def retraction_discard(self, backup: RetractionBackup) -> None:
|
||||||
self.host_pool_group.free(backup.host_indices)
|
self.host_pool_group.free(backup.host_indices)
|
||||||
for transfer in backup.pool_transfers or []:
|
self.host_pool_group.release_transfers(backup.pool_transfers)
|
||||||
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)
|
|
||||||
|
|
||||||
# ---- HiCache: Backup / LoadBack ----
|
# ---- HiCache: Backup / LoadBack ----
|
||||||
|
|
||||||
@@ -2259,9 +2251,9 @@ class UnifiedRadixCache(BasePrefixCache):
|
|||||||
host_indices_list.append(host_indices)
|
host_indices_list.append(host_indices)
|
||||||
released_tokens += len(host_indices)
|
released_tokens += len(host_indices)
|
||||||
if host_indices_list:
|
if host_indices_list:
|
||||||
entry = cc.mem_pool_host.entry_map.get(pool_name)
|
cc.mem_pool_host.free(
|
||||||
if entry is not None:
|
torch.cat(host_indices_list, dim=0), pool=pool_name
|
||||||
entry.host_pool.free(torch.cat(host_indices_list, dim=0))
|
)
|
||||||
drained[pool_name] = (len(host_indices_list), released_tokens)
|
drained[pool_name] = (len(host_indices_list), released_tokens)
|
||||||
return drained
|
return drained
|
||||||
|
|
||||||
|
|||||||
@@ -117,7 +117,6 @@ class TestDecodeRetractionBackup(unittest.TestCase):
|
|||||||
device_pools=(draft_pool,),
|
device_pools=(draft_pool,),
|
||||||
),
|
),
|
||||||
server_args=server_args,
|
server_args=server_args,
|
||||||
page_size=1,
|
|
||||||
)
|
)
|
||||||
self.assertIn(PoolName.DRAFT, cache.host_pool_group.entry_map)
|
self.assertIn(PoolName.DRAFT, cache.host_pool_group.entry_map)
|
||||||
cache.validate_retraction_host_capacity()
|
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 (
|
from sglang.srt.mem_cache.memory_pool_host import (
|
||||||
DeepSeekV4PagedHostPool,
|
DeepSeekV4PagedHostPool,
|
||||||
DeepSeekV4StateHostPool,
|
DeepSeekV4StateHostPool,
|
||||||
HostPoolGroup,
|
|
||||||
LogicalHostPool,
|
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.dsa import DSAIndexerPoolHost
|
||||||
from sglang.srt.mem_cache.pool_host.mamba import MambaPoolHost
|
from sglang.srt.mem_cache.pool_host.mamba import MambaPoolHost
|
||||||
from sglang.srt.mem_cache.pool_host.mha import MHATokenToKVPoolHost
|
from sglang.srt.mem_cache.pool_host.mha import MHATokenToKVPoolHost
|
||||||
@@ -222,8 +221,6 @@ class TestHiCacheStagedWriteBackDispatch(CustomTestCase):
|
|||||||
op.pool_transfers,
|
op.pool_transfers,
|
||||||
)
|
)
|
||||||
controller.mem_pool_host = _host_group_stub([], can_use_write_back_jit=False)
|
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: (
|
controller._l2_transfers.side_effect = lambda *args: (
|
||||||
HybridCacheController._l2_transfers(controller, *args)
|
HybridCacheController._l2_transfers(controller, *args)
|
||||||
)
|
)
|
||||||
@@ -287,20 +284,19 @@ class TestHiCacheStagedWriteBackDispatch(CustomTestCase):
|
|||||||
def test_packed_draft_load_is_flattened_into_l2_transfers(self):
|
def test_packed_draft_load_is_flattened_into_l2_transfers(self):
|
||||||
host_pool = mock.Mock()
|
host_pool = mock.Mock()
|
||||||
controller = HybridCacheController.__new__(HybridCacheController)
|
controller = HybridCacheController.__new__(HybridCacheController)
|
||||||
controller.mem_pool_host = SimpleNamespace(
|
entry = PoolEntry(
|
||||||
anchor_entry=PoolEntry(
|
|
||||||
name=PoolName.KV,
|
name=PoolName.KV,
|
||||||
host_pool=host_pool,
|
host_pool=host_pool,
|
||||||
device_pool=mock.sentinel.target_device_pool,
|
device_pool=mock.sentinel.target_device_pool,
|
||||||
layer_mapper={0: 0, 1: 1, 2: 2}.get,
|
layer_mapper={0: 0, 1: 1, 2: 2}.get,
|
||||||
is_primary_index_anchor=True,
|
is_primary_index_anchor=True,
|
||||||
),
|
packed_draft_device_pools=(mock.sentinel.draft_device_pool,),
|
||||||
entry_map={},
|
)
|
||||||
|
controller.mem_pool_host = SimpleNamespace(
|
||||||
|
anchor_entry=entry,
|
||||||
|
entry_map={entry.name: entry},
|
||||||
)
|
)
|
||||||
controller.layer_num = 2
|
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(
|
self.assertEqual(
|
||||||
len(controller._l2_transfers(_indices(0, 2), _indices(2, 4))), 1
|
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
|
captured, can_use_write_back_jit=True
|
||||||
)
|
)
|
||||||
controller.mem_pool_device = None
|
controller.mem_pool_device = None
|
||||||
controller.has_draft = False
|
|
||||||
controller.ack_write_queue = []
|
controller.ack_write_queue = []
|
||||||
controller.move_hybrid_indices = mock.Mock(
|
controller.move_hybrid_indices = mock.Mock(
|
||||||
side_effect=AssertionError(
|
side_effect=AssertionError(
|
||||||
@@ -972,7 +967,6 @@ class TestHiCacheStagedWriteBackDispatch(CustomTestCase):
|
|||||||
captured, can_use_write_back_jit=False
|
captured, can_use_write_back_jit=False
|
||||||
)
|
)
|
||||||
controller.mem_pool_device = None
|
controller.mem_pool_device = None
|
||||||
controller.has_draft = False
|
|
||||||
controller.ack_write_queue = []
|
controller.ack_write_queue = []
|
||||||
controller.move_hybrid_indices = mock.Mock(
|
controller.move_hybrid_indices = mock.Mock(
|
||||||
return_value=(op.host_indices, op.device_indices, op.pool_transfers)
|
return_value=(op.host_indices, op.device_indices, op.pool_transfers)
|
||||||
@@ -1007,7 +1001,6 @@ class TestHiCacheStagedWriteBackDispatch(CustomTestCase):
|
|||||||
controller.io_backend = "kernel"
|
controller.io_backend = "kernel"
|
||||||
controller.mem_pool_host = FakeHostPool()
|
controller.mem_pool_host = FakeHostPool()
|
||||||
controller.mem_pool_device = None
|
controller.mem_pool_device = None
|
||||||
controller.has_draft = False
|
|
||||||
controller.device = "cuda"
|
controller.device = "cuda"
|
||||||
controller.ack_write_queue = []
|
controller.ack_write_queue = []
|
||||||
controller.move_indices = mock.Mock(
|
controller.move_indices = mock.Mock(
|
||||||
@@ -1044,7 +1037,6 @@ class TestHiCacheStagedWriteBackDispatch(CustomTestCase):
|
|||||||
controller.io_backend = "kernel"
|
controller.io_backend = "kernel"
|
||||||
controller.mem_pool_host = FakeHostPool()
|
controller.mem_pool_host = FakeHostPool()
|
||||||
controller.mem_pool_device = None
|
controller.mem_pool_device = None
|
||||||
controller.has_draft = False
|
|
||||||
controller.device = "cuda"
|
controller.device = "cuda"
|
||||||
controller.ack_write_queue = []
|
controller.ack_write_queue = []
|
||||||
controller.move_indices = mock.Mock(
|
controller.move_indices = mock.Mock(
|
||||||
|
|||||||
@@ -6,12 +6,13 @@ import unittest.mock
|
|||||||
|
|
||||||
import torch
|
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 import MHATokenToKVPool
|
||||||
from sglang.srt.mem_cache.memory_pool_host import (
|
from sglang.srt.mem_cache.memory_pool_host import (
|
||||||
DeepSeekV4PagedHostPool,
|
DeepSeekV4PagedHostPool,
|
||||||
LogicalHostPool,
|
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.mamba import MambaPoolHost
|
||||||
from sglang.srt.mem_cache.pool_host.mha import MHATokenToKVPoolHost
|
from sglang.srt.mem_cache.pool_host.mha import MHATokenToKVPoolHost
|
||||||
from sglang.srt.runtime_context import get_context
|
from sglang.srt.runtime_context import get_context
|
||||||
@@ -238,5 +239,54 @@ class TestHostMemoryBudget(CustomTestCase):
|
|||||||
self.assertEqual(base.ranks_per_host(), 8)
|
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__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
@@ -5518,7 +5518,7 @@ class UnifiedRadixCacheSuite:
|
|||||||
self.assertEqual(xfer.nodes_to_load, [n.id for n in loaded_nodes])
|
self.assertEqual(xfer.nodes_to_load, [n.id for n in loaded_nodes])
|
||||||
|
|
||||||
# Allocate SWA device slots from the inner allocator (mirrors how
|
# 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).
|
# swa_attn_allocator.alloc on the load-back path).
|
||||||
n_swa = int(xfer.host_indices.numel())
|
n_swa = int(xfer.host_indices.numel())
|
||||||
new_swa = allocator.swa_attn_allocator.alloc(n_swa)
|
new_swa = allocator.swa_attn_allocator.alloc(n_swa)
|
||||||
|
|||||||
@@ -219,7 +219,6 @@ _OVERRIDDEN_AND_READ = {
|
|||||||
("weight_cache/daemon.py", "model_path"),
|
("weight_cache/daemon.py", "model_path"),
|
||||||
("configs/model_config.py", "dtype"),
|
("configs/model_config.py", "dtype"),
|
||||||
("configs/model_config.py", "model_path"),
|
("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"),
|
||||||
("mem_cache/pool_host/common.py", "hicache_storage_backend_extra_config"),
|
("mem_cache/pool_host/common.py", "hicache_storage_backend_extra_config"),
|
||||||
("mem_cache/unified_radix_cache.py", "hicache_storage_backend"),
|
("mem_cache/unified_radix_cache.py", "hicache_storage_backend"),
|
||||||
|
|||||||
Reference in New Issue
Block a user