[HiCache] Support packed and sidecar draft caches for MTP/EAGLE/DSpark (#30393)

Co-authored-by: hjzhang <hjzhang89.gmail.com>
Co-authored-by: Zhangheng <hzh0425@apache.org>
Co-authored-by: shuwenn <47200617+alphabetc1@users.noreply.github.com>
This commit is contained in:
hjzhang
2026-08-06 14:31:11 +08:00
committed by GitHub
co-authored by hjzhang Zhangheng shuwenn
parent c84ddc0e76
commit 8e11feb68e
19 changed files with 992 additions and 206 deletions
@@ -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 (
@@ -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
+10 -15
View File
@@ -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():
@@ -53,3 +53,5 @@ class CacheInitParams:
component_registry_override: Optional[dict[ComponentType, type[TreeComponent]]] = (
None
)
mtp_draft_device_pools: tuple[object, ...] = ()
@@ -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
@@ -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.
@@ -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"),
+54 -41
View File
@@ -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)
+120 -32
View File
@@ -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,
+28 -8
View File
@@ -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.
+139 -56
View File
@@ -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,
+85 -34
View File
@@ -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,
@@ -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,
@@ -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()
@@ -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
@@ -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,