[HiCache] support memory_pool_host page head layout (#11644)

This commit is contained in:
huangtingwei
2025-11-17 13:45:17 +08:00
committed by GitHub
parent 15bc1f5cd7
commit 1dcde53928
3 changed files with 50 additions and 2 deletions

View File

@@ -17,6 +17,7 @@ if not (_is_npu or _is_xpu):
transfer_kv_all_layer,
transfer_kv_all_layer_direct_lf_pf,
transfer_kv_all_layer_lf_pf,
transfer_kv_all_layer_lf_ph,
transfer_kv_all_layer_mla,
transfer_kv_all_layer_mla_lf_pf,
transfer_kv_direct,
@@ -25,6 +26,7 @@ if not (_is_npu or _is_xpu):
transfer_kv_per_layer_mla,
transfer_kv_per_layer_mla_pf_lf,
transfer_kv_per_layer_pf_lf,
transfer_kv_per_layer_ph_lf,
)
if _is_npu:
from sgl_kernel_npu.kvcacheio import TransferDirection, transfer_kv_dim_exchange
@@ -238,6 +240,15 @@ class MHATokenToKVPoolHost(HostKVCache):
self.head_num,
self.head_dim,
)
elif self.layout == "page_head":
dims = (
2,
self.page_num,
self.head_num,
self.page_size,
self.layer_num,
self.head_dim,
)
else:
raise ValueError(f"Unsupported layout: {self.layout}")
self.token_stride_size = self.head_num * self.head_dim * self.dtype.itemsize
@@ -292,6 +303,20 @@ class MHATokenToKVPoolHost(HostKVCache):
item_size=self.token_stride_size,
src_layout_dim=self.layout_dim,
)
elif self.layout == "page_head":
transfer_kv_per_layer_ph_lf(
src_k=self.k_buffer,
dst_k=device_pool.k_buffer[layer_id],
src_v=self.v_buffer,
dst_v=device_pool.v_buffer[layer_id],
src_indices=host_indices,
dst_indices=device_indices,
layer_id=layer_id,
item_size=self.token_stride_size,
src_layout_dim=self.layout_dim,
page_size=self.page_size,
head_num=self.head_num,
)
else:
raise ValueError(f"Unsupported layout: {self.layout}")
elif io_backend == "direct":
@@ -366,6 +391,20 @@ class MHATokenToKVPoolHost(HostKVCache):
dst_layout_dim=self.layout_dim,
num_layers=self.layer_num,
)
elif self.layout == "page_head":
transfer_kv_all_layer_lf_ph(
src_k_layers=device_pool.k_data_ptrs,
dst_k=self.k_buffer,
src_v_layers=device_pool.v_data_ptrs,
dst_v=self.v_buffer,
src_indices=device_indices,
dst_indices=host_indices,
item_size=self.token_stride_size,
dst_layout_dim=self.layout_dim,
num_layers=self.layer_num,
page_size=self.page_size,
head_num=self.head_num,
)
else:
raise ValueError(f"Unsupported layout: {self.layout}")
elif io_backend == "direct":
@@ -409,7 +448,7 @@ class MHATokenToKVPoolHost(HostKVCache):
data_page = self.kv_buffer[:, :, index : index + self.page_size, :, :]
elif self.layout == "page_first":
data_page = self.kv_buffer[:, index : index + self.page_size, :, :, :]
elif self.layout == "page_first_direct":
elif self.layout in ["page_first_direct", "page_head"]:
real_index = index // self.page_size
data_page = self.kv_buffer[:, real_index : real_index + 1, :, :, :, :]
else:
@@ -450,6 +489,13 @@ class MHATokenToKVPoolHost(HostKVCache):
2, 1, self.layer_num, self.page_size, self.head_num, self.head_dim
)
)
elif self.layout == "page_head":
real_index = index // self.page_size
self.kv_buffer[:, real_index : real_index + 1, :, :, :, :] = (
data_page.reshape(
2, 1, self.head_num, self.page_size, self.layer_num, self.head_dim
)
)
else:
raise ValueError(f"Unsupported layout: {self.layout}")
@@ -490,7 +536,7 @@ class MHATokenToKVPoolHost(HostKVCache):
self.dtype.itemsize * self.page_size * self.head_num * self.head_dim
)
element_size_list = [element_size] * len(ptr_list)
elif self.layout in ["page_first", "page_first_direct"]:
elif self.layout in ["page_first", "page_first_direct", "page_head"]:
for index in range(0, len(indices), self.page_size):
k_ptr = (
kv_buffer_data_ptr

View File

@@ -265,6 +265,7 @@ class MooncakeStore(HiCacheStorage):
assert self.mem_pool_host.layout in [
"page_first",
"page_first_direct",
"page_head",
], "mooncake store storage backend only support page first or page first direct layout"
buffer = self.mem_pool_host.kv_buffer
try:

View File

@@ -3074,6 +3074,7 @@ class ServerArgs:
"page_first",
"page_first_direct",
"page_first_kv_split",
"page_head",
],
default=ServerArgs.hicache_mem_layout,
help="The layout of host memory pool for hierarchical cache.",