From 8e11feb68e057eb5f8fb8ababa97d4b4d9d40f07 Mon Sep 17 00:00:00 2001 From: hjzhang <499213894@qq.com> Date: Thu, 6 Aug 2026 14:31:11 +0800 Subject: [PATCH] [HiCache] Support packed and sidecar draft caches for MTP/EAGLE/DSpark (#30393) Co-authored-by: hjzhang Co-authored-by: Zhangheng Co-authored-by: shuwenn <47200617+alphabetc1@users.noreply.github.com> --- python/sglang/srt/configs/model_config.py | 1 + .../sglang/srt/managers/cache_controller.py | 7 + python/sglang/srt/managers/scheduler.py | 25 +- .../sglang/srt/mem_cache/cache_init_params.py | 2 + .../sglang/srt/mem_cache/hicache_storage.py | 2 + .../hybrid_cache/hybrid_cache_controller.py | 78 ++++- .../hybrid_cache/hybrid_pool_assembler.py | 278 +++++++++++++++++- .../sglang/srt/mem_cache/kv_cache_builder.py | 95 +++--- .../sglang/srt/mem_cache/memory_pool_host.py | 152 ++++++++-- python/sglang/srt/mem_cache/pool_host/base.py | 36 ++- python/sglang/srt/mem_cache/pool_host/mha.py | 195 ++++++++---- python/sglang/srt/mem_cache/pool_host/mla.py | 119 +++++--- .../storage/mooncake_store/mooncake_store.py | 17 ++ .../srt/mem_cache/unified_radix_cache.py | 10 + .../sglang/srt/model_executor/model_runner.py | 1 + .../srt/speculative/base_spec_worker.py | 94 ++++++ .../test_unified_radix_cache_kl_dsv4.py | 62 +++- .../test_unified_radix_cache_kl_mamba.py | 8 + .../test_unified_radix_cache_kl_nightly.py | 16 + 19 files changed, 992 insertions(+), 206 deletions(-) diff --git a/python/sglang/srt/configs/model_config.py b/python/sglang/srt/configs/model_config.py index ded01222c..4b2172937 100644 --- a/python/sglang/srt/configs/model_config.py +++ b/python/sglang/srt/configs/model_config.py @@ -632,6 +632,7 @@ class ModelConfig: self.hf_config.architectures[0] = "MiMoMTP" if is_draft_model and self.hf_config.architectures[0] in MIMO_V2_MODEL_ARCHS: self.hf_config.architectures[0] = "MiMoV2MTP" + self.hf_config.num_nextn_predict_layers = 1 if is_draft_model and self.hf_config.architectures[0] == "Step3p5ForCausalLM": self.hf_config.architectures[0] = "Step3p5MTP" if ( diff --git a/python/sglang/srt/managers/cache_controller.py b/python/sglang/srt/managers/cache_controller.py index beb4b3980..1a41be3dd 100644 --- a/python/sglang/srt/managers/cache_controller.py +++ b/python/sglang/srt/managers/cache_controller.py @@ -274,6 +274,8 @@ class HiCacheController: self.mem_pool_host_draft = None self.draft_page_get_func = None self.draft_page_set_func = None + self.has_mtp_draft = False + self.mtp_draft_device_pools = () # Default storage page IO functions (may be overridden by attach). self.page_get_func = self._generic_page_get @@ -886,6 +888,11 @@ class HiCacheController: # Otherwise this will be deferred until attach_storage_backend(). self._maybe_register_draft_with_storage() + def set_mtp_draft_pools(self, device_pools) -> None: + """Register MTP device pools used for L2 load-back.""" + self.mtp_draft_device_pools = tuple(device_pools) + self.has_mtp_draft = bool(self.mtp_draft_device_pools) + def _maybe_register_draft_with_storage(self) -> None: """Pick the draft L3 IO implementation.""" self.draft_page_get_func = None diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 4f023f7a8..3aed2f55e 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -527,6 +527,11 @@ class Scheduler( tp_group=self.tp_group, pp_group=self.pp_group, enable_hierarchical_cache=self.enable_hierarchical_cache, + hicache_draft_plan=( + self.draft_worker.hicache_draft_plan + if self.draft_worker is not None + else None + ), ) self.is_hybrid_swa = result.is_hybrid_swa self.is_hybrid_ssm = result.is_hybrid_ssm @@ -563,16 +568,6 @@ class Scheduler( else: self.decode_offload_manager = None - # Register draft KV pool (when spec + HiCache co-enabled). - kv_cache_builder.maybe_register_hicache_draft( - tree_cache=self.tree_cache, - draft_worker=self.draft_worker, - spec_algorithm=self.spec_algorithm, - server_args=self.server_args, - enable_hierarchical_cache=self.enable_hierarchical_cache, - page_size=self.page_size, - ) - # Init running status self.init_running_status() @@ -958,6 +953,7 @@ class Scheduler( req_to_token_pool=pool, token_to_kv_pool_allocator=allocator, ) + self.draft_worker.init_hicache_draft_plan() def init_all_attention_backends(self): """Initialize attention backends for all workers.""" @@ -1284,11 +1280,10 @@ class Scheduler( transfer_backend=self.transfer_backend, ) - # todo: should we fix this when enabling mtp or it doesn't matter since we only enable mtp in decode node thus we don't transfer draft kvs between P and D? - draft_token_to_kv_pool = kv_cache_builder.get_draft_kv_pool( - draft_worker=self.draft_worker, - spec_algorithm=self.spec_algorithm, - server_args=self.server_args, + draft_token_to_kv_pool = ( + self.draft_worker.primary_draft_kv_pool + if self.draft_worker is not None + else None ) if self.spec_algorithm.carries_draft_hidden_states(): diff --git a/python/sglang/srt/mem_cache/cache_init_params.py b/python/sglang/srt/mem_cache/cache_init_params.py index c402cec4a..a5eb8122f 100644 --- a/python/sglang/srt/mem_cache/cache_init_params.py +++ b/python/sglang/srt/mem_cache/cache_init_params.py @@ -53,3 +53,5 @@ class CacheInitParams: component_registry_override: Optional[dict[ComponentType, type[TreeComponent]]] = ( None ) + + mtp_draft_device_pools: tuple[object, ...] = () diff --git a/python/sglang/srt/mem_cache/hicache_storage.py b/python/sglang/srt/mem_cache/hicache_storage.py index 28de214ab..c9829e93a 100644 --- a/python/sglang/srt/mem_cache/hicache_storage.py +++ b/python/sglang/srt/mem_cache/hicache_storage.py @@ -73,6 +73,8 @@ class PoolName(str, Enum): # Draft KV pool DRAFT = "draft" + DRAFT_INDEXER = "draft_indexer" + DRAFT_SWA = "draft_swa" def __str__(self) -> str: return self.value 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 191cd8a90..bc19336d4 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 @@ -33,7 +33,8 @@ from sglang.srt.mem_cache.hicache_storage import ( PoolTransfer, PoolTransferResult, ) -from sglang.srt.mem_cache.memory_pool_host import PoolEntry +from sglang.srt.mem_cache.memory_pool_host import HostPoolGroup, PoolEntry +from sglang.srt.mem_cache.pool_host.mha import MHATokenToKVPoolHost from sglang.srt.utils import get_device_module if TYPE_CHECKING: @@ -232,6 +233,15 @@ class HybridCacheController(BaseHiCacheController): for entry in host_pools or []: self.storage_backend.register_mem_host_pool_v2(entry.host_pool, entry.name) + def register_host_pool_entry(self, entry: PoolEntry) -> None: + if not isinstance(self.mem_pool_host, HostPoolGroup): + raise TypeError("Dynamic HiCache sidecars require HostPoolGroup.") + self.mem_pool_host.add_entry(entry) + if not entry.is_primary_index_anchor: + self.extra_host_mem_release_queues.setdefault(entry.name, Queue()) + if self.enable_storage and self.storage_backend is not None: + self.storage_backend.register_mem_host_pool_v2(entry.host_pool, entry.name) + @staticmethod def parse_storage_backend_extra_config( storage_backend_extra_config: Optional[str], @@ -553,9 +563,10 @@ class HybridCacheController(BaseHiCacheController): with device_module.stream(self.load_stream): producer_event.start_event.wait(self.load_stream) ack_start_event.record() + target_device_pool = self.mem_pool_host.anchor_entry.device_pool for i in range(self.layer_num): self.mem_pool_host.load_to_device_per_layer( - self.mem_pool_device, + target_device_pool, host_indices, device_indices, i, @@ -574,6 +585,31 @@ class HybridCacheController(BaseHiCacheController): i, self.io_backend, ) + + # HiCache now supports draft caches through two paths: + # + # - Packed: standard NextN/MTP models (DeepSeek-V3.2, GLM-5.x, + # DeepSeek-V4, MiMo-V2.5) and DeepSeek-V4 DSpark. Draft KV/indexer/SWA + # buffers are appended to the matching target host pools as tail layers + # and share their slot mappings. D2H/H2D therefore moves target and draft + # in the same cache operation; the branch below restores the tail layers. + # + # - Sidecar: standalone EAGLE/EAGLE3 (for example Llama-2/Llama-3.1), + # DFlash (for example Gemma-4), and non-DeepSeek-V4 DSpark. Draft + # KV/indexer/SWA gets a separate host-pool entry sized to its source target + # pool. Its PoolTransfer follows the target KV or SWA indices and is + # attached to the same cache operation. + + if self.has_mtp_draft and i < len(self.mtp_draft_device_pools): + self.mem_pool_host.load_to_device_per_layer( + self.mtp_draft_device_pools[i], + host_indices, + device_indices, + self.layer_num + i, + self.io_backend, + pool_transfers=resolved_pool_transfers, + is_draft=True, + ) producer_event.complete(i) ack_finish_event.record() self._record_transfer_indices_on_stream( @@ -725,16 +761,12 @@ class HybridCacheController(BaseHiCacheController): def _page_backup(self, operation): # MLA KV is replicated across TP ranks and should still be written only - # by TP0. On follower ranks, only the rank-sharded Mamba/KDA pool is - # owned by the rank and must be written here. Do not replicate other - # sidecar pools (for example SWA or indexer state) accidentally. - backup_transfers = operation.pool_transfers - if self.backup_skip: - backup_transfers = [ - transfer - for transfer in operation.pool_transfers or [] - if transfer.name == PoolName.MAMBA - ] + # by TP0. Rank-sharded sidecars still need every TP rank. + backup_transfers = [ + transfer + for transfer in operation.pool_transfers or [] + if self.should_backup(transfer) + ] if backup_transfers: self._resolve_sidecar_derived_pool_transfers(operation) @@ -764,6 +796,28 @@ class HybridCacheController(BaseHiCacheController): len(operation.hash_value) * self.page_size if sidecar_ok else 0 ) + def should_backup(self, transfer: PoolTransfer) -> bool: + if not self.backup_skip: + return True + + # Kimi-K3 Mamba/KDA state is TP-sharded even when the primary MLA KV + # pool is replicated. + if transfer.name == PoolName.MAMBA: + return True + + # Mooncake gives MHA draft and draft-SWA objects rank-specific keys. + # MLA/DeepSeek-V4 draft pools remain TP0-only. + if self.storage_backend_type == "mooncake" and transfer.name in ( + PoolName.DRAFT, + PoolName.DRAFT_SWA, + ): + entry = self.mem_pool_host.entry_map.get(transfer.name) + return entry is not None and isinstance( + entry.host_pool, MHATokenToKVPoolHost + ) + + return False + def backup_thread_func(self): """Back up rank-sharded sidecars on every TP rank. 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 21990fff5..7e28bb048 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 @@ -58,6 +58,19 @@ def _make_layer_mapper( return mapper +def _with_mtp_layer_mapping( + layer_mapping: dict[int, int], + *, + transfer_layer_start: int, + target_device_layer_num: int, + draft_layer_num: int, +) -> dict[int, int]: + return layer_mapping | { + transfer_layer_start + depth: target_device_layer_num + depth + for depth in range(draft_layer_num) + } + + def build_kv_host_pool( *, kv_pool: Any, @@ -66,6 +79,7 @@ def build_kv_host_pool( use_mla: bool, override_kv_cache_dim: Optional[int] = None, host_size: Optional[float] = None, + mtp_draft_device_pools: tuple[Any, ...] = (), pool_label: str = "kv", ): kv_host_pool_cls = ( @@ -74,6 +88,8 @@ def build_kv_host_pool( kwargs = {} if override_kv_cache_dim is not None: kwargs["override_kv_cache_dim"] = override_kv_cache_dim + if mtp_draft_device_pools: + kwargs["mtp_draft_device_pools"] = mtp_draft_device_pools parallel = get_parallel() if parallel.dcp_enabled: assert use_mla, ( @@ -158,14 +174,23 @@ def build_kv_only_stack( server_args=server_args, use_mla=use_mla, override_kv_cache_dim=override_kv_cache_dim, + mtp_draft_device_pools=params.mtp_draft_device_pools, ) + if params.mtp_draft_device_pools: + full_layer_mapping = _with_mtp_layer_mapping( + full_layer_mapping, + transfer_layer_start=transfer_layer_num, + target_device_layer_num=kv_pool.layer_num, + draft_layer_num=len(params.mtp_draft_device_pools), + ) + entries = [ build_pool_entry( name=PoolName.KV, host_pool=kv_host_pool, device_pool=kv_pool, layer_mapping=full_layer_mapping, - transfer_layer_num=transfer_layer_num, + transfer_layer_num=transfer_layer_num + len(params.mtp_draft_device_pools), is_anchor=True, ) ] @@ -188,6 +213,9 @@ def build_kv_only_stack( transfer_layer_num=transfer_layer_num, enable_storage_metrics=enable_storage_metrics, ) + if params.mtp_draft_device_pools: + cache_controller.set_mtp_draft_pools(params.mtp_draft_device_pools) + return host_pool_group, cache_controller @@ -210,11 +238,17 @@ def build_hybrid_swa_stack( enable_storage_metrics: bool = False, ) -> tuple[HostPoolGroup, HybridCacheController]: transfer_layer_num = len(full_layer_mapping | swa_layer_mapping) + # 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 + ) + kv_host_size = swa_host_size = None if server_args.hicache_size > 0: kv_host_size, swa_host_size = _split_hicache_size( server_args.hicache_size, (full_kv_pool, swa_kv_pool) ) + kv_host_pool = build_kv_host_pool( kv_pool=full_kv_pool, page_size=params.page_size, @@ -229,9 +263,18 @@ def build_hybrid_swa_stack( server_args=server_args, use_mla=use_mla, host_size=swa_host_size, + mtp_draft_device_pools=mtp_swa_device_pools, pool_label="swa", ) + if mtp_swa_device_pools: + swa_layer_mapping = _with_mtp_layer_mapping( + swa_layer_mapping, + transfer_layer_start=transfer_layer_num, + target_device_layer_num=swa_kv_pool.layer_num, + draft_layer_num=len(mtp_swa_device_pools), + ) + # For SWA hybrid, the device alloc/free goes through the inner swa_attn_allocator swa_attn_allocator = params.token_to_kv_pool_allocator.swa_attn_allocator entries = [ @@ -248,7 +291,7 @@ def build_hybrid_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_num=transfer_layer_num + len(mtp_swa_device_pools), host_evict_fn=host_swa_evict_fn, device_evict_fn=device_swa_evict_fn, device_alloc_fn=swa_attn_allocator.alloc, @@ -274,6 +317,8 @@ def build_hybrid_swa_stack( transfer_layer_num=transfer_layer_num, enable_storage_metrics=enable_storage_metrics, ) + if mtp_swa_device_pools: + cache_controller.set_mtp_draft_pools(mtp_swa_device_pools) return host_pool_group, cache_controller @@ -332,6 +377,7 @@ def build_deepseek_v4_hicache_stack( full_layer_mapping = {layer_id: layer_id for layer_id in range(transfer_layer_num)} is_unified_kv = getattr(kvcache, "_unified_kv", False) + mtp_swa_device_buffers = [] if is_unified_kv: # unified_kv keeps the SWA ring inside the unified pool and never offloads it, # so there is no separate SWA host pool to map. @@ -346,6 +392,19 @@ def build_deepseek_v4_hicache_stack( swa_layer_mapping = { layer_id: layer_id for layer_id in range(transfer_layer_num) } + # Keep every uncompressed draft SWA layer after the target SWA layers. + # NextN has one layer per pool, while DSpark keeps all stages in one pool. + mtp_swa_device_buffers = [ + buffer + for pool in params.mtp_draft_device_pools + for buffer in pool.swa_kv_pool.kv_buffer + ] + swa_layer_mapping = _with_mtp_layer_mapping( + swa_layer_mapping, + transfer_layer_start=transfer_layer_num, + target_device_layer_num=transfer_layer_num, + draft_layer_num=len(mtp_swa_device_buffers), + ) c4_layer_mapping = {} c128_layer_mapping = {} @@ -390,7 +449,10 @@ def build_deepseek_v4_hicache_stack( if not is_unified_kv: swa_host_pool = DeepSeekV4PagedHostPool( pool_name=str(PoolName.SWA), - device_buffers=kvcache.swa_kv_pool.kv_buffer, + device_buffers=[ + *kvcache.swa_kv_pool.kv_buffer, + *mtp_swa_device_buffers, + ], item_bytes=kvcache.swa_kv_pool.bytes_per_page_padded, num_host_pages=swa_num_host_pages, slot_page_size=kvcache.swa_page_size, @@ -404,7 +466,7 @@ 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, + transfer_layer_num=transfer_layer_num + 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, @@ -542,6 +604,8 @@ def build_deepseek_v4_hicache_stack( transfer_layer_num=transfer_layer_num, enable_storage_metrics=enable_storage_metrics, ) + if mtp_swa_device_buffers: + cache_controller.set_mtp_draft_pools(mtp_swa_device_buffers) return host_pool_group, cache_controller @@ -565,6 +629,9 @@ def build_hybrid_mamba_stack( ) -> tuple[HostPoolGroup, HybridCacheController]: transfer_layer_num = len(full_layer_mapping | mamba_layer_mapping) mamba_allocator = params.req_to_token_pool.mamba_allocator + mtp_draft_device_pools = tuple( + pool.full_kv_pool for pool in params.mtp_draft_device_pools + ) kv_host_size, mamba_host_size = None, 0 if server_args.hicache_size > 0: kv_host_size, mamba_host_size = _split_hicache_size( @@ -576,7 +643,15 @@ def build_hybrid_mamba_stack( server_args=server_args, use_mla=use_mla, host_size=kv_host_size, + mtp_draft_device_pools=mtp_draft_device_pools, ) + if mtp_draft_device_pools: + full_layer_mapping = _with_mtp_layer_mapping( + full_layer_mapping, + transfer_layer_start=transfer_layer_num, + target_device_layer_num=kv_pool.layer_num, + draft_layer_num=len(mtp_draft_device_pools), + ) mamba_host_pool = MambaPoolHost( mamba_pool, server_args.hicache_ratio, @@ -590,7 +665,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, + transfer_layer_num=transfer_layer_num + len(mtp_draft_device_pools), is_anchor=True, ), build_pool_entry( @@ -624,6 +699,8 @@ def build_hybrid_mamba_stack( transfer_layer_num=transfer_layer_num, enable_storage_metrics=enable_storage_metrics, ) + if mtp_draft_device_pools: + cache_controller.set_mtp_draft_pools(mtp_draft_device_pools) return host_pool_group, cache_controller @@ -758,21 +835,33 @@ def build_anchor_sidecar_stack( enable_storage_metrics: bool = False, ) -> tuple[HostPoolGroup, HybridCacheController]: transfer_layer_num = 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 + ) kv_host_pool = build_kv_host_pool( kv_pool=kv_pool, page_size=params.page_size, server_args=server_args, use_mla=use_mla, override_kv_cache_dim=override_kv_cache_dim, + mtp_draft_device_pools=mtp_draft_device_pools, ) sidecar_host_pool = sidecar_host_pool_factory(kv_host_pool) + # Let HostPoolGroup dispatch packed MTP tail layers through the normal path. + if mtp_draft_device_pools: + full_layer_mapping = _with_mtp_layer_mapping( + full_layer_mapping, + transfer_layer_start=transfer_layer_num, + target_device_layer_num=kv_pool.layer_num, + draft_layer_num=len(mtp_draft_device_pools), + ) entries = [ build_pool_entry( name=PoolName.KV, host_pool=kv_host_pool, device_pool=kv_pool, layer_mapping=full_layer_mapping, - transfer_layer_num=transfer_layer_num, + transfer_layer_num=transfer_layer_num + len(mtp_draft_device_pools), is_anchor=True, ), build_pool_entry( @@ -780,7 +869,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, + transfer_layer_num=transfer_layer_num + len(mtp_draft_device_pools), ), ] host_pool_group = HostPoolGroup(entries) @@ -802,9 +891,184 @@ def build_anchor_sidecar_stack( transfer_layer_num=transfer_layer_num, enable_storage_metrics=enable_storage_metrics, ) + if mtp_draft_device_pools: + cache_controller.set_mtp_draft_pools(mtp_draft_device_pools) return host_pool_group, cache_controller +def _build_mha_mla_host_pool( + *, + pool: Any, + host_to_device_ratio: float, + page_size: int, + layout: str, + allocator_type: str, + pool_label: str, +): + from sglang.srt.mem_cache.memory_pool import MHATokenToKVPool + + kwargs = dict( + host_to_device_ratio=host_to_device_ratio, + host_size=0, + page_size=page_size, + layout=layout, + allocator_type=allocator_type, + pool_label=pool_label, + ) + if isinstance(pool, MHATokenToKVPool): + return get_mha_host_pool_cls(pool)(pool, **kwargs) + return MLATokenToKVPoolHost( + pool, + override_kv_cache_dim=pool.kv_cache_dim, + **kwargs, + ) + + +def build_full_draft_pools( + *, + draft_kv_pool: Any, + tree_cache: Any, + server_args: ServerArgs, +) -> tuple[list[SidecarPoolSpec], list[PoolEntry]]: + """Build draft KV/DSA sidecars whose indices follow target full KV.""" + from sglang.srt.mem_cache.memory_pool import DSATokenToKVPool + + pool = draft_kv_pool + if pool.layer_num == 0: + return [], [] + + controller = tree_cache.cache_controller + host_pool_group = controller.mem_pool_host + + draft_host_pool = _build_mha_mla_host_pool( + pool=pool, + host_to_device_ratio=host_pool_group.size / pool.size, + page_size=controller.page_size, + layout=server_args.hicache_mem_layout, + allocator_type=_get_allocator_type(server_args), + pool_label="draft", + ) + draft_layer_mapping = {i: i for i in range(pool.layer_num)} + + specs = [ + SidecarPoolSpec( + pool_name=PoolName.DRAFT, + indices_from_pool=PoolName.KV, + ) + ] + entries = [ + build_pool_entry( + name=PoolName.DRAFT, + host_pool=draft_host_pool, + device_pool=pool, + layer_mapping=draft_layer_mapping, + transfer_layer_num=draft_host_pool.layer_num, + ) + ] + + if isinstance(pool, DSATokenToKVPool) and pool.index_k_with_scale_buffer: + indexer_host_pool = DSAIndexerPoolHost( + pool, + draft_host_pool, + server_args.hicache_mem_layout, + allocator_type=_get_allocator_type(server_args), + ) + specs.append( + SidecarPoolSpec( + pool_name=PoolName.DRAFT_INDEXER, + indices_from_pool=PoolName.KV, + ) + ) + entries.append( + build_pool_entry( + name=PoolName.DRAFT_INDEXER, + host_pool=indexer_host_pool, + device_pool=pool, + layer_mapping=draft_layer_mapping, + transfer_layer_num=indexer_host_pool.layer_num, + ) + ) + + return specs, entries + + +def build_swa_draft_pools( + *, + draft_kv_pool: Any, + tree_cache: Any, + server_args: ServerArgs, +) -> tuple[list[SidecarPoolSpec], list[PoolEntry]]: + """Build a draft SWA sidecar whose indices follow target SWA.""" + draft_swa_pool = draft_kv_pool.swa_kv_pool + if draft_swa_pool is None: + raise NotImplementedError( + "HiCache draft SWA sidecar requires a non-unified draft SWA pool." + ) + if draft_swa_pool.layer_num == 0: + return [], [] + controller = tree_cache.cache_controller + host_pool_group = controller.mem_pool_host + target_swa_host_pool = host_pool_group.entry_map[PoolName.SWA].host_pool + + if isinstance(target_swa_host_pool, DeepSeekV4PagedHostPool): + host_pool = DeepSeekV4PagedHostPool( + pool_name=str(PoolName.DRAFT_SWA), + device_buffers=draft_swa_pool.kv_buffer, + item_bytes=draft_swa_pool.bytes_per_page_padded, + num_host_pages=target_swa_host_pool.num_host_pages, + slot_page_size=draft_swa_pool.page_size, + layout=target_swa_host_pool.layout, + allocator_type=_get_allocator_type(server_args), + ) + else: + host_pool = _build_mha_mla_host_pool( + pool=draft_swa_pool, + host_to_device_ratio=target_swa_host_pool.size / draft_swa_pool.size, + page_size=target_swa_host_pool.page_size, + layout=target_swa_host_pool.layout, + allocator_type=_get_allocator_type(server_args), + pool_label="draft_swa", + ) + + layer_mapping = {i: i for i in range(draft_swa_pool.layer_num)} + spec = SidecarPoolSpec( + pool_name=PoolName.DRAFT_SWA, + indices_from_pool=PoolName.SWA, + hit_policy=PoolHitPolicy.TRAILING_PAGES, + ) + entry = build_pool_entry( + name=PoolName.DRAFT_SWA, + host_pool=host_pool, + device_pool=draft_swa_pool, + layer_mapping=layer_mapping, + transfer_layer_num=host_pool.layer_num, + ) + return [spec], [entry] + + +def build_hicache_draft_sidecars( + *, + draft_device_pools: tuple[Any, ...], + tree_cache: Any, + server_args: ServerArgs, +) -> tuple[list[SidecarPoolSpec], list[PoolEntry]]: + """Compose the full and SWA draft-sidecar paths.""" + from sglang.srt.mem_cache.base_swa_memory_pool import BaseSWAKVPool + + assert len(draft_device_pools) == 1 + draft_kv_pool = draft_device_pools[0] + builder = ( + build_swa_draft_pools + if isinstance(draft_kv_pool, BaseSWAKVPool) + else build_full_draft_pools + ) + return builder( + draft_kv_pool=draft_kv_pool, + tree_cache=tree_cache, + server_args=server_args, + ) + + _COMPONENT_HOST_ATTR: dict[ComponentType, tuple[str, str]] = { ComponentType.FULL: ("full_kv_pool_host", "_full_kv_pool_host"), ComponentType.SWA: ("swa_kv_pool_host", "_swa_kv_pool_host"), diff --git a/python/sglang/srt/mem_cache/kv_cache_builder.py b/python/sglang/srt/mem_cache/kv_cache_builder.py index 2a396aed8..cdcc2b5c7 100644 --- a/python/sglang/srt/mem_cache/kv_cache_builder.py +++ b/python/sglang/srt/mem_cache/kv_cache_builder.py @@ -46,68 +46,69 @@ if TYPE_CHECKING: from sglang.srt.distributed.parallel_state import GroupCoordinator from sglang.srt.distributed.parallel_state_wrapper import ParallelState from sglang.srt.managers.tp_worker import BaseTpWorker - from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache from sglang.srt.server_args import ServerArgs + from sglang.srt.speculative.base_spec_worker import HiCacheDraftPlan from sglang.srt.speculative.spec_info import SpeculativeAlgorithm -def get_draft_kv_pool( - *, - draft_worker: BaseTpWorker, - spec_algorithm: SpeculativeAlgorithm, - server_args: ServerArgs, -): - """Return the draft token-to-KV pool for the current draft worker, - or None when no draft KV pool is available.""" - if draft_worker is None or spec_algorithm.is_ngram(): - return None - - # V2 workers nest the draft runner under `.draft_worker`. - if server_args.enable_multi_layer_eagle: - draft_runner = draft_worker.draft_worker.draft_runner_list[0] - else: - draft_runner = draft_worker.draft_worker.draft_runner - return draft_runner.token_to_kv_pool - - def maybe_register_hicache_draft( *, - tree_cache: BasePrefixCache, - draft_worker: BaseTpWorker, - spec_algorithm: SpeculativeAlgorithm, + tree_cache, + draft_plan: HiCacheDraftPlan, server_args: ServerArgs, - enable_hierarchical_cache: bool, page_size: int, ) -> None: - """Register draft KV pool with HiCacheController for piggyback L2/L3 ops.""" - if not enable_hierarchical_cache: + from sglang.srt.speculative.base_spec_worker import HiCacheDraftMode + + if draft_plan.mode != HiCacheDraftMode.SIDECAR: return - draft_kv_pool = get_draft_kv_pool( - draft_worker=draft_worker, - spec_algorithm=spec_algorithm, + from sglang.srt.mem_cache.unified_radix_cache import UnifiedRadixCache + + if not isinstance(tree_cache, UnifiedRadixCache): + _register_legacy_hicache_draft( + tree_cache=tree_cache, + draft_pool=draft_plan.device_pools[0], + server_args=server_args, + page_size=page_size, + ) + return + + from sglang.srt.mem_cache.hybrid_cache.hybrid_pool_assembler import ( + build_hicache_draft_sidecars, + ) + + specs, entries = build_hicache_draft_sidecars( + draft_device_pools=draft_plan.device_pools, + tree_cache=tree_cache, server_args=server_args, ) - if draft_kv_pool is None: - return + tree_cache.register_hicache_draft_pools(specs, entries) + +def _register_legacy_hicache_draft( + *, + tree_cache, + draft_pool, + server_args: ServerArgs, + page_size: int, +) -> None: from sglang.srt.mem_cache.memory_pool import ( - HybridLinearKVPool, MHATokenToKVPool, MLATokenToKVPool, ) from sglang.srt.mem_cache.pool_host.mha import get_mha_host_pool_cls from sglang.srt.mem_cache.pool_host.mla import MLATokenToKVPoolHost - pool = draft_kv_pool - if isinstance(pool, HybridLinearKVPool): - pool = pool.full_kv_pool + pool = draft_pool + if pool.layer_num == 0: + return # Create host pool for draft with the same slot count as the target host pool, # so that host indices stay 1-to-1 between target and draft KV caches. - primary = tree_cache.cache_controller.mem_pool_host - kw = dict( - host_to_device_ratio=primary.size / pool.size, + primary_host_pool = tree_cache.cache_controller.mem_pool_host + host_pool_kwargs = dict( + host_to_device_ratio=primary_host_pool.size / pool.size, host_size=0, page_size=page_size, layout=server_args.hicache_mem_layout, @@ -115,12 +116,13 @@ def maybe_register_hicache_draft( pool_label="draft", ) if isinstance(pool, MHATokenToKVPool): - draft_host_pool = get_mha_host_pool_cls(pool)(pool, **kw) + draft_host_pool = get_mha_host_pool_cls(pool)(pool, **host_pool_kwargs) elif isinstance(pool, MLATokenToKVPool): - draft_host_pool = MLATokenToKVPoolHost(pool, **kw) + draft_host_pool = MLATokenToKVPoolHost(pool, **host_pool_kwargs) else: logger.warning( - "Draft pool type %s not supported for HiCache, skipping.", + "Draft pool type %s is not supported by the legacy HiCache path; " + "skipping draft KV registration.", type(pool).__name__, ) return @@ -144,6 +146,7 @@ def build_kv_cache( tp_group: GroupCoordinator, pp_group: GroupCoordinator, enable_hierarchical_cache: bool, + hicache_draft_plan: Optional[HiCacheDraftPlan] = None, ) -> KVCacheBuildResult: sliding_window_size: Optional[int] = None full_tokens_per_layer: Optional[int] = None @@ -173,6 +176,7 @@ def build_kv_cache( ) req_to_token_pool, token_to_kv_pool_allocator = tp_worker.get_memory_pool() + mtp_draft_device_pools = tp_worker.model_runner.mtp_draft_device_pools disable_radix_cache = server_args.disable_radix_cache or ( model_config.is_multimodal and uses_transformers_backend @@ -234,6 +238,7 @@ def build_kv_cache( pp_size=ps.pp_size, chunked_prefill_size=effective_chunked_prefill_size, sliding_window_size=sliding_window_size, + mtp_draft_device_pools=mtp_draft_device_pools, ) tree_cache = create_tree_cache( @@ -255,6 +260,14 @@ def build_kv_cache( ) ) + if enable_hierarchical_cache and hicache_draft_plan is not None: + maybe_register_hicache_draft( + tree_cache=tree_cache, + draft_plan=hicache_draft_plan, + server_args=server_args, + page_size=page_size, + ) + embedding_cache_size = envs.SGLANG_VLM_CACHE_SIZE_MB.get() init_mm_embedding_cache(embedding_cache_size * 1024 * 1024) diff --git a/python/sglang/srt/mem_cache/memory_pool_host.py b/python/sglang/srt/mem_cache/memory_pool_host.py index 2670de52c..51aec93d4 100644 --- a/python/sglang/srt/mem_cache/memory_pool_host.py +++ b/python/sglang/srt/mem_cache/memory_pool_host.py @@ -64,7 +64,6 @@ from sglang.srt.mem_cache.pool_host.hisparse import HiSparseHostPoolMixin class MambaPoolHost(HostKVCache): - def __init__( self, device_pool: MambaPool, @@ -432,6 +431,8 @@ class MambaPoolHost(HostKVCache): device_indices, layer_id, io_backend="kernel", + *, + is_draft: bool = False, ): if self.layout in ["page_first", "page_first_direct"]: # no ssm state on conv-only models: nothing to transfer @@ -704,7 +705,14 @@ class LogicalHostPool: pass def load_to_device_per_layer( - self, device_pool, host_indices, device_indices, layer_id, io_backend + self, + device_pool, + host_indices, + device_indices, + layer_id, + io_backend, + *, + is_draft: bool = False, ): pass @@ -988,7 +996,14 @@ class DeepSeekV4PagedHostPool(HiSparseHostPoolMixin, HostKVCache): ) def load_to_device_per_layer( - self, device_pool, host_indices, device_indices, layer_id, io_backend + self, + device_pool, + host_indices, + device_indices, + layer_id, + io_backend, + *, + is_draft: bool = False, ): if not self._has_transfer_indices(host_indices, device_indices): return @@ -1374,7 +1389,14 @@ class DeepSeekV4StateHostPool(HostKVCache): ) def load_to_device_per_layer( - self, device_pool, host_indices, device_indices, layer_id, io_backend + self, + device_pool, + host_indices, + device_indices, + layer_id, + io_backend, + *, + is_draft: bool = False, ): if host_indices is None or device_indices is None: return @@ -1538,6 +1560,19 @@ class HostPoolGroup: self.can_use_write_back_jit = all(child_write_back_jit) self.supports_per_pool_backup_indices = any(child_write_back_jit) + def add_entry(self, entry: PoolEntry) -> None: + if entry.name in self.entry_map: + raise ValueError(f"Host pool {entry.name} is already registered.") + self.entries.append(entry) + self.entry_map[entry.name] = entry + self.can_use_write_back_jit = ( + self.can_use_write_back_jit and entry.host_pool.can_use_write_back_jit + ) + self.supports_per_pool_backup_indices = ( + self.supports_per_pool_backup_indices + or entry.host_pool.can_use_write_back_jit + ) + @property def kv_buffer(self): return self.anchor_entry.host_pool.kv_buffer @@ -1608,17 +1643,20 @@ class HostPoolGroup: layer_id, io_backend, pool_transfers: Optional[list] = None, + *, + is_draft: bool = False, ) -> None: # 1. Anchor (KV) transfer anchor = self.anchor_entry local_layer_id = anchor.layer_mapper(layer_id) if local_layer_id is not None and host_indices.numel() > 0: anchor.host_pool.load_to_device_per_layer( - anchor.device_pool, + device_pool if is_draft else anchor.device_pool, host_indices, device_indices, local_layer_id, io_backend, + is_draft=is_draft, ) # 2. Extra pool transfers @@ -1630,11 +1668,12 @@ class HostPoolGroup: if local_layer_id is None: continue entry.host_pool.load_to_device_per_layer( - entry.device_pool, + device_pool if is_draft else entry.device_pool, transfer.host_indices, transfer.device_indices, local_layer_id, io_backend, + is_draft=is_draft, ) def _backup_uses_cpu_host_indices(self, host_pool, io_backend) -> bool: @@ -1732,7 +1771,9 @@ class DSAIndexerPoolHost(HostKVCache): self.dtype = device_pool.store_dtype self.start_layer = device_pool.start_layer self.end_layer = device_pool.end_layer - self.layer_num = self._effective_host_layer_num() + self.target_layer_num = self._effective_host_layer_num() + self.mtp_draft_device_pools = anchor_host.mtp_draft_device_pools + self.layer_num = self.target_layer_num + len(self.mtp_draft_device_pools) self.index_head_dim = device_pool.index_head_dim self.indexer_quant_block_size = device_pool.quant_block_size @@ -1763,11 +1804,24 @@ class DSAIndexerPoolHost(HostKVCache): f"Requesting {requested_bytes / 1e9:.2f} GB but only have " f"{available_bytes / 1e9:.2f} GB free." ) - logger.info( - "Allocating %.2f GB host memory for DSA indexer (layout=%s).", - requested_bytes / 1e9, - layout, - ) + draft_layer_num = self.layer_num - self.target_layer_num + if draft_layer_num > 0: + logger.info( + "Allocating %.2f GB host memory for DSA indexer (layout=%s), " + "packed MTP layers: " + "target_layers=%d, draft_layers=%d, total_layers=%d.", + requested_bytes / 1e9, + layout, + self.target_layer_num, + draft_layer_num, + self.layer_num, + ) + else: + logger.info( + "Allocating %.2f GB host memory for DSA indexer (layout=%s).", + requested_bytes / 1e9, + layout, + ) self.init_kv_buffer() self.can_use_jit = False self.can_use_write_back_jit = False @@ -1785,8 +1839,12 @@ class DSAIndexerPoolHost(HostKVCache): def init_kv_buffer(self): alloc_func = ALLOC_MEMORY_FUNCS[self.device_pool.device] + device_pools = (self.device_pool, *self.mtp_draft_device_pools) + self.packed_device_index_buffers = [ + buffer for pool in device_pools for buffer in pool.index_k_with_scale_buffer + ] self.index_k_device_ptrs = torch.tensor( - [x.data_ptr() for x in self.device_pool.index_k_with_scale_buffer], + [x.data_ptr() for x in self.packed_device_index_buffers], dtype=torch.uint64, device=self.device_pool.device, ) @@ -1863,11 +1921,20 @@ class DSAIndexerPoolHost(HostKVCache): return host_page_indices, device_page_indices def load_to_device_per_layer( - self, device_pool, host_indices, device_indices, layer_id, io_backend + self, + device_pool, + host_indices, + device_indices, + layer_id, + io_backend, + *, + is_draft: bool = False, ): - if not self._is_device_layer_owned(device_pool, layer_id): + if not is_draft and not self._is_device_layer_owned(device_pool, layer_id): return - host_layer = self._host_layer_index(layer_id) + # MTP draft layers do not participate in CP layer sharding. + host_layer_id = layer_id if is_draft else self._host_layer_index(layer_id) + device_layer_id = 0 if is_draft else layer_id host_page_indices, device_page_indices = self._get_indexer_page_indices( host_indices, device_indices @@ -1876,8 +1943,8 @@ class DSAIndexerPoolHost(HostKVCache): if use_kernel: if self.layout == "layer_first": transfer_kv_per_layer_mla( - src=self.index_k_with_scale_buffer[host_layer], - dst=device_pool.index_k_with_scale_buffer[layer_id], + src=self.index_k_with_scale_buffer[host_layer_id], + dst=device_pool.index_k_with_scale_buffer[device_layer_id], src_indices=host_page_indices, dst_indices=device_page_indices, item_size=self.indexer_page_stride_size, @@ -1885,10 +1952,10 @@ class DSAIndexerPoolHost(HostKVCache): elif self.layout == "page_first": transfer_kv_per_layer_mla_pf_lf( src=self.index_k_with_scale_buffer, - dst=device_pool.index_k_with_scale_buffer[layer_id], + dst=device_pool.index_k_with_scale_buffer[device_layer_id], src_indices=host_page_indices, dst_indices=device_page_indices, - layer_id=host_layer, + layer_id=host_layer_id, item_size=self.indexer_page_stride_size, src_layout_dim=self.indexer_layout_dim, ) @@ -1897,8 +1964,8 @@ class DSAIndexerPoolHost(HostKVCache): elif io_backend == "direct": if self.layout == "layer_first": transfer_kv_direct( - src_layers=[self.index_k_with_scale_buffer[host_layer]], - dst_layers=[device_pool.index_k_with_scale_buffer[layer_id]], + src_layers=[self.index_k_with_scale_buffer[host_layer_id]], + dst_layers=[device_pool.index_k_with_scale_buffer[device_layer_id]], src_indices=host_page_indices, dst_indices=device_page_indices, page_size=1, @@ -1906,10 +1973,10 @@ class DSAIndexerPoolHost(HostKVCache): elif self.layout == "page_first_direct": transfer_kv_per_layer_direct_pf_lf( src_ptrs=[self.index_k_with_scale_buffer], - dst_ptrs=[device_pool.index_k_with_scale_buffer[layer_id]], + dst_ptrs=[device_pool.index_k_with_scale_buffer[device_layer_id]], src_indices=host_page_indices, dst_indices=device_page_indices, - layer_id=host_layer, + layer_id=host_layer_id, page_size=1, ) else: @@ -1918,9 +1985,19 @@ class DSAIndexerPoolHost(HostKVCache): raise ValueError(f"Unsupported IO backend: {io_backend}") def _backup_from_device_per_layer( - self, device_pool, host_indices, device_indices, layer_id, io_backend + self, + device_pool, + host_indices, + device_indices, + layer_id, + io_backend, + *, + is_draft: bool = False, ): - host_layer = self._host_layer_index(layer_id) + # MTP draft layers do not participate in CP layer sharding. + host_layer_id = layer_id if is_draft else self._host_layer_index(layer_id) + device_layer_id = 0 if is_draft else layer_id + host_page_indices, device_page_indices = self._get_indexer_page_indices( host_indices, device_indices ) @@ -1928,8 +2005,8 @@ class DSAIndexerPoolHost(HostKVCache): if use_kernel: if self.layout == "layer_first": transfer_kv_per_layer_mla( - src=device_pool.index_k_with_scale_buffer[layer_id], - dst=self.index_k_with_scale_buffer[host_layer], + src=device_pool.index_k_with_scale_buffer[device_layer_id], + dst=self.index_k_with_scale_buffer[host_layer_id], src_indices=device_page_indices, dst_indices=host_page_indices, item_size=self.indexer_page_stride_size, @@ -1944,8 +2021,8 @@ class DSAIndexerPoolHost(HostKVCache): elif io_backend == "direct": if self.layout == "layer_first": transfer_kv_direct( - src_layers=[device_pool.index_k_with_scale_buffer[layer_id]], - dst_layers=[self.index_k_with_scale_buffer[host_layer]], + src_layers=[device_pool.index_k_with_scale_buffer[device_layer_id]], + dst_layers=[self.index_k_with_scale_buffer[host_layer_id]], src_indices=device_page_indices, dst_indices=host_page_indices, page_size=1, @@ -1966,6 +2043,17 @@ class DSAIndexerPoolHost(HostKVCache): self._backup_from_device_per_layer( device_pool, host_indices, device_indices, layer_id, io_backend ) + for draft_layer_id, draft_device_pool in enumerate( + self.mtp_draft_device_pools + ): + self._backup_from_device_per_layer( + draft_device_pool, + host_indices, + device_indices, + self.device_pool.layer_num + draft_layer_id, + io_backend, + is_draft=True, + ) return host_page_indices, device_page_indices = self._get_indexer_page_indices( @@ -2008,7 +2096,7 @@ class DSAIndexerPoolHost(HostKVCache): elif io_backend == "direct": if self.layout == "layer_first": transfer_kv_direct( - src_layers=device_pool.index_k_with_scale_buffer, + src_layers=self.packed_device_index_buffers, dst_layers=self.index_k_data_refs, src_indices=device_page_indices, dst_indices=host_page_indices, @@ -2016,7 +2104,7 @@ class DSAIndexerPoolHost(HostKVCache): ) elif self.layout == "page_first_direct": transfer_kv_all_layer_direct_lf_pf( - src_ptrs=device_pool.index_k_with_scale_buffer, + src_ptrs=self.packed_device_index_buffers, dst_ptrs=[self.index_k_with_scale_buffer], src_indices=device_page_indices, dst_indices=host_page_indices, diff --git a/python/sglang/srt/mem_cache/pool_host/base.py b/python/sglang/srt/mem_cache/pool_host/base.py index 57dc8c1a2..fe2916b7f 100644 --- a/python/sglang/srt/mem_cache/pool_host/base.py +++ b/python/sglang/srt/mem_cache/pool_host/base.py @@ -150,12 +150,26 @@ class HostKVCache(abc.ABC): f"size of the hierarchical cache." ) else: - logger.info( - "Allocating %s hierarchical KV host pool: %d tokens, %.2f GB host memory.", - pool_label, - self.size, - requested_bytes / 1e9, - ) + draft_layer_num = self.layer_num - self.target_layer_num + if draft_layer_num > 0: + logger.info( + "Allocating %s hierarchical KV host pool: %d tokens, " + "%.2f GB host memory, packed MTP KV layers: " + "target_layers=%d, draft_layers=%d, total_layers=%d.", + pool_label, + self.size, + requested_bytes / 1e9, + self.target_layer_num, + draft_layer_num, + self.layer_num, + ) + else: + logger.info( + "Allocating %s hierarchical KV host pool: %d tokens, %.2f GB host memory.", + pool_label, + self.size, + requested_bytes / 1e9, + ) self.kv_buffer = self.init_kv_buffer() self.fd = getattr(self.allocator, "fd", None) @@ -215,7 +229,6 @@ class HostKVCache(abc.ABC): return start <= layer_id < end def _host_layer_index(self, layer_id: int, device_pool=None) -> int: - """Map a full local device layer id to its compacted host-buffer slot.""" start, _ = self._device_owned_layer_range(device_pool) return layer_id - start @@ -229,7 +242,14 @@ class HostKVCache(abc.ABC): @abc.abstractmethod def load_to_device_per_layer( - self, device_pool, host_indices, device_indices, layer_id, io_backend + self, + device_pool, + host_indices, + device_indices, + layer_id, + io_backend, + *, + is_draft: bool = False, ) -> None: """ Load KV data from the host memory pool to the device memory pool for a specific layer. diff --git a/python/sglang/srt/mem_cache/pool_host/mha.py b/python/sglang/srt/mem_cache/pool_host/mha.py index 150c37b72..2d71a8a08 100644 --- a/python/sglang/srt/mem_cache/pool_host/mha.py +++ b/python/sglang/srt/mem_cache/pool_host/mha.py @@ -2,6 +2,7 @@ from __future__ import annotations import logging import threading +from typing import Sequence import psutil import torch @@ -67,7 +68,8 @@ logger = logging.getLogger(__name__) class MHATokenToKVPoolHost(HostKVCache): - device_pool: MHATokenToKVPool + device_pool: MHATokenToKVPool | None = None + mtp_draft_device_pools: tuple[MHATokenToKVPool, ...] = () def __init__( self, @@ -80,8 +82,11 @@ class MHATokenToKVPoolHost(HostKVCache): device: str = "cpu", allocator_type: str = "default", *, + mtp_draft_device_pools: Sequence[MHATokenToKVPool] = (), pool_label: str = "kv", ): + self.mtp_draft_device_pools = tuple(mtp_draft_device_pools) + self.target_layer_num = device_pool.layer_num super().__init__( device_pool, host_to_device_ratio, @@ -122,12 +127,30 @@ class MHATokenToKVPoolHost(HostKVCache): dtype=torch.uint64, device=self.device_pool.device, ) + if self.mtp_draft_device_pools: + device_pools = (self.device_pool, *self.mtp_draft_device_pools) + self.packed_device_k_data_ptrs = torch.cat( + [pool.k_data_ptrs for pool in device_pools] + ) + self.packed_device_v_data_ptrs = torch.cat( + [pool.v_data_ptrs for pool in device_pools] + ) + self.packed_device_k_buffers = [ + buffer for pool in device_pools for buffer in pool.k_buffer + ] + self.packed_device_v_buffers = [ + buffer for pool in device_pools for buffer in pool.v_buffer + ] + self.packed_device_kv_buffers = ( + self.packed_device_k_buffers + self.packed_device_v_buffers + ) + self.host_kv_data_refs = self.k_data_refs + self.v_data_refs self._init_write_back_staging_buffers() def get_size_per_token(self): self.head_num = self.device_pool.head_num self.head_dim = self.device_pool.head_dim - self.layer_num = self.device_pool.layer_num + self.layer_num = self.target_layer_num + len(self.mtp_draft_device_pools) return self.head_dim * self.head_num * self.layer_num * self.dtype.itemsize * 2 def get_ksize_per_token(self): @@ -219,25 +242,36 @@ class MHATokenToKVPoolHost(HostKVCache): device_indices, layer_id, io_backend, + *, + is_draft: bool = False, ): + if self.device_pool is not None: + if not is_draft and not self._is_device_layer_owned(device_pool, layer_id): + return + # MTP draft layers do not participate in CP layer sharding. + host_layer_id = layer_id if is_draft else self._host_layer_index(layer_id) + device_layer_id = 0 if is_draft else layer_id + else: + host_layer_id = device_layer_id = layer_id + if io_backend == "kernel": if self.layout == "layer_first": if self.can_use_jit: jit_transfer_hicache_one_layer( - k_cache_dst=device_pool.k_buffer[layer_id], - v_cache_dst=device_pool.v_buffer[layer_id], - k_cache_src=self.k_buffer[layer_id], - v_cache_src=self.v_buffer[layer_id], + k_cache_dst=device_pool.k_buffer[device_layer_id], + v_cache_dst=device_pool.v_buffer[device_layer_id], + k_cache_src=self.k_buffer[host_layer_id], + v_cache_src=self.v_buffer[host_layer_id], indices_dst=device_indices, indices_src=host_indices, element_dim=self.element_dim, ) else: transfer_kv_per_layer( - src_k=self.k_buffer[layer_id], - dst_k=device_pool.k_buffer[layer_id], - src_v=self.v_buffer[layer_id], - dst_v=device_pool.v_buffer[layer_id], + src_k=self.k_buffer[host_layer_id], + dst_k=device_pool.k_buffer[device_layer_id], + src_v=self.v_buffer[host_layer_id], + dst_v=device_pool.v_buffer[device_layer_id], src_indices=host_indices, dst_indices=device_indices, item_size=self.token_stride_size, @@ -248,10 +282,10 @@ class MHATokenToKVPoolHost(HostKVCache): # index by layer_id to get a per-layer view with strided layout. # The kernel handles different src/dst strides automatically. jit_transfer_hicache_one_layer( - k_cache_dst=device_pool.k_buffer[layer_id], - v_cache_dst=device_pool.v_buffer[layer_id], - k_cache_src=self.k_data_refs[layer_id], - v_cache_src=self.v_data_refs[layer_id], + k_cache_dst=device_pool.k_buffer[device_layer_id], + v_cache_dst=device_pool.v_buffer[device_layer_id], + k_cache_src=self.k_data_refs[host_layer_id], + v_cache_src=self.v_data_refs[host_layer_id], indices_dst=device_indices, indices_src=host_indices, element_dim=self.element_dim, @@ -259,24 +293,24 @@ class MHATokenToKVPoolHost(HostKVCache): else: transfer_kv_per_layer_pf_lf( src_k=self.k_buffer, - dst_k=device_pool.k_buffer[layer_id], + dst_k=device_pool.k_buffer[device_layer_id], src_v=self.v_buffer, - dst_v=device_pool.v_buffer[layer_id], + dst_v=device_pool.v_buffer[device_layer_id], src_indices=host_indices, dst_indices=device_indices, - layer_id=layer_id, + layer_id=host_layer_id, item_size=self.token_stride_size, src_layout_dim=self.layout_dim, ) elif self.layout == "page_head": transfer_kv_per_layer_ph_lf( src_k=self.k_buffer, - dst_k=device_pool.k_buffer[layer_id], + dst_k=device_pool.k_buffer[device_layer_id], src_v=self.v_buffer, - dst_v=device_pool.v_buffer[layer_id], + dst_v=device_pool.v_buffer[device_layer_id], src_indices=host_indices, dst_indices=device_indices, - layer_id=layer_id, + layer_id=host_layer_id, item_size=self.token_stride_size, src_layout_dim=self.layout_dim, page_size=self.page_size, @@ -287,10 +321,13 @@ class MHATokenToKVPoolHost(HostKVCache): elif io_backend == "direct": if self.layout == "layer_first": transfer_kv_direct( - src_layers=[self.k_buffer[layer_id], self.v_buffer[layer_id]], + src_layers=[ + self.k_buffer[host_layer_id], + self.v_buffer[host_layer_id], + ], dst_layers=[ - device_pool.k_buffer[layer_id], - device_pool.v_buffer[layer_id], + device_pool.k_buffer[device_layer_id], + device_pool.v_buffer[device_layer_id], ], src_indices=host_indices, dst_indices=device_indices, @@ -300,12 +337,12 @@ class MHATokenToKVPoolHost(HostKVCache): transfer_kv_per_layer_direct_pf_lf( src_ptrs=[self.k_buffer, self.v_buffer], dst_ptrs=[ - device_pool.k_buffer[layer_id], - device_pool.v_buffer[layer_id], + device_pool.k_buffer[device_layer_id], + device_pool.v_buffer[device_layer_id], ], src_indices=host_indices, dst_indices=device_indices, - layer_id=layer_id, + layer_id=host_layer_id, page_size=self.page_size, ) else: @@ -313,7 +350,7 @@ class MHATokenToKVPoolHost(HostKVCache): elif io_backend == "kernel_ascend": if self.layout == "page_first_direct": # Ascend-specific: transfer KV data for all layers when layer_id == 0 - if layer_id == 0: + if host_layer_id == 0: transfer_kv_dim_exchange( device_indices=device_indices, host_indices=host_indices, @@ -329,9 +366,31 @@ class MHATokenToKVPoolHost(HostKVCache): else: raise ValueError(f"Unsupported IO backend: {io_backend}") + def _resolve_device_transfer_buffers(self, device_pool): + if self.mtp_draft_device_pools: + return ( + self.packed_device_k_data_ptrs, + self.packed_device_v_data_ptrs, + self.packed_device_k_buffers, + self.packed_device_v_buffers, + ) + return ( + device_pool.k_data_ptrs, + device_pool.v_data_ptrs, + device_pool.k_buffer, + device_pool.v_buffer, + ) + def backup_from_device_all_layer( self, device_pool, host_indices, device_indices, io_backend ): + ( + device_k_data_ptrs, + device_v_data_ptrs, + device_k_buffers, + device_v_buffers, + ) = self._resolve_device_transfer_buffers(device_pool) + device_kv_buffers = device_k_buffers + device_v_buffers if io_backend == "kernel": if self.layout == "layer_first": if self.can_use_jit: @@ -339,8 +398,8 @@ class MHATokenToKVPoolHost(HostKVCache): k_ptr_dst=self.k_data_ptrs, v_ptr_dst=self.v_data_ptrs, indices_dst=host_indices, - k_ptr_src=device_pool.k_data_ptrs, - v_ptr_src=device_pool.v_data_ptrs, + k_ptr_src=device_k_data_ptrs, + v_ptr_src=device_v_data_ptrs, indices_src=device_indices, kv_cache_dst_stride_bytes=self.token_stride_size, kv_cache_src_stride_bytes=self.token_stride_size, @@ -348,9 +407,9 @@ class MHATokenToKVPoolHost(HostKVCache): ) else: transfer_kv_all_layer( - src_k_layers=device_pool.k_data_ptrs, + src_k_layers=device_k_data_ptrs, dst_k_layers=self.k_data_ptrs, - src_v_layers=device_pool.v_data_ptrs, + src_v_layers=device_v_data_ptrs, dst_v_layers=self.v_data_ptrs, src_indices=device_indices, dst_indices=host_indices, @@ -360,8 +419,8 @@ class MHATokenToKVPoolHost(HostKVCache): elif self.layout == "page_first": if self.can_use_write_back_jit: jit_transfer_hicache_all_layer_staged_lf_pf( - k_ptr_src=device_pool.k_data_ptrs, - v_ptr_src=device_pool.v_data_ptrs, + k_ptr_src=device_k_data_ptrs, + v_ptr_src=device_v_data_ptrs, src_indices=device_indices, dst_indices=host_indices, staging_k=self.staging_k_buffer, @@ -372,9 +431,9 @@ class MHATokenToKVPoolHost(HostKVCache): ) else: transfer_kv_all_layer_lf_pf( - src_k_layers=device_pool.k_data_ptrs, + src_k_layers=device_k_data_ptrs, dst_k=self.k_buffer, - src_v_layers=device_pool.v_data_ptrs, + src_v_layers=device_v_data_ptrs, dst_v=self.v_buffer, src_indices=device_indices, dst_indices=host_indices, @@ -384,9 +443,9 @@ class MHATokenToKVPoolHost(HostKVCache): ) elif self.layout == "page_head": transfer_kv_all_layer_lf_ph( - src_k_layers=device_pool.k_data_ptrs, + src_k_layers=device_k_data_ptrs, dst_k=self.k_buffer, - src_v_layers=device_pool.v_data_ptrs, + src_v_layers=device_v_data_ptrs, dst_v=self.v_buffer, src_indices=device_indices, dst_indices=host_indices, @@ -401,15 +460,15 @@ class MHATokenToKVPoolHost(HostKVCache): elif io_backend == "direct": if self.layout == "layer_first": transfer_kv_direct( - src_layers=device_pool.k_buffer + device_pool.v_buffer, - dst_layers=self.k_data_refs + self.v_data_refs, + src_layers=device_kv_buffers, + dst_layers=self.host_kv_data_refs, src_indices=device_indices, dst_indices=host_indices, page_size=self.page_size, ) elif self.layout == "page_first_direct": transfer_kv_all_layer_direct_lf_pf( - src_ptrs=device_pool.k_buffer + device_pool.v_buffer, + src_ptrs=device_kv_buffers, dst_ptrs=[self.k_buffer, self.v_buffer], src_indices=device_indices, dst_indices=host_indices, @@ -734,7 +793,14 @@ class MHATokenToKOnlyPoolHost(HostKVCache): return [self.k_buffer] def load_to_device_per_layer( - self, device_pool, host_indices, device_indices, layer_id, io_backend + self, + device_pool, + host_indices, + device_indices, + layer_id, + io_backend, + *, + is_draft: bool = False, ): if io_backend == "kernel": if self.layout == "layer_first": @@ -995,7 +1061,7 @@ class AsymmetricMHATokenToKVPoolHost(MHATokenToKVPoolHost): def get_size_per_token(self): self.head_num = self.device_pool.head_num self.head_dim = self.device_pool.head_dim - self.layer_num = self.device_pool.layer_num + self.layer_num = self.target_layer_num + len(self.mtp_draft_device_pools) self.v_head_dim = self.device_pool.v_head_dim return ( (self.head_dim + self.v_head_dim) @@ -1080,7 +1146,18 @@ class AsymmetricMHATokenToKVPoolHost(MHATokenToKVPoolHost): device_indices, layer_id, io_backend, + *, + is_draft: bool = False, ): + if self.device_pool is not None: + if not is_draft and not self._is_device_layer_owned(device_pool, layer_id): + return + # MTP draft layers do not participate in CP layer sharding. + host_layer_id = layer_id if is_draft else self._host_layer_index(layer_id) + device_layer_id = 0 if is_draft else layer_id + else: + host_layer_id = device_layer_id = layer_id + if io_backend == "kernel": if self.layout != "page_first": raise ValueError( @@ -1089,19 +1166,19 @@ class AsymmetricMHATokenToKVPoolHost(MHATokenToKVPoolHost): ) transfer_kv_per_layer_mla_pf_lf( src=self.k_buffer, - dst=device_pool.k_buffer[layer_id], + dst=device_pool.k_buffer[device_layer_id], src_indices=host_indices, dst_indices=device_indices, - layer_id=layer_id, + layer_id=host_layer_id, item_size=self._k_token_stride_size(), src_layout_dim=self._k_layout_dim(), ) transfer_kv_per_layer_mla_pf_lf( src=self.v_buffer, - dst=device_pool.v_buffer[layer_id], + dst=device_pool.v_buffer[device_layer_id], src_indices=host_indices, dst_indices=device_indices, - layer_id=layer_id, + layer_id=host_layer_id, item_size=self._v_token_stride_size(), src_layout_dim=self._v_layout_dim(), ) @@ -1114,18 +1191,18 @@ class AsymmetricMHATokenToKVPoolHost(MHATokenToKVPoolHost): ) transfer_kv_per_layer_direct_pf_lf( src_ptrs=[self.k_buffer], - dst_ptrs=[device_pool.k_buffer[layer_id]], + dst_ptrs=[device_pool.k_buffer[device_layer_id]], src_indices=host_indices, dst_indices=device_indices, - layer_id=layer_id, + layer_id=host_layer_id, page_size=self.page_size, ) transfer_kv_per_layer_direct_pf_lf( src_ptrs=[self.v_buffer], - dst_ptrs=[device_pool.v_buffer[layer_id]], + dst_ptrs=[device_pool.v_buffer[device_layer_id]], src_indices=host_indices, dst_indices=device_indices, - layer_id=layer_id, + layer_id=host_layer_id, page_size=self.page_size, ) else: @@ -1137,6 +1214,12 @@ class AsymmetricMHATokenToKVPoolHost(MHATokenToKVPoolHost): def backup_from_device_all_layer( self, device_pool, host_indices, device_indices, io_backend ): + ( + device_k_data_ptrs, + device_v_data_ptrs, + device_k_buffers, + device_v_buffers, + ) = self._resolve_device_transfer_buffers(device_pool) if io_backend == "kernel": if self.layout != "page_first": raise ValueError( @@ -1145,7 +1228,7 @@ class AsymmetricMHATokenToKVPoolHost(MHATokenToKVPoolHost): ) if self.can_use_write_back_jit: jit_transfer_hicache_all_layer_mla_staged_lf_pf( - ptr_src=device_pool.k_data_ptrs, + ptr_src=device_k_data_ptrs, src_indices=device_indices, dst_indices=host_indices, staging=self.staging_k_buffer, @@ -1153,7 +1236,7 @@ class AsymmetricMHATokenToKVPoolHost(MHATokenToKVPoolHost): page_size=self.page_size, ) jit_transfer_hicache_all_layer_mla_staged_lf_pf( - ptr_src=device_pool.v_data_ptrs, + ptr_src=device_v_data_ptrs, src_indices=device_indices, dst_indices=host_indices, staging=self.staging_v_buffer, @@ -1162,7 +1245,7 @@ class AsymmetricMHATokenToKVPoolHost(MHATokenToKVPoolHost): ) else: transfer_kv_all_layer_mla_lf_pf( - src_layers=device_pool.k_data_ptrs, + src_layers=device_k_data_ptrs, dst=self.k_buffer, src_indices=device_indices, dst_indices=host_indices, @@ -1171,7 +1254,7 @@ class AsymmetricMHATokenToKVPoolHost(MHATokenToKVPoolHost): num_layers=self.layer_num, ) transfer_kv_all_layer_mla_lf_pf( - src_layers=device_pool.v_data_ptrs, + src_layers=device_v_data_ptrs, dst=self.v_buffer, src_indices=device_indices, dst_indices=host_indices, @@ -1187,14 +1270,14 @@ class AsymmetricMHATokenToKVPoolHost(MHATokenToKVPoolHost): "'page_first_direct'." ) transfer_kv_all_layer_direct_lf_pf( - src_ptrs=device_pool.k_buffer, + src_ptrs=device_k_buffers, dst_ptrs=[self.k_buffer], src_indices=device_indices, dst_indices=host_indices, page_size=self.page_size, ) transfer_kv_all_layer_direct_lf_pf( - src_ptrs=device_pool.v_buffer, + src_ptrs=device_v_buffers, dst_ptrs=[self.v_buffer], src_indices=device_indices, dst_indices=host_indices, diff --git a/python/sglang/srt/mem_cache/pool_host/mla.py b/python/sglang/srt/mem_cache/pool_host/mla.py index e31152072..d1b65c2db 100644 --- a/python/sglang/srt/mem_cache/pool_host/mla.py +++ b/python/sglang/srt/mem_cache/pool_host/mla.py @@ -1,7 +1,7 @@ from __future__ import annotations import logging -from typing import Optional +from typing import Optional, Sequence import torch @@ -50,6 +50,7 @@ logger = logging.getLogger(__name__) class MLATokenToKVPoolHost(HiSparseHostPoolMixin, HostKVCache): device_pool: MLATokenToKVPool + mtp_draft_device_pools: tuple[MLATokenToKVPool, ...] = () def __init__( self, @@ -62,12 +63,14 @@ class MLATokenToKVPoolHost(HiSparseHostPoolMixin, HostKVCache): device: str = "cpu", allocator_type: str = "default", override_kv_cache_dim: Optional[int] = None, + mtp_draft_device_pools: Sequence[MLATokenToKVPool] = (), dcp_size: int = 1, dcp_rank: int = 0, *, pool_label: str = "kv", ): self.override_kv_cache_dim = override_kv_cache_dim + self.mtp_draft_device_pools = tuple(mtp_draft_device_pools) super().__init__( device_pool, host_to_device_ratio, @@ -101,6 +104,14 @@ class MLATokenToKVPoolHost(HiSparseHostPoolMixin, HostKVCache): dtype=torch.uint64, device=self.device_pool.device, ) + if self.mtp_draft_device_pools: + device_pools = (self.device_pool, *self.mtp_draft_device_pools) + self.packed_device_data_ptrs = torch.cat( + [pool.data_ptrs for pool in device_pools] + ) + self.packed_device_kv_buffers = [ + buffer for pool in device_pools for buffer in pool.kv_buffer + ] self._init_write_back_staging_buffers() def get_contiguous_buf_infos(self): @@ -114,7 +125,8 @@ class MLATokenToKVPoolHost(HiSparseHostPoolMixin, HostKVCache): def get_size_per_token(self): self.kv_lora_rank = self.device_pool.kv_lora_rank self.qk_rope_head_dim = self.device_pool.qk_rope_head_dim - self.layer_num = self._effective_host_layer_num() + self.target_layer_num = self._effective_host_layer_num() + self.layer_num = self.target_layer_num + len(self.mtp_draft_device_pools) self.kv_cache_dim = self.override_kv_cache_dim or ( self.kv_lora_rank + self.qk_rope_head_dim ) @@ -229,28 +241,37 @@ class MLATokenToKVPoolHost(HiSparseHostPoolMixin, HostKVCache): ) def load_to_device_per_layer( - self, device_pool, host_indices, device_indices, layer_id, io_backend + self, + device_pool, + host_indices, + device_indices, + layer_id, + io_backend, + *, + is_draft: bool = False, ): - if not self._is_device_layer_owned(device_pool, layer_id): + if not is_draft and not self._is_device_layer_owned(device_pool, layer_id): return host_indices = self.dcp_kernel_indices(host_indices) device_indices = self.dcp_kernel_indices(device_indices) - host_layer = self._host_layer_index(layer_id) + # MTP draft layers do not participate in CP layer sharding. + host_layer_id = layer_id if is_draft else self._host_layer_index(layer_id) + device_layer_id = 0 if is_draft else layer_id if io_backend == "kernel": if self.layout == "layer_first": if self.can_use_jit: jit_transfer_hicache_one_layer_mla( - cache_dst=device_pool.kv_buffer[layer_id], - cache_src=self.kv_buffer[host_layer], + cache_dst=device_pool.kv_buffer[device_layer_id], + cache_src=self.kv_buffer[host_layer_id], indices_dst=device_indices, indices_src=host_indices, element_dim=self.kv_cache_dim, ) else: transfer_kv_per_layer_mla( - src=self.kv_buffer[host_layer], - dst=device_pool.kv_buffer[layer_id], + src=self.kv_buffer[host_layer_id], + dst=device_pool.kv_buffer[device_layer_id], src_indices=host_indices, dst_indices=device_indices, item_size=self.token_stride_size, @@ -258,8 +279,8 @@ class MLATokenToKVPoolHost(HiSparseHostPoolMixin, HostKVCache): elif self.layout == "page_first": if self.can_use_jit: jit_transfer_hicache_one_layer_mla( - cache_dst=device_pool.kv_buffer[layer_id], - cache_src=self.data_refs[host_layer], + cache_dst=device_pool.kv_buffer[device_layer_id], + cache_src=self.data_refs[host_layer_id], indices_dst=device_indices, indices_src=host_indices, element_dim=self.kv_cache_dim, @@ -267,10 +288,10 @@ class MLATokenToKVPoolHost(HiSparseHostPoolMixin, HostKVCache): else: transfer_kv_per_layer_mla_pf_lf( src=self.kv_buffer, - dst=device_pool.kv_buffer[layer_id], + dst=device_pool.kv_buffer[device_layer_id], src_indices=host_indices, dst_indices=device_indices, - layer_id=host_layer, + layer_id=host_layer_id, item_size=self.token_stride_size, src_layout_dim=self.layout_dim, ) @@ -279,8 +300,8 @@ class MLATokenToKVPoolHost(HiSparseHostPoolMixin, HostKVCache): elif io_backend == "direct": if self.layout == "layer_first": transfer_kv_direct( - src_layers=[self.kv_buffer[host_layer]], - dst_layers=[device_pool.kv_buffer[layer_id]], + src_layers=[self.kv_buffer[host_layer_id]], + dst_layers=[device_pool.kv_buffer[device_layer_id]], src_indices=host_indices, dst_indices=device_indices, page_size=self.page_size, @@ -288,10 +309,10 @@ class MLATokenToKVPoolHost(HiSparseHostPoolMixin, HostKVCache): elif self.layout == "page_first_direct": transfer_kv_per_layer_direct_pf_lf( src_ptrs=[self.kv_buffer], - dst_ptrs=[device_pool.kv_buffer[layer_id]], + dst_ptrs=[device_pool.kv_buffer[device_layer_id]], src_indices=host_indices, dst_indices=device_indices, - layer_id=host_layer, + layer_id=host_layer_id, page_size=self.page_size, ) else: @@ -299,7 +320,7 @@ class MLATokenToKVPoolHost(HiSparseHostPoolMixin, HostKVCache): elif io_backend == "kernel_ascend": if self.layout == "page_first_kv_split": # Ascend-specific: transfer KV data for all layers when layer_id == 0 - if layer_id == 0: + if device_layer_id == 0: transfer_kv_dim_exchange( device_indices=device_indices, host_indices=host_indices, @@ -318,24 +339,34 @@ class MLATokenToKVPoolHost(HiSparseHostPoolMixin, HostKVCache): raise ValueError(f"Unsupported IO backend: {io_backend}") def _backup_from_device_per_layer( - self, device_pool, host_indices, device_indices, layer_id, io_backend + self, + device_pool, + host_indices, + device_indices, + layer_id, + io_backend, + *, + is_draft: bool = False, ): # Indices arrive already translated by backup_from_device_all_layer. - host_layer = self._host_layer_index(layer_id) + # MTP draft layers do not participate in CP layer sharding. + host_layer_id = layer_id if is_draft else self._host_layer_index(layer_id) + device_layer_id = 0 if is_draft else layer_id + if io_backend == "kernel": if self.layout == "layer_first": if self.can_use_jit: jit_transfer_hicache_one_layer_mla( - cache_dst=self.kv_buffer[host_layer], - cache_src=device_pool.kv_buffer[layer_id], + cache_dst=self.kv_buffer[host_layer_id], + cache_src=device_pool.kv_buffer[device_layer_id], indices_dst=host_indices, indices_src=device_indices, element_dim=self.kv_cache_dim, ) else: transfer_kv_per_layer_mla( - src=device_pool.kv_buffer[layer_id], - dst=self.kv_buffer[host_layer], + src=device_pool.kv_buffer[device_layer_id], + dst=self.kv_buffer[host_layer_id], src_indices=device_indices, dst_indices=host_indices, item_size=self.token_stride_size, @@ -343,8 +374,8 @@ class MLATokenToKVPoolHost(HiSparseHostPoolMixin, HostKVCache): elif self.layout == "page_first": if self.can_use_jit: jit_transfer_hicache_one_layer_mla( - cache_dst=self.data_refs[host_layer], - cache_src=device_pool.kv_buffer[layer_id], + cache_dst=self.data_refs[host_layer_id], + cache_src=device_pool.kv_buffer[device_layer_id], indices_dst=host_indices, indices_src=device_indices, element_dim=self.kv_cache_dim, @@ -361,8 +392,8 @@ class MLATokenToKVPoolHost(HiSparseHostPoolMixin, HostKVCache): elif io_backend == "direct": if self.layout == "layer_first": transfer_kv_direct( - src_layers=[device_pool.kv_buffer[layer_id]], - dst_layers=[self.kv_buffer[host_layer]], + src_layers=[device_pool.kv_buffer[device_layer_id]], + dst_layers=[self.kv_buffer[host_layer_id]], src_indices=device_indices, dst_indices=host_indices, page_size=self.page_size, @@ -377,6 +408,11 @@ class MLATokenToKVPoolHost(HiSparseHostPoolMixin, HostKVCache): f"Layer-sharded HiCache backup does not support IO backend: {io_backend}" ) + def _resolve_device_transfer_buffers(self, device_pool): + if self.mtp_draft_device_pools: + return self.packed_device_data_ptrs, self.packed_device_kv_buffers + return device_pool.data_ptrs, device_pool.kv_buffer + def backup_from_device_all_layer( self, device_pool, host_indices, device_indices, io_backend ): @@ -387,15 +423,30 @@ class MLATokenToKVPoolHost(HiSparseHostPoolMixin, HostKVCache): self._backup_from_device_per_layer( device_pool, host_indices, device_indices, layer_id, io_backend ) + for draft_layer_id, draft_device_pool in enumerate( + self.mtp_draft_device_pools + ): + self._backup_from_device_per_layer( + draft_device_pool, + host_indices, + device_indices, + self.device_pool.layer_num + draft_layer_id, + io_backend, + is_draft=True, + ) return + device_data_ptrs, device_kv_buffers = self._resolve_device_transfer_buffers( + device_pool + ) + if io_backend == "kernel": if self.layout == "layer_first": if self.can_use_jit: jit_transfer_hicache_all_layer_mla( ptr_dst=self.data_ptrs, indices_dst=host_indices, - ptr_src=device_pool.data_ptrs, + ptr_src=device_data_ptrs, indices_src=device_indices, cache_dst_stride_bytes=self.token_stride_size, cache_src_stride_bytes=self.token_stride_size, @@ -403,7 +454,7 @@ class MLATokenToKVPoolHost(HiSparseHostPoolMixin, HostKVCache): ) else: transfer_kv_all_layer_mla( - src_layers=device_pool.data_ptrs, + src_layers=device_data_ptrs, dst_layers=self.data_ptrs, src_indices=device_indices, dst_indices=host_indices, @@ -413,7 +464,7 @@ class MLATokenToKVPoolHost(HiSparseHostPoolMixin, HostKVCache): elif self.layout == "page_first": if self.can_use_write_back_jit: jit_transfer_hicache_all_layer_mla_staged_lf_pf( - ptr_src=device_pool.data_ptrs, + ptr_src=device_data_ptrs, src_indices=device_indices, dst_indices=host_indices, staging=self.staging_buffer, @@ -422,7 +473,7 @@ class MLATokenToKVPoolHost(HiSparseHostPoolMixin, HostKVCache): ) else: transfer_kv_all_layer_mla_lf_pf( - src_layers=device_pool.data_ptrs, + src_layers=device_data_ptrs, dst=self.kv_buffer, src_indices=device_indices, dst_indices=host_indices, @@ -435,7 +486,7 @@ class MLATokenToKVPoolHost(HiSparseHostPoolMixin, HostKVCache): elif io_backend == "direct": if self.layout == "layer_first": transfer_kv_direct( - src_layers=device_pool.kv_buffer, + src_layers=device_kv_buffers, dst_layers=self.data_refs, src_indices=device_indices, dst_indices=host_indices, @@ -443,7 +494,7 @@ class MLATokenToKVPoolHost(HiSparseHostPoolMixin, HostKVCache): ) elif self.layout == "page_first_direct": transfer_kv_all_layer_direct_lf_pf( - src_ptrs=device_pool.kv_buffer, + src_ptrs=device_kv_buffers, dst_ptrs=[self.kv_buffer], src_indices=device_indices, dst_indices=host_indices, diff --git a/python/sglang/srt/mem_cache/storage/mooncake_store/mooncake_store.py b/python/sglang/srt/mem_cache/storage/mooncake_store/mooncake_store.py index 0051f1241..53bb5a8ac 100644 --- a/python/sglang/srt/mem_cache/storage/mooncake_store/mooncake_store.py +++ b/python/sglang/srt/mem_cache/storage/mooncake_store/mooncake_store.py @@ -772,8 +772,25 @@ class MooncakeStore(HiCacheStorage, MooncakeBaseStore): f"_{self.mha_suffix}_{PoolName.DRAFT}_k", f"_{self.mha_suffix}_{PoolName.DRAFT}_v", ] + elif pool_name == PoolName.DRAFT_SWA: + from sglang.srt.mem_cache.memory_pool_host import ( + DeepSeekV4PagedHostPool, + ) + from sglang.srt.mem_cache.pool_host.mha import MHATokenToKVPoolHost + + if isinstance( + host_pool, + (DeepSeekV4PagedHostPool, MLATokenToKVPoolHost), + ): + suffixes = [f"_{self.mla_suffix}_{pool_name}"] + elif isinstance(host_pool, MHATokenToKVPoolHost): + suffixes = [ + f"_{self.mha_suffix}_{pool_name}_k", + f"_{self.mha_suffix}_{pool_name}_v", + ] elif pool_name in ( PoolName.INDEXER, + PoolName.DRAFT_INDEXER, PoolName.DEEPSEEK_V4_C4, PoolName.DEEPSEEK_V4_C4_INDEXER, PoolName.DEEPSEEK_V4_C128, diff --git a/python/sglang/srt/mem_cache/unified_radix_cache.py b/python/sglang/srt/mem_cache/unified_radix_cache.py index a85a7157e..6fe61bade 100644 --- a/python/sglang/srt/mem_cache/unified_radix_cache.py +++ b/python/sglang/srt/mem_cache/unified_radix_cache.py @@ -78,6 +78,7 @@ if TYPE_CHECKING: from sglang.srt.mem_cache.hybrid_cache.hybrid_cache_controller import ( PrefetchOperation, ) + from sglang.srt.mem_cache.memory_pool_host import PoolEntry from sglang.srt.server_args import ServerArgs @@ -395,6 +396,15 @@ class UnifiedRadixCache(BasePrefixCache): def register_sidecar_pool(self, spec: SidecarPoolSpec) -> None: self.sidecar_pool_specs.append(spec) + def register_hicache_draft_pools( + self, specs: list[SidecarPoolSpec], entries: list[PoolEntry] + ) -> None: + if self.cache_controller is None: + raise RuntimeError("HiCache controller is not attached.") + for spec, entry in zip(specs, entries, strict=True): + self.cache_controller.register_host_pool_entry(entry) + self.register_sidecar_pool(spec) + def release_host_resources(self) -> None: if self.host_pool_group is not None: self.host_pool_group.destroy() diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index 12bbd4fba..cceb6bfb1 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -335,6 +335,7 @@ class ModelRunner: self.page_size = server_args.page_size self.req_to_token_pool = req_to_token_pool self.token_to_kv_pool_allocator = token_to_kv_pool_allocator + self.mtp_draft_device_pools = () self.is_hybrid_swa = model_config.is_hybrid_swa self.is_hybrid_swa_compress = model_config.is_hybrid_swa_compress self.use_mla_backend = self.model_config.attention_arch == AttentionArch.MLA diff --git a/python/sglang/srt/speculative/base_spec_worker.py b/python/sglang/srt/speculative/base_spec_worker.py index 8af7540ed..259f9847f 100644 --- a/python/sglang/srt/speculative/base_spec_worker.py +++ b/python/sglang/srt/speculative/base_spec_worker.py @@ -1,6 +1,8 @@ from __future__ import annotations from abc import ABC, abstractmethod +from dataclasses import dataclass +from enum import Enum from typing import TYPE_CHECKING, Optional import torch @@ -18,6 +20,38 @@ if TYPE_CHECKING: ) from sglang.srt.managers.tp_worker import TpModelWorker from sglang.srt.model_executor.model_runner import ModelRunner + from sglang.srt.speculative.spec_info import SpeculativeAlgorithm + + +class HiCacheDraftMode(str, Enum): + NONE = "none" + PACKED = "packed" + SIDECAR = "sidecar" + + +@dataclass(frozen=True, slots=True) +class HiCacheDraftPlan: + mode: HiCacheDraftMode = HiCacheDraftMode.NONE + device_pools: tuple[object, ...] = () + + +def _can_pack_hicache_mtp( + spec_algorithm: SpeculativeAlgorithm, + draft_runners: tuple[ModelRunner, ...], +) -> bool: + is_nextn_mtp = ( + spec_algorithm.is_eagle() + and not spec_algorithm.is_eagle3() + and all( + runner.model_config.num_nextn_predict_layers for runner in draft_runners + ) + ) + is_dspark_dsv4 = ( + spec_algorithm.is_dspark() + and draft_runners[0].model_config.hf_config.architectures[0] + == "DeepseekV4ForCausalLMDSpark" + ) + return is_nextn_mtp or is_dspark_dsv4 class EagleDraftWorkerBase(ABC): @@ -111,10 +145,34 @@ class EagleDraftWorkerBase(ABC): class BaseSpecWorker(ABC): + _hicache_draft_plan = HiCacheDraftPlan() + def __init__(self) -> None: self._additional_graph_memory_usage: dict[str, float] = {} self._additional_graph_time_usage: dict[str, float] = {} + @property + def hicache_draft_plan(self) -> HiCacheDraftPlan: + return self._hicache_draft_plan + + def _draft_model_runners(self) -> tuple[ModelRunner, ...]: + spec_algorithm = self.target_worker.model_runner.spec_algorithm + draft_worker = self.draft_worker + if ( + draft_worker is None + or spec_algorithm.is_ngram() + or spec_algorithm.is_frozen_kv_mtp() + ): + return () + if spec_algorithm.is_dflash_family(): + return (draft_worker.model_runner,) + return tuple(draft_worker.draft_runners) + + @property + def primary_draft_kv_pool(self) -> Optional[object]: + draft_runners = self._draft_model_runners() + return draft_runners[0].token_to_kv_pool if draft_runners else None + @property def target_worker(self) -> TpModelWorker: return self._target_worker @@ -173,6 +231,42 @@ class BaseSpecWorker(ABC): # TODO: move this method to BaseTpWorker and call through self.model_runner pass + def _build_hicache_draft_plan(self) -> HiCacheDraftPlan: + target_model_runner = self.target_worker.model_runner + target_model_runner.mtp_draft_device_pools = () + spec_algorithm = target_model_runner.spec_algorithm + if not self.server_args.enable_hierarchical_cache: + return HiCacheDraftPlan() + + draft_runners = self._draft_model_runners() + if not draft_runners: + return HiCacheDraftPlan() + draft_pools = tuple(runner.token_to_kv_pool for runner in draft_runners) + if ( + "InklingForConditionalGenerationMTP" + in draft_runners[0].model_config.hf_config.architectures + ): + raise NotImplementedError( + "HiCache does not support Inkling MTP draft state yet." + ) + + if _can_pack_hicache_mtp(spec_algorithm, draft_runners): + target_model_runner.mtp_draft_device_pools = draft_pools + return HiCacheDraftPlan( + mode=HiCacheDraftMode.PACKED, + device_pools=draft_pools, + ) + + return HiCacheDraftPlan( + mode=HiCacheDraftMode.SIDECAR, + # Preserve the legacy non-packed HiCache behavior: multi-layer + # EAGLE registers only the first draft runner as the sidecar. + device_pools=draft_pools[:1], + ) + + def init_hicache_draft_plan(self) -> None: + self._hicache_draft_plan = self._build_hicache_draft_plan() + def alloc_memory_pool( self, memory_pool_config=None, diff --git a/test/registered/radix_cache/unified_radix_tree/test_unified_radix_cache_kl_dsv4.py b/test/registered/radix_cache/unified_radix_tree/test_unified_radix_cache_kl_dsv4.py index 05960c090..7e840e94a 100644 --- a/test/registered/radix_cache/unified_radix_tree/test_unified_radix_cache_kl_dsv4.py +++ b/test/registered/radix_cache/unified_radix_tree/test_unified_radix_cache_kl_dsv4.py @@ -19,9 +19,10 @@ from sglang.test.test_utils import ( ) DSV4_FLASH_MODEL = "sgl-project/DeepSeek-V4-Flash-FP8" +DSV4_DSPARK_MODEL = "deepseek-ai/DeepSeek-V4-Flash-DSpark" DSV4_FLASH_LAUNCH_TIMEOUT = 3600 -register_cuda_ci(est_time=1000, stage="extra-b", runner_config="4-gpu-h100") +register_cuda_ci(est_time=1500, stage="extra-b", runner_config="4-gpu-h100") def _assert_dsv4_decode_cached_tokens(result, history_len, output_len, label): @@ -330,5 +331,64 @@ class TestUnifiedDeepSeekV4FlashEagleHiCacheL3(AccuracyTwoPassMixin, CustomTestC self.assertEqual(cached_details.get("storage_backend"), "HiCacheFile") +class TestUnifiedDeepSeekV4FlashDSparkHiCacheL3( + TestUnifiedDeepSeekV4FlashEagleHiCacheL3 +): + """DeepSeek V4 Flash DSpark + HiCache L3 should load from storage.""" + + @classmethod + def setUpClass(cls): + cls.model = DSV4_DSPARK_MODEL + cls.base_url = DEFAULT_URL_FOR_TEST + cls.hicache_dir = tempfile.mkdtemp(prefix="hicache_l3_dspark_dsv4_") + cls.process = popen_launch_server( + cls.model, + cls.base_url, + timeout=DSV4_FLASH_LAUNCH_TIMEOUT, + other_args=[ + "--trust-remote-code", + "--tp-size", + "4", + "--attention-backend", + "compressed", + "--page-size", + str(cls.page_size), + "--chunked-prefill-size", + "8192", + "--mem-fraction-static", + "0.95", + "--disable-shared-experts-fusion", + "--enable-hierarchical-cache", + "--hicache-ratio", + "2", + "--hicache-write-policy", + "write_through", + "--hicache-storage-prefetch-policy", + "wait_complete", + "--hicache-io-backend", + "kernel", + "--hicache-mem-layout", + "page_first", + "--hicache-storage-backend", + "file", + "--enable-cache-report", + "--swa-full-tokens-ratio", + "0.25", + "--max-total-tokens", + "20000", + "--max-running-requests", + "4", + "--moe-runner-backend", + "marlin", + "--speculative-algorithm", + "DSPARK", + ], + env={ + "SGLANG_ENABLE_UNIFIED_RADIX_TREE": "1", + "SGLANG_HICACHE_FILE_BACKEND_STORAGE_DIR": cls.hicache_dir, + }, + ) + + if __name__ == "__main__": unittest.main() diff --git a/test/registered/radix_cache/unified_radix_tree/test_unified_radix_cache_kl_mamba.py b/test/registered/radix_cache/unified_radix_tree/test_unified_radix_cache_kl_mamba.py index 3ee63f27f..cc6b27dcc 100644 --- a/test/registered/radix_cache/unified_radix_tree/test_unified_radix_cache_kl_mamba.py +++ b/test/registered/radix_cache/unified_radix_tree/test_unified_radix_cache_kl_mamba.py @@ -179,6 +179,14 @@ class TestUnifiedMambaHiCacheL3(AccuracyTwoPassMixin, CustomTestCase): "--max-mamba-cache-size", "500", "--weight-loader-prefetch-checkpoints", + "--speculative-algorithm", + "NEXTN", + "--speculative-num-steps", + "3", + "--speculative-eagle-topk", + "1", + "--speculative-num-draft-tokens", + "4", ], env={ "SGLANG_ENABLE_UNIFIED_RADIX_TREE": "1", diff --git a/test/registered/radix_cache/unified_radix_tree/test_unified_radix_cache_kl_nightly.py b/test/registered/radix_cache/unified_radix_tree/test_unified_radix_cache_kl_nightly.py index b6f039c27..7f1dc2175 100644 --- a/test/registered/radix_cache/unified_radix_tree/test_unified_radix_cache_kl_nightly.py +++ b/test/registered/radix_cache/unified_radix_tree/test_unified_radix_cache_kl_nightly.py @@ -175,6 +175,14 @@ class TestGLM5HiRadixCacheL3Accuracy(AccuracyTwoPassMixin, CustomTestCase): "page_first", "--hicache-storage-backend", "file", + "--speculative-algorithm", + "EAGLE", + "--speculative-num-steps", + "3", + "--speculative-eagle-topk", + "1", + "--speculative-num-draft-tokens", + "4", ], env={ "SGLANG_HICACHE_FILE_BACKEND_STORAGE_DIR": cls.hicache_dir, @@ -223,6 +231,14 @@ class TestGLM5UnifiedRadixCacheL3Accuracy(AccuracyTwoPassMixin, CustomTestCase): "page_first", "--hicache-storage-backend", "file", + "--speculative-algorithm", + "EAGLE", + "--speculative-num-steps", + "3", + "--speculative-eagle-topk", + "1", + "--speculative-num-draft-tokens", + "4", ], env={ "SGLANG_HICACHE_FILE_BACKEND_STORAGE_DIR": cls.hicache_dir,