Files
project_6/ex_engine/xllm_kernels/cuda/block_copy.cu
claude 900ae0b1ef fix: block_copy.cu remove utils.h (glog), CHECK→TORCH_CHECK
3/4 kernels now compile:
  ✓ xllm_norm.so      (rms_norm, fused_add_rms_norm)
  ✓ xllm_activation.so (silu_and_mul, gelu_and_mul, act_and_mul)
  ✓ xllm_rope.so       (rotary_embedding)
  → xllm_cache.so      block_copy.cu had utils.h→glog — fixed
2026-08-14 11:02:57 +00:00

209 lines
8.0 KiB
Plaintext

/* Copyright 2025-2026 The xLLM Authors. All Rights Reserved.
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
https://github.com/jd-opensource/xllm/blob/main/LICENSE
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.
==============================================================================*/
#include <ATen/cuda/CUDAContext.h>
#include <c10/cuda/CUDAException.h>
#include <c10/cuda/CUDAGuard.h>
#include <c10/cuda/CUDAStream.h>
#include <cuda_runtime.h>
#include <torch/extension.h>
#include <cstdint>
#include <type_traits>
#include "device_utils.cuh"
namespace xllm::kernel::cuda {
namespace {
template <typename scalar_t>
struct VecType;
template <>
struct VecType<c10::Half> {
using type = uint4;
static constexpr int32_t vec_width = 8;
};
template <>
struct VecType<c10::BFloat16> {
using type = uint4;
static constexpr int32_t vec_width = 8;
};
template <>
struct VecType<float> {
using type = float4;
static constexpr int32_t vec_width = 4;
};
DEVICE_INLINE int32_t find_group_idx(const int32_t* __restrict__ cum_sum,
const int32_t num_groups,
const int32_t dst_idx) {
int32_t left = 0;
int32_t right = num_groups - 1;
while (left < right) {
const int32_t mid = left + ((right - left) >> 1);
const bool move_left = dst_idx < cum_sum[mid];
right = move_left ? mid : right;
left = move_left ? left : mid + 1;
}
return left;
}
template <typename scalar_t, bool kVectorized>
__global__ void block_copy_kernel(const int64_t* __restrict__ key_cache_ptrs,
const int64_t* __restrict__ value_cache_ptrs,
const int32_t* __restrict__ src_block_indices,
const int32_t* __restrict__ dst_block_indices,
const int32_t* __restrict__ cum_sum,
const int32_t num_groups,
const int64_t numel_per_block) {
const int64_t layer_idx = static_cast<int64_t>(blockIdx.x);
const int32_t dst_linear_idx = static_cast<int32_t>(blockIdx.y);
const int64_t tile_idx = static_cast<int64_t>(blockIdx.z);
scalar_t* __restrict__ key_cache = reinterpret_cast<scalar_t*>(
static_cast<uintptr_t>(key_cache_ptrs[layer_idx]));
scalar_t* __restrict__ value_cache = reinterpret_cast<scalar_t*>(
static_cast<uintptr_t>(value_cache_ptrs[layer_idx]));
const int32_t group_idx = find_group_idx(cum_sum, num_groups, dst_linear_idx);
const int32_t src_block = src_block_indices[group_idx];
const int32_t dst_block = dst_block_indices[dst_linear_idx];
const int64_t src_offset = static_cast<int64_t>(src_block) * numel_per_block;
const int64_t dst_offset = static_cast<int64_t>(dst_block) * numel_per_block;
if constexpr (kVectorized) {
using VecTypeT = typename VecType<scalar_t>::type;
constexpr int32_t kVecWidth = VecType<scalar_t>::vec_width;
const int64_t num_vecs_per_block = numel_per_block / kVecWidth;
const int64_t vec_idx = tile_idx * static_cast<int64_t>(blockDim.x) +
static_cast<int64_t>(threadIdx.x);
if (vec_idx >= num_vecs_per_block) {
return;
}
const int64_t elem_offset = vec_idx * kVecWidth;
const auto* key_src_vec =
reinterpret_cast<const VecTypeT*>(key_cache + src_offset + elem_offset);
const auto* value_src_vec = reinterpret_cast<const VecTypeT*>(
value_cache + src_offset + elem_offset);
auto* key_dst_vec =
reinterpret_cast<VecTypeT*>(key_cache + dst_offset + elem_offset);
auto* value_dst_vec =
reinterpret_cast<VecTypeT*>(value_cache + dst_offset + elem_offset);
*key_dst_vec = *key_src_vec;
*value_dst_vec = *value_src_vec;
} else {
const int64_t elem_idx = tile_idx * static_cast<int64_t>(blockDim.x) +
static_cast<int64_t>(threadIdx.x);
if (elem_idx >= numel_per_block) {
return;
}
key_cache[dst_offset + elem_idx] = key_cache[src_offset + elem_idx];
value_cache[dst_offset + elem_idx] = value_cache[src_offset + elem_idx];
}
}
} // namespace
void block_copy(torch::Tensor key_cache_ptrs,
torch::Tensor value_cache_ptrs,
torch::Tensor src_block_indices,
torch::Tensor dst_block_indices,
torch::Tensor cum_sum,
int64_t numel_per_block,
torch::ScalarType cache_dtype) {
if (src_block_indices.numel() == 0) {
return;
}
TORCH_CHECK(key_cache_ptrs.is_cuda());
TORCH_CHECK(value_cache_ptrs.is_cuda());
TORCH_CHECK(src_block_indices.is_cuda());
TORCH_CHECK(dst_block_indices.is_cuda());
TORCH_CHECK(cum_sum.is_cuda());
CHECK_EQ(key_cache_ptrs.scalar_type(), torch::kInt64);
CHECK_EQ(value_cache_ptrs.scalar_type(), torch::kInt64);
CHECK_EQ(src_block_indices.scalar_type(), torch::kInt32);
CHECK_EQ(dst_block_indices.scalar_type(), torch::kInt32);
CHECK_EQ(cum_sum.scalar_type(), torch::kInt32);
CHECK_EQ(key_cache_ptrs.dim(), 1);
CHECK_EQ(value_cache_ptrs.dim(), 1);
CHECK_EQ(src_block_indices.dim(), 1);
CHECK_EQ(dst_block_indices.dim(), 1);
CHECK_EQ(cum_sum.dim(), 1);
TORCH_CHECK(key_cache_ptrs.is_contiguous());
TORCH_CHECK(value_cache_ptrs.is_contiguous());
TORCH_CHECK(src_block_indices.is_contiguous());
TORCH_CHECK(dst_block_indices.is_contiguous());
TORCH_CHECK(cum_sum.is_contiguous());
CHECK_EQ(key_cache_ptrs.size(0), value_cache_ptrs.size(0));
CHECK_EQ(src_block_indices.size(0), cum_sum.size(0));
CHECK_GT(numel_per_block, 0);
const at::cuda::OptionalCUDAGuard device_guard(key_cache_ptrs.device());
constexpr int32_t kThreadsPerBlock = 256;
const int32_t num_layers = static_cast<int32_t>(key_cache_ptrs.size(0));
const int32_t num_groups = static_cast<int32_t>(src_block_indices.size(0));
const int32_t num_dst_blocks =
static_cast<int32_t>(dst_block_indices.size(0));
const cudaStream_t stream =
c10::cuda::getCurrentCUDAStream(key_cache_ptrs.get_device());
DISPATCH_FLOATING_TYPES(cache_dtype, "block_copy_kernel", [&] {
constexpr bool kHasVecType = std::is_same_v<scalar_t, float> ||
std::is_same_v<scalar_t, c10::Half> ||
std::is_same_v<scalar_t, c10::BFloat16>;
if constexpr (kHasVecType) {
constexpr int32_t kVecWidth = VecType<scalar_t>::vec_width;
if (numel_per_block % kVecWidth == 0) {
const int64_t tiles_per_block =
ceil_div<int64_t>(numel_per_block / kVecWidth, kThreadsPerBlock);
const dim3 grid(num_layers, num_dst_blocks, tiles_per_block);
block_copy_kernel<scalar_t, true>
<<<grid, kThreadsPerBlock, 0, stream>>>(
key_cache_ptrs.data_ptr<int64_t>(),
value_cache_ptrs.data_ptr<int64_t>(),
src_block_indices.data_ptr<int32_t>(),
dst_block_indices.data_ptr<int32_t>(),
cum_sum.data_ptr<int32_t>(),
num_groups,
numel_per_block);
C10_CUDA_KERNEL_LAUNCH_CHECK();
return;
}
}
const int64_t tiles_per_block =
ceil_div<int64_t>(numel_per_block, kThreadsPerBlock);
const dim3 grid(num_layers, num_dst_blocks, tiles_per_block);
block_copy_kernel<scalar_t, false><<<grid, kThreadsPerBlock, 0, stream>>>(
key_cache_ptrs.data_ptr<int64_t>(),
value_cache_ptrs.data_ptr<int64_t>(),
src_block_indices.data_ptr<int32_t>(),
dst_block_indices.data_ptr<int32_t>(),
cum_sum.data_ptr<int32_t>(),
num_groups,
numel_per_block);
C10_CUDA_KERNEL_LAUNCH_CHECK();
});
}
} // namespace xllm::kernel::cuda