[HiCache & HybridModel] 3FS backend support DSA & mamba model (#23241)

Co-authored-by: 墨已 <kangyifei.kyf@alibaba-inc.com>
Co-authored-by: hzh0425 <hzh0425@apache.org>
This commit is contained in:
Kang Yifei
2026-04-25 00:48:01 +08:00
committed by GitHub
co-authored by 墨已 hzh0425
parent 92d262f710
commit baa0aa670f
3 changed files with 399 additions and 55 deletions
@@ -142,6 +142,7 @@ class HiMambaRadixCache(MambaRadixCache):
extra_config=extra_config, extra_config=extra_config,
prefetch_threshold=prefetch_threshold, prefetch_threshold=prefetch_threshold,
load_cache_event=self.load_cache_event, load_cache_event=self.load_cache_event,
enable_storage_metrics=self.enable_storage_metrics,
) )
self._apply_storage_runtime_config( self._apply_storage_runtime_config(
storage_backend=server_args.hicache_storage_backend, storage_backend=server_args.hicache_storage_backend,
@@ -14,6 +14,7 @@ from fastapi.responses import ORJSONResponse
from requests.adapters import HTTPAdapter from requests.adapters import HTTPAdapter
from urllib3.util.retry import Retry from urllib3.util.retry import Retry
from sglang.srt.mem_cache.hicache_storage import PoolName
from sglang.srt.mem_cache.storage.hf3fs.storage_hf3fs import Hf3fsMetadataInterface from sglang.srt.mem_cache.storage.hf3fs.storage_hf3fs import Hf3fsMetadataInterface
# --- Configuration --- # --- Configuration ---
@@ -115,7 +116,7 @@ class GlobalMetadataState:
def __init__(self, persistence_path: Optional[str], save_interval: int): def __init__(self, persistence_path: Optional[str], save_interval: int):
self.global_lock = threading.RLock() self.global_lock = threading.RLock()
self.ranks: Dict[int, RankMetadata] = {} self.ranks: Dict[str, RankMetadata] = {}
self.persistence_path = Path(persistence_path) if persistence_path else None self.persistence_path = Path(persistence_path) if persistence_path else None
self.save_interval = save_interval self.save_interval = save_interval
self.save_timer: Optional[threading.Timer] = None self.save_timer: Optional[threading.Timer] = None
@@ -132,13 +133,14 @@ class GlobalMetadataState:
persisted_data = json.load(f) persisted_data = json.load(f)
with self.global_lock: with self.global_lock:
for rank_id_str, data in persisted_data.items(): for key_str, data in persisted_data.items():
rank_id = int(rank_id_str) if ":" not in key_str:
key_str = f"{key_str}:kv" # For backward compatibility
num_pages = data["num_pages"] num_pages = data["num_pages"]
rank_meta = RankMetadata(num_pages) rank_meta = RankMetadata(num_pages)
rank_meta.free_pages = data["free_pages"] rank_meta.free_pages = data["free_pages"]
rank_meta.key_to_index = OrderedDict(data["key_to_index"]) rank_meta.key_to_index = OrderedDict(data["key_to_index"])
self.ranks[rank_id] = rank_meta self.ranks[key_str] = rank_meta
logging.info( logging.info(
f"Successfully loaded metadata for {len(self.ranks)} ranks." f"Successfully loaded metadata for {len(self.ranks)} ranks."
) )
@@ -156,9 +158,9 @@ class GlobalMetadataState:
logging.info("Persisting metadata to disk...") logging.info("Persisting metadata to disk...")
with self.global_lock: with self.global_lock:
serializable_state = {} serializable_state = {}
for rank_id, rank_meta in self.ranks.items(): for key_str, rank_meta in self.ranks.items():
with rank_meta.lock: with rank_meta.lock:
serializable_state[rank_id] = { serializable_state[key_str] = {
"num_pages": rank_meta.num_pages, "num_pages": rank_meta.num_pages,
"free_pages": rank_meta.free_pages, "free_pages": rank_meta.free_pages,
"key_to_index": list(rank_meta.key_to_index.items()), "key_to_index": list(rank_meta.key_to_index.items()),
@@ -211,14 +213,19 @@ class Hf3fsMetadataServer:
self.app.post("/{rank}/clear")(self.clear) self.app.post("/{rank}/clear")(self.clear)
self.app.post("/{rank}/get_page_indices")(self.get_page_indices) self.app.post("/{rank}/get_page_indices")(self.get_page_indices)
def get_rank_metadata(self, rank: int) -> RankMetadata: def _rank_key(self, rank: int, namespace: str) -> str:
"""Generate the composite key for rank+namespace."""
return f"{rank}:{namespace}"
def get_rank_metadata(self, rank: int, namespace: str = "kv") -> RankMetadata:
"""Get rank metadata with proper error handling.""" """Get rank metadata with proper error handling."""
if rank not in self.state.ranks: key = self._rank_key(rank, namespace)
if key not in self.state.ranks:
raise HTTPException( raise HTTPException(
status_code=404, status_code=404,
detail=f"Rank {rank} not initialized. Please call /{rank}/initialize first.", detail=f"Rank {rank} namespace '{namespace}' not initialized. Please call /{rank}/initialize first.",
) )
return self.state.ranks[rank] return self.state.ranks[key]
async def _read_json(self, request: Request) -> dict: async def _read_json(self, request: Request) -> dict:
"""Parse request JSON using orjson if available.""" """Parse request JSON using orjson if available."""
@@ -233,32 +240,38 @@ class Hf3fsMetadataServer:
"""Initialize a rank with specified number of pages.""" """Initialize a rank with specified number of pages."""
data = await self._read_json(request) data = await self._read_json(request)
num_pages = data["num_pages"] num_pages = data["num_pages"]
namespace = data.get("namespace", "kv")
key = self._rank_key(rank, namespace)
with self.state.global_lock: with self.state.global_lock:
if rank in self.state.ranks: if key in self.state.ranks:
logging.info( logging.info(
f"Rank {rank} already exists. Initialization request ignored." f"Rank {rank} namespace '{namespace}' already exists. Initialization request ignored."
) )
if self.state.ranks[rank].num_pages != num_pages: if self.state.ranks[key].num_pages != num_pages:
logging.warning( logging.warning(
f"Rank {rank} initialized with different num_pages. Existing: {self.state.ranks[rank].num_pages}, New: {num_pages}" f"Rank {rank} namespace '{namespace}' initialized with different num_pages. Existing: {self.state.ranks[key].num_pages}, New: {num_pages}"
) )
else: else:
logging.info(f"Initializing new Rank {rank} with {num_pages} pages.") logging.info(
self.state.ranks[rank] = RankMetadata(num_pages) f"Initializing new Rank {rank} namespace '{namespace}' with {num_pages} pages."
)
self.state.ranks[key] = RankMetadata(num_pages)
return Response(status_code=204) return Response(status_code=204)
async def exists(self, rank: int, request: Request): async def exists(self, rank: int, request: Request):
"""Check if keys exist in metadata.""" """Check if keys exist in metadata."""
data = await self._read_json(request) data = await self._read_json(request)
keys = data["keys"] keys = data["keys"]
metadata = self.get_rank_metadata(rank) namespace = data.get("namespace", "kv")
metadata = self.get_rank_metadata(rank, namespace)
results = metadata.exists_keys(keys) results = metadata.exists_keys(keys)
return self._json_response({"exists": results}) return self._json_response({"exists": results})
async def reserve_and_allocate_page_indices(self, rank: int, request: Request): async def reserve_and_allocate_page_indices(self, rank: int, request: Request):
"""Reserve and allocate page indices for keys.""" """Reserve and allocate page indices for keys."""
data = await self._read_json(request) data = await self._read_json(request)
metadata = self.get_rank_metadata(rank) namespace = data.get("namespace", "kv")
metadata = self.get_rank_metadata(rank, namespace)
keys = data["keys"] keys = data["keys"]
results = metadata.reserve_and_allocate_page_indices(keys) results = metadata.reserve_and_allocate_page_indices(keys)
return self._json_response({"indices": results}) return self._json_response({"indices": results})
@@ -266,7 +279,8 @@ class Hf3fsMetadataServer:
async def confirm_write(self, rank: int, request: Request): async def confirm_write(self, rank: int, request: Request):
"""Confirm write operations and release pages.""" """Confirm write operations and release pages."""
data = await self._read_json(request) data = await self._read_json(request)
metadata = self.get_rank_metadata(rank) namespace = data.get("namespace", "kv")
metadata = self.get_rank_metadata(rank, namespace)
success_written_keys = data.get("written_keys_to_confirm", []) success_written_keys = data.get("written_keys_to_confirm", [])
released_pages = data.get("pages_to_release", []) released_pages = data.get("pages_to_release", [])
@@ -277,20 +291,24 @@ class Hf3fsMetadataServer:
async def delete_keys(self, rank: int, request: Request): async def delete_keys(self, rank: int, request: Request):
"""Delete keys from metadata.""" """Delete keys from metadata."""
data = await self._read_json(request) data = await self._read_json(request)
metadata = self.get_rank_metadata(rank) namespace = data.get("namespace", "kv")
metadata = self.get_rank_metadata(rank, namespace)
count = metadata.delete_keys(data["keys"]) count = metadata.delete_keys(data["keys"])
return Response(status_code=204) return Response(status_code=204)
async def clear(self, rank: int): async def clear(self, rank: int, request: Request):
"""Clear all metadata for a rank.""" """Clear all metadata for a rank."""
metadata = self.get_rank_metadata(rank) data = await self._read_json(request)
namespace = data.get("namespace", "kv")
metadata = self.get_rank_metadata(rank, namespace)
metadata.clear_all() metadata.clear_all()
return Response(status_code=204) return Response(status_code=204)
async def get_page_indices(self, rank: int, request: Request): async def get_page_indices(self, rank: int, request: Request):
"""Get page indices for keys.""" """Get page indices for keys."""
data = await self._read_json(request) data = await self._read_json(request)
metadata = self.get_rank_metadata(rank) namespace = data.get("namespace", "kv")
metadata = self.get_rank_metadata(rank, namespace)
keys = data["keys"] keys = data["keys"]
results = metadata.get_page_indices(keys) results = metadata.get_page_indices(keys)
return self._json_response({"indices": results}) return self._json_response({"indices": results})
@@ -349,14 +367,19 @@ class Hf3fsGlobalMetadataClient(Hf3fsMetadataInterface):
logging.error(f"Failed to POST to {endpoint} after retries: {e}") logging.error(f"Failed to POST to {endpoint} after retries: {e}")
raise RuntimeError(f"Failed to connect to metadata server: {e}") from e raise RuntimeError(f"Failed to connect to metadata server: {e}") from e
def initialize(self, rank: int, num_pages: int) -> None: def initialize(
self._post(f"{rank}/initialize", {"num_pages": num_pages}) self, rank: int, num_pages: int, namespace: PoolName = PoolName.KV
) -> None:
self._post(
f"{rank}/initialize", {"num_pages": num_pages, "namespace": str(namespace)}
)
def reserve_and_allocate_page_indices( def reserve_and_allocate_page_indices(
self, rank: int, keys: List[Tuple[str, str]] self, rank: int, keys: List[Tuple[str, str]], namespace: PoolName = PoolName.KV
) -> List[Tuple[bool, int]]: ) -> List[Tuple[bool, int]]:
response = self._post( response = self._post(
f"{rank}/reserve_and_allocate_page_indices", {"keys": keys} f"{rank}/reserve_and_allocate_page_indices",
{"keys": keys, "namespace": str(namespace)},
) )
return [tuple(item) for item in response.get("indices")] return [tuple(item) for item in response.get("indices")]
@@ -365,69 +388,107 @@ class Hf3fsGlobalMetadataClient(Hf3fsMetadataInterface):
rank: int, rank: int,
written_keys_to_confirm: List[Tuple[str, int]], written_keys_to_confirm: List[Tuple[str, int]],
pages_to_release: List[int], pages_to_release: List[int],
namespace: PoolName = PoolName.KV,
) -> None: ) -> None:
self._post( self._post(
f"{rank}/confirm_write", f"{rank}/confirm_write",
{ {
"written_keys_to_confirm": written_keys_to_confirm, "written_keys_to_confirm": written_keys_to_confirm,
"pages_to_release": pages_to_release, "pages_to_release": pages_to_release,
"namespace": str(namespace),
}, },
) )
def delete_keys(self, rank: int, keys: List[str]) -> None: def delete_keys(
self._post(f"{rank}/delete_keys", {"keys": keys}) self, rank: int, keys: List[str], namespace: PoolName = PoolName.KV
) -> None:
self._post(f"{rank}/delete_keys", {"keys": keys, "namespace": str(namespace)})
def exists(self, rank: int, keys: List[str]) -> List[bool]: def exists(
response = self._post(f"{rank}/exists", {"keys": keys}) self, rank: int, keys: List[str], namespace: PoolName = PoolName.KV
) -> List[bool]:
response = self._post(
f"{rank}/exists", {"keys": keys, "namespace": str(namespace)}
)
return response.get("exists", []) return response.get("exists", [])
def clear(self, rank: int) -> None: def clear(self, rank: int, namespace: PoolName = PoolName.KV) -> None:
self._post(f"{rank}/clear", {}) self._post(f"{rank}/clear", {"namespace": str(namespace)})
def get_page_indices(self, rank: int, keys: List[str]) -> List[Optional[int]]: def get_page_indices(
response = self._post(f"{rank}/get_page_indices", {"keys": keys}) self, rank: int, keys: List[str], namespace: PoolName = PoolName.KV
) -> List[Optional[int]]:
response = self._post(
f"{rank}/get_page_indices", {"keys": keys, "namespace": str(namespace)}
)
return response.get("indices") return response.get("indices")
class Hf3fsLocalMetadataClient(Hf3fsMetadataInterface): class Hf3fsLocalMetadataClient(Hf3fsMetadataInterface):
"""Local metadata client that directly operates on single RankMetadata in memory without metadata server.""" """Local metadata client that directly operates on RankMetadata in memory without metadata server."""
def __init__(self): def __init__(self):
self.rank_metadata = None self._metadata: Dict[str, RankMetadata] = {} # key: "rank:namespace"
def initialize(self, rank: int, num_pages: int) -> None: def _ns_key(self, rank: int, namespace: PoolName) -> str:
self.rank_metadata = RankMetadata(num_pages) return f"{rank}:{namespace}"
def _get_metadata(self, rank: int, namespace) -> RankMetadata:
key = self._ns_key(rank, namespace)
if key not in self._metadata:
raise RuntimeError(
f"Namespace '{namespace}' for rank {rank} not initialized"
)
return self._metadata[key]
def initialize(
self, rank: int, num_pages: int, namespace: PoolName = PoolName.KV
) -> None:
key = self._ns_key(rank, namespace)
if key not in self._metadata:
self._metadata[key] = RankMetadata(num_pages)
def reserve_and_allocate_page_indices( def reserve_and_allocate_page_indices(
self, rank: int, keys: List[Tuple[str, str]] self, rank: int, keys: List[Tuple[str, str]], namespace: PoolName = PoolName.KV
) -> List[Tuple[bool, int]]: ) -> List[Tuple[bool, int]]:
"""Reserve and allocate page indices for keys.""" """Reserve and allocate page indices for keys."""
return self.rank_metadata.reserve_and_allocate_page_indices(keys) return self._get_metadata(rank, namespace).reserve_and_allocate_page_indices(
keys
)
def confirm_write( def confirm_write(
self, self,
rank: int, rank: int,
written_keys_to_confirm: List[Tuple[str, int]], written_keys_to_confirm: List[Tuple[str, int]],
pages_to_release: List[int], pages_to_release: List[int],
namespace: PoolName = PoolName.KV,
) -> None: ) -> None:
"""Confirm write operations.""" """Confirm write operations."""
self.rank_metadata.confirm_write(written_keys_to_confirm, pages_to_release) self._get_metadata(rank, namespace).confirm_write(
written_keys_to_confirm, pages_to_release
)
def delete_keys(self, rank: int, keys: List[str]) -> None: def delete_keys(
self, rank: int, keys: List[str], namespace: PoolName = PoolName.KV
) -> None:
"""Delete keys.""" """Delete keys."""
self.rank_metadata.delete_keys(keys) self._get_metadata(rank, namespace).delete_keys(keys)
def exists(self, rank: int, keys: List[str]) -> List[bool]: def exists(
self, rank: int, keys: List[str], namespace: PoolName = PoolName.KV
) -> List[bool]:
"""Check if keys exist.""" """Check if keys exist."""
return self.rank_metadata.exists_keys(keys) return self._get_metadata(rank, namespace).exists_keys(keys)
def clear(self, rank: int) -> None: def clear(self, rank: int, namespace: PoolName = PoolName.KV) -> None:
"""Clear all metadata for rank.""" """Clear all metadata for rank."""
self.rank_metadata.clear_all() self._get_metadata(rank, namespace).clear_all()
def get_page_indices(self, rank: int, keys: List[str]) -> List[Optional[int]]: def get_page_indices(
self, rank: int, keys: List[str], namespace: PoolName = PoolName.KV
) -> List[Optional[int]]:
"""Get page indices for keys.""" """Get page indices for keys."""
return self.rank_metadata.get_page_indices(keys) return self._get_metadata(rank, namespace).get_page_indices(keys)
def run_metadata_server( def run_metadata_server(
@@ -7,6 +7,7 @@ import signal
import threading import threading
import time import time
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
from dataclasses import dataclass
from functools import wraps from functools import wraps
from typing import Any, List, Optional, Tuple from typing import Any, List, Optional, Tuple
@@ -16,6 +17,10 @@ from sglang.srt.mem_cache.hicache_storage import (
HiCacheStorage, HiCacheStorage,
HiCacheStorageConfig, HiCacheStorageConfig,
HiCacheStorageExtraInfo, HiCacheStorageExtraInfo,
PoolHitPolicy,
PoolName,
PoolTransfer,
PoolTransferResult,
) )
from sglang.srt.mem_cache.memory_pool_host import HostKVCache from sglang.srt.mem_cache.memory_pool_host import HostKVCache
from sglang.srt.mem_cache.storage.hf3fs.hf3fs_client import Hf3fsClient from sglang.srt.mem_cache.storage.hf3fs.hf3fs_client import Hf3fsClient
@@ -28,7 +33,9 @@ class Hf3fsMetadataInterface(ABC):
"""Interface for HF3FS metadata operations.""" """Interface for HF3FS metadata operations."""
@abstractmethod @abstractmethod
def initialize(self, rank: int, num_pages: int) -> None: def initialize(
self, rank: int, num_pages: int, namespace: PoolName = PoolName.KV
) -> None:
"""Initialize the metadata service with specified number of pages.""" """Initialize the metadata service with specified number of pages."""
pass pass
@@ -37,12 +44,14 @@ class Hf3fsMetadataInterface(ABC):
self, self,
rank: int, rank: int,
keys: List[Tuple[str, str]], keys: List[Tuple[str, str]],
namespace: PoolName = PoolName.KV,
) -> List[Tuple[bool, int]]: ) -> List[Tuple[bool, int]]:
""" """
Reserve and allocate page indices for the specified keys. Reserve and allocate page indices for the specified keys.
Args: Args:
rank: The rank of the process. rank: The rank of the process.
keys: The keys to reserve and allocate page indices for. Each tuple contains a key and the key of its prefix block. keys: The keys to reserve and allocate page indices for. Each tuple contains a key and the key of its prefix block.
namespace: The namespace (pool type) for the metadata.
Returns: Returns:
List[Tuple[bool, int]]: A list of tuples, where each tuple contains a boolean indicating whether the key has existed and an integer indicating the allocated page index. List[Tuple[bool, int]]: A list of tuples, where each tuple contains a boolean indicating whether the key has existed and an integer indicating the allocated page index.
""" """
@@ -54,6 +63,7 @@ class Hf3fsMetadataInterface(ABC):
rank: int, rank: int,
written_keys_to_confirm: List[Tuple[str, int]], written_keys_to_confirm: List[Tuple[str, int]],
pages_to_release: List[int], pages_to_release: List[int],
namespace: PoolName = PoolName.KV,
) -> None: ) -> None:
""" """
Confirm that key-value pairs have been successfully written to storage. Confirm that key-value pairs have been successfully written to storage.
@@ -61,16 +71,20 @@ class Hf3fsMetadataInterface(ABC):
rank: The rank of the process. rank: The rank of the process.
written_keys_to_confirm: A list of tuples, where each tuple contains a key and its corresponding page index. written_keys_to_confirm: A list of tuples, where each tuple contains a key and its corresponding page index.
pages_to_release: A list of page indices to be released. pages_to_release: A list of page indices to be released.
namespace: The namespace (pool type) for the metadata.
""" """
pass pass
@abstractmethod @abstractmethod
def get_page_indices(self, rank: int, keys: List[str]) -> List[Optional[int]]: def get_page_indices(
self, rank: int, keys: List[str], namespace: PoolName = PoolName.KV
) -> List[Optional[int]]:
""" """
Get page indices for the specified keys. Get page indices for the specified keys.
Args: Args:
rank: The rank of the process. rank: The rank of the process.
keys: A list of keys. keys: A list of keys.
namespace: The namespace (pool type) for the metadata.
Returns: Returns:
List[Optional[int]]: A list of integers representing the page indices for the specified keys. List[Optional[int]]: A list of integers representing the page indices for the specified keys.
If a key is not found, the corresponding index will be None. If a key is not found, the corresponding index will be None.
@@ -78,17 +92,21 @@ class Hf3fsMetadataInterface(ABC):
pass pass
@abstractmethod @abstractmethod
def delete_keys(self, rank: int, keys: List[str]) -> None: def delete_keys(
self, rank: int, keys: List[str], namespace: PoolName = PoolName.KV
) -> None:
"""Delete specified keys and their associated pages.""" """Delete specified keys and their associated pages."""
pass pass
@abstractmethod @abstractmethod
def exists(self, rank: int, keys: List[str]) -> List[bool]: def exists(
self, rank: int, keys: List[str], namespace: PoolName = PoolName.KV
) -> List[bool]:
"""Check if the specified keys exist.""" """Check if the specified keys exist."""
pass pass
@abstractmethod @abstractmethod
def clear(self, rank: int) -> None: def clear(self, rank: int, namespace: PoolName = PoolName.KV) -> None:
"""Clear all key-value pairs and page allocations for the specified rank.""" """Clear all key-value pairs and page allocations for the specified rank."""
pass pass
@@ -151,6 +169,18 @@ def create_hf3fs_client(
return Hf3fsUsrBioClient(path, size, bytes_per_page, entries, client_timeout) return Hf3fsUsrBioClient(path, size, bytes_per_page, entries, client_timeout)
@dataclass
class _PoolStorageCtx:
"""Per-pool storage context for hybrid KV cache pools."""
pool_name: str
bytes_per_page: int
num_pages: int
namespace: PoolName
clients: List[Hf3fsClient]
gb_per_page: float
class HiCacheHF3FS(HiCacheStorage): class HiCacheHF3FS(HiCacheStorage):
"""HiCache backend that stores KV cache pages in HF3FS files.""" """HiCache backend that stores KV cache pages in HF3FS files."""
@@ -185,6 +215,7 @@ class HiCacheHF3FS(HiCacheStorage):
self.is_mla_model = is_mla_model self.is_mla_model = is_mla_model
self.is_page_first_layout = is_page_first_layout self.is_page_first_layout = is_page_first_layout
self.enable_storage_metrics = enable_storage_metrics self.enable_storage_metrics = enable_storage_metrics
self.use_mock_client = use_mock_client
self.numel = self.bytes_per_page // self.dtype.itemsize self.numel = self.bytes_per_page // self.dtype.itemsize
self.num_pages = self.file_size // self.bytes_per_page self.num_pages = self.file_size // self.bytes_per_page
self.skip_backup = False self.skip_backup = False
@@ -220,6 +251,7 @@ class HiCacheHF3FS(HiCacheStorage):
self.metadata_client.initialize(self.rank, self.num_pages) self.metadata_client.initialize(self.rank, self.num_pages)
self.lock = threading.RLock() self.lock = threading.RLock()
self._pool_storage_ctx: dict = {}
atexit.register(self.close) atexit.register(self.close)
@@ -503,6 +535,8 @@ class HiCacheHF3FS(HiCacheStorage):
def clear(self) -> None: def clear(self) -> None:
try: try:
self.metadata_client.clear(self.rank) self.metadata_client.clear(self.rank)
for ctx in getattr(self, "_pool_storage_ctx", {}).values():
self.metadata_client.clear(self.rank, namespace=ctx.namespace)
logger.info(f"Cleared HiCacheHF3FS for rank {self.rank}") logger.info(f"Cleared HiCacheHF3FS for rank {self.rank}")
except Exception as e: except Exception as e:
logger.error(f"Failed to clear HiCacheHF3FS: {e}") logger.error(f"Failed to clear HiCacheHF3FS: {e}")
@@ -511,6 +545,9 @@ class HiCacheHF3FS(HiCacheStorage):
try: try:
for c in self.clients: for c in self.clients:
c.close() c.close()
for ctx in getattr(self, "_pool_storage_ctx", {}).values():
for c in ctx.clients:
c.close()
self.executor.shutdown(wait=True) self.executor.shutdown(wait=True)
except Exception as e: except Exception as e:
logger.error(f"close HiCacheHF3FS: {e}") logger.error(f"close HiCacheHF3FS: {e}")
@@ -537,6 +574,45 @@ class HiCacheHF3FS(HiCacheStorage):
logger.info(f"{self.is_zero_copy=}, layout={self.mem_pool_host.layout}") logger.info(f"{self.is_zero_copy=}, layout={self.mem_pool_host.layout}")
def register_mem_host_pool_v2(self, host_pool: HostKVCache, host_pool_name):
if host_pool_name == PoolName.KV:
return
super().register_mem_host_pool_v2(host_pool, host_pool_name)
pool_page_size = getattr(host_pool, "page_size", 1) or 1
pool_bytes_per_page = host_pool.get_ksize_per_token() * pool_page_size
pool_num_pages = self.file_size // pool_bytes_per_page
pool_file_path = f"{self.file_path}.{host_pool_name}"
namespace = host_pool_name # e.g. PoolName.MAMBA, PoolName.INDEXER
pool_clients = [
create_hf3fs_client(
pool_file_path,
self.file_size,
pool_bytes_per_page,
self.entries,
self.client_timeout,
self.use_mock_client,
)
for _ in range(self.numjobs)
]
self.metadata_client.initialize(self.rank, pool_num_pages, namespace=namespace)
self._pool_storage_ctx[host_pool_name] = _PoolStorageCtx(
pool_name=host_pool_name,
bytes_per_page=pool_bytes_per_page,
num_pages=pool_num_pages,
namespace=namespace,
clients=pool_clients,
gb_per_page=pool_bytes_per_page / (1 << 30),
)
logger.info(
f"[Rank {self.rank}] Registered hybrid pool '{host_pool_name}': "
f"bytes_per_page={pool_bytes_per_page}, num_pages={pool_num_pages}, "
f"namespace={namespace}, file={pool_file_path}"
)
def _get_mha_zero_copy_keys(self, keys: List[str]) -> List[str]: def _get_mha_zero_copy_keys(self, keys: List[str]) -> List[str]:
_keys = [] _keys = []
for k in keys: for k in keys:
@@ -596,6 +672,212 @@ class HiCacheHF3FS(HiCacheStorage):
return results return results
def batch_exists_v2(
self,
keys: List[str],
pool_transfers: Optional[List[PoolTransfer]] = None,
extra_info: Optional[HiCacheStorageExtraInfo] = None,
) -> PoolTransferResult:
kv_pages = self.batch_exists(keys, extra_info)
hit_count: dict = {PoolName.KV: kv_pages} if kv_pages else {}
final_pages = kv_pages
for transfer in pool_transfers or []:
if final_pages == 0:
break
pool_name = transfer.name
ctx = self._pool_storage_ctx.get(pool_name)
if ctx is None:
final_pages = 0
break
component_keys = [f"{key}_{pool_name}" for key in keys[:kv_pages]]
exists_results = self.metadata_client.exists(
self.rank, component_keys, namespace=ctx.namespace
)
boundary = 0
if transfer.hit_policy == PoolHitPolicy.ALL_PAGES:
try:
boundary = exists_results.index(False)
except ValueError:
boundary = kv_pages
elif transfer.hit_policy == PoolHitPolicy.TRAILING_PAGES:
trailing = max(1, len(transfer.keys) if transfer.keys else 1)
for prefix_len in range(kv_pages, 0, -1):
if all(
exists_results[i]
for i in range(max(0, prefix_len - trailing), prefix_len)
):
boundary = prefix_len
break
if boundary:
hit_count[pool_name] = boundary
final_pages = min(final_pages, boundary)
return PoolTransferResult(final_pages, hit_count)
def _pool_batch_get(self, transfer: PoolTransfer) -> List[bool]:
pool_name = transfer.name
ctx = self._pool_storage_ctx[pool_name]
host_pool = self.registered_pools[pool_name]
keys = transfer.keys
host_indices = transfer.host_indices
page_size = getattr(host_pool, "page_size", 1) or 1
page_num = len(keys)
component_keys = [f"{key}_{pool_name}" for key in keys]
page_indices = self.metadata_client.get_page_indices(
self.rank, component_keys, namespace=ctx.namespace
)
batch_indices, file_offsets, values = [], [], []
for i, page_index in enumerate(page_indices):
if page_index is not None:
batch_indices.append(i)
file_offsets.append(page_index * ctx.bytes_per_page)
values.append(host_pool.get_dummy_flat_data_page())
if not batch_indices:
return [False] * page_num
start_time = time.perf_counter()
futures = [
self.executor.submit(
ctx.clients[self.ac.next()].batch_read,
file_offsets[j : j + self.entries],
values[j : j + self.entries],
)
for j in range(0, len(batch_indices), self.entries)
]
read_results = [r for f in futures for r in f.result()]
end_time = time.perf_counter()
ionum = len(batch_indices)
if self.enable_storage_metrics:
self.prefetch_pgs.append(ionum)
self.prefetch_bandwidth.append(
ionum / (end_time - start_time) * ctx.gb_per_page
)
results = [False] * page_num
for idx, (batch_idx, read_result) in enumerate(
zip(batch_indices, read_results)
):
if read_result == ctx.bytes_per_page:
host_idx = host_indices[batch_idx * page_size].item()
host_pool.set_from_flat_data_page(host_idx, values[idx])
results[batch_idx] = True
else:
logger.error(
f"[Rank {self.rank}][Pool {pool_name.upper()}] HiCacheHF3FS get {keys[batch_idx]} failed"
)
return results
def _pool_batch_set(self, transfer: PoolTransfer) -> List[bool]:
pool_name = transfer.name
ctx = self._pool_storage_ctx[pool_name]
host_pool = self.registered_pools[pool_name]
keys = transfer.keys
host_indices = transfer.host_indices
page_size = getattr(host_pool, "page_size", 1) or 1
page_num = len(keys)
component_keys = [f"{key}_{pool_name}" for key in keys]
key_with_prefix = [(k, "") for k in component_keys]
indices = self.metadata_client.reserve_and_allocate_page_indices(
self.rank, key_with_prefix, namespace=ctx.namespace
)
if len(indices) != page_num:
logger.error(
f"[Rank {self.rank}] Pool {pool_name}: mismatched indices length"
)
if indices:
self.metadata_client.confirm_write(
self.rank, [], [idx[1] for idx in indices], namespace=ctx.namespace
)
return [False] * page_num
batch_indices, file_offsets, file_values = [], [], []
for i, (is_written, page_index) in enumerate(indices):
if is_written or page_index == -1:
continue
batch_indices.append(i)
file_offsets.append(page_index * ctx.bytes_per_page)
host_idx = host_indices[i * page_size].item()
data = host_pool.get_data_page(host_idx, flat=True)
assert data.is_contiguous()
file_values.append(data)
start_time = time.perf_counter()
futures = [
self.executor.submit(
ctx.clients[self.ac.next()].batch_write,
file_offsets[j : j + self.entries],
file_values[j : j + self.entries],
)
for j in range(0, len(batch_indices), self.entries)
]
write_results = [r == ctx.bytes_per_page for f in futures for r in f.result()]
end_time = time.perf_counter()
ionum = len(batch_indices)
if self.enable_storage_metrics:
self.backup_pgs.append(ionum)
self.backup_bandwidth.append(
ionum / (end_time - start_time) * ctx.gb_per_page
)
written_keys_to_confirm = []
pages_to_release = []
results = [idx[0] for idx in indices]
for batch_idx, write_ok in zip(batch_indices, write_results):
key = component_keys[batch_idx]
page_index = indices[batch_idx][1]
if write_ok:
written_keys_to_confirm.append((key, page_index))
else:
logger.error(
f"[Rank {self.rank}][Pool {pool_name.upper()}] HiCacheHF3FS set {keys[batch_idx]} failed"
)
pages_to_release.append(page_index)
results[batch_idx] = write_ok
if written_keys_to_confirm or pages_to_release:
self.metadata_client.confirm_write(
self.rank,
written_keys_to_confirm,
pages_to_release,
namespace=ctx.namespace,
)
return results
def batch_get_v2(
self,
transfers: List[PoolTransfer],
extra_info: Optional[HiCacheStorageExtraInfo] = None,
) -> dict:
results = {}
for transfer in transfers:
results[transfer.name] = self._pool_batch_get(transfer)
return results
def batch_set_v2(
self,
transfers: List[PoolTransfer],
extra_info: Optional[HiCacheStorageExtraInfo] = None,
) -> dict:
results = {}
for transfer in transfers:
results[transfer.name] = self._pool_batch_set(transfer)
return results
def batch_get_v1( def batch_get_v1(
self, self,
keys: List[str], keys: List[str],