[HiCache]Asymmetric pool support direct backend (#28446)

This commit is contained in:
huangtingwei
2026-06-16 13:17:57 -07:00
committed by GitHub
parent c0a6c3ce66
commit 9b4432fe18
6 changed files with 228 additions and 67 deletions
+82 -23
View File
@@ -936,9 +936,8 @@ class AsymmetricMHATokenToKVPoolHost(MHATokenToKVPoolHost):
``self.v_buffer``) instead of a single ``(2, ...)`` tensor, so each side
keeps its native stride. The kernel transfer path dispatches K and V as
independent single-buffer copies so each side uses its own ``item_size``.
Direct transfer and the flat-page L3 storage interface assume a single
shared ``item_size`` in paths that are not safe for asymmetric K/V, so they
raise instead of silently corrupting V copies.
K/V direct transfers must be dispatched separately because the direct
kernels derive copy sizes from each call's first tensor.
"""
def get_size_per_token(self):
@@ -960,10 +959,25 @@ class AsymmetricMHATokenToKVPoolHost(MHATokenToKVPoolHost):
if self.layout == "page_first":
k_dims = (self.size, self.layer_num, self.head_num, self.head_dim)
v_dims = (self.size, self.layer_num, self.head_num, self.v_head_dim)
elif self.layout == "page_first_direct":
k_dims = (
self.page_num,
self.layer_num,
self.page_size,
self.head_num,
self.head_dim,
)
v_dims = (
self.page_num,
self.layer_num,
self.page_size,
self.head_num,
self.v_head_dim,
)
else:
raise ValueError(
f"Unsupported layout for models with head_dim != v_head_dim: "
f"{self.layout}; expected 'page_first'."
f"{self.layout}; expected 'page_first' or 'page_first_direct'."
)
# token_stride_size / layout_dim are intentionally NOT set: K and V
@@ -1039,10 +1053,33 @@ class AsymmetricMHATokenToKVPoolHost(MHATokenToKVPoolHost):
item_size=self._v_token_stride_size(),
src_layout_dim=self._v_layout_dim(),
)
elif io_backend == "direct":
if self.layout != "page_first_direct":
raise ValueError(
f"Unsupported layout for models with head_dim != v_head_dim "
f"and io_backend='direct': {self.layout}; expected "
"'page_first_direct'."
)
transfer_kv_per_layer_direct_pf_lf(
src_ptrs=[self.k_buffer],
dst_ptrs=[device_pool.k_buffer[layer_id]],
src_indices=host_indices,
dst_indices=device_indices,
layer_id=layer_id,
page_size=self.page_size,
)
transfer_kv_per_layer_direct_pf_lf(
src_ptrs=[self.v_buffer],
dst_ptrs=[device_pool.v_buffer[layer_id]],
src_indices=host_indices,
dst_indices=device_indices,
layer_id=layer_id,
page_size=self.page_size,
)
else:
raise ValueError(
f"Unsupported IO backend for models with head_dim != v_head_dim: "
f"{io_backend}; expected 'kernel'."
f"{io_backend}; expected 'kernel' or 'direct'."
)
def backup_from_device_all_layer(
@@ -1072,10 +1109,31 @@ class AsymmetricMHATokenToKVPoolHost(MHATokenToKVPoolHost):
dst_layout_dim=self._v_layout_dim(),
num_layers=self.layer_num,
)
elif io_backend == "direct":
if self.layout != "page_first_direct":
raise ValueError(
f"Unsupported layout for models with head_dim != v_head_dim "
f"and io_backend='direct': {self.layout}; expected "
"'page_first_direct'."
)
transfer_kv_all_layer_direct_lf_pf(
src_ptrs=device_pool.k_buffer,
dst_ptrs=[self.k_buffer],
src_indices=device_indices,
dst_indices=host_indices,
page_size=self.page_size,
)
transfer_kv_all_layer_direct_lf_pf(
src_ptrs=device_pool.v_buffer,
dst_ptrs=[self.v_buffer],
src_indices=device_indices,
dst_indices=host_indices,
page_size=self.page_size,
)
else:
raise ValueError(
f"Unsupported IO backend for models with head_dim != v_head_dim: "
f"{io_backend}; expected 'kernel'."
f"{io_backend}; expected 'kernel' or 'direct'."
)
def get_data_page(self, index, flat: bool = True) -> torch.Tensor:
@@ -1097,7 +1155,7 @@ class AsymmetricMHATokenToKVPoolHost(MHATokenToKVPoolHost):
def get_page_buffer_meta(self, indices):
assert len(indices) % self.page_size == 0
if self.layout != "page_first":
if self.layout not in ("page_first", "page_first_direct"):
raise ValueError(
f"Unsupported layout for models with head_dim != v_head_dim: "
f"{self.layout}"
@@ -1121,29 +1179,30 @@ class AsymmetricMHATokenToKVPoolHost(MHATokenToKVPoolHost):
)
ptr_list = []
element_size_list = []
if self.layout == "page_first_direct":
k_index_stride = (
self.layer_num * self.page_size * self.head_num * self.head_dim
)
v_index_stride = (
self.layer_num * self.page_size * self.head_num * self.v_head_dim
)
else:
k_index_stride = self.layer_num * self.head_num * self.head_dim
v_index_stride = self.layer_num * self.head_num * self.v_head_dim
for index in range(0, len(indices), self.page_size):
k_ptr = (
k_base_ptr
+ indices[index]
* self.layer_num
* self.head_num
* self.head_dim
* self.dtype.itemsize
)
v_ptr = (
v_base_ptr
+ indices[index]
* self.layer_num
* self.head_num
* self.v_head_dim
* self.dtype.itemsize
buffer_index = (
indices[index] // self.page_size
if self.layout == "page_first_direct"
else indices[index]
)
k_ptr = k_base_ptr + buffer_index * k_index_stride * self.dtype.itemsize
v_ptr = v_base_ptr + buffer_index * v_index_stride * self.dtype.itemsize
ptr_list.extend([k_ptr, v_ptr])
element_size_list.extend([k_element_size, v_element_size])
return ptr_list, element_size_list
def is_stride_page_aligned(self, page_size_bytes: int = 4096) -> bool:
if self.layout != "page_first":
if self.layout not in ("page_first", "page_first_direct"):
return False
k_stride = (
self.page_size
+2 -15
View File
@@ -2440,21 +2440,8 @@ class ServerArgs:
)
# MiMoV2 has head_dim != v_head_dim, so the host KV pool uses
# asymmetric K/V allocation. Only the kernel/page_first transfer
# path has a safe split K/V implementation.
if self.hicache_io_backend != "kernel":
logger.warning(
f"Force hicache_io_backend to 'kernel' for MiMoV2 model "
f"(was {self.hicache_io_backend!r})."
)
self.hicache_io_backend = "kernel"
if self.hicache_mem_layout != "page_first":
logger.warning(
f"Force hicache_mem_layout to 'page_first' for "
f"MiMoV2 model (was {self.hicache_mem_layout!r}); "
f"asymmetric K/V HiCache requires kernel/page_first."
)
self.hicache_mem_layout = "page_first"
# asymmetric K/V allocation. Both kernel/page_first and
# direct/page_first_direct have split K/V transfer paths.
elif (
"Step3p5ForCausalLM" in model_arch
or "Step3p7ForConditionalGeneration" in model_arch