Reuse prepared HiCache load descriptors across CP prefill layers

CP shared-KV bs>1 cache-hit loads already merge request load ops, but the host pool still rebuilt layer-invariant mapping work from the same host/device indices. Introduce a PreparedLoadDescriptor lifecycle around begin/end load, wire MLA KV and NSA index H2D loads through tai-kernel prepared submit when available, and add timing hooks plus regression coverage for descriptor reuse and explicit fallback logging. Record the P4/P6b design and benchmark results in the advanced feature notes.

Constraint: Radix residency and allocator decisions remain synchronous; only the data-transfer descriptor is prepared for per-layer async submit.

Constraint: Production fast path must not silently fall back when tai prepared H2D support is missing.

Rejected: Cross-batch descriptor reuse | descriptor lifetime and tensor ownership are only safe within one load operation.

Rejected: Change L2->L1 scheduling to layer-ahead prefetch in this commit | that is a separate lifecycle change after descriptor reuse is stable.

Confidence: medium

Scope-risk: moderate

Directive: Keep LayerDoneCounter per-layer readiness semantics; do not replace with all-layer waits.

Tested: python -m py_compile python/sglang/srt/mem_cache/memory_pool_host.py python/sglang/srt/managers/cache_controller.py

Tested: Remote g0034:cjy-glm5-new PYTHONPATH=python python -m pytest -q test/registered/unit/managers/test_hicache_controller_cp.py (88 passed)

Tested: Remote tai-kernel prepared descriptor CUDA test (6 passed) and P4 benchmark full matrix (90 rows)

Not-tested: ETE replay/GSM8K cache-hit correctness after this commit

Not-tested: Layer-ahead L2->L1 prefetch scheduling

Co-authored-by: OmX <omx@oh-my-codex.dev>
This commit is contained in:
laoyao0822
2026-06-11 05:09:41 +08:00
parent adf357b02c
commit 7284a469a2
4 changed files with 1375 additions and 10 deletions

View File

@@ -78,6 +78,23 @@ def _cp_shared_kv_bs_gt1_cache_timing(
)
def _cp_shared_kv_bs_gt1_cache_timing_start() -> Optional[float]:
if not envs.SGLANG_CP_SHARED_KV_BS_GT1_TIMING.get():
return None
return time.perf_counter()
def _cp_shared_kv_bs_gt1_cache_timing_log(
key: str,
start_time: Optional[float],
message: str,
*args,
) -> None:
if start_time is None:
return
_cp_shared_kv_bs_gt1_cache_timing(key, start_time, message, *args)
class LayerLoadingEvent:
def __init__(self, num_layers: int):
self._num_layers = num_layers
@@ -1752,6 +1769,7 @@ class HiCacheController:
try:
with device_module.stream(self.load_stream):
producer_event.start_event.wait(self.load_stream)
prepare_start = _cp_shared_kv_bs_gt1_cache_timing_start()
if op is not None:
target_load_op_started = begin_load_op(
self.mem_pool_host, host_indices, device_indices
@@ -1762,10 +1780,25 @@ class HiCacheController:
draft_host_indices,
draft_device_indices,
)
_cp_shared_kv_bs_gt1_cache_timing_log(
"prepare_load_descriptor",
prepare_start,
"target_tokens=%s draft_tokens=%s target_started=%s draft_started=%s",
int(host_indices.numel()) if host_indices is not None else 0,
int(draft_host_indices.numel())
if draft_host_indices is not None
else 0,
target_load_op_started,
draft_load_op_started,
)
draft_layer_num = (
self.draft_mem_pool_device.layer_num if draft_op is not None else 0
)
layer_loop_start = _cp_shared_kv_bs_gt1_cache_timing_start()
for i in range(max(self.layer_num, draft_layer_num)):
layer_submit_start = _cp_shared_kv_bs_gt1_cache_timing_start()
submitted_target = False
submitted_draft = False
if draft_op is not None and i < draft_layer_num:
if len(draft_host_indices) > 0:
self.draft_mem_pool_host.load_to_device_per_layer(
@@ -1775,6 +1808,7 @@ class HiCacheController:
i,
self.io_backend,
)
submitted_draft = True
if op is not None and i < self.layer_num:
if len(host_indices) > 0:
self.mem_pool_host.load_to_device_per_layer(
@@ -1784,9 +1818,32 @@ class HiCacheController:
i,
self.io_backend,
)
submitted_target = True
producer_event.complete(i)
elif op is None and i < self.layer_num:
producer_event.complete(i)
_cp_shared_kv_bs_gt1_cache_timing_log(
"submit_h2d_layer_per_call_slow",
layer_submit_start,
"layer_id=%s target=%s draft=%s target_tokens=%s draft_tokens=%s",
i,
submitted_target,
submitted_draft,
int(host_indices.numel()) if host_indices is not None else 0,
int(draft_host_indices.numel())
if draft_host_indices is not None
else 0,
)
_cp_shared_kv_bs_gt1_cache_timing_log(
"submit_h2d_layer_loop",
layer_loop_start,
"layers=%s target_tokens=%s draft_tokens=%s",
max(self.layer_num, draft_layer_num),
int(host_indices.numel()) if host_indices is not None else 0,
int(draft_host_indices.numel())
if draft_host_indices is not None
else 0,
)
# NOTE: We must save the host indices and device indices here,
# this is because we need to guarantee that these tensors are
# still alive when the load stream is executing.
@@ -1799,10 +1856,18 @@ class HiCacheController:
if draft_device_indices is not None and draft_device_indices.is_cuda:
draft_device_indices.record_stream(self.load_stream)
finally:
end_start = _cp_shared_kv_bs_gt1_cache_timing_start()
if draft_load_op_started:
end_load_op(self.draft_mem_pool_host)
if target_load_op_started:
end_load_op(self.mem_pool_host)
_cp_shared_kv_bs_gt1_cache_timing_log(
"end_load_descriptor",
end_start,
"target_started=%s draft_started=%s",
target_load_op_started,
draft_load_op_started,
)
self.ack_load_queue.append(
HiCacheAck(

View File

@@ -5,8 +5,9 @@ import logging
import threading
import weakref
from collections import defaultdict
from dataclasses import dataclass
from functools import lru_cache, wraps
from typing import Optional
from typing import Optional, Tuple
import psutil
import torch
@@ -54,6 +55,26 @@ if _is_npu:
logger = logging.getLogger(__name__)
@dataclass
class PreparedLoadDescriptor:
"""Layer-invariant metadata for one host->device HiCache load op."""
host_indices: torch.Tensor
device_indices: torch.Tensor
host_page_indices: Optional[torch.Tensor]
device_page_indices: Optional[torch.Tensor]
page_size: int
io_backend: str
layout: str
num_tokens: int
num_pages: int
index_host_page_indices: Optional[torch.Tensor] = None
index_device_page_indices: Optional[torch.Tensor] = None
index_active_layer_ids: Tuple[int, ...] = ()
tai_h2d_descriptor: Optional[object] = None
tai_index_h2d_descriptor: Optional[object] = None
@lru_cache(maxsize=1)
def _load_tai_transfer_kv_per_layer_mla_lf_pf():
try:
@@ -132,6 +153,47 @@ def _load_tai_transfer_kv_per_layer_direct_lpf_lf():
) from exc
@lru_cache(maxsize=1)
def _load_tai_prepare_h2d_page_descriptor():
try:
from tai_kernel.nsa_prefill import prepare_h2d_page_descriptor
return prepare_h2d_page_descriptor
except Exception as exc:
raise RuntimeError(
"[CP_HICACHE_FALLBACK][missing_tai_prepared_h2d_descriptor] "
"tai_kernel.nsa_prefill.prepare_h2d_page_descriptor is unavailable. "
"Falling back to legacy per-layer direct H2D submit."
) from exc
@lru_cache(maxsize=1)
def _load_tai_submit_h2d_layer():
try:
from tai_kernel.nsa_prefill import submit_h2d_layer
return submit_h2d_layer
except Exception as exc:
raise RuntimeError(
"[CP_HICACHE_FALLBACK][missing_tai_submit_h2d_layer] "
"tai_kernel.nsa_prefill.submit_h2d_layer is unavailable. "
"Falling back to legacy per-layer direct H2D submit."
) from exc
@lru_cache(maxsize=1)
def _load_tai_destroy_h2d_page_descriptor():
try:
from tai_kernel.nsa_prefill import destroy_h2d_page_descriptor
return destroy_h2d_page_descriptor
except Exception as exc:
raise RuntimeError(
"[CP_HICACHE_FALLBACK][missing_tai_destroy_h2d_page_descriptor] "
"tai_kernel.nsa_prefill.destroy_h2d_page_descriptor is unavailable."
) from exc
def _raise_layer_page_first_page_buffer_meta_unsupported() -> None:
raise RuntimeError(
"[CP_HICACHE_FAILFAST][layer_page_first_page_buffer_meta_unsupported] "
@@ -315,6 +377,9 @@ class HostKVCache(abc.ABC):
)
self.kv_buffer = self.init_kv_buffer()
self._active_load_descriptor: Optional[PreparedLoadDescriptor] = None
self._missing_load_descriptor_warning_emitted = False
self._missing_tai_prepared_h2d_warning_emitted = False
# A lock for synchronized operations on memory allocation and state transitions.
self.lock = threading.RLock()
@@ -341,9 +406,151 @@ class HostKVCache(abc.ABC):
self, host_indices: torch.Tensor, device_indices: torch.Tensor, io_backend: str
) -> None:
"""Prepare layer-invariant metadata for one host->device load op."""
host_indices = host_indices.contiguous()
device_indices = device_indices.contiguous()
host_page_indices, device_page_indices = self._prepare_load_page_indices(
host_indices, device_indices
)
num_tokens = int(host_indices.numel())
self._active_load_descriptor = PreparedLoadDescriptor(
host_indices=host_indices,
device_indices=device_indices,
host_page_indices=host_page_indices,
device_page_indices=device_page_indices,
page_size=int(self.page_size),
io_backend=io_backend,
layout=self.layout,
num_tokens=num_tokens,
num_pages=(num_tokens // int(self.page_size))
if num_tokens % int(self.page_size) == 0
else 0,
)
self._active_load_descriptor.tai_h2d_descriptor = (
self._try_prepare_tai_h2d_descriptor(
self._active_load_descriptor,
host_indices=host_indices,
device_indices=device_indices,
page_size=int(self.page_size),
)
)
def end_load_to_device_op(self) -> None:
"""Release metadata prepared by ``begin_load_to_device_op``."""
descriptor = getattr(self, "_active_load_descriptor", None)
if descriptor is not None:
self._destroy_tai_h2d_descriptor(descriptor.tai_h2d_descriptor)
if descriptor.tai_index_h2d_descriptor is not descriptor.tai_h2d_descriptor:
self._destroy_tai_h2d_descriptor(descriptor.tai_index_h2d_descriptor)
self._active_load_descriptor = None
def _prepare_load_page_indices(
self, host_indices: torch.Tensor, device_indices: torch.Tensor
) -> tuple[Optional[torch.Tensor], Optional[torch.Tensor]]:
if self.layout not in ("page_first_direct", "layer_page_first"):
return None, None
validate_page_aligned_token_indices(host_indices, self.page_size, "host_indices")
validate_page_aligned_token_indices(
device_indices, self.page_size, "device_indices"
)
host_page_indices = (
host_indices.reshape(-1, self.page_size)[:, 0] // self.page_size
)
device_page_indices = (
device_indices.reshape(-1, self.page_size)[:, 0] // self.page_size
)
return host_page_indices.contiguous(), device_page_indices.contiguous()
def _get_active_load_descriptor(
self, io_backend: str
) -> Optional[PreparedLoadDescriptor]:
descriptor = getattr(self, "_active_load_descriptor", None)
if descriptor is None:
self._log_missing_load_descriptor_once(io_backend)
return None
if descriptor.io_backend != io_backend or descriptor.layout != self.layout:
raise RuntimeError(
"[CP_HICACHE_FAILFAST][load_descriptor_contract_mismatch] "
f"active descriptor io/layout=({descriptor.io_backend}, {descriptor.layout}) "
f"but current load uses ({io_backend}, {self.layout})"
)
return descriptor
def _log_missing_load_descriptor_once(self, io_backend: str) -> None:
if io_backend != "direct":
return
if self.layout not in ("page_first_direct", "layer_page_first"):
return
if getattr(self, "_missing_load_descriptor_warning_emitted", False):
return
self._missing_load_descriptor_warning_emitted = True
logger.warning(
"[CP_HICACHE_FALLBACK][load_descriptor] "
"reason=missing_prepared_descriptor pool=%s layout=%s io_backend=%s",
type(self).__name__,
self.layout,
io_backend,
)
def _try_prepare_tai_h2d_descriptor(
self,
descriptor: PreparedLoadDescriptor,
*,
host_indices: torch.Tensor,
device_indices: torch.Tensor,
page_size: int,
) -> Optional[object]:
if descriptor.io_backend != "direct":
return None
if descriptor.layout not in ("page_first_direct", "layer_page_first"):
return None
try:
prepare = _load_tai_prepare_h2d_page_descriptor()
except RuntimeError as exc:
self._log_missing_tai_prepared_h2d_once(str(exc))
return None
try:
tai_descriptor = prepare(
host_indices,
device_indices,
page_size=page_size,
layout=descriptor.layout,
)
if not bool(getattr(tai_descriptor, "cpp_prepared", False)):
self._log_missing_tai_prepared_h2d_once(
"tai prepared descriptor Python API is present, but the "
"C++ prepared submit data is unavailable; legacy per-layer "
"direct submit remains active."
)
return tai_descriptor
except Exception as exc:
raise RuntimeError(
"[CP_HICACHE_FAILFAST][tai_prepared_h2d_descriptor_failed] "
f"pool={type(self).__name__} layout={descriptor.layout} "
f"page_size={page_size} tokens={int(host_indices.numel())}"
) from exc
def _log_missing_tai_prepared_h2d_once(self, message: str) -> None:
if getattr(self, "_missing_tai_prepared_h2d_warning_emitted", False):
return
self._missing_tai_prepared_h2d_warning_emitted = True
logger.warning(
"[CP_HICACHE_FALLBACK][tai_prepared_h2d_descriptor] "
"reason=missing_tai_api %s",
message,
)
def _destroy_tai_h2d_descriptor(self, tai_descriptor: Optional[object]) -> None:
if tai_descriptor is None:
return
try:
destroy = _load_tai_destroy_h2d_page_descriptor()
destroy(tai_descriptor)
except Exception as exc:
logger.warning(
"[CP_HICACHE_FALLBACK][tai_prepared_h2d_descriptor] "
"reason=destroy_failed error=%s",
exc,
)
@abc.abstractmethod
def backup_from_device_per_layer(
@@ -1503,6 +1710,10 @@ class MLATokenToKVPoolHost(HostKVCache):
def load_to_device_per_layer(
self, device_pool, host_indices, device_indices, layer_id, io_backend
):
descriptor = self._get_active_load_descriptor(io_backend)
if descriptor is not None:
host_indices = descriptor.host_indices
device_indices = descriptor.device_indices
if io_backend == "kernel":
if self.layout == "layer_first":
transfer_kv_per_layer_mla(
@@ -1534,6 +1745,14 @@ class MLATokenToKVPoolHost(HostKVCache):
page_size=self.page_size,
)
elif self.layout == "page_first_direct":
if descriptor is not None and descriptor.tai_h2d_descriptor is not None:
_load_tai_submit_h2d_layer()(
descriptor.tai_h2d_descriptor,
[self.kv_buffer],
[device_pool.kv_buffer[layer_id]],
layer_id=layer_id,
)
return
_load_tai_transfer_kv_per_layer_direct_pf_lf()(
src_ptrs=[self.kv_buffer],
dst_ptrs=[device_pool.kv_buffer[layer_id]],
@@ -1543,6 +1762,14 @@ class MLATokenToKVPoolHost(HostKVCache):
page_size=self.page_size,
)
elif self.layout == "layer_page_first":
if descriptor is not None and descriptor.tai_h2d_descriptor is not None:
_load_tai_submit_h2d_layer()(
descriptor.tai_h2d_descriptor,
[self.kv_buffer],
[device_pool.kv_buffer[layer_id]],
layer_id=layer_id,
)
return
_load_tai_transfer_kv_per_layer_direct_lpf_lf()(
src_ptrs=[self.kv_buffer],
dst_ptrs=[device_pool.kv_buffer[layer_id]],
@@ -2033,13 +2260,32 @@ class NSATokenToKVPoolHost(MLATokenToKVPoolHost):
def begin_load_to_device_op(
self, host_indices: torch.Tensor, device_indices: torch.Tensor, io_backend: str
) -> None:
super().begin_load_to_device_op(host_indices, device_indices, io_backend)
self._active_load_indexer_page_indices = None
self._active_load_indexer_page_indices = self._get_indexer_page_indices(
host_indices, device_indices
descriptor = self._get_active_load_descriptor(io_backend)
host_page_indices, device_page_indices = self._get_indexer_page_indices(
descriptor.host_indices,
descriptor.device_indices,
)
descriptor.index_host_page_indices = host_page_indices
descriptor.index_device_page_indices = device_page_indices
descriptor.index_active_layer_ids = tuple(
getattr(self, "index_active_layer_ids", ())
)
descriptor.tai_index_h2d_descriptor = self._try_prepare_tai_h2d_descriptor(
descriptor,
host_indices=host_page_indices,
device_indices=device_page_indices,
page_size=1,
)
self._active_load_indexer_page_indices = (
descriptor.index_host_page_indices,
descriptor.index_device_page_indices,
)
def end_load_to_device_op(self) -> None:
self._active_load_indexer_page_indices = None
super().end_load_to_device_op()
def _load_indexer_to_device_per_layer(
self, device_pool, host_indices, device_indices, layer_id, io_backend
@@ -2048,13 +2294,26 @@ class NSATokenToKVPoolHost(MLATokenToKVPoolHost):
return
device_layer_slot = self._device_index_layer_slot(device_pool, layer_id)
host_layer_slot = self._host_index_layer_slot(layer_id)
prepared_indices = getattr(self, "_active_load_indexer_page_indices", None)
if prepared_indices is None:
descriptor = self._get_active_load_descriptor(io_backend)
if (
descriptor is not None
and descriptor.index_host_page_indices is not None
and descriptor.index_device_page_indices is not None
):
host_page_indices = descriptor.index_host_page_indices
device_page_indices = descriptor.index_device_page_indices
else:
prepared_indices = getattr(self, "_active_load_indexer_page_indices", None)
if prepared_indices is None:
host_page_indices, device_page_indices = self._get_indexer_page_indices(
host_indices, device_indices
)
else:
host_page_indices, device_page_indices = prepared_indices
if host_page_indices is None or device_page_indices is None:
host_page_indices, device_page_indices = self._get_indexer_page_indices(
host_indices, device_indices
)
else:
host_page_indices, device_page_indices = prepared_indices
use_kernel = io_backend == "kernel" and self.indexer_page_stride_size % 8 == 0
if use_kernel:
if self.layout == "layer_first":
@@ -2089,6 +2348,17 @@ class NSATokenToKVPoolHost(MLATokenToKVPoolHost):
page_size=1,
)
elif self.layout == "page_first_direct":
if (
descriptor is not None
and descriptor.tai_index_h2d_descriptor is not None
):
_load_tai_submit_h2d_layer()(
descriptor.tai_index_h2d_descriptor,
[self.index_k_with_scale_buffer],
[device_pool.index_k_with_scale_buffer[device_layer_slot]],
layer_id=host_layer_slot,
)
return
_load_tai_transfer_kv_per_layer_direct_pf_lf()(
src_ptrs=[self.index_k_with_scale_buffer],
dst_ptrs=[
@@ -2100,6 +2370,17 @@ class NSATokenToKVPoolHost(MLATokenToKVPoolHost):
page_size=1,
)
elif self.layout == "layer_page_first":
if (
descriptor is not None
and descriptor.tai_index_h2d_descriptor is not None
):
_load_tai_submit_h2d_layer()(
descriptor.tai_index_h2d_descriptor,
[self.index_k_with_scale_buffer],
[device_pool.index_k_with_scale_buffer[device_layer_slot]],
layer_id=host_layer_slot,
)
return
_load_tai_transfer_kv_per_layer_direct_lpf_lf()(
src_ptrs=[self.index_k_with_scale_buffer],
dst_ptrs=[