[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"
|
self.hf_config.architectures[0] = "MiMoMTP"
|
||||||
if is_draft_model and self.hf_config.architectures[0] in MIMO_V2_MODEL_ARCHS:
|
if is_draft_model and self.hf_config.architectures[0] in MIMO_V2_MODEL_ARCHS:
|
||||||
self.hf_config.architectures[0] = "MiMoV2MTP"
|
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":
|
if is_draft_model and self.hf_config.architectures[0] == "Step3p5ForCausalLM":
|
||||||
self.hf_config.architectures[0] = "Step3p5MTP"
|
self.hf_config.architectures[0] = "Step3p5MTP"
|
||||||
if (
|
if (
|
||||||
|
|||||||
@@ -274,6 +274,8 @@ class HiCacheController:
|
|||||||
self.mem_pool_host_draft = None
|
self.mem_pool_host_draft = None
|
||||||
self.draft_page_get_func = None
|
self.draft_page_get_func = None
|
||||||
self.draft_page_set_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).
|
# Default storage page IO functions (may be overridden by attach).
|
||||||
self.page_get_func = self._generic_page_get
|
self.page_get_func = self._generic_page_get
|
||||||
@@ -886,6 +888,11 @@ class HiCacheController:
|
|||||||
# Otherwise this will be deferred until attach_storage_backend().
|
# Otherwise this will be deferred until attach_storage_backend().
|
||||||
self._maybe_register_draft_with_storage()
|
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:
|
def _maybe_register_draft_with_storage(self) -> None:
|
||||||
"""Pick the draft L3 IO implementation."""
|
"""Pick the draft L3 IO implementation."""
|
||||||
self.draft_page_get_func = None
|
self.draft_page_get_func = None
|
||||||
|
|||||||
@@ -527,6 +527,11 @@ class Scheduler(
|
|||||||
tp_group=self.tp_group,
|
tp_group=self.tp_group,
|
||||||
pp_group=self.pp_group,
|
pp_group=self.pp_group,
|
||||||
enable_hierarchical_cache=self.enable_hierarchical_cache,
|
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_swa = result.is_hybrid_swa
|
||||||
self.is_hybrid_ssm = result.is_hybrid_ssm
|
self.is_hybrid_ssm = result.is_hybrid_ssm
|
||||||
@@ -563,16 +568,6 @@ class Scheduler(
|
|||||||
else:
|
else:
|
||||||
self.decode_offload_manager = None
|
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
|
# Init running status
|
||||||
self.init_running_status()
|
self.init_running_status()
|
||||||
|
|
||||||
@@ -958,6 +953,7 @@ class Scheduler(
|
|||||||
req_to_token_pool=pool,
|
req_to_token_pool=pool,
|
||||||
token_to_kv_pool_allocator=allocator,
|
token_to_kv_pool_allocator=allocator,
|
||||||
)
|
)
|
||||||
|
self.draft_worker.init_hicache_draft_plan()
|
||||||
|
|
||||||
def init_all_attention_backends(self):
|
def init_all_attention_backends(self):
|
||||||
"""Initialize attention backends for all workers."""
|
"""Initialize attention backends for all workers."""
|
||||||
@@ -1284,11 +1280,10 @@ class Scheduler(
|
|||||||
transfer_backend=self.transfer_backend,
|
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 = (
|
||||||
draft_token_to_kv_pool = kv_cache_builder.get_draft_kv_pool(
|
self.draft_worker.primary_draft_kv_pool
|
||||||
draft_worker=self.draft_worker,
|
if self.draft_worker is not None
|
||||||
spec_algorithm=self.spec_algorithm,
|
else None
|
||||||
server_args=self.server_args,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
if self.spec_algorithm.carries_draft_hidden_states():
|
if self.spec_algorithm.carries_draft_hidden_states():
|
||||||
|
|||||||
@@ -53,3 +53,5 @@ class CacheInitParams:
|
|||||||
component_registry_override: Optional[dict[ComponentType, type[TreeComponent]]] = (
|
component_registry_override: Optional[dict[ComponentType, type[TreeComponent]]] = (
|
||||||
None
|
None
|
||||||
)
|
)
|
||||||
|
|
||||||
|
mtp_draft_device_pools: tuple[object, ...] = ()
|
||||||
|
|||||||
@@ -73,6 +73,8 @@ class PoolName(str, Enum):
|
|||||||
|
|
||||||
# Draft KV pool
|
# Draft KV pool
|
||||||
DRAFT = "draft"
|
DRAFT = "draft"
|
||||||
|
DRAFT_INDEXER = "draft_indexer"
|
||||||
|
DRAFT_SWA = "draft_swa"
|
||||||
|
|
||||||
def __str__(self) -> str:
|
def __str__(self) -> str:
|
||||||
return self.value
|
return self.value
|
||||||
|
|||||||
@@ -33,7 +33,8 @@ from sglang.srt.mem_cache.hicache_storage import (
|
|||||||
PoolTransfer,
|
PoolTransfer,
|
||||||
PoolTransferResult,
|
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
|
from sglang.srt.utils import get_device_module
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
@@ -232,6 +233,15 @@ class HybridCacheController(BaseHiCacheController):
|
|||||||
for entry in host_pools or []:
|
for entry in host_pools or []:
|
||||||
self.storage_backend.register_mem_host_pool_v2(entry.host_pool, entry.name)
|
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
|
@staticmethod
|
||||||
def parse_storage_backend_extra_config(
|
def parse_storage_backend_extra_config(
|
||||||
storage_backend_extra_config: Optional[str],
|
storage_backend_extra_config: Optional[str],
|
||||||
@@ -553,9 +563,10 @@ class HybridCacheController(BaseHiCacheController):
|
|||||||
with device_module.stream(self.load_stream):
|
with device_module.stream(self.load_stream):
|
||||||
producer_event.start_event.wait(self.load_stream)
|
producer_event.start_event.wait(self.load_stream)
|
||||||
ack_start_event.record()
|
ack_start_event.record()
|
||||||
|
target_device_pool = self.mem_pool_host.anchor_entry.device_pool
|
||||||
for i in range(self.layer_num):
|
for i in range(self.layer_num):
|
||||||
self.mem_pool_host.load_to_device_per_layer(
|
self.mem_pool_host.load_to_device_per_layer(
|
||||||
self.mem_pool_device,
|
target_device_pool,
|
||||||
host_indices,
|
host_indices,
|
||||||
device_indices,
|
device_indices,
|
||||||
i,
|
i,
|
||||||
@@ -574,6 +585,31 @@ class HybridCacheController(BaseHiCacheController):
|
|||||||
i,
|
i,
|
||||||
self.io_backend,
|
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)
|
producer_event.complete(i)
|
||||||
ack_finish_event.record()
|
ack_finish_event.record()
|
||||||
self._record_transfer_indices_on_stream(
|
self._record_transfer_indices_on_stream(
|
||||||
@@ -725,16 +761,12 @@ class HybridCacheController(BaseHiCacheController):
|
|||||||
|
|
||||||
def _page_backup(self, operation):
|
def _page_backup(self, operation):
|
||||||
# MLA KV is replicated across TP ranks and should still be written only
|
# 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
|
# by TP0. Rank-sharded sidecars still need every TP rank.
|
||||||
# owned by the rank and must be written here. Do not replicate other
|
backup_transfers = [
|
||||||
# sidecar pools (for example SWA or indexer state) accidentally.
|
transfer
|
||||||
backup_transfers = operation.pool_transfers
|
for transfer in operation.pool_transfers or []
|
||||||
if self.backup_skip:
|
if self.should_backup(transfer)
|
||||||
backup_transfers = [
|
]
|
||||||
transfer
|
|
||||||
for transfer in operation.pool_transfers or []
|
|
||||||
if transfer.name == PoolName.MAMBA
|
|
||||||
]
|
|
||||||
|
|
||||||
if backup_transfers:
|
if backup_transfers:
|
||||||
self._resolve_sidecar_derived_pool_transfers(operation)
|
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
|
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):
|
def backup_thread_func(self):
|
||||||
"""Back up rank-sharded sidecars on every TP rank.
|
"""Back up rank-sharded sidecars on every TP rank.
|
||||||
|
|
||||||
|
|||||||
@@ -58,6 +58,19 @@ def _make_layer_mapper(
|
|||||||
return 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(
|
def build_kv_host_pool(
|
||||||
*,
|
*,
|
||||||
kv_pool: Any,
|
kv_pool: Any,
|
||||||
@@ -66,6 +79,7 @@ def build_kv_host_pool(
|
|||||||
use_mla: bool,
|
use_mla: bool,
|
||||||
override_kv_cache_dim: Optional[int] = None,
|
override_kv_cache_dim: Optional[int] = None,
|
||||||
host_size: Optional[float] = None,
|
host_size: Optional[float] = None,
|
||||||
|
mtp_draft_device_pools: tuple[Any, ...] = (),
|
||||||
pool_label: str = "kv",
|
pool_label: str = "kv",
|
||||||
):
|
):
|
||||||
kv_host_pool_cls = (
|
kv_host_pool_cls = (
|
||||||
@@ -74,6 +88,8 @@ def build_kv_host_pool(
|
|||||||
kwargs = {}
|
kwargs = {}
|
||||||
if override_kv_cache_dim is not None:
|
if override_kv_cache_dim is not None:
|
||||||
kwargs["override_kv_cache_dim"] = override_kv_cache_dim
|
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()
|
parallel = get_parallel()
|
||||||
if parallel.dcp_enabled:
|
if parallel.dcp_enabled:
|
||||||
assert use_mla, (
|
assert use_mla, (
|
||||||
@@ -158,14 +174,23 @@ def build_kv_only_stack(
|
|||||||
server_args=server_args,
|
server_args=server_args,
|
||||||
use_mla=use_mla,
|
use_mla=use_mla,
|
||||||
override_kv_cache_dim=override_kv_cache_dim,
|
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 = [
|
entries = [
|
||||||
build_pool_entry(
|
build_pool_entry(
|
||||||
name=PoolName.KV,
|
name=PoolName.KV,
|
||||||
host_pool=kv_host_pool,
|
host_pool=kv_host_pool,
|
||||||
device_pool=kv_pool,
|
device_pool=kv_pool,
|
||||||
layer_mapping=full_layer_mapping,
|
layer_mapping=full_layer_mapping,
|
||||||
transfer_layer_num=transfer_layer_num,
|
transfer_layer_num=transfer_layer_num + len(params.mtp_draft_device_pools),
|
||||||
is_anchor=True,
|
is_anchor=True,
|
||||||
)
|
)
|
||||||
]
|
]
|
||||||
@@ -188,6 +213,9 @@ def build_kv_only_stack(
|
|||||||
transfer_layer_num=transfer_layer_num,
|
transfer_layer_num=transfer_layer_num,
|
||||||
enable_storage_metrics=enable_storage_metrics,
|
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
|
return host_pool_group, cache_controller
|
||||||
|
|
||||||
|
|
||||||
@@ -210,11 +238,17 @@ def build_hybrid_swa_stack(
|
|||||||
enable_storage_metrics: bool = False,
|
enable_storage_metrics: bool = False,
|
||||||
) -> tuple[HostPoolGroup, HybridCacheController]:
|
) -> tuple[HostPoolGroup, HybridCacheController]:
|
||||||
transfer_layer_num = len(full_layer_mapping | swa_layer_mapping)
|
transfer_layer_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
|
kv_host_size = swa_host_size = None
|
||||||
if server_args.hicache_size > 0:
|
if server_args.hicache_size > 0:
|
||||||
kv_host_size, swa_host_size = _split_hicache_size(
|
kv_host_size, swa_host_size = _split_hicache_size(
|
||||||
server_args.hicache_size, (full_kv_pool, swa_kv_pool)
|
server_args.hicache_size, (full_kv_pool, swa_kv_pool)
|
||||||
)
|
)
|
||||||
|
|
||||||
kv_host_pool = build_kv_host_pool(
|
kv_host_pool = build_kv_host_pool(
|
||||||
kv_pool=full_kv_pool,
|
kv_pool=full_kv_pool,
|
||||||
page_size=params.page_size,
|
page_size=params.page_size,
|
||||||
@@ -229,9 +263,18 @@ def build_hybrid_swa_stack(
|
|||||||
server_args=server_args,
|
server_args=server_args,
|
||||||
use_mla=use_mla,
|
use_mla=use_mla,
|
||||||
host_size=swa_host_size,
|
host_size=swa_host_size,
|
||||||
|
mtp_draft_device_pools=mtp_swa_device_pools,
|
||||||
pool_label="swa",
|
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
|
# 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
|
swa_attn_allocator = params.token_to_kv_pool_allocator.swa_attn_allocator
|
||||||
entries = [
|
entries = [
|
||||||
@@ -248,7 +291,7 @@ def build_hybrid_swa_stack(
|
|||||||
host_pool=swa_host_pool,
|
host_pool=swa_host_pool,
|
||||||
device_pool=swa_kv_pool,
|
device_pool=swa_kv_pool,
|
||||||
layer_mapping=swa_layer_mapping,
|
layer_mapping=swa_layer_mapping,
|
||||||
transfer_layer_num=transfer_layer_num,
|
transfer_layer_num=transfer_layer_num + len(mtp_swa_device_pools),
|
||||||
host_evict_fn=host_swa_evict_fn,
|
host_evict_fn=host_swa_evict_fn,
|
||||||
device_evict_fn=device_swa_evict_fn,
|
device_evict_fn=device_swa_evict_fn,
|
||||||
device_alloc_fn=swa_attn_allocator.alloc,
|
device_alloc_fn=swa_attn_allocator.alloc,
|
||||||
@@ -274,6 +317,8 @@ def build_hybrid_swa_stack(
|
|||||||
transfer_layer_num=transfer_layer_num,
|
transfer_layer_num=transfer_layer_num,
|
||||||
enable_storage_metrics=enable_storage_metrics,
|
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
|
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)}
|
full_layer_mapping = {layer_id: layer_id for layer_id in range(transfer_layer_num)}
|
||||||
|
|
||||||
is_unified_kv = getattr(kvcache, "_unified_kv", False)
|
is_unified_kv = getattr(kvcache, "_unified_kv", False)
|
||||||
|
mtp_swa_device_buffers = []
|
||||||
if is_unified_kv:
|
if is_unified_kv:
|
||||||
# unified_kv keeps the SWA ring inside the unified pool and never offloads it,
|
# unified_kv keeps the SWA ring inside the unified pool and never offloads it,
|
||||||
# so there is no separate SWA host pool to map.
|
# so there is no separate SWA host pool to map.
|
||||||
@@ -346,6 +392,19 @@ def build_deepseek_v4_hicache_stack(
|
|||||||
swa_layer_mapping = {
|
swa_layer_mapping = {
|
||||||
layer_id: layer_id for layer_id in range(transfer_layer_num)
|
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 = {}
|
c4_layer_mapping = {}
|
||||||
c128_layer_mapping = {}
|
c128_layer_mapping = {}
|
||||||
@@ -390,7 +449,10 @@ def build_deepseek_v4_hicache_stack(
|
|||||||
if not is_unified_kv:
|
if not is_unified_kv:
|
||||||
swa_host_pool = DeepSeekV4PagedHostPool(
|
swa_host_pool = DeepSeekV4PagedHostPool(
|
||||||
pool_name=str(PoolName.SWA),
|
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,
|
item_bytes=kvcache.swa_kv_pool.bytes_per_page_padded,
|
||||||
num_host_pages=swa_num_host_pages,
|
num_host_pages=swa_num_host_pages,
|
||||||
slot_page_size=kvcache.swa_page_size,
|
slot_page_size=kvcache.swa_page_size,
|
||||||
@@ -404,7 +466,7 @@ def build_deepseek_v4_hicache_stack(
|
|||||||
host_pool=swa_host_pool,
|
host_pool=swa_host_pool,
|
||||||
device_pool=kvcache.swa_kv_pool,
|
device_pool=kvcache.swa_kv_pool,
|
||||||
layer_mapping=swa_layer_mapping,
|
layer_mapping=swa_layer_mapping,
|
||||||
transfer_layer_num=transfer_layer_num,
|
transfer_layer_num=transfer_layer_num + len(mtp_swa_device_buffers),
|
||||||
host_evict_fn=host_swa_evict_fn,
|
host_evict_fn=host_swa_evict_fn,
|
||||||
device_evict_fn=device_swa_evict_fn,
|
device_evict_fn=device_swa_evict_fn,
|
||||||
device_alloc_fn=swa_attn_allocator.alloc,
|
device_alloc_fn=swa_attn_allocator.alloc,
|
||||||
@@ -542,6 +604,8 @@ def build_deepseek_v4_hicache_stack(
|
|||||||
transfer_layer_num=transfer_layer_num,
|
transfer_layer_num=transfer_layer_num,
|
||||||
enable_storage_metrics=enable_storage_metrics,
|
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
|
return host_pool_group, cache_controller
|
||||||
|
|
||||||
|
|
||||||
@@ -565,6 +629,9 @@ def build_hybrid_mamba_stack(
|
|||||||
) -> tuple[HostPoolGroup, HybridCacheController]:
|
) -> tuple[HostPoolGroup, HybridCacheController]:
|
||||||
transfer_layer_num = len(full_layer_mapping | mamba_layer_mapping)
|
transfer_layer_num = len(full_layer_mapping | mamba_layer_mapping)
|
||||||
mamba_allocator = params.req_to_token_pool.mamba_allocator
|
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
|
kv_host_size, mamba_host_size = None, 0
|
||||||
if server_args.hicache_size > 0:
|
if server_args.hicache_size > 0:
|
||||||
kv_host_size, mamba_host_size = _split_hicache_size(
|
kv_host_size, mamba_host_size = _split_hicache_size(
|
||||||
@@ -576,7 +643,15 @@ def build_hybrid_mamba_stack(
|
|||||||
server_args=server_args,
|
server_args=server_args,
|
||||||
use_mla=use_mla,
|
use_mla=use_mla,
|
||||||
host_size=kv_host_size,
|
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_host_pool = MambaPoolHost(
|
||||||
mamba_pool,
|
mamba_pool,
|
||||||
server_args.hicache_ratio,
|
server_args.hicache_ratio,
|
||||||
@@ -590,7 +665,7 @@ def build_hybrid_mamba_stack(
|
|||||||
host_pool=kv_host_pool,
|
host_pool=kv_host_pool,
|
||||||
device_pool=kv_pool,
|
device_pool=kv_pool,
|
||||||
layer_mapping=full_layer_mapping,
|
layer_mapping=full_layer_mapping,
|
||||||
transfer_layer_num=transfer_layer_num,
|
transfer_layer_num=transfer_layer_num + len(mtp_draft_device_pools),
|
||||||
is_anchor=True,
|
is_anchor=True,
|
||||||
),
|
),
|
||||||
build_pool_entry(
|
build_pool_entry(
|
||||||
@@ -624,6 +699,8 @@ def build_hybrid_mamba_stack(
|
|||||||
transfer_layer_num=transfer_layer_num,
|
transfer_layer_num=transfer_layer_num,
|
||||||
enable_storage_metrics=enable_storage_metrics,
|
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
|
return host_pool_group, cache_controller
|
||||||
|
|
||||||
|
|
||||||
@@ -758,21 +835,33 @@ def build_anchor_sidecar_stack(
|
|||||||
enable_storage_metrics: bool = False,
|
enable_storage_metrics: bool = False,
|
||||||
) -> tuple[HostPoolGroup, HybridCacheController]:
|
) -> tuple[HostPoolGroup, HybridCacheController]:
|
||||||
transfer_layer_num = len(full_layer_mapping)
|
transfer_layer_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_host_pool = build_kv_host_pool(
|
||||||
kv_pool=kv_pool,
|
kv_pool=kv_pool,
|
||||||
page_size=params.page_size,
|
page_size=params.page_size,
|
||||||
server_args=server_args,
|
server_args=server_args,
|
||||||
use_mla=use_mla,
|
use_mla=use_mla,
|
||||||
override_kv_cache_dim=override_kv_cache_dim,
|
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)
|
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 = [
|
entries = [
|
||||||
build_pool_entry(
|
build_pool_entry(
|
||||||
name=PoolName.KV,
|
name=PoolName.KV,
|
||||||
host_pool=kv_host_pool,
|
host_pool=kv_host_pool,
|
||||||
device_pool=kv_pool,
|
device_pool=kv_pool,
|
||||||
layer_mapping=full_layer_mapping,
|
layer_mapping=full_layer_mapping,
|
||||||
transfer_layer_num=transfer_layer_num,
|
transfer_layer_num=transfer_layer_num + len(mtp_draft_device_pools),
|
||||||
is_anchor=True,
|
is_anchor=True,
|
||||||
),
|
),
|
||||||
build_pool_entry(
|
build_pool_entry(
|
||||||
@@ -780,7 +869,7 @@ def build_anchor_sidecar_stack(
|
|||||||
host_pool=sidecar_host_pool,
|
host_pool=sidecar_host_pool,
|
||||||
device_pool=kv_pool,
|
device_pool=kv_pool,
|
||||||
layer_mapping=full_layer_mapping,
|
layer_mapping=full_layer_mapping,
|
||||||
transfer_layer_num=transfer_layer_num,
|
transfer_layer_num=transfer_layer_num + len(mtp_draft_device_pools),
|
||||||
),
|
),
|
||||||
]
|
]
|
||||||
host_pool_group = HostPoolGroup(entries)
|
host_pool_group = HostPoolGroup(entries)
|
||||||
@@ -802,9 +891,184 @@ def build_anchor_sidecar_stack(
|
|||||||
transfer_layer_num=transfer_layer_num,
|
transfer_layer_num=transfer_layer_num,
|
||||||
enable_storage_metrics=enable_storage_metrics,
|
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
|
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]] = {
|
_COMPONENT_HOST_ATTR: dict[ComponentType, tuple[str, str]] = {
|
||||||
ComponentType.FULL: ("full_kv_pool_host", "_full_kv_pool_host"),
|
ComponentType.FULL: ("full_kv_pool_host", "_full_kv_pool_host"),
|
||||||
ComponentType.SWA: ("swa_kv_pool_host", "_swa_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 import GroupCoordinator
|
||||||
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
|
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
|
||||||
from sglang.srt.managers.tp_worker import BaseTpWorker
|
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.server_args import ServerArgs
|
||||||
|
from sglang.srt.speculative.base_spec_worker import HiCacheDraftPlan
|
||||||
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
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(
|
def maybe_register_hicache_draft(
|
||||||
*,
|
*,
|
||||||
tree_cache: BasePrefixCache,
|
tree_cache,
|
||||||
draft_worker: BaseTpWorker,
|
draft_plan: HiCacheDraftPlan,
|
||||||
spec_algorithm: SpeculativeAlgorithm,
|
|
||||||
server_args: ServerArgs,
|
server_args: ServerArgs,
|
||||||
enable_hierarchical_cache: bool,
|
|
||||||
page_size: int,
|
page_size: int,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Register draft KV pool with HiCacheController for piggyback L2/L3 ops."""
|
from sglang.srt.speculative.base_spec_worker import HiCacheDraftMode
|
||||||
if not enable_hierarchical_cache:
|
|
||||||
|
if draft_plan.mode != HiCacheDraftMode.SIDECAR:
|
||||||
return
|
return
|
||||||
|
|
||||||
draft_kv_pool = get_draft_kv_pool(
|
from sglang.srt.mem_cache.unified_radix_cache import UnifiedRadixCache
|
||||||
draft_worker=draft_worker,
|
|
||||||
spec_algorithm=spec_algorithm,
|
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,
|
server_args=server_args,
|
||||||
)
|
)
|
||||||
if draft_kv_pool is None:
|
tree_cache.register_hicache_draft_pools(specs, entries)
|
||||||
return
|
|
||||||
|
|
||||||
|
|
||||||
|
def _register_legacy_hicache_draft(
|
||||||
|
*,
|
||||||
|
tree_cache,
|
||||||
|
draft_pool,
|
||||||
|
server_args: ServerArgs,
|
||||||
|
page_size: int,
|
||||||
|
) -> None:
|
||||||
from sglang.srt.mem_cache.memory_pool import (
|
from sglang.srt.mem_cache.memory_pool import (
|
||||||
HybridLinearKVPool,
|
|
||||||
MHATokenToKVPool,
|
MHATokenToKVPool,
|
||||||
MLATokenToKVPool,
|
MLATokenToKVPool,
|
||||||
)
|
)
|
||||||
from sglang.srt.mem_cache.pool_host.mha import get_mha_host_pool_cls
|
from sglang.srt.mem_cache.pool_host.mha import get_mha_host_pool_cls
|
||||||
from sglang.srt.mem_cache.pool_host.mla import MLATokenToKVPoolHost
|
from sglang.srt.mem_cache.pool_host.mla import MLATokenToKVPoolHost
|
||||||
|
|
||||||
pool = draft_kv_pool
|
pool = draft_pool
|
||||||
if isinstance(pool, HybridLinearKVPool):
|
if pool.layer_num == 0:
|
||||||
pool = pool.full_kv_pool
|
return
|
||||||
|
|
||||||
# Create host pool for draft with the same slot count as the target host pool,
|
# 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.
|
# so that host indices stay 1-to-1 between target and draft KV caches.
|
||||||
primary = tree_cache.cache_controller.mem_pool_host
|
primary_host_pool = tree_cache.cache_controller.mem_pool_host
|
||||||
kw = dict(
|
host_pool_kwargs = dict(
|
||||||
host_to_device_ratio=primary.size / pool.size,
|
host_to_device_ratio=primary_host_pool.size / pool.size,
|
||||||
host_size=0,
|
host_size=0,
|
||||||
page_size=page_size,
|
page_size=page_size,
|
||||||
layout=server_args.hicache_mem_layout,
|
layout=server_args.hicache_mem_layout,
|
||||||
@@ -115,12 +116,13 @@ def maybe_register_hicache_draft(
|
|||||||
pool_label="draft",
|
pool_label="draft",
|
||||||
)
|
)
|
||||||
if isinstance(pool, MHATokenToKVPool):
|
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):
|
elif isinstance(pool, MLATokenToKVPool):
|
||||||
draft_host_pool = MLATokenToKVPoolHost(pool, **kw)
|
draft_host_pool = MLATokenToKVPoolHost(pool, **host_pool_kwargs)
|
||||||
else:
|
else:
|
||||||
logger.warning(
|
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__,
|
type(pool).__name__,
|
||||||
)
|
)
|
||||||
return
|
return
|
||||||
@@ -144,6 +146,7 @@ def build_kv_cache(
|
|||||||
tp_group: GroupCoordinator,
|
tp_group: GroupCoordinator,
|
||||||
pp_group: GroupCoordinator,
|
pp_group: GroupCoordinator,
|
||||||
enable_hierarchical_cache: bool,
|
enable_hierarchical_cache: bool,
|
||||||
|
hicache_draft_plan: Optional[HiCacheDraftPlan] = None,
|
||||||
) -> KVCacheBuildResult:
|
) -> KVCacheBuildResult:
|
||||||
sliding_window_size: Optional[int] = None
|
sliding_window_size: Optional[int] = None
|
||||||
full_tokens_per_layer: 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()
|
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 (
|
disable_radix_cache = server_args.disable_radix_cache or (
|
||||||
model_config.is_multimodal and uses_transformers_backend
|
model_config.is_multimodal and uses_transformers_backend
|
||||||
@@ -234,6 +238,7 @@ def build_kv_cache(
|
|||||||
pp_size=ps.pp_size,
|
pp_size=ps.pp_size,
|
||||||
chunked_prefill_size=effective_chunked_prefill_size,
|
chunked_prefill_size=effective_chunked_prefill_size,
|
||||||
sliding_window_size=sliding_window_size,
|
sliding_window_size=sliding_window_size,
|
||||||
|
mtp_draft_device_pools=mtp_draft_device_pools,
|
||||||
)
|
)
|
||||||
|
|
||||||
tree_cache = create_tree_cache(
|
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()
|
embedding_cache_size = envs.SGLANG_VLM_CACHE_SIZE_MB.get()
|
||||||
init_mm_embedding_cache(embedding_cache_size * 1024 * 1024)
|
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):
|
class MambaPoolHost(HostKVCache):
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
device_pool: MambaPool,
|
device_pool: MambaPool,
|
||||||
@@ -432,6 +431,8 @@ class MambaPoolHost(HostKVCache):
|
|||||||
device_indices,
|
device_indices,
|
||||||
layer_id,
|
layer_id,
|
||||||
io_backend="kernel",
|
io_backend="kernel",
|
||||||
|
*,
|
||||||
|
is_draft: bool = False,
|
||||||
):
|
):
|
||||||
if self.layout in ["page_first", "page_first_direct"]:
|
if self.layout in ["page_first", "page_first_direct"]:
|
||||||
# no ssm state on conv-only models: nothing to transfer
|
# no ssm state on conv-only models: nothing to transfer
|
||||||
@@ -704,7 +705,14 @@ class LogicalHostPool:
|
|||||||
pass
|
pass
|
||||||
|
|
||||||
def load_to_device_per_layer(
|
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
|
pass
|
||||||
|
|
||||||
@@ -988,7 +996,14 @@ class DeepSeekV4PagedHostPool(HiSparseHostPoolMixin, HostKVCache):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def load_to_device_per_layer(
|
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):
|
if not self._has_transfer_indices(host_indices, device_indices):
|
||||||
return
|
return
|
||||||
@@ -1374,7 +1389,14 @@ class DeepSeekV4StateHostPool(HostKVCache):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def load_to_device_per_layer(
|
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:
|
if host_indices is None or device_indices is None:
|
||||||
return
|
return
|
||||||
@@ -1538,6 +1560,19 @@ class HostPoolGroup:
|
|||||||
self.can_use_write_back_jit = all(child_write_back_jit)
|
self.can_use_write_back_jit = all(child_write_back_jit)
|
||||||
self.supports_per_pool_backup_indices = any(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
|
@property
|
||||||
def kv_buffer(self):
|
def kv_buffer(self):
|
||||||
return self.anchor_entry.host_pool.kv_buffer
|
return self.anchor_entry.host_pool.kv_buffer
|
||||||
@@ -1608,17 +1643,20 @@ class HostPoolGroup:
|
|||||||
layer_id,
|
layer_id,
|
||||||
io_backend,
|
io_backend,
|
||||||
pool_transfers: Optional[list] = None,
|
pool_transfers: Optional[list] = None,
|
||||||
|
*,
|
||||||
|
is_draft: bool = False,
|
||||||
) -> None:
|
) -> None:
|
||||||
# 1. Anchor (KV) transfer
|
# 1. Anchor (KV) transfer
|
||||||
anchor = self.anchor_entry
|
anchor = self.anchor_entry
|
||||||
local_layer_id = anchor.layer_mapper(layer_id)
|
local_layer_id = anchor.layer_mapper(layer_id)
|
||||||
if local_layer_id is not None and host_indices.numel() > 0:
|
if local_layer_id is not None and host_indices.numel() > 0:
|
||||||
anchor.host_pool.load_to_device_per_layer(
|
anchor.host_pool.load_to_device_per_layer(
|
||||||
anchor.device_pool,
|
device_pool if is_draft else anchor.device_pool,
|
||||||
host_indices,
|
host_indices,
|
||||||
device_indices,
|
device_indices,
|
||||||
local_layer_id,
|
local_layer_id,
|
||||||
io_backend,
|
io_backend,
|
||||||
|
is_draft=is_draft,
|
||||||
)
|
)
|
||||||
|
|
||||||
# 2. Extra pool transfers
|
# 2. Extra pool transfers
|
||||||
@@ -1630,11 +1668,12 @@ class HostPoolGroup:
|
|||||||
if local_layer_id is None:
|
if local_layer_id is None:
|
||||||
continue
|
continue
|
||||||
entry.host_pool.load_to_device_per_layer(
|
entry.host_pool.load_to_device_per_layer(
|
||||||
entry.device_pool,
|
device_pool if is_draft else entry.device_pool,
|
||||||
transfer.host_indices,
|
transfer.host_indices,
|
||||||
transfer.device_indices,
|
transfer.device_indices,
|
||||||
local_layer_id,
|
local_layer_id,
|
||||||
io_backend,
|
io_backend,
|
||||||
|
is_draft=is_draft,
|
||||||
)
|
)
|
||||||
|
|
||||||
def _backup_uses_cpu_host_indices(self, host_pool, io_backend) -> bool:
|
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.dtype = device_pool.store_dtype
|
||||||
self.start_layer = device_pool.start_layer
|
self.start_layer = device_pool.start_layer
|
||||||
self.end_layer = device_pool.end_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.index_head_dim = device_pool.index_head_dim
|
||||||
self.indexer_quant_block_size = device_pool.quant_block_size
|
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"Requesting {requested_bytes / 1e9:.2f} GB but only have "
|
||||||
f"{available_bytes / 1e9:.2f} GB free."
|
f"{available_bytes / 1e9:.2f} GB free."
|
||||||
)
|
)
|
||||||
logger.info(
|
draft_layer_num = self.layer_num - self.target_layer_num
|
||||||
"Allocating %.2f GB host memory for DSA indexer (layout=%s).",
|
if draft_layer_num > 0:
|
||||||
requested_bytes / 1e9,
|
logger.info(
|
||||||
layout,
|
"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.init_kv_buffer()
|
||||||
self.can_use_jit = False
|
self.can_use_jit = False
|
||||||
self.can_use_write_back_jit = False
|
self.can_use_write_back_jit = False
|
||||||
@@ -1785,8 +1839,12 @@ class DSAIndexerPoolHost(HostKVCache):
|
|||||||
|
|
||||||
def init_kv_buffer(self):
|
def init_kv_buffer(self):
|
||||||
alloc_func = ALLOC_MEMORY_FUNCS[self.device_pool.device]
|
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(
|
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,
|
dtype=torch.uint64,
|
||||||
device=self.device_pool.device,
|
device=self.device_pool.device,
|
||||||
)
|
)
|
||||||
@@ -1863,11 +1921,20 @@ class DSAIndexerPoolHost(HostKVCache):
|
|||||||
return host_page_indices, device_page_indices
|
return host_page_indices, device_page_indices
|
||||||
|
|
||||||
def load_to_device_per_layer(
|
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
|
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_page_indices, device_page_indices = self._get_indexer_page_indices(
|
||||||
host_indices, device_indices
|
host_indices, device_indices
|
||||||
@@ -1876,8 +1943,8 @@ class DSAIndexerPoolHost(HostKVCache):
|
|||||||
if use_kernel:
|
if use_kernel:
|
||||||
if self.layout == "layer_first":
|
if self.layout == "layer_first":
|
||||||
transfer_kv_per_layer_mla(
|
transfer_kv_per_layer_mla(
|
||||||
src=self.index_k_with_scale_buffer[host_layer],
|
src=self.index_k_with_scale_buffer[host_layer_id],
|
||||||
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,
|
src_indices=host_page_indices,
|
||||||
dst_indices=device_page_indices,
|
dst_indices=device_page_indices,
|
||||||
item_size=self.indexer_page_stride_size,
|
item_size=self.indexer_page_stride_size,
|
||||||
@@ -1885,10 +1952,10 @@ class DSAIndexerPoolHost(HostKVCache):
|
|||||||
elif self.layout == "page_first":
|
elif self.layout == "page_first":
|
||||||
transfer_kv_per_layer_mla_pf_lf(
|
transfer_kv_per_layer_mla_pf_lf(
|
||||||
src=self.index_k_with_scale_buffer,
|
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,
|
src_indices=host_page_indices,
|
||||||
dst_indices=device_page_indices,
|
dst_indices=device_page_indices,
|
||||||
layer_id=host_layer,
|
layer_id=host_layer_id,
|
||||||
item_size=self.indexer_page_stride_size,
|
item_size=self.indexer_page_stride_size,
|
||||||
src_layout_dim=self.indexer_layout_dim,
|
src_layout_dim=self.indexer_layout_dim,
|
||||||
)
|
)
|
||||||
@@ -1897,8 +1964,8 @@ class DSAIndexerPoolHost(HostKVCache):
|
|||||||
elif io_backend == "direct":
|
elif io_backend == "direct":
|
||||||
if self.layout == "layer_first":
|
if self.layout == "layer_first":
|
||||||
transfer_kv_direct(
|
transfer_kv_direct(
|
||||||
src_layers=[self.index_k_with_scale_buffer[host_layer]],
|
src_layers=[self.index_k_with_scale_buffer[host_layer_id]],
|
||||||
dst_layers=[device_pool.index_k_with_scale_buffer[layer_id]],
|
dst_layers=[device_pool.index_k_with_scale_buffer[device_layer_id]],
|
||||||
src_indices=host_page_indices,
|
src_indices=host_page_indices,
|
||||||
dst_indices=device_page_indices,
|
dst_indices=device_page_indices,
|
||||||
page_size=1,
|
page_size=1,
|
||||||
@@ -1906,10 +1973,10 @@ class DSAIndexerPoolHost(HostKVCache):
|
|||||||
elif self.layout == "page_first_direct":
|
elif self.layout == "page_first_direct":
|
||||||
transfer_kv_per_layer_direct_pf_lf(
|
transfer_kv_per_layer_direct_pf_lf(
|
||||||
src_ptrs=[self.index_k_with_scale_buffer],
|
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,
|
src_indices=host_page_indices,
|
||||||
dst_indices=device_page_indices,
|
dst_indices=device_page_indices,
|
||||||
layer_id=host_layer,
|
layer_id=host_layer_id,
|
||||||
page_size=1,
|
page_size=1,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
@@ -1918,9 +1985,19 @@ class DSAIndexerPoolHost(HostKVCache):
|
|||||||
raise ValueError(f"Unsupported IO backend: {io_backend}")
|
raise ValueError(f"Unsupported IO backend: {io_backend}")
|
||||||
|
|
||||||
def _backup_from_device_per_layer(
|
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_page_indices, device_page_indices = self._get_indexer_page_indices(
|
||||||
host_indices, device_indices
|
host_indices, device_indices
|
||||||
)
|
)
|
||||||
@@ -1928,8 +2005,8 @@ class DSAIndexerPoolHost(HostKVCache):
|
|||||||
if use_kernel:
|
if use_kernel:
|
||||||
if self.layout == "layer_first":
|
if self.layout == "layer_first":
|
||||||
transfer_kv_per_layer_mla(
|
transfer_kv_per_layer_mla(
|
||||||
src=device_pool.index_k_with_scale_buffer[layer_id],
|
src=device_pool.index_k_with_scale_buffer[device_layer_id],
|
||||||
dst=self.index_k_with_scale_buffer[host_layer],
|
dst=self.index_k_with_scale_buffer[host_layer_id],
|
||||||
src_indices=device_page_indices,
|
src_indices=device_page_indices,
|
||||||
dst_indices=host_page_indices,
|
dst_indices=host_page_indices,
|
||||||
item_size=self.indexer_page_stride_size,
|
item_size=self.indexer_page_stride_size,
|
||||||
@@ -1944,8 +2021,8 @@ class DSAIndexerPoolHost(HostKVCache):
|
|||||||
elif io_backend == "direct":
|
elif io_backend == "direct":
|
||||||
if self.layout == "layer_first":
|
if self.layout == "layer_first":
|
||||||
transfer_kv_direct(
|
transfer_kv_direct(
|
||||||
src_layers=[device_pool.index_k_with_scale_buffer[layer_id]],
|
src_layers=[device_pool.index_k_with_scale_buffer[device_layer_id]],
|
||||||
dst_layers=[self.index_k_with_scale_buffer[host_layer]],
|
dst_layers=[self.index_k_with_scale_buffer[host_layer_id]],
|
||||||
src_indices=device_page_indices,
|
src_indices=device_page_indices,
|
||||||
dst_indices=host_page_indices,
|
dst_indices=host_page_indices,
|
||||||
page_size=1,
|
page_size=1,
|
||||||
@@ -1966,6 +2043,17 @@ class DSAIndexerPoolHost(HostKVCache):
|
|||||||
self._backup_from_device_per_layer(
|
self._backup_from_device_per_layer(
|
||||||
device_pool, host_indices, device_indices, layer_id, io_backend
|
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
|
return
|
||||||
|
|
||||||
host_page_indices, device_page_indices = self._get_indexer_page_indices(
|
host_page_indices, device_page_indices = self._get_indexer_page_indices(
|
||||||
@@ -2008,7 +2096,7 @@ class DSAIndexerPoolHost(HostKVCache):
|
|||||||
elif io_backend == "direct":
|
elif io_backend == "direct":
|
||||||
if self.layout == "layer_first":
|
if self.layout == "layer_first":
|
||||||
transfer_kv_direct(
|
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,
|
dst_layers=self.index_k_data_refs,
|
||||||
src_indices=device_page_indices,
|
src_indices=device_page_indices,
|
||||||
dst_indices=host_page_indices,
|
dst_indices=host_page_indices,
|
||||||
@@ -2016,7 +2104,7 @@ class DSAIndexerPoolHost(HostKVCache):
|
|||||||
)
|
)
|
||||||
elif self.layout == "page_first_direct":
|
elif self.layout == "page_first_direct":
|
||||||
transfer_kv_all_layer_direct_lf_pf(
|
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],
|
dst_ptrs=[self.index_k_with_scale_buffer],
|
||||||
src_indices=device_page_indices,
|
src_indices=device_page_indices,
|
||||||
dst_indices=host_page_indices,
|
dst_indices=host_page_indices,
|
||||||
|
|||||||
@@ -150,12 +150,26 @@ class HostKVCache(abc.ABC):
|
|||||||
f"size of the hierarchical cache."
|
f"size of the hierarchical cache."
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
logger.info(
|
draft_layer_num = self.layer_num - self.target_layer_num
|
||||||
"Allocating %s hierarchical KV host pool: %d tokens, %.2f GB host memory.",
|
if draft_layer_num > 0:
|
||||||
pool_label,
|
logger.info(
|
||||||
self.size,
|
"Allocating %s hierarchical KV host pool: %d tokens, "
|
||||||
requested_bytes / 1e9,
|
"%.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.kv_buffer = self.init_kv_buffer()
|
||||||
self.fd = getattr(self.allocator, "fd", None)
|
self.fd = getattr(self.allocator, "fd", None)
|
||||||
@@ -215,7 +229,6 @@ class HostKVCache(abc.ABC):
|
|||||||
return start <= layer_id < end
|
return start <= layer_id < end
|
||||||
|
|
||||||
def _host_layer_index(self, layer_id: int, device_pool=None) -> int:
|
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)
|
start, _ = self._device_owned_layer_range(device_pool)
|
||||||
return layer_id - start
|
return layer_id - start
|
||||||
|
|
||||||
@@ -229,7 +242,14 @@ class HostKVCache(abc.ABC):
|
|||||||
|
|
||||||
@abc.abstractmethod
|
@abc.abstractmethod
|
||||||
def load_to_device_per_layer(
|
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:
|
) -> None:
|
||||||
"""
|
"""
|
||||||
Load KV data from the host memory pool to the device memory pool for a specific layer.
|
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 logging
|
||||||
import threading
|
import threading
|
||||||
|
from typing import Sequence
|
||||||
|
|
||||||
import psutil
|
import psutil
|
||||||
import torch
|
import torch
|
||||||
@@ -67,7 +68,8 @@ logger = logging.getLogger(__name__)
|
|||||||
|
|
||||||
|
|
||||||
class MHATokenToKVPoolHost(HostKVCache):
|
class MHATokenToKVPoolHost(HostKVCache):
|
||||||
device_pool: MHATokenToKVPool
|
device_pool: MHATokenToKVPool | None = None
|
||||||
|
mtp_draft_device_pools: tuple[MHATokenToKVPool, ...] = ()
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
@@ -80,8 +82,11 @@ class MHATokenToKVPoolHost(HostKVCache):
|
|||||||
device: str = "cpu",
|
device: str = "cpu",
|
||||||
allocator_type: str = "default",
|
allocator_type: str = "default",
|
||||||
*,
|
*,
|
||||||
|
mtp_draft_device_pools: Sequence[MHATokenToKVPool] = (),
|
||||||
pool_label: str = "kv",
|
pool_label: str = "kv",
|
||||||
):
|
):
|
||||||
|
self.mtp_draft_device_pools = tuple(mtp_draft_device_pools)
|
||||||
|
self.target_layer_num = device_pool.layer_num
|
||||||
super().__init__(
|
super().__init__(
|
||||||
device_pool,
|
device_pool,
|
||||||
host_to_device_ratio,
|
host_to_device_ratio,
|
||||||
@@ -122,12 +127,30 @@ class MHATokenToKVPoolHost(HostKVCache):
|
|||||||
dtype=torch.uint64,
|
dtype=torch.uint64,
|
||||||
device=self.device_pool.device,
|
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()
|
self._init_write_back_staging_buffers()
|
||||||
|
|
||||||
def get_size_per_token(self):
|
def get_size_per_token(self):
|
||||||
self.head_num = self.device_pool.head_num
|
self.head_num = self.device_pool.head_num
|
||||||
self.head_dim = self.device_pool.head_dim
|
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
|
return self.head_dim * self.head_num * self.layer_num * self.dtype.itemsize * 2
|
||||||
|
|
||||||
def get_ksize_per_token(self):
|
def get_ksize_per_token(self):
|
||||||
@@ -219,25 +242,36 @@ class MHATokenToKVPoolHost(HostKVCache):
|
|||||||
device_indices,
|
device_indices,
|
||||||
layer_id,
|
layer_id,
|
||||||
io_backend,
|
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 io_backend == "kernel":
|
||||||
if self.layout == "layer_first":
|
if self.layout == "layer_first":
|
||||||
if self.can_use_jit:
|
if self.can_use_jit:
|
||||||
jit_transfer_hicache_one_layer(
|
jit_transfer_hicache_one_layer(
|
||||||
k_cache_dst=device_pool.k_buffer[layer_id],
|
k_cache_dst=device_pool.k_buffer[device_layer_id],
|
||||||
v_cache_dst=device_pool.v_buffer[layer_id],
|
v_cache_dst=device_pool.v_buffer[device_layer_id],
|
||||||
k_cache_src=self.k_buffer[layer_id],
|
k_cache_src=self.k_buffer[host_layer_id],
|
||||||
v_cache_src=self.v_buffer[layer_id],
|
v_cache_src=self.v_buffer[host_layer_id],
|
||||||
indices_dst=device_indices,
|
indices_dst=device_indices,
|
||||||
indices_src=host_indices,
|
indices_src=host_indices,
|
||||||
element_dim=self.element_dim,
|
element_dim=self.element_dim,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
transfer_kv_per_layer(
|
transfer_kv_per_layer(
|
||||||
src_k=self.k_buffer[layer_id],
|
src_k=self.k_buffer[host_layer_id],
|
||||||
dst_k=device_pool.k_buffer[layer_id],
|
dst_k=device_pool.k_buffer[device_layer_id],
|
||||||
src_v=self.v_buffer[layer_id],
|
src_v=self.v_buffer[host_layer_id],
|
||||||
dst_v=device_pool.v_buffer[layer_id],
|
dst_v=device_pool.v_buffer[device_layer_id],
|
||||||
src_indices=host_indices,
|
src_indices=host_indices,
|
||||||
dst_indices=device_indices,
|
dst_indices=device_indices,
|
||||||
item_size=self.token_stride_size,
|
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.
|
# index by layer_id to get a per-layer view with strided layout.
|
||||||
# The kernel handles different src/dst strides automatically.
|
# The kernel handles different src/dst strides automatically.
|
||||||
jit_transfer_hicache_one_layer(
|
jit_transfer_hicache_one_layer(
|
||||||
k_cache_dst=device_pool.k_buffer[layer_id],
|
k_cache_dst=device_pool.k_buffer[device_layer_id],
|
||||||
v_cache_dst=device_pool.v_buffer[layer_id],
|
v_cache_dst=device_pool.v_buffer[device_layer_id],
|
||||||
k_cache_src=self.k_data_refs[layer_id],
|
k_cache_src=self.k_data_refs[host_layer_id],
|
||||||
v_cache_src=self.v_data_refs[layer_id],
|
v_cache_src=self.v_data_refs[host_layer_id],
|
||||||
indices_dst=device_indices,
|
indices_dst=device_indices,
|
||||||
indices_src=host_indices,
|
indices_src=host_indices,
|
||||||
element_dim=self.element_dim,
|
element_dim=self.element_dim,
|
||||||
@@ -259,24 +293,24 @@ class MHATokenToKVPoolHost(HostKVCache):
|
|||||||
else:
|
else:
|
||||||
transfer_kv_per_layer_pf_lf(
|
transfer_kv_per_layer_pf_lf(
|
||||||
src_k=self.k_buffer,
|
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,
|
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,
|
src_indices=host_indices,
|
||||||
dst_indices=device_indices,
|
dst_indices=device_indices,
|
||||||
layer_id=layer_id,
|
layer_id=host_layer_id,
|
||||||
item_size=self.token_stride_size,
|
item_size=self.token_stride_size,
|
||||||
src_layout_dim=self.layout_dim,
|
src_layout_dim=self.layout_dim,
|
||||||
)
|
)
|
||||||
elif self.layout == "page_head":
|
elif self.layout == "page_head":
|
||||||
transfer_kv_per_layer_ph_lf(
|
transfer_kv_per_layer_ph_lf(
|
||||||
src_k=self.k_buffer,
|
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,
|
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,
|
src_indices=host_indices,
|
||||||
dst_indices=device_indices,
|
dst_indices=device_indices,
|
||||||
layer_id=layer_id,
|
layer_id=host_layer_id,
|
||||||
item_size=self.token_stride_size,
|
item_size=self.token_stride_size,
|
||||||
src_layout_dim=self.layout_dim,
|
src_layout_dim=self.layout_dim,
|
||||||
page_size=self.page_size,
|
page_size=self.page_size,
|
||||||
@@ -287,10 +321,13 @@ class MHATokenToKVPoolHost(HostKVCache):
|
|||||||
elif io_backend == "direct":
|
elif io_backend == "direct":
|
||||||
if self.layout == "layer_first":
|
if self.layout == "layer_first":
|
||||||
transfer_kv_direct(
|
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=[
|
dst_layers=[
|
||||||
device_pool.k_buffer[layer_id],
|
device_pool.k_buffer[device_layer_id],
|
||||||
device_pool.v_buffer[layer_id],
|
device_pool.v_buffer[device_layer_id],
|
||||||
],
|
],
|
||||||
src_indices=host_indices,
|
src_indices=host_indices,
|
||||||
dst_indices=device_indices,
|
dst_indices=device_indices,
|
||||||
@@ -300,12 +337,12 @@ class MHATokenToKVPoolHost(HostKVCache):
|
|||||||
transfer_kv_per_layer_direct_pf_lf(
|
transfer_kv_per_layer_direct_pf_lf(
|
||||||
src_ptrs=[self.k_buffer, self.v_buffer],
|
src_ptrs=[self.k_buffer, self.v_buffer],
|
||||||
dst_ptrs=[
|
dst_ptrs=[
|
||||||
device_pool.k_buffer[layer_id],
|
device_pool.k_buffer[device_layer_id],
|
||||||
device_pool.v_buffer[layer_id],
|
device_pool.v_buffer[device_layer_id],
|
||||||
],
|
],
|
||||||
src_indices=host_indices,
|
src_indices=host_indices,
|
||||||
dst_indices=device_indices,
|
dst_indices=device_indices,
|
||||||
layer_id=layer_id,
|
layer_id=host_layer_id,
|
||||||
page_size=self.page_size,
|
page_size=self.page_size,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
@@ -313,7 +350,7 @@ class MHATokenToKVPoolHost(HostKVCache):
|
|||||||
elif io_backend == "kernel_ascend":
|
elif io_backend == "kernel_ascend":
|
||||||
if self.layout == "page_first_direct":
|
if self.layout == "page_first_direct":
|
||||||
# Ascend-specific: transfer KV data for all layers when layer_id == 0
|
# 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(
|
transfer_kv_dim_exchange(
|
||||||
device_indices=device_indices,
|
device_indices=device_indices,
|
||||||
host_indices=host_indices,
|
host_indices=host_indices,
|
||||||
@@ -329,9 +366,31 @@ class MHATokenToKVPoolHost(HostKVCache):
|
|||||||
else:
|
else:
|
||||||
raise ValueError(f"Unsupported IO backend: {io_backend}")
|
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(
|
def backup_from_device_all_layer(
|
||||||
self, device_pool, host_indices, device_indices, io_backend
|
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 io_backend == "kernel":
|
||||||
if self.layout == "layer_first":
|
if self.layout == "layer_first":
|
||||||
if self.can_use_jit:
|
if self.can_use_jit:
|
||||||
@@ -339,8 +398,8 @@ class MHATokenToKVPoolHost(HostKVCache):
|
|||||||
k_ptr_dst=self.k_data_ptrs,
|
k_ptr_dst=self.k_data_ptrs,
|
||||||
v_ptr_dst=self.v_data_ptrs,
|
v_ptr_dst=self.v_data_ptrs,
|
||||||
indices_dst=host_indices,
|
indices_dst=host_indices,
|
||||||
k_ptr_src=device_pool.k_data_ptrs,
|
k_ptr_src=device_k_data_ptrs,
|
||||||
v_ptr_src=device_pool.v_data_ptrs,
|
v_ptr_src=device_v_data_ptrs,
|
||||||
indices_src=device_indices,
|
indices_src=device_indices,
|
||||||
kv_cache_dst_stride_bytes=self.token_stride_size,
|
kv_cache_dst_stride_bytes=self.token_stride_size,
|
||||||
kv_cache_src_stride_bytes=self.token_stride_size,
|
kv_cache_src_stride_bytes=self.token_stride_size,
|
||||||
@@ -348,9 +407,9 @@ class MHATokenToKVPoolHost(HostKVCache):
|
|||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
transfer_kv_all_layer(
|
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,
|
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,
|
dst_v_layers=self.v_data_ptrs,
|
||||||
src_indices=device_indices,
|
src_indices=device_indices,
|
||||||
dst_indices=host_indices,
|
dst_indices=host_indices,
|
||||||
@@ -360,8 +419,8 @@ class MHATokenToKVPoolHost(HostKVCache):
|
|||||||
elif self.layout == "page_first":
|
elif self.layout == "page_first":
|
||||||
if self.can_use_write_back_jit:
|
if self.can_use_write_back_jit:
|
||||||
jit_transfer_hicache_all_layer_staged_lf_pf(
|
jit_transfer_hicache_all_layer_staged_lf_pf(
|
||||||
k_ptr_src=device_pool.k_data_ptrs,
|
k_ptr_src=device_k_data_ptrs,
|
||||||
v_ptr_src=device_pool.v_data_ptrs,
|
v_ptr_src=device_v_data_ptrs,
|
||||||
src_indices=device_indices,
|
src_indices=device_indices,
|
||||||
dst_indices=host_indices,
|
dst_indices=host_indices,
|
||||||
staging_k=self.staging_k_buffer,
|
staging_k=self.staging_k_buffer,
|
||||||
@@ -372,9 +431,9 @@ class MHATokenToKVPoolHost(HostKVCache):
|
|||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
transfer_kv_all_layer_lf_pf(
|
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,
|
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,
|
dst_v=self.v_buffer,
|
||||||
src_indices=device_indices,
|
src_indices=device_indices,
|
||||||
dst_indices=host_indices,
|
dst_indices=host_indices,
|
||||||
@@ -384,9 +443,9 @@ class MHATokenToKVPoolHost(HostKVCache):
|
|||||||
)
|
)
|
||||||
elif self.layout == "page_head":
|
elif self.layout == "page_head":
|
||||||
transfer_kv_all_layer_lf_ph(
|
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,
|
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,
|
dst_v=self.v_buffer,
|
||||||
src_indices=device_indices,
|
src_indices=device_indices,
|
||||||
dst_indices=host_indices,
|
dst_indices=host_indices,
|
||||||
@@ -401,15 +460,15 @@ class MHATokenToKVPoolHost(HostKVCache):
|
|||||||
elif io_backend == "direct":
|
elif io_backend == "direct":
|
||||||
if self.layout == "layer_first":
|
if self.layout == "layer_first":
|
||||||
transfer_kv_direct(
|
transfer_kv_direct(
|
||||||
src_layers=device_pool.k_buffer + device_pool.v_buffer,
|
src_layers=device_kv_buffers,
|
||||||
dst_layers=self.k_data_refs + self.v_data_refs,
|
dst_layers=self.host_kv_data_refs,
|
||||||
src_indices=device_indices,
|
src_indices=device_indices,
|
||||||
dst_indices=host_indices,
|
dst_indices=host_indices,
|
||||||
page_size=self.page_size,
|
page_size=self.page_size,
|
||||||
)
|
)
|
||||||
elif self.layout == "page_first_direct":
|
elif self.layout == "page_first_direct":
|
||||||
transfer_kv_all_layer_direct_lf_pf(
|
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],
|
dst_ptrs=[self.k_buffer, self.v_buffer],
|
||||||
src_indices=device_indices,
|
src_indices=device_indices,
|
||||||
dst_indices=host_indices,
|
dst_indices=host_indices,
|
||||||
@@ -734,7 +793,14 @@ class MHATokenToKOnlyPoolHost(HostKVCache):
|
|||||||
return [self.k_buffer]
|
return [self.k_buffer]
|
||||||
|
|
||||||
def load_to_device_per_layer(
|
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 io_backend == "kernel":
|
||||||
if self.layout == "layer_first":
|
if self.layout == "layer_first":
|
||||||
@@ -995,7 +1061,7 @@ class AsymmetricMHATokenToKVPoolHost(MHATokenToKVPoolHost):
|
|||||||
def get_size_per_token(self):
|
def get_size_per_token(self):
|
||||||
self.head_num = self.device_pool.head_num
|
self.head_num = self.device_pool.head_num
|
||||||
self.head_dim = self.device_pool.head_dim
|
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
|
self.v_head_dim = self.device_pool.v_head_dim
|
||||||
return (
|
return (
|
||||||
(self.head_dim + self.v_head_dim)
|
(self.head_dim + self.v_head_dim)
|
||||||
@@ -1080,7 +1146,18 @@ class AsymmetricMHATokenToKVPoolHost(MHATokenToKVPoolHost):
|
|||||||
device_indices,
|
device_indices,
|
||||||
layer_id,
|
layer_id,
|
||||||
io_backend,
|
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 io_backend == "kernel":
|
||||||
if self.layout != "page_first":
|
if self.layout != "page_first":
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
@@ -1089,19 +1166,19 @@ class AsymmetricMHATokenToKVPoolHost(MHATokenToKVPoolHost):
|
|||||||
)
|
)
|
||||||
transfer_kv_per_layer_mla_pf_lf(
|
transfer_kv_per_layer_mla_pf_lf(
|
||||||
src=self.k_buffer,
|
src=self.k_buffer,
|
||||||
dst=device_pool.k_buffer[layer_id],
|
dst=device_pool.k_buffer[device_layer_id],
|
||||||
src_indices=host_indices,
|
src_indices=host_indices,
|
||||||
dst_indices=device_indices,
|
dst_indices=device_indices,
|
||||||
layer_id=layer_id,
|
layer_id=host_layer_id,
|
||||||
item_size=self._k_token_stride_size(),
|
item_size=self._k_token_stride_size(),
|
||||||
src_layout_dim=self._k_layout_dim(),
|
src_layout_dim=self._k_layout_dim(),
|
||||||
)
|
)
|
||||||
transfer_kv_per_layer_mla_pf_lf(
|
transfer_kv_per_layer_mla_pf_lf(
|
||||||
src=self.v_buffer,
|
src=self.v_buffer,
|
||||||
dst=device_pool.v_buffer[layer_id],
|
dst=device_pool.v_buffer[device_layer_id],
|
||||||
src_indices=host_indices,
|
src_indices=host_indices,
|
||||||
dst_indices=device_indices,
|
dst_indices=device_indices,
|
||||||
layer_id=layer_id,
|
layer_id=host_layer_id,
|
||||||
item_size=self._v_token_stride_size(),
|
item_size=self._v_token_stride_size(),
|
||||||
src_layout_dim=self._v_layout_dim(),
|
src_layout_dim=self._v_layout_dim(),
|
||||||
)
|
)
|
||||||
@@ -1114,18 +1191,18 @@ class AsymmetricMHATokenToKVPoolHost(MHATokenToKVPoolHost):
|
|||||||
)
|
)
|
||||||
transfer_kv_per_layer_direct_pf_lf(
|
transfer_kv_per_layer_direct_pf_lf(
|
||||||
src_ptrs=[self.k_buffer],
|
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,
|
src_indices=host_indices,
|
||||||
dst_indices=device_indices,
|
dst_indices=device_indices,
|
||||||
layer_id=layer_id,
|
layer_id=host_layer_id,
|
||||||
page_size=self.page_size,
|
page_size=self.page_size,
|
||||||
)
|
)
|
||||||
transfer_kv_per_layer_direct_pf_lf(
|
transfer_kv_per_layer_direct_pf_lf(
|
||||||
src_ptrs=[self.v_buffer],
|
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,
|
src_indices=host_indices,
|
||||||
dst_indices=device_indices,
|
dst_indices=device_indices,
|
||||||
layer_id=layer_id,
|
layer_id=host_layer_id,
|
||||||
page_size=self.page_size,
|
page_size=self.page_size,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
@@ -1137,6 +1214,12 @@ class AsymmetricMHATokenToKVPoolHost(MHATokenToKVPoolHost):
|
|||||||
def backup_from_device_all_layer(
|
def backup_from_device_all_layer(
|
||||||
self, device_pool, host_indices, device_indices, io_backend
|
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 io_backend == "kernel":
|
||||||
if self.layout != "page_first":
|
if self.layout != "page_first":
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
@@ -1145,7 +1228,7 @@ class AsymmetricMHATokenToKVPoolHost(MHATokenToKVPoolHost):
|
|||||||
)
|
)
|
||||||
if self.can_use_write_back_jit:
|
if self.can_use_write_back_jit:
|
||||||
jit_transfer_hicache_all_layer_mla_staged_lf_pf(
|
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,
|
src_indices=device_indices,
|
||||||
dst_indices=host_indices,
|
dst_indices=host_indices,
|
||||||
staging=self.staging_k_buffer,
|
staging=self.staging_k_buffer,
|
||||||
@@ -1153,7 +1236,7 @@ class AsymmetricMHATokenToKVPoolHost(MHATokenToKVPoolHost):
|
|||||||
page_size=self.page_size,
|
page_size=self.page_size,
|
||||||
)
|
)
|
||||||
jit_transfer_hicache_all_layer_mla_staged_lf_pf(
|
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,
|
src_indices=device_indices,
|
||||||
dst_indices=host_indices,
|
dst_indices=host_indices,
|
||||||
staging=self.staging_v_buffer,
|
staging=self.staging_v_buffer,
|
||||||
@@ -1162,7 +1245,7 @@ class AsymmetricMHATokenToKVPoolHost(MHATokenToKVPoolHost):
|
|||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
transfer_kv_all_layer_mla_lf_pf(
|
transfer_kv_all_layer_mla_lf_pf(
|
||||||
src_layers=device_pool.k_data_ptrs,
|
src_layers=device_k_data_ptrs,
|
||||||
dst=self.k_buffer,
|
dst=self.k_buffer,
|
||||||
src_indices=device_indices,
|
src_indices=device_indices,
|
||||||
dst_indices=host_indices,
|
dst_indices=host_indices,
|
||||||
@@ -1171,7 +1254,7 @@ class AsymmetricMHATokenToKVPoolHost(MHATokenToKVPoolHost):
|
|||||||
num_layers=self.layer_num,
|
num_layers=self.layer_num,
|
||||||
)
|
)
|
||||||
transfer_kv_all_layer_mla_lf_pf(
|
transfer_kv_all_layer_mla_lf_pf(
|
||||||
src_layers=device_pool.v_data_ptrs,
|
src_layers=device_v_data_ptrs,
|
||||||
dst=self.v_buffer,
|
dst=self.v_buffer,
|
||||||
src_indices=device_indices,
|
src_indices=device_indices,
|
||||||
dst_indices=host_indices,
|
dst_indices=host_indices,
|
||||||
@@ -1187,14 +1270,14 @@ class AsymmetricMHATokenToKVPoolHost(MHATokenToKVPoolHost):
|
|||||||
"'page_first_direct'."
|
"'page_first_direct'."
|
||||||
)
|
)
|
||||||
transfer_kv_all_layer_direct_lf_pf(
|
transfer_kv_all_layer_direct_lf_pf(
|
||||||
src_ptrs=device_pool.k_buffer,
|
src_ptrs=device_k_buffers,
|
||||||
dst_ptrs=[self.k_buffer],
|
dst_ptrs=[self.k_buffer],
|
||||||
src_indices=device_indices,
|
src_indices=device_indices,
|
||||||
dst_indices=host_indices,
|
dst_indices=host_indices,
|
||||||
page_size=self.page_size,
|
page_size=self.page_size,
|
||||||
)
|
)
|
||||||
transfer_kv_all_layer_direct_lf_pf(
|
transfer_kv_all_layer_direct_lf_pf(
|
||||||
src_ptrs=device_pool.v_buffer,
|
src_ptrs=device_v_buffers,
|
||||||
dst_ptrs=[self.v_buffer],
|
dst_ptrs=[self.v_buffer],
|
||||||
src_indices=device_indices,
|
src_indices=device_indices,
|
||||||
dst_indices=host_indices,
|
dst_indices=host_indices,
|
||||||
|
|||||||
@@ -1,7 +1,7 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import logging
|
import logging
|
||||||
from typing import Optional
|
from typing import Optional, Sequence
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
@@ -50,6 +50,7 @@ logger = logging.getLogger(__name__)
|
|||||||
|
|
||||||
class MLATokenToKVPoolHost(HiSparseHostPoolMixin, HostKVCache):
|
class MLATokenToKVPoolHost(HiSparseHostPoolMixin, HostKVCache):
|
||||||
device_pool: MLATokenToKVPool
|
device_pool: MLATokenToKVPool
|
||||||
|
mtp_draft_device_pools: tuple[MLATokenToKVPool, ...] = ()
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
@@ -62,12 +63,14 @@ class MLATokenToKVPoolHost(HiSparseHostPoolMixin, HostKVCache):
|
|||||||
device: str = "cpu",
|
device: str = "cpu",
|
||||||
allocator_type: str = "default",
|
allocator_type: str = "default",
|
||||||
override_kv_cache_dim: Optional[int] = None,
|
override_kv_cache_dim: Optional[int] = None,
|
||||||
|
mtp_draft_device_pools: Sequence[MLATokenToKVPool] = (),
|
||||||
dcp_size: int = 1,
|
dcp_size: int = 1,
|
||||||
dcp_rank: int = 0,
|
dcp_rank: int = 0,
|
||||||
*,
|
*,
|
||||||
pool_label: str = "kv",
|
pool_label: str = "kv",
|
||||||
):
|
):
|
||||||
self.override_kv_cache_dim = override_kv_cache_dim
|
self.override_kv_cache_dim = override_kv_cache_dim
|
||||||
|
self.mtp_draft_device_pools = tuple(mtp_draft_device_pools)
|
||||||
super().__init__(
|
super().__init__(
|
||||||
device_pool,
|
device_pool,
|
||||||
host_to_device_ratio,
|
host_to_device_ratio,
|
||||||
@@ -101,6 +104,14 @@ class MLATokenToKVPoolHost(HiSparseHostPoolMixin, HostKVCache):
|
|||||||
dtype=torch.uint64,
|
dtype=torch.uint64,
|
||||||
device=self.device_pool.device,
|
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()
|
self._init_write_back_staging_buffers()
|
||||||
|
|
||||||
def get_contiguous_buf_infos(self):
|
def get_contiguous_buf_infos(self):
|
||||||
@@ -114,7 +125,8 @@ class MLATokenToKVPoolHost(HiSparseHostPoolMixin, HostKVCache):
|
|||||||
def get_size_per_token(self):
|
def get_size_per_token(self):
|
||||||
self.kv_lora_rank = self.device_pool.kv_lora_rank
|
self.kv_lora_rank = self.device_pool.kv_lora_rank
|
||||||
self.qk_rope_head_dim = self.device_pool.qk_rope_head_dim
|
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_cache_dim = self.override_kv_cache_dim or (
|
||||||
self.kv_lora_rank + self.qk_rope_head_dim
|
self.kv_lora_rank + self.qk_rope_head_dim
|
||||||
)
|
)
|
||||||
@@ -229,28 +241,37 @@ class MLATokenToKVPoolHost(HiSparseHostPoolMixin, HostKVCache):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def load_to_device_per_layer(
|
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
|
return
|
||||||
host_indices = self.dcp_kernel_indices(host_indices)
|
host_indices = self.dcp_kernel_indices(host_indices)
|
||||||
device_indices = self.dcp_kernel_indices(device_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 io_backend == "kernel":
|
||||||
if self.layout == "layer_first":
|
if self.layout == "layer_first":
|
||||||
if self.can_use_jit:
|
if self.can_use_jit:
|
||||||
jit_transfer_hicache_one_layer_mla(
|
jit_transfer_hicache_one_layer_mla(
|
||||||
cache_dst=device_pool.kv_buffer[layer_id],
|
cache_dst=device_pool.kv_buffer[device_layer_id],
|
||||||
cache_src=self.kv_buffer[host_layer],
|
cache_src=self.kv_buffer[host_layer_id],
|
||||||
indices_dst=device_indices,
|
indices_dst=device_indices,
|
||||||
indices_src=host_indices,
|
indices_src=host_indices,
|
||||||
element_dim=self.kv_cache_dim,
|
element_dim=self.kv_cache_dim,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
transfer_kv_per_layer_mla(
|
transfer_kv_per_layer_mla(
|
||||||
src=self.kv_buffer[host_layer],
|
src=self.kv_buffer[host_layer_id],
|
||||||
dst=device_pool.kv_buffer[layer_id],
|
dst=device_pool.kv_buffer[device_layer_id],
|
||||||
src_indices=host_indices,
|
src_indices=host_indices,
|
||||||
dst_indices=device_indices,
|
dst_indices=device_indices,
|
||||||
item_size=self.token_stride_size,
|
item_size=self.token_stride_size,
|
||||||
@@ -258,8 +279,8 @@ class MLATokenToKVPoolHost(HiSparseHostPoolMixin, HostKVCache):
|
|||||||
elif self.layout == "page_first":
|
elif self.layout == "page_first":
|
||||||
if self.can_use_jit:
|
if self.can_use_jit:
|
||||||
jit_transfer_hicache_one_layer_mla(
|
jit_transfer_hicache_one_layer_mla(
|
||||||
cache_dst=device_pool.kv_buffer[layer_id],
|
cache_dst=device_pool.kv_buffer[device_layer_id],
|
||||||
cache_src=self.data_refs[host_layer],
|
cache_src=self.data_refs[host_layer_id],
|
||||||
indices_dst=device_indices,
|
indices_dst=device_indices,
|
||||||
indices_src=host_indices,
|
indices_src=host_indices,
|
||||||
element_dim=self.kv_cache_dim,
|
element_dim=self.kv_cache_dim,
|
||||||
@@ -267,10 +288,10 @@ class MLATokenToKVPoolHost(HiSparseHostPoolMixin, HostKVCache):
|
|||||||
else:
|
else:
|
||||||
transfer_kv_per_layer_mla_pf_lf(
|
transfer_kv_per_layer_mla_pf_lf(
|
||||||
src=self.kv_buffer,
|
src=self.kv_buffer,
|
||||||
dst=device_pool.kv_buffer[layer_id],
|
dst=device_pool.kv_buffer[device_layer_id],
|
||||||
src_indices=host_indices,
|
src_indices=host_indices,
|
||||||
dst_indices=device_indices,
|
dst_indices=device_indices,
|
||||||
layer_id=host_layer,
|
layer_id=host_layer_id,
|
||||||
item_size=self.token_stride_size,
|
item_size=self.token_stride_size,
|
||||||
src_layout_dim=self.layout_dim,
|
src_layout_dim=self.layout_dim,
|
||||||
)
|
)
|
||||||
@@ -279,8 +300,8 @@ class MLATokenToKVPoolHost(HiSparseHostPoolMixin, HostKVCache):
|
|||||||
elif io_backend == "direct":
|
elif io_backend == "direct":
|
||||||
if self.layout == "layer_first":
|
if self.layout == "layer_first":
|
||||||
transfer_kv_direct(
|
transfer_kv_direct(
|
||||||
src_layers=[self.kv_buffer[host_layer]],
|
src_layers=[self.kv_buffer[host_layer_id]],
|
||||||
dst_layers=[device_pool.kv_buffer[layer_id]],
|
dst_layers=[device_pool.kv_buffer[device_layer_id]],
|
||||||
src_indices=host_indices,
|
src_indices=host_indices,
|
||||||
dst_indices=device_indices,
|
dst_indices=device_indices,
|
||||||
page_size=self.page_size,
|
page_size=self.page_size,
|
||||||
@@ -288,10 +309,10 @@ class MLATokenToKVPoolHost(HiSparseHostPoolMixin, HostKVCache):
|
|||||||
elif self.layout == "page_first_direct":
|
elif self.layout == "page_first_direct":
|
||||||
transfer_kv_per_layer_direct_pf_lf(
|
transfer_kv_per_layer_direct_pf_lf(
|
||||||
src_ptrs=[self.kv_buffer],
|
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,
|
src_indices=host_indices,
|
||||||
dst_indices=device_indices,
|
dst_indices=device_indices,
|
||||||
layer_id=host_layer,
|
layer_id=host_layer_id,
|
||||||
page_size=self.page_size,
|
page_size=self.page_size,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
@@ -299,7 +320,7 @@ class MLATokenToKVPoolHost(HiSparseHostPoolMixin, HostKVCache):
|
|||||||
elif io_backend == "kernel_ascend":
|
elif io_backend == "kernel_ascend":
|
||||||
if self.layout == "page_first_kv_split":
|
if self.layout == "page_first_kv_split":
|
||||||
# Ascend-specific: transfer KV data for all layers when layer_id == 0
|
# 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(
|
transfer_kv_dim_exchange(
|
||||||
device_indices=device_indices,
|
device_indices=device_indices,
|
||||||
host_indices=host_indices,
|
host_indices=host_indices,
|
||||||
@@ -318,24 +339,34 @@ class MLATokenToKVPoolHost(HiSparseHostPoolMixin, HostKVCache):
|
|||||||
raise ValueError(f"Unsupported IO backend: {io_backend}")
|
raise ValueError(f"Unsupported IO backend: {io_backend}")
|
||||||
|
|
||||||
def _backup_from_device_per_layer(
|
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.
|
# 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 io_backend == "kernel":
|
||||||
if self.layout == "layer_first":
|
if self.layout == "layer_first":
|
||||||
if self.can_use_jit:
|
if self.can_use_jit:
|
||||||
jit_transfer_hicache_one_layer_mla(
|
jit_transfer_hicache_one_layer_mla(
|
||||||
cache_dst=self.kv_buffer[host_layer],
|
cache_dst=self.kv_buffer[host_layer_id],
|
||||||
cache_src=device_pool.kv_buffer[layer_id],
|
cache_src=device_pool.kv_buffer[device_layer_id],
|
||||||
indices_dst=host_indices,
|
indices_dst=host_indices,
|
||||||
indices_src=device_indices,
|
indices_src=device_indices,
|
||||||
element_dim=self.kv_cache_dim,
|
element_dim=self.kv_cache_dim,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
transfer_kv_per_layer_mla(
|
transfer_kv_per_layer_mla(
|
||||||
src=device_pool.kv_buffer[layer_id],
|
src=device_pool.kv_buffer[device_layer_id],
|
||||||
dst=self.kv_buffer[host_layer],
|
dst=self.kv_buffer[host_layer_id],
|
||||||
src_indices=device_indices,
|
src_indices=device_indices,
|
||||||
dst_indices=host_indices,
|
dst_indices=host_indices,
|
||||||
item_size=self.token_stride_size,
|
item_size=self.token_stride_size,
|
||||||
@@ -343,8 +374,8 @@ class MLATokenToKVPoolHost(HiSparseHostPoolMixin, HostKVCache):
|
|||||||
elif self.layout == "page_first":
|
elif self.layout == "page_first":
|
||||||
if self.can_use_jit:
|
if self.can_use_jit:
|
||||||
jit_transfer_hicache_one_layer_mla(
|
jit_transfer_hicache_one_layer_mla(
|
||||||
cache_dst=self.data_refs[host_layer],
|
cache_dst=self.data_refs[host_layer_id],
|
||||||
cache_src=device_pool.kv_buffer[layer_id],
|
cache_src=device_pool.kv_buffer[device_layer_id],
|
||||||
indices_dst=host_indices,
|
indices_dst=host_indices,
|
||||||
indices_src=device_indices,
|
indices_src=device_indices,
|
||||||
element_dim=self.kv_cache_dim,
|
element_dim=self.kv_cache_dim,
|
||||||
@@ -361,8 +392,8 @@ class MLATokenToKVPoolHost(HiSparseHostPoolMixin, HostKVCache):
|
|||||||
elif io_backend == "direct":
|
elif io_backend == "direct":
|
||||||
if self.layout == "layer_first":
|
if self.layout == "layer_first":
|
||||||
transfer_kv_direct(
|
transfer_kv_direct(
|
||||||
src_layers=[device_pool.kv_buffer[layer_id]],
|
src_layers=[device_pool.kv_buffer[device_layer_id]],
|
||||||
dst_layers=[self.kv_buffer[host_layer]],
|
dst_layers=[self.kv_buffer[host_layer_id]],
|
||||||
src_indices=device_indices,
|
src_indices=device_indices,
|
||||||
dst_indices=host_indices,
|
dst_indices=host_indices,
|
||||||
page_size=self.page_size,
|
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}"
|
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(
|
def backup_from_device_all_layer(
|
||||||
self, device_pool, host_indices, device_indices, io_backend
|
self, device_pool, host_indices, device_indices, io_backend
|
||||||
):
|
):
|
||||||
@@ -387,15 +423,30 @@ class MLATokenToKVPoolHost(HiSparseHostPoolMixin, HostKVCache):
|
|||||||
self._backup_from_device_per_layer(
|
self._backup_from_device_per_layer(
|
||||||
device_pool, host_indices, device_indices, layer_id, io_backend
|
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
|
return
|
||||||
|
|
||||||
|
device_data_ptrs, device_kv_buffers = self._resolve_device_transfer_buffers(
|
||||||
|
device_pool
|
||||||
|
)
|
||||||
|
|
||||||
if io_backend == "kernel":
|
if io_backend == "kernel":
|
||||||
if self.layout == "layer_first":
|
if self.layout == "layer_first":
|
||||||
if self.can_use_jit:
|
if self.can_use_jit:
|
||||||
jit_transfer_hicache_all_layer_mla(
|
jit_transfer_hicache_all_layer_mla(
|
||||||
ptr_dst=self.data_ptrs,
|
ptr_dst=self.data_ptrs,
|
||||||
indices_dst=host_indices,
|
indices_dst=host_indices,
|
||||||
ptr_src=device_pool.data_ptrs,
|
ptr_src=device_data_ptrs,
|
||||||
indices_src=device_indices,
|
indices_src=device_indices,
|
||||||
cache_dst_stride_bytes=self.token_stride_size,
|
cache_dst_stride_bytes=self.token_stride_size,
|
||||||
cache_src_stride_bytes=self.token_stride_size,
|
cache_src_stride_bytes=self.token_stride_size,
|
||||||
@@ -403,7 +454,7 @@ class MLATokenToKVPoolHost(HiSparseHostPoolMixin, HostKVCache):
|
|||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
transfer_kv_all_layer_mla(
|
transfer_kv_all_layer_mla(
|
||||||
src_layers=device_pool.data_ptrs,
|
src_layers=device_data_ptrs,
|
||||||
dst_layers=self.data_ptrs,
|
dst_layers=self.data_ptrs,
|
||||||
src_indices=device_indices,
|
src_indices=device_indices,
|
||||||
dst_indices=host_indices,
|
dst_indices=host_indices,
|
||||||
@@ -413,7 +464,7 @@ class MLATokenToKVPoolHost(HiSparseHostPoolMixin, HostKVCache):
|
|||||||
elif self.layout == "page_first":
|
elif self.layout == "page_first":
|
||||||
if self.can_use_write_back_jit:
|
if self.can_use_write_back_jit:
|
||||||
jit_transfer_hicache_all_layer_mla_staged_lf_pf(
|
jit_transfer_hicache_all_layer_mla_staged_lf_pf(
|
||||||
ptr_src=device_pool.data_ptrs,
|
ptr_src=device_data_ptrs,
|
||||||
src_indices=device_indices,
|
src_indices=device_indices,
|
||||||
dst_indices=host_indices,
|
dst_indices=host_indices,
|
||||||
staging=self.staging_buffer,
|
staging=self.staging_buffer,
|
||||||
@@ -422,7 +473,7 @@ class MLATokenToKVPoolHost(HiSparseHostPoolMixin, HostKVCache):
|
|||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
transfer_kv_all_layer_mla_lf_pf(
|
transfer_kv_all_layer_mla_lf_pf(
|
||||||
src_layers=device_pool.data_ptrs,
|
src_layers=device_data_ptrs,
|
||||||
dst=self.kv_buffer,
|
dst=self.kv_buffer,
|
||||||
src_indices=device_indices,
|
src_indices=device_indices,
|
||||||
dst_indices=host_indices,
|
dst_indices=host_indices,
|
||||||
@@ -435,7 +486,7 @@ class MLATokenToKVPoolHost(HiSparseHostPoolMixin, HostKVCache):
|
|||||||
elif io_backend == "direct":
|
elif io_backend == "direct":
|
||||||
if self.layout == "layer_first":
|
if self.layout == "layer_first":
|
||||||
transfer_kv_direct(
|
transfer_kv_direct(
|
||||||
src_layers=device_pool.kv_buffer,
|
src_layers=device_kv_buffers,
|
||||||
dst_layers=self.data_refs,
|
dst_layers=self.data_refs,
|
||||||
src_indices=device_indices,
|
src_indices=device_indices,
|
||||||
dst_indices=host_indices,
|
dst_indices=host_indices,
|
||||||
@@ -443,7 +494,7 @@ class MLATokenToKVPoolHost(HiSparseHostPoolMixin, HostKVCache):
|
|||||||
)
|
)
|
||||||
elif self.layout == "page_first_direct":
|
elif self.layout == "page_first_direct":
|
||||||
transfer_kv_all_layer_direct_lf_pf(
|
transfer_kv_all_layer_direct_lf_pf(
|
||||||
src_ptrs=device_pool.kv_buffer,
|
src_ptrs=device_kv_buffers,
|
||||||
dst_ptrs=[self.kv_buffer],
|
dst_ptrs=[self.kv_buffer],
|
||||||
src_indices=device_indices,
|
src_indices=device_indices,
|
||||||
dst_indices=host_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}_k",
|
||||||
f"_{self.mha_suffix}_{PoolName.DRAFT}_v",
|
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 (
|
elif pool_name in (
|
||||||
PoolName.INDEXER,
|
PoolName.INDEXER,
|
||||||
|
PoolName.DRAFT_INDEXER,
|
||||||
PoolName.DEEPSEEK_V4_C4,
|
PoolName.DEEPSEEK_V4_C4,
|
||||||
PoolName.DEEPSEEK_V4_C4_INDEXER,
|
PoolName.DEEPSEEK_V4_C4_INDEXER,
|
||||||
PoolName.DEEPSEEK_V4_C128,
|
PoolName.DEEPSEEK_V4_C128,
|
||||||
|
|||||||
@@ -78,6 +78,7 @@ if TYPE_CHECKING:
|
|||||||
from sglang.srt.mem_cache.hybrid_cache.hybrid_cache_controller import (
|
from sglang.srt.mem_cache.hybrid_cache.hybrid_cache_controller import (
|
||||||
PrefetchOperation,
|
PrefetchOperation,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.mem_cache.memory_pool_host import PoolEntry
|
||||||
from sglang.srt.server_args import ServerArgs
|
from sglang.srt.server_args import ServerArgs
|
||||||
|
|
||||||
|
|
||||||
@@ -395,6 +396,15 @@ class UnifiedRadixCache(BasePrefixCache):
|
|||||||
def register_sidecar_pool(self, spec: SidecarPoolSpec) -> None:
|
def register_sidecar_pool(self, spec: SidecarPoolSpec) -> None:
|
||||||
self.sidecar_pool_specs.append(spec)
|
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:
|
def release_host_resources(self) -> None:
|
||||||
if self.host_pool_group is not None:
|
if self.host_pool_group is not None:
|
||||||
self.host_pool_group.destroy()
|
self.host_pool_group.destroy()
|
||||||
|
|||||||
@@ -335,6 +335,7 @@ class ModelRunner:
|
|||||||
self.page_size = server_args.page_size
|
self.page_size = server_args.page_size
|
||||||
self.req_to_token_pool = req_to_token_pool
|
self.req_to_token_pool = req_to_token_pool
|
||||||
self.token_to_kv_pool_allocator = token_to_kv_pool_allocator
|
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 = model_config.is_hybrid_swa
|
||||||
self.is_hybrid_swa_compress = model_config.is_hybrid_swa_compress
|
self.is_hybrid_swa_compress = model_config.is_hybrid_swa_compress
|
||||||
self.use_mla_backend = self.model_config.attention_arch == AttentionArch.MLA
|
self.use_mla_backend = self.model_config.attention_arch == AttentionArch.MLA
|
||||||
|
|||||||
@@ -1,6 +1,8 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from enum import Enum
|
||||||
from typing import TYPE_CHECKING, Optional
|
from typing import TYPE_CHECKING, Optional
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
@@ -18,6 +20,38 @@ if TYPE_CHECKING:
|
|||||||
)
|
)
|
||||||
from sglang.srt.managers.tp_worker import TpModelWorker
|
from sglang.srt.managers.tp_worker import TpModelWorker
|
||||||
from sglang.srt.model_executor.model_runner import ModelRunner
|
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):
|
class EagleDraftWorkerBase(ABC):
|
||||||
@@ -111,10 +145,34 @@ class EagleDraftWorkerBase(ABC):
|
|||||||
|
|
||||||
|
|
||||||
class BaseSpecWorker(ABC):
|
class BaseSpecWorker(ABC):
|
||||||
|
_hicache_draft_plan = HiCacheDraftPlan()
|
||||||
|
|
||||||
def __init__(self) -> None:
|
def __init__(self) -> None:
|
||||||
self._additional_graph_memory_usage: dict[str, float] = {}
|
self._additional_graph_memory_usage: dict[str, float] = {}
|
||||||
self._additional_graph_time_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
|
@property
|
||||||
def target_worker(self) -> TpModelWorker:
|
def target_worker(self) -> TpModelWorker:
|
||||||
return self._target_worker
|
return self._target_worker
|
||||||
@@ -173,6 +231,42 @@ class BaseSpecWorker(ABC):
|
|||||||
# TODO: move this method to BaseTpWorker and call through self.model_runner
|
# TODO: move this method to BaseTpWorker and call through self.model_runner
|
||||||
pass
|
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(
|
def alloc_memory_pool(
|
||||||
self,
|
self,
|
||||||
memory_pool_config=None,
|
memory_pool_config=None,
|
||||||
|
|||||||
@@ -19,9 +19,10 @@ from sglang.test.test_utils import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
DSV4_FLASH_MODEL = "sgl-project/DeepSeek-V4-Flash-FP8"
|
DSV4_FLASH_MODEL = "sgl-project/DeepSeek-V4-Flash-FP8"
|
||||||
|
DSV4_DSPARK_MODEL = "deepseek-ai/DeepSeek-V4-Flash-DSpark"
|
||||||
DSV4_FLASH_LAUNCH_TIMEOUT = 3600
|
DSV4_FLASH_LAUNCH_TIMEOUT = 3600
|
||||||
|
|
||||||
register_cuda_ci(est_time=1000, stage="extra-b", runner_config="4-gpu-h100")
|
register_cuda_ci(est_time=1500, stage="extra-b", runner_config="4-gpu-h100")
|
||||||
|
|
||||||
|
|
||||||
def _assert_dsv4_decode_cached_tokens(result, history_len, output_len, label):
|
def _assert_dsv4_decode_cached_tokens(result, history_len, output_len, label):
|
||||||
@@ -330,5 +331,64 @@ class TestUnifiedDeepSeekV4FlashEagleHiCacheL3(AccuracyTwoPassMixin, CustomTestC
|
|||||||
self.assertEqual(cached_details.get("storage_backend"), "HiCacheFile")
|
self.assertEqual(cached_details.get("storage_backend"), "HiCacheFile")
|
||||||
|
|
||||||
|
|
||||||
|
class TestUnifiedDeepSeekV4FlashDSparkHiCacheL3(
|
||||||
|
TestUnifiedDeepSeekV4FlashEagleHiCacheL3
|
||||||
|
):
|
||||||
|
"""DeepSeek V4 Flash DSpark + HiCache L3 should load from storage."""
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def setUpClass(cls):
|
||||||
|
cls.model = DSV4_DSPARK_MODEL
|
||||||
|
cls.base_url = DEFAULT_URL_FOR_TEST
|
||||||
|
cls.hicache_dir = tempfile.mkdtemp(prefix="hicache_l3_dspark_dsv4_")
|
||||||
|
cls.process = popen_launch_server(
|
||||||
|
cls.model,
|
||||||
|
cls.base_url,
|
||||||
|
timeout=DSV4_FLASH_LAUNCH_TIMEOUT,
|
||||||
|
other_args=[
|
||||||
|
"--trust-remote-code",
|
||||||
|
"--tp-size",
|
||||||
|
"4",
|
||||||
|
"--attention-backend",
|
||||||
|
"compressed",
|
||||||
|
"--page-size",
|
||||||
|
str(cls.page_size),
|
||||||
|
"--chunked-prefill-size",
|
||||||
|
"8192",
|
||||||
|
"--mem-fraction-static",
|
||||||
|
"0.95",
|
||||||
|
"--disable-shared-experts-fusion",
|
||||||
|
"--enable-hierarchical-cache",
|
||||||
|
"--hicache-ratio",
|
||||||
|
"2",
|
||||||
|
"--hicache-write-policy",
|
||||||
|
"write_through",
|
||||||
|
"--hicache-storage-prefetch-policy",
|
||||||
|
"wait_complete",
|
||||||
|
"--hicache-io-backend",
|
||||||
|
"kernel",
|
||||||
|
"--hicache-mem-layout",
|
||||||
|
"page_first",
|
||||||
|
"--hicache-storage-backend",
|
||||||
|
"file",
|
||||||
|
"--enable-cache-report",
|
||||||
|
"--swa-full-tokens-ratio",
|
||||||
|
"0.25",
|
||||||
|
"--max-total-tokens",
|
||||||
|
"20000",
|
||||||
|
"--max-running-requests",
|
||||||
|
"4",
|
||||||
|
"--moe-runner-backend",
|
||||||
|
"marlin",
|
||||||
|
"--speculative-algorithm",
|
||||||
|
"DSPARK",
|
||||||
|
],
|
||||||
|
env={
|
||||||
|
"SGLANG_ENABLE_UNIFIED_RADIX_TREE": "1",
|
||||||
|
"SGLANG_HICACHE_FILE_BACKEND_STORAGE_DIR": cls.hicache_dir,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
@@ -179,6 +179,14 @@ class TestUnifiedMambaHiCacheL3(AccuracyTwoPassMixin, CustomTestCase):
|
|||||||
"--max-mamba-cache-size",
|
"--max-mamba-cache-size",
|
||||||
"500",
|
"500",
|
||||||
"--weight-loader-prefetch-checkpoints",
|
"--weight-loader-prefetch-checkpoints",
|
||||||
|
"--speculative-algorithm",
|
||||||
|
"NEXTN",
|
||||||
|
"--speculative-num-steps",
|
||||||
|
"3",
|
||||||
|
"--speculative-eagle-topk",
|
||||||
|
"1",
|
||||||
|
"--speculative-num-draft-tokens",
|
||||||
|
"4",
|
||||||
],
|
],
|
||||||
env={
|
env={
|
||||||
"SGLANG_ENABLE_UNIFIED_RADIX_TREE": "1",
|
"SGLANG_ENABLE_UNIFIED_RADIX_TREE": "1",
|
||||||
|
|||||||
@@ -175,6 +175,14 @@ class TestGLM5HiRadixCacheL3Accuracy(AccuracyTwoPassMixin, CustomTestCase):
|
|||||||
"page_first",
|
"page_first",
|
||||||
"--hicache-storage-backend",
|
"--hicache-storage-backend",
|
||||||
"file",
|
"file",
|
||||||
|
"--speculative-algorithm",
|
||||||
|
"EAGLE",
|
||||||
|
"--speculative-num-steps",
|
||||||
|
"3",
|
||||||
|
"--speculative-eagle-topk",
|
||||||
|
"1",
|
||||||
|
"--speculative-num-draft-tokens",
|
||||||
|
"4",
|
||||||
],
|
],
|
||||||
env={
|
env={
|
||||||
"SGLANG_HICACHE_FILE_BACKEND_STORAGE_DIR": cls.hicache_dir,
|
"SGLANG_HICACHE_FILE_BACKEND_STORAGE_DIR": cls.hicache_dir,
|
||||||
@@ -223,6 +231,14 @@ class TestGLM5UnifiedRadixCacheL3Accuracy(AccuracyTwoPassMixin, CustomTestCase):
|
|||||||
"page_first",
|
"page_first",
|
||||||
"--hicache-storage-backend",
|
"--hicache-storage-backend",
|
||||||
"file",
|
"file",
|
||||||
|
"--speculative-algorithm",
|
||||||
|
"EAGLE",
|
||||||
|
"--speculative-num-steps",
|
||||||
|
"3",
|
||||||
|
"--speculative-eagle-topk",
|
||||||
|
"1",
|
||||||
|
"--speculative-num-draft-tokens",
|
||||||
|
"4",
|
||||||
],
|
],
|
||||||
env={
|
env={
|
||||||
"SGLANG_HICACHE_FILE_BACKEND_STORAGE_DIR": cls.hicache_dir,
|
"SGLANG_HICACHE_FILE_BACKEND_STORAGE_DIR": cls.hicache_dir,
|
||||||
|
|||||||
Reference in New Issue
Block a user