[HiCache]Support hybrid pool staged H2D kernel (#28434)
Co-authored-by: hzh0425 <hzh0425@apache.org>
This commit is contained in:
@@ -26,14 +26,33 @@ def _jit_hicache_module(*, element_size: int, unroll: int, block_quota: int) ->
|
||||
*args,
|
||||
cuda_files=[
|
||||
"kvcacheio/hicache.cuh",
|
||||
"kvcacheio/relayout.cuh",
|
||||
"kvcacheio/staged_write_back.cuh",
|
||||
],
|
||||
cuda_wrappers=[
|
||||
("launch_one", f"&HiCacheKernel<{args}>::run_one"),
|
||||
("launch_all", f"&HiCacheKernel<{args}>::run_all"),
|
||||
("launch_one_mla", f"&HiCacheKernel<{args}>::run_one_mla"),
|
||||
("launch_all_mla", f"&HiCacheKernel<{args}>::run_all_mla"),
|
||||
],
|
||||
)
|
||||
|
||||
|
||||
@cache_once
|
||||
def _jit_hicache_staged_module(
|
||||
*, element_size: int, unroll: int, block_quota: int
|
||||
) -> Module:
|
||||
args = make_cpp_args(
|
||||
element_size,
|
||||
unroll,
|
||||
block_quota,
|
||||
1024, # num_threads, kept for template compatibility
|
||||
)
|
||||
return load_jit(
|
||||
"hicache_staged",
|
||||
*args,
|
||||
cuda_files=[
|
||||
"kvcacheio/staged_write_back.cuh",
|
||||
],
|
||||
cuda_wrappers=[
|
||||
(
|
||||
"launch_all_lf_pf_staged",
|
||||
f"&HiCacheStagedWriteBackKernel<{args}>::run_all_lf_pf_staged",
|
||||
@@ -70,6 +89,30 @@ def can_use_hicache_jit_kernel(
|
||||
return False
|
||||
|
||||
|
||||
def can_use_write_back_jit_kernel(
|
||||
*,
|
||||
element_size: int,
|
||||
unroll: int | None = None, # can be tuned for performance
|
||||
block_quota: int | None = None, # can be tuned for less interference
|
||||
) -> bool:
|
||||
logger = logging.getLogger(__name__)
|
||||
if element_size % 16 != 0:
|
||||
logger.warning(f"Unsupported {element_size = } for staged JIT HiCache kernel")
|
||||
return False
|
||||
try:
|
||||
unroll = unroll or _default_unroll(element_size)
|
||||
block_quota = block_quota or DEFAULT_BLOCK_QUOTA
|
||||
_jit_hicache_staged_module(
|
||||
element_size=element_size,
|
||||
unroll=unroll,
|
||||
block_quota=block_quota,
|
||||
)
|
||||
return True
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to load staged JIT HiCache kernel: {e}")
|
||||
return False
|
||||
|
||||
|
||||
def _default_unroll(element_size: int) -> int:
|
||||
if element_size <= 512:
|
||||
return 4
|
||||
@@ -238,7 +281,7 @@ def transfer_hicache_all_layer_staged_lf_pf(
|
||||
block_quota = block_quota or DEFAULT_BLOCK_QUOTA
|
||||
unroll = unroll or _default_unroll(element_size)
|
||||
src_page_indices = src_indices[::page_size].contiguous()
|
||||
module = _jit_hicache_module(
|
||||
module = _jit_hicache_staged_module(
|
||||
element_size=element_size,
|
||||
unroll=unroll,
|
||||
block_quota=block_quota,
|
||||
@@ -284,7 +327,7 @@ def transfer_hicache_all_layer_mla_staged_lf_pf(
|
||||
block_quota = block_quota or DEFAULT_BLOCK_QUOTA
|
||||
unroll = unroll or _default_unroll(element_size)
|
||||
src_page_indices = src_indices[::page_size].contiguous()
|
||||
module = _jit_hicache_module(
|
||||
module = _jit_hicache_staged_module(
|
||||
element_size=element_size,
|
||||
unroll=unroll,
|
||||
block_quota=block_quota,
|
||||
|
||||
@@ -726,10 +726,12 @@ class HiCacheController:
|
||||
return
|
||||
|
||||
op = CacheOperation.merge_ops(self.write_queue)
|
||||
# For now, kernel write-back keeps host indices on CPU only for page_first.
|
||||
# More layouts can use this path once their write-back kernels accept CPU
|
||||
# destination indices.
|
||||
if self.io_backend == "kernel" and self.mem_pool_host.layout == "page_first":
|
||||
# Page-first write-back JIT kernels can keep destination host indices on CPU.
|
||||
if (
|
||||
self.io_backend == "kernel"
|
||||
and self.mem_pool_host.layout == "page_first"
|
||||
and getattr(self.mem_pool_host, "can_use_write_back_jit", False)
|
||||
):
|
||||
host_indices, device_indices = op.host_indices, op.device_indices
|
||||
else:
|
||||
host_indices, device_indices = self.move_indices(
|
||||
|
||||
@@ -394,10 +394,12 @@ class HybridCacheController(BaseHiCacheController):
|
||||
if not self.write_queue:
|
||||
return
|
||||
op = CacheOperation.merge_ops(self.write_queue)
|
||||
# For now, kernel write-back keeps host indices on CPU only for page_first.
|
||||
# More layouts can use this path once their write-back kernels accept CPU
|
||||
# destination indices.
|
||||
if self.io_backend == "kernel" and self.mem_pool_host.layout == "page_first":
|
||||
# Page-first write-back JIT kernels can keep destination host indices on CPU.
|
||||
if (
|
||||
self.io_backend == "kernel"
|
||||
and self.mem_pool_host.layout == "page_first"
|
||||
and getattr(self.mem_pool_host, "can_use_write_back_jit", False)
|
||||
):
|
||||
host_indices = op.host_indices
|
||||
device_indices = op.device_indices
|
||||
resolved_pool_transfers = op.pool_transfers
|
||||
|
||||
@@ -321,7 +321,9 @@ def build_deepseek_v4_hicache_stack(
|
||||
swa_page_size=kvcache.swa_page_size,
|
||||
)
|
||||
|
||||
logical_host_pool = LogicalHostPool(num_host_pages * page_size, page_size)
|
||||
logical_host_pool = LogicalHostPool(
|
||||
num_host_pages * page_size, page_size, layout=server_args.hicache_mem_layout
|
||||
)
|
||||
swa_host_pool = DeepSeekV4PagedHostPool(
|
||||
pool_name=str(PoolName.SWA),
|
||||
device_buffers=kvcache.swa_kv_pool.kv_buffer,
|
||||
|
||||
@@ -17,6 +17,7 @@ import torch
|
||||
|
||||
from sglang.jit_kernel.hicache import (
|
||||
can_use_hicache_jit_kernel,
|
||||
can_use_write_back_jit_kernel,
|
||||
)
|
||||
from sglang.jit_kernel.hicache import (
|
||||
transfer_hicache_all_layer as jit_transfer_hicache_all_layer,
|
||||
@@ -257,6 +258,7 @@ class HostKVCache(abc.ABC):
|
||||
self.pin_memory = pin_memory
|
||||
self.device = device
|
||||
self.allocator = get_allocator_from_storage(allocator_type)
|
||||
self.can_use_write_back_jit = False
|
||||
|
||||
self.dtype = device_pool.store_dtype
|
||||
self.size_per_token = self.get_size_per_token()
|
||||
@@ -490,9 +492,16 @@ class MHATokenToKVPoolHost(HostKVCache):
|
||||
self.staging_token_capacity = 0
|
||||
self.staging_k_buffer = None
|
||||
self.staging_v_buffer = None
|
||||
self.can_use_write_back_jit = False
|
||||
if self.layout != "page_first" or (_is_npu or _is_xpu or _is_mps):
|
||||
return
|
||||
|
||||
self.can_use_write_back_jit = _is_cuda and can_use_write_back_jit_kernel(
|
||||
element_size=self.element_dim * self.dtype.itemsize,
|
||||
)
|
||||
if not self.can_use_write_back_jit:
|
||||
return
|
||||
|
||||
self.staging_page_capacity = min(self.page_num, _WRITE_BACK_STAGING_PAGE_CHUNK)
|
||||
self.staging_token_capacity = self.staging_page_capacity * self.page_size
|
||||
self.staging_k_buffer = torch.empty(
|
||||
@@ -661,7 +670,7 @@ class MHATokenToKVPoolHost(HostKVCache):
|
||||
num_layers=self.layer_num,
|
||||
)
|
||||
elif self.layout == "page_first":
|
||||
if self.can_use_jit:
|
||||
if self.can_use_write_back_jit:
|
||||
jit_transfer_hicache_all_layer_staged_lf_pf(
|
||||
k_ptr_src=device_pool.k_data_ptrs,
|
||||
v_ptr_src=device_pool.v_data_ptrs,
|
||||
@@ -1270,7 +1279,7 @@ class MLATokenToKVPoolHost(HiSparseHostPoolMixin, HostKVCache):
|
||||
element_size=self.kv_cache_dim * self.dtype.itemsize
|
||||
)
|
||||
|
||||
if self.layout == "page_first" and self.can_use_jit:
|
||||
if self.layout == "page_first":
|
||||
# Transpose [page, layer, ...] -> [layer, page, ...] to get per-layer views
|
||||
# This swaps strides without copying data
|
||||
transposed = self.kv_buffer.transpose(0, 1)
|
||||
@@ -1382,9 +1391,16 @@ class MLATokenToKVPoolHost(HiSparseHostPoolMixin, HostKVCache):
|
||||
self.staging_page_capacity = 0
|
||||
self.staging_token_capacity = 0
|
||||
self.staging_buffer = None
|
||||
self.can_use_write_back_jit = False
|
||||
if self.layout != "page_first" or (_is_npu or _is_xpu or _is_mps):
|
||||
return
|
||||
|
||||
self.can_use_write_back_jit = _is_cuda and can_use_write_back_jit_kernel(
|
||||
element_size=self.kv_cache_dim * self.dtype.itemsize,
|
||||
)
|
||||
if not self.can_use_write_back_jit:
|
||||
return
|
||||
|
||||
self.staging_page_capacity = min(self.page_num, _WRITE_BACK_STAGING_PAGE_CHUNK)
|
||||
self.staging_token_capacity = self.staging_page_capacity * self.page_size
|
||||
self.staging_buffer = torch.empty(
|
||||
@@ -1506,7 +1522,7 @@ class MLATokenToKVPoolHost(HiSparseHostPoolMixin, HostKVCache):
|
||||
num_layers=self.layer_num,
|
||||
)
|
||||
elif self.layout == "page_first":
|
||||
if self.can_use_jit:
|
||||
if self.can_use_write_back_jit:
|
||||
jit_transfer_hicache_all_layer_mla_staged_lf_pf(
|
||||
ptr_src=device_pool.data_ptrs,
|
||||
src_indices=device_indices,
|
||||
@@ -1747,7 +1763,25 @@ class MambaPoolHost(HostKVCache):
|
||||
self.layout,
|
||||
)
|
||||
|
||||
self.temporal_device_ptrs = torch.tensor(
|
||||
[
|
||||
device_pool.mamba_cache.temporal[i].data_ptr()
|
||||
for i in range(self.num_mamba_layers)
|
||||
],
|
||||
dtype=torch.uint64,
|
||||
device=self.device_pool.device,
|
||||
)
|
||||
self.conv_device_ptrs = [
|
||||
torch.tensor(
|
||||
[conv_state[i].data_ptr() for i in range(self.num_mamba_layers)],
|
||||
dtype=torch.uint64,
|
||||
device=self.device_pool.device,
|
||||
)
|
||||
for conv_state in device_pool.mamba_cache.conv
|
||||
]
|
||||
|
||||
self.init_kv_buffer()
|
||||
self._init_write_back_staging_buffers()
|
||||
self.lock = threading.RLock()
|
||||
self.clear()
|
||||
|
||||
@@ -1806,6 +1840,54 @@ class MambaPoolHost(HostKVCache):
|
||||
)
|
||||
)
|
||||
|
||||
def _init_write_back_staging_buffers(self):
|
||||
self.temporal_staging_buffer = None
|
||||
self.conv_staging_buffers = [None] * len(self.conv_buffer)
|
||||
self.can_use_write_back_jit = False
|
||||
self._temporal_can_use_jit = False
|
||||
self._conv_can_use_jit = [False] * len(self.conv_buffer)
|
||||
if self.layout != "page_first" or (_is_npu or _is_xpu or _is_mps):
|
||||
return
|
||||
|
||||
self._temporal_can_use_jit = _is_cuda and can_use_write_back_jit_kernel(
|
||||
element_size=self._item_size_per_index(self.temporal_buffer[0]),
|
||||
)
|
||||
self._conv_can_use_jit = [
|
||||
_is_cuda
|
||||
and can_use_write_back_jit_kernel(
|
||||
element_size=self._item_size_per_index(buf[0]),
|
||||
)
|
||||
for buf in self.conv_buffer
|
||||
]
|
||||
self.can_use_write_back_jit = self._temporal_can_use_jit and all(
|
||||
self._conv_can_use_jit
|
||||
)
|
||||
self.staging_page_capacity = min(self.page_num, _WRITE_BACK_STAGING_PAGE_CHUNK)
|
||||
self.staging_token_capacity = self.staging_page_capacity * self.page_size
|
||||
self.temporal_staging_buffer = torch.empty(
|
||||
(
|
||||
self.staging_token_capacity,
|
||||
self.num_mamba_layers,
|
||||
1,
|
||||
*self.temporal_state_shape,
|
||||
),
|
||||
dtype=self.temporal_dtype,
|
||||
device=self.device_pool.device,
|
||||
)
|
||||
self.conv_staging_buffers = [
|
||||
torch.empty(
|
||||
(
|
||||
self.staging_token_capacity,
|
||||
self.num_mamba_layers,
|
||||
1,
|
||||
*conv_shape,
|
||||
),
|
||||
dtype=self.conv_dtype,
|
||||
device=self.device_pool.device,
|
||||
)
|
||||
for conv_shape in self.conv_state_shapes
|
||||
]
|
||||
|
||||
def get_hybrid_pool_buffer(self):
|
||||
# Expose all mamba host tensors that need Mooncake buffer registration.
|
||||
return [self.temporal_buffer, *self.conv_buffer]
|
||||
@@ -1941,27 +2023,35 @@ class MambaPoolHost(HostKVCache):
|
||||
src_indices: torch.Tensor,
|
||||
dst_indices: torch.Tensor,
|
||||
num_layers: int,
|
||||
device: str,
|
||||
io_backend: str,
|
||||
src_ptrs: torch.Tensor,
|
||||
staging: Optional[torch.Tensor] = None,
|
||||
can_use_jit: bool = False,
|
||||
) -> None:
|
||||
if src_indices.numel() == 0:
|
||||
return
|
||||
if io_backend == "kernel":
|
||||
item_size = MambaPoolHost._item_size_per_index(src_layers[0])
|
||||
src_ptrs = torch.tensor(
|
||||
[src_layers[i].data_ptr() for i in range(num_layers)],
|
||||
dtype=torch.uint64,
|
||||
device=device,
|
||||
)
|
||||
transfer_kv_all_layer_mla_lf_pf(
|
||||
src_layers=src_ptrs,
|
||||
dst=dst,
|
||||
src_indices=src_indices,
|
||||
dst_indices=dst_indices,
|
||||
item_size=item_size,
|
||||
dst_layout_dim=item_size * num_layers,
|
||||
num_layers=num_layers,
|
||||
)
|
||||
if can_use_jit:
|
||||
jit_transfer_hicache_all_layer_mla_staged_lf_pf(
|
||||
ptr_src=src_ptrs,
|
||||
src_indices=src_indices,
|
||||
dst_indices=dst_indices,
|
||||
staging=staging,
|
||||
dst=dst,
|
||||
page_size=1,
|
||||
element_size=item_size,
|
||||
)
|
||||
else:
|
||||
transfer_kv_all_layer_mla_lf_pf(
|
||||
src_layers=src_ptrs,
|
||||
dst=dst,
|
||||
src_indices=src_indices,
|
||||
dst_indices=dst_indices,
|
||||
item_size=item_size,
|
||||
dst_layout_dim=item_size * num_layers,
|
||||
num_layers=num_layers,
|
||||
)
|
||||
elif io_backend == "direct":
|
||||
src_ptrs = [src_layers[i] for i in range(num_layers)]
|
||||
transfer_kv_all_layer_direct_lf_pf(
|
||||
@@ -2029,8 +2119,10 @@ class MambaPoolHost(HostKVCache):
|
||||
src_indices=device_indices,
|
||||
dst_indices=host_indices,
|
||||
num_layers=self.num_mamba_layers,
|
||||
device=self.device_pool.device,
|
||||
io_backend=io_backend,
|
||||
staging=self.temporal_staging_buffer,
|
||||
can_use_jit=self._temporal_can_use_jit,
|
||||
src_ptrs=self.temporal_device_ptrs,
|
||||
)
|
||||
for conv_idx in range(len(self.conv_state_shapes)):
|
||||
self._copy_tensor_all_layers_lf_pf(
|
||||
@@ -2039,8 +2131,10 @@ class MambaPoolHost(HostKVCache):
|
||||
src_indices=device_indices,
|
||||
dst_indices=host_indices,
|
||||
num_layers=self.num_mamba_layers,
|
||||
device=self.device_pool.device,
|
||||
io_backend=io_backend,
|
||||
staging=self.conv_staging_buffers[conv_idx],
|
||||
can_use_jit=self._conv_can_use_jit[conv_idx],
|
||||
src_ptrs=self.conv_device_ptrs[conv_idx],
|
||||
)
|
||||
else:
|
||||
for layer_id in range(self.num_mamba_layers):
|
||||
@@ -2161,7 +2255,7 @@ class LogicalHostPool:
|
||||
compressed side pools use these logical FULL indices as stable page anchors.
|
||||
"""
|
||||
|
||||
def __init__(self, size: int, page_size: int):
|
||||
def __init__(self, size: int, page_size: int, layout: str = "layer_first"):
|
||||
if size % page_size != 0:
|
||||
raise ValueError(
|
||||
"LogicalHostPool size must be page-aligned, "
|
||||
@@ -2170,7 +2264,7 @@ class LogicalHostPool:
|
||||
self.size = size
|
||||
self.page_size = page_size
|
||||
self.device = "cpu"
|
||||
self.layout = "layer_first"
|
||||
self.layout = layout
|
||||
self.dtype = torch.uint8
|
||||
self.layer_num = 0
|
||||
self.start_layer = 0
|
||||
@@ -2178,6 +2272,7 @@ class LogicalHostPool:
|
||||
self.kv_buffer = None
|
||||
self.size_per_token = 0
|
||||
self.allocator = None
|
||||
self.can_use_write_back_jit = True
|
||||
self.lock = threading.RLock()
|
||||
self.clear()
|
||||
|
||||
@@ -2342,8 +2437,26 @@ class DeepSeekV4PagedHostPool(HiSparseHostPoolMixin, HostKVCache):
|
||||
if self.data_refs
|
||||
else None
|
||||
)
|
||||
self.can_use_jit = False
|
||||
self.can_use_write_back_jit = False
|
||||
self._init_write_back_staging_buffers()
|
||||
self.clear()
|
||||
|
||||
def _init_write_back_staging_buffers(self):
|
||||
self.staging_buffer = None
|
||||
if self.layout != "page_first" or (_is_npu or _is_xpu or _is_mps):
|
||||
return
|
||||
|
||||
self.can_use_write_back_jit = _is_cuda and can_use_write_back_jit_kernel(
|
||||
element_size=self.item_bytes * self.dtype.itemsize,
|
||||
)
|
||||
staging_page_capacity = min(self.num_host_pages, _WRITE_BACK_STAGING_PAGE_CHUNK)
|
||||
self.staging_buffer = torch.empty(
|
||||
(staging_page_capacity, self.layer_num, self.item_bytes),
|
||||
dtype=self.dtype,
|
||||
device=self.gpu_device,
|
||||
)
|
||||
|
||||
def get_contiguous_buf_infos(self):
|
||||
"""Return per-layer page-row buffers for PD direct-to-host transfer."""
|
||||
data_ptrs = [int(self.data_ptrs[i].item()) for i in range(self.layer_num)]
|
||||
@@ -2434,15 +2547,26 @@ class DeepSeekV4PagedHostPool(HiSparseHostPoolMixin, HostKVCache):
|
||||
num_layers=self.layer_num,
|
||||
)
|
||||
elif io_backend == "kernel" and self.layout == "page_first":
|
||||
transfer_kv_all_layer_mla_lf_pf(
|
||||
src_layers=self.device_ptrs,
|
||||
dst=self.kv_buffer,
|
||||
src_indices=device_rows,
|
||||
dst_indices=host_rows,
|
||||
item_size=self.item_bytes,
|
||||
dst_layout_dim=self.layer_num * self.item_bytes,
|
||||
num_layers=self.layer_num,
|
||||
)
|
||||
if self.can_use_write_back_jit:
|
||||
jit_transfer_hicache_all_layer_mla_staged_lf_pf(
|
||||
ptr_src=self.device_ptrs,
|
||||
src_indices=device_rows,
|
||||
dst_indices=host_rows,
|
||||
staging=self.staging_buffer,
|
||||
dst=self.kv_buffer,
|
||||
page_size=1,
|
||||
element_size=self.item_bytes,
|
||||
)
|
||||
else:
|
||||
transfer_kv_all_layer_mla_lf_pf(
|
||||
src_layers=self.device_ptrs,
|
||||
dst=self.kv_buffer,
|
||||
src_indices=device_rows,
|
||||
dst_indices=host_rows,
|
||||
item_size=self.item_bytes,
|
||||
dst_layout_dim=self.layer_num * self.item_bytes,
|
||||
num_layers=self.layer_num,
|
||||
)
|
||||
elif io_backend == "direct" and self.layout == "layer_first":
|
||||
transfer_kv_direct(
|
||||
src_layers=self.device_buffers,
|
||||
@@ -2690,6 +2814,9 @@ class DeepSeekV4StateHostPool(HostKVCache):
|
||||
if self.data_refs
|
||||
else None
|
||||
)
|
||||
self.can_use_jit = False
|
||||
self.can_use_write_back_jit = False
|
||||
self._init_write_back_staging_buffers()
|
||||
|
||||
def _init_device_page_views(self) -> None:
|
||||
expected_ring_size = None
|
||||
@@ -2724,6 +2851,21 @@ class DeepSeekV4StateHostPool(HostKVCache):
|
||||
self.ring_size = expected_ring_size or 0
|
||||
self.state_page_bytes = expected_state_page_bytes or 0
|
||||
|
||||
def _init_write_back_staging_buffers(self):
|
||||
self.staging_buffer = None
|
||||
if self.layout != "page_first" or (_is_npu or _is_xpu or _is_mps):
|
||||
return
|
||||
|
||||
self.can_use_write_back_jit = _is_cuda and can_use_write_back_jit_kernel(
|
||||
element_size=self.state_page_bytes * self.dtype.itemsize,
|
||||
)
|
||||
staging_page_capacity = min(self.num_host_pages, _WRITE_BACK_STAGING_PAGE_CHUNK)
|
||||
self.staging_buffer = torch.empty(
|
||||
(staging_page_capacity, self.layer_num, self.state_page_bytes),
|
||||
dtype=self.dtype,
|
||||
device=self.gpu_device,
|
||||
)
|
||||
|
||||
def _to_page_indices(self, indices: torch.Tensor) -> torch.Tensor:
|
||||
if indices.numel() % self.swa_page_size != 0:
|
||||
raise ValueError(
|
||||
@@ -2782,15 +2924,26 @@ class DeepSeekV4StateHostPool(HostKVCache):
|
||||
num_layers=self.layer_num,
|
||||
)
|
||||
elif io_backend == "kernel" and self.layout == "page_first":
|
||||
transfer_kv_all_layer_mla_lf_pf(
|
||||
src_layers=self.device_ptrs,
|
||||
dst=self.kv_buffer,
|
||||
src_indices=device_rows,
|
||||
dst_indices=host_rows,
|
||||
item_size=self.state_page_bytes,
|
||||
dst_layout_dim=self.layer_num * self.state_page_bytes,
|
||||
num_layers=self.layer_num,
|
||||
)
|
||||
if self.can_use_write_back_jit:
|
||||
jit_transfer_hicache_all_layer_mla_staged_lf_pf(
|
||||
ptr_src=self.device_ptrs,
|
||||
src_indices=device_rows,
|
||||
dst_indices=host_rows,
|
||||
staging=self.staging_buffer,
|
||||
dst=self.kv_buffer,
|
||||
page_size=1,
|
||||
element_size=self.state_page_bytes,
|
||||
)
|
||||
else:
|
||||
transfer_kv_all_layer_mla_lf_pf(
|
||||
src_layers=self.device_ptrs,
|
||||
dst=self.kv_buffer,
|
||||
src_indices=device_rows,
|
||||
dst_indices=host_rows,
|
||||
item_size=self.state_page_bytes,
|
||||
dst_layout_dim=self.layer_num * self.state_page_bytes,
|
||||
num_layers=self.layer_num,
|
||||
)
|
||||
elif io_backend == "direct" and self.layout == "layer_first":
|
||||
transfer_kv_direct(
|
||||
src_layers=self.device_page_views,
|
||||
@@ -2960,6 +3113,10 @@ class HostPoolGroup:
|
||||
self.page_size = self.anchor_entry.host_pool.page_size
|
||||
self.device = self.anchor_entry.host_pool.device
|
||||
self.size = self.anchor_entry.host_pool.size
|
||||
self.can_use_write_back_jit = all(
|
||||
getattr(entry.host_pool, "can_use_write_back_jit", False)
|
||||
for entry in entries
|
||||
)
|
||||
|
||||
@property
|
||||
def kv_buffer(self):
|
||||
@@ -3141,6 +3298,9 @@ class DSAIndexerPoolHost(HostKVCache):
|
||||
layout,
|
||||
)
|
||||
self.init_kv_buffer()
|
||||
self.can_use_jit = False
|
||||
self.can_use_write_back_jit = False
|
||||
self._init_write_back_staging_buffers()
|
||||
self.lock = threading.RLock()
|
||||
self.clear()
|
||||
|
||||
@@ -3191,6 +3351,28 @@ class DSAIndexerPoolHost(HostKVCache):
|
||||
else:
|
||||
raise ValueError(f"Unsupported layout: {self.layout}")
|
||||
|
||||
def _init_write_back_staging_buffers(self):
|
||||
self.staging_buffer = None
|
||||
if self.layout != "page_first" or (_is_npu or _is_xpu or _is_mps):
|
||||
return
|
||||
|
||||
self.can_use_write_back_jit = _is_cuda and can_use_write_back_jit_kernel(
|
||||
element_size=self.indexer_page_stride_size * self.indexer_dtype.itemsize,
|
||||
)
|
||||
staging_page_capacity = min(
|
||||
self.indexer_page_num, _WRITE_BACK_STAGING_PAGE_CHUNK
|
||||
)
|
||||
self.staging_buffer = torch.empty(
|
||||
(
|
||||
staging_page_capacity,
|
||||
self.layer_num,
|
||||
1,
|
||||
self.indexer_page_stride_size,
|
||||
),
|
||||
dtype=self.indexer_dtype,
|
||||
device=self.device_pool.device,
|
||||
)
|
||||
|
||||
def get_hybrid_pool_buffer(self):
|
||||
return [self.index_k_with_scale_buffer]
|
||||
|
||||
@@ -3278,15 +3460,26 @@ class DSAIndexerPoolHost(HostKVCache):
|
||||
num_layers=self.layer_num,
|
||||
)
|
||||
elif self.layout == "page_first":
|
||||
transfer_kv_all_layer_mla_lf_pf(
|
||||
src_layers=self.index_k_device_ptrs,
|
||||
dst=self.index_k_with_scale_buffer,
|
||||
src_indices=device_page_indices,
|
||||
dst_indices=host_page_indices,
|
||||
item_size=self.indexer_page_stride_size,
|
||||
dst_layout_dim=self.indexer_layout_dim,
|
||||
num_layers=self.layer_num,
|
||||
)
|
||||
if self.can_use_write_back_jit:
|
||||
jit_transfer_hicache_all_layer_mla_staged_lf_pf(
|
||||
ptr_src=self.index_k_device_ptrs,
|
||||
src_indices=device_page_indices,
|
||||
dst_indices=host_page_indices,
|
||||
staging=self.staging_buffer,
|
||||
dst=self.index_k_with_scale_buffer,
|
||||
page_size=1,
|
||||
element_size=self.indexer_page_stride_size,
|
||||
)
|
||||
else:
|
||||
transfer_kv_all_layer_mla_lf_pf(
|
||||
src_layers=self.index_k_device_ptrs,
|
||||
dst=self.index_k_with_scale_buffer,
|
||||
src_indices=device_page_indices,
|
||||
dst_indices=host_page_indices,
|
||||
item_size=self.indexer_page_stride_size,
|
||||
dst_layout_dim=self.indexer_layout_dim,
|
||||
num_layers=self.layer_num,
|
||||
)
|
||||
else:
|
||||
raise ValueError(f"Unsupported layout: {self.layout}")
|
||||
elif io_backend == "direct":
|
||||
|
||||
Reference in New Issue
Block a user