[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,
|
||||||
transfer_kv_all_layer_direct_lf_pf,
|
transfer_kv_all_layer_direct_lf_pf,
|
||||||
transfer_kv_all_layer_lf_pf,
|
transfer_kv_all_layer_lf_pf,
|
||||||
|
transfer_kv_all_layer_lf_ph,
|
||||||
transfer_kv_all_layer_mla,
|
transfer_kv_all_layer_mla,
|
||||||
transfer_kv_all_layer_mla_lf_pf,
|
transfer_kv_all_layer_mla_lf_pf,
|
||||||
transfer_kv_direct,
|
transfer_kv_direct,
|
||||||
@@ -25,6 +26,7 @@ if not (_is_npu or _is_xpu):
|
|||||||
transfer_kv_per_layer_mla,
|
transfer_kv_per_layer_mla,
|
||||||
transfer_kv_per_layer_mla_pf_lf,
|
transfer_kv_per_layer_mla_pf_lf,
|
||||||
transfer_kv_per_layer_pf_lf,
|
transfer_kv_per_layer_pf_lf,
|
||||||
|
transfer_kv_per_layer_ph_lf,
|
||||||
)
|
)
|
||||||
if _is_npu:
|
if _is_npu:
|
||||||
from sgl_kernel_npu.kvcacheio import TransferDirection, transfer_kv_dim_exchange
|
from sgl_kernel_npu.kvcacheio import TransferDirection, transfer_kv_dim_exchange
|
||||||
@@ -238,6 +240,15 @@ class MHATokenToKVPoolHost(HostKVCache):
|
|||||||
self.head_num,
|
self.head_num,
|
||||||
self.head_dim,
|
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:
|
else:
|
||||||
raise ValueError(f"Unsupported layout: {self.layout}")
|
raise ValueError(f"Unsupported layout: {self.layout}")
|
||||||
self.token_stride_size = self.head_num * self.head_dim * self.dtype.itemsize
|
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,
|
item_size=self.token_stride_size,
|
||||||
src_layout_dim=self.layout_dim,
|
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:
|
else:
|
||||||
raise ValueError(f"Unsupported layout: {self.layout}")
|
raise ValueError(f"Unsupported layout: {self.layout}")
|
||||||
elif io_backend == "direct":
|
elif io_backend == "direct":
|
||||||
@@ -366,6 +391,20 @@ class MHATokenToKVPoolHost(HostKVCache):
|
|||||||
dst_layout_dim=self.layout_dim,
|
dst_layout_dim=self.layout_dim,
|
||||||
num_layers=self.layer_num,
|
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:
|
else:
|
||||||
raise ValueError(f"Unsupported layout: {self.layout}")
|
raise ValueError(f"Unsupported layout: {self.layout}")
|
||||||
elif io_backend == "direct":
|
elif io_backend == "direct":
|
||||||
@@ -409,7 +448,7 @@ class MHATokenToKVPoolHost(HostKVCache):
|
|||||||
data_page = self.kv_buffer[:, :, index : index + self.page_size, :, :]
|
data_page = self.kv_buffer[:, :, index : index + self.page_size, :, :]
|
||||||
elif self.layout == "page_first":
|
elif self.layout == "page_first":
|
||||||
data_page = self.kv_buffer[:, index : index + self.page_size, :, :, :]
|
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
|
real_index = index // self.page_size
|
||||||
data_page = self.kv_buffer[:, real_index : real_index + 1, :, :, :, :]
|
data_page = self.kv_buffer[:, real_index : real_index + 1, :, :, :, :]
|
||||||
else:
|
else:
|
||||||
@@ -450,6 +489,13 @@ class MHATokenToKVPoolHost(HostKVCache):
|
|||||||
2, 1, self.layer_num, self.page_size, self.head_num, self.head_dim
|
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:
|
else:
|
||||||
raise ValueError(f"Unsupported layout: {self.layout}")
|
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
|
self.dtype.itemsize * self.page_size * self.head_num * self.head_dim
|
||||||
)
|
)
|
||||||
element_size_list = [element_size] * len(ptr_list)
|
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):
|
for index in range(0, len(indices), self.page_size):
|
||||||
k_ptr = (
|
k_ptr = (
|
||||||
kv_buffer_data_ptr
|
kv_buffer_data_ptr
|
||||||
|
|||||||
@@ -265,6 +265,7 @@ class MooncakeStore(HiCacheStorage):
|
|||||||
assert self.mem_pool_host.layout in [
|
assert self.mem_pool_host.layout in [
|
||||||
"page_first",
|
"page_first",
|
||||||
"page_first_direct",
|
"page_first_direct",
|
||||||
|
"page_head",
|
||||||
], "mooncake store storage backend only support page first or page first direct layout"
|
], "mooncake store storage backend only support page first or page first direct layout"
|
||||||
buffer = self.mem_pool_host.kv_buffer
|
buffer = self.mem_pool_host.kv_buffer
|
||||||
try:
|
try:
|
||||||
|
|||||||
@@ -3074,6 +3074,7 @@ class ServerArgs:
|
|||||||
"page_first",
|
"page_first",
|
||||||
"page_first_direct",
|
"page_first_direct",
|
||||||
"page_first_kv_split",
|
"page_first_kv_split",
|
||||||
|
"page_head",
|
||||||
],
|
],
|
||||||
default=ServerArgs.hicache_mem_layout,
|
default=ServerArgs.hicache_mem_layout,
|
||||||
help="The layout of host memory pool for hierarchical cache.",
|
help="The layout of host memory pool for hierarchical cache.",
|
||||||
|
|||||||
Reference in New Issue
Block a user