[HiCache] Fix sparse hybrid transfer layer IDs (#37870)

Co-authored-by: Seokhoon Kang <sh.kang@postech.ac.kr>
This commit is contained in:
Shuwen Wang
2026-09-20 16:24:24 +08:00
committed by GitHub
co-authored by Seokhoon Kang
parent e54009240a
commit 9f3d275940
10 changed files with 161 additions and 103 deletions
@@ -343,8 +343,8 @@ class HiCacheController:
self.host_mem_release_queue: Optional[Queue[torch.Tensor]] = None
self.device = self.mem_pool_device.device
self.layer_num = self.mem_pool_device.layer_num
self.layer_done_counter = LayerDoneCounter(self.layer_num)
self.transfer_layer_id_max = self.mem_pool_device.layer_num
self.layer_done_counter = LayerDoneCounter(self.transfer_layer_id_max)
self.mem_pool_device.register_layer_transfer_counter(self.layer_done_counter)
if write_policy not in [
@@ -961,7 +961,7 @@ class HiCacheController:
self._l2_load_transfers(host_indices, device_indices, pool_transfers),
start_event=producer_event.start_event,
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(
@@ -181,7 +181,7 @@ class HybridCacheController(BaseHiCacheController):
prefetch_threshold: int = 256,
model_name: Optional[str] = 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,
host_memory_mode: str = "cache",
):
@@ -211,11 +211,14 @@ class HybridCacheController(BaseHiCacheController):
enable_storage_metrics=enable_storage_metrics,
host_memory_mode=host_memory_mode,
)
# 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.
if transfer_layer_num is not None and transfer_layer_num != self.layer_num:
self.layer_num = transfer_layer_num
self.layer_done_counter = LayerDoneCounter(self.layer_num)
# Hybrid transfer IDs span every component pool, including holes for
# uncached layers that the anchor pool alone cannot describe.
if (
transfer_layer_id_max is not None
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
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:
continue
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:
continue
@@ -66,10 +66,12 @@ def _evict_mamba_for_device_alloc(cache: UnifiedRadixCache, required_size: int)
def _make_layer_mapper(
layer_mapping: dict[int, int],
transfer_layer_num: int,
transfer_layer_id_max: 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]:
if not 0 <= layer_id < transfer_layer_num:
if not 0 <= layer_id < transfer_layer_id_max:
return None
return layer_mapping.get(layer_id)
@@ -99,7 +101,7 @@ def _with_mtp_layer_mapping(
class _DeepSeekV4LayerMappings(NamedTuple):
transfer_layer_num: int
transfer_layer_id_max: int
full: dict[int, int]
swa: dict[int, int]
c4: dict[int, int]
@@ -111,8 +113,8 @@ class _DeepSeekV4LayerMappings(NamedTuple):
def _resolve_deepseek_v4_layer_mappings(
kvcache: Any,
) -> _DeepSeekV4LayerMappings:
transfer_layer_num = kvcache.end_layer - kvcache.start_layer
full = {layer: layer for layer in range(transfer_layer_num)}
transfer_layer_id_max = kvcache.end_layer - kvcache.start_layer
full = {layer: layer for layer in range(transfer_layer_id_max)}
swa = full.copy() if kvcache.swa_kv_pool is not None else {}
c4, c128, c4_state_global_layers = {}, {}, []
@@ -126,7 +128,7 @@ def _resolve_deepseek_v4_layer_mappings(
c128[local_layer] = item.compress_layer_id
return _DeepSeekV4LayerMappings(
transfer_layer_num=transfer_layer_num,
transfer_layer_id_max=transfer_layer_id_max,
full=full,
swa=swa,
c4=c4,
@@ -196,7 +198,7 @@ def build_pool_entry(
host_pool: Any,
device_pool: Any,
layer_mapping: dict[int, int],
transfer_layer_num: int,
transfer_layer_id_max: int,
is_anchor: bool = False,
host_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,
host_pool=host_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,
host_evict_fn=host_evict_fn,
device_evict_fn=device_evict_fn,
@@ -229,7 +231,7 @@ def build_kv_only_group(
mtp_draft_device_pools: tuple[Any, ...] = (),
) -> HostPoolGroup:
"""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_pool=kv_pool,
page_size=page_size,
@@ -241,7 +243,7 @@ def build_kv_only_group(
if mtp_draft_device_pools:
full_layer_mapping = _with_mtp_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,
draft_layer_num=len(mtp_draft_device_pools),
)
@@ -252,7 +254,8 @@ def build_kv_only_group(
host_pool=kv_host_pool,
device_pool=kv_pool,
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,
packed_draft_device_pools=mtp_draft_device_pools,
)
@@ -276,7 +279,9 @@ def build_hybrid_swa_group(
mtp_swa_device_pools: tuple[Any, ...] = (),
) -> HostPoolGroup:
"""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_pool=full_kv_pool,
page_size=page_size,
@@ -295,7 +300,7 @@ def build_hybrid_swa_group(
if mtp_swa_device_pools:
swa_layer_mapping = _with_mtp_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,
draft_layer_num=len(mtp_swa_device_pools),
)
@@ -306,7 +311,7 @@ def build_hybrid_swa_group(
host_pool=kv_host_pool,
device_pool=full_kv_pool,
layer_mapping=full_layer_mapping,
transfer_layer_num=transfer_layer_num,
transfer_layer_id_max=transfer_layer_id_max,
is_anchor=True,
),
build_pool_entry(
@@ -314,7 +319,7 @@ def build_hybrid_swa_group(
host_pool=swa_host_pool,
device_pool=swa_kv_pool,
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,
device_evict_fn=device_swa_evict_fn,
device_alloc_fn=(
@@ -343,7 +348,7 @@ def build_kv_only_stack(
storage_backend_extra_config: Optional[dict] = None,
enable_storage_metrics: bool = False,
) -> 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(
page_size=params.page_size,
kv_pool=kv_pool,
@@ -367,7 +372,7 @@ def build_kv_only_stack(
prefetch_threshold=prefetch_threshold,
model_name=model_name,
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,
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,
enable_storage_metrics: bool = False,
) -> 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_swa_device_pools = tuple(
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,
model_name=model_name,
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,
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(
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
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,
layer_mapping=layer_mapping,
transfer_layer_num=transfer_layer_num,
transfer_layer_id_max=transfer_layer_id_max,
)
)
return entries
@@ -705,7 +712,7 @@ def _build_dsv4_rope_entry(
layer_mapping: dict[int, int],
num_host_pages: int,
slot_page_size: int,
transfer_layer_num: int,
transfer_layer_id_max: int,
) -> Optional[PoolEntry]:
sibling = _dsv4_rope_sibling(kvcache, ratio)
if sibling is None:
@@ -725,7 +732,7 @@ def _build_dsv4_rope_entry(
),
device_pool=device_pool,
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]:
page_size = params.page_size
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
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,
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(
"DeepSeek V4 SWA KV pool must be PP-stage-local: "
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
# 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,
transfer_layer_start=transfer_layer_num,
target_device_layer_num=transfer_layer_num,
transfer_layer_start=transfer_layer_id_max,
target_device_layer_num=transfer_layer_id_max,
draft_layer_num=len(mtp_swa_device_buffers),
)
@@ -806,7 +813,7 @@ def build_deepseek_v4_hicache_stack(
host_pool=logical_host_pool,
device_pool=kvcache,
layer_mapping=full_layer_mapping,
transfer_layer_num=transfer_layer_num,
transfer_layer_id_max=transfer_layer_id_max,
is_anchor=True,
),
]
@@ -832,7 +839,8 @@ def build_deepseek_v4_hicache_stack(
host_pool=swa_host_pool,
device_pool=kvcache.swa_kv_pool,
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,
device_evict_fn=device_swa_evict_fn,
device_alloc_fn=swa_attn_allocator.alloc,
@@ -863,7 +871,7 @@ def build_deepseek_v4_hicache_stack(
host_pool=c4_host_pool,
device_pool=kvcache.c4_kv_pool,
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):
@@ -882,7 +890,7 @@ def build_deepseek_v4_hicache_stack(
),
device_pool=kvcache.c4_indexer_kv_pool,
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,
num_host_pages=num_host_pages,
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:
entries.append(c4_rope_entry)
@@ -929,14 +937,14 @@ def build_deepseek_v4_hicache_stack(
host_pool=c4_state_host_pool,
device_pool=None,
layer_mapping=c4_state_mapping,
transfer_layer_num=transfer_layer_num,
transfer_layer_id_max=transfer_layer_id_max,
),
build_pool_entry(
name=PoolName.DEEPSEEK_V4_C4_INDEXER_STATE,
host_pool=c4_indexer_state_host_pool,
device_pool=None,
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,
device_pool=kvcache.c128_kv_pool,
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.
# FULL leaf eviction also frees each attached C128 host value.
# GPU keeps its KV-derived sidecar callbacks unset.
@@ -999,13 +1007,15 @@ def build_deepseek_v4_hicache_stack(
layer_mapping=c128_layer_mapping,
num_host_pages=c128_num_host_pages,
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:
entries.append(c128_rope_entry)
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)
@@ -1024,7 +1034,7 @@ def build_deepseek_v4_hicache_stack(
prefetch_threshold=prefetch_threshold,
model_name=model_name,
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,
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,
enable_storage_metrics: bool = False,
) -> 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
from sglang.srt.mem_cache.memory_pool import HybridLinearKVPool
@@ -1071,7 +1083,7 @@ def build_hybrid_mamba_stack(
if mtp_draft_device_pools:
full_layer_mapping = _with_mtp_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,
draft_layer_num=len(mtp_draft_device_pools),
)
@@ -1097,7 +1109,7 @@ def build_hybrid_mamba_stack(
host_pool=kv_host_pool,
device_pool=kv_pool,
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,
packed_draft_device_pools=mtp_draft_device_pools,
),
@@ -1106,7 +1118,7 @@ def build_hybrid_mamba_stack(
host_pool=mamba_host_pool,
device_pool=mamba_pool,
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,
device_evict_fn=device_mamba_evict_fn,
device_alloc_fn=mamba_allocator.alloc,
@@ -1129,7 +1141,7 @@ def build_hybrid_mamba_stack(
prefetch_threshold=prefetch_threshold,
model_name=model_name,
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,
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,
enable_storage_metrics: bool = False,
) -> tuple[HostPoolGroup, HybridCacheController]:
transfer_layer_num = len(
full_layer_mapping | swa_layer_mapping | mamba_layer_mapping
transfer_layer_id_max = (
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
mamba_allocator = params.req_to_token_pool.mamba_allocator
@@ -1198,7 +1215,7 @@ def build_hybrid_mamba_swa_stack(
host_pool=kv_host_pool,
device_pool=full_kv_pool,
layer_mapping=full_layer_mapping,
transfer_layer_num=transfer_layer_num,
transfer_layer_id_max=transfer_layer_id_max,
is_anchor=True,
),
build_pool_entry(
@@ -1206,7 +1223,7 @@ def build_hybrid_mamba_swa_stack(
host_pool=swa_host_pool,
device_pool=swa_kv_pool,
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,
device_evict_fn=device_swa_evict_fn,
device_alloc_fn=swa_attn_allocator.alloc,
@@ -1217,7 +1234,7 @@ def build_hybrid_mamba_swa_stack(
host_pool=mamba_host_pool,
device_pool=mamba_pool,
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,
device_evict_fn=device_mamba_evict_fn,
device_alloc_fn=mamba_allocator.alloc,
@@ -1240,7 +1257,7 @@ def build_hybrid_mamba_swa_stack(
prefetch_threshold=prefetch_threshold,
model_name=model_name,
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,
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,
enable_storage_metrics: bool = False,
) -> tuple[HostPoolGroup, HybridCacheController]:
transfer_layer_num = len(full_layer_mapping)
transfer_layer_id_max = len(full_layer_mapping)
mtp_draft_device_pools = tuple(
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:
full_layer_mapping = _with_mtp_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,
draft_layer_num=len(mtp_draft_device_pools),
)
@@ -1289,7 +1306,7 @@ def build_anchor_sidecar_stack(
host_pool=kv_host_pool,
device_pool=kv_pool,
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,
packed_draft_device_pools=mtp_draft_device_pools,
),
@@ -1298,7 +1315,7 @@ def build_anchor_sidecar_stack(
host_pool=sidecar_host_pool,
device_pool=kv_pool,
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,
),
]
@@ -1318,7 +1335,7 @@ def build_anchor_sidecar_stack(
prefetch_threshold=prefetch_threshold,
model_name=model_name,
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,
host_memory_mode=get_memory().hicache_host_memory_mode,
)
@@ -1401,7 +1418,7 @@ def build_full_draft_pools(
host_pool=draft_host_pool,
device_pool=pool,
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,
device_pool=pool,
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,
device_pool=draft_swa_pool,
layer_mapping=layer_mapping,
transfer_layer_num=host_pool.layer_num,
transfer_layer_id_max=host_pool.layer_num,
)
return [spec], [entry]
@@ -1527,7 +1544,6 @@ class StackBuildResult:
# Mamba state lives in req_to_token_pool, not in kvcache, so its
# layer_transfer_counter has to be wired separately.
register_req_to_token_counter: bool = False
transfer_layer_num: int = 0
pools_desc: str = ""
@@ -1681,7 +1697,6 @@ class _DeepSeekV4Strategy(StackStrategy):
cache_controller=cache_controller,
component_host_pools=component_host_pools,
sidecars=sidecars,
transfer_layer_num=kvcache.end_layer - kvcache.start_layer,
pools_desc="KV + SWA + DeepSeekV4 sidecars",
)
@@ -1739,7 +1754,6 @@ class _MambaStrategy(StackStrategy):
ComponentType.MAMBA: host_pool_group.get_pool(PoolName.MAMBA),
},
register_req_to_token_counter=True,
transfer_layer_num=len(full_layer_mapping | mamba_layer_mapping),
pools_desc="KV + MAMBA",
)
@@ -1807,7 +1821,6 @@ class _SwaStrategy(StackStrategy):
ComponentType.FULL: host_pool_group.get_pool(PoolName.KV),
ComponentType.SWA: host_pool_group.get_pool(PoolName.SWA),
},
transfer_layer_num=len(full_layer_mapping | swa_layer_mapping),
pools_desc="Full + SWA",
)
@@ -1877,9 +1890,6 @@ class _MambaSwaStrategy(StackStrategy):
ComponentType.MAMBA: host_pool_group.get_pool(PoolName.MAMBA),
},
register_req_to_token_counter=True,
transfer_layer_num=len(
full_layer_mapping | swa_layer_mapping | mamba_layer_mapping
),
pools_desc="KV + SWA + MAMBA",
)
@@ -1952,7 +1962,6 @@ class _DsaStrategy(StackStrategy):
indices_from_pool=PoolName.KV,
),
],
transfer_layer_num=len(full_layer_mapping),
pools_desc="KV + INDEXER",
)
@@ -2006,7 +2015,6 @@ class _MiniMaxSparseStrategy(StackStrategy):
ComponentType.FULL: host_pool_group.get_pool(PoolName.KV),
},
sidecars=sidecars,
transfer_layer_num=kvcache.main_pool.layer_num,
pools_desc=pools_desc,
)
@@ -2071,7 +2079,6 @@ class _PlainKvStrategy(StackStrategy):
component_host_pools={
ComponentType.FULL: host_pool_group.get_pool(PoolName.KV),
},
transfer_layer_num=len(full_layer_mapping),
pools_desc="KV",
)
@@ -2128,9 +2135,9 @@ def _apply_stack_result(
)
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.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,
) -> 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).
# 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.
if params.pp_size > 1:
raise NotImplementedError(
@@ -2194,10 +2201,12 @@ def build_minimax_sparse_hicache_stack(
)
main_pool = sparse_pool.main_pool
start_layer = main_pool.start_layer
transfer_layer_num = main_pool.layer_num
# Stage-local keys (0..transfer_layer_num) match the controller's per-layer
transfer_layer_id_max = main_pool.layer_num
# 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.
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_pool=main_pool,
@@ -2210,7 +2219,7 @@ def build_minimax_sparse_hicache_stack(
host_pool=kv_host_pool,
device_pool=main_pool,
layer_mapping=full_layer_mapping,
transfer_layer_num=transfer_layer_num,
transfer_layer_id_max=transfer_layer_id_max,
is_anchor=True,
),
]
@@ -2232,7 +2241,7 @@ def build_minimax_sparse_hicache_stack(
gid - start_layer: sub_id
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,
storage_backend_extra_config=storage_backend_extra_config,
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,
)
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
logger.info(
"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,
main_pool.layer_num,
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
logger.info(
"Attached hybrid DSA pool stack to HiRadixCache: pools=KV + INDEXER, "
"transfer_layer_num=%s",
"transfer_layer_id_max=%s",
len(layer_mapping),
)
except Exception:
@@ -368,7 +368,7 @@ def _build_deepseek_v4_device_pool_group(
)
return DevicePoolGroup(
entries,
mappings.transfer_layer_num,
mappings.transfer_layer_id_max,
page_size,
rank_replicated=True,
)
+2 -2
View File
@@ -75,7 +75,7 @@ class L2TransferEngine:
self,
transfers: list[L2Transfer],
*,
layer_num: int,
transfer_layer_id_max: int,
start_event=None,
on_layer_done=None,
) -> TransferCompletion:
@@ -85,7 +85,7 @@ class L2TransferEngine:
with device_module.stream(self.host_to_device_stream):
start_event.wait(self.host_to_device_stream)
ack_start.record()
for layer_id in range(layer_num):
for layer_id in range(transfer_layer_id_max):
for transfer in transfers:
local_layer_id = (
transfer.layer_mapper(layer_id)
@@ -1527,7 +1527,7 @@ class UnifiedRadixCache(BasePrefixCache):
self.cache_controller._l2_load_transfers(
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()
self.retraction_discard(backup)