[HiCache] Fix sparse hybrid transfer layer IDs (#37870)
Co-authored-by: Seokhoon Kang <sh.kang@postech.ac.kr>
This commit is contained in:
co-authored by
Seokhoon Kang
parent
e54009240a
commit
9f3d275940
@@ -343,8 +343,8 @@ class HiCacheController:
|
|||||||
self.host_mem_release_queue: Optional[Queue[torch.Tensor]] = None
|
self.host_mem_release_queue: Optional[Queue[torch.Tensor]] = None
|
||||||
|
|
||||||
self.device = self.mem_pool_device.device
|
self.device = self.mem_pool_device.device
|
||||||
self.layer_num = self.mem_pool_device.layer_num
|
self.transfer_layer_id_max = self.mem_pool_device.layer_num
|
||||||
self.layer_done_counter = LayerDoneCounter(self.layer_num)
|
self.layer_done_counter = LayerDoneCounter(self.transfer_layer_id_max)
|
||||||
self.mem_pool_device.register_layer_transfer_counter(self.layer_done_counter)
|
self.mem_pool_device.register_layer_transfer_counter(self.layer_done_counter)
|
||||||
|
|
||||||
if write_policy not in [
|
if write_policy not in [
|
||||||
@@ -961,7 +961,7 @@ class HiCacheController:
|
|||||||
self._l2_load_transfers(host_indices, device_indices, pool_transfers),
|
self._l2_load_transfers(host_indices, device_indices, pool_transfers),
|
||||||
start_event=producer_event.start_event,
|
start_event=producer_event.start_event,
|
||||||
on_layer_done=producer_event.complete,
|
on_layer_done=producer_event.complete,
|
||||||
layer_num=self.layer_num,
|
transfer_layer_id_max=self.transfer_layer_id_max,
|
||||||
)
|
)
|
||||||
|
|
||||||
self.ack_load_queue.append(
|
self.ack_load_queue.append(
|
||||||
|
|||||||
@@ -181,7 +181,7 @@ class HybridCacheController(BaseHiCacheController):
|
|||||||
prefetch_threshold: int = 256,
|
prefetch_threshold: int = 256,
|
||||||
model_name: Optional[str] = None,
|
model_name: Optional[str] = None,
|
||||||
storage_backend_extra_config: Optional[dict] = None,
|
storage_backend_extra_config: Optional[dict] = None,
|
||||||
transfer_layer_num: Optional[int] = None,
|
transfer_layer_id_max: Optional[int] = None,
|
||||||
enable_storage_metrics: bool = False,
|
enable_storage_metrics: bool = False,
|
||||||
host_memory_mode: str = "cache",
|
host_memory_mode: str = "cache",
|
||||||
):
|
):
|
||||||
@@ -211,11 +211,14 @@ class HybridCacheController(BaseHiCacheController):
|
|||||||
enable_storage_metrics=enable_storage_metrics,
|
enable_storage_metrics=enable_storage_metrics,
|
||||||
host_memory_mode=host_memory_mode,
|
host_memory_mode=host_memory_mode,
|
||||||
)
|
)
|
||||||
# Override layer_num: hybrid models transfer all layers (For example, Linear Model (KV + Mamba)),
|
# Hybrid transfer IDs span every component pool, including holes for
|
||||||
# not just the full attention layers reported by full_kv_pool.
|
# uncached layers that the anchor pool alone cannot describe.
|
||||||
if transfer_layer_num is not None and transfer_layer_num != self.layer_num:
|
if (
|
||||||
self.layer_num = transfer_layer_num
|
transfer_layer_id_max is not None
|
||||||
self.layer_done_counter = LayerDoneCounter(self.layer_num)
|
and transfer_layer_id_max != self.transfer_layer_id_max
|
||||||
|
):
|
||||||
|
self.transfer_layer_id_max = transfer_layer_id_max
|
||||||
|
self.layer_done_counter = LayerDoneCounter(self.transfer_layer_id_max)
|
||||||
|
|
||||||
self.storage_host_pool = mem_pool_host.anchor_entry.host_pool
|
self.storage_host_pool = mem_pool_host.anchor_entry.host_pool
|
||||||
if startup_storage_backend is not None:
|
if startup_storage_backend is not None:
|
||||||
@@ -582,7 +585,9 @@ class HybridCacheController(BaseHiCacheController):
|
|||||||
if target_transfer is None or target_transfer.layer_mapper is None:
|
if target_transfer is None or target_transfer.layer_mapper is None:
|
||||||
continue
|
continue
|
||||||
for depth, draft_device_pool in enumerate(entry.packed_draft_device_pools):
|
for depth, draft_device_pool in enumerate(entry.packed_draft_device_pools):
|
||||||
draft_host_layer = target_transfer.layer_mapper(self.layer_num + depth)
|
draft_host_layer = target_transfer.layer_mapper(
|
||||||
|
self.transfer_layer_id_max + depth
|
||||||
|
)
|
||||||
if draft_host_layer is None:
|
if draft_host_layer is None:
|
||||||
continue
|
continue
|
||||||
|
|
||||||
|
|||||||
@@ -66,10 +66,12 @@ def _evict_mamba_for_device_alloc(cache: UnifiedRadixCache, required_size: int)
|
|||||||
|
|
||||||
def _make_layer_mapper(
|
def _make_layer_mapper(
|
||||||
layer_mapping: dict[int, int],
|
layer_mapping: dict[int, int],
|
||||||
transfer_layer_num: int,
|
transfer_layer_id_max: int,
|
||||||
) -> Callable[[int], Optional[int]]:
|
) -> Callable[[int], Optional[int]]:
|
||||||
|
# The exclusive transfer-ID bound includes holes for uncached layers;
|
||||||
|
# each pool skips IDs absent from its mapping.
|
||||||
def mapper(layer_id: int) -> Optional[int]:
|
def mapper(layer_id: int) -> Optional[int]:
|
||||||
if not 0 <= layer_id < transfer_layer_num:
|
if not 0 <= layer_id < transfer_layer_id_max:
|
||||||
return None
|
return None
|
||||||
return layer_mapping.get(layer_id)
|
return layer_mapping.get(layer_id)
|
||||||
|
|
||||||
@@ -99,7 +101,7 @@ def _with_mtp_layer_mapping(
|
|||||||
|
|
||||||
|
|
||||||
class _DeepSeekV4LayerMappings(NamedTuple):
|
class _DeepSeekV4LayerMappings(NamedTuple):
|
||||||
transfer_layer_num: int
|
transfer_layer_id_max: int
|
||||||
full: dict[int, int]
|
full: dict[int, int]
|
||||||
swa: dict[int, int]
|
swa: dict[int, int]
|
||||||
c4: dict[int, int]
|
c4: dict[int, int]
|
||||||
@@ -111,8 +113,8 @@ class _DeepSeekV4LayerMappings(NamedTuple):
|
|||||||
def _resolve_deepseek_v4_layer_mappings(
|
def _resolve_deepseek_v4_layer_mappings(
|
||||||
kvcache: Any,
|
kvcache: Any,
|
||||||
) -> _DeepSeekV4LayerMappings:
|
) -> _DeepSeekV4LayerMappings:
|
||||||
transfer_layer_num = kvcache.end_layer - kvcache.start_layer
|
transfer_layer_id_max = kvcache.end_layer - kvcache.start_layer
|
||||||
full = {layer: layer for layer in range(transfer_layer_num)}
|
full = {layer: layer for layer in range(transfer_layer_id_max)}
|
||||||
swa = full.copy() if kvcache.swa_kv_pool is not None else {}
|
swa = full.copy() if kvcache.swa_kv_pool is not None else {}
|
||||||
|
|
||||||
c4, c128, c4_state_global_layers = {}, {}, []
|
c4, c128, c4_state_global_layers = {}, {}, []
|
||||||
@@ -126,7 +128,7 @@ def _resolve_deepseek_v4_layer_mappings(
|
|||||||
c128[local_layer] = item.compress_layer_id
|
c128[local_layer] = item.compress_layer_id
|
||||||
|
|
||||||
return _DeepSeekV4LayerMappings(
|
return _DeepSeekV4LayerMappings(
|
||||||
transfer_layer_num=transfer_layer_num,
|
transfer_layer_id_max=transfer_layer_id_max,
|
||||||
full=full,
|
full=full,
|
||||||
swa=swa,
|
swa=swa,
|
||||||
c4=c4,
|
c4=c4,
|
||||||
@@ -196,7 +198,7 @@ def build_pool_entry(
|
|||||||
host_pool: Any,
|
host_pool: Any,
|
||||||
device_pool: Any,
|
device_pool: Any,
|
||||||
layer_mapping: dict[int, int],
|
layer_mapping: dict[int, int],
|
||||||
transfer_layer_num: int,
|
transfer_layer_id_max: int,
|
||||||
is_anchor: bool = False,
|
is_anchor: bool = False,
|
||||||
host_evict_fn: Optional[Callable[[int], Any]] = None,
|
host_evict_fn: Optional[Callable[[int], Any]] = None,
|
||||||
device_evict_fn: Optional[Callable[[int], Any]] = None,
|
device_evict_fn: Optional[Callable[[int], Any]] = None,
|
||||||
@@ -208,7 +210,7 @@ def build_pool_entry(
|
|||||||
name=name,
|
name=name,
|
||||||
host_pool=host_pool,
|
host_pool=host_pool,
|
||||||
device_pool=device_pool,
|
device_pool=device_pool,
|
||||||
layer_mapper=_make_layer_mapper(layer_mapping, transfer_layer_num),
|
layer_mapper=_make_layer_mapper(layer_mapping, transfer_layer_id_max),
|
||||||
is_primary_index_anchor=is_anchor,
|
is_primary_index_anchor=is_anchor,
|
||||||
host_evict_fn=host_evict_fn,
|
host_evict_fn=host_evict_fn,
|
||||||
device_evict_fn=device_evict_fn,
|
device_evict_fn=device_evict_fn,
|
||||||
@@ -229,7 +231,7 @@ def build_kv_only_group(
|
|||||||
mtp_draft_device_pools: tuple[Any, ...] = (),
|
mtp_draft_device_pools: tuple[Any, ...] = (),
|
||||||
) -> HostPoolGroup:
|
) -> HostPoolGroup:
|
||||||
"""Anchor-only host pool group for a flat MHA/MLA device pool."""
|
"""Anchor-only host pool group for a flat MHA/MLA device pool."""
|
||||||
transfer_layer_num = len(full_layer_mapping)
|
transfer_layer_id_max = len(full_layer_mapping)
|
||||||
kv_host_pool = build_kv_host_pool(
|
kv_host_pool = build_kv_host_pool(
|
||||||
kv_pool=kv_pool,
|
kv_pool=kv_pool,
|
||||||
page_size=page_size,
|
page_size=page_size,
|
||||||
@@ -241,7 +243,7 @@ def build_kv_only_group(
|
|||||||
if mtp_draft_device_pools:
|
if mtp_draft_device_pools:
|
||||||
full_layer_mapping = _with_mtp_layer_mapping(
|
full_layer_mapping = _with_mtp_layer_mapping(
|
||||||
full_layer_mapping,
|
full_layer_mapping,
|
||||||
transfer_layer_start=transfer_layer_num,
|
transfer_layer_start=transfer_layer_id_max,
|
||||||
target_device_layer_num=kv_pool.layer_num,
|
target_device_layer_num=kv_pool.layer_num,
|
||||||
draft_layer_num=len(mtp_draft_device_pools),
|
draft_layer_num=len(mtp_draft_device_pools),
|
||||||
)
|
)
|
||||||
@@ -252,7 +254,8 @@ def build_kv_only_group(
|
|||||||
host_pool=kv_host_pool,
|
host_pool=kv_host_pool,
|
||||||
device_pool=kv_pool,
|
device_pool=kv_pool,
|
||||||
layer_mapping=full_layer_mapping,
|
layer_mapping=full_layer_mapping,
|
||||||
transfer_layer_num=transfer_layer_num + len(mtp_draft_device_pools),
|
transfer_layer_id_max=transfer_layer_id_max
|
||||||
|
+ len(mtp_draft_device_pools),
|
||||||
is_anchor=True,
|
is_anchor=True,
|
||||||
packed_draft_device_pools=mtp_draft_device_pools,
|
packed_draft_device_pools=mtp_draft_device_pools,
|
||||||
)
|
)
|
||||||
@@ -276,7 +279,9 @@ def build_hybrid_swa_group(
|
|||||||
mtp_swa_device_pools: tuple[Any, ...] = (),
|
mtp_swa_device_pools: tuple[Any, ...] = (),
|
||||||
) -> HostPoolGroup:
|
) -> HostPoolGroup:
|
||||||
"""Anchor (full) + SWA host pool group for a hybrid-SWA device pool."""
|
"""Anchor (full) + SWA host pool group for a hybrid-SWA device pool."""
|
||||||
transfer_layer_num = len(full_layer_mapping | swa_layer_mapping)
|
transfer_layer_id_max = (
|
||||||
|
max(full_layer_mapping.keys() | swa_layer_mapping.keys()) + 1
|
||||||
|
)
|
||||||
kv_host_pool = build_kv_host_pool(
|
kv_host_pool = build_kv_host_pool(
|
||||||
kv_pool=full_kv_pool,
|
kv_pool=full_kv_pool,
|
||||||
page_size=page_size,
|
page_size=page_size,
|
||||||
@@ -295,7 +300,7 @@ def build_hybrid_swa_group(
|
|||||||
if mtp_swa_device_pools:
|
if mtp_swa_device_pools:
|
||||||
swa_layer_mapping = _with_mtp_layer_mapping(
|
swa_layer_mapping = _with_mtp_layer_mapping(
|
||||||
swa_layer_mapping,
|
swa_layer_mapping,
|
||||||
transfer_layer_start=transfer_layer_num,
|
transfer_layer_start=transfer_layer_id_max,
|
||||||
target_device_layer_num=swa_kv_pool.layer_num,
|
target_device_layer_num=swa_kv_pool.layer_num,
|
||||||
draft_layer_num=len(mtp_swa_device_pools),
|
draft_layer_num=len(mtp_swa_device_pools),
|
||||||
)
|
)
|
||||||
@@ -306,7 +311,7 @@ def build_hybrid_swa_group(
|
|||||||
host_pool=kv_host_pool,
|
host_pool=kv_host_pool,
|
||||||
device_pool=full_kv_pool,
|
device_pool=full_kv_pool,
|
||||||
layer_mapping=full_layer_mapping,
|
layer_mapping=full_layer_mapping,
|
||||||
transfer_layer_num=transfer_layer_num,
|
transfer_layer_id_max=transfer_layer_id_max,
|
||||||
is_anchor=True,
|
is_anchor=True,
|
||||||
),
|
),
|
||||||
build_pool_entry(
|
build_pool_entry(
|
||||||
@@ -314,7 +319,7 @@ def build_hybrid_swa_group(
|
|||||||
host_pool=swa_host_pool,
|
host_pool=swa_host_pool,
|
||||||
device_pool=swa_kv_pool,
|
device_pool=swa_kv_pool,
|
||||||
layer_mapping=swa_layer_mapping,
|
layer_mapping=swa_layer_mapping,
|
||||||
transfer_layer_num=transfer_layer_num + len(mtp_swa_device_pools),
|
transfer_layer_id_max=transfer_layer_id_max + len(mtp_swa_device_pools),
|
||||||
host_evict_fn=host_swa_evict_fn,
|
host_evict_fn=host_swa_evict_fn,
|
||||||
device_evict_fn=device_swa_evict_fn,
|
device_evict_fn=device_swa_evict_fn,
|
||||||
device_alloc_fn=(
|
device_alloc_fn=(
|
||||||
@@ -343,7 +348,7 @@ def build_kv_only_stack(
|
|||||||
storage_backend_extra_config: Optional[dict] = None,
|
storage_backend_extra_config: Optional[dict] = None,
|
||||||
enable_storage_metrics: bool = False,
|
enable_storage_metrics: bool = False,
|
||||||
) -> tuple[HostPoolGroup, HybridCacheController]:
|
) -> tuple[HostPoolGroup, HybridCacheController]:
|
||||||
transfer_layer_num = len(full_layer_mapping)
|
transfer_layer_id_max = len(full_layer_mapping)
|
||||||
host_pool_group = build_kv_only_group(
|
host_pool_group = build_kv_only_group(
|
||||||
page_size=params.page_size,
|
page_size=params.page_size,
|
||||||
kv_pool=kv_pool,
|
kv_pool=kv_pool,
|
||||||
@@ -367,7 +372,7 @@ def build_kv_only_stack(
|
|||||||
prefetch_threshold=prefetch_threshold,
|
prefetch_threshold=prefetch_threshold,
|
||||||
model_name=model_name,
|
model_name=model_name,
|
||||||
storage_backend_extra_config=storage_backend_extra_config,
|
storage_backend_extra_config=storage_backend_extra_config,
|
||||||
transfer_layer_num=transfer_layer_num,
|
transfer_layer_id_max=transfer_layer_id_max,
|
||||||
enable_storage_metrics=enable_storage_metrics,
|
enable_storage_metrics=enable_storage_metrics,
|
||||||
host_memory_mode=get_memory().hicache_host_memory_mode,
|
host_memory_mode=get_memory().hicache_host_memory_mode,
|
||||||
)
|
)
|
||||||
@@ -391,7 +396,9 @@ def build_hybrid_swa_stack(
|
|||||||
storage_backend_extra_config: Optional[dict] = None,
|
storage_backend_extra_config: Optional[dict] = None,
|
||||||
enable_storage_metrics: bool = False,
|
enable_storage_metrics: bool = False,
|
||||||
) -> tuple[HostPoolGroup, HybridCacheController]:
|
) -> tuple[HostPoolGroup, HybridCacheController]:
|
||||||
transfer_layer_num = len(full_layer_mapping | swa_layer_mapping)
|
transfer_layer_id_max = (
|
||||||
|
max(full_layer_mapping.keys() | swa_layer_mapping.keys()) + 1
|
||||||
|
)
|
||||||
# MTP draft pools follow the target SWA layout; select their SWA storage.
|
# MTP draft pools follow the target SWA layout; select their SWA storage.
|
||||||
mtp_swa_device_pools = tuple(
|
mtp_swa_device_pools = tuple(
|
||||||
pool.swa_kv_pool for pool in params.mtp_draft_device_pools
|
pool.swa_kv_pool for pool in params.mtp_draft_device_pools
|
||||||
@@ -433,7 +440,7 @@ def build_hybrid_swa_stack(
|
|||||||
prefetch_threshold=prefetch_threshold,
|
prefetch_threshold=prefetch_threshold,
|
||||||
model_name=model_name,
|
model_name=model_name,
|
||||||
storage_backend_extra_config=storage_backend_extra_config,
|
storage_backend_extra_config=storage_backend_extra_config,
|
||||||
transfer_layer_num=transfer_layer_num,
|
transfer_layer_id_max=transfer_layer_id_max,
|
||||||
enable_storage_metrics=enable_storage_metrics,
|
enable_storage_metrics=enable_storage_metrics,
|
||||||
host_memory_mode=get_memory().hicache_host_memory_mode,
|
host_memory_mode=get_memory().hicache_host_memory_mode,
|
||||||
)
|
)
|
||||||
@@ -599,7 +606,7 @@ def _dsv4_indexer_regions(kvcache: Any, page_size: int) -> list[_IndexerRegion]:
|
|||||||
|
|
||||||
|
|
||||||
def _dsv4_low_ratio_entries(
|
def _dsv4_low_ratio_entries(
|
||||||
kvcache: Any, page_size: int, num_host_pages: int, transfer_layer_num: int
|
kvcache: Any, page_size: int, num_host_pages: int, transfer_layer_id_max: int
|
||||||
):
|
):
|
||||||
"""Mirror each shared source once, in FULL-page units. Prefixes end on an even
|
"""Mirror each shared source once, in FULL-page units. Prefixes end on an even
|
||||||
page boundary, so ratio-2's request-scoped ring is rebuilt, not cached."""
|
page boundary, so ratio-2's request-scoped ring is rebuilt, not cached."""
|
||||||
@@ -671,7 +678,7 @@ def _dsv4_low_ratio_entries(
|
|||||||
),
|
),
|
||||||
device_pool=device_pool,
|
device_pool=device_pool,
|
||||||
layer_mapping=layer_mapping,
|
layer_mapping=layer_mapping,
|
||||||
transfer_layer_num=transfer_layer_num,
|
transfer_layer_id_max=transfer_layer_id_max,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
return entries
|
return entries
|
||||||
@@ -705,7 +712,7 @@ def _build_dsv4_rope_entry(
|
|||||||
layer_mapping: dict[int, int],
|
layer_mapping: dict[int, int],
|
||||||
num_host_pages: int,
|
num_host_pages: int,
|
||||||
slot_page_size: int,
|
slot_page_size: int,
|
||||||
transfer_layer_num: int,
|
transfer_layer_id_max: int,
|
||||||
) -> Optional[PoolEntry]:
|
) -> Optional[PoolEntry]:
|
||||||
sibling = _dsv4_rope_sibling(kvcache, ratio)
|
sibling = _dsv4_rope_sibling(kvcache, ratio)
|
||||||
if sibling is None:
|
if sibling is None:
|
||||||
@@ -725,7 +732,7 @@ def _build_dsv4_rope_entry(
|
|||||||
),
|
),
|
||||||
device_pool=device_pool,
|
device_pool=device_pool,
|
||||||
layer_mapping=layer_mapping,
|
layer_mapping=layer_mapping,
|
||||||
transfer_layer_num=transfer_layer_num,
|
transfer_layer_id_max=transfer_layer_id_max,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@@ -746,7 +753,7 @@ def build_deepseek_v4_hicache_stack(
|
|||||||
) -> tuple[HostPoolGroup, HybridCacheController]:
|
) -> tuple[HostPoolGroup, HybridCacheController]:
|
||||||
page_size = params.page_size
|
page_size = params.page_size
|
||||||
layer_mappings = layer_mappings or _resolve_deepseek_v4_layer_mappings(kvcache)
|
layer_mappings = layer_mappings or _resolve_deepseek_v4_layer_mappings(kvcache)
|
||||||
transfer_layer_num = layer_mappings.transfer_layer_num
|
transfer_layer_id_max = layer_mappings.transfer_layer_id_max
|
||||||
full_layer_mapping = layer_mappings.full
|
full_layer_mapping = layer_mappings.full
|
||||||
|
|
||||||
is_unified_kv = getattr(kvcache, "_unified_kv", False)
|
is_unified_kv = getattr(kvcache, "_unified_kv", False)
|
||||||
@@ -761,11 +768,11 @@ def build_deepseek_v4_hicache_stack(
|
|||||||
physical_page_size=kvcache.swa_kv_pool.page_size,
|
physical_page_size=kvcache.swa_kv_pool.page_size,
|
||||||
consumer="DeepSeek-V4 HiCache",
|
consumer="DeepSeek-V4 HiCache",
|
||||||
)
|
)
|
||||||
if len(kvcache.swa_kv_pool.kv_buffer) != transfer_layer_num:
|
if len(kvcache.swa_kv_pool.kv_buffer) != transfer_layer_id_max:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
"DeepSeek V4 SWA KV pool must be PP-stage-local: "
|
"DeepSeek V4 SWA KV pool must be PP-stage-local: "
|
||||||
f"got {len(kvcache.swa_kv_pool.kv_buffer)} buffers for "
|
f"got {len(kvcache.swa_kv_pool.kv_buffer)} buffers for "
|
||||||
f"{transfer_layer_num} local layers"
|
f"{transfer_layer_id_max} local layers"
|
||||||
)
|
)
|
||||||
swa_layer_mapping = layer_mappings.swa
|
swa_layer_mapping = layer_mappings.swa
|
||||||
# Keep every uncompressed draft SWA layer after the target SWA layers.
|
# Keep every uncompressed draft SWA layer after the target SWA layers.
|
||||||
@@ -777,8 +784,8 @@ def build_deepseek_v4_hicache_stack(
|
|||||||
]
|
]
|
||||||
swa_layer_mapping = _with_mtp_layer_mapping(
|
swa_layer_mapping = _with_mtp_layer_mapping(
|
||||||
swa_layer_mapping,
|
swa_layer_mapping,
|
||||||
transfer_layer_start=transfer_layer_num,
|
transfer_layer_start=transfer_layer_id_max,
|
||||||
target_device_layer_num=transfer_layer_num,
|
target_device_layer_num=transfer_layer_id_max,
|
||||||
draft_layer_num=len(mtp_swa_device_buffers),
|
draft_layer_num=len(mtp_swa_device_buffers),
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -806,7 +813,7 @@ def build_deepseek_v4_hicache_stack(
|
|||||||
host_pool=logical_host_pool,
|
host_pool=logical_host_pool,
|
||||||
device_pool=kvcache,
|
device_pool=kvcache,
|
||||||
layer_mapping=full_layer_mapping,
|
layer_mapping=full_layer_mapping,
|
||||||
transfer_layer_num=transfer_layer_num,
|
transfer_layer_id_max=transfer_layer_id_max,
|
||||||
is_anchor=True,
|
is_anchor=True,
|
||||||
),
|
),
|
||||||
]
|
]
|
||||||
@@ -832,7 +839,8 @@ def build_deepseek_v4_hicache_stack(
|
|||||||
host_pool=swa_host_pool,
|
host_pool=swa_host_pool,
|
||||||
device_pool=kvcache.swa_kv_pool,
|
device_pool=kvcache.swa_kv_pool,
|
||||||
layer_mapping=swa_layer_mapping,
|
layer_mapping=swa_layer_mapping,
|
||||||
transfer_layer_num=transfer_layer_num + len(mtp_swa_device_buffers),
|
transfer_layer_id_max=transfer_layer_id_max
|
||||||
|
+ len(mtp_swa_device_buffers),
|
||||||
host_evict_fn=host_swa_evict_fn,
|
host_evict_fn=host_swa_evict_fn,
|
||||||
device_evict_fn=device_swa_evict_fn,
|
device_evict_fn=device_swa_evict_fn,
|
||||||
device_alloc_fn=swa_attn_allocator.alloc,
|
device_alloc_fn=swa_attn_allocator.alloc,
|
||||||
@@ -863,7 +871,7 @@ def build_deepseek_v4_hicache_stack(
|
|||||||
host_pool=c4_host_pool,
|
host_pool=c4_host_pool,
|
||||||
device_pool=kvcache.c4_kv_pool,
|
device_pool=kvcache.c4_kv_pool,
|
||||||
layer_mapping=c4_layer_mapping,
|
layer_mapping=c4_layer_mapping,
|
||||||
transfer_layer_num=transfer_layer_num,
|
transfer_layer_id_max=transfer_layer_id_max,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
for region in _dsv4_indexer_regions(kvcache, page_size):
|
for region in _dsv4_indexer_regions(kvcache, page_size):
|
||||||
@@ -882,7 +890,7 @@ def build_deepseek_v4_hicache_stack(
|
|||||||
),
|
),
|
||||||
device_pool=kvcache.c4_indexer_kv_pool,
|
device_pool=kvcache.c4_indexer_kv_pool,
|
||||||
layer_mapping=c4_layer_mapping,
|
layer_mapping=c4_layer_mapping,
|
||||||
transfer_layer_num=transfer_layer_num,
|
transfer_layer_id_max=transfer_layer_id_max,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -894,7 +902,7 @@ def build_deepseek_v4_hicache_stack(
|
|||||||
layer_mapping=c4_layer_mapping,
|
layer_mapping=c4_layer_mapping,
|
||||||
num_host_pages=num_host_pages,
|
num_host_pages=num_host_pages,
|
||||||
slot_page_size=page_size,
|
slot_page_size=page_size,
|
||||||
transfer_layer_num=transfer_layer_num,
|
transfer_layer_id_max=transfer_layer_id_max,
|
||||||
)
|
)
|
||||||
if c4_rope_entry is not None:
|
if c4_rope_entry is not None:
|
||||||
entries.append(c4_rope_entry)
|
entries.append(c4_rope_entry)
|
||||||
@@ -929,14 +937,14 @@ def build_deepseek_v4_hicache_stack(
|
|||||||
host_pool=c4_state_host_pool,
|
host_pool=c4_state_host_pool,
|
||||||
device_pool=None,
|
device_pool=None,
|
||||||
layer_mapping=c4_state_mapping,
|
layer_mapping=c4_state_mapping,
|
||||||
transfer_layer_num=transfer_layer_num,
|
transfer_layer_id_max=transfer_layer_id_max,
|
||||||
),
|
),
|
||||||
build_pool_entry(
|
build_pool_entry(
|
||||||
name=PoolName.DEEPSEEK_V4_C4_INDEXER_STATE,
|
name=PoolName.DEEPSEEK_V4_C4_INDEXER_STATE,
|
||||||
host_pool=c4_indexer_state_host_pool,
|
host_pool=c4_indexer_state_host_pool,
|
||||||
device_pool=None,
|
device_pool=None,
|
||||||
layer_mapping=c4_state_mapping,
|
layer_mapping=c4_state_mapping,
|
||||||
transfer_layer_num=transfer_layer_num,
|
transfer_layer_id_max=transfer_layer_id_max,
|
||||||
),
|
),
|
||||||
]
|
]
|
||||||
)
|
)
|
||||||
@@ -973,7 +981,7 @@ def build_deepseek_v4_hicache_stack(
|
|||||||
host_pool=c128_host_pool,
|
host_pool=c128_host_pool,
|
||||||
device_pool=kvcache.c128_kv_pool,
|
device_pool=kvcache.c128_kv_pool,
|
||||||
layer_mapping=c128_layer_mapping,
|
layer_mapping=c128_layer_mapping,
|
||||||
transfer_layer_num=transfer_layer_num,
|
transfer_layer_id_max=transfer_layer_id_max,
|
||||||
# NPU C128 uses bare allocator callbacks for independent indices.
|
# NPU C128 uses bare allocator callbacks for independent indices.
|
||||||
# FULL leaf eviction also frees each attached C128 host value.
|
# FULL leaf eviction also frees each attached C128 host value.
|
||||||
# GPU keeps its KV-derived sidecar callbacks unset.
|
# GPU keeps its KV-derived sidecar callbacks unset.
|
||||||
@@ -999,13 +1007,15 @@ def build_deepseek_v4_hicache_stack(
|
|||||||
layer_mapping=c128_layer_mapping,
|
layer_mapping=c128_layer_mapping,
|
||||||
num_host_pages=c128_num_host_pages,
|
num_host_pages=c128_num_host_pages,
|
||||||
slot_page_size=c128_slot_page_size,
|
slot_page_size=c128_slot_page_size,
|
||||||
transfer_layer_num=transfer_layer_num,
|
transfer_layer_id_max=transfer_layer_id_max,
|
||||||
)
|
)
|
||||||
if c128_rope_entry is not None:
|
if c128_rope_entry is not None:
|
||||||
entries.append(c128_rope_entry)
|
entries.append(c128_rope_entry)
|
||||||
|
|
||||||
entries.extend(
|
entries.extend(
|
||||||
_dsv4_low_ratio_entries(kvcache, page_size, num_host_pages, transfer_layer_num)
|
_dsv4_low_ratio_entries(
|
||||||
|
kvcache, page_size, num_host_pages, transfer_layer_id_max
|
||||||
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
host_pool_group = HostPoolGroup(entries)
|
host_pool_group = HostPoolGroup(entries)
|
||||||
@@ -1024,7 +1034,7 @@ def build_deepseek_v4_hicache_stack(
|
|||||||
prefetch_threshold=prefetch_threshold,
|
prefetch_threshold=prefetch_threshold,
|
||||||
model_name=model_name,
|
model_name=model_name,
|
||||||
storage_backend_extra_config=storage_backend_extra_config,
|
storage_backend_extra_config=storage_backend_extra_config,
|
||||||
transfer_layer_num=transfer_layer_num,
|
transfer_layer_id_max=transfer_layer_id_max,
|
||||||
enable_storage_metrics=enable_storage_metrics,
|
enable_storage_metrics=enable_storage_metrics,
|
||||||
host_memory_mode=get_memory().hicache_host_memory_mode,
|
host_memory_mode=get_memory().hicache_host_memory_mode,
|
||||||
)
|
)
|
||||||
@@ -1048,7 +1058,9 @@ def build_hybrid_mamba_stack(
|
|||||||
storage_backend_extra_config: Optional[dict] = None,
|
storage_backend_extra_config: Optional[dict] = None,
|
||||||
enable_storage_metrics: bool = False,
|
enable_storage_metrics: bool = False,
|
||||||
) -> tuple[HostPoolGroup, HybridCacheController]:
|
) -> tuple[HostPoolGroup, HybridCacheController]:
|
||||||
transfer_layer_num = len(full_layer_mapping | mamba_layer_mapping)
|
transfer_layer_id_max = (
|
||||||
|
max(full_layer_mapping.keys() | mamba_layer_mapping.keys()) + 1
|
||||||
|
)
|
||||||
mamba_allocator = params.req_to_token_pool.mamba_allocator
|
mamba_allocator = params.req_to_token_pool.mamba_allocator
|
||||||
from sglang.srt.mem_cache.memory_pool import HybridLinearKVPool
|
from sglang.srt.mem_cache.memory_pool import HybridLinearKVPool
|
||||||
|
|
||||||
@@ -1071,7 +1083,7 @@ def build_hybrid_mamba_stack(
|
|||||||
if mtp_draft_device_pools:
|
if mtp_draft_device_pools:
|
||||||
full_layer_mapping = _with_mtp_layer_mapping(
|
full_layer_mapping = _with_mtp_layer_mapping(
|
||||||
full_layer_mapping,
|
full_layer_mapping,
|
||||||
transfer_layer_start=transfer_layer_num,
|
transfer_layer_start=transfer_layer_id_max,
|
||||||
target_device_layer_num=kv_pool.layer_num,
|
target_device_layer_num=kv_pool.layer_num,
|
||||||
draft_layer_num=len(mtp_draft_device_pools),
|
draft_layer_num=len(mtp_draft_device_pools),
|
||||||
)
|
)
|
||||||
@@ -1097,7 +1109,7 @@ def build_hybrid_mamba_stack(
|
|||||||
host_pool=kv_host_pool,
|
host_pool=kv_host_pool,
|
||||||
device_pool=kv_pool,
|
device_pool=kv_pool,
|
||||||
layer_mapping=full_layer_mapping,
|
layer_mapping=full_layer_mapping,
|
||||||
transfer_layer_num=transfer_layer_num + len(mtp_draft_device_pools),
|
transfer_layer_id_max=transfer_layer_id_max + len(mtp_draft_device_pools),
|
||||||
is_anchor=True,
|
is_anchor=True,
|
||||||
packed_draft_device_pools=mtp_draft_device_pools,
|
packed_draft_device_pools=mtp_draft_device_pools,
|
||||||
),
|
),
|
||||||
@@ -1106,7 +1118,7 @@ def build_hybrid_mamba_stack(
|
|||||||
host_pool=mamba_host_pool,
|
host_pool=mamba_host_pool,
|
||||||
device_pool=mamba_pool,
|
device_pool=mamba_pool,
|
||||||
layer_mapping=mamba_layer_mapping,
|
layer_mapping=mamba_layer_mapping,
|
||||||
transfer_layer_num=transfer_layer_num,
|
transfer_layer_id_max=transfer_layer_id_max,
|
||||||
host_evict_fn=host_mamba_evict_fn,
|
host_evict_fn=host_mamba_evict_fn,
|
||||||
device_evict_fn=device_mamba_evict_fn,
|
device_evict_fn=device_mamba_evict_fn,
|
||||||
device_alloc_fn=mamba_allocator.alloc,
|
device_alloc_fn=mamba_allocator.alloc,
|
||||||
@@ -1129,7 +1141,7 @@ def build_hybrid_mamba_stack(
|
|||||||
prefetch_threshold=prefetch_threshold,
|
prefetch_threshold=prefetch_threshold,
|
||||||
model_name=model_name,
|
model_name=model_name,
|
||||||
storage_backend_extra_config=storage_backend_extra_config,
|
storage_backend_extra_config=storage_backend_extra_config,
|
||||||
transfer_layer_num=transfer_layer_num,
|
transfer_layer_id_max=transfer_layer_id_max,
|
||||||
enable_storage_metrics=enable_storage_metrics,
|
enable_storage_metrics=enable_storage_metrics,
|
||||||
host_memory_mode=get_memory().hicache_host_memory_mode,
|
host_memory_mode=get_memory().hicache_host_memory_mode,
|
||||||
)
|
)
|
||||||
@@ -1161,8 +1173,13 @@ def build_hybrid_mamba_swa_stack(
|
|||||||
storage_backend_extra_config: Optional[dict] = None,
|
storage_backend_extra_config: Optional[dict] = None,
|
||||||
enable_storage_metrics: bool = False,
|
enable_storage_metrics: bool = False,
|
||||||
) -> tuple[HostPoolGroup, HybridCacheController]:
|
) -> tuple[HostPoolGroup, HybridCacheController]:
|
||||||
transfer_layer_num = len(
|
transfer_layer_id_max = (
|
||||||
full_layer_mapping | swa_layer_mapping | mamba_layer_mapping
|
max(
|
||||||
|
full_layer_mapping.keys()
|
||||||
|
| swa_layer_mapping.keys()
|
||||||
|
| mamba_layer_mapping.keys()
|
||||||
|
)
|
||||||
|
+ 1
|
||||||
)
|
)
|
||||||
swa_attn_allocator = params.token_to_kv_pool_allocator.swa_attn_allocator
|
swa_attn_allocator = params.token_to_kv_pool_allocator.swa_attn_allocator
|
||||||
mamba_allocator = params.req_to_token_pool.mamba_allocator
|
mamba_allocator = params.req_to_token_pool.mamba_allocator
|
||||||
@@ -1198,7 +1215,7 @@ def build_hybrid_mamba_swa_stack(
|
|||||||
host_pool=kv_host_pool,
|
host_pool=kv_host_pool,
|
||||||
device_pool=full_kv_pool,
|
device_pool=full_kv_pool,
|
||||||
layer_mapping=full_layer_mapping,
|
layer_mapping=full_layer_mapping,
|
||||||
transfer_layer_num=transfer_layer_num,
|
transfer_layer_id_max=transfer_layer_id_max,
|
||||||
is_anchor=True,
|
is_anchor=True,
|
||||||
),
|
),
|
||||||
build_pool_entry(
|
build_pool_entry(
|
||||||
@@ -1206,7 +1223,7 @@ def build_hybrid_mamba_swa_stack(
|
|||||||
host_pool=swa_host_pool,
|
host_pool=swa_host_pool,
|
||||||
device_pool=swa_kv_pool,
|
device_pool=swa_kv_pool,
|
||||||
layer_mapping=swa_layer_mapping,
|
layer_mapping=swa_layer_mapping,
|
||||||
transfer_layer_num=transfer_layer_num,
|
transfer_layer_id_max=transfer_layer_id_max,
|
||||||
host_evict_fn=host_swa_evict_fn,
|
host_evict_fn=host_swa_evict_fn,
|
||||||
device_evict_fn=device_swa_evict_fn,
|
device_evict_fn=device_swa_evict_fn,
|
||||||
device_alloc_fn=swa_attn_allocator.alloc,
|
device_alloc_fn=swa_attn_allocator.alloc,
|
||||||
@@ -1217,7 +1234,7 @@ def build_hybrid_mamba_swa_stack(
|
|||||||
host_pool=mamba_host_pool,
|
host_pool=mamba_host_pool,
|
||||||
device_pool=mamba_pool,
|
device_pool=mamba_pool,
|
||||||
layer_mapping=mamba_layer_mapping,
|
layer_mapping=mamba_layer_mapping,
|
||||||
transfer_layer_num=transfer_layer_num,
|
transfer_layer_id_max=transfer_layer_id_max,
|
||||||
host_evict_fn=host_mamba_evict_fn,
|
host_evict_fn=host_mamba_evict_fn,
|
||||||
device_evict_fn=device_mamba_evict_fn,
|
device_evict_fn=device_mamba_evict_fn,
|
||||||
device_alloc_fn=mamba_allocator.alloc,
|
device_alloc_fn=mamba_allocator.alloc,
|
||||||
@@ -1240,7 +1257,7 @@ def build_hybrid_mamba_swa_stack(
|
|||||||
prefetch_threshold=prefetch_threshold,
|
prefetch_threshold=prefetch_threshold,
|
||||||
model_name=model_name,
|
model_name=model_name,
|
||||||
storage_backend_extra_config=storage_backend_extra_config,
|
storage_backend_extra_config=storage_backend_extra_config,
|
||||||
transfer_layer_num=transfer_layer_num,
|
transfer_layer_id_max=transfer_layer_id_max,
|
||||||
enable_storage_metrics=enable_storage_metrics,
|
enable_storage_metrics=enable_storage_metrics,
|
||||||
host_memory_mode=get_memory().hicache_host_memory_mode,
|
host_memory_mode=get_memory().hicache_host_memory_mode,
|
||||||
)
|
)
|
||||||
@@ -1263,7 +1280,7 @@ def build_anchor_sidecar_stack(
|
|||||||
storage_backend_extra_config: Optional[dict] = None,
|
storage_backend_extra_config: Optional[dict] = None,
|
||||||
enable_storage_metrics: bool = False,
|
enable_storage_metrics: bool = False,
|
||||||
) -> tuple[HostPoolGroup, HybridCacheController]:
|
) -> tuple[HostPoolGroup, HybridCacheController]:
|
||||||
transfer_layer_num = len(full_layer_mapping)
|
transfer_layer_id_max = len(full_layer_mapping)
|
||||||
mtp_draft_device_pools = tuple(
|
mtp_draft_device_pools = tuple(
|
||||||
pool for pool in params.mtp_draft_device_pools if pool.index_k_with_scale_buffer
|
pool for pool in params.mtp_draft_device_pools if pool.index_k_with_scale_buffer
|
||||||
)
|
)
|
||||||
@@ -1279,7 +1296,7 @@ def build_anchor_sidecar_stack(
|
|||||||
if mtp_draft_device_pools:
|
if mtp_draft_device_pools:
|
||||||
full_layer_mapping = _with_mtp_layer_mapping(
|
full_layer_mapping = _with_mtp_layer_mapping(
|
||||||
full_layer_mapping,
|
full_layer_mapping,
|
||||||
transfer_layer_start=transfer_layer_num,
|
transfer_layer_start=transfer_layer_id_max,
|
||||||
target_device_layer_num=kv_pool.layer_num,
|
target_device_layer_num=kv_pool.layer_num,
|
||||||
draft_layer_num=len(mtp_draft_device_pools),
|
draft_layer_num=len(mtp_draft_device_pools),
|
||||||
)
|
)
|
||||||
@@ -1289,7 +1306,7 @@ def build_anchor_sidecar_stack(
|
|||||||
host_pool=kv_host_pool,
|
host_pool=kv_host_pool,
|
||||||
device_pool=kv_pool,
|
device_pool=kv_pool,
|
||||||
layer_mapping=full_layer_mapping,
|
layer_mapping=full_layer_mapping,
|
||||||
transfer_layer_num=transfer_layer_num + len(mtp_draft_device_pools),
|
transfer_layer_id_max=transfer_layer_id_max + len(mtp_draft_device_pools),
|
||||||
is_anchor=True,
|
is_anchor=True,
|
||||||
packed_draft_device_pools=mtp_draft_device_pools,
|
packed_draft_device_pools=mtp_draft_device_pools,
|
||||||
),
|
),
|
||||||
@@ -1298,7 +1315,7 @@ def build_anchor_sidecar_stack(
|
|||||||
host_pool=sidecar_host_pool,
|
host_pool=sidecar_host_pool,
|
||||||
device_pool=kv_pool,
|
device_pool=kv_pool,
|
||||||
layer_mapping=full_layer_mapping,
|
layer_mapping=full_layer_mapping,
|
||||||
transfer_layer_num=transfer_layer_num + len(mtp_draft_device_pools),
|
transfer_layer_id_max=transfer_layer_id_max + len(mtp_draft_device_pools),
|
||||||
packed_draft_device_pools=mtp_draft_device_pools,
|
packed_draft_device_pools=mtp_draft_device_pools,
|
||||||
),
|
),
|
||||||
]
|
]
|
||||||
@@ -1318,7 +1335,7 @@ def build_anchor_sidecar_stack(
|
|||||||
prefetch_threshold=prefetch_threshold,
|
prefetch_threshold=prefetch_threshold,
|
||||||
model_name=model_name,
|
model_name=model_name,
|
||||||
storage_backend_extra_config=storage_backend_extra_config,
|
storage_backend_extra_config=storage_backend_extra_config,
|
||||||
transfer_layer_num=transfer_layer_num,
|
transfer_layer_id_max=transfer_layer_id_max,
|
||||||
enable_storage_metrics=enable_storage_metrics,
|
enable_storage_metrics=enable_storage_metrics,
|
||||||
host_memory_mode=get_memory().hicache_host_memory_mode,
|
host_memory_mode=get_memory().hicache_host_memory_mode,
|
||||||
)
|
)
|
||||||
@@ -1401,7 +1418,7 @@ def build_full_draft_pools(
|
|||||||
host_pool=draft_host_pool,
|
host_pool=draft_host_pool,
|
||||||
device_pool=pool,
|
device_pool=pool,
|
||||||
layer_mapping=draft_layer_mapping,
|
layer_mapping=draft_layer_mapping,
|
||||||
transfer_layer_num=draft_host_pool.layer_num,
|
transfer_layer_id_max=draft_host_pool.layer_num,
|
||||||
)
|
)
|
||||||
]
|
]
|
||||||
|
|
||||||
@@ -1424,7 +1441,7 @@ def build_full_draft_pools(
|
|||||||
host_pool=indexer_host_pool,
|
host_pool=indexer_host_pool,
|
||||||
device_pool=pool,
|
device_pool=pool,
|
||||||
layer_mapping=draft_layer_mapping,
|
layer_mapping=draft_layer_mapping,
|
||||||
transfer_layer_num=indexer_host_pool.layer_num,
|
transfer_layer_id_max=indexer_host_pool.layer_num,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -1484,7 +1501,7 @@ def build_swa_draft_pools(
|
|||||||
host_pool=host_pool,
|
host_pool=host_pool,
|
||||||
device_pool=draft_swa_pool,
|
device_pool=draft_swa_pool,
|
||||||
layer_mapping=layer_mapping,
|
layer_mapping=layer_mapping,
|
||||||
transfer_layer_num=host_pool.layer_num,
|
transfer_layer_id_max=host_pool.layer_num,
|
||||||
)
|
)
|
||||||
return [spec], [entry]
|
return [spec], [entry]
|
||||||
|
|
||||||
@@ -1527,7 +1544,6 @@ class StackBuildResult:
|
|||||||
# Mamba state lives in req_to_token_pool, not in kvcache, so its
|
# Mamba state lives in req_to_token_pool, not in kvcache, so its
|
||||||
# layer_transfer_counter has to be wired separately.
|
# layer_transfer_counter has to be wired separately.
|
||||||
register_req_to_token_counter: bool = False
|
register_req_to_token_counter: bool = False
|
||||||
transfer_layer_num: int = 0
|
|
||||||
pools_desc: str = ""
|
pools_desc: str = ""
|
||||||
|
|
||||||
|
|
||||||
@@ -1681,7 +1697,6 @@ class _DeepSeekV4Strategy(StackStrategy):
|
|||||||
cache_controller=cache_controller,
|
cache_controller=cache_controller,
|
||||||
component_host_pools=component_host_pools,
|
component_host_pools=component_host_pools,
|
||||||
sidecars=sidecars,
|
sidecars=sidecars,
|
||||||
transfer_layer_num=kvcache.end_layer - kvcache.start_layer,
|
|
||||||
pools_desc="KV + SWA + DeepSeekV4 sidecars",
|
pools_desc="KV + SWA + DeepSeekV4 sidecars",
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -1739,7 +1754,6 @@ class _MambaStrategy(StackStrategy):
|
|||||||
ComponentType.MAMBA: host_pool_group.get_pool(PoolName.MAMBA),
|
ComponentType.MAMBA: host_pool_group.get_pool(PoolName.MAMBA),
|
||||||
},
|
},
|
||||||
register_req_to_token_counter=True,
|
register_req_to_token_counter=True,
|
||||||
transfer_layer_num=len(full_layer_mapping | mamba_layer_mapping),
|
|
||||||
pools_desc="KV + MAMBA",
|
pools_desc="KV + MAMBA",
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -1807,7 +1821,6 @@ class _SwaStrategy(StackStrategy):
|
|||||||
ComponentType.FULL: host_pool_group.get_pool(PoolName.KV),
|
ComponentType.FULL: host_pool_group.get_pool(PoolName.KV),
|
||||||
ComponentType.SWA: host_pool_group.get_pool(PoolName.SWA),
|
ComponentType.SWA: host_pool_group.get_pool(PoolName.SWA),
|
||||||
},
|
},
|
||||||
transfer_layer_num=len(full_layer_mapping | swa_layer_mapping),
|
|
||||||
pools_desc="Full + SWA",
|
pools_desc="Full + SWA",
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -1877,9 +1890,6 @@ class _MambaSwaStrategy(StackStrategy):
|
|||||||
ComponentType.MAMBA: host_pool_group.get_pool(PoolName.MAMBA),
|
ComponentType.MAMBA: host_pool_group.get_pool(PoolName.MAMBA),
|
||||||
},
|
},
|
||||||
register_req_to_token_counter=True,
|
register_req_to_token_counter=True,
|
||||||
transfer_layer_num=len(
|
|
||||||
full_layer_mapping | swa_layer_mapping | mamba_layer_mapping
|
|
||||||
),
|
|
||||||
pools_desc="KV + SWA + MAMBA",
|
pools_desc="KV + SWA + MAMBA",
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -1952,7 +1962,6 @@ class _DsaStrategy(StackStrategy):
|
|||||||
indices_from_pool=PoolName.KV,
|
indices_from_pool=PoolName.KV,
|
||||||
),
|
),
|
||||||
],
|
],
|
||||||
transfer_layer_num=len(full_layer_mapping),
|
|
||||||
pools_desc="KV + INDEXER",
|
pools_desc="KV + INDEXER",
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -2006,7 +2015,6 @@ class _MiniMaxSparseStrategy(StackStrategy):
|
|||||||
ComponentType.FULL: host_pool_group.get_pool(PoolName.KV),
|
ComponentType.FULL: host_pool_group.get_pool(PoolName.KV),
|
||||||
},
|
},
|
||||||
sidecars=sidecars,
|
sidecars=sidecars,
|
||||||
transfer_layer_num=kvcache.main_pool.layer_num,
|
|
||||||
pools_desc=pools_desc,
|
pools_desc=pools_desc,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -2071,7 +2079,6 @@ class _PlainKvStrategy(StackStrategy):
|
|||||||
component_host_pools={
|
component_host_pools={
|
||||||
ComponentType.FULL: host_pool_group.get_pool(PoolName.KV),
|
ComponentType.FULL: host_pool_group.get_pool(PoolName.KV),
|
||||||
},
|
},
|
||||||
transfer_layer_num=len(full_layer_mapping),
|
|
||||||
pools_desc="KV",
|
pools_desc="KV",
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -2128,9 +2135,9 @@ def _apply_stack_result(
|
|||||||
)
|
)
|
||||||
|
|
||||||
logger.info(
|
logger.info(
|
||||||
"Attached hybrid pool stack to UnifiedRadixCache: pools=%s, transfer_layer_num=%s",
|
"Attached hybrid pool stack to UnifiedRadixCache: pools=%s, transfer_layer_id_max=%s",
|
||||||
result.pools_desc,
|
result.pools_desc,
|
||||||
result.transfer_layer_num,
|
result.cache_controller.transfer_layer_id_max,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@@ -2179,7 +2186,7 @@ def build_minimax_sparse_hicache_stack(
|
|||||||
enable_storage_metrics: bool = False,
|
enable_storage_metrics: bool = False,
|
||||||
) -> tuple[HostPoolGroup, HybridCacheController]:
|
) -> tuple[HostPoolGroup, HybridCacheController]:
|
||||||
"""KV (main_pool) + INDEXER (index_k_pool) host stack for MiniMax M3 sparse."""
|
"""KV (main_pool) + INDEXER (index_k_pool) host stack for MiniMax M3 sparse."""
|
||||||
# Mappings are stage-local keyed (controller iterates 0..transfer_layer_num).
|
# Mappings are stage-local keyed (controller iterates 0..transfer_layer_id_max).
|
||||||
# PP>1 stays gated below pending end-to-end validation of the sparse host path.
|
# PP>1 stays gated below pending end-to-end validation of the sparse host path.
|
||||||
if params.pp_size > 1:
|
if params.pp_size > 1:
|
||||||
raise NotImplementedError(
|
raise NotImplementedError(
|
||||||
@@ -2194,10 +2201,12 @@ def build_minimax_sparse_hicache_stack(
|
|||||||
)
|
)
|
||||||
main_pool = sparse_pool.main_pool
|
main_pool = sparse_pool.main_pool
|
||||||
start_layer = main_pool.start_layer
|
start_layer = main_pool.start_layer
|
||||||
transfer_layer_num = main_pool.layer_num
|
transfer_layer_id_max = main_pool.layer_num
|
||||||
# Stage-local keys (0..transfer_layer_num) match the controller's per-layer
|
# Stage-local keys (0..transfer_layer_id_max) match the controller's per-layer
|
||||||
# load loop; values index the host pool's local layer buffer.
|
# load loop; values index the host pool's local layer buffer.
|
||||||
full_layer_mapping = {layer_id: layer_id for layer_id in range(transfer_layer_num)}
|
full_layer_mapping = {
|
||||||
|
layer_id: layer_id for layer_id in range(transfer_layer_id_max)
|
||||||
|
}
|
||||||
|
|
||||||
kv_host_pool = build_kv_host_pool(
|
kv_host_pool = build_kv_host_pool(
|
||||||
kv_pool=main_pool,
|
kv_pool=main_pool,
|
||||||
@@ -2210,7 +2219,7 @@ def build_minimax_sparse_hicache_stack(
|
|||||||
host_pool=kv_host_pool,
|
host_pool=kv_host_pool,
|
||||||
device_pool=main_pool,
|
device_pool=main_pool,
|
||||||
layer_mapping=full_layer_mapping,
|
layer_mapping=full_layer_mapping,
|
||||||
transfer_layer_num=transfer_layer_num,
|
transfer_layer_id_max=transfer_layer_id_max,
|
||||||
is_anchor=True,
|
is_anchor=True,
|
||||||
),
|
),
|
||||||
]
|
]
|
||||||
@@ -2232,7 +2241,7 @@ def build_minimax_sparse_hicache_stack(
|
|||||||
gid - start_layer: sub_id
|
gid - start_layer: sub_id
|
||||||
for gid, sub_id in sparse_pool.index_k_layer_id_mapping.items()
|
for gid, sub_id in sparse_pool.index_k_layer_id_mapping.items()
|
||||||
},
|
},
|
||||||
transfer_layer_num=transfer_layer_num,
|
transfer_layer_id_max=transfer_layer_id_max,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -2252,7 +2261,7 @@ def build_minimax_sparse_hicache_stack(
|
|||||||
model_name=model_name,
|
model_name=model_name,
|
||||||
storage_backend_extra_config=storage_backend_extra_config,
|
storage_backend_extra_config=storage_backend_extra_config,
|
||||||
pp_group=params.pp_cache_group,
|
pp_group=params.pp_cache_group,
|
||||||
transfer_layer_num=transfer_layer_num,
|
transfer_layer_id_max=transfer_layer_id_max,
|
||||||
enable_storage_metrics=enable_storage_metrics,
|
enable_storage_metrics=enable_storage_metrics,
|
||||||
)
|
)
|
||||||
return host_pool_group, cache_controller
|
return host_pool_group, cache_controller
|
||||||
@@ -2320,7 +2329,7 @@ def attach_hybrid_minimax_sparse_pool_to_hiradix_cache(
|
|||||||
radix_cache.cache_controller = cache_controller
|
radix_cache.cache_controller = cache_controller
|
||||||
logger.info(
|
logger.info(
|
||||||
"Attached hybrid MiniMax sparse pool stack to HiRadixCache: pools=%s, "
|
"Attached hybrid MiniMax sparse pool stack to HiRadixCache: pools=%s, "
|
||||||
"transfer_layer_num=%s, sparse_index_k_layers=%s",
|
"transfer_layer_id_max=%s, sparse_index_k_layers=%s",
|
||||||
pools_desc,
|
pools_desc,
|
||||||
main_pool.layer_num,
|
main_pool.layer_num,
|
||||||
len(sparse_pool.index_k_layer_id_mapping),
|
len(sparse_pool.index_k_layer_id_mapping),
|
||||||
@@ -2371,7 +2380,7 @@ def attach_hybrid_dsa_pool_to_hiradix_cache(
|
|||||||
radix_cache.cache_controller = cache_controller
|
radix_cache.cache_controller = cache_controller
|
||||||
logger.info(
|
logger.info(
|
||||||
"Attached hybrid DSA pool stack to HiRadixCache: pools=KV + INDEXER, "
|
"Attached hybrid DSA pool stack to HiRadixCache: pools=KV + INDEXER, "
|
||||||
"transfer_layer_num=%s",
|
"transfer_layer_id_max=%s",
|
||||||
len(layer_mapping),
|
len(layer_mapping),
|
||||||
)
|
)
|
||||||
except Exception:
|
except Exception:
|
||||||
|
|||||||
@@ -368,7 +368,7 @@ def _build_deepseek_v4_device_pool_group(
|
|||||||
)
|
)
|
||||||
return DevicePoolGroup(
|
return DevicePoolGroup(
|
||||||
entries,
|
entries,
|
||||||
mappings.transfer_layer_num,
|
mappings.transfer_layer_id_max,
|
||||||
page_size,
|
page_size,
|
||||||
rank_replicated=True,
|
rank_replicated=True,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -75,7 +75,7 @@ class L2TransferEngine:
|
|||||||
self,
|
self,
|
||||||
transfers: list[L2Transfer],
|
transfers: list[L2Transfer],
|
||||||
*,
|
*,
|
||||||
layer_num: int,
|
transfer_layer_id_max: int,
|
||||||
start_event=None,
|
start_event=None,
|
||||||
on_layer_done=None,
|
on_layer_done=None,
|
||||||
) -> TransferCompletion:
|
) -> TransferCompletion:
|
||||||
@@ -85,7 +85,7 @@ class L2TransferEngine:
|
|||||||
with device_module.stream(self.host_to_device_stream):
|
with device_module.stream(self.host_to_device_stream):
|
||||||
start_event.wait(self.host_to_device_stream)
|
start_event.wait(self.host_to_device_stream)
|
||||||
ack_start.record()
|
ack_start.record()
|
||||||
for layer_id in range(layer_num):
|
for layer_id in range(transfer_layer_id_max):
|
||||||
for transfer in transfers:
|
for transfer in transfers:
|
||||||
local_layer_id = (
|
local_layer_id = (
|
||||||
transfer.layer_mapper(layer_id)
|
transfer.layer_mapper(layer_id)
|
||||||
|
|||||||
@@ -1527,7 +1527,7 @@ class UnifiedRadixCache(BasePrefixCache):
|
|||||||
self.cache_controller._l2_load_transfers(
|
self.cache_controller._l2_load_transfers(
|
||||||
load_host, load_device, load_pools
|
load_host, load_device, load_pools
|
||||||
),
|
),
|
||||||
layer_num=self.cache_controller.layer_num,
|
transfer_layer_id_max=self.cache_controller.transfer_layer_id_max,
|
||||||
)
|
)
|
||||||
completion.finish_event.synchronize()
|
completion.finish_event.synchronize()
|
||||||
self.retraction_discard(backup)
|
self.retraction_discard(backup)
|
||||||
|
|||||||
@@ -254,7 +254,7 @@ class TestDSV4PoolAssembly(CustomTestCase):
|
|||||||
pp_cache_group=None,
|
pp_cache_group=None,
|
||||||
)
|
)
|
||||||
mappings = assembler._DeepSeekV4LayerMappings(
|
mappings = assembler._DeepSeekV4LayerMappings(
|
||||||
transfer_layer_num=1,
|
transfer_layer_id_max=1,
|
||||||
full={0: 0},
|
full={0: 0},
|
||||||
swa={},
|
swa={},
|
||||||
c4={0: 0},
|
c4={0: 0},
|
||||||
|
|||||||
@@ -238,7 +238,7 @@ class TestHiCacheStagedWriteBackDispatch(CustomTestCase):
|
|||||||
start_event=object(), finish_event=object(), timing_enabled=False
|
start_event=object(), finish_event=object(), timing_enabled=False
|
||||||
)
|
)
|
||||||
controller.l2_transfer_engine.submit_host_to_device.return_value = completion
|
controller.l2_transfer_engine.submit_host_to_device.return_value = completion
|
||||||
controller.layer_num = 2
|
controller.transfer_layer_id_max = 2
|
||||||
controller.ack_load_queue = []
|
controller.ack_load_queue = []
|
||||||
|
|
||||||
self.assertEqual(HybridCacheController.start_loading(controller), 0)
|
self.assertEqual(HybridCacheController.start_loading(controller), 0)
|
||||||
@@ -344,7 +344,9 @@ class TestHiCacheStagedWriteBackDispatch(CustomTestCase):
|
|||||||
layer_mapper={1: 0, 3: 1}.get,
|
layer_mapper={1: 0, 3: 1}.get,
|
||||||
)
|
)
|
||||||
with mock.patch.object(transfer_module, "device_module", _FakeDeviceModule):
|
with mock.patch.object(transfer_module, "device_module", _FakeDeviceModule):
|
||||||
L2TransferEngine("kernel").submit_host_to_device([transfer], layer_num=4)
|
L2TransferEngine("kernel").submit_host_to_device(
|
||||||
|
[transfer], transfer_layer_id_max=4
|
||||||
|
)
|
||||||
|
|
||||||
self.assertEqual(
|
self.assertEqual(
|
||||||
[
|
[
|
||||||
@@ -369,7 +371,7 @@ class TestHiCacheStagedWriteBackDispatch(CustomTestCase):
|
|||||||
anchor_entry=entry,
|
anchor_entry=entry,
|
||||||
entry_map={entry.name: entry},
|
entry_map={entry.name: entry},
|
||||||
)
|
)
|
||||||
controller.layer_num = 2
|
controller.transfer_layer_id_max = 2
|
||||||
|
|
||||||
self.assertEqual(
|
self.assertEqual(
|
||||||
len(controller._l2_transfers(_indices(0, 2), _indices(2, 4))), 1
|
len(controller._l2_transfers(_indices(0, 2), _indices(2, 4))), 1
|
||||||
@@ -380,7 +382,9 @@ class TestHiCacheStagedWriteBackDispatch(CustomTestCase):
|
|||||||
self.assertFalse(transfers[0].is_draft)
|
self.assertFalse(transfers[0].is_draft)
|
||||||
self.assertTrue(transfers[1].is_draft)
|
self.assertTrue(transfers[1].is_draft)
|
||||||
with mock.patch.object(transfer_module, "device_module", _FakeDeviceModule):
|
with mock.patch.object(transfer_module, "device_module", _FakeDeviceModule):
|
||||||
L2TransferEngine("kernel").submit_host_to_device(transfers, layer_num=2)
|
L2TransferEngine("kernel").submit_host_to_device(
|
||||||
|
transfers, transfer_layer_id_max=2
|
||||||
|
)
|
||||||
self.assertEqual(
|
self.assertEqual(
|
||||||
[
|
[
|
||||||
call.args[3]
|
call.args[3]
|
||||||
|
|||||||
@@ -15,6 +15,7 @@ from sglang.srt.mem_cache.hybrid_cache.hybrid_pool_assembler import (
|
|||||||
_split_hicache_size,
|
_split_hicache_size,
|
||||||
_SwaStrategy,
|
_SwaStrategy,
|
||||||
build_full_draft_pools,
|
build_full_draft_pools,
|
||||||
|
build_hybrid_swa_group,
|
||||||
)
|
)
|
||||||
from sglang.srt.mem_cache.memory_pool import HybridLinearKVPool
|
from sglang.srt.mem_cache.memory_pool import HybridLinearKVPool
|
||||||
from sglang.test.ci.ci_register import register_cpu_ci
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
@@ -153,7 +154,10 @@ class TestHybridStageLayerMappings(CustomTestCase):
|
|||||||
with patch.object(
|
with patch.object(
|
||||||
hybrid_pool_assembler,
|
hybrid_pool_assembler,
|
||||||
builder_name,
|
builder_name,
|
||||||
return_value=(MagicMock(), object()),
|
return_value=(
|
||||||
|
MagicMock(),
|
||||||
|
SimpleNamespace(transfer_layer_id_max=4),
|
||||||
|
),
|
||||||
) as build_stack:
|
) as build_stack:
|
||||||
result = strategy_cls().build(
|
result = strategy_cls().build(
|
||||||
cache=SimpleNamespace(page_size=1),
|
cache=SimpleNamespace(page_size=1),
|
||||||
@@ -169,7 +173,7 @@ class TestHybridStageLayerMappings(CustomTestCase):
|
|||||||
build_stack.call_args.kwargs[f"{name}_layer_mapping"],
|
build_stack.call_args.kwargs[f"{name}_layer_mapping"],
|
||||||
mapping,
|
mapping,
|
||||||
)
|
)
|
||||||
self.assertEqual(result.transfer_layer_num, 4)
|
self.assertEqual(result.cache_controller.transfer_layer_id_max, 4)
|
||||||
self.assertEqual(
|
self.assertEqual(
|
||||||
kvcache.full_attention_layer_id_mapping, global_maps["full"]
|
kvcache.full_attention_layer_id_mapping, global_maps["full"]
|
||||||
)
|
)
|
||||||
@@ -229,5 +233,44 @@ class TestDraftSidecarPoolDispatch(CustomTestCase):
|
|||||||
self.assertIs(entries[0].host_pool, draft_host_pool)
|
self.assertIs(entries[0].host_pool, draft_host_pool)
|
||||||
|
|
||||||
|
|
||||||
|
_ASSEMBLER = "sglang.srt.mem_cache.hybrid_cache.hybrid_pool_assembler."
|
||||||
|
|
||||||
|
|
||||||
|
class TestTransferLayerSpan(CustomTestCase):
|
||||||
|
"""``transfer_layer_id_max`` must span global layer ids, not count the mapped ones.
|
||||||
|
|
||||||
|
A hybrid model with an uncached layer type keys its mappings non-contiguously,
|
||||||
|
and the per-layer transfer loop then never reaches the high layer ids.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def test_pool_entries_span_the_highest_global_layer_id(self):
|
||||||
|
# Global ids 0/2/4/6 with holes between them, the shape NemotronH's
|
||||||
|
# cache-ineligible MLP layers produce: 4 mapped layers spanning 7 ids.
|
||||||
|
full_layer_mapping = {0: 0, 6: 1}
|
||||||
|
swa_layer_mapping = {2: 0, 4: 1}
|
||||||
|
|
||||||
|
with (
|
||||||
|
patch(_ASSEMBLER + "build_kv_host_pool"),
|
||||||
|
patch(_ASSEMBLER + "HostPoolGroup"),
|
||||||
|
patch(_ASSEMBLER + "build_pool_entry") as build_pool_entry,
|
||||||
|
):
|
||||||
|
build_hybrid_swa_group(
|
||||||
|
page_size=64,
|
||||||
|
full_kv_pool=MagicMock(),
|
||||||
|
swa_kv_pool=MagicMock(),
|
||||||
|
full_layer_mapping=full_layer_mapping,
|
||||||
|
swa_layer_mapping=swa_layer_mapping,
|
||||||
|
use_mla=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(
|
||||||
|
[
|
||||||
|
c.kwargs["transfer_layer_id_max"]
|
||||||
|
for c in build_pool_entry.call_args_list
|
||||||
|
],
|
||||||
|
[7, 7],
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
@@ -115,7 +115,6 @@ class TestUnifiedRadixHiCacheDispatch(unittest.TestCase):
|
|||||||
self.assertIs(result.cache_controller, cache_controller)
|
self.assertIs(result.cache_controller, cache_controller)
|
||||||
self.assertIs(result.component_host_pools[FULL], kv_host_pool)
|
self.assertIs(result.component_host_pools[FULL], kv_host_pool)
|
||||||
self.assertEqual(result.pools_desc, "KV + INDEXER(k-only)")
|
self.assertEqual(result.pools_desc, "KV + INDEXER(k-only)")
|
||||||
self.assertEqual(result.transfer_layer_num, 8)
|
|
||||||
self.assertEqual(len(result.sidecars), 1)
|
self.assertEqual(len(result.sidecars), 1)
|
||||||
self.assertEqual(result.sidecars[0].pool_name, PoolName.INDEXER)
|
self.assertEqual(result.sidecars[0].pool_name, PoolName.INDEXER)
|
||||||
self.assertEqual(result.sidecars[0].indices_from_pool, PoolName.KV)
|
self.assertEqual(result.sidecars[0].indices_from_pool, PoolName.KV)
|
||||||
@@ -189,7 +188,6 @@ class TestApplyStackResult(unittest.TestCase):
|
|||||||
component_host_pools={FULL: full_host, SWA: swa_host, MAMBA: mamba_host},
|
component_host_pools={FULL: full_host, SWA: swa_host, MAMBA: mamba_host},
|
||||||
sidecars=[sidecar],
|
sidecars=[sidecar],
|
||||||
register_req_to_token_counter=True,
|
register_req_to_token_counter=True,
|
||||||
transfer_layer_num=8,
|
|
||||||
pools_desc="KV + SWA + MAMBA",
|
pools_desc="KV + SWA + MAMBA",
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -221,7 +219,6 @@ class TestApplyStackResult(unittest.TestCase):
|
|||||||
component_host_pools={FULL: MagicMock()},
|
component_host_pools={FULL: MagicMock()},
|
||||||
sidecars=[],
|
sidecars=[],
|
||||||
register_req_to_token_counter=False,
|
register_req_to_token_counter=False,
|
||||||
transfer_layer_num=1,
|
|
||||||
pools_desc="KV",
|
pools_desc="KV",
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user