diff --git a/python/sglang/jit_kernel/csrc/hicache.cuh b/python/sglang/jit_kernel/csrc/hicache.cuh index 04f093a02ba6..ae297061136e 100644 --- a/python/sglang/jit_kernel/csrc/hicache.cuh +++ b/python/sglang/jit_kernel/csrc/hicache.cuh @@ -14,6 +14,11 @@ namespace device { namespace details { +template +struct LocalStorage { + T data[N]; +}; + template 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; + using Storage = details::LocalStorage; const auto src_packed = static_cast(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 +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(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(src_k); - const auto vec_v = load_vec(src_v); store_vec(dst_k, vec_k); - store_vec(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(src_v); + store_vec(dst_v, vec_v); + } } } -template +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(indices_dst)[i]; for (uint32_t layer = 0; layer < num_layers; ++layer) { const auto k_cache_src = static_cast(k_ptr_src)[layer]; - const auto v_cache_src = static_cast(v_ptr_src)[layer]; const auto k_cache_dst = static_cast(k_ptr_dst)[layer]; - const auto v_cache_dst = static_cast(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(src_k); - const auto vec_v = load_vec(src_v); store_vec(dst_k, vec_k); - store_vec(dst_v, vec_v); + if constexpr (!kIsMLA) { + const auto v_cache_src = static_cast(v_ptr_src)[layer]; + const auto v_cache_dst = static_cast(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(src_v); + store_vec(dst_v, vec_v); + } } } } @@ -206,6 +227,12 @@ struct HiCacheKernel { static constexpr auto kernel_one = hicache_transfer_per_layer; template static constexpr auto kernel_all = hicache_transfer_all_layer; + template + static constexpr auto kernel_one_mla = + hicache_transfer_per_layer; + template + static constexpr auto kernel_all_mla = + hicache_transfer_all_layer; static void run_one( const tvm::ffi::TensorView k_cache_dst, @@ -333,6 +360,119 @@ struct HiCacheKernel { const auto kernel = use_int32 ? kernel_all : kernel_all; 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() + .verify(cache_src); + TensorMatcher({-1, D}) // + .with_strides({M, 1}) + .with_dtype(cache_dtype) + .with_device() + .verify(cache_dst); + TensorMatcher({L}) // + .with_dtype(indices_dtype) + .with_device(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(L.unwrap()); + const auto cache_src_stride = static_cast(N.unwrap() * dtype_size); + const auto cache_dst_stride = static_cast(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 : kernel_one_mla; + 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() + .with_device(device_) + .verify(ptr_src) + .verify(ptr_dst); + TensorMatcher({L}) // + .with_dtype(dtype_) + .with_device(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(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(N.unwrap()), + }; + const auto kernel = use_int32 ? kernel_all_mla : kernel_all_mla; + LaunchKernel(num_blocks, kBlockSize, device)(kernel, params); + } }; #undef SGL_HICACHE_KERNEL diff --git a/python/sglang/jit_kernel/hicache.py b/python/sglang/jit_kernel/hicache.py index 7a357790147b..0e5ed5802fc2 100644 --- a/python/sglang/jit_kernel/hicache.py +++ b/python/sglang/jit_kernel/hicache.py @@ -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, + ) diff --git a/python/sglang/jit_kernel/tests/test_hicache.py b/python/sglang/jit_kernel/tests/test_hicache.py new file mode 100644 index 000000000000..b6059b2c1d50 --- /dev/null +++ b/python/sglang/jit_kernel/tests/test_hicache.py @@ -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"])) diff --git a/python/sglang/srt/mem_cache/memory_pool_host.py b/python/sglang/srt/mem_cache/memory_pool_host.py index 10fc35239d89..9666080d3f72 100644 --- a/python/sglang/srt/mem_cache/memory_pool_host.py +++ b/python/sglang/srt/mem_cache/memory_pool_host.py @@ -21,9 +21,15 @@ 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 @@ def __init__( 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 @@ def load_to_device_per_layer( 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 @@ def backup_from_device_all_layer( 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 @@ def __init__( 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 @@ def load_to_device_per_layer( ): 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 @@ def backup_from_device_all_layer( ): 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":