[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,
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],