[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:
@@ -142,6 +142,7 @@ class HiMambaRadixCache(MambaRadixCache):
|
||||
extra_config=extra_config,
|
||||
prefetch_threshold=prefetch_threshold,
|
||||
load_cache_event=self.load_cache_event,
|
||||
enable_storage_metrics=self.enable_storage_metrics,
|
||||
)
|
||||
self._apply_storage_runtime_config(
|
||||
storage_backend=server_args.hicache_storage_backend,
|
||||
|
||||
@@ -14,6 +14,7 @@ from fastapi.responses import ORJSONResponse
|
||||
from requests.adapters import HTTPAdapter
|
||||
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
|
||||
|
||||
# --- Configuration ---
|
||||
@@ -115,7 +116,7 @@ class GlobalMetadataState:
|
||||
|
||||
def __init__(self, persistence_path: Optional[str], save_interval: int):
|
||||
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.save_interval = save_interval
|
||||
self.save_timer: Optional[threading.Timer] = None
|
||||
@@ -132,13 +133,14 @@ class GlobalMetadataState:
|
||||
persisted_data = json.load(f)
|
||||
|
||||
with self.global_lock:
|
||||
for rank_id_str, data in persisted_data.items():
|
||||
rank_id = int(rank_id_str)
|
||||
for key_str, data in persisted_data.items():
|
||||
if ":" not in key_str:
|
||||
key_str = f"{key_str}:kv" # For backward compatibility
|
||||
num_pages = data["num_pages"]
|
||||
rank_meta = RankMetadata(num_pages)
|
||||
rank_meta.free_pages = data["free_pages"]
|
||||
rank_meta.key_to_index = OrderedDict(data["key_to_index"])
|
||||
self.ranks[rank_id] = rank_meta
|
||||
self.ranks[key_str] = rank_meta
|
||||
logging.info(
|
||||
f"Successfully loaded metadata for {len(self.ranks)} ranks."
|
||||
)
|
||||
@@ -156,9 +158,9 @@ class GlobalMetadataState:
|
||||
logging.info("Persisting metadata to disk...")
|
||||
with self.global_lock:
|
||||
serializable_state = {}
|
||||
for rank_id, rank_meta in self.ranks.items():
|
||||
for key_str, rank_meta in self.ranks.items():
|
||||
with rank_meta.lock:
|
||||
serializable_state[rank_id] = {
|
||||
serializable_state[key_str] = {
|
||||
"num_pages": rank_meta.num_pages,
|
||||
"free_pages": rank_meta.free_pages,
|
||||
"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}/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."""
|
||||
if rank not in self.state.ranks:
|
||||
key = self._rank_key(rank, namespace)
|
||||
if key not in self.state.ranks:
|
||||
raise HTTPException(
|
||||
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:
|
||||
"""Parse request JSON using orjson if available."""
|
||||
@@ -233,32 +240,38 @@ class Hf3fsMetadataServer:
|
||||
"""Initialize a rank with specified number of pages."""
|
||||
data = await self._read_json(request)
|
||||
num_pages = data["num_pages"]
|
||||
namespace = data.get("namespace", "kv")
|
||||
key = self._rank_key(rank, namespace)
|
||||
with self.state.global_lock:
|
||||
if rank in self.state.ranks:
|
||||
if key in self.state.ranks:
|
||||
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(
|
||||
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:
|
||||
logging.info(f"Initializing new Rank {rank} with {num_pages} pages.")
|
||||
self.state.ranks[rank] = RankMetadata(num_pages)
|
||||
logging.info(
|
||||
f"Initializing new Rank {rank} namespace '{namespace}' with {num_pages} pages."
|
||||
)
|
||||
self.state.ranks[key] = RankMetadata(num_pages)
|
||||
return Response(status_code=204)
|
||||
|
||||
async def exists(self, rank: int, request: Request):
|
||||
"""Check if keys exist in metadata."""
|
||||
data = await self._read_json(request)
|
||||
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)
|
||||
return self._json_response({"exists": results})
|
||||
|
||||
async def reserve_and_allocate_page_indices(self, rank: int, request: Request):
|
||||
"""Reserve and allocate page indices for keys."""
|
||||
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"]
|
||||
results = metadata.reserve_and_allocate_page_indices(keys)
|
||||
return self._json_response({"indices": results})
|
||||
@@ -266,7 +279,8 @@ class Hf3fsMetadataServer:
|
||||
async def confirm_write(self, rank: int, request: Request):
|
||||
"""Confirm write operations and release pages."""
|
||||
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", [])
|
||||
released_pages = data.get("pages_to_release", [])
|
||||
|
||||
@@ -277,20 +291,24 @@ class Hf3fsMetadataServer:
|
||||
async def delete_keys(self, rank: int, request: Request):
|
||||
"""Delete keys from metadata."""
|
||||
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"])
|
||||
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."""
|
||||
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()
|
||||
return Response(status_code=204)
|
||||
|
||||
async def get_page_indices(self, rank: int, request: Request):
|
||||
"""Get page indices for keys."""
|
||||
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"]
|
||||
results = metadata.get_page_indices(keys)
|
||||
return self._json_response({"indices": results})
|
||||
@@ -349,14 +367,19 @@ class Hf3fsGlobalMetadataClient(Hf3fsMetadataInterface):
|
||||
logging.error(f"Failed to POST to {endpoint} after retries: {e}")
|
||||
raise RuntimeError(f"Failed to connect to metadata server: {e}") from e
|
||||
|
||||
def initialize(self, rank: int, num_pages: int) -> None:
|
||||
self._post(f"{rank}/initialize", {"num_pages": num_pages})
|
||||
def initialize(
|
||||
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(
|
||||
self, rank: int, keys: List[Tuple[str, str]]
|
||||
self, rank: int, keys: List[Tuple[str, str]], namespace: PoolName = PoolName.KV
|
||||
) -> List[Tuple[bool, int]]:
|
||||
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")]
|
||||
|
||||
@@ -365,69 +388,107 @@ class Hf3fsGlobalMetadataClient(Hf3fsMetadataInterface):
|
||||
rank: int,
|
||||
written_keys_to_confirm: List[Tuple[str, int]],
|
||||
pages_to_release: List[int],
|
||||
namespace: PoolName = PoolName.KV,
|
||||
) -> None:
|
||||
self._post(
|
||||
f"{rank}/confirm_write",
|
||||
{
|
||||
"written_keys_to_confirm": written_keys_to_confirm,
|
||||
"pages_to_release": pages_to_release,
|
||||
"namespace": str(namespace),
|
||||
},
|
||||
)
|
||||
|
||||
def delete_keys(self, rank: int, keys: List[str]) -> None:
|
||||
self._post(f"{rank}/delete_keys", {"keys": keys})
|
||||
def delete_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]:
|
||||
response = self._post(f"{rank}/exists", {"keys": keys})
|
||||
def exists(
|
||||
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", [])
|
||||
|
||||
def clear(self, rank: int) -> None:
|
||||
self._post(f"{rank}/clear", {})
|
||||
def clear(self, rank: int, namespace: PoolName = PoolName.KV) -> None:
|
||||
self._post(f"{rank}/clear", {"namespace": str(namespace)})
|
||||
|
||||
def get_page_indices(self, rank: int, keys: List[str]) -> List[Optional[int]]:
|
||||
response = self._post(f"{rank}/get_page_indices", {"keys": keys})
|
||||
def get_page_indices(
|
||||
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")
|
||||
|
||||
|
||||
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):
|
||||
self.rank_metadata = None
|
||||
self._metadata: Dict[str, RankMetadata] = {} # key: "rank:namespace"
|
||||
|
||||
def initialize(self, rank: int, num_pages: int) -> None:
|
||||
self.rank_metadata = RankMetadata(num_pages)
|
||||
def _ns_key(self, rank: int, namespace: PoolName) -> str:
|
||||
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(
|
||||
self, rank: int, keys: List[Tuple[str, str]]
|
||||
self, rank: int, keys: List[Tuple[str, str]], namespace: PoolName = PoolName.KV
|
||||
) -> List[Tuple[bool, int]]:
|
||||
"""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(
|
||||
self,
|
||||
rank: int,
|
||||
written_keys_to_confirm: List[Tuple[str, int]],
|
||||
pages_to_release: List[int],
|
||||
namespace: PoolName = PoolName.KV,
|
||||
) -> None:
|
||||
"""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."""
|
||||
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."""
|
||||
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."""
|
||||
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."""
|
||||
return self.rank_metadata.get_page_indices(keys)
|
||||
return self._get_metadata(rank, namespace).get_page_indices(keys)
|
||||
|
||||
|
||||
def run_metadata_server(
|
||||
|
||||
@@ -7,6 +7,7 @@ import signal
|
||||
import threading
|
||||
import time
|
||||
from abc import ABC, abstractmethod
|
||||
from dataclasses import dataclass
|
||||
from functools import wraps
|
||||
from typing import Any, List, Optional, Tuple
|
||||
|
||||
@@ -16,6 +17,10 @@ from sglang.srt.mem_cache.hicache_storage import (
|
||||
HiCacheStorage,
|
||||
HiCacheStorageConfig,
|
||||
HiCacheStorageExtraInfo,
|
||||
PoolHitPolicy,
|
||||
PoolName,
|
||||
PoolTransfer,
|
||||
PoolTransferResult,
|
||||
)
|
||||
from sglang.srt.mem_cache.memory_pool_host import HostKVCache
|
||||
from sglang.srt.mem_cache.storage.hf3fs.hf3fs_client import Hf3fsClient
|
||||
@@ -28,7 +33,9 @@ class Hf3fsMetadataInterface(ABC):
|
||||
"""Interface for HF3FS metadata operations."""
|
||||
|
||||
@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."""
|
||||
pass
|
||||
|
||||
@@ -37,12 +44,14 @@ class Hf3fsMetadataInterface(ABC):
|
||||
self,
|
||||
rank: int,
|
||||
keys: List[Tuple[str, str]],
|
||||
namespace: PoolName = PoolName.KV,
|
||||
) -> List[Tuple[bool, int]]:
|
||||
"""
|
||||
Reserve and allocate page indices for the specified keys.
|
||||
Args:
|
||||
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.
|
||||
namespace: The namespace (pool type) for the metadata.
|
||||
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.
|
||||
"""
|
||||
@@ -54,6 +63,7 @@ class Hf3fsMetadataInterface(ABC):
|
||||
rank: int,
|
||||
written_keys_to_confirm: List[Tuple[str, int]],
|
||||
pages_to_release: List[int],
|
||||
namespace: PoolName = PoolName.KV,
|
||||
) -> None:
|
||||
"""
|
||||
Confirm that key-value pairs have been successfully written to storage.
|
||||
@@ -61,16 +71,20 @@ class Hf3fsMetadataInterface(ABC):
|
||||
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.
|
||||
pages_to_release: A list of page indices to be released.
|
||||
namespace: The namespace (pool type) for the metadata.
|
||||
"""
|
||||
pass
|
||||
|
||||
@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.
|
||||
Args:
|
||||
rank: The rank of the process.
|
||||
keys: A list of keys.
|
||||
namespace: The namespace (pool type) for the metadata.
|
||||
Returns:
|
||||
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.
|
||||
@@ -78,17 +92,21 @@ class Hf3fsMetadataInterface(ABC):
|
||||
pass
|
||||
|
||||
@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."""
|
||||
pass
|
||||
|
||||
@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."""
|
||||
pass
|
||||
|
||||
@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."""
|
||||
pass
|
||||
|
||||
@@ -151,6 +169,18 @@ def create_hf3fs_client(
|
||||
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):
|
||||
"""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_page_first_layout = is_page_first_layout
|
||||
self.enable_storage_metrics = enable_storage_metrics
|
||||
self.use_mock_client = use_mock_client
|
||||
self.numel = self.bytes_per_page // self.dtype.itemsize
|
||||
self.num_pages = self.file_size // self.bytes_per_page
|
||||
self.skip_backup = False
|
||||
@@ -220,6 +251,7 @@ class HiCacheHF3FS(HiCacheStorage):
|
||||
|
||||
self.metadata_client.initialize(self.rank, self.num_pages)
|
||||
self.lock = threading.RLock()
|
||||
self._pool_storage_ctx: dict = {}
|
||||
|
||||
atexit.register(self.close)
|
||||
|
||||
@@ -503,6 +535,8 @@ class HiCacheHF3FS(HiCacheStorage):
|
||||
def clear(self) -> None:
|
||||
try:
|
||||
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}")
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to clear HiCacheHF3FS: {e}")
|
||||
@@ -511,6 +545,9 @@ class HiCacheHF3FS(HiCacheStorage):
|
||||
try:
|
||||
for c in self.clients:
|
||||
c.close()
|
||||
for ctx in getattr(self, "_pool_storage_ctx", {}).values():
|
||||
for c in ctx.clients:
|
||||
c.close()
|
||||
self.executor.shutdown(wait=True)
|
||||
except Exception as 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}")
|
||||
|
||||
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]:
|
||||
_keys = []
|
||||
for k in keys:
|
||||
@@ -596,6 +672,212 @@ class HiCacheHF3FS(HiCacheStorage):
|
||||
|
||||
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(
|
||||
self,
|
||||
keys: List[str],
|
||||
|
||||
Reference in New Issue
Block a user