diff --git a/python/sglang/srt/managers/cache_controller.py b/python/sglang/srt/managers/cache_controller.py index 0cb2f607e..3d5b5594a 100644 --- a/python/sglang/srt/managers/cache_controller.py +++ b/python/sglang/srt/managers/cache_controller.py @@ -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( diff --git a/python/sglang/srt/mem_cache/hybrid_cache/hybrid_cache_controller.py b/python/sglang/srt/mem_cache/hybrid_cache/hybrid_cache_controller.py index eb19e6e3e..798646d86 100644 --- a/python/sglang/srt/mem_cache/hybrid_cache/hybrid_cache_controller.py +++ b/python/sglang/srt/mem_cache/hybrid_cache/hybrid_cache_controller.py @@ -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 diff --git a/python/sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py b/python/sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py index 9cb0d6c1b..1e91f0364 100644 --- a/python/sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py +++ b/python/sglang/srt/mem_cache/hybrid_cache/hybrid_pool_assembler.py @@ -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: diff --git a/python/sglang/srt/mem_cache/hybrid_cache/linker_pool_assembler.py b/python/sglang/srt/mem_cache/hybrid_cache/linker_pool_assembler.py index 0df7c4413..9050152d2 100644 --- a/python/sglang/srt/mem_cache/hybrid_cache/linker_pool_assembler.py +++ b/python/sglang/srt/mem_cache/hybrid_cache/linker_pool_assembler.py @@ -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, ) diff --git a/python/sglang/srt/mem_cache/l2_transfer.py b/python/sglang/srt/mem_cache/l2_transfer.py index f80afd9c7..abb4fa898 100644 --- a/python/sglang/srt/mem_cache/l2_transfer.py +++ b/python/sglang/srt/mem_cache/l2_transfer.py @@ -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) diff --git a/python/sglang/srt/mem_cache/unified_radix_cache.py b/python/sglang/srt/mem_cache/unified_radix_cache.py index d2fa662cf..e78514b1d 100644 --- a/python/sglang/srt/mem_cache/unified_radix_cache.py +++ b/python/sglang/srt/mem_cache/unified_radix_cache.py @@ -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) diff --git a/test/registered/unit/mem_cache/test_dsv4_hicache_l2.py b/test/registered/unit/mem_cache/test_dsv4_hicache_l2.py index bff743331..b151a166a 100644 --- a/test/registered/unit/mem_cache/test_dsv4_hicache_l2.py +++ b/test/registered/unit/mem_cache/test_dsv4_hicache_l2.py @@ -254,7 +254,7 @@ class TestDSV4PoolAssembly(CustomTestCase): pp_cache_group=None, ) mappings = assembler._DeepSeekV4LayerMappings( - transfer_layer_num=1, + transfer_layer_id_max=1, full={0: 0}, swa={}, c4={0: 0}, diff --git a/test/registered/unit/mem_cache/test_hicache_staged_write_back_dispatch.py b/test/registered/unit/mem_cache/test_hicache_staged_write_back_dispatch.py index 119d2766e..df1c6c74e 100644 --- a/test/registered/unit/mem_cache/test_hicache_staged_write_back_dispatch.py +++ b/test/registered/unit/mem_cache/test_hicache_staged_write_back_dispatch.py @@ -238,7 +238,7 @@ class TestHiCacheStagedWriteBackDispatch(CustomTestCase): start_event=object(), finish_event=object(), timing_enabled=False ) 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 = [] self.assertEqual(HybridCacheController.start_loading(controller), 0) @@ -344,7 +344,9 @@ class TestHiCacheStagedWriteBackDispatch(CustomTestCase): layer_mapper={1: 0, 3: 1}.get, ) 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( [ @@ -369,7 +371,7 @@ class TestHiCacheStagedWriteBackDispatch(CustomTestCase): anchor_entry=entry, entry_map={entry.name: entry}, ) - controller.layer_num = 2 + controller.transfer_layer_id_max = 2 self.assertEqual( 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.assertTrue(transfers[1].is_draft) 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( [ call.args[3] diff --git a/test/registered/unit/mem_cache/test_hybrid_pool_assembler.py b/test/registered/unit/mem_cache/test_hybrid_pool_assembler.py index 2e958ffea..4181e6f27 100644 --- a/test/registered/unit/mem_cache/test_hybrid_pool_assembler.py +++ b/test/registered/unit/mem_cache/test_hybrid_pool_assembler.py @@ -15,6 +15,7 @@ from sglang.srt.mem_cache.hybrid_cache.hybrid_pool_assembler import ( _split_hicache_size, _SwaStrategy, build_full_draft_pools, + build_hybrid_swa_group, ) from sglang.srt.mem_cache.memory_pool import HybridLinearKVPool from sglang.test.ci.ci_register import register_cpu_ci @@ -153,7 +154,10 @@ class TestHybridStageLayerMappings(CustomTestCase): with patch.object( hybrid_pool_assembler, builder_name, - return_value=(MagicMock(), object()), + return_value=( + MagicMock(), + SimpleNamespace(transfer_layer_id_max=4), + ), ) as build_stack: result = strategy_cls().build( cache=SimpleNamespace(page_size=1), @@ -169,7 +173,7 @@ class TestHybridStageLayerMappings(CustomTestCase): build_stack.call_args.kwargs[f"{name}_layer_mapping"], mapping, ) - self.assertEqual(result.transfer_layer_num, 4) + self.assertEqual(result.cache_controller.transfer_layer_id_max, 4) self.assertEqual( 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) +_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__": unittest.main() diff --git a/test/registered/unit/mem_cache/test_unified_radix_hicache_dispatch.py b/test/registered/unit/mem_cache/test_unified_radix_hicache_dispatch.py index cbab1e06e..f13234817 100644 --- a/test/registered/unit/mem_cache/test_unified_radix_hicache_dispatch.py +++ b/test/registered/unit/mem_cache/test_unified_radix_hicache_dispatch.py @@ -115,7 +115,6 @@ class TestUnifiedRadixHiCacheDispatch(unittest.TestCase): self.assertIs(result.cache_controller, cache_controller) self.assertIs(result.component_host_pools[FULL], kv_host_pool) self.assertEqual(result.pools_desc, "KV + INDEXER(k-only)") - self.assertEqual(result.transfer_layer_num, 8) self.assertEqual(len(result.sidecars), 1) self.assertEqual(result.sidecars[0].pool_name, PoolName.INDEXER) 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}, sidecars=[sidecar], register_req_to_token_counter=True, - transfer_layer_num=8, pools_desc="KV + SWA + MAMBA", ) @@ -221,7 +219,6 @@ class TestApplyStackResult(unittest.TestCase): component_host_pools={FULL: MagicMock()}, sidecars=[], register_req_to_token_counter=False, - transfer_layer_num=1, pools_desc="KV", )