[HiCache] feat: support draft offload for mooncake (#24984)
Co-authored-by: huangtingwei9988 <141888744+huangtingwei9988@users.noreply.github.com> Co-authored-by: stmatengss <11641725+stmatengss@users.noreply.github.com>
This commit is contained in:
co-authored by
huangtingwei9988
stmatengss
parent
c3aaafc5f2
commit
f4e7a98fe5
@@ -25,6 +25,8 @@ from sglang.srt.mem_cache.hicache_storage import (
|
||||
STORAGE_BATCH_SIZE,
|
||||
HiCacheStorageConfig,
|
||||
HiCacheStorageExtraInfo,
|
||||
PoolName,
|
||||
PoolTransfer,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
@@ -288,6 +290,8 @@ class HiCacheController:
|
||||
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
|
||||
|
||||
# Default storage page IO functions (may be overridden by attach).
|
||||
self.page_get_func = self._generic_page_get
|
||||
@@ -529,6 +533,8 @@ 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()
|
||||
@@ -553,6 +559,8 @@ 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):
|
||||
@@ -594,6 +602,8 @@ 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()
|
||||
|
||||
@@ -844,6 +854,47 @@ class HiCacheController:
|
||||
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 _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,
|
||||
@@ -1075,44 +1126,71 @@ class HiCacheController:
|
||||
)
|
||||
|
||||
def _draft_page_set(self, hash_values, host_indices) -> None:
|
||||
"""Best-effort write draft KV pages to L3 with 'd:' prefixed keys.
|
||||
|
||||
TODO: support batch_set_v1 (zero-copy) for high-performance backends.
|
||||
"""
|
||||
"""Best-effort write draft KV pages to L3 alongside the target backup."""
|
||||
if self.draft_page_set_func is None:
|
||||
return
|
||||
try:
|
||||
draft_keys = [f"d:{h}" 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)
|
||||
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 with 'd:' prefixed keys.
|
||||
|
||||
TODO: support batch_get_v1 (zero-copy) for high-performance backends.
|
||||
"""
|
||||
"""Best-effort read draft KV pages from L3 (mirrors `_draft_page_set`)."""
|
||||
if self.draft_page_get_func is None:
|
||||
return
|
||||
try:
|
||||
draft_keys = [f"d:{h}" 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
|
||||
)
|
||||
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
|
||||
|
||||
Reference in New Issue
Block a user