[Hicache & JIT_kernel] Support page first layout & mla jit kernel (#18311)
This commit is contained in:
@@ -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 vec_k = load_vec<kElementSize, kNumThreads>(src_k);
|
||||
store_vec<kElementSize, kNumThreads>(dst_k, vec_k);
|
||||
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_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);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
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,20 +204,22 @@ 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 vec_k = load_vec<kElementSize, kNumThreads>(src_k);
|
||||
store_vec<kElementSize, kNumThreads>(dst_k, vec_k);
|
||||
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_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);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template <int64_t kElementSize, uint32_t kUnroll, uint32_t kBlockQuota, uint32_t kBlockSize>
|
||||
struct HiCacheKernel {
|
||||
@@ -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
|
||||
|
||||
@@ -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"]))
|
||||
@@ -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,6 +315,14 @@ class MHATokenToKVPoolHost(HostKVCache):
|
||||
element_size=self.element_dim * self.dtype.itemsize
|
||||
)
|
||||
|
||||
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(
|
||||
@@ -409,6 +423,20 @@ class MHATokenToKVPoolHost(HostKVCache):
|
||||
item_size=self.token_stride_size,
|
||||
)
|
||||
elif self.layout == "page_first":
|
||||
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],
|
||||
@@ -510,6 +538,21 @@ class MHATokenToKVPoolHost(HostKVCache):
|
||||
num_layers=self.layer_num,
|
||||
)
|
||||
elif self.layout == "page_first":
|
||||
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,
|
||||
@@ -766,6 +809,16 @@ class MLATokenToKVPoolHost(HostKVCache):
|
||||
device,
|
||||
allocator_type,
|
||||
)
|
||||
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],
|
||||
@@ -864,6 +917,15 @@ class MLATokenToKVPoolHost(HostKVCache):
|
||||
):
|
||||
if io_backend == "kernel":
|
||||
if self.layout == "layer_first":
|
||||
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],
|
||||
@@ -872,6 +934,15 @@ class MLATokenToKVPoolHost(HostKVCache):
|
||||
item_size=self.token_stride_size,
|
||||
)
|
||||
elif self.layout == "page_first":
|
||||
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],
|
||||
@@ -929,6 +1000,17 @@ class MLATokenToKVPoolHost(HostKVCache):
|
||||
):
|
||||
if io_backend == "kernel":
|
||||
if self.layout == "layer_first":
|
||||
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,
|
||||
@@ -938,6 +1020,17 @@ class MLATokenToKVPoolHost(HostKVCache):
|
||||
num_layers=self.layer_num,
|
||||
)
|
||||
elif self.layout == "page_first":
|
||||
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,
|
||||
|
||||
Reference in New Issue
Block a user