Skip to content
Merged
Show file tree
Hide file tree
Changes from 1 commit
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
249 changes: 249 additions & 0 deletions python/sglang/jit_kernel/csrc/hicache.cuh
Original file line number Diff line number Diff line change
Expand Up @@ -130,6 +130,22 @@ struct HicacheKernelParams {
std::size_t num_layers = 0; // only used in all_layer transfer
};

struct HicachePfKernelParams {
void* __restrict__ k_cache_dst;
void* __restrict__ v_cache_dst;
const void* __restrict__ indices_dst;
void* __restrict__ k_cache_src;
void* __restrict__ v_cache_src;
const void* __restrict__ indices_src;
std::size_t length;
std::size_t kv_cache_src_stride;
std::size_t kv_cache_dst_stride;
std::size_t num_layers = 0; // only used in all_layer transfer
std::size_t layer_id = 0; // only used in per_layer transfer
std::size_t src_layout_dim = 0; // only used in per_layer transfer
std::size_t dst_layout_dim = 0; // only used in all_layer transfer
};
Comment thread
DarkSharpness marked this conversation as resolved.
Outdated

template <
std::integral T,
std::size_t kElementSize,
Expand Down Expand Up @@ -223,6 +239,97 @@ __global__ __launch_bounds__(kNumThreads, kMaxOccupancy) void hicache_transfer_a
}
}

// Page first to layer first (per layer)
template <
std::integral T,
std::size_t kElementSize,
std::size_t kUnroll,
std::size_t kBlockQuota,
std::size_t kNumThreads,
std::size_t kMaxOccupancy>
__global__ __launch_bounds__(kNumThreads, kMaxOccupancy) void hicache_transfer_per_layer_pf_lf(
const __grid_constant__ HicachePfKernelParams params) {
using namespace device;
static_assert(kNumThreads % kWarpThreads == 0);
static_assert(kWarpThreads % kUnroll == 0);

constexpr auto kWarpThreads = device::kWarpThreads / kUnroll;
constexpr auto kWarpsPerBlock = kNumThreads / kWarpThreads;
constexpr auto kWorkers = kWarpsPerBlock * kBlockQuota;

const auto& [
k_cache_dst, v_cache_dst, indices_dst, // dst
k_cache_src, v_cache_src, indices_src, // src
length, kv_cache_src_stride, kv_cache_dst_stride, num_layers_unused, layer_id, src_layout_dim, dst_layout_dim_unused // metadata
] = params;
const auto warp_id = blockIdx.x * kWarpsPerBlock + threadIdx.x / kWarpThreads;

constexpr auto kGranularity = 128 / kWarpThreads;

for (auto i = warp_id; i < length; i += kWorkers) {
const auto pos_src = static_cast<const T*>(indices_src)[i];
const auto pos_dst = static_cast<const T*>(indices_dst)[i];
// Page first: base + page_id * page_dim + layer_id * item_size_bytes
const auto src_k = pointer::offset(k_cache_src, pos_src * src_layout_dim + layer_id * kv_cache_src_stride);
// Layer first: base + layer_id * layer_dim + page_id * item_size_bytes
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 * src_layout_dim + layer_id * kv_cache_src_stride);
const auto dst_v = pointer::offset(v_cache_dst, pos_dst * kv_cache_dst_stride);
const auto vec_k = warp::load_vec<kElementSize, kGranularity, kWarpThreads>(src_k);
const auto vec_v = warp::load_vec<kElementSize, kGranularity, kWarpThreads>(src_v);
warp::store_vec<kElementSize, kGranularity, kWarpThreads>(dst_k, vec_k);
warp::store_vec<kElementSize, kGranularity, kWarpThreads>(dst_v, vec_v);
}
}

// Layer first to page first (all layers)
template <
std::integral T,
std::size_t kElementSize,
std::size_t kUnroll,
std::size_t kBlockQuota,
std::size_t kNumThreads,
std::size_t kMaxOccupancy>
__global__ __launch_bounds__(kNumThreads, kMaxOccupancy) void hicache_transfer_all_layer_lf_pf(
const __grid_constant__ HicachePfKernelParams params) {
using namespace device;
using src_ptr_t = std::add_pointer_t<const void* const>;
using dst_ptr_t = std::add_pointer_t<void* const>;

static_assert(kNumThreads % kWarpThreads == 0);
constexpr auto kWarpThreads = device::kWarpThreads / kUnroll;
constexpr auto kWarpsPerBlock = static_cast<uint32_t>(kNumThreads) / kWarpThreads;
constexpr auto kWorkers = kWarpsPerBlock * kBlockQuota;

const auto& [
k_ptr_dst, v_ptr_dst, indices_dst, // dst
k_ptr_src, v_ptr_src, indices_src, // src
length, kv_cache_src_stride, kv_cache_dst_stride, num_layers, layer_id_unused, src_layout_dim_unused, dst_layout_dim // metadata
] = params;
const auto warp_id = blockIdx.x * kWarpsPerBlock + threadIdx.x / kWarpThreads;

constexpr auto kGranularity = 128 / kWarpThreads;

for (auto i = warp_id; i < length; i += kWorkers) {
const auto pos_src = static_cast<const T*>(indices_src)[i];
const auto pos_dst = static_cast<const T*>(indices_dst)[i];
for (std::size_t layer = 0; layer < num_layers; ++layer) {
const auto k_cache_src = static_cast<src_ptr_t>(k_ptr_src)[layer];
const auto v_cache_src = static_cast<src_ptr_t>(v_ptr_src)[layer];
// Layer first: base + layer_id * layer_dim + page_id * item_size_bytes
const auto src_k = pointer::offset(k_cache_src, pos_src * kv_cache_src_stride);
// Page first: base + page_id * page_dim + layer_id * item_size_bytes
const auto dst_k = pointer::offset(k_ptr_dst, pos_dst * dst_layout_dim + layer * 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_ptr_dst, pos_dst * dst_layout_dim + layer * kv_cache_dst_stride);
const auto vec_k = warp::load_vec<kElementSize, kGranularity, kWarpThreads>(src_k);
const auto vec_v = warp::load_vec<kElementSize, kGranularity, kWarpThreads>(src_v);
warp::store_vec<kElementSize, kGranularity, kWarpThreads>(dst_k, vec_k);
warp::store_vec<kElementSize, kGranularity, kWarpThreads>(dst_v, vec_v);
}
}
}

template <
std::size_t kElementSize,
std::size_t kUnroll,
Expand All @@ -236,6 +343,12 @@ struct HiCacheKernel {
template <typename T>
static constexpr auto _kernel_all =
hicache_transfer_all_layer<T, kElementSize, kUnroll, kBlockQuota, kNumThreads, kMaxOccupancy>;
template <typename T>
static constexpr auto _kernel_pf_lf =
hicache_transfer_per_layer_pf_lf<T, kElementSize, kUnroll, kBlockQuota, kNumThreads, kMaxOccupancy>;
template <typename T>
static constexpr auto _kernel_lf_pf =
hicache_transfer_all_layer_lf_pf<T, kElementSize, kUnroll, kBlockQuota, kNumThreads, kMaxOccupancy>;

static void run_one(
const tvm::ffi::TensorView k_cache_dst,
Expand Down Expand Up @@ -363,6 +476,142 @@ struct HiCacheKernel {
const auto kernel = use_int32 ? _kernel_all<int32_t> : _kernel_all<int64_t>;
LaunchKernel(num_blocks, kNumThreads, device)(kernel, params);
}

static void run_pf_lf(
const tvm::ffi::TensorView k_cache_dst,
const tvm::ffi::TensorView v_cache_dst,
const tvm::ffi::TensorView indices_dst,
const tvm::ffi::TensorView k_cache_src,
const tvm::ffi::TensorView v_cache_src,
const tvm::ffi::TensorView indices_src,
const std::size_t layer_id,
const std::size_t src_layout_dim) {
using namespace host;

auto D = SymbolicSize{"head dimension"};
auto N = SymbolicSize{"src kv stride"};
auto M = SymbolicSize{"dst kv 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(k_cache_src)
.verify(v_cache_src);
TensorMatcher({-1, D}) //
.with_strides({M, 1})
.with_dtype(cache_dtype)
.with_device<kDLCUDA, kDLCUDAHost, kDLCPU>()
.verify(k_cache_dst)
.verify(v_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: cache dimension mismatch.");

const auto k_cache_dst_ptr = k_cache_dst.data_ptr();
const auto v_cache_dst_ptr = v_cache_dst.data_ptr();
const auto k_cache_src_ptr = k_cache_src.data_ptr();
const auto v_cache_src_ptr = v_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<std::size_t>(L.unwrap());
const auto kv_cache_src_stride = static_cast<std::size_t>(D.unwrap()) * dtype_size;
const auto kv_cache_dst_stride = static_cast<std::size_t>(M.unwrap()) * dtype_size;
const auto use_int32 = indices_dtype.unwrap().bits == 32;
const auto device = indices_device.unwrap();

constexpr auto kWorkersPerBlock = kNumThreads / (device::kWarpThreads / kUnroll);
const auto num_blocks = std::min(div_ceil(length, kWorkersPerBlock), kBlockQuota);
const auto params = HicachePfKernelParams{
.k_cache_dst = k_cache_dst_ptr,
.v_cache_dst = v_cache_dst_ptr,
.indices_dst = indices_dst_ptr,
.k_cache_src = k_cache_src_ptr,
.v_cache_src = v_cache_src_ptr,
.indices_src = indices_src_ptr,
.length = length,
.kv_cache_src_stride = kv_cache_src_stride,
.kv_cache_dst_stride = kv_cache_dst_stride,
.layer_id = layer_id,
.src_layout_dim = src_layout_dim,
};
const auto kernel = use_int32 ? _kernel_pf_lf<int32_t> : _kernel_pf_lf<int64_t>;
LaunchKernel(num_blocks, kNumThreads, device)(kernel, params);
}
Comment thread
DarkSharpness marked this conversation as resolved.
Outdated

static void run_lf_pf(
const tvm::ffi::TensorView k_ptr_dst,
const tvm::ffi::TensorView v_ptr_dst,
const tvm::ffi::TensorView indices_dst,
const tvm::ffi::TensorView k_ptr_src,
const tvm::ffi::TensorView v_ptr_src,
const tvm::ffi::TensorView indices_src,
const std::size_t kv_src_stride,
const std::size_t kv_dst_stride,
const std::size_t dst_layout_dim) {
using namespace host;

auto N = SymbolicSize{"num_layers"};
auto L = SymbolicSize{"indices length"};
auto dtype_ = SymbolicDType{};
auto cache_dtype = SymbolicDType{};
auto src_device = SymbolicDevice{};
auto dst_device = SymbolicDevice{};

TensorMatcher({N}) //
.with_dtype<uint64_t>()
.with_device<kDLCUDA>(src_device)
.verify(k_ptr_src)
.verify(v_ptr_src);
TensorMatcher({-1}) //
.with_dtype(cache_dtype)
.with_device<kDLCUDA, kDLCUDAHost, kDLCPU>(dst_device)
.verify(k_ptr_dst)
.verify(v_ptr_dst);
TensorMatcher({L}) //
.with_dtype<int32_t, int64_t>(dtype_)
.with_device<kDLCUDA>(src_device)
.verify(indices_src)
.verify(indices_dst);

const auto k_cache_dst_ptr = k_ptr_dst.data_ptr();
const auto v_cache_dst_ptr = v_ptr_dst.data_ptr();
const auto k_cache_src_ptr = k_ptr_src.data_ptr();
const auto v_cache_src_ptr = v_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<std::size_t>(L.unwrap());
const auto use_int32 = dtype_.unwrap().bits == 32;
const auto device = src_device.unwrap();

constexpr auto kWorkersPerBlock = kNumThreads / (device::kWarpThreads / kUnroll);
const auto num_blocks = std::min(div_ceil(length, kWorkersPerBlock), kBlockQuota);
const auto params = HicachePfKernelParams{
.k_cache_dst = k_cache_dst_ptr,
.v_cache_dst = v_cache_dst_ptr,
.indices_dst = indices_dst_ptr,
.k_cache_src = k_cache_src_ptr,
.v_cache_src = v_cache_src_ptr,
.indices_src = indices_src_ptr,
.length = length,
.kv_cache_src_stride = kv_src_stride,
.kv_cache_dst_stride = kv_dst_stride,
.num_layers = static_cast<std::size_t>(N.unwrap()),
.dst_layout_dim = dst_layout_dim,
};
const auto kernel = use_int32 ? _kernel_lf_pf<int32_t> : _kernel_lf_pf<int64_t>;
LaunchKernel(num_blocks, kNumThreads, device)(kernel, params);
}
};

} // namespace
82 changes: 82 additions & 0 deletions python/sglang/jit_kernel/hicache.py
Original file line number Diff line number Diff line change
Expand Up @@ -29,6 +29,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_pf_lf", f"&HiCacheKernel<{args}>::run_pf_lf"),
("launch_lf_pf", f"&HiCacheKernel<{args}>::run_lf_pf"),
],
)

Expand Down Expand Up @@ -135,3 +137,83 @@ def transfer_hicache_all_layer(
kv_cache_src_stride_bytes,
kv_cache_dst_stride_bytes,
)


def transfer_hicache_per_layer_pf_lf(
k_cache_dst: torch.Tensor,
v_cache_dst: torch.Tensor,
indices_dst: torch.Tensor,
k_cache_src: torch.Tensor,
v_cache_src: torch.Tensor,
indices_src: torch.Tensor,
*,
layer_id: int,
src_layout_dim: int,
element_dim: int | None = None,
unroll: int | None = None,
block_quota: int | None = None,
) -> None:
element_dim = element_dim or k_cache_dst.size(-1)
k_cache_src = k_cache_src.view(-1, element_dim)
v_cache_src = v_cache_src.view(-1, element_dim)
k_cache_dst = k_cache_dst.view(-1, element_dim)
v_cache_dst = v_cache_dst.view(-1, element_dim)
element_size = element_dim * k_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_pf_lf(
k_cache_dst,
v_cache_dst,
indices_dst,
k_cache_src,
v_cache_src,
indices_src,
layer_id,
src_layout_dim,
)


def transfer_hicache_all_layer_lf_pf(
k_ptr_dst: torch.Tensor,
v_ptr_dst: torch.Tensor,
indices_dst: torch.Tensor,
k_ptr_src: torch.Tensor,
v_ptr_src: torch.Tensor,
indices_src: torch.Tensor,
*,
kv_cache_src_stride_bytes: int,
kv_cache_dst_stride_bytes: int,
dst_layout_dim: int,
element_size: int | None = None,
unroll: int | None = None,
block_quota: int | None = None,
) -> None:
if element_size is None:
assert kv_cache_dst_stride_bytes == kv_cache_src_stride_bytes
element_size = kv_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_lf_pf(
k_ptr_dst.view(-1),
v_ptr_dst.view(-1),
indices_dst,
k_ptr_src,
v_ptr_src,
indices_src,
kv_cache_src_stride_bytes,
kv_cache_dst_stride_bytes,
dst_layout_dim,
)


Loading