|
|
|
@@ -109,12 +109,7 @@ def build_kv_only_stack(
|
|
|
|
|
server_args: ServerArgs,
|
|
|
|
|
kv_pool: Any,
|
|
|
|
|
full_layer_mapping: dict[int, int],
|
|
|
|
|
page_size: int,
|
|
|
|
|
tp_group,
|
|
|
|
|
load_cache_event,
|
|
|
|
|
attn_cp_group: Optional[torch.distributed.ProcessGroup] = None,
|
|
|
|
|
attn_tp_group: Optional[torch.distributed.ProcessGroup] = None,
|
|
|
|
|
pp_group: Optional[torch.distributed.ProcessGroup] = None,
|
|
|
|
|
storage_backend: Optional[str],
|
|
|
|
|
use_mla: bool,
|
|
|
|
|
override_kv_cache_dim: Optional[int] = None,
|
|
|
|
@@ -126,7 +121,7 @@ def build_kv_only_stack(
|
|
|
|
|
transfer_layer_num = len(full_layer_mapping)
|
|
|
|
|
kv_host_pool = build_kv_host_pool(
|
|
|
|
|
kv_pool=kv_pool,
|
|
|
|
|
page_size=page_size,
|
|
|
|
|
page_size=params.page_size,
|
|
|
|
|
server_args=server_args,
|
|
|
|
|
use_mla=use_mla,
|
|
|
|
|
override_kv_cache_dim=override_kv_cache_dim,
|
|
|
|
@@ -145,12 +140,12 @@ def build_kv_only_stack(
|
|
|
|
|
cache_controller = HybridCacheController(
|
|
|
|
|
params.token_to_kv_pool_allocator,
|
|
|
|
|
host_pool_group,
|
|
|
|
|
page_size,
|
|
|
|
|
tp_group,
|
|
|
|
|
params.page_size,
|
|
|
|
|
params.tp_cache_group,
|
|
|
|
|
load_cache_event=load_cache_event,
|
|
|
|
|
attn_cp_group=attn_cp_group,
|
|
|
|
|
attn_tp_group=attn_tp_group,
|
|
|
|
|
pp_group=pp_group,
|
|
|
|
|
attn_cp_group=params.attn_cp_cache_group,
|
|
|
|
|
attn_tp_group=params.attn_tp_cache_group,
|
|
|
|
|
pp_group=params.pp_cache_group,
|
|
|
|
|
write_policy=server_args.hicache_write_policy,
|
|
|
|
|
io_backend=server_args.hicache_io_backend,
|
|
|
|
|
storage_backend=storage_backend,
|
|
|
|
@@ -171,12 +166,7 @@ def build_hybrid_swa_stack(
|
|
|
|
|
swa_kv_pool: Any,
|
|
|
|
|
full_layer_mapping: dict[int, int],
|
|
|
|
|
swa_layer_mapping: dict[int, int],
|
|
|
|
|
page_size: int,
|
|
|
|
|
tp_group,
|
|
|
|
|
load_cache_event,
|
|
|
|
|
attn_cp_group: Optional[torch.distributed.ProcessGroup] = None,
|
|
|
|
|
attn_tp_group: Optional[torch.distributed.ProcessGroup] = None,
|
|
|
|
|
pp_group: Optional[torch.distributed.ProcessGroup] = None,
|
|
|
|
|
storage_backend: Optional[str],
|
|
|
|
|
use_mla: bool,
|
|
|
|
|
host_swa_evict_fn: Optional[Callable[[int], Any]] = None,
|
|
|
|
@@ -189,13 +179,13 @@ def build_hybrid_swa_stack(
|
|
|
|
|
transfer_layer_num = len(full_layer_mapping | swa_layer_mapping)
|
|
|
|
|
kv_host_pool = build_kv_host_pool(
|
|
|
|
|
kv_pool=full_kv_pool,
|
|
|
|
|
page_size=page_size,
|
|
|
|
|
page_size=params.page_size,
|
|
|
|
|
server_args=server_args,
|
|
|
|
|
use_mla=use_mla,
|
|
|
|
|
)
|
|
|
|
|
swa_host_pool = build_kv_host_pool(
|
|
|
|
|
kv_pool=swa_kv_pool,
|
|
|
|
|
page_size=page_size,
|
|
|
|
|
page_size=params.page_size,
|
|
|
|
|
server_args=server_args,
|
|
|
|
|
use_mla=use_mla,
|
|
|
|
|
)
|
|
|
|
@@ -227,12 +217,12 @@ def build_hybrid_swa_stack(
|
|
|
|
|
cache_controller = HybridCacheController(
|
|
|
|
|
params.token_to_kv_pool_allocator,
|
|
|
|
|
host_pool_group,
|
|
|
|
|
page_size,
|
|
|
|
|
tp_group,
|
|
|
|
|
params.page_size,
|
|
|
|
|
params.tp_cache_group,
|
|
|
|
|
load_cache_event=load_cache_event,
|
|
|
|
|
attn_cp_group=attn_cp_group,
|
|
|
|
|
attn_tp_group=attn_tp_group,
|
|
|
|
|
pp_group=pp_group,
|
|
|
|
|
attn_cp_group=params.attn_cp_cache_group,
|
|
|
|
|
attn_tp_group=params.attn_tp_cache_group,
|
|
|
|
|
pp_group=params.pp_cache_group,
|
|
|
|
|
write_policy=server_args.hicache_write_policy,
|
|
|
|
|
io_backend=server_args.hicache_io_backend,
|
|
|
|
|
storage_backend=storage_backend,
|
|
|
|
@@ -286,12 +276,7 @@ def build_deepseek_v4_hicache_stack(
|
|
|
|
|
params: CacheInitParams,
|
|
|
|
|
server_args: ServerArgs,
|
|
|
|
|
kvcache: Any,
|
|
|
|
|
page_size: int,
|
|
|
|
|
tp_group,
|
|
|
|
|
load_cache_event,
|
|
|
|
|
attn_cp_group: Optional[torch.distributed.ProcessGroup] = None,
|
|
|
|
|
attn_tp_group: Optional[torch.distributed.ProcessGroup] = None,
|
|
|
|
|
pp_group: Optional[torch.distributed.ProcessGroup] = None,
|
|
|
|
|
storage_backend: Optional[str],
|
|
|
|
|
host_swa_evict_fn: Optional[Callable[[int], Any]] = None,
|
|
|
|
|
device_swa_evict_fn: Optional[Callable[[int], Any]] = None,
|
|
|
|
@@ -300,6 +285,7 @@ def build_deepseek_v4_hicache_stack(
|
|
|
|
|
storage_backend_extra_config: Optional[dict] = None,
|
|
|
|
|
enable_storage_metrics: bool = False,
|
|
|
|
|
) -> tuple[HostPoolGroup, HybridCacheController]:
|
|
|
|
|
page_size = params.page_size
|
|
|
|
|
transfer_layer_num = kvcache.end_layer - kvcache.start_layer
|
|
|
|
|
full_layer_mapping = {layer_id: layer_id for layer_id in range(transfer_layer_num)}
|
|
|
|
|
|
|
|
|
@@ -500,11 +486,11 @@ def build_deepseek_v4_hicache_stack(
|
|
|
|
|
params.token_to_kv_pool_allocator,
|
|
|
|
|
host_pool_group,
|
|
|
|
|
page_size,
|
|
|
|
|
tp_group,
|
|
|
|
|
params.tp_cache_group,
|
|
|
|
|
load_cache_event=load_cache_event,
|
|
|
|
|
attn_cp_group=attn_cp_group,
|
|
|
|
|
attn_tp_group=attn_tp_group,
|
|
|
|
|
pp_group=pp_group,
|
|
|
|
|
attn_cp_group=params.attn_cp_cache_group,
|
|
|
|
|
attn_tp_group=params.attn_tp_cache_group,
|
|
|
|
|
pp_group=params.pp_cache_group,
|
|
|
|
|
write_policy=server_args.hicache_write_policy,
|
|
|
|
|
io_backend=server_args.hicache_io_backend,
|
|
|
|
|
storage_backend=storage_backend,
|
|
|
|
@@ -525,12 +511,7 @@ def build_hybrid_mamba_stack(
|
|
|
|
|
mamba_pool: Any,
|
|
|
|
|
full_layer_mapping: dict[int, int],
|
|
|
|
|
mamba_layer_mapping: dict[int, int],
|
|
|
|
|
page_size: int,
|
|
|
|
|
tp_group,
|
|
|
|
|
load_cache_event,
|
|
|
|
|
attn_cp_group: Optional[torch.distributed.ProcessGroup] = None,
|
|
|
|
|
attn_tp_group: Optional[torch.distributed.ProcessGroup] = None,
|
|
|
|
|
pp_group: Optional[torch.distributed.ProcessGroup] = None,
|
|
|
|
|
storage_backend: Optional[str],
|
|
|
|
|
use_mla: bool,
|
|
|
|
|
host_mamba_evict_fn: Optional[Callable[[int], Any]] = None,
|
|
|
|
@@ -544,7 +525,7 @@ def build_hybrid_mamba_stack(
|
|
|
|
|
mamba_allocator = params.req_to_token_pool.mamba_allocator
|
|
|
|
|
kv_host_pool = build_kv_host_pool(
|
|
|
|
|
kv_pool=kv_pool,
|
|
|
|
|
page_size=page_size,
|
|
|
|
|
page_size=params.page_size,
|
|
|
|
|
server_args=server_args,
|
|
|
|
|
use_mla=use_mla,
|
|
|
|
|
)
|
|
|
|
@@ -580,12 +561,12 @@ def build_hybrid_mamba_stack(
|
|
|
|
|
cache_controller = HybridCacheController(
|
|
|
|
|
params.token_to_kv_pool_allocator,
|
|
|
|
|
host_pool_group,
|
|
|
|
|
page_size,
|
|
|
|
|
tp_group,
|
|
|
|
|
params.page_size,
|
|
|
|
|
params.tp_cache_group,
|
|
|
|
|
load_cache_event=load_cache_event,
|
|
|
|
|
attn_cp_group=attn_cp_group,
|
|
|
|
|
attn_tp_group=attn_tp_group,
|
|
|
|
|
pp_group=pp_group,
|
|
|
|
|
attn_cp_group=params.attn_cp_cache_group,
|
|
|
|
|
attn_tp_group=params.attn_tp_cache_group,
|
|
|
|
|
pp_group=params.pp_cache_group,
|
|
|
|
|
write_policy=server_args.hicache_write_policy,
|
|
|
|
|
io_backend=server_args.hicache_io_backend,
|
|
|
|
|
storage_backend=storage_backend,
|
|
|
|
@@ -709,12 +690,7 @@ def build_anchor_sidecar_stack(
|
|
|
|
|
kv_pool: Any,
|
|
|
|
|
sidecar_pool_name: PoolName,
|
|
|
|
|
full_layer_mapping: dict[int, int],
|
|
|
|
|
page_size: int,
|
|
|
|
|
tp_group,
|
|
|
|
|
load_cache_event,
|
|
|
|
|
attn_cp_group: Optional[torch.distributed.ProcessGroup] = None,
|
|
|
|
|
attn_tp_group: Optional[torch.distributed.ProcessGroup] = None,
|
|
|
|
|
pp_group: Optional[torch.distributed.ProcessGroup] = None,
|
|
|
|
|
storage_backend: Optional[str],
|
|
|
|
|
use_mla: bool,
|
|
|
|
|
override_kv_cache_dim: Optional[int] = None,
|
|
|
|
@@ -727,7 +703,7 @@ def build_anchor_sidecar_stack(
|
|
|
|
|
transfer_layer_num = len(full_layer_mapping)
|
|
|
|
|
kv_host_pool = build_kv_host_pool(
|
|
|
|
|
kv_pool=kv_pool,
|
|
|
|
|
page_size=page_size,
|
|
|
|
|
page_size=params.page_size,
|
|
|
|
|
server_args=server_args,
|
|
|
|
|
use_mla=use_mla,
|
|
|
|
|
override_kv_cache_dim=override_kv_cache_dim,
|
|
|
|
@@ -754,12 +730,12 @@ def build_anchor_sidecar_stack(
|
|
|
|
|
cache_controller = HybridCacheController(
|
|
|
|
|
params.token_to_kv_pool_allocator,
|
|
|
|
|
host_pool_group,
|
|
|
|
|
page_size,
|
|
|
|
|
tp_group,
|
|
|
|
|
params.page_size,
|
|
|
|
|
params.tp_cache_group,
|
|
|
|
|
load_cache_event=load_cache_event,
|
|
|
|
|
attn_cp_group=attn_cp_group,
|
|
|
|
|
attn_tp_group=attn_tp_group,
|
|
|
|
|
pp_group=pp_group,
|
|
|
|
|
attn_cp_group=params.attn_cp_cache_group,
|
|
|
|
|
attn_tp_group=params.attn_tp_cache_group,
|
|
|
|
|
pp_group=params.pp_cache_group,
|
|
|
|
|
write_policy=server_args.hicache_write_policy,
|
|
|
|
|
io_backend=server_args.hicache_io_backend,
|
|
|
|
|
storage_backend=storage_backend,
|
|
|
|
@@ -804,8 +780,6 @@ class StackStrategy:
|
|
|
|
|
params: CacheInitParams,
|
|
|
|
|
server_args: ServerArgs,
|
|
|
|
|
load_cache_event,
|
|
|
|
|
attn_cp_group: Optional[torch.distributed.ProcessGroup] = None,
|
|
|
|
|
attn_tp_group: Optional[torch.distributed.ProcessGroup] = None,
|
|
|
|
|
storage_backend: Optional[str] = None,
|
|
|
|
|
storage_backend_extra_config: Optional[dict] = None,
|
|
|
|
|
prefetch_threshold: int = 256,
|
|
|
|
@@ -834,8 +808,6 @@ class _DeepSeekV4Strategy(StackStrategy):
|
|
|
|
|
params,
|
|
|
|
|
server_args,
|
|
|
|
|
load_cache_event,
|
|
|
|
|
attn_cp_group=None,
|
|
|
|
|
attn_tp_group=None,
|
|
|
|
|
storage_backend=None,
|
|
|
|
|
storage_backend_extra_config=None,
|
|
|
|
|
prefetch_threshold=256,
|
|
|
|
@@ -848,12 +820,7 @@ class _DeepSeekV4Strategy(StackStrategy):
|
|
|
|
|
params=params,
|
|
|
|
|
server_args=server_args,
|
|
|
|
|
kvcache=kvcache,
|
|
|
|
|
page_size=cache.page_size,
|
|
|
|
|
tp_group=params.tp_cache_group,
|
|
|
|
|
load_cache_event=load_cache_event,
|
|
|
|
|
attn_cp_group=attn_cp_group,
|
|
|
|
|
attn_tp_group=attn_tp_group,
|
|
|
|
|
pp_group=params.pp_cache_group,
|
|
|
|
|
storage_backend=storage_backend,
|
|
|
|
|
host_swa_evict_fn=lambda n: cache.evict_host(n, ComponentType.SWA),
|
|
|
|
|
device_swa_evict_fn=lambda n: cache.evict(EvictParams(swa_num_tokens=n)),
|
|
|
|
@@ -917,8 +884,6 @@ class _MambaStrategy(StackStrategy):
|
|
|
|
|
params,
|
|
|
|
|
server_args,
|
|
|
|
|
load_cache_event,
|
|
|
|
|
attn_cp_group=None,
|
|
|
|
|
attn_tp_group=None,
|
|
|
|
|
storage_backend=None,
|
|
|
|
|
storage_backend_extra_config=None,
|
|
|
|
|
prefetch_threshold=256,
|
|
|
|
@@ -936,12 +901,7 @@ class _MambaStrategy(StackStrategy):
|
|
|
|
|
mamba_pool=params.req_to_token_pool.mamba_pool,
|
|
|
|
|
full_layer_mapping=full_layer_mapping,
|
|
|
|
|
mamba_layer_mapping=mamba_layer_mapping,
|
|
|
|
|
page_size=cache.page_size,
|
|
|
|
|
tp_group=params.tp_cache_group,
|
|
|
|
|
load_cache_event=load_cache_event,
|
|
|
|
|
attn_cp_group=attn_cp_group,
|
|
|
|
|
attn_tp_group=attn_tp_group,
|
|
|
|
|
pp_group=params.pp_cache_group,
|
|
|
|
|
storage_backend=storage_backend,
|
|
|
|
|
use_mla=kvcache.use_mla,
|
|
|
|
|
host_mamba_evict_fn=lambda n: cache.evict_host(n, ComponentType.MAMBA),
|
|
|
|
@@ -993,8 +953,6 @@ class _SwaStrategy(StackStrategy):
|
|
|
|
|
params,
|
|
|
|
|
server_args,
|
|
|
|
|
load_cache_event,
|
|
|
|
|
attn_cp_group=None,
|
|
|
|
|
attn_tp_group=None,
|
|
|
|
|
storage_backend=None,
|
|
|
|
|
storage_backend_extra_config=None,
|
|
|
|
|
prefetch_threshold=256,
|
|
|
|
@@ -1011,12 +969,7 @@ class _SwaStrategy(StackStrategy):
|
|
|
|
|
swa_kv_pool=kvcache.swa_kv_pool,
|
|
|
|
|
full_layer_mapping=full_layer_mapping,
|
|
|
|
|
swa_layer_mapping=swa_layer_mapping,
|
|
|
|
|
page_size=cache.page_size,
|
|
|
|
|
tp_group=params.tp_cache_group,
|
|
|
|
|
load_cache_event=load_cache_event,
|
|
|
|
|
attn_cp_group=attn_cp_group,
|
|
|
|
|
attn_tp_group=attn_tp_group,
|
|
|
|
|
pp_group=params.pp_cache_group,
|
|
|
|
|
storage_backend=storage_backend,
|
|
|
|
|
use_mla=False,
|
|
|
|
|
host_swa_evict_fn=lambda n: cache.evict_host(n, ComponentType.SWA),
|
|
|
|
@@ -1129,8 +1082,6 @@ class _DsaStrategy(StackStrategy):
|
|
|
|
|
params,
|
|
|
|
|
server_args,
|
|
|
|
|
load_cache_event,
|
|
|
|
|
attn_cp_group=None,
|
|
|
|
|
attn_tp_group=None,
|
|
|
|
|
storage_backend=None,
|
|
|
|
|
storage_backend_extra_config=None,
|
|
|
|
|
prefetch_threshold=256,
|
|
|
|
@@ -1148,11 +1099,7 @@ class _DsaStrategy(StackStrategy):
|
|
|
|
|
kv_pool=full_kv_pool,
|
|
|
|
|
sidecar_pool_name=PoolName.INDEXER,
|
|
|
|
|
full_layer_mapping=full_layer_mapping,
|
|
|
|
|
page_size=cache.page_size,
|
|
|
|
|
tp_group=params.tp_cache_group,
|
|
|
|
|
load_cache_event=load_cache_event,
|
|
|
|
|
attn_cp_group=attn_cp_group,
|
|
|
|
|
attn_tp_group=attn_tp_group,
|
|
|
|
|
storage_backend=storage_backend,
|
|
|
|
|
use_mla=use_mla,
|
|
|
|
|
override_kv_cache_dim=full_kv_pool.kv_cache_dim,
|
|
|
|
@@ -1200,8 +1147,6 @@ class _MiniMaxSparseStrategy(StackStrategy):
|
|
|
|
|
params,
|
|
|
|
|
server_args,
|
|
|
|
|
load_cache_event,
|
|
|
|
|
attn_cp_group=None,
|
|
|
|
|
attn_tp_group=None,
|
|
|
|
|
storage_backend=None,
|
|
|
|
|
storage_backend_extra_config=None,
|
|
|
|
|
prefetch_threshold=256,
|
|
|
|
@@ -1212,17 +1157,11 @@ class _MiniMaxSparseStrategy(StackStrategy):
|
|
|
|
|
params=params,
|
|
|
|
|
server_args=server_args,
|
|
|
|
|
sparse_pool=kvcache,
|
|
|
|
|
page_size=cache.page_size,
|
|
|
|
|
tp_group=params.tp_cache_group,
|
|
|
|
|
load_cache_event=load_cache_event,
|
|
|
|
|
attn_cp_group=attn_cp_group,
|
|
|
|
|
attn_tp_group=attn_tp_group,
|
|
|
|
|
storage_backend=storage_backend,
|
|
|
|
|
prefetch_threshold=prefetch_threshold,
|
|
|
|
|
model_name=model_name,
|
|
|
|
|
storage_backend_extra_config=storage_backend_extra_config,
|
|
|
|
|
pp_rank=params.pp_rank,
|
|
|
|
|
pp_size=params.pp_size,
|
|
|
|
|
enable_storage_metrics=enable_storage_metrics,
|
|
|
|
|
)
|
|
|
|
|
sidecars = []
|
|
|
|
@@ -1280,8 +1219,6 @@ class _PlainKvStrategy(StackStrategy):
|
|
|
|
|
params,
|
|
|
|
|
server_args,
|
|
|
|
|
load_cache_event,
|
|
|
|
|
attn_cp_group=None,
|
|
|
|
|
attn_tp_group=None,
|
|
|
|
|
storage_backend=None,
|
|
|
|
|
storage_backend_extra_config=None,
|
|
|
|
|
prefetch_threshold=256,
|
|
|
|
@@ -1298,12 +1235,7 @@ class _PlainKvStrategy(StackStrategy):
|
|
|
|
|
server_args=server_args,
|
|
|
|
|
kv_pool=full_kv_pool,
|
|
|
|
|
full_layer_mapping=full_layer_mapping,
|
|
|
|
|
page_size=cache.page_size,
|
|
|
|
|
tp_group=params.tp_cache_group,
|
|
|
|
|
load_cache_event=load_cache_event,
|
|
|
|
|
attn_cp_group=attn_cp_group,
|
|
|
|
|
attn_tp_group=attn_tp_group,
|
|
|
|
|
pp_group=params.pp_cache_group,
|
|
|
|
|
storage_backend=storage_backend,
|
|
|
|
|
use_mla=use_mla,
|
|
|
|
|
prefetch_threshold=prefetch_threshold,
|
|
|
|
@@ -1386,8 +1318,6 @@ def attach_hybrid_pool_to_unified_cache(
|
|
|
|
|
server_args: ServerArgs,
|
|
|
|
|
*,
|
|
|
|
|
load_cache_event,
|
|
|
|
|
attn_cp_group: Optional[torch.distributed.ProcessGroup] = None,
|
|
|
|
|
attn_tp_group: Optional[torch.distributed.ProcessGroup] = None,
|
|
|
|
|
storage_backend: Optional[str] = None,
|
|
|
|
|
storage_extra_config: Optional[dict] = None,
|
|
|
|
|
storage_prefetch_threshold: int = 256,
|
|
|
|
@@ -1403,8 +1333,6 @@ def attach_hybrid_pool_to_unified_cache(
|
|
|
|
|
params=params,
|
|
|
|
|
server_args=server_args,
|
|
|
|
|
load_cache_event=load_cache_event,
|
|
|
|
|
attn_cp_group=attn_cp_group,
|
|
|
|
|
attn_tp_group=attn_tp_group,
|
|
|
|
|
storage_backend=storage_backend,
|
|
|
|
|
storage_backend_extra_config=storage_extra_config,
|
|
|
|
|
prefetch_threshold=storage_prefetch_threshold,
|
|
|
|
@@ -1422,23 +1350,17 @@ def build_minimax_sparse_hicache_stack(
|
|
|
|
|
params: CacheInitParams,
|
|
|
|
|
server_args: ServerArgs,
|
|
|
|
|
sparse_pool: Any,
|
|
|
|
|
page_size: int,
|
|
|
|
|
tp_group,
|
|
|
|
|
load_cache_event,
|
|
|
|
|
attn_cp_group: Optional[torch.distributed.ProcessGroup] = None,
|
|
|
|
|
attn_tp_group: Optional[torch.distributed.ProcessGroup] = None,
|
|
|
|
|
storage_backend: Optional[str],
|
|
|
|
|
prefetch_threshold: int = 256,
|
|
|
|
|
model_name: Optional[str] = None,
|
|
|
|
|
storage_backend_extra_config: Optional[dict] = None,
|
|
|
|
|
pp_rank: int = 0,
|
|
|
|
|
pp_size: int = 1,
|
|
|
|
|
enable_storage_metrics: bool = False,
|
|
|
|
|
) -> tuple[HostPoolGroup, HybridCacheController]:
|
|
|
|
|
"""KV (main_pool) + INDEXER (index_k_pool) host stack for MiniMax M3 sparse."""
|
|
|
|
|
# Mappings are stage-local keyed (controller iterates 0..transfer_layer_num).
|
|
|
|
|
# PP>1 stays gated below pending end-to-end validation of the sparse host path.
|
|
|
|
|
if pp_size > 1:
|
|
|
|
|
if params.pp_size > 1:
|
|
|
|
|
raise NotImplementedError(
|
|
|
|
|
"MiniMax-M3 sparse HiCache does not support pipeline parallelism "
|
|
|
|
|
"(pp_size>1) yet."
|
|
|
|
@@ -1458,7 +1380,7 @@ def build_minimax_sparse_hicache_stack(
|
|
|
|
|
|
|
|
|
|
kv_host_pool = build_kv_host_pool(
|
|
|
|
|
kv_pool=main_pool,
|
|
|
|
|
page_size=page_size,
|
|
|
|
|
page_size=params.page_size,
|
|
|
|
|
server_args=server_args,
|
|
|
|
|
use_mla=False,
|
|
|
|
|
)
|
|
|
|
@@ -1498,11 +1420,11 @@ def build_minimax_sparse_hicache_stack(
|
|
|
|
|
cache_controller = HybridCacheController(
|
|
|
|
|
params.token_to_kv_pool_allocator,
|
|
|
|
|
host_pool_group,
|
|
|
|
|
page_size,
|
|
|
|
|
tp_group,
|
|
|
|
|
params.page_size,
|
|
|
|
|
params.tp_cache_group,
|
|
|
|
|
load_cache_event=load_cache_event,
|
|
|
|
|
attn_cp_group=attn_cp_group,
|
|
|
|
|
attn_tp_group=attn_tp_group,
|
|
|
|
|
attn_cp_group=params.attn_cp_cache_group,
|
|
|
|
|
attn_tp_group=params.attn_tp_cache_group,
|
|
|
|
|
write_policy=server_args.hicache_write_policy,
|
|
|
|
|
io_backend=server_args.hicache_io_backend,
|
|
|
|
|
storage_backend=storage_backend,
|
|
|
|
@@ -1525,8 +1447,6 @@ def attach_hybrid_minimax_sparse_pool_to_hiradix_cache(
|
|
|
|
|
prefetch_threshold: int,
|
|
|
|
|
enable_storage_metrics: bool,
|
|
|
|
|
load_cache_event,
|
|
|
|
|
attn_cp_group: Optional[torch.distributed.ProcessGroup] = None,
|
|
|
|
|
attn_tp_group: Optional[torch.distributed.ProcessGroup] = None,
|
|
|
|
|
) -> None:
|
|
|
|
|
"""Attach HostPoolGroup (KV + index K) + HybridCacheController for HiRadixCache."""
|
|
|
|
|
from sglang.srt.mem_cache.memory_pool import MiniMaxSparseKVPool
|
|
|
|
@@ -1554,18 +1474,12 @@ def attach_hybrid_minimax_sparse_pool_to_hiradix_cache(
|
|
|
|
|
full_layer_mapping={
|
|
|
|
|
layer_id: layer_id for layer_id in range(main_pool.layer_num)
|
|
|
|
|
},
|
|
|
|
|
page_size=radix_cache.page_size,
|
|
|
|
|
tp_group=radix_cache.tp_group,
|
|
|
|
|
load_cache_event=load_cache_event,
|
|
|
|
|
attn_cp_group=attn_cp_group,
|
|
|
|
|
attn_tp_group=attn_tp_group,
|
|
|
|
|
storage_backend=server_args.hicache_storage_backend,
|
|
|
|
|
use_mla=False,
|
|
|
|
|
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,
|
|
|
|
|
enable_storage_metrics=enable_storage_metrics,
|
|
|
|
|
)
|
|
|
|
|
pools_desc = "KV"
|
|
|
|
@@ -1574,17 +1488,11 @@ def attach_hybrid_minimax_sparse_pool_to_hiradix_cache(
|
|
|
|
|
params=params,
|
|
|
|
|
server_args=server_args,
|
|
|
|
|
sparse_pool=sparse_pool,
|
|
|
|
|
page_size=radix_cache.page_size,
|
|
|
|
|
tp_group=radix_cache.tp_group,
|
|
|
|
|
load_cache_event=load_cache_event,
|
|
|
|
|
attn_cp_group=attn_cp_group,
|
|
|
|
|
attn_tp_group=attn_tp_group,
|
|
|
|
|
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,
|
|
|
|
|
enable_storage_metrics=enable_storage_metrics,
|
|
|
|
|
)
|
|
|
|
|
pools_desc = "KV + INDEXER(k-only)"
|
|
|
|
@@ -1614,8 +1522,6 @@ def attach_hybrid_dsa_pool_to_hiradix_cache(
|
|
|
|
|
prefetch_threshold: int,
|
|
|
|
|
enable_storage_metrics: bool,
|
|
|
|
|
load_cache_event,
|
|
|
|
|
attn_cp_group: Optional[torch.distributed.ProcessGroup] = None,
|
|
|
|
|
attn_tp_group: Optional[torch.distributed.ProcessGroup] = None,
|
|
|
|
|
) -> None:
|
|
|
|
|
"""Attach HostPoolGroup (KV + indexer) + HybridCacheController for HiRadixCache.
|
|
|
|
|
|
|
|
|
@@ -1630,12 +1536,7 @@ def attach_hybrid_dsa_pool_to_hiradix_cache(
|
|
|
|
|
kv_pool=kv,
|
|
|
|
|
sidecar_pool_name=PoolName.INDEXER,
|
|
|
|
|
full_layer_mapping=layer_mapping,
|
|
|
|
|
page_size=radix_cache.page_size,
|
|
|
|
|
tp_group=radix_cache.tp_group,
|
|
|
|
|
load_cache_event=load_cache_event,
|
|
|
|
|
attn_cp_group=attn_cp_group,
|
|
|
|
|
attn_tp_group=attn_tp_group,
|
|
|
|
|
pp_group=radix_cache.pp_group,
|
|
|
|
|
storage_backend=server_args.hicache_storage_backend,
|
|
|
|
|
use_mla=True,
|
|
|
|
|
override_kv_cache_dim=kv.kv_cache_dim,
|
|
|
|
@@ -1672,8 +1573,6 @@ def attach_hybrid_pool_to_mamba_cache(
|
|
|
|
|
prefetch_threshold: int,
|
|
|
|
|
load_cache_event,
|
|
|
|
|
enable_storage_metrics: bool = False,
|
|
|
|
|
attn_cp_group: Optional[torch.distributed.ProcessGroup] = None,
|
|
|
|
|
attn_tp_group: Optional[torch.distributed.ProcessGroup] = None,
|
|
|
|
|
) -> None:
|
|
|
|
|
"""Attach HostPoolGroup (KV + Mamba) + HybridCacheController for HiMambaRadixCache.
|
|
|
|
|
|
|
|
|
@@ -1691,12 +1590,7 @@ def attach_hybrid_pool_to_mamba_cache(
|
|
|
|
|
mamba_pool=params.req_to_token_pool.mamba_pool,
|
|
|
|
|
full_layer_mapping=full_layer_mapping,
|
|
|
|
|
mamba_layer_mapping=mamba_layer_mapping,
|
|
|
|
|
page_size=params.page_size,
|
|
|
|
|
tp_group=params.tp_cache_group,
|
|
|
|
|
load_cache_event=load_cache_event,
|
|
|
|
|
attn_cp_group=attn_cp_group,
|
|
|
|
|
attn_tp_group=attn_tp_group,
|
|
|
|
|
pp_group=params.pp_cache_group,
|
|
|
|
|
storage_backend=server_args.hicache_storage_backend,
|
|
|
|
|
use_mla=hybrid_kv.use_mla,
|
|
|
|
|
host_mamba_evict_fn=mamba_cache.evict_mamba_host,
|
|
|
|
|