[Hicache & JIT_kernel] Support page first layout & mla jit kernel (#18311)

This commit is contained in:
huangtingwei
2026-03-27 08:54:36 -07:00
committed by GitHub
parent 30397e0a1e
commit d864622a68
4 changed files with 615 additions and 71 deletions
+153 -13
View File
@@ -14,6 +14,11 @@ namespace device {
namespace details {
template <typename T, uint32_t N>
struct LocalStorage {
T data[N];
};
template <int kUnit>
inline constexpr auto get_mem_package() {
if constexpr (kUnit == 16) {
@@ -78,7 +83,7 @@ SGL_DEVICE auto load_vec(const void* __restrict__ src) {
static_assert(128 % kNumThreads == 0, "kNumThreads must divide 128 bytes");
constexpr uint32_t kLoopCount = kBytes / 128;
using Package = details::PackageType<128 / kNumThreads>;
using Storage = AlignedStorage<Package, kLoopCount>;
using Storage = details::LocalStorage<Package, kLoopCount>;
const auto src_packed = static_cast<const Package*>(src);
const auto lane_id = threadIdx.x % kNumThreads;
@@ -129,7 +134,13 @@ struct HicacheKernelParams {
uint32_t num_layers = 0; // only used in all_layer transfer
};
template <typename T, int64_t kElementSize, uint32_t kUnroll, uint32_t kBlockQuota, uint32_t kBlockSize>
template <
typename T,
int64_t kElementSize,
uint32_t kUnroll,
uint32_t kBlockQuota,
uint32_t kBlockSize,
bool kIsMLA = false>
SGL_HICACHE_KERNEL void hicache_transfer_per_layer(const __grid_constant__ HicacheKernelParams params) {
using namespace device;
static_assert(kBlockSize % kWarpThreads == 0);
@@ -151,16 +162,24 @@ SGL_HICACHE_KERNEL void hicache_transfer_per_layer(const __grid_constant__ Hicac
const auto pos_dst = static_cast<const T*>(indices_dst)[i];
const auto src_k = pointer::offset(k_cache_src, pos_src * kv_cache_src_stride);
const auto dst_k = pointer::offset(k_cache_dst, pos_dst * kv_cache_dst_stride);
const auto src_v = pointer::offset(v_cache_src, pos_src * kv_cache_src_stride);
const auto dst_v = pointer::offset(v_cache_dst, pos_dst * kv_cache_dst_stride);
const auto vec_k = load_vec<kElementSize, kNumThreads>(src_k);
const auto vec_v = load_vec<kElementSize, kNumThreads>(src_v);
store_vec<kElementSize, kNumThreads>(dst_k, vec_k);
store_vec<kElementSize, kNumThreads>(dst_v, vec_v);
if constexpr (!kIsMLA) {
const auto src_v = pointer::offset(v_cache_src, pos_src * kv_cache_src_stride);
const auto dst_v = pointer::offset(v_cache_dst, pos_dst * kv_cache_dst_stride);
const auto vec_v = load_vec<kElementSize, kNumThreads>(src_v);
store_vec<kElementSize, kNumThreads>(dst_v, vec_v);
}
}
}
template <typename T, int64_t kElementSize, uint32_t kUnroll, uint32_t kBlockQuota, uint32_t kBlockSize>
template <
typename T,
int64_t kElementSize,
uint32_t kUnroll,
uint32_t kBlockQuota,
uint32_t kBlockSize,
bool kIsMLA = false>
SGL_HICACHE_KERNEL void hicache_transfer_all_layer(const __grid_constant__ HicacheKernelParams params) {
using namespace device;
using src_ptr_t = const void*;
@@ -185,17 +204,19 @@ SGL_HICACHE_KERNEL void hicache_transfer_all_layer(const __grid_constant__ Hicac
const auto pos_dst = static_cast<const T*>(indices_dst)[i];
for (uint32_t layer = 0; layer < num_layers; ++layer) {
const auto k_cache_src = static_cast<const src_ptr_t*>(k_ptr_src)[layer];
const auto v_cache_src = static_cast<const src_ptr_t*>(v_ptr_src)[layer];
const auto k_cache_dst = static_cast<const dst_ptr_t*>(k_ptr_dst)[layer];
const auto v_cache_dst = static_cast<const dst_ptr_t*>(v_ptr_dst)[layer];
const auto src_k = pointer::offset(k_cache_src, pos_src * kv_cache_src_stride);
const auto dst_k = pointer::offset(k_cache_dst, pos_dst * kv_cache_dst_stride);
const auto src_v = pointer::offset(v_cache_src, pos_src * kv_cache_src_stride);
const auto dst_v = pointer::offset(v_cache_dst, pos_dst * kv_cache_dst_stride);
const auto vec_k = load_vec<kElementSize, kNumThreads>(src_k);
const auto vec_v = load_vec<kElementSize, kNumThreads>(src_v);
store_vec<kElementSize, kNumThreads>(dst_k, vec_k);
store_vec<kElementSize, kNumThreads>(dst_v, vec_v);
if constexpr (!kIsMLA) {
const auto v_cache_src = static_cast<const src_ptr_t*>(v_ptr_src)[layer];
const auto v_cache_dst = static_cast<const dst_ptr_t*>(v_ptr_dst)[layer];
const auto src_v = pointer::offset(v_cache_src, pos_src * kv_cache_src_stride);
const auto dst_v = pointer::offset(v_cache_dst, pos_dst * kv_cache_dst_stride);
const auto vec_v = load_vec<kElementSize, kNumThreads>(src_v);
store_vec<kElementSize, kNumThreads>(dst_v, vec_v);
}
}
}
}
@@ -206,6 +227,12 @@ struct HiCacheKernel {
static constexpr auto kernel_one = hicache_transfer_per_layer<T, kElementSize, kUnroll, kBlockQuota, kBlockSize>;
template <typename T>
static constexpr auto kernel_all = hicache_transfer_all_layer<T, kElementSize, kUnroll, kBlockQuota, kBlockSize>;
template <typename T>
static constexpr auto kernel_one_mla =
hicache_transfer_per_layer<T, kElementSize, kUnroll, kBlockQuota, kBlockSize, true>;
template <typename T>
static constexpr auto kernel_all_mla =
hicache_transfer_all_layer<T, kElementSize, kUnroll, kBlockQuota, kBlockSize, true>;
static void run_one(
const tvm::ffi::TensorView k_cache_dst,
@@ -333,6 +360,119 @@ struct HiCacheKernel {
const auto kernel = use_int32 ? kernel_all<int32_t> : kernel_all<int64_t>;
LaunchKernel(num_blocks, kBlockSize, device)(kernel, params);
}
static void run_one_mla(
const tvm::ffi::TensorView cache_dst,
const tvm::ffi::TensorView indices_dst,
const tvm::ffi::TensorView cache_src,
const tvm::ffi::TensorView indices_src) {
using namespace host;
auto D = SymbolicSize{"head dimension"};
auto N = SymbolicSize{"src stride"};
auto M = SymbolicSize{"dst stride"};
auto L = SymbolicSize{"indices length"};
auto cache_dtype = SymbolicDType{};
auto indices_dtype = SymbolicDType{};
auto indices_device = SymbolicDevice{};
TensorMatcher({-1, D}) //
.with_strides({N, 1})
.with_dtype(cache_dtype)
.with_device<kDLCUDA, kDLCUDAHost, kDLCPU>()
.verify(cache_src);
TensorMatcher({-1, D}) //
.with_strides({M, 1})
.with_dtype(cache_dtype)
.with_device<kDLCUDA, kDLCUDAHost, kDLCPU>()
.verify(cache_dst);
TensorMatcher({L}) //
.with_dtype<int32_t, int64_t>(indices_dtype)
.with_device<kDLCUDA>(indices_device)
.verify(indices_src)
.verify(indices_dst);
const auto dtype_size = dtype_bytes(cache_dtype.unwrap());
const auto element_bytes = D.unwrap() * dtype_size;
RuntimeCheck(kElementSize == element_bytes, "HicacheKernel MLA: cache dimension mismatch.");
const auto cache_dst_ptr = cache_dst.data_ptr();
const auto cache_src_ptr = cache_src.data_ptr();
const auto indices_dst_ptr = indices_dst.data_ptr();
const auto indices_src_ptr = indices_src.data_ptr();
const auto length = static_cast<uint32_t>(L.unwrap());
const auto cache_src_stride = static_cast<int64_t>(N.unwrap() * dtype_size);
const auto cache_dst_stride = static_cast<int64_t>(M.unwrap() * dtype_size);
const auto use_int32 = indices_dtype.unwrap().bits == 32;
const auto device = indices_device.unwrap();
constexpr auto kWorkersPerBlock = kBlockSize / (device::kWarpThreads / kUnroll);
const auto num_blocks = std::min(div_ceil(length, kWorkersPerBlock), kBlockQuota);
const auto params = HicacheKernelParams{
.k_cache_dst = cache_dst_ptr,
.v_cache_dst = nullptr,
.indices_dst = indices_dst_ptr,
.k_cache_src = cache_src_ptr,
.v_cache_src = nullptr,
.indices_src = indices_src_ptr,
.kv_cache_src_stride = cache_src_stride,
.kv_cache_dst_stride = cache_dst_stride,
.length = length,
};
const auto kernel = use_int32 ? kernel_one_mla<int32_t> : kernel_one_mla<int64_t>;
LaunchKernel(num_blocks, kBlockSize, device)(kernel, params);
}
static void run_all_mla(
const tvm::ffi::TensorView ptr_dst,
const tvm::ffi::TensorView indices_dst,
const tvm::ffi::TensorView ptr_src,
const tvm::ffi::TensorView indices_src,
const int64_t src_stride_bytes,
const int64_t dst_stride_bytes) {
using namespace host;
auto N = SymbolicSize{"num_layers"};
auto L = SymbolicSize{"indices length"};
auto dtype_ = SymbolicDType{};
auto device_ = SymbolicDevice{};
TensorMatcher({N}) //
.with_dtype<uint64_t>()
.with_device<kDLCUDA>(device_)
.verify(ptr_src)
.verify(ptr_dst);
TensorMatcher({L}) //
.with_dtype<int32_t, int64_t>(dtype_)
.with_device<kDLCUDA>(device_)
.verify(indices_src)
.verify(indices_dst);
const auto cache_dst_ptr = ptr_dst.data_ptr();
const auto cache_src_ptr = ptr_src.data_ptr();
const auto indices_dst_ptr = indices_dst.data_ptr();
const auto indices_src_ptr = indices_src.data_ptr();
const auto length = static_cast<uint32_t>(L.unwrap());
const auto use_int32 = dtype_.unwrap().bits == 32;
const auto device = device_.unwrap();
constexpr auto kWorkersPerBlock = kBlockSize / (device::kWarpThreads / kUnroll);
const auto num_blocks = std::min(div_ceil(length, kWorkersPerBlock), kBlockQuota);
const auto params = HicacheKernelParams{
.k_cache_dst = cache_dst_ptr,
.v_cache_dst = nullptr,
.indices_dst = indices_dst_ptr,
.k_cache_src = cache_src_ptr,
.v_cache_src = nullptr,
.indices_src = indices_src_ptr,
.kv_cache_src_stride = src_stride_bytes,
.kv_cache_dst_stride = dst_stride_bytes,
.length = length,
.num_layers = static_cast<uint32_t>(N.unwrap()),
};
const auto kernel = use_int32 ? kernel_all_mla<int32_t> : kernel_all_mla<int64_t>;
LaunchKernel(num_blocks, kBlockSize, device)(kernel, params);
}
};
#undef SGL_HICACHE_KERNEL
+64
View File
@@ -28,6 +28,8 @@ def _jit_hicache_module(*, element_size: int, unroll: int, block_quota: int) ->
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"),
],
)
@@ -139,3 +141,65 @@ def transfer_hicache_all_layer(
kv_cache_src_stride_bytes,
kv_cache_dst_stride_bytes,
)
def transfer_hicache_one_layer_mla(
cache_dst: torch.Tensor,
indices_dst: torch.Tensor,
cache_src: torch.Tensor,
indices_src: torch.Tensor,
*,
element_dim: int | None = None,
unroll: int | None = None,
block_quota: int | None = None,
) -> None:
element_dim = element_dim or cache_dst.size(-1)
cache_src = cache_src.view(-1, element_dim)
cache_dst = cache_dst.view(-1, element_dim)
element_size = element_dim * cache_dst.element_size()
block_quota = block_quota or DEFAULT_BLOCK_QUOTA
unroll = unroll or _default_unroll(element_size)
module = _jit_hicache_module(
element_size=element_size,
unroll=unroll,
block_quota=block_quota,
)
module.launch_one_mla(
cache_dst,
indices_dst,
cache_src,
indices_src,
)
def transfer_hicache_all_layer_mla(
ptr_dst: torch.Tensor,
indices_dst: torch.Tensor,
ptr_src: torch.Tensor,
indices_src: torch.Tensor,
*,
cache_src_stride_bytes: int,
cache_dst_stride_bytes: int,
element_size: int | None = None,
unroll: int | None = None,
block_quota: int | None = None,
) -> None:
if element_size is None:
assert cache_dst_stride_bytes == cache_src_stride_bytes
element_size = cache_dst_stride_bytes
block_quota = block_quota or DEFAULT_BLOCK_QUOTA
unroll = unroll or _default_unroll(element_size)
module = _jit_hicache_module(
element_size=element_size,
unroll=unroll,
block_quota=block_quota,
)
module.launch_all_mla(
ptr_dst,
indices_dst,
ptr_src,
indices_src,
cache_src_stride_bytes,
cache_dst_stride_bytes,
)
@@ -0,0 +1,247 @@
import sys
import pytest
import torch
from sglang.srt.mem_cache.memory_pool import MHATokenToKVPool, MLATokenToKVPool
from sglang.srt.mem_cache.memory_pool_host import (
ALLOC_MEMORY_FUNCS,
MHATokenToKVPoolHost,
MLATokenToKVPoolHost,
alloc_with_pin_memory,
)
from sglang.srt.utils import is_cuda, is_hip, is_npu, is_xpu
from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci(est_time=10, suite="stage-b-kernel-unit-1-gpu-large")
register_cuda_ci(est_time=120, suite="nightly-kernel-1-gpu", nightly=True)
pytestmark = pytest.mark.skipif(
not torch.cuda.is_available()
or is_npu()
or is_xpu()
or not (is_cuda() or is_hip()),
reason="HiCache JIT tests require CUDA/ROCm.",
)
DEVICE = "cuda"
PAGE_SIZE = 1 if is_hip() else 16
NUM_LAYERS = 2
POOL_SIZE = PAGE_SIZE * 8
MHA_ELEMENT_DIMS = [128, 256, 512, 1024]
MLA_ELEMENT_DIMS = [576]
LAYOUTS = ["layer_first", "page_first"]
def _token_indices_for_pages(
pages: torch.Tensor, page_size: int = PAGE_SIZE, device: str = DEVICE
) -> torch.Tensor:
parts = [
torch.arange(
int(page) * page_size,
(int(page) + 1) * page_size,
device=device,
dtype=torch.int64,
)
for page in pages.tolist()
]
return torch.cat(parts, dim=0)
def _pinned_host_pool(host_pool_cls, **kwargs):
original_alloc = ALLOC_MEMORY_FUNCS[DEVICE]
ALLOC_MEMORY_FUNCS[DEVICE] = alloc_with_pin_memory
try:
return host_pool_cls(
host_to_device_ratio=2.0,
host_size=0,
page_size=PAGE_SIZE,
pin_memory=True,
device="cpu",
**kwargs,
)
finally:
ALLOC_MEMORY_FUNCS[DEVICE] = original_alloc
def _copy_tensor_with_offset(tensor: torch.Tensor, offset: int) -> None:
data = torch.arange(
tensor.numel(), device=tensor.device, dtype=tensor.dtype
).view_as(tensor)
tensor.copy_(data + offset)
def _run_transfer_roundtrip_mha(layout: str, element_dim: int) -> None:
device_pool = MHATokenToKVPool(
size=POOL_SIZE,
page_size=PAGE_SIZE,
head_num=element_dim // 128,
head_dim=128,
dtype=torch.bfloat16,
layer_num=NUM_LAYERS,
device=DEVICE,
enable_memory_saver=False,
)
host_pool = _pinned_host_pool(
MHATokenToKVPoolHost,
device_pool=device_pool,
layout=layout,
)
assert (
host_pool.can_use_jit
), f"Expected JIT HiCache kernel for MHA dim={element_dim}"
for layer_id in range(NUM_LAYERS):
_copy_tensor_with_offset(device_pool.k_buffer[layer_id], layer_id)
_copy_tensor_with_offset(device_pool.v_buffer[layer_id], layer_id + 100)
device_pages = torch.tensor([1, 2, 3], device=DEVICE, dtype=torch.int64)
host_pages = torch.tensor([0, 1, 2], device=DEVICE, dtype=torch.int64)
device_indices = _token_indices_for_pages(device_pages)
host_indices = _token_indices_for_pages(host_pages)
host_pool.backup_from_device_all_layer(
device_pool, host_indices, device_indices, "kernel"
)
torch.cuda.synchronize()
for layer_id in range(NUM_LAYERS):
for host_page, device_page in zip(host_pages.tolist(), device_pages.tolist()):
host_start = host_page * PAGE_SIZE
device_start = device_page * PAGE_SIZE
assert torch.equal(
host_pool.k_data_refs[layer_id][
host_start : host_start + PAGE_SIZE
].cpu(),
device_pool.k_buffer[layer_id][
device_start : device_start + PAGE_SIZE
].cpu(),
)
assert torch.equal(
host_pool.v_data_refs[layer_id][
host_start : host_start + PAGE_SIZE
].cpu(),
device_pool.v_buffer[layer_id][
device_start : device_start + PAGE_SIZE
].cpu(),
)
for layer_id in range(NUM_LAYERS):
device_pool.k_buffer[layer_id].zero_()
device_pool.v_buffer[layer_id].zero_()
load_pages = torch.tensor([4, 5, 6], device=DEVICE, dtype=torch.int64)
load_indices = _token_indices_for_pages(load_pages)
for layer_id in range(NUM_LAYERS):
host_pool.load_to_device_per_layer(
device_pool, host_indices, load_indices, layer_id, "kernel"
)
torch.cuda.synchronize()
for layer_id in range(NUM_LAYERS):
for host_page, device_page in zip(host_pages.tolist(), load_pages.tolist()):
host_start = host_page * PAGE_SIZE
device_start = device_page * PAGE_SIZE
assert torch.equal(
device_pool.k_buffer[layer_id][
device_start : device_start + PAGE_SIZE
].cpu(),
host_pool.k_data_refs[layer_id][
host_start : host_start + PAGE_SIZE
].cpu(),
)
assert torch.equal(
device_pool.v_buffer[layer_id][
device_start : device_start + PAGE_SIZE
].cpu(),
host_pool.v_data_refs[layer_id][
host_start : host_start + PAGE_SIZE
].cpu(),
)
def _run_transfer_roundtrip_mla(layout: str, element_dim: int) -> None:
device_pool = MLATokenToKVPool(
size=POOL_SIZE,
page_size=PAGE_SIZE,
kv_lora_rank=element_dim - 64,
qk_rope_head_dim=64,
dtype=torch.bfloat16,
layer_num=NUM_LAYERS,
device=DEVICE,
enable_memory_saver=False,
)
host_pool = _pinned_host_pool(
MLATokenToKVPoolHost,
device_pool=device_pool,
layout=layout,
)
assert (
host_pool.can_use_jit
), f"Expected JIT HiCache kernel for MLA dim={element_dim}"
for layer_id in range(NUM_LAYERS):
_copy_tensor_with_offset(device_pool.kv_buffer[layer_id], layer_id)
device_pages = torch.tensor([1, 2, 3], device=DEVICE, dtype=torch.int64)
host_pages = torch.tensor([0, 1, 2], device=DEVICE, dtype=torch.int64)
device_indices = _token_indices_for_pages(device_pages)
host_indices = _token_indices_for_pages(host_pages)
host_pool.backup_from_device_all_layer(
device_pool, host_indices, device_indices, "kernel"
)
torch.cuda.synchronize()
for layer_id in range(NUM_LAYERS):
for host_page, device_page in zip(host_pages.tolist(), device_pages.tolist()):
host_start = host_page * PAGE_SIZE
device_start = device_page * PAGE_SIZE
assert torch.equal(
host_pool.data_refs[layer_id][
host_start : host_start + PAGE_SIZE
].cpu(),
device_pool.kv_buffer[layer_id][
device_start : device_start + PAGE_SIZE
].cpu(),
)
for layer_id in range(NUM_LAYERS):
device_pool.kv_buffer[layer_id].zero_()
load_pages = torch.tensor([4, 5, 6], device=DEVICE, dtype=torch.int64)
load_indices = _token_indices_for_pages(load_pages)
for layer_id in range(NUM_LAYERS):
host_pool.load_to_device_per_layer(
device_pool, host_indices, load_indices, layer_id, "kernel"
)
torch.cuda.synchronize()
for layer_id in range(NUM_LAYERS):
for host_page, device_page in zip(host_pages.tolist(), load_pages.tolist()):
host_start = host_page * PAGE_SIZE
device_start = device_page * PAGE_SIZE
assert torch.equal(
device_pool.kv_buffer[layer_id][
device_start : device_start + PAGE_SIZE
].cpu(),
host_pool.data_refs[layer_id][
host_start : host_start + PAGE_SIZE
].cpu(),
)
@pytest.mark.parametrize("layout", LAYOUTS)
@pytest.mark.parametrize("element_dim", MHA_ELEMENT_DIMS)
def test_hicache_transfer_mha(layout: str, element_dim: int) -> None:
_run_transfer_roundtrip_mha(layout, element_dim)
@pytest.mark.parametrize("layout", LAYOUTS)
@pytest.mark.parametrize("element_dim", MLA_ELEMENT_DIMS)
def test_hicache_transfer_mla(layout: str, element_dim: int) -> None:
_run_transfer_roundtrip_mla(layout, element_dim)
if __name__ == "__main__":
sys.exit(pytest.main([__file__, "-v", "-s"]))
+151 -58
View File
@@ -21,9 +21,15 @@ from sglang.jit_kernel.hicache import (
from sglang.jit_kernel.hicache import (
transfer_hicache_all_layer as jit_transfer_hicache_all_layer,
)
from sglang.jit_kernel.hicache import (
transfer_hicache_all_layer_mla as jit_transfer_hicache_all_layer_mla,
)
from sglang.jit_kernel.hicache import (
transfer_hicache_one_layer as jit_transfer_hicache_one_layer,
)
from sglang.jit_kernel.hicache import (
transfer_hicache_one_layer_mla as jit_transfer_hicache_one_layer_mla,
)
from sglang.srt.mem_cache.memory_pool import (
KVCache,
MambaPool,
@@ -309,8 +315,16 @@ class MHATokenToKVPoolHost(HostKVCache):
element_size=self.element_dim * self.dtype.itemsize
)
self.k_data_refs = [self.k_buffer[i] for i in range(self.layer_num)]
self.v_data_refs = [self.v_buffer[i] for i in range(self.layer_num)]
if self.layout == "page_first":
# Transpose [page, layer, ...] -> [layer, page, ...] to get per-layer views
# This swaps strides without copying data
k_transposed = self.k_buffer.transpose(0, 1)
v_transposed = self.v_buffer.transpose(0, 1)
self.k_data_refs = [k_transposed[i] for i in range(self.layer_num)]
self.v_data_refs = [v_transposed[i] for i in range(self.layer_num)]
else:
self.k_data_refs = [self.k_buffer[i] for i in range(self.layer_num)]
self.v_data_refs = [self.v_buffer[i] for i in range(self.layer_num)]
self.k_data_ptrs = torch.tensor(
[x.data_ptr() for x in self.k_data_refs],
dtype=torch.uint64,
@@ -409,17 +423,31 @@ class MHATokenToKVPoolHost(HostKVCache):
item_size=self.token_stride_size,
)
elif self.layout == "page_first":
transfer_kv_per_layer_pf_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,
)
if self.can_use_jit:
# Transpose [page, layer, ...] -> [layer, page, ...] then
# index by layer_id to get a per-layer view with strided layout.
# The kernel handles different src/dst strides automatically.
jit_transfer_hicache_one_layer(
k_cache_dst=device_pool.k_buffer[layer_id],
v_cache_dst=device_pool.v_buffer[layer_id],
k_cache_src=self.k_data_refs[layer_id],
v_cache_src=self.v_data_refs[layer_id],
indices_dst=device_indices,
indices_src=host_indices,
element_dim=self.element_dim,
)
else:
transfer_kv_per_layer_pf_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,
)
elif self.layout == "page_head":
transfer_kv_per_layer_ph_lf(
src_k=self.k_buffer,
@@ -510,17 +538,32 @@ class MHATokenToKVPoolHost(HostKVCache):
num_layers=self.layer_num,
)
elif self.layout == "page_first":
transfer_kv_all_layer_lf_pf(
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,
)
if self.can_use_jit:
# Use transposed data ptrs so the kernel writes to
# [layer, page, item] view with stride layout_dim per token.
jit_transfer_hicache_all_layer(
k_ptr_dst=self.k_data_ptrs,
v_ptr_dst=self.v_data_ptrs,
indices_dst=host_indices,
k_ptr_src=device_pool.k_data_ptrs,
v_ptr_src=device_pool.v_data_ptrs,
indices_src=device_indices,
kv_cache_src_stride_bytes=self.token_stride_size,
kv_cache_dst_stride_bytes=self.layout_dim,
element_size=self.element_dim * self.dtype.itemsize,
)
else:
transfer_kv_all_layer_lf_pf(
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,
)
elif self.layout == "page_head":
transfer_kv_all_layer_lf_ph(
src_k_layers=device_pool.k_data_ptrs,
@@ -766,7 +809,17 @@ class MLATokenToKVPoolHost(HostKVCache):
device,
allocator_type,
)
self.data_refs = [self.kv_buffer[i] for i in range(self.layer_num)]
self.can_use_jit = _is_cuda and can_use_hicache_jit_kernel(
element_size=self.kv_cache_dim * self.dtype.itemsize
)
if self.layout == "page_first" and self.can_use_jit:
# Transpose [page, layer, ...] -> [layer, page, ...] to get per-layer views
# This swaps strides without copying data
transposed = self.kv_buffer.transpose(0, 1)
self.data_refs = [transposed[i] for i in range(self.layer_num)]
else:
self.data_refs = [self.kv_buffer[i] for i in range(self.layer_num)]
self.data_ptrs = torch.tensor(
[x.data_ptr() for x in self.data_refs],
dtype=torch.uint64,
@@ -864,23 +917,41 @@ class MLATokenToKVPoolHost(HostKVCache):
):
if io_backend == "kernel":
if self.layout == "layer_first":
transfer_kv_per_layer_mla(
src=self.kv_buffer[layer_id],
dst=device_pool.kv_buffer[layer_id],
src_indices=host_indices,
dst_indices=device_indices,
item_size=self.token_stride_size,
)
if self.can_use_jit:
jit_transfer_hicache_one_layer_mla(
cache_dst=device_pool.kv_buffer[layer_id],
cache_src=self.kv_buffer[layer_id],
indices_dst=device_indices,
indices_src=host_indices,
element_dim=self.kv_cache_dim,
)
else:
transfer_kv_per_layer_mla(
src=self.kv_buffer[layer_id],
dst=device_pool.kv_buffer[layer_id],
src_indices=host_indices,
dst_indices=device_indices,
item_size=self.token_stride_size,
)
elif self.layout == "page_first":
transfer_kv_per_layer_mla_pf_lf(
src=self.kv_buffer,
dst=device_pool.kv_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,
)
if self.can_use_jit:
jit_transfer_hicache_one_layer_mla(
cache_dst=device_pool.kv_buffer[layer_id],
cache_src=self.data_refs[layer_id],
indices_dst=device_indices,
indices_src=host_indices,
element_dim=self.kv_cache_dim,
)
else:
transfer_kv_per_layer_mla_pf_lf(
src=self.kv_buffer,
dst=device_pool.kv_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,
)
else:
raise ValueError(f"Unsupported layout: {self.layout}")
elif io_backend == "direct":
@@ -929,24 +1000,46 @@ class MLATokenToKVPoolHost(HostKVCache):
):
if io_backend == "kernel":
if self.layout == "layer_first":
transfer_kv_all_layer_mla(
src_layers=device_pool.data_ptrs,
dst_layers=self.data_ptrs,
src_indices=device_indices,
dst_indices=host_indices,
item_size=self.token_stride_size,
num_layers=self.layer_num,
)
if self.can_use_jit:
jit_transfer_hicache_all_layer_mla(
ptr_dst=self.data_ptrs,
indices_dst=host_indices,
ptr_src=device_pool.data_ptrs,
indices_src=device_indices,
cache_dst_stride_bytes=self.token_stride_size,
cache_src_stride_bytes=self.token_stride_size,
element_size=self.kv_cache_dim * self.dtype.itemsize,
)
else:
transfer_kv_all_layer_mla(
src_layers=device_pool.data_ptrs,
dst_layers=self.data_ptrs,
src_indices=device_indices,
dst_indices=host_indices,
item_size=self.token_stride_size,
num_layers=self.layer_num,
)
elif self.layout == "page_first":
transfer_kv_all_layer_mla_lf_pf(
src_layers=device_pool.data_ptrs,
dst=self.kv_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,
)
if self.can_use_jit:
jit_transfer_hicache_all_layer_mla(
ptr_dst=self.data_ptrs,
indices_dst=host_indices,
ptr_src=device_pool.data_ptrs,
indices_src=device_indices,
cache_src_stride_bytes=self.token_stride_size,
cache_dst_stride_bytes=self.layout_dim,
element_size=self.kv_cache_dim * self.dtype.itemsize,
)
else:
transfer_kv_all_layer_mla_lf_pf(
src_layers=device_pool.data_ptrs,
dst=self.kv_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,
)
else:
raise ValueError(f"Unsupported layout: {self.layout}")
elif io_backend == "direct":