#include #include #include #include #include #include #include #include #include namespace device::warp { template struct device_vec { T data[N]; }; namespace details { template inline constexpr auto get_mem_package() { if constexpr (kUnit == 16) { return uint4{}; } else if constexpr (kUnit == 8) { return uint2{}; } else if constexpr (kUnit == 4) { return uint1{}; } else { static_assert(kUnit == 16 || kUnit == 8 || kUnit == 4, "Unsupported memory package size"); } } template using mem_package_t = decltype(get_mem_package()); __always_inline __device__ auto load_nc(const uint1* __restrict__ src) -> uint1 { uint32_t tmp; asm volatile("ld.global.cs.b32 %0,[%1];" : "=r"(tmp) : "l"(src)); return uint1{tmp}; } __always_inline __device__ auto load_nc(const uint2* __restrict__ src) -> uint2 { uint32_t tmp0, tmp1; asm volatile("ld.global.cs.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 { 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)); return uint4{tmp0, tmp1, tmp2, tmp3}; } __always_inline __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)); } __always_inline __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)); } __always_inline __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)); } } // namespace details template __always_inline __device__ auto load_vec(const void* __restrict__ src) { using Package = details::mem_package_t; constexpr auto kBytesPerLoop = sizeof(Package) * kThreads; constexpr auto kLoopCount = kBytes / kBytesPerLoop; static_assert(kBytes % kBytesPerLoop == 0, "kBytes must be multiple of 128 bytes"); const auto src_packed = static_cast(src); const auto lane_id = threadIdx.x % kThreads; device_vec 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); } return vec; } template __always_inline __device__ void store_vec(void* __restrict__ dst, const Tp& vec) { using Package = details::mem_package_t; 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>); const auto dst_packed = static_cast(dst); const auto lane_id = threadIdx.x % kThreads; #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]); } } } // namespace device::warp namespace { struct HicacheKernelParams { 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 }; 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 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, _ // 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 auto pos_src = static_cast(indices_src)[i]; 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 = warp::load_vec(src_k); const auto vec_v = warp::load_vec(src_v); warp::store_vec(dst_k, vec_k); warp::store_vec(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 using namespace device; using src_ptr_t = std::add_pointer_t; using dst_ptr_t = std::add_pointer_t; static_assert(kNumThreads % kWarpThreads == 0); constexpr auto kWarpThreads = device::kWarpThreads / kUnroll; constexpr auto kWarpsPerBlock = static_cast(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 // 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 auto pos_src = static_cast(indices_src)[i]; const auto pos_dst = static_cast(indices_dst)[i]; for (std::size_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 = warp::load_vec(src_k); const auto vec_v = warp::load_vec(src_v); warp::store_vec(dst_k, vec_k); warp::store_vec(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> struct HiCacheKernel { template static constexpr auto _kernel_one = hicache_transfer_per_layer; template static constexpr auto _kernel_all = hicache_transfer_all_layer; static void run_one( 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) { 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() .verify(k_cache_src) .verify(v_cache_src); TensorMatcher({-1, D}) // .with_strides({M, 1}) .with_dtype(cache_dtype) .with_device() .verify(k_cache_dst) .verify(v_cache_dst); TensorMatcher({L}) // .with_dtype(indices_dtype) .with_device(indices_device) .verify(indices_src) .verify(indices_dst); // verify dimension match 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(L.unwrap()); const auto kv_cache_src_stride = static_cast(N.unwrap()) * dtype_size; const auto kv_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 = kNumThreads / (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, .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, }; const auto kernel = use_int32 ? _kernel_one : _kernel_one; LaunchKernel(num_blocks, kNumThreads, device)(kernel, params); } static void run_all( 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) { 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(k_ptr_src) .verify(v_ptr_src) .verify(k_ptr_dst) .verify(v_ptr_dst); TensorMatcher({L}) // .with_dtype(dtype_) .with_device(device_) .verify(indices_src) .verify(indices_dst); // verify dimension match 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(L.unwrap()); const auto use_int32 = dtype_.unwrap().bits == 32; const auto device = device_.unwrap(); constexpr auto kWorkersPerBlock = kNumThreads / (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, .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(N.unwrap()), }; const auto kernel = use_int32 ? _kernel_all : _kernel_all; LaunchKernel(num_blocks, kNumThreads, device)(kernel, params); } }; } // namespace