[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.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,
) )
+2 -2
View File
@@ -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",
) )