diff --git a/norch/csrc/cuda.cu b/norch/csrc/cuda.cu index 45f5f28..bf413e3 100644 --- a/norch/csrc/cuda.cu +++ b/norch/csrc/cuda.cu @@ -4,7 +4,6 @@ #include #define THREADS_PER_BLOCK 128 -#define THREADS_PER_BLOCK_SUM 1024 #define TILE_SIZE 32 #define SHMEM_SIZE THREADS_PER_BLOCK_SUM * sizeof(float) @@ -88,15 +87,15 @@ __global__ void sum_tensor_cuda_kernel(float* data, float* result_data, int size __host__ void sum_tensor_cuda(Tensor* tensor, float* result_data) { cudaMemcpy(result_data, tensor->data, tensor->size * sizeof(float), cudaMemcpyHostToDevice); - int num_blocks = (tensor->size + THREADS_PER_BLOCK_SUM - 1) / THREADS_PER_BLOCK_SUM; + int num_blocks = (tensor->size + THREADS_PER_BLOCK - 1) / THREADS_PER_BLOCK; // First-level reduction - sum_tensor_cuda_kernel<<>>(tensor->data, result_data, tensor->size); + sum_tensor_cuda_kernel<<>>(tensor->data, result_data, tensor->size); // If necessary, perform multiple levels of reduction while (num_blocks > 1) { - int num_blocks_next = (num_blocks + THREADS_PER_BLOCK_SUM - 1) / THREADS_PER_BLOCK_SUM; - sum_tensor_cuda_kernel<<>>(result_data, result_data, num_blocks); + int num_blocks_next = (num_blocks + THREADS_PER_BLOCK - 1) / THREADS_PER_BLOCK; + sum_tensor_cuda_kernel<<>>(result_data, result_data, num_blocks); num_blocks = num_blocks_next; }