[Refactor] Fix test and clean up hicache code (#18555)

This commit is contained in:
DarkSharpness
2026-02-18 14:37:46 +08:00
committed by GitHub
parent 95c44cea29
commit 9d138685c1
3 changed files with 252 additions and 260 deletions
+105 -133
View File
@@ -2,25 +2,19 @@
#include <sgl_kernel/utils.h>
#include <sgl_kernel/utils.cuh>
#include <sgl_kernel/vec.cuh>
#include <dlpack/dlpack.h>
#include <algorithm>
#include <concepts>
#include <cstddef>
#include <cstdint>
#include <type_traits>
namespace device::warp {
template <typename T, std::size_t N>
struct device_vec {
T data[N];
};
namespace device {
namespace details {
template <std::size_t kUnit>
template <int kUnit>
inline constexpr auto get_mem_package() {
if constexpr (kUnit == 16) {
return uint4{};
@@ -33,90 +27,95 @@ inline constexpr auto get_mem_package() {
}
}
template <std::size_t kBytes, std::size_t kUnit>
using mem_package_t = decltype(get_mem_package<kUnit>());
template <int kUnit>
using PackageType = decltype(get_mem_package<kUnit>());
__always_inline __device__ auto load_nc(const uint1* __restrict__ src) -> uint1 {
SGL_DEVICE uint1 load_nc(const uint1* __restrict__ src) {
uint32_t tmp;
asm volatile("ld.global.cs.b32 %0,[%1];" : "=r"(tmp) : "l"(src));
asm volatile("ld.global.L1::no_allocate.b32 %0,[%1];" : "=r"(tmp) : "l"(src));
return uint1{tmp};
}
__always_inline __device__ auto load_nc(const uint2* __restrict__ src) -> uint2 {
SGL_DEVICE uint2 load_nc(const uint2* __restrict__ src) {
uint32_t tmp0, tmp1;
asm volatile("ld.global.cs.v2.b32 {%0,%1},[%2];" : "=r"(tmp0), "=r"(tmp1) : "l"(src));
asm volatile("ld.global.L1::no_allocate.v2.b32 {%0,%1},[%2];" : "=r"(tmp0), "=r"(tmp1) : "l"(src));
return uint2{tmp0, tmp1};
}
__always_inline __device__ auto load_nc(const uint4* __restrict__ src) -> uint4 {
SGL_DEVICE uint4 load_nc(const uint4* __restrict__ src) {
uint32_t tmp0, tmp1, tmp2, tmp3;
asm volatile("ld.global.cs.v4.b32 {%0,%1,%2,%3},[%4];" : "=r"(tmp0), "=r"(tmp1), "=r"(tmp2), "=r"(tmp3) : "l"(src));
asm volatile("ld.global.L1::no_allocate.v4.b32 {%0,%1,%2,%3},[%4];"
: "=r"(tmp0), "=r"(tmp1), "=r"(tmp2), "=r"(tmp3)
: "l"(src));
return uint4{tmp0, tmp1, tmp2, tmp3};
}
__always_inline __device__ void store_nc(uint1* __restrict__ dst, const uint1& value) {
SGL_DEVICE void store_nc(uint1* __restrict__ dst, const uint1& value) {
uint32_t tmp = value.x;
asm volatile("st.global.cs.b32 [%0],%1;" ::"l"(dst), "r"(tmp));
asm volatile("st.global.L1::no_allocate.b32 [%0],%1;" ::"l"(dst), "r"(tmp));
}
__always_inline __device__ void store_nc(uint2* __restrict__ dst, const uint2& value) {
SGL_DEVICE void store_nc(uint2* __restrict__ dst, const uint2& value) {
uint32_t tmp0 = value.x;
uint32_t tmp1 = value.y;
asm volatile("st.global.cs.v2.b32 [%0],{%1,%2};" ::"l"(dst), "r"(tmp0), "r"(tmp1));
asm volatile("st.global.L1::no_allocate.v2.b32 [%0],{%1,%2};" ::"l"(dst), "r"(tmp0), "r"(tmp1));
}
__always_inline __device__ void store_nc(uint4* __restrict__ dst, const uint4& value) {
SGL_DEVICE void store_nc(uint4* __restrict__ dst, const uint4& value) {
uint32_t tmp0 = value.x;
uint32_t tmp1 = value.y;
uint32_t tmp2 = value.z;
uint32_t tmp3 = value.w;
asm volatile("st.global.cs.v4.b32 [%0],{%1,%2,%3,%4};" ::"l"(dst), "r"(tmp0), "r"(tmp1), "r"(tmp2), "r"(tmp3));
asm volatile(
"st.global.L1::no_allocate.v4.b32 [%0],{%1,%2,%3,%4};" ::"l"(dst), "r"(tmp0), "r"(tmp1), "r"(tmp2), "r"(tmp3));
}
} // namespace details
template <std::size_t kBytes, std::size_t kUnit, std::size_t kThreads>
__always_inline __device__ auto load_vec(const void* __restrict__ src) {
using Package = details::mem_package_t<kBytes, kUnit>;
constexpr auto kBytesPerLoop = sizeof(Package) * kThreads;
constexpr auto kLoopCount = kBytes / kBytesPerLoop;
static_assert(kBytes % kBytesPerLoop == 0, "kBytes must be multiple of 128 bytes");
template <int64_t kBytes, uint32_t kNumThreads>
SGL_DEVICE auto load_vec(const void* __restrict__ src) {
static_assert(kBytes % 128 == 0, "kBytes must be multiple of 128 bytes");
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>;
const auto src_packed = static_cast<const Package*>(src);
const auto lane_id = threadIdx.x % kThreads;
device_vec<Package, kLoopCount> vec;
const auto lane_id = threadIdx.x % kNumThreads;
Storage vec;
#pragma unroll kLoopCount
for (std::size_t i = 0; i < kLoopCount; ++i) {
const auto j = i * kThreads + lane_id;
vec.data[i] = details::load_nc(src_packed + j);
for (uint32_t i = 0; i < kLoopCount; ++i) {
const auto j = i * kNumThreads + lane_id;
vec.data[i] = details::load_nc(&src_packed[j]);
}
return vec;
}
template <std::size_t kBytes, std::size_t kUnit, std::size_t kThreads, typename Tp>
__always_inline __device__ void store_vec(void* __restrict__ dst, const Tp& vec) {
using Package = details::mem_package_t<kBytes, kUnit>;
constexpr auto kBytesPerLoop = sizeof(Package) * kThreads;
constexpr auto kLoopCount = kBytes / kBytesPerLoop;
static_assert(kBytes % kBytesPerLoop == 0, "kBytes must be multiple of 128 bytes");
static_assert(std::is_same_v<Tp, device_vec<Package, kLoopCount>>);
template <int64_t kBytes, uint32_t kNumThreads, typename Storage>
SGL_DEVICE void store_vec(void* __restrict__ dst, const Storage& vec) {
using Package = std::decay_t<decltype(vec.data[0])>;
constexpr uint32_t kBytesPerLoop = sizeof(Package) * kNumThreads;
constexpr uint32_t kLoopCount = kBytes / kBytesPerLoop;
static_assert(kBytes % kBytesPerLoop == 0, "Invalid Storage configuration");
const auto dst_packed = static_cast<Package*>(dst);
const auto lane_id = threadIdx.x % kThreads;
const auto lane_id = threadIdx.x % kNumThreads;
#pragma unroll kLoopCount
for (std::size_t i = 0; i < kLoopCount; ++i) {
const auto j = i * kThreads + lane_id;
details::store_nc(dst_packed + j, vec.data[i]);
for (uint32_t i = 0; i < kLoopCount; ++i) {
const auto j = i * kNumThreads + lane_id;
details::store_nc(&dst_packed[j], vec.data[i]);
}
}
} // namespace device::warp
} // namespace device
namespace {
#define SGL_HICACHE_KERNEL __global__ __launch_bounds__(kBlockSize, 1)
struct HicacheKernelParams {
void* __restrict__ k_cache_dst;
void* __restrict__ v_cache_dst;
@@ -124,118 +123,89 @@ struct HicacheKernelParams {
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
int64_t kv_cache_src_stride;
int64_t kv_cache_dst_stride;
uint32_t length;
uint32_t num_layers = 0; // only used in all_layer transfer
};
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(
const __grid_constant__ HicacheKernelParams params) {
// each warp acts as a worker
template <typename T, int64_t kElementSize, uint32_t kUnroll, uint32_t kBlockQuota, uint32_t kBlockSize>
SGL_HICACHE_KERNEL void hicache_transfer_per_layer(const __grid_constant__ HicacheKernelParams params) {
using namespace device;
static_assert(kNumThreads % kWarpThreads == 0);
static_assert(kBlockSize % kWarpThreads == 0);
static_assert(kWarpThreads % kUnroll == 0);
constexpr auto kWarpThreads = device::kWarpThreads / kUnroll;
constexpr auto kWarpsPerBlock = kNumThreads / kWarpThreads;
constexpr auto kWorkers = kWarpsPerBlock * kBlockQuota;
constexpr uint32_t kNumThreads = kWarpThreads / kUnroll;
constexpr uint32_t kWorkersPerBlock = kBlockSize / kNumThreads;
constexpr uint32_t kNumWorkers = kWorkersPerBlock * 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, _ // metadata
kv_cache_src_stride, kv_cache_dst_stride, length, _ // metadata
] = params;
const auto warp_id = blockIdx.x * kWarpsPerBlock + threadIdx.x / kWarpThreads;
// force to transfer 128 bytes per iteration
// since the PCIe transaction size is 128 bytes aligned
constexpr auto kGranularity = 128 / kWarpThreads;
for (auto i = warp_id; i < length; i += kWorkers) {
const uint32_t work_id = blockIdx.x * kWorkersPerBlock + threadIdx.x / kNumThreads;
for (uint32_t i = work_id; i < length; i += kNumWorkers) {
const auto pos_src = static_cast<const T*>(indices_src)[i];
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 = 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);
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 <
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(
const __grid_constant__ HicacheKernelParams params) {
// each warp acts as a worker
template <typename T, int64_t kElementSize, uint32_t kUnroll, uint32_t kBlockQuota, uint32_t kBlockSize>
SGL_HICACHE_KERNEL void hicache_transfer_all_layer(const __grid_constant__ HicacheKernelParams 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>;
using src_ptr_t = const void*;
using dst_ptr_t = void*;
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;
static_assert(kBlockSize % kWarpThreads == 0);
static_assert(kWarpThreads % kUnroll == 0);
constexpr uint32_t kNumThreads = kWarpThreads / kUnroll;
constexpr uint32_t kWorkersPerBlock = kBlockSize / kNumThreads;
constexpr uint32_t kNumWorkers = kWorkersPerBlock * 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 // metadata
kv_cache_src_stride, kv_cache_dst_stride, length, num_layers // metadata
] = params;
const auto warp_id = blockIdx.x * kWarpsPerBlock + threadIdx.x / kWarpThreads;
// force to transfer 128 bytes per iteration
// since the PCIe transaction size is 128 bytes aligned
constexpr auto kGranularity = 128 / kWarpThreads;
for (auto i = warp_id; i < length; i += kWorkers) {
const uint32_t work_id = blockIdx.x * kWorkersPerBlock + threadIdx.x / kNumThreads;
for (uint32_t i = work_id; i < length; i += kNumWorkers) {
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];
const auto k_cache_dst = static_cast<dst_ptr_t>(k_ptr_dst)[layer];
const auto v_cache_dst = static_cast<dst_ptr_t>(v_ptr_dst)[layer];
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 = 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);
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 <
std::size_t kElementSize,
std::size_t kUnroll,
std::size_t kBlockQuota,
std::size_t kNumThreads,
std::size_t kMaxOccupancy>
template <int64_t kElementSize, uint32_t kUnroll, uint32_t kBlockQuota, uint32_t kBlockSize>
struct HiCacheKernel {
template <typename T>
static constexpr auto _kernel_one =
hicache_transfer_per_layer<T, kElementSize, kUnroll, kBlockQuota, kNumThreads, kMaxOccupancy>;
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, kNumThreads, kMaxOccupancy>;
static constexpr auto kernel_all = hicache_transfer_all_layer<T, kElementSize, kUnroll, kBlockQuota, kBlockSize>;
static void run_one(
const tvm::ffi::TensorView k_cache_dst,
@@ -283,13 +253,13 @@ struct HiCacheKernel {
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>(N.unwrap()) * dtype_size;
const auto kv_cache_dst_stride = static_cast<std::size_t>(M.unwrap()) * dtype_size;
const auto length = static_cast<uint32_t>(L.unwrap());
const auto kv_cache_src_stride = static_cast<int64_t>(N.unwrap() * dtype_size);
const auto kv_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 = kNumThreads / (device::kWarpThreads / kUnroll);
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 = k_cache_dst_ptr,
@@ -298,12 +268,12 @@ struct HiCacheKernel {
.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,
.length = length,
};
const auto kernel = use_int32 ? _kernel_one<int32_t> : _kernel_one<int64_t>;
LaunchKernel(num_blocks, kNumThreads, device)(kernel, params);
const auto kernel = use_int32 ? kernel_one<int32_t> : kernel_one<int64_t>;
LaunchKernel(num_blocks, kBlockSize, device)(kernel, params);
}
static void run_all(
@@ -313,8 +283,8 @@ struct HiCacheKernel {
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 int64_t kv_src_stride_bytes,
const int64_t kv_dst_stride_bytes) {
using namespace host;
auto N = SymbolicSize{"num_layers"};
@@ -342,11 +312,11 @@ struct HiCacheKernel {
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 length = static_cast<uint32_t>(L.unwrap());
const auto use_int32 = dtype_.unwrap().bits == 32;
const auto device = device_.unwrap();
constexpr auto kWorkersPerBlock = kNumThreads / (device::kWarpThreads / kUnroll);
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 = k_cache_dst_ptr,
@@ -355,14 +325,16 @@ struct HiCacheKernel {
.k_cache_src = k_cache_src_ptr,
.v_cache_src = v_cache_src_ptr,
.indices_src = indices_src_ptr,
.kv_cache_src_stride = kv_src_stride_bytes,
.kv_cache_dst_stride = kv_dst_stride_bytes,
.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()),
.num_layers = static_cast<uint32_t>(N.unwrap()),
};
const auto kernel = use_int32 ? _kernel_all<int32_t> : _kernel_all<int64_t>;
LaunchKernel(num_blocks, kNumThreads, device)(kernel, params);
const auto kernel = use_int32 ? kernel_all<int32_t> : kernel_all<int64_t>;
LaunchKernel(num_blocks, kBlockSize, device)(kernel, params);
}
};
#undef SGL_HICACHE_KERNEL
} // namespace