[HiCache] support memory_pool_host page head layout (#11644)
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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.",
|
||||
|
||||
Reference in New Issue
Block a user