[HiCache & HybridModel] mooncake backend support DSA & mamba model (#21259)
Signed-off-by: Vladislav Nosivskoy <vladnosiv@gmail.com> Co-authored-by: hzh0425 <hzh0425@apache.org> Co-authored-by: pansicheng <sicheng.pan.chn@gmail.com> Co-authored-by: ispobock <ispobaoke@gmail.com> Co-authored-by: Vladislav Nosivskoy <vladnosiv@gmail.com>
This commit is contained in:
co-authored by
hzh0425
pansicheng
ispobock
Vladislav Nosivskoy
parent
cc449ac4e5
commit
e9d6b9eb2d
@@ -24,9 +24,11 @@ from sglang.srt.mem_cache.base_prefix_cache import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.mem_cache.hicache_storage import PoolHitPolicy, PoolName, PoolTransfer
|
from sglang.srt.mem_cache.hicache_storage import PoolHitPolicy, PoolName, PoolTransfer
|
||||||
from sglang.srt.mem_cache.hybrid_cache.hybrid_cache_controller import (
|
from sglang.srt.mem_cache.hybrid_cache.hybrid_cache_controller import (
|
||||||
HybridCacheController,
|
|
||||||
PrefetchOperation,
|
PrefetchOperation,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.mem_cache.hybrid_cache.hybrid_pool_assembler import (
|
||||||
|
build_mamba_hybrid_stack,
|
||||||
|
)
|
||||||
from sglang.srt.mem_cache.mamba_radix_cache import (
|
from sglang.srt.mem_cache.mamba_radix_cache import (
|
||||||
LRUList,
|
LRUList,
|
||||||
MambaRadixCache,
|
MambaRadixCache,
|
||||||
@@ -34,13 +36,6 @@ from sglang.srt.mem_cache.mamba_radix_cache import (
|
|||||||
get_last_access_time,
|
get_last_access_time,
|
||||||
)
|
)
|
||||||
from sglang.srt.mem_cache.memory_pool import HybridLinearKVPool, HybridReqToTokenPool
|
from sglang.srt.mem_cache.memory_pool import HybridLinearKVPool, HybridReqToTokenPool
|
||||||
from sglang.srt.mem_cache.memory_pool_host import (
|
|
||||||
HostPoolGroup,
|
|
||||||
MambaPoolHost,
|
|
||||||
MHATokenToKVPoolHost,
|
|
||||||
MLATokenToKVPoolHost,
|
|
||||||
PoolEntry,
|
|
||||||
)
|
|
||||||
from sglang.srt.mem_cache.radix_cache import (
|
from sglang.srt.mem_cache.radix_cache import (
|
||||||
RadixKey,
|
RadixKey,
|
||||||
compute_node_hash_values,
|
compute_node_hash_values,
|
||||||
@@ -115,66 +110,6 @@ class HiMambaRadixCache(MambaRadixCache):
|
|||||||
)
|
)
|
||||||
|
|
||||||
self.kvcache = self.hybrid_kv_cache.full_kv_pool
|
self.kvcache = self.hybrid_kv_cache.full_kv_pool
|
||||||
kv_host_pool_cls = (
|
|
||||||
MLATokenToKVPoolHost
|
|
||||||
if self.hybrid_kv_cache.use_mla
|
|
||||||
else MHATokenToKVPoolHost
|
|
||||||
)
|
|
||||||
self.full_kv_pool_host = kv_host_pool_cls(
|
|
||||||
self.kvcache,
|
|
||||||
server_args.hicache_ratio,
|
|
||||||
server_args.hicache_size,
|
|
||||||
params.page_size,
|
|
||||||
server_args.hicache_mem_layout,
|
|
||||||
allocator_type=server_args.hicache_storage_backend,
|
|
||||||
)
|
|
||||||
self.mamba_pool_host = MambaPoolHost(
|
|
||||||
params.req_to_token_pool.mamba_pool,
|
|
||||||
server_args.hicache_ratio,
|
|
||||||
server_args.hicache_size,
|
|
||||||
allocator_type=server_args.hicache_storage_backend,
|
|
||||||
layout=server_args.hicache_mem_layout,
|
|
||||||
)
|
|
||||||
|
|
||||||
full_layer_ids = sorted(
|
|
||||||
self.hybrid_kv_cache.full_attention_layer_id_mapping.keys()
|
|
||||||
)
|
|
||||||
mamba_layer_ids = sorted(params.req_to_token_pool.mamba_map.keys())
|
|
||||||
self.transfer_layer_num = len(set(full_layer_ids) | set(mamba_layer_ids))
|
|
||||||
|
|
||||||
full_layer_mapping = dict(self.hybrid_kv_cache.full_attention_layer_id_mapping)
|
|
||||||
mamba_layer_mapping = dict(params.req_to_token_pool.mamba_map)
|
|
||||||
transfer_layer_num = self.transfer_layer_num
|
|
||||||
|
|
||||||
def kv_layer_mapper(layer_id: int) -> Optional[int]:
|
|
||||||
if not 0 <= layer_id < transfer_layer_num:
|
|
||||||
return None
|
|
||||||
return full_layer_mapping.get(layer_id)
|
|
||||||
|
|
||||||
def mamba_layer_mapper(layer_id: int) -> Optional[int]:
|
|
||||||
if not 0 <= layer_id < transfer_layer_num:
|
|
||||||
return None
|
|
||||||
return mamba_layer_mapping.get(layer_id)
|
|
||||||
|
|
||||||
self.host_pool_group = HostPoolGroup(
|
|
||||||
[
|
|
||||||
PoolEntry(
|
|
||||||
name=PoolName.KV,
|
|
||||||
host_pool=self.full_kv_pool_host,
|
|
||||||
device_pool=self.kvcache,
|
|
||||||
layer_mapper=kv_layer_mapper,
|
|
||||||
is_primary_index_anchor=True,
|
|
||||||
),
|
|
||||||
PoolEntry(
|
|
||||||
name=PoolName.MAMBA,
|
|
||||||
host_pool=self.mamba_pool_host,
|
|
||||||
device_pool=params.req_to_token_pool.mamba_pool,
|
|
||||||
layer_mapper=mamba_layer_mapper,
|
|
||||||
host_evict_fn=self.evict_mamba_host,
|
|
||||||
device_evict_fn=self.evict_mamba,
|
|
||||||
),
|
|
||||||
]
|
|
||||||
)
|
|
||||||
|
|
||||||
self.tp_group = params.tp_cache_group
|
self.tp_group = params.tp_cache_group
|
||||||
self.tp_world_size = (
|
self.tp_world_size = (
|
||||||
@@ -200,27 +135,13 @@ class HiMambaRadixCache(MambaRadixCache):
|
|||||||
self.prefetch_stop_policy = server_args.hicache_storage_prefetch_policy
|
self.prefetch_stop_policy = server_args.hicache_storage_prefetch_policy
|
||||||
|
|
||||||
self.load_cache_event = threading.Event()
|
self.load_cache_event = threading.Event()
|
||||||
self.cache_controller = HybridCacheController(
|
build_mamba_hybrid_stack(
|
||||||
params.token_to_kv_pool_allocator,
|
self,
|
||||||
self.host_pool_group,
|
params,
|
||||||
params.page_size,
|
server_args,
|
||||||
self.tp_group,
|
extra_config=extra_config,
|
||||||
load_cache_event=self.load_cache_event,
|
|
||||||
write_policy=server_args.hicache_write_policy,
|
|
||||||
io_backend=server_args.hicache_io_backend,
|
|
||||||
storage_backend=server_args.hicache_storage_backend,
|
|
||||||
prefetch_threshold=prefetch_threshold,
|
prefetch_threshold=prefetch_threshold,
|
||||||
model_name=server_args.served_model_name,
|
load_cache_event=self.load_cache_event,
|
||||||
storage_backend_extra_config=extra_config,
|
|
||||||
pp_rank=params.pp_rank,
|
|
||||||
pp_size=params.pp_size,
|
|
||||||
transfer_layer_num=self.transfer_layer_num,
|
|
||||||
)
|
|
||||||
params.req_to_token_pool.register_layer_transfer_counter(
|
|
||||||
self.cache_controller.layer_done_counter
|
|
||||||
)
|
|
||||||
self.hybrid_kv_cache.register_layer_transfer_counter(
|
|
||||||
self.cache_controller.layer_done_counter
|
|
||||||
)
|
)
|
||||||
self._apply_storage_runtime_config(
|
self._apply_storage_runtime_config(
|
||||||
storage_backend=server_args.hicache_storage_backend,
|
storage_backend=server_args.hicache_storage_backend,
|
||||||
|
|||||||
@@ -41,12 +41,13 @@ class PoolName(str, Enum):
|
|||||||
|
|
||||||
KV = "kv"
|
KV = "kv"
|
||||||
MAMBA = "mamba"
|
MAMBA = "mamba"
|
||||||
|
INDEXER = "indexer"
|
||||||
|
|
||||||
|
|
||||||
class PoolHitPolicy(str, Enum):
|
class PoolHitPolicy(str, Enum):
|
||||||
"""Hit policy for batch_exists_v2 per-pool prefix matching.
|
"""Hit policy for batch_exists_v2 per-pool prefix matching.
|
||||||
|
|
||||||
ALL_PAGES : every page in [0, kv_hit) must exist (default).
|
ALL_PAGES : every page in [0, kv_hit) must exist (e.g. DSA).
|
||||||
TRAILING_PAGES : only the last N pages must exist (e.g. Mamba/SWA states).
|
TRAILING_PAGES : only the last N pages must exist (e.g. Mamba/SWA states).
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
|||||||
@@ -25,6 +25,17 @@ from sglang.srt.mem_cache.base_prefix_cache import (
|
|||||||
MatchPrefixParams,
|
MatchPrefixParams,
|
||||||
MatchResult,
|
MatchResult,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.mem_cache.hicache_storage import (
|
||||||
|
PoolHitPolicy,
|
||||||
|
PoolName,
|
||||||
|
PoolTransfer,
|
||||||
|
)
|
||||||
|
from sglang.srt.mem_cache.hybrid_cache.hybrid_cache_controller import (
|
||||||
|
HybridCacheController,
|
||||||
|
)
|
||||||
|
from sglang.srt.mem_cache.hybrid_cache.hybrid_pool_assembler import (
|
||||||
|
build_nsa_hybrid_stack,
|
||||||
|
)
|
||||||
from sglang.srt.mem_cache.memory_pool import (
|
from sglang.srt.mem_cache.memory_pool import (
|
||||||
MHATokenToKVPool,
|
MHATokenToKVPool,
|
||||||
MLATokenToKVPool,
|
MLATokenToKVPool,
|
||||||
@@ -33,7 +44,6 @@ from sglang.srt.mem_cache.memory_pool import (
|
|||||||
from sglang.srt.mem_cache.memory_pool_host import (
|
from sglang.srt.mem_cache.memory_pool_host import (
|
||||||
MHATokenToKVPoolHost,
|
MHATokenToKVPoolHost,
|
||||||
MLATokenToKVPoolHost,
|
MLATokenToKVPoolHost,
|
||||||
NSATokenToKVPoolHost,
|
|
||||||
)
|
)
|
||||||
from sglang.srt.mem_cache.radix_cache import (
|
from sglang.srt.mem_cache.radix_cache import (
|
||||||
RadixCache,
|
RadixCache,
|
||||||
@@ -70,14 +80,8 @@ class HiRadixCache(RadixCache):
|
|||||||
allocator_type=server_args.hicache_storage_backend,
|
allocator_type=server_args.hicache_storage_backend,
|
||||||
)
|
)
|
||||||
elif isinstance(self.kv_cache, NSATokenToKVPool):
|
elif isinstance(self.kv_cache, NSATokenToKVPool):
|
||||||
self.token_to_kv_pool_host = NSATokenToKVPoolHost(
|
# Filled by build_nsa_hybrid_stack after storage extra_config is parsed.
|
||||||
self.kv_cache,
|
self.token_to_kv_pool_host = None
|
||||||
server_args.hicache_ratio,
|
|
||||||
server_args.hicache_size,
|
|
||||||
self.page_size,
|
|
||||||
server_args.hicache_mem_layout,
|
|
||||||
allocator_type=server_args.hicache_storage_backend,
|
|
||||||
)
|
|
||||||
elif isinstance(self.kv_cache, MLATokenToKVPool):
|
elif isinstance(self.kv_cache, MLATokenToKVPool):
|
||||||
self.token_to_kv_pool_host = MLATokenToKVPoolHost(
|
self.token_to_kv_pool_host = MLATokenToKVPoolHost(
|
||||||
self.kv_cache,
|
self.kv_cache,
|
||||||
@@ -88,7 +92,9 @@ class HiRadixCache(RadixCache):
|
|||||||
allocator_type=server_args.hicache_storage_backend,
|
allocator_type=server_args.hicache_storage_backend,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
raise ValueError(f"HiRadixCache only supports MHA and MLA yet")
|
raise ValueError(
|
||||||
|
"HiRadixCache only supports MHA, MLA, and NSA (DSA) models"
|
||||||
|
)
|
||||||
|
|
||||||
self.tp_group = params.tp_cache_group
|
self.tp_group = params.tp_cache_group
|
||||||
self.tp_world_size = torch.distributed.get_world_size(group=self.tp_group)
|
self.tp_world_size = torch.distributed.get_world_size(group=self.tp_group)
|
||||||
@@ -114,24 +120,35 @@ class HiRadixCache(RadixCache):
|
|||||||
self.prefetch_stop_policy = server_args.hicache_storage_prefetch_policy
|
self.prefetch_stop_policy = server_args.hicache_storage_prefetch_policy
|
||||||
|
|
||||||
self.load_cache_event = threading.Event()
|
self.load_cache_event = threading.Event()
|
||||||
self.cache_controller = HiCacheController(
|
if isinstance(self.kv_cache, NSATokenToKVPool):
|
||||||
params.token_to_kv_pool_allocator,
|
build_nsa_hybrid_stack(
|
||||||
self.token_to_kv_pool_host,
|
self,
|
||||||
self.page_size,
|
params,
|
||||||
self.tp_group,
|
server_args,
|
||||||
load_cache_event=self.load_cache_event,
|
extra_config=extra_config,
|
||||||
write_policy=server_args.hicache_write_policy,
|
prefetch_threshold=prefetch_threshold,
|
||||||
io_backend=server_args.hicache_io_backend,
|
enable_storage_metrics=self.enable_storage_metrics,
|
||||||
storage_backend=server_args.hicache_storage_backend,
|
load_cache_event=self.load_cache_event,
|
||||||
prefetch_threshold=prefetch_threshold,
|
)
|
||||||
model_name=server_args.served_model_name,
|
else:
|
||||||
storage_backend_extra_config=extra_config,
|
self.cache_controller = HiCacheController(
|
||||||
pp_rank=self.pp_rank,
|
params.token_to_kv_pool_allocator,
|
||||||
pp_size=self.pp_size,
|
self.token_to_kv_pool_host,
|
||||||
attn_cp_rank=self.attn_cp_rank,
|
self.page_size,
|
||||||
attn_cp_size=self.attn_cp_size,
|
self.tp_group,
|
||||||
enable_storage_metrics=self.enable_storage_metrics,
|
load_cache_event=self.load_cache_event,
|
||||||
)
|
write_policy=server_args.hicache_write_policy,
|
||||||
|
io_backend=server_args.hicache_io_backend,
|
||||||
|
storage_backend=server_args.hicache_storage_backend,
|
||||||
|
prefetch_threshold=prefetch_threshold,
|
||||||
|
model_name=server_args.served_model_name,
|
||||||
|
storage_backend_extra_config=extra_config,
|
||||||
|
pp_rank=self.pp_rank,
|
||||||
|
pp_size=self.pp_size,
|
||||||
|
attn_cp_rank=self.attn_cp_rank,
|
||||||
|
attn_cp_size=self.attn_cp_size,
|
||||||
|
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,
|
||||||
prefetch_threshold=prefetch_threshold,
|
prefetch_threshold=prefetch_threshold,
|
||||||
@@ -326,6 +343,7 @@ class HiRadixCache(RadixCache):
|
|||||||
prefetch_threshold=prefetch_threshold,
|
prefetch_threshold=prefetch_threshold,
|
||||||
model_name=served_model_name,
|
model_name=served_model_name,
|
||||||
storage_backend_extra_config=extra_config,
|
storage_backend_extra_config=extra_config,
|
||||||
|
**self._get_hybrid_storage_attach_kwargs(),
|
||||||
)
|
)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.exception(
|
logger.exception(
|
||||||
@@ -591,6 +609,24 @@ class HiRadixCache(RadixCache):
|
|||||||
height += 1
|
height += 1
|
||||||
return height
|
return height
|
||||||
|
|
||||||
|
def _get_extra_pools(self) -> dict:
|
||||||
|
if not isinstance(self.cache_controller, HybridCacheController):
|
||||||
|
return {}
|
||||||
|
if isinstance(self.kv_cache, NSATokenToKVPool):
|
||||||
|
pool = PoolTransfer(
|
||||||
|
name=PoolName.INDEXER,
|
||||||
|
hit_policy=PoolHitPolicy.ALL_PAGES,
|
||||||
|
)
|
||||||
|
return {"extra_pools": [pool]}
|
||||||
|
else:
|
||||||
|
return {}
|
||||||
|
|
||||||
|
def _get_hybrid_storage_attach_kwargs(self) -> dict:
|
||||||
|
"""Extra kwargs for attach_storage_backend when controller is HybridCacheController."""
|
||||||
|
if isinstance(self.cache_controller, HybridCacheController):
|
||||||
|
return {"host_pools": self.cache_controller.mem_pool_host.entries}
|
||||||
|
return {}
|
||||||
|
|
||||||
def clear_storage_backend(self) -> bool:
|
def clear_storage_backend(self) -> bool:
|
||||||
if self.enable_storage:
|
if self.enable_storage:
|
||||||
try:
|
try:
|
||||||
@@ -625,12 +661,14 @@ class HiRadixCache(RadixCache):
|
|||||||
host_indices = self.cache_controller.write(
|
host_indices = self.cache_controller.write(
|
||||||
device_indices=node.value,
|
device_indices=node.value,
|
||||||
node_id=node.id,
|
node_id=node.id,
|
||||||
|
**self._get_extra_pools(),
|
||||||
)
|
)
|
||||||
if host_indices is None:
|
if host_indices is None:
|
||||||
self.evict_host(len(node.value))
|
self.evict_host(len(node.value))
|
||||||
host_indices = self.cache_controller.write(
|
host_indices = self.cache_controller.write(
|
||||||
device_indices=node.value,
|
device_indices=node.value,
|
||||||
node_id=node.id,
|
node_id=node.id,
|
||||||
|
**self._get_extra_pools(),
|
||||||
)
|
)
|
||||||
if host_indices is not None:
|
if host_indices is not None:
|
||||||
node.host_value = host_indices.clone()
|
node.host_value = host_indices.clone()
|
||||||
@@ -652,7 +690,11 @@ class HiRadixCache(RadixCache):
|
|||||||
)
|
)
|
||||||
|
|
||||||
operation_id = self.cache_controller.write_storage(
|
operation_id = self.cache_controller.write_storage(
|
||||||
node.host_value, node.key, node.hash_value, prefix_keys
|
node.host_value,
|
||||||
|
node.key,
|
||||||
|
node.hash_value,
|
||||||
|
prefix_keys,
|
||||||
|
**self._get_extra_pools(),
|
||||||
)
|
)
|
||||||
self.ongoing_backup[operation_id] = node
|
self.ongoing_backup[operation_id] = node
|
||||||
node.protect_host()
|
node.protect_host()
|
||||||
@@ -925,12 +967,16 @@ class HiRadixCache(RadixCache):
|
|||||||
return None
|
return None
|
||||||
|
|
||||||
device_indices = self.cache_controller.load(
|
device_indices = self.cache_controller.load(
|
||||||
host_indices=host_indices, node_id=last_hit_node.id
|
host_indices=host_indices,
|
||||||
|
node_id=last_hit_node.id,
|
||||||
|
**self._get_extra_pools(),
|
||||||
)
|
)
|
||||||
if device_indices is None:
|
if device_indices is None:
|
||||||
self.evict(EvictParams(num_tokens=len(host_indices)))
|
self.evict(EvictParams(num_tokens=len(host_indices)))
|
||||||
device_indices = self.cache_controller.load(
|
device_indices = self.cache_controller.load(
|
||||||
host_indices=host_indices, node_id=last_hit_node.id
|
host_indices=host_indices,
|
||||||
|
node_id=last_hit_node.id,
|
||||||
|
**self._get_extra_pools(),
|
||||||
)
|
)
|
||||||
self.dec_lock_ref(ancester_node)
|
self.dec_lock_ref(ancester_node)
|
||||||
if device_indices is None:
|
if device_indices is None:
|
||||||
@@ -1231,7 +1277,12 @@ class HiRadixCache(RadixCache):
|
|||||||
# no sufficient host memory for prefetch
|
# no sufficient host memory for prefetch
|
||||||
return
|
return
|
||||||
operation = self.cache_controller.prefetch(
|
operation = self.cache_controller.prefetch(
|
||||||
req_id, host_indices, new_input_tokens, last_hash, prefix_keys
|
req_id,
|
||||||
|
host_indices,
|
||||||
|
new_input_tokens,
|
||||||
|
last_hash,
|
||||||
|
prefix_keys,
|
||||||
|
**self._get_extra_pools(),
|
||||||
)
|
)
|
||||||
self.ongoing_prefetch[req_id] = (
|
self.ongoing_prefetch[req_id] = (
|
||||||
last_host_node,
|
last_host_node,
|
||||||
|
|||||||
@@ -163,6 +163,7 @@ class HybridCacheController(BaseHiCacheController):
|
|||||||
pp_rank: int = 0,
|
pp_rank: int = 0,
|
||||||
pp_size: int = 1,
|
pp_size: int = 1,
|
||||||
transfer_layer_num: Optional[int] = None,
|
transfer_layer_num: Optional[int] = None,
|
||||||
|
enable_storage_metrics: bool = False,
|
||||||
):
|
):
|
||||||
startup_storage_backend = storage_backend
|
startup_storage_backend = storage_backend
|
||||||
super().__init__(
|
super().__init__(
|
||||||
@@ -179,6 +180,7 @@ class HybridCacheController(BaseHiCacheController):
|
|||||||
storage_backend_extra_config=storage_backend_extra_config,
|
storage_backend_extra_config=storage_backend_extra_config,
|
||||||
pp_rank=pp_rank,
|
pp_rank=pp_rank,
|
||||||
pp_size=pp_size,
|
pp_size=pp_size,
|
||||||
|
enable_storage_metrics=enable_storage_metrics,
|
||||||
)
|
)
|
||||||
# Override layer_num: hybrid models transfer all layers (For example, Linear Model (KV + Mamba)),
|
# Override layer_num: hybrid models transfer all layers (For example, Linear Model (KV + Mamba)),
|
||||||
# not just the full attention layers reported by full_kv_pool.
|
# not just the full attention layers reported by full_kv_pool.
|
||||||
@@ -230,7 +232,10 @@ class HybridCacheController(BaseHiCacheController):
|
|||||||
if host_indices is None:
|
if host_indices is None:
|
||||||
return None
|
return None
|
||||||
pool_transfers = self._resolve_pool_transfers_allocation(
|
pool_transfers = self._resolve_pool_transfers_allocation(
|
||||||
extra_pools, alloc_host=True
|
extra_pools,
|
||||||
|
alloc_host=True,
|
||||||
|
kv_device_indices=device_indices,
|
||||||
|
kv_host_indices=host_indices,
|
||||||
)
|
)
|
||||||
if pool_transfers is None and extra_pools:
|
if pool_transfers is None and extra_pools:
|
||||||
self.mem_pool_host.free(host_indices)
|
self.mem_pool_host.free(host_indices)
|
||||||
@@ -289,7 +294,10 @@ class HybridCacheController(BaseHiCacheController):
|
|||||||
device_indices = torch.empty((0,), dtype=torch.int64, device=self.device)
|
device_indices = torch.empty((0,), dtype=torch.int64, device=self.device)
|
||||||
|
|
||||||
pool_transfers = self._resolve_pool_transfers_allocation(
|
pool_transfers = self._resolve_pool_transfers_allocation(
|
||||||
extra_pools, alloc_host=False
|
extra_pools,
|
||||||
|
alloc_host=False,
|
||||||
|
kv_device_indices=device_indices,
|
||||||
|
kv_host_indices=host_indices,
|
||||||
)
|
)
|
||||||
if pool_transfers is None and extra_pools:
|
if pool_transfers is None and extra_pools:
|
||||||
if need_load_kv:
|
if need_load_kv:
|
||||||
@@ -415,6 +423,7 @@ class HybridCacheController(BaseHiCacheController):
|
|||||||
def _page_transfer(self, operation):
|
def _page_transfer(self, operation):
|
||||||
# Transfer extra pools
|
# Transfer extra pools
|
||||||
if operation.pool_transfers and not operation.is_terminated():
|
if operation.pool_transfers and not operation.is_terminated():
|
||||||
|
self._resolve_shared_pool_transfers(operation)
|
||||||
results = self.storage_backend.batch_get_v2(operation.pool_transfers)
|
results = self.storage_backend.batch_get_v2(operation.pool_transfers)
|
||||||
operation.pool_storage_result.update_extra_pool_hit_pages(results)
|
operation.pool_storage_result.update_extra_pool_hit_pages(results)
|
||||||
|
|
||||||
@@ -424,12 +433,20 @@ class HybridCacheController(BaseHiCacheController):
|
|||||||
def _page_backup(self, operation):
|
def _page_backup(self, operation):
|
||||||
# Backup extra pools
|
# Backup extra pools
|
||||||
if operation.pool_transfers:
|
if operation.pool_transfers:
|
||||||
|
self._resolve_shared_pool_transfers(operation)
|
||||||
results = self.storage_backend.batch_set_v2(operation.pool_transfers)
|
results = self.storage_backend.batch_set_v2(operation.pool_transfers)
|
||||||
operation.pool_storage_result.update_extra_pool_hit_pages(results)
|
operation.pool_storage_result.update_extra_pool_hit_pages(results)
|
||||||
|
|
||||||
# Backup kv pools
|
# Backup kv pools
|
||||||
super()._page_backup(operation)
|
super()._page_backup(operation)
|
||||||
|
|
||||||
|
def _resolve_shared_pool_transfers(self, operation):
|
||||||
|
for transfer in operation.pool_transfers:
|
||||||
|
entry = self.mem_pool_host.entry_map.get(transfer.name)
|
||||||
|
if entry.share_indices_with_anchor:
|
||||||
|
transfer.keys = operation.hash_value
|
||||||
|
transfer.host_indices = operation.host_indices
|
||||||
|
|
||||||
def _sync_trailing_keys(
|
def _sync_trailing_keys(
|
||||||
self,
|
self,
|
||||||
pool_transfers: list[PoolTransfer],
|
pool_transfers: list[PoolTransfer],
|
||||||
@@ -454,6 +471,8 @@ class HybridCacheController(BaseHiCacheController):
|
|||||||
self,
|
self,
|
||||||
extra_pools: Optional[list[PoolTransfer]],
|
extra_pools: Optional[list[PoolTransfer]],
|
||||||
alloc_host: bool,
|
alloc_host: bool,
|
||||||
|
kv_device_indices: Optional[torch.Tensor] = None,
|
||||||
|
kv_host_indices: Optional[torch.Tensor] = None,
|
||||||
) -> Optional[list[PoolTransfer]]:
|
) -> Optional[list[PoolTransfer]]:
|
||||||
"""Auto-alloc host or device indices for PoolTransfers where they are None."""
|
"""Auto-alloc host or device indices for PoolTransfers where they are None."""
|
||||||
if not extra_pools:
|
if not extra_pools:
|
||||||
@@ -463,6 +482,10 @@ class HybridCacheController(BaseHiCacheController):
|
|||||||
entry = self.mem_pool_host.entry_map.get(pool.name)
|
entry = self.mem_pool_host.entry_map.get(pool.name)
|
||||||
if entry is None:
|
if entry is None:
|
||||||
continue
|
continue
|
||||||
|
if entry.share_indices_with_anchor:
|
||||||
|
pool.device_indices = kv_device_indices
|
||||||
|
pool.host_indices = kv_host_indices
|
||||||
|
continue
|
||||||
if alloc_host:
|
if alloc_host:
|
||||||
if pool.host_indices is not None or pool.device_indices is None:
|
if pool.host_indices is not None or pool.device_indices is None:
|
||||||
continue
|
continue
|
||||||
|
|||||||
@@ -0,0 +1,212 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import logging
|
||||||
|
from typing import TYPE_CHECKING, Optional
|
||||||
|
|
||||||
|
from sglang.srt.mem_cache.hicache_storage import PoolName
|
||||||
|
from sglang.srt.mem_cache.hybrid_cache.hybrid_cache_controller import (
|
||||||
|
HybridCacheController,
|
||||||
|
)
|
||||||
|
from sglang.srt.mem_cache.memory_pool_host import (
|
||||||
|
HostPoolGroup,
|
||||||
|
MambaPoolHost,
|
||||||
|
MHATokenToKVPoolHost,
|
||||||
|
MLATokenToKVPoolHost,
|
||||||
|
NSAIndexerPoolHost,
|
||||||
|
PoolEntry,
|
||||||
|
)
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from sglang.srt.mem_cache.cache_init_params import CacheInitParams
|
||||||
|
from sglang.srt.mem_cache.hi_mamba_radix_cache import HiMambaRadixCache
|
||||||
|
from sglang.srt.mem_cache.hiradix_cache import HiRadixCache
|
||||||
|
from sglang.srt.server_args import ServerArgs
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
def build_nsa_hybrid_stack(
|
||||||
|
radix_cache: "HiRadixCache",
|
||||||
|
params: "CacheInitParams",
|
||||||
|
server_args: "ServerArgs",
|
||||||
|
*,
|
||||||
|
extra_config: dict,
|
||||||
|
prefetch_threshold: int,
|
||||||
|
enable_storage_metrics: bool,
|
||||||
|
load_cache_event,
|
||||||
|
) -> None:
|
||||||
|
"""HostPoolGroup (KV + indexer) + HybridCacheController for NSA (DSA)."""
|
||||||
|
try:
|
||||||
|
kv = radix_cache.kv_cache
|
||||||
|
mla_host = MLATokenToKVPoolHost(
|
||||||
|
kv,
|
||||||
|
server_args.hicache_ratio,
|
||||||
|
server_args.hicache_size,
|
||||||
|
radix_cache.page_size,
|
||||||
|
server_args.hicache_mem_layout,
|
||||||
|
allocator_type=server_args.hicache_storage_backend,
|
||||||
|
override_kv_cache_dim=kv.kv_cache_dim,
|
||||||
|
)
|
||||||
|
indexer_host = NSAIndexerPoolHost(
|
||||||
|
kv,
|
||||||
|
mla_host,
|
||||||
|
server_args.hicache_mem_layout,
|
||||||
|
allocator_type=server_args.hicache_storage_backend,
|
||||||
|
)
|
||||||
|
layer_num = kv.layer_num
|
||||||
|
|
||||||
|
def layer_mapper(layer_id: int):
|
||||||
|
if 0 <= layer_id < layer_num:
|
||||||
|
return layer_id
|
||||||
|
return None
|
||||||
|
|
||||||
|
host_pool_group = HostPoolGroup(
|
||||||
|
[
|
||||||
|
PoolEntry(
|
||||||
|
name=PoolName.KV,
|
||||||
|
host_pool=mla_host,
|
||||||
|
device_pool=kv,
|
||||||
|
layer_mapper=layer_mapper,
|
||||||
|
is_primary_index_anchor=True,
|
||||||
|
),
|
||||||
|
PoolEntry(
|
||||||
|
name=PoolName.INDEXER,
|
||||||
|
host_pool=indexer_host,
|
||||||
|
device_pool=kv,
|
||||||
|
layer_mapper=layer_mapper,
|
||||||
|
share_indices_with_anchor=True,
|
||||||
|
),
|
||||||
|
]
|
||||||
|
)
|
||||||
|
cache_controller = HybridCacheController(
|
||||||
|
params.token_to_kv_pool_allocator,
|
||||||
|
host_pool_group,
|
||||||
|
radix_cache.page_size,
|
||||||
|
radix_cache.tp_group,
|
||||||
|
load_cache_event=load_cache_event,
|
||||||
|
write_policy=server_args.hicache_write_policy,
|
||||||
|
io_backend=server_args.hicache_io_backend,
|
||||||
|
storage_backend=server_args.hicache_storage_backend,
|
||||||
|
prefetch_threshold=prefetch_threshold,
|
||||||
|
model_name=server_args.served_model_name,
|
||||||
|
storage_backend_extra_config=extra_config,
|
||||||
|
pp_rank=radix_cache.pp_rank,
|
||||||
|
pp_size=radix_cache.pp_size,
|
||||||
|
transfer_layer_num=layer_num,
|
||||||
|
enable_storage_metrics=enable_storage_metrics,
|
||||||
|
)
|
||||||
|
radix_cache.full_kv_pool_host = mla_host
|
||||||
|
radix_cache.token_to_kv_pool_host = host_pool_group
|
||||||
|
radix_cache.cache_controller = cache_controller
|
||||||
|
logger.info(
|
||||||
|
"Hybrid hierarchical cache: HostPoolGroup(KV + INDEXER), HybridCacheController, "
|
||||||
|
"transfer_layer_num=%s",
|
||||||
|
layer_num,
|
||||||
|
)
|
||||||
|
except Exception:
|
||||||
|
logger.exception("build_nsa_hybrid_stack failed")
|
||||||
|
raise
|
||||||
|
|
||||||
|
|
||||||
|
def build_mamba_hybrid_stack(
|
||||||
|
mamba_cache: "HiMambaRadixCache",
|
||||||
|
params: "CacheInitParams",
|
||||||
|
server_args: "ServerArgs",
|
||||||
|
*,
|
||||||
|
extra_config: dict,
|
||||||
|
prefetch_threshold: int,
|
||||||
|
load_cache_event,
|
||||||
|
enable_storage_metrics: bool = False,
|
||||||
|
) -> None:
|
||||||
|
"""HostPoolGroup (KV + Mamba) + HybridCacheController for hybrid SSM models."""
|
||||||
|
try:
|
||||||
|
hybrid_kv = mamba_cache.hybrid_kv_cache
|
||||||
|
kvcache = mamba_cache.kvcache
|
||||||
|
kv_host_pool_cls = (
|
||||||
|
MLATokenToKVPoolHost if hybrid_kv.use_mla else MHATokenToKVPoolHost
|
||||||
|
)
|
||||||
|
full_kv_pool_host = kv_host_pool_cls(
|
||||||
|
kvcache,
|
||||||
|
server_args.hicache_ratio,
|
||||||
|
server_args.hicache_size,
|
||||||
|
params.page_size,
|
||||||
|
server_args.hicache_mem_layout,
|
||||||
|
allocator_type=server_args.hicache_storage_backend,
|
||||||
|
)
|
||||||
|
mamba_pool_host = MambaPoolHost(
|
||||||
|
params.req_to_token_pool.mamba_pool,
|
||||||
|
server_args.hicache_ratio,
|
||||||
|
server_args.hicache_size,
|
||||||
|
allocator_type=server_args.hicache_storage_backend,
|
||||||
|
layout=server_args.hicache_mem_layout,
|
||||||
|
)
|
||||||
|
|
||||||
|
full_layer_ids = sorted(hybrid_kv.full_attention_layer_id_mapping.keys())
|
||||||
|
mamba_layer_ids = sorted(params.req_to_token_pool.mamba_map.keys())
|
||||||
|
transfer_layer_num = len(set(full_layer_ids) | set(mamba_layer_ids))
|
||||||
|
full_layer_mapping = dict(hybrid_kv.full_attention_layer_id_mapping)
|
||||||
|
mamba_layer_mapping = dict(params.req_to_token_pool.mamba_map)
|
||||||
|
|
||||||
|
def kv_layer_mapper(layer_id: int) -> Optional[int]:
|
||||||
|
if not 0 <= layer_id < transfer_layer_num:
|
||||||
|
return None
|
||||||
|
return full_layer_mapping.get(layer_id)
|
||||||
|
|
||||||
|
def mamba_layer_mapper(layer_id: int) -> Optional[int]:
|
||||||
|
if not 0 <= layer_id < transfer_layer_num:
|
||||||
|
return None
|
||||||
|
return mamba_layer_mapping.get(layer_id)
|
||||||
|
|
||||||
|
host_pool_group = HostPoolGroup(
|
||||||
|
[
|
||||||
|
PoolEntry(
|
||||||
|
name=PoolName.KV,
|
||||||
|
host_pool=full_kv_pool_host,
|
||||||
|
device_pool=kvcache,
|
||||||
|
layer_mapper=kv_layer_mapper,
|
||||||
|
is_primary_index_anchor=True,
|
||||||
|
),
|
||||||
|
PoolEntry(
|
||||||
|
name=PoolName.MAMBA,
|
||||||
|
host_pool=mamba_pool_host,
|
||||||
|
device_pool=params.req_to_token_pool.mamba_pool,
|
||||||
|
layer_mapper=mamba_layer_mapper,
|
||||||
|
host_evict_fn=mamba_cache.evict_mamba_host,
|
||||||
|
device_evict_fn=mamba_cache.evict_mamba,
|
||||||
|
),
|
||||||
|
]
|
||||||
|
)
|
||||||
|
cache_controller = HybridCacheController(
|
||||||
|
params.token_to_kv_pool_allocator,
|
||||||
|
host_pool_group,
|
||||||
|
params.page_size,
|
||||||
|
params.tp_cache_group,
|
||||||
|
load_cache_event=load_cache_event,
|
||||||
|
write_policy=server_args.hicache_write_policy,
|
||||||
|
io_backend=server_args.hicache_io_backend,
|
||||||
|
storage_backend=server_args.hicache_storage_backend,
|
||||||
|
prefetch_threshold=prefetch_threshold,
|
||||||
|
model_name=server_args.served_model_name,
|
||||||
|
storage_backend_extra_config=extra_config,
|
||||||
|
pp_rank=params.pp_rank,
|
||||||
|
pp_size=params.pp_size,
|
||||||
|
transfer_layer_num=transfer_layer_num,
|
||||||
|
enable_storage_metrics=enable_storage_metrics,
|
||||||
|
)
|
||||||
|
mamba_cache.full_kv_pool_host = full_kv_pool_host
|
||||||
|
mamba_cache.mamba_pool_host = mamba_pool_host
|
||||||
|
mamba_cache.transfer_layer_num = transfer_layer_num
|
||||||
|
mamba_cache.host_pool_group = host_pool_group
|
||||||
|
mamba_cache.cache_controller = cache_controller
|
||||||
|
params.req_to_token_pool.register_layer_transfer_counter(
|
||||||
|
cache_controller.layer_done_counter
|
||||||
|
)
|
||||||
|
hybrid_kv.register_layer_transfer_counter(cache_controller.layer_done_counter)
|
||||||
|
logger.info(
|
||||||
|
"Hybrid hierarchical cache: HostPoolGroup(KV + MAMBA), HybridCacheController, "
|
||||||
|
"transfer_layer_num=%s",
|
||||||
|
transfer_layer_num,
|
||||||
|
)
|
||||||
|
except Exception:
|
||||||
|
logger.exception("build_mamba_hybrid_stack failed")
|
||||||
|
raise
|
||||||
@@ -64,6 +64,9 @@ if _is_npu:
|
|||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
# Host RAM to leave free when sizing HiCache pools (OS, other processes).
|
||||||
|
HICACHE_HOST_MEMORY_RESERVE_BYTES: int = 10 * (1024**3)
|
||||||
|
|
||||||
|
|
||||||
def synchronized(func):
|
def synchronized(func):
|
||||||
@wraps(func)
|
@wraps(func)
|
||||||
@@ -187,9 +190,7 @@ class HostKVCache(abc.ABC):
|
|||||||
# Verify there is enough available host memory.
|
# Verify there is enough available host memory.
|
||||||
host_mem = psutil.virtual_memory()
|
host_mem = psutil.virtual_memory()
|
||||||
requested_bytes = self.size * self.size_per_token
|
requested_bytes = self.size * self.size_per_token
|
||||||
# preserve at least 10GB for other usage
|
available_bytes = host_mem.available - HICACHE_HOST_MEMORY_RESERVE_BYTES
|
||||||
ten_gb = 10 * (1024**3)
|
|
||||||
available_bytes = host_mem.available - ten_gb
|
|
||||||
if requested_bytes > available_bytes:
|
if requested_bytes > available_bytes:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"Not enough host memory available. Requesting "
|
f"Not enough host memory available. Requesting "
|
||||||
@@ -1213,6 +1214,10 @@ class MambaPoolHost(HostKVCache):
|
|||||||
conv_state.shape[2:] for conv_state in device_pool.mamba_cache.conv
|
conv_state.shape[2:] for conv_state in device_pool.mamba_cache.conv
|
||||||
]
|
]
|
||||||
self.temporal_state_shape = device_pool.mamba_cache.temporal.shape[2:]
|
self.temporal_state_shape = device_pool.mamba_cache.temporal.shape[2:]
|
||||||
|
self.temporal_state_elem_size = int(np.prod(self.temporal_state_shape))
|
||||||
|
self.conv_state_elem_sizes = [
|
||||||
|
int(np.prod(conv_shape)) for conv_shape in self.conv_state_shapes
|
||||||
|
]
|
||||||
self.conv_dtype = device_pool.mamba_cache.conv[0].dtype
|
self.conv_dtype = device_pool.mamba_cache.conv[0].dtype
|
||||||
self.temporal_dtype = device_pool.mamba_cache.temporal.dtype
|
self.temporal_dtype = device_pool.mamba_cache.temporal.dtype
|
||||||
self.dtype = self.conv_dtype
|
self.dtype = self.conv_dtype
|
||||||
@@ -1232,8 +1237,7 @@ class MambaPoolHost(HostKVCache):
|
|||||||
|
|
||||||
host_mem = psutil.virtual_memory()
|
host_mem = psutil.virtual_memory()
|
||||||
requested_bytes = self.size * self.size_per_token
|
requested_bytes = self.size * self.size_per_token
|
||||||
ten_gb = 10 * (1024**3)
|
available_bytes = host_mem.available - HICACHE_HOST_MEMORY_RESERVE_BYTES
|
||||||
available_bytes = host_mem.available - ten_gb
|
|
||||||
if requested_bytes > available_bytes:
|
if requested_bytes > available_bytes:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"Not enough host memory available. Requesting "
|
f"Not enough host memory available. Requesting "
|
||||||
@@ -1306,6 +1310,10 @@ class MambaPoolHost(HostKVCache):
|
|||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def get_hybrid_pool_buffer(self):
|
||||||
|
# Expose all mamba host tensors that need Mooncake buffer registration.
|
||||||
|
return [self.temporal_buffer, *self.conv_buffer]
|
||||||
|
|
||||||
def _iter_page_tensors(self, index: int):
|
def _iter_page_tensors(self, index: int):
|
||||||
if self.layout in ["page_first", "page_first_direct"]:
|
if self.layout in ["page_first", "page_first_direct"]:
|
||||||
yield self.temporal_buffer[index]
|
yield self.temporal_buffer[index]
|
||||||
@@ -1347,12 +1355,11 @@ class MambaPoolHost(HostKVCache):
|
|||||||
return len(indices)
|
return len(indices)
|
||||||
|
|
||||||
def get_size_per_token(self):
|
def get_size_per_token(self):
|
||||||
conv_total_size = 0
|
conv_total_size = sum(
|
||||||
for conv_shape in self.conv_state_shapes:
|
conv_elem_size * self.conv_dtype.itemsize
|
||||||
conv_total_size += int(np.prod(conv_shape)) * self.conv_dtype.itemsize
|
for conv_elem_size in self.conv_state_elem_sizes
|
||||||
temporal_size = (
|
|
||||||
int(np.prod(self.temporal_state_shape)) * self.temporal_dtype.itemsize
|
|
||||||
)
|
)
|
||||||
|
temporal_size = self.temporal_state_elem_size * self.temporal_dtype.itemsize
|
||||||
return (conv_total_size + temporal_size) * self.num_mamba_layers
|
return (conv_total_size + temporal_size) * self.num_mamba_layers
|
||||||
|
|
||||||
def get_ksize_per_token(self):
|
def get_ksize_per_token(self):
|
||||||
@@ -1588,6 +1595,65 @@ class MambaPoolHost(HostKVCache):
|
|||||||
restored = tensor_bytes.view(dtype=tensor.dtype).reshape(tensor.shape)
|
restored = tensor_bytes.view(dtype=tensor.dtype).reshape(tensor.shape)
|
||||||
tensor.copy_(restored)
|
tensor.copy_(restored)
|
||||||
|
|
||||||
|
def get_page_buffer_meta(self, indices):
|
||||||
|
"""Meta data for zero-copy storage I/O.
|
||||||
|
|
||||||
|
Only page-first layouts are supported for mamba storage zero-copy because
|
||||||
|
each page slot in temporal/conv buffers is directly addressable.
|
||||||
|
"""
|
||||||
|
assert len(indices) % self.page_size == 0
|
||||||
|
if self.layout not in ["page_first", "page_first_direct"]:
|
||||||
|
raise ValueError(
|
||||||
|
f"Mamba storage zero-copy requires page_first layout, got {self.layout}"
|
||||||
|
)
|
||||||
|
indices = indices.tolist()
|
||||||
|
ptr_list = []
|
||||||
|
element_size_list = []
|
||||||
|
|
||||||
|
# Compute base pointers once; each page pointer is offset from these bases.
|
||||||
|
temporal_base_ptr = self.temporal_buffer.data_ptr()
|
||||||
|
conv_base_ptrs = [buf.data_ptr() for buf in self.conv_buffer]
|
||||||
|
# Component sizes are constant across pages, so precompute once as well.
|
||||||
|
temporal_element_size = (
|
||||||
|
self.page_size
|
||||||
|
* self.num_mamba_layers
|
||||||
|
* self.temporal_dtype.itemsize
|
||||||
|
* self.temporal_state_elem_size
|
||||||
|
)
|
||||||
|
conv_element_sizes = [
|
||||||
|
(
|
||||||
|
self.page_size
|
||||||
|
* self.num_mamba_layers
|
||||||
|
* self.conv_dtype.itemsize
|
||||||
|
* self.conv_state_elem_sizes[i]
|
||||||
|
)
|
||||||
|
for i in range(len(self.conv_state_shapes))
|
||||||
|
]
|
||||||
|
|
||||||
|
for i in range(0, len(indices), self.page_size):
|
||||||
|
# Emit component pointers in stable order:
|
||||||
|
# temporal first, then conv_0..conv_n for this page.
|
||||||
|
temporal_ptr = (
|
||||||
|
temporal_base_ptr
|
||||||
|
+ indices[i]
|
||||||
|
* self.num_mamba_layers
|
||||||
|
* self.temporal_state_elem_size
|
||||||
|
* self.temporal_dtype.itemsize
|
||||||
|
)
|
||||||
|
ptr_list.append(temporal_ptr)
|
||||||
|
element_size_list.append(temporal_element_size)
|
||||||
|
for j in range(len(self.conv_buffer)):
|
||||||
|
conv_ptr = (
|
||||||
|
conv_base_ptrs[j]
|
||||||
|
+ indices[i]
|
||||||
|
* self.num_mamba_layers
|
||||||
|
* self.conv_state_elem_sizes[j]
|
||||||
|
* self.conv_dtype.itemsize
|
||||||
|
)
|
||||||
|
ptr_list.append(conv_ptr)
|
||||||
|
element_size_list.append(conv_element_sizes[j])
|
||||||
|
return ptr_list, element_size_list
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class PoolEntry:
|
class PoolEntry:
|
||||||
@@ -1596,6 +1662,9 @@ class PoolEntry:
|
|||||||
device_pool: Any
|
device_pool: Any
|
||||||
layer_mapper: Callable[[int], Optional[int]]
|
layer_mapper: Callable[[int], Optional[int]]
|
||||||
is_primary_index_anchor: bool = False
|
is_primary_index_anchor: bool = False
|
||||||
|
# When True, host_pool uses the same logical slot indices as the anchor pool
|
||||||
|
# (e.g. DSA indexer); HostPoolGroup.free mirrors frees to this pool.
|
||||||
|
share_indices_with_anchor: bool = False
|
||||||
# Optional eviction callbacks for auto-alloc in HybridCacheController.
|
# Optional eviction callbacks for auto-alloc in HybridCacheController.
|
||||||
# host_evict_fn(n): evict n slots from the host pool (used by write()).
|
# host_evict_fn(n): evict n slots from the host pool (used by write()).
|
||||||
# device_evict_fn(n): evict n slots from the device pool (used by load()).
|
# device_evict_fn(n): evict n slots from the device pool (used by load()).
|
||||||
@@ -1619,15 +1688,54 @@ class HostPoolGroup:
|
|||||||
self.device = self.anchor_entry.host_pool.device
|
self.device = self.anchor_entry.host_pool.device
|
||||||
self.size = self.anchor_entry.host_pool.size
|
self.size = self.anchor_entry.host_pool.size
|
||||||
|
|
||||||
|
@property
|
||||||
|
def kv_buffer(self):
|
||||||
|
return self.anchor_entry.host_pool.kv_buffer
|
||||||
|
|
||||||
|
@property
|
||||||
|
def size_per_token(self):
|
||||||
|
return self.anchor_entry.host_pool.size_per_token
|
||||||
|
|
||||||
|
@property
|
||||||
|
def allocator(self):
|
||||||
|
return self.anchor_entry.host_pool.allocator
|
||||||
|
|
||||||
|
@property
|
||||||
|
def dtype(self):
|
||||||
|
return self.anchor_entry.host_pool.dtype
|
||||||
|
|
||||||
|
@property
|
||||||
|
def start_layer(self):
|
||||||
|
return self.anchor_entry.host_pool.start_layer
|
||||||
|
|
||||||
|
@property
|
||||||
|
def end_layer(self):
|
||||||
|
return self.anchor_entry.host_pool.end_layer
|
||||||
|
|
||||||
|
def get_ksize_per_token(self):
|
||||||
|
return self.anchor_entry.host_pool.get_ksize_per_token()
|
||||||
|
|
||||||
|
def get_page_buffer_meta(self, indices):
|
||||||
|
return self.anchor_entry.host_pool.get_page_buffer_meta(indices)
|
||||||
|
|
||||||
def clear(self) -> None:
|
def clear(self) -> None:
|
||||||
for entry in self.entries:
|
for entry in self.entries:
|
||||||
entry.host_pool.clear()
|
entry.host_pool.clear()
|
||||||
|
|
||||||
|
def available_size(self):
|
||||||
|
return self.anchor_entry.host_pool.available_size()
|
||||||
|
|
||||||
def alloc(self, need_size: int) -> Optional[torch.Tensor]:
|
def alloc(self, need_size: int) -> Optional[torch.Tensor]:
|
||||||
return self.anchor_entry.host_pool.alloc(need_size)
|
return self.anchor_entry.host_pool.alloc(need_size)
|
||||||
|
|
||||||
def free(self, indices: torch.Tensor) -> int:
|
def free(self, indices: torch.Tensor) -> int:
|
||||||
return self.anchor_entry.host_pool.free(indices)
|
n = self.anchor_entry.host_pool.free(indices)
|
||||||
|
for entry in self.entries:
|
||||||
|
if entry is self.anchor_entry:
|
||||||
|
continue
|
||||||
|
if getattr(entry, "share_indices_with_anchor", False):
|
||||||
|
entry.host_pool.free(indices)
|
||||||
|
return n
|
||||||
|
|
||||||
def get_data_page(self, index, flat: bool = True):
|
def get_data_page(self, index, flat: bool = True):
|
||||||
return self.anchor_entry.host_pool.get_data_page(index, flat)
|
return self.anchor_entry.host_pool.get_data_page(index, flat)
|
||||||
@@ -1703,21 +1811,31 @@ class HostPoolGroup:
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
class NSATokenToKVPoolHost(MLATokenToKVPoolHost):
|
class NSAIndexerPoolHost(HostKVCache):
|
||||||
|
"""Host-side NSA index buffers only. Slot layout matches the anchor MLA host pool."""
|
||||||
|
|
||||||
device_pool: NSATokenToKVPool
|
device_pool: NSATokenToKVPool
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
device_pool: NSATokenToKVPool,
|
device_pool: NSATokenToKVPool,
|
||||||
host_to_device_ratio: float,
|
anchor_host: MLATokenToKVPoolHost,
|
||||||
host_size: int,
|
|
||||||
page_size: int,
|
|
||||||
layout: str,
|
layout: str,
|
||||||
pin_memory: bool = True,
|
pin_memory: bool = True,
|
||||||
device: str = "cpu",
|
device: str = "cpu",
|
||||||
allocator_type: str = "default",
|
allocator_type: str = "default",
|
||||||
):
|
):
|
||||||
# Initialize indexer metadata before HostKVCache.__init__ calls get_size_per_token.
|
self.device_pool = device_pool
|
||||||
|
self.page_size = anchor_host.page_size
|
||||||
|
self.layout = layout
|
||||||
|
self.pin_memory = pin_memory
|
||||||
|
self.device = device
|
||||||
|
self.allocator = get_allocator_from_storage(allocator_type)
|
||||||
|
self.dtype = device_pool.store_dtype
|
||||||
|
self.start_layer = device_pool.start_layer
|
||||||
|
self.end_layer = device_pool.end_layer
|
||||||
|
self.layer_num = device_pool.layer_num
|
||||||
|
|
||||||
self.index_head_dim = device_pool.index_head_dim
|
self.index_head_dim = device_pool.index_head_dim
|
||||||
self.indexer_quant_block_size = device_pool.quant_block_size
|
self.indexer_quant_block_size = device_pool.quant_block_size
|
||||||
self.indexer_dtype = NSATokenToKVPool.index_k_with_scale_buffer_dtype
|
self.indexer_dtype = NSATokenToKVPool.index_k_with_scale_buffer_dtype
|
||||||
@@ -1725,35 +1843,46 @@ class NSATokenToKVPoolHost(MLATokenToKVPoolHost):
|
|||||||
self.index_head_dim
|
self.index_head_dim
|
||||||
+ self.index_head_dim // self.indexer_quant_block_size * 4
|
+ self.index_head_dim // self.indexer_quant_block_size * 4
|
||||||
)
|
)
|
||||||
super().__init__(
|
self.size = anchor_host.size
|
||||||
device_pool,
|
self.page_num = anchor_host.page_num
|
||||||
host_to_device_ratio,
|
|
||||||
host_size,
|
|
||||||
page_size,
|
|
||||||
layout,
|
|
||||||
pin_memory,
|
|
||||||
device,
|
|
||||||
allocator_type,
|
|
||||||
override_kv_cache_dim=device_pool.kv_cache_dim,
|
|
||||||
)
|
|
||||||
self.indexer_page_stride_size = (
|
self.indexer_page_stride_size = (
|
||||||
self.indexer_size_per_token * self.page_size * self.indexer_dtype.itemsize
|
self.indexer_size_per_token * self.page_size * self.indexer_dtype.itemsize
|
||||||
)
|
)
|
||||||
self.indexer_layout_dim = self.indexer_page_stride_size * self.layer_num
|
self.indexer_layout_dim = self.indexer_page_stride_size * self.layer_num
|
||||||
self.indexer_page_num = (self.size + self.page_size + 1) // self.page_size
|
self.indexer_page_num = (self.size + self.page_size + 1) // self.page_size
|
||||||
self._init_indexer_buffers()
|
self.size_per_token = (
|
||||||
logger.info(
|
self.indexer_size_per_token * self.layer_num * self.indexer_dtype.itemsize
|
||||||
f"NSATokenToKVPoolHost initialized with indexer page stride size: {self.indexer_page_stride_size}, page num: {self.indexer_page_num}"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
|
buf_elem_size = self.page_num * self.layer_num * self.indexer_page_stride_size
|
||||||
|
requested_bytes = buf_elem_size * self.indexer_dtype.itemsize
|
||||||
|
host_mem = psutil.virtual_memory()
|
||||||
|
available_bytes = host_mem.available - HICACHE_HOST_MEMORY_RESERVE_BYTES
|
||||||
|
if requested_bytes > available_bytes:
|
||||||
|
raise ValueError(
|
||||||
|
f"Not enough host memory for NSA indexer hierarchical cache. "
|
||||||
|
f"Requesting {requested_bytes / 1e9:.2f} GB but only have "
|
||||||
|
f"{available_bytes / 1e9:.2f} GB free."
|
||||||
|
)
|
||||||
|
logger.info(
|
||||||
|
"Allocating %.2f GB host memory for NSA indexer (layout=%s).",
|
||||||
|
requested_bytes / 1e9,
|
||||||
|
layout,
|
||||||
|
)
|
||||||
|
self.init_kv_buffer()
|
||||||
|
self.lock = threading.RLock()
|
||||||
|
self.clear()
|
||||||
|
|
||||||
def get_size_per_token(self):
|
def get_size_per_token(self):
|
||||||
base = super().get_size_per_token()
|
|
||||||
return (
|
return (
|
||||||
base
|
self.indexer_size_per_token * self.layer_num * self.indexer_dtype.itemsize
|
||||||
+ self.indexer_size_per_token * self.layer_num * self.indexer_dtype.itemsize
|
|
||||||
)
|
)
|
||||||
|
|
||||||
def _init_indexer_buffers(self):
|
def get_ksize_per_token(self):
|
||||||
|
return self.get_size_per_token()
|
||||||
|
|
||||||
|
def init_kv_buffer(self):
|
||||||
alloc_func = ALLOC_MEMORY_FUNCS[self.device_pool.device]
|
alloc_func = ALLOC_MEMORY_FUNCS[self.device_pool.device]
|
||||||
self.index_k_device_ptrs = torch.tensor(
|
self.index_k_device_ptrs = torch.tensor(
|
||||||
[x.data_ptr() for x in self.device_pool.index_k_with_scale_buffer],
|
[x.data_ptr() for x in self.device_pool.index_k_with_scale_buffer],
|
||||||
@@ -1761,16 +1890,13 @@ class NSATokenToKVPoolHost(MLATokenToKVPoolHost):
|
|||||||
device=self.device_pool.device,
|
device=self.device_pool.device,
|
||||||
)
|
)
|
||||||
if self.layout == "layer_first":
|
if self.layout == "layer_first":
|
||||||
self.index_k_with_scale_buffer = [
|
self.index_k_with_scale_buffer = alloc_func(
|
||||||
alloc_func(
|
(self.layer_num, self.indexer_page_num, self.indexer_page_stride_size),
|
||||||
(self.indexer_page_num, self.indexer_page_stride_size),
|
dtype=self.indexer_dtype,
|
||||||
dtype=self.indexer_dtype,
|
device=self.device,
|
||||||
device=self.device,
|
pin_memory=self.pin_memory,
|
||||||
pin_memory=self.pin_memory,
|
allocator=self.allocator,
|
||||||
allocator=self.allocator,
|
)
|
||||||
)
|
|
||||||
for _ in range(self.layer_num)
|
|
||||||
]
|
|
||||||
self.index_k_data_refs = [
|
self.index_k_data_refs = [
|
||||||
self.index_k_with_scale_buffer[i] for i in range(self.layer_num)
|
self.index_k_with_scale_buffer[i] for i in range(self.layer_num)
|
||||||
]
|
]
|
||||||
@@ -1795,6 +1921,9 @@ class NSATokenToKVPoolHost(MLATokenToKVPoolHost):
|
|||||||
else:
|
else:
|
||||||
raise ValueError(f"Unsupported layout: {self.layout}")
|
raise ValueError(f"Unsupported layout: {self.layout}")
|
||||||
|
|
||||||
|
def get_hybrid_pool_buffer(self):
|
||||||
|
return [self.index_k_with_scale_buffer]
|
||||||
|
|
||||||
def _get_indexer_page_indices(self, host_indices, device_indices):
|
def _get_indexer_page_indices(self, host_indices, device_indices):
|
||||||
if host_indices.numel() == 0:
|
if host_indices.numel() == 0:
|
||||||
return host_indices, device_indices
|
return host_indices, device_indices
|
||||||
@@ -1810,7 +1939,7 @@ class NSATokenToKVPoolHost(MLATokenToKVPoolHost):
|
|||||||
)
|
)
|
||||||
return host_page_indices, device_page_indices
|
return host_page_indices, device_page_indices
|
||||||
|
|
||||||
def _load_indexer_to_device_per_layer(
|
def load_to_device_per_layer(
|
||||||
self, device_pool, host_indices, device_indices, layer_id, io_backend
|
self, device_pool, host_indices, device_indices, layer_id, io_backend
|
||||||
):
|
):
|
||||||
host_page_indices, device_page_indices = self._get_indexer_page_indices(
|
host_page_indices, device_page_indices = self._get_indexer_page_indices(
|
||||||
@@ -1861,7 +1990,7 @@ class NSATokenToKVPoolHost(MLATokenToKVPoolHost):
|
|||||||
else:
|
else:
|
||||||
raise ValueError(f"Unsupported IO backend: {io_backend}")
|
raise ValueError(f"Unsupported IO backend: {io_backend}")
|
||||||
|
|
||||||
def _backup_indexer_from_device_all_layer(
|
def backup_from_device_all_layer(
|
||||||
self, device_pool, host_indices, device_indices, io_backend
|
self, device_pool, host_indices, device_indices, io_backend
|
||||||
):
|
):
|
||||||
host_page_indices, device_page_indices = self._get_indexer_page_indices(
|
host_page_indices, device_page_indices = self._get_indexer_page_indices(
|
||||||
@@ -1894,7 +2023,7 @@ class NSATokenToKVPoolHost(MLATokenToKVPoolHost):
|
|||||||
if self.layout == "layer_first":
|
if self.layout == "layer_first":
|
||||||
transfer_kv_direct(
|
transfer_kv_direct(
|
||||||
src_layers=device_pool.index_k_with_scale_buffer,
|
src_layers=device_pool.index_k_with_scale_buffer,
|
||||||
dst_layers=self.index_k_with_scale_buffer,
|
dst_layers=self.index_k_data_refs,
|
||||||
src_indices=device_page_indices,
|
src_indices=device_page_indices,
|
||||||
dst_indices=host_page_indices,
|
dst_indices=host_page_indices,
|
||||||
page_size=1,
|
page_size=1,
|
||||||
@@ -1912,27 +2041,60 @@ class NSATokenToKVPoolHost(MLATokenToKVPoolHost):
|
|||||||
else:
|
else:
|
||||||
raise ValueError(f"Unsupported IO backend: {io_backend}")
|
raise ValueError(f"Unsupported IO backend: {io_backend}")
|
||||||
|
|
||||||
def load_to_device_per_layer(
|
def get_data_page(self, index, flat: bool = True) -> torch.Tensor:
|
||||||
self,
|
page_idx = int(index) // self.page_size
|
||||||
device_pool,
|
if self.layout == "layer_first":
|
||||||
host_indices,
|
data_page = self.index_k_with_scale_buffer[:, page_idx : page_idx + 1, :]
|
||||||
device_indices,
|
elif self.layout in ["page_first", "page_first_direct"]:
|
||||||
layer_id,
|
data_page = self.index_k_with_scale_buffer[page_idx : page_idx + 1, :, :, :]
|
||||||
io_backend,
|
else:
|
||||||
):
|
raise ValueError(f"Unsupported layout: {self.layout}")
|
||||||
super().load_to_device_per_layer(
|
if flat:
|
||||||
device_pool, host_indices, device_indices, layer_id, io_backend
|
data_page = data_page.flatten()
|
||||||
)
|
return data_page
|
||||||
self._load_indexer_to_device_per_layer(
|
|
||||||
device_pool, host_indices, device_indices, layer_id, io_backend
|
|
||||||
)
|
|
||||||
|
|
||||||
def backup_from_device_all_layer(
|
def get_dummy_flat_data_page(self) -> torch.Tensor:
|
||||||
self, device_pool, host_indices, device_indices, io_backend
|
return torch.zeros(
|
||||||
):
|
(self.layer_num, self.indexer_page_stride_size),
|
||||||
super().backup_from_device_all_layer(
|
dtype=self.indexer_dtype,
|
||||||
device_pool, host_indices, device_indices, io_backend
|
device=self.device,
|
||||||
)
|
pin_memory=self.pin_memory,
|
||||||
self._backup_indexer_from_device_all_layer(
|
).flatten()
|
||||||
device_pool, host_indices, device_indices, io_backend
|
|
||||||
|
def set_from_flat_data_page(self, index: int, data_page: torch.Tensor) -> None:
|
||||||
|
page_idx = int(index) // self.page_size
|
||||||
|
if self.layout == "layer_first":
|
||||||
|
self.index_k_with_scale_buffer[:, page_idx : page_idx + 1, :] = (
|
||||||
|
data_page.reshape(
|
||||||
|
self.layer_num,
|
||||||
|
1,
|
||||||
|
self.indexer_page_stride_size,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
elif self.layout in ["page_first", "page_first_direct"]:
|
||||||
|
self.index_k_with_scale_buffer[page_idx : page_idx + 1, :, :, :] = (
|
||||||
|
data_page.reshape(
|
||||||
|
1,
|
||||||
|
self.layer_num,
|
||||||
|
1,
|
||||||
|
self.indexer_page_stride_size,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
raise ValueError(f"Unsupported layout: {self.layout}")
|
||||||
|
|
||||||
|
def get_page_buffer_meta(self, indices):
|
||||||
|
"""Meta data for zero-copy storage I/O."""
|
||||||
|
assert len(indices) % self.page_size == 0
|
||||||
|
if self.layout not in ["page_first", "page_first_direct"]:
|
||||||
|
raise ValueError(f"Unsupported layout: {self.layout}")
|
||||||
|
ptr_list = []
|
||||||
|
indices = indices.tolist()
|
||||||
|
page_stride_bytes = (
|
||||||
|
self.layer_num * self.indexer_page_stride_size * self.indexer_dtype.itemsize
|
||||||
)
|
)
|
||||||
|
base_ptr = self.index_k_with_scale_buffer.data_ptr()
|
||||||
|
for i in range(0, len(indices), self.page_size):
|
||||||
|
page_index = int(indices[i]) // self.page_size
|
||||||
|
ptr_list.append(base_ptr + page_index * page_stride_bytes)
|
||||||
|
return ptr_list, [page_stride_bytes] * len(ptr_list)
|
||||||
|
|||||||
@@ -5,7 +5,7 @@ import os
|
|||||||
import time
|
import time
|
||||||
import uuid
|
import uuid
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from typing import Any, List, Optional
|
from typing import Any, List, Optional, Tuple
|
||||||
|
|
||||||
import requests
|
import requests
|
||||||
import torch
|
import torch
|
||||||
@@ -15,6 +15,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, HostTensorAllocator
|
from sglang.srt.mem_cache.memory_pool_host import HostKVCache, HostTensorAllocator
|
||||||
from sglang.srt.observability.metrics_collector import StorageMetrics
|
from sglang.srt.observability.metrics_collector import StorageMetrics
|
||||||
@@ -432,6 +436,8 @@ class MooncakeStore(HiCacheStorage, MooncakeBaseStore):
|
|||||||
else:
|
else:
|
||||||
self.mha_suffix = [f"{rank}" for rank in target_ranks]
|
self.mha_suffix = [f"{rank}" for rank in target_ranks]
|
||||||
|
|
||||||
|
self.registered_pools = {}
|
||||||
|
|
||||||
self.gb_per_page = None
|
self.gb_per_page = None
|
||||||
self.prefetch_pgs = []
|
self.prefetch_pgs = []
|
||||||
self.backup_pgs = []
|
self.backup_pgs = []
|
||||||
@@ -502,6 +508,154 @@ class MooncakeStore(HiCacheStorage, MooncakeBaseStore):
|
|||||||
bytes_per_page = mem_pool_host.get_ksize_per_token() * mem_pool_host.page_size
|
bytes_per_page = mem_pool_host.get_ksize_per_token() * mem_pool_host.page_size
|
||||||
self.gb_per_page = bytes_per_page / (1 << 30)
|
self.gb_per_page = bytes_per_page / (1 << 30)
|
||||||
|
|
||||||
|
def register_mem_host_pool_v2(self, host_pool: HostKVCache, host_pool_name):
|
||||||
|
# KV anchor memory is already registered via register_mem_pool_host().
|
||||||
|
# v2 here only registers additional hybrid pools.
|
||||||
|
if host_pool_name == PoolName.KV:
|
||||||
|
return
|
||||||
|
# Keep a name->pool mapping so batch v2 can resolve PoolTransfer.name to
|
||||||
|
# the corresponding host pool implementation at runtime.
|
||||||
|
self.registered_pools[host_pool_name] = host_pool
|
||||||
|
|
||||||
|
# Hybrid pools expose the tensors that Mooncake needs for zero-copy I/O.
|
||||||
|
# The storage backend only depends on this accessor, not concrete fields.
|
||||||
|
buf_list = host_pool.get_hybrid_pool_buffer()
|
||||||
|
for buf in buf_list:
|
||||||
|
super().register_buffer(buf)
|
||||||
|
|
||||||
|
def _tag_keys(self, keys: List[str]) -> List[str]:
|
||||||
|
if self.extra_backend_tag is None:
|
||||||
|
return keys
|
||||||
|
return [f"{ self.extra_backend_tag}_{key}" for key in keys]
|
||||||
|
|
||||||
|
def _get_hybrid_page_component_keys(
|
||||||
|
self, page_keys: List[str], transfer: PoolTransfer
|
||||||
|
) -> Tuple[List[str], int]:
|
||||||
|
# A logical "page" may map to multiple physical objects in storage.
|
||||||
|
# - INDEXER: one key per page
|
||||||
|
# - MAMBA : one temporal key + N conv keys per page
|
||||||
|
# key_multiplier records how many component keys are generated per page.
|
||||||
|
name = transfer.name
|
||||||
|
suffixes = []
|
||||||
|
if name == PoolName.INDEXER:
|
||||||
|
suffixes = [f"_{self.mla_suffix}_{PoolName.INDEXER}"]
|
||||||
|
elif name == PoolName.MAMBA:
|
||||||
|
pools = getattr(self, "registered_pools", {})
|
||||||
|
mamba_pool = pools.get(PoolName.MAMBA)
|
||||||
|
conv_num = len(getattr(mamba_pool, "conv_buffer", None) or [])
|
||||||
|
base_suffix = f"_{self.mha_suffix}"
|
||||||
|
suffixes = [f"{base_suffix}_temporal"] + [
|
||||||
|
f"{base_suffix}_conv_{i}" for i in range(conv_num)
|
||||||
|
]
|
||||||
|
key_multiplier = len(suffixes)
|
||||||
|
component_keys = [
|
||||||
|
f"{page_key}{suffix}" for page_key in page_keys for suffix in suffixes
|
||||||
|
]
|
||||||
|
return component_keys, key_multiplier
|
||||||
|
|
||||||
|
def batch_exists_v2(
|
||||||
|
self,
|
||||||
|
keys: List[str],
|
||||||
|
pool_transfers: Optional[List[PoolTransfer]] = None,
|
||||||
|
extra_info: Optional[HiCacheStorageExtraInfo] = None,
|
||||||
|
) -> PoolTransferResult:
|
||||||
|
qkeys = self._tag_keys(keys)
|
||||||
|
kv_pages = self.batch_exists(qkeys, 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
|
||||||
|
component_keys, key_multiplier = self._get_hybrid_page_component_keys(
|
||||||
|
qkeys, transfer
|
||||||
|
)
|
||||||
|
ex = self._batch_exist(component_keys)
|
||||||
|
if key_multiplier > 0:
|
||||||
|
page_exists = [
|
||||||
|
all(
|
||||||
|
r == 1
|
||||||
|
for r in ex[i * key_multiplier : (i + 1) * key_multiplier]
|
||||||
|
)
|
||||||
|
for i in range(kv_pages)
|
||||||
|
]
|
||||||
|
else:
|
||||||
|
page_exists = [False] * kv_pages
|
||||||
|
boundary = 0
|
||||||
|
if transfer.hit_policy == PoolHitPolicy.ALL_PAGES:
|
||||||
|
try:
|
||||||
|
boundary = page_exists.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(
|
||||||
|
page_exists[i]
|
||||||
|
for i in range(max(0, prefix_len - trailing), prefix_len)
|
||||||
|
):
|
||||||
|
boundary = prefix_len
|
||||||
|
break
|
||||||
|
if boundary:
|
||||||
|
hit_count[transfer.name] = boundary
|
||||||
|
final_pages = min(final_pages, boundary)
|
||||||
|
|
||||||
|
return PoolTransferResult(final_pages, hit_count)
|
||||||
|
|
||||||
|
def _batch_io_v2(self, transfers: List[PoolTransfer], is_set: bool):
|
||||||
|
# Unified v2 I/O path: each PoolTransfer can expand to one or more
|
||||||
|
# storage objects per logical page, but API still reports page-level result.
|
||||||
|
results: dict = {}
|
||||||
|
for transfer in transfers:
|
||||||
|
host_pool = getattr(self, "registered_pools", {}).get(transfer.name)
|
||||||
|
keys = transfer.keys
|
||||||
|
page_size = getattr(host_pool, "page_size", 1) or 1
|
||||||
|
host_indices = transfer.host_indices
|
||||||
|
assert len(keys) > 0
|
||||||
|
assert len(keys) == len(host_indices) // page_size
|
||||||
|
|
||||||
|
ptr_list, element_size_list = host_pool.get_page_buffer_meta(host_indices)
|
||||||
|
key_strs, key_multiplier = self._get_hybrid_page_component_keys(
|
||||||
|
keys, transfer
|
||||||
|
)
|
||||||
|
key_strs = self._tag_keys(key_strs)
|
||||||
|
|
||||||
|
if is_set:
|
||||||
|
exist_result = self._batch_exist(key_strs)
|
||||||
|
io_results = [0 if state == 1 else -1 for state in exist_result]
|
||||||
|
missing_idx = [i for i, state in enumerate(exist_result) if state != 1]
|
||||||
|
if missing_idx:
|
||||||
|
put_results = self._put_batch_zero_copy_impl(
|
||||||
|
[key_strs[i] for i in missing_idx],
|
||||||
|
[ptr_list[i] for i in missing_idx],
|
||||||
|
[element_size_list[i] for i in missing_idx],
|
||||||
|
)
|
||||||
|
for i, res in zip(missing_idx, put_results):
|
||||||
|
io_results[i] = res
|
||||||
|
else:
|
||||||
|
io_results = self._get_batch_zero_copy_impl(
|
||||||
|
key_strs, ptr_list, element_size_list
|
||||||
|
)
|
||||||
|
results[transfer.name] = self._batch_postprocess(
|
||||||
|
io_results, is_set_operate=is_set, key_multiplier=key_multiplier
|
||||||
|
)
|
||||||
|
return results
|
||||||
|
|
||||||
|
def batch_get_v2(
|
||||||
|
self,
|
||||||
|
transfers: List[PoolTransfer],
|
||||||
|
extra_info: Optional[HiCacheStorageExtraInfo] = None,
|
||||||
|
) -> dict:
|
||||||
|
return self._batch_io_v2(transfers, is_set=False)
|
||||||
|
|
||||||
|
def batch_set_v2(
|
||||||
|
self,
|
||||||
|
transfers: List[PoolTransfer],
|
||||||
|
extra_info: Optional[HiCacheStorageExtraInfo] = None,
|
||||||
|
) -> dict:
|
||||||
|
return self._batch_io_v2(transfers, is_set=True)
|
||||||
|
|
||||||
def _get_mha_split_heads_buffer_meta(self, keys, indices):
|
def _get_mha_split_heads_buffer_meta(self, keys, indices):
|
||||||
ptr_list, element_size_list = (
|
ptr_list, element_size_list = (
|
||||||
self.mem_pool_host.get_split_heads_page_buffer_meta(
|
self.mem_pool_host.get_split_heads_page_buffer_meta(
|
||||||
@@ -544,7 +698,9 @@ class MooncakeStore(HiCacheStorage, MooncakeBaseStore):
|
|||||||
else:
|
else:
|
||||||
return self._get_mha_buffer_meta(keys, host_indices)
|
return self._get_mha_buffer_meta(keys, host_indices)
|
||||||
|
|
||||||
def _batch_postprocess(self, results: List[int], is_set_operate=False):
|
def _batch_postprocess(
|
||||||
|
self, results: List[int], is_set_operate=False, key_multiplier=None
|
||||||
|
):
|
||||||
"""
|
"""
|
||||||
refer to https://github.com/kvcache-ai/Mooncake/blob/main/mooncake-store/include/pybind_client.h
|
refer to https://github.com/kvcache-ai/Mooncake/blob/main/mooncake-store/include/pybind_client.h
|
||||||
for batch_get_into, results is Vector of integers,
|
for batch_get_into, results is Vector of integers,
|
||||||
@@ -552,32 +708,26 @@ class MooncakeStore(HiCacheStorage, MooncakeBaseStore):
|
|||||||
for batch_put_from, results is Vector of integers,
|
for batch_put_from, results is Vector of integers,
|
||||||
where each element is 0 on success, or a negative value on error
|
where each element is 0 on success, or a negative value on error
|
||||||
"""
|
"""
|
||||||
if self.is_mla_backend:
|
if key_multiplier is None:
|
||||||
return [k_res == 0 if is_set_operate else k_res > 0 for k_res in results]
|
if self.is_mla_backend:
|
||||||
else:
|
key_multiplier = 1
|
||||||
if self.storage_config.should_split_heads:
|
|
||||||
kv_groups = [
|
|
||||||
results[i : i + self.split_factor * 2]
|
|
||||||
for i in range(0, len(results), self.split_factor * 2)
|
|
||||||
]
|
|
||||||
return [
|
|
||||||
(
|
|
||||||
all(res == 0 for res in kv_group)
|
|
||||||
if is_set_operate
|
|
||||||
else all(res > 0 for res in kv_group)
|
|
||||||
)
|
|
||||||
for kv_group in kv_groups
|
|
||||||
]
|
|
||||||
else:
|
else:
|
||||||
kv_pairs = zip(results[::2], results[1::2])
|
key_multiplier = 2
|
||||||
return [
|
if self.storage_config.should_split_heads:
|
||||||
(
|
key_multiplier *= self.split_factor
|
||||||
(k_res == 0 and v_res == 0)
|
|
||||||
if is_set_operate
|
result_groups = [
|
||||||
else (k_res > 0 and v_res > 0)
|
results[i : i + key_multiplier]
|
||||||
)
|
for i in range(0, len(results), key_multiplier)
|
||||||
for k_res, v_res in kv_pairs
|
]
|
||||||
]
|
return [
|
||||||
|
(
|
||||||
|
all(res == 0 for res in group)
|
||||||
|
if is_set_operate
|
||||||
|
else all(res > 0 for res in group)
|
||||||
|
)
|
||||||
|
for group in result_groups
|
||||||
|
]
|
||||||
|
|
||||||
def batch_get_v1(
|
def batch_get_v1(
|
||||||
self,
|
self,
|
||||||
@@ -586,9 +736,7 @@ class MooncakeStore(HiCacheStorage, MooncakeBaseStore):
|
|||||||
extra_info: Optional[HiCacheStorageExtraInfo] = None,
|
extra_info: Optional[HiCacheStorageExtraInfo] = None,
|
||||||
) -> List[bool]:
|
) -> List[bool]:
|
||||||
# Apply extra_backend_tag prefix if available
|
# Apply extra_backend_tag prefix if available
|
||||||
if self.extra_backend_tag is not None:
|
keys = self._tag_keys(keys)
|
||||||
prefix = self.extra_backend_tag
|
|
||||||
keys = [f"{prefix}_{key}" for key in keys]
|
|
||||||
|
|
||||||
key_strs, buffer_ptrs, buffer_sizes = self._batch_preprocess(keys, host_indices)
|
key_strs, buffer_ptrs, buffer_sizes = self._batch_preprocess(keys, host_indices)
|
||||||
|
|
||||||
@@ -613,9 +761,7 @@ class MooncakeStore(HiCacheStorage, MooncakeBaseStore):
|
|||||||
extra_info: Optional[HiCacheStorageExtraInfo] = None,
|
extra_info: Optional[HiCacheStorageExtraInfo] = None,
|
||||||
) -> List[bool]:
|
) -> List[bool]:
|
||||||
# Apply extra_backend_tag prefix if available
|
# Apply extra_backend_tag prefix if available
|
||||||
if self.extra_backend_tag is not None:
|
keys = self._tag_keys(keys)
|
||||||
prefix = self.extra_backend_tag
|
|
||||||
keys = [f"{prefix}_{key}" for key in keys]
|
|
||||||
|
|
||||||
key_strs, buffer_ptrs, buffer_sizes = self._batch_preprocess(keys, host_indices)
|
key_strs, buffer_ptrs, buffer_sizes = self._batch_preprocess(keys, host_indices)
|
||||||
exist_result = self._batch_exist(key_strs)
|
exist_result = self._batch_exist(key_strs)
|
||||||
@@ -780,9 +926,7 @@ class MooncakeStore(HiCacheStorage, MooncakeBaseStore):
|
|||||||
self, keys, extra_info: Optional[HiCacheStorageExtraInfo] = None
|
self, keys, extra_info: Optional[HiCacheStorageExtraInfo] = None
|
||||||
) -> int:
|
) -> int:
|
||||||
# Apply extra_backend_tag prefix if available
|
# Apply extra_backend_tag prefix if available
|
||||||
if self.extra_backend_tag is not None:
|
keys = self._tag_keys(keys)
|
||||||
prefix = self.extra_backend_tag
|
|
||||||
keys = [f"{prefix}_{key}" for key in keys]
|
|
||||||
|
|
||||||
if self.is_mla_backend:
|
if self.is_mla_backend:
|
||||||
query_keys = [f"{key}_{self.mla_suffix}_k" for key in keys]
|
query_keys = [f"{key}_{self.mla_suffix}_k" for key in keys]
|
||||||
|
|||||||
@@ -5,7 +5,8 @@ import torch
|
|||||||
from sglang.srt.mem_cache.memory_pool import NSATokenToKVPool
|
from sglang.srt.mem_cache.memory_pool import NSATokenToKVPool
|
||||||
from sglang.srt.mem_cache.memory_pool_host import (
|
from sglang.srt.mem_cache.memory_pool_host import (
|
||||||
ALLOC_MEMORY_FUNCS,
|
ALLOC_MEMORY_FUNCS,
|
||||||
NSATokenToKVPoolHost,
|
MLATokenToKVPoolHost,
|
||||||
|
NSAIndexerPoolHost,
|
||||||
alloc_with_pin_memory,
|
alloc_with_pin_memory,
|
||||||
)
|
)
|
||||||
from sglang.srt.utils import is_cuda, is_hip, is_npu, is_xpu
|
from sglang.srt.utils import is_cuda, is_hip, is_npu, is_xpu
|
||||||
@@ -58,7 +59,7 @@ class TestNSAHiCacheTransfer(unittest.TestCase):
|
|||||||
if pin_memory:
|
if pin_memory:
|
||||||
ALLOC_MEMORY_FUNCS["cuda"] = alloc_with_pin_memory
|
ALLOC_MEMORY_FUNCS["cuda"] = alloc_with_pin_memory
|
||||||
try:
|
try:
|
||||||
host_pool = NSATokenToKVPoolHost(
|
mla_host = MLATokenToKVPoolHost(
|
||||||
device_pool=device_pool,
|
device_pool=device_pool,
|
||||||
host_to_device_ratio=2.0,
|
host_to_device_ratio=2.0,
|
||||||
host_size=0,
|
host_size=0,
|
||||||
@@ -66,6 +67,16 @@ class TestNSAHiCacheTransfer(unittest.TestCase):
|
|||||||
layout="layer_first",
|
layout="layer_first",
|
||||||
pin_memory=pin_memory,
|
pin_memory=pin_memory,
|
||||||
device="cpu",
|
device="cpu",
|
||||||
|
allocator_type="default",
|
||||||
|
override_kv_cache_dim=device_pool.kv_cache_dim,
|
||||||
|
)
|
||||||
|
indexer_host = NSAIndexerPoolHost(
|
||||||
|
device_pool=device_pool,
|
||||||
|
anchor_host=mla_host,
|
||||||
|
layout="layer_first",
|
||||||
|
pin_memory=pin_memory,
|
||||||
|
device="cpu",
|
||||||
|
allocator_type="default",
|
||||||
)
|
)
|
||||||
finally:
|
finally:
|
||||||
ALLOC_MEMORY_FUNCS["cuda"] = original_alloc
|
ALLOC_MEMORY_FUNCS["cuda"] = original_alloc
|
||||||
@@ -97,7 +108,10 @@ class TestNSAHiCacheTransfer(unittest.TestCase):
|
|||||||
device="cuda" if io_backend == "kernel" else "cpu",
|
device="cuda" if io_backend == "kernel" else "cpu",
|
||||||
)
|
)
|
||||||
|
|
||||||
host_pool.backup_from_device_all_layer(
|
mla_host.backup_from_device_all_layer(
|
||||||
|
device_pool, host_indices, device_indices, io_backend
|
||||||
|
)
|
||||||
|
indexer_host.backup_from_device_all_layer(
|
||||||
device_pool, host_indices, device_indices, io_backend
|
device_pool, host_indices, device_indices, io_backend
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -105,14 +119,14 @@ class TestNSAHiCacheTransfer(unittest.TestCase):
|
|||||||
for host_page, device_page in zip(
|
for host_page, device_page in zip(
|
||||||
host_pages.tolist(), device_pages.tolist()
|
host_pages.tolist(), device_pages.tolist()
|
||||||
):
|
):
|
||||||
got = host_pool.index_k_with_scale_buffer[layer_id][host_page].cpu()
|
got = indexer_host.index_k_with_scale_buffer[layer_id][host_page].cpu()
|
||||||
expected = device_pool.index_k_with_scale_buffer[layer_id][
|
expected = device_pool.index_k_with_scale_buffer[layer_id][
|
||||||
device_page
|
device_page
|
||||||
].cpu()
|
].cpu()
|
||||||
self.assertTrue(torch.equal(got, expected))
|
self.assertTrue(torch.equal(got, expected))
|
||||||
host_start = host_page * page_size
|
host_start = host_page * page_size
|
||||||
device_start = device_page * page_size
|
device_start = device_page * page_size
|
||||||
got_kv = host_pool.kv_buffer[layer_id][
|
got_kv = mla_host.kv_buffer[layer_id][
|
||||||
host_start : host_start + page_size
|
host_start : host_start + page_size
|
||||||
].cpu()
|
].cpu()
|
||||||
expected_kv = device_pool.kv_buffer[layer_id][
|
expected_kv = device_pool.kv_buffer[layer_id][
|
||||||
|
|||||||
Reference in New Issue
Block a user