[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:
co-authored by
hjzhang
Zhangheng
shuwenn
parent
c84ddc0e76
commit
8e11feb68e
@@ -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
|
||||
|
||||
@@ -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"),
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user