[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 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
|
||||
|
||||
Reference in New Issue
Block a user