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:
@@ -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(
|
||||
|
||||
@@ -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=[
|
||||
|
||||
Reference in New Issue
Block a user