[HiCache][HA 1/N] Support HiCache storage runtime attach/detach (#15892)

This commit is contained in:
shuwenn
2026-01-26 19:33:19 -08:00
committed by GitHub
parent 1b56a886bb
commit fd3b179ffd
10 changed files with 1488 additions and 124 deletions
+272 -77
View File
@@ -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: