fix: block_copy.cu — DEVICE_INLINE, CHECK_EQ→TORCH_CHECK, cstdint
Remaining glog dependencies removed: - DEVICE_INLINE macro defined inline - CHECK_EQ(a,b) → TORCH_CHECK(a == b) - CHECK_GT(a,b) → TORCH_CHECK(a > b) - #include <cstdint> for int32_t
This commit is contained in:
@@ -25,6 +25,12 @@ limitations under the License.
|
|||||||
|
|
||||||
#include "device_utils.cuh"
|
#include "device_utils.cuh"
|
||||||
|
|
||||||
|
#ifndef DEVICE_INLINE
|
||||||
|
#define DEVICE_INLINE __device__ __forceinline__
|
||||||
|
#define HOST_DEVICE_INLINE __host__ __device__ __forceinline__
|
||||||
|
#endif
|
||||||
|
|
||||||
|
|
||||||
namespace xllm::kernel::cuda {
|
namespace xllm::kernel::cuda {
|
||||||
namespace {
|
namespace {
|
||||||
|
|
||||||
@@ -137,24 +143,24 @@ void block_copy(torch::Tensor key_cache_ptrs,
|
|||||||
TORCH_CHECK(src_block_indices.is_cuda());
|
TORCH_CHECK(src_block_indices.is_cuda());
|
||||||
TORCH_CHECK(dst_block_indices.is_cuda());
|
TORCH_CHECK(dst_block_indices.is_cuda());
|
||||||
TORCH_CHECK(cum_sum.is_cuda());
|
TORCH_CHECK(cum_sum.is_cuda());
|
||||||
CHECK_EQ(key_cache_ptrs.scalar_type(), torch::kInt64);
|
TORCH_CHECK(key_cache_ptrs.scalar_type() == torch::kInt64);
|
||||||
CHECK_EQ(value_cache_ptrs.scalar_type(), torch::kInt64);
|
TORCH_CHECK(value_cache_ptrs.scalar_type() == torch::kInt64);
|
||||||
CHECK_EQ(src_block_indices.scalar_type(), torch::kInt32);
|
TORCH_CHECK(src_block_indices.scalar_type() == torch::kInt32);
|
||||||
CHECK_EQ(dst_block_indices.scalar_type(), torch::kInt32);
|
TORCH_CHECK(dst_block_indices.scalar_type() == torch::kInt32);
|
||||||
CHECK_EQ(cum_sum.scalar_type(), torch::kInt32);
|
TORCH_CHECK(cum_sum.scalar_type() == torch::kInt32);
|
||||||
CHECK_EQ(key_cache_ptrs.dim(), 1);
|
TORCH_CHECK(key_cache_ptrs.dim() == 1);
|
||||||
CHECK_EQ(value_cache_ptrs.dim(), 1);
|
TORCH_CHECK(value_cache_ptrs.dim() == 1);
|
||||||
CHECK_EQ(src_block_indices.dim(), 1);
|
TORCH_CHECK(src_block_indices.dim() == 1);
|
||||||
CHECK_EQ(dst_block_indices.dim(), 1);
|
TORCH_CHECK(dst_block_indices.dim() == 1);
|
||||||
CHECK_EQ(cum_sum.dim(), 1);
|
TORCH_CHECK(cum_sum.dim() == 1);
|
||||||
TORCH_CHECK(key_cache_ptrs.is_contiguous());
|
TORCH_CHECK(key_cache_ptrs.is_contiguous());
|
||||||
TORCH_CHECK(value_cache_ptrs.is_contiguous());
|
TORCH_CHECK(value_cache_ptrs.is_contiguous());
|
||||||
TORCH_CHECK(src_block_indices.is_contiguous());
|
TORCH_CHECK(src_block_indices.is_contiguous());
|
||||||
TORCH_CHECK(dst_block_indices.is_contiguous());
|
TORCH_CHECK(dst_block_indices.is_contiguous());
|
||||||
TORCH_CHECK(cum_sum.is_contiguous());
|
TORCH_CHECK(cum_sum.is_contiguous());
|
||||||
CHECK_EQ(key_cache_ptrs.size(0), value_cache_ptrs.size(0));
|
TORCH_CHECK(key_cache_ptrs.size(0) == value_cache_ptrs.size(0));
|
||||||
CHECK_EQ(src_block_indices.size(0), cum_sum.size(0));
|
TORCH_CHECK(src_block_indices.size(0) == cum_sum.size(0));
|
||||||
CHECK_GT(numel_per_block, 0);
|
TORCH_CHECK(numel_per_block > 0);
|
||||||
|
|
||||||
const at::cuda::OptionalCUDAGuard device_guard(key_cache_ptrs.device());
|
const at::cuda::OptionalCUDAGuard device_guard(key_cache_ptrs.device());
|
||||||
constexpr int32_t kThreadsPerBlock = 256;
|
constexpr int32_t kThreadsPerBlock = 256;
|
||||||
|
|||||||
Reference in New Issue
Block a user