[HiCache][HA 1/N] Support HiCache storage runtime attach/detach (#15892)
This commit is contained in:
@@ -262,6 +262,7 @@ class HiCacheController:
|
||||
pp_rank: int = 0,
|
||||
pp_size: int = 1,
|
||||
):
|
||||
self.tp_group = tp_group
|
||||
self.mem_pool_device_allocator = token_to_kv_pool_allocator
|
||||
self.mem_pool_device = token_to_kv_pool_allocator.get_kvcache()
|
||||
self.mem_pool_host = mem_pool_host
|
||||
@@ -269,69 +270,20 @@ class HiCacheController:
|
||||
self.page_size = page_size
|
||||
self.io_backend = io_backend
|
||||
self.enable_storage = False
|
||||
self.storage_backend = None
|
||||
self.storage_backend_type = None
|
||||
self.pp_rank = pp_rank
|
||||
self.pp_size = pp_size
|
||||
|
||||
if storage_backend is not None:
|
||||
self.storage_backend_type = storage_backend
|
||||
from sglang.srt.mem_cache.hicache_storage import get_hash_str
|
||||
# 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
|
||||
|
||||
self.get_hash_str = get_hash_str
|
||||
self.storage_config = self._generate_storage_config(
|
||||
model_name, storage_backend_extra_config
|
||||
)
|
||||
# for MLA models, only one rank needs to backup the KV cache
|
||||
self.backup_skip = (
|
||||
self.storage_config.is_mla_model
|
||||
# todo: load balancing
|
||||
and self.storage_config.tp_rank != 0
|
||||
)
|
||||
|
||||
# Use storage backend factory for dynamic backend creation
|
||||
from sglang.srt.mem_cache.storage import StorageBackendFactory
|
||||
|
||||
try:
|
||||
self.storage_backend = StorageBackendFactory.create_backend(
|
||||
storage_backend, self.storage_config, self.mem_pool_host
|
||||
)
|
||||
except ValueError as e:
|
||||
raise ValueError(f"Failed to create storage backend: {e}") from e
|
||||
|
||||
self.storage_backend.register_mem_pool_host(self.mem_pool_host)
|
||||
|
||||
self.enable_storage = True
|
||||
# todo: threshold policy for prefetching
|
||||
self.prefetch_threshold = max(prefetch_threshold, self.page_size)
|
||||
self.prefetch_capacity_limit = int(
|
||||
0.8 * (self.mem_pool_host.size - self.mem_pool_device.size)
|
||||
)
|
||||
# granularity of batch storage IO operations, in number of pages
|
||||
self.storage_batch_size = 128
|
||||
# tracking the number of tokens locked in prefetching, updated by the main scheduler thread
|
||||
self.prefetch_tokens_occupied = 0
|
||||
|
||||
# create a new communication group for synchronizing storage operations across TP workers
|
||||
self.tp_world_size = torch.distributed.get_world_size(group=tp_group)
|
||||
if self.tp_world_size > 1:
|
||||
from sglang.srt.distributed.parallel_state import (
|
||||
create_custom_parallel_group,
|
||||
)
|
||||
|
||||
group_ranks = torch.distributed.get_process_group_ranks(tp_group)
|
||||
self.prefetch_tp_group = create_custom_parallel_group(
|
||||
group_ranks=group_ranks, backend="gloo"
|
||||
)
|
||||
|
||||
# Select the get and set functions
|
||||
self.page_get_func = self._generic_page_get
|
||||
self.page_set_func = self._generic_page_set
|
||||
|
||||
if (self.storage_backend_type in ["hf3fs", "mooncake", "eic"]) or (
|
||||
self.storage_backend_type == "dynamic"
|
||||
and bool(self.storage_config.extra_config.get("interface_v1", 0))
|
||||
):
|
||||
self.page_get_func = self._page_get_zero_copy
|
||||
self.page_set_func = self._page_set_zero_copy
|
||||
# Dedicated stop event for storage background threads (prefetch/backup).
|
||||
# NOTE: Do NOT reuse `self.stop_event` here since it also guards core HiCache
|
||||
# transfer buffers (CPU<->GPU). We want to allow runtime attach/detach of
|
||||
# storage without stopping the whole controller.
|
||||
self.storage_stop_event = threading.Event()
|
||||
|
||||
self.device = self.mem_pool_device.device
|
||||
self.layer_num = self.mem_pool_device.layer_num
|
||||
@@ -360,22 +312,259 @@ class HiCacheController:
|
||||
self.write_stream = device_module.Stream()
|
||||
self.load_stream = device_module.Stream()
|
||||
|
||||
# If a storage backend is provided at startup, treat it as an implicit attach,
|
||||
# so init/runtime share the same lifecycle semantics and code paths.
|
||||
if storage_backend is not None:
|
||||
try:
|
||||
self.attach_storage_backend(
|
||||
storage_backend=storage_backend,
|
||||
prefetch_threshold=prefetch_threshold,
|
||||
model_name=model_name,
|
||||
storage_backend_extra_config=storage_backend_extra_config,
|
||||
)
|
||||
except ValueError as e:
|
||||
# Preserve the historical error shape on init for unknown backends.
|
||||
raise ValueError(f"Failed to create storage backend: {e}") from e
|
||||
|
||||
def _start_storage_threads(self):
|
||||
"""Start storage prefetch/backup threads and their queues.
|
||||
|
||||
This is used by runtime attach, and also by reset when storage is enabled.
|
||||
"""
|
||||
assert self.enable_storage
|
||||
assert not self.storage_stop_event.is_set()
|
||||
|
||||
self.prefetch_thread = threading.Thread(
|
||||
target=self.prefetch_thread_func, daemon=True
|
||||
)
|
||||
self.backup_thread = threading.Thread(
|
||||
target=self.backup_thread_func, daemon=True
|
||||
)
|
||||
self.prefetch_queue = Queue()
|
||||
self.backup_queue = Queue()
|
||||
|
||||
self.prefetch_revoke_queue = Queue()
|
||||
self.ack_backup_queue = Queue()
|
||||
self.host_mem_release_queue = Queue()
|
||||
|
||||
self.prefetch_thread.start()
|
||||
self.backup_thread.start()
|
||||
|
||||
def _stop_storage_threads(self):
|
||||
"""Stop storage prefetch/backup threads and drain internal queues.
|
||||
|
||||
Caller should ensure no in-flight requests.
|
||||
"""
|
||||
# Always request stop. This is safe even when storage is already disabled,
|
||||
# and makes detach truly idempotent (previous partial detach may have left
|
||||
# threads alive).
|
||||
# NOTE: do NOT clear stop_event unless threads have fully stopped; otherwise
|
||||
# a still-alive thread may resume and touch released state.
|
||||
self.storage_stop_event.set()
|
||||
|
||||
# Best-effort wakeups so threads exit promptly even if blocked on queues.
|
||||
try:
|
||||
if hasattr(self, "prefetch_queue"):
|
||||
self.prefetch_queue.put_nowait(None)
|
||||
if hasattr(self, "backup_queue"):
|
||||
self.backup_queue.put_nowait(None)
|
||||
if hasattr(self, "prefetch_buffer"):
|
||||
self.prefetch_buffer.put_nowait(None)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# Best-effort joins (threads are daemon, but join keeps state clean).
|
||||
threads = []
|
||||
if hasattr(self, "prefetch_thread"):
|
||||
threads.append(self.prefetch_thread)
|
||||
if hasattr(self, "backup_thread"):
|
||||
threads.append(self.backup_thread)
|
||||
if hasattr(self, "prefetch_io_aux_thread"):
|
||||
threads.append(self.prefetch_io_aux_thread)
|
||||
|
||||
for t in threads:
|
||||
try:
|
||||
t.join(timeout=10)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
alive = [t for t in threads if getattr(t, "is_alive", lambda: False)()]
|
||||
if alive:
|
||||
logger.error(
|
||||
"Failed to stop HiCache storage threads cleanly: %s",
|
||||
[getattr(t, "name", repr(t)) for t in alive],
|
||||
)
|
||||
raise RuntimeError("Failed to stop HiCache storage threads cleanly.")
|
||||
|
||||
def attach_storage_backend(
|
||||
self,
|
||||
storage_backend: str,
|
||||
prefetch_threshold: int = 256,
|
||||
model_name: Optional[str] = None,
|
||||
storage_backend_extra_config: Optional[dict] = None,
|
||||
):
|
||||
"""Attach (enable) storage backend at runtime.
|
||||
|
||||
Requirement: no in-flight requests. This call is expected to run on the scheduler
|
||||
thread (control path), not concurrently with prefetch/backup.
|
||||
"""
|
||||
if self.enable_storage:
|
||||
self.prefetch_thread = threading.Thread(
|
||||
target=self.prefetch_thread_func, daemon=True
|
||||
)
|
||||
self.backup_thread = threading.Thread(
|
||||
target=self.backup_thread_func, daemon=True
|
||||
)
|
||||
self.prefetch_queue = Queue()
|
||||
self.backup_queue = Queue()
|
||||
raise RuntimeError("Storage backend already attached.")
|
||||
|
||||
self.prefetch_revoke_queue = Queue()
|
||||
self.ack_backup_queue = Queue()
|
||||
self.host_mem_release_queue = Queue()
|
||||
# Defensive: a previous partial detach may have flipped `enable_storage` but
|
||||
# left background threads alive. Attaching on top of them is unsafe.
|
||||
try:
|
||||
self._stop_storage_threads()
|
||||
except Exception as e:
|
||||
raise RuntimeError(
|
||||
"Cannot attach storage backend: previous detach did not stop storage threads cleanly."
|
||||
) from e
|
||||
|
||||
self.prefetch_thread.start()
|
||||
self.backup_thread.start()
|
||||
# Rollback-safe init: if creation fails, keep controller state consistent
|
||||
# for future attach attempts.
|
||||
self.storage_backend_type = storage_backend
|
||||
from sglang.srt.mem_cache.hicache_storage import get_hash_str
|
||||
|
||||
self.get_hash_str = get_hash_str
|
||||
self.storage_config = self._generate_storage_config(
|
||||
model_name, storage_backend_extra_config
|
||||
)
|
||||
# for MLA models, only one rank needs to backup the KV cache
|
||||
self.backup_skip = (
|
||||
self.storage_config.is_mla_model
|
||||
# todo: load balancing
|
||||
and self.storage_config.tp_rank != 0
|
||||
)
|
||||
|
||||
# Use storage backend factory for dynamic backend creation
|
||||
from sglang.srt.mem_cache.storage import StorageBackendFactory
|
||||
|
||||
try:
|
||||
self.storage_backend = StorageBackendFactory.create_backend(
|
||||
storage_backend, self.storage_config, self.mem_pool_host
|
||||
)
|
||||
self.storage_backend.register_mem_pool_host(self.mem_pool_host)
|
||||
|
||||
self.enable_storage = True
|
||||
# todo: threshold policy for prefetching
|
||||
self.prefetch_threshold = max(prefetch_threshold, self.page_size)
|
||||
self.prefetch_capacity_limit = max(
|
||||
0, int(0.8 * (self.mem_pool_host.size - self.mem_pool_device.size))
|
||||
)
|
||||
# granularity of batch storage IO operations, in number of pages
|
||||
self.storage_batch_size = 128
|
||||
# tracking the number of tokens locked in prefetching, updated by the main scheduler thread
|
||||
self.prefetch_tokens_occupied = 0
|
||||
|
||||
# create a new communication group for synchronizing storage operations across TP workers
|
||||
self.tp_world_size = torch.distributed.get_world_size(group=self.tp_group)
|
||||
if self.tp_world_size > 1:
|
||||
from sglang.srt.distributed.parallel_state import (
|
||||
create_custom_parallel_group,
|
||||
)
|
||||
|
||||
group_ranks = torch.distributed.get_process_group_ranks(self.tp_group)
|
||||
self.prefetch_tp_group = create_custom_parallel_group(
|
||||
group_ranks=group_ranks, backend="gloo"
|
||||
)
|
||||
|
||||
# Select the get and set functions
|
||||
self.page_get_func = self._generic_page_get
|
||||
self.page_set_func = self._generic_page_set
|
||||
|
||||
if (self.storage_backend_type in ["hf3fs", "mooncake", "eic"]) or (
|
||||
self.storage_backend_type == "dynamic"
|
||||
and bool(self.storage_config.extra_config.get("interface_v1", 0))
|
||||
):
|
||||
self.page_get_func = self._page_get_zero_copy
|
||||
self.page_set_func = self._page_set_zero_copy
|
||||
|
||||
# Ensure stop_event is clear before starting threads.
|
||||
self.storage_stop_event.clear()
|
||||
self._start_storage_threads()
|
||||
except Exception:
|
||||
# Best-effort cleanup for partial init.
|
||||
try:
|
||||
self._stop_storage_threads()
|
||||
except Exception:
|
||||
pass
|
||||
try:
|
||||
if hasattr(self, "prefetch_tp_group"):
|
||||
try:
|
||||
torch.distributed.destroy_process_group(self.prefetch_tp_group)
|
||||
except Exception:
|
||||
pass
|
||||
self.prefetch_tp_group = None
|
||||
except Exception:
|
||||
pass
|
||||
try:
|
||||
if (
|
||||
hasattr(self, "storage_backend")
|
||||
and self.storage_backend is not None
|
||||
):
|
||||
if hasattr(self.storage_backend, "close"):
|
||||
self.storage_backend.close()
|
||||
except Exception:
|
||||
pass
|
||||
self.storage_backend = None
|
||||
self.storage_backend_type = None
|
||||
self.enable_storage = False
|
||||
self.page_get_func = self._generic_page_get
|
||||
self.page_set_func = self._generic_page_set
|
||||
raise
|
||||
|
||||
def detach_storage_backend(self):
|
||||
"""Detach (disable) storage backend at runtime.
|
||||
|
||||
Requirement: no in-flight requests. This will stop storage threads and release
|
||||
the backend instance (best-effort close).
|
||||
"""
|
||||
# Idempotent cleanup: even if `enable_storage` is already False,
|
||||
# we may still have leftover resources (threads/backend/process group) from a
|
||||
# previous partial detach. We attempt cleanup whenever possible.
|
||||
try:
|
||||
self._stop_storage_threads()
|
||||
except Exception as e:
|
||||
# Do not proceed tearing down backend/process group if threads are not
|
||||
# fully stopped; otherwise still-alive threads may touch released state.
|
||||
# Caller can retry detach.
|
||||
logger.exception("Stop storage threads failed: %s", e)
|
||||
# IMPORTANT: Do not silently succeed. Upper layers rely on exceptions here
|
||||
# to avoid flipping `enable_storage` flags while threads are still alive.
|
||||
raise RuntimeError("Stop storage threads failed; detach aborted.") from e
|
||||
|
||||
# Best-effort destroy process group created for storage ops.
|
||||
try:
|
||||
if (
|
||||
hasattr(self, "prefetch_tp_group")
|
||||
and self.prefetch_tp_group is not None
|
||||
):
|
||||
try:
|
||||
torch.distributed.destroy_process_group(self.prefetch_tp_group)
|
||||
except Exception:
|
||||
pass
|
||||
self.prefetch_tp_group = None
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# Best-effort close (some backends rely on GC/destructor).
|
||||
try:
|
||||
if (
|
||||
hasattr(self, "storage_backend")
|
||||
and self.storage_backend is not None
|
||||
and hasattr(self.storage_backend, "close")
|
||||
):
|
||||
self.storage_backend.close()
|
||||
except Exception:
|
||||
logger.exception("Failed to close storage backend cleanly.")
|
||||
|
||||
self.storage_backend = None
|
||||
self.storage_backend_type = None
|
||||
self.enable_storage = False
|
||||
self.page_get_func = self._generic_page_get
|
||||
self.page_set_func = self._generic_page_set
|
||||
# Now it's safe to clear the stop event for future re-attach.
|
||||
self.storage_stop_event.clear()
|
||||
|
||||
def _generate_storage_config(
|
||||
self,
|
||||
@@ -408,6 +597,7 @@ class HiCacheController:
|
||||
|
||||
def reset(self):
|
||||
self.stop_event.set()
|
||||
self.storage_stop_event.set()
|
||||
|
||||
self.write_queue.clear()
|
||||
self.load_queue.clear()
|
||||
@@ -424,6 +614,7 @@ class HiCacheController:
|
||||
self.ack_backup_queue.queue.clear()
|
||||
|
||||
self.stop_event.clear()
|
||||
self.storage_stop_event.clear()
|
||||
|
||||
if self.enable_storage:
|
||||
self.prefetch_thread = threading.Thread(
|
||||
@@ -661,9 +852,11 @@ class HiCacheController:
|
||||
"""
|
||||
Auxiliary function conducting IO operations for prefetching.
|
||||
"""
|
||||
while not self.stop_event.is_set():
|
||||
while not self.storage_stop_event.is_set():
|
||||
try:
|
||||
operation = self.prefetch_buffer.get(block=True, timeout=1)
|
||||
if operation is None:
|
||||
continue
|
||||
self._page_transfer(operation)
|
||||
# operation terminated by controller, release pre-allocated memory
|
||||
self.append_host_mem_release(
|
||||
@@ -719,9 +912,11 @@ class HiCacheController:
|
||||
Manage prefetching operations from storage backend to host memory.
|
||||
"""
|
||||
self.prefetch_buffer = Queue()
|
||||
aux_thread = threading.Thread(target=self.prefetch_io_aux_func, daemon=True)
|
||||
aux_thread.start()
|
||||
while (not self.stop_event.is_set()) or not self.prefetch_queue.empty():
|
||||
self.prefetch_io_aux_thread = threading.Thread(
|
||||
target=self.prefetch_io_aux_func, daemon=True
|
||||
)
|
||||
self.prefetch_io_aux_thread.start()
|
||||
while (not self.storage_stop_event.is_set()) or not self.prefetch_queue.empty():
|
||||
try:
|
||||
operation = self.prefetch_queue.get(block=True, timeout=1)
|
||||
if operation is None:
|
||||
@@ -818,7 +1013,7 @@ class HiCacheController:
|
||||
"""
|
||||
Manage backup operations from host memory to storage backend.
|
||||
"""
|
||||
while not self.stop_event.is_set():
|
||||
while not self.storage_stop_event.is_set():
|
||||
try:
|
||||
operation = self.backup_queue.get(block=True, timeout=1)
|
||||
if operation is None:
|
||||
|
||||
Reference in New Issue
Block a user