diff --git a/build/cuda.cu.o b/build/cuda.cu.o index e6f0304..4ed5a3e 100644 Binary files a/build/cuda.cu.o and b/build/cuda.cu.o differ diff --git a/build/tensor.o b/build/tensor.o index 43889bd..47e5749 100644 Binary files a/build/tensor.o and b/build/tensor.o differ diff --git a/norch/csrc/cuda.cu b/norch/csrc/cuda.cu index c9de11f..6908780 100644 --- a/norch/csrc/cuda.cu +++ b/norch/csrc/cuda.cu @@ -141,40 +141,27 @@ __global__ void sum_tensor_cuda_kernel(float* data, float* result_data, int size } } -__global__ void sum_tensor_axis_cuda_kernel(float* data, float* result_data, int size, int axis_size) { - __shared__ float partial_sum[THREADS_PER_BLOCK]; +__global__ void sum_tensor_cuda_kernel_axis(float* data, float* result_data, int* strides, int target_axis, int inner_size, int outer_size) { + int tid = blockIdx.x * blockDim.x + threadIdx.x; - int tid = threadIdx.x; - int block_offset = blockIdx.x * axis_size; - int i = block_offset + tid; + if (tid < outer_size) { + int outer_index = tid / inner_size; + int inner_index = tid % inner_size; - partial_sum[tid] = 0.0f; - while (i < block_offset + axis_size && i < size) { - partial_sum[tid] += data[i]; - i += blockDim.x; - } + int offset = outer_index * strides[0] + inner_index; - __syncthreads(); - - // Perform block-wise reduction - for (int s = blockDim.x / 2; s > 0; s >>= 1) { - if (tid < s && i < block_offset + axis_size) { - partial_sum[tid] += partial_sum[tid + s]; + for (int i = 0; i < strides[target_axis]; ++i) { + int index = offset + i * strides[target_axis + 1]; + atomicAdd(&result_data[outer_index * inner_size + inner_index], data[index]); } - __syncthreads(); - } - - // Write block sum to global memory - if (tid == 0) { - result_data[blockIdx.x] = partial_sum[0]; } } __host__ void sum_tensor_cuda(Tensor* tensor, float* result_data, int axis) { - cudaMemcpy(result_data, tensor->data, tensor->size * sizeof(float), cudaMemcpyHostToDevice); if (axis == -1) { + cudaMemcpy(result_data, tensor->data, tensor->size * sizeof(float), cudaMemcpyHostToDevice); int num_blocks = (tensor->size + THREADS_PER_BLOCK - 1) / THREADS_PER_BLOCK; @@ -197,19 +184,20 @@ __host__ void sum_tensor_cuda(Tensor* tensor, float* result_data, int axis) { cudaDeviceSynchronize(); } else { - int axis_size = tensor->shape[axis]; - int num_blocks = (tensor->size + THREADS_PER_BLOCK - 1) / THREADS_PER_BLOCK; + int target_axis_stride = tensor->strides[axis]; + int inner_size = tensor->strides[axis + 1]; + int outer_size = tensor->size / target_axis_stride; - // First-level reduction - sum_tensor_axis_cuda_kernel<<>>(result_data, result_data, tensor->size, axis_size); + int* d_strides; + cudaMalloc(&d_strides, (tensor->ndim + 1) * sizeof(int)); + cudaMemcpy(d_strides, tensor->strides, (tensor->ndim + 1) * sizeof(int), cudaMemcpyHostToDevice); - // If necessary, perform multiple levels of reduction - while (num_blocks > 1) { - 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; - } + cudaMemset(result_data, 0, outer_size * sizeof(float)); + + int num_threads = outer_size; + int num_blocks = (num_threads + THREADS_PER_BLOCK - 1) / THREADS_PER_BLOCK; + sum_tensor_cuda_kernel_axis<<>>(tensor->data, result_data, d_strides, axis, inner_size, outer_size); cudaError_t error = cudaGetLastError(); if (error != cudaSuccess) { @@ -218,6 +206,9 @@ __host__ void sum_tensor_cuda(Tensor* tensor, float* result_data, int axis) { } cudaDeviceSynchronize(); + + // Free allocated memory + cudaFree(d_strides); } } diff --git a/norch/csrc/cuda.h b/norch/csrc/cuda.h index 5e180cf..822d661 100644 --- a/norch/csrc/cuda.h +++ b/norch/csrc/cuda.h @@ -14,7 +14,7 @@ __host__ void sub_broadcasted_tensor_cuda(Tensor* tensor1, Tensor* tensor2, float* result_data, int* broadcasted_shape, int broadcasted_size); __global__ void sum_tensor_cuda_kernel(float* data, float* result_data); - __global__ void sum_tensor_axis_cuda_kernel(float* data, float* result_data, int size, int axis_size); + __global__ void sum_tensor_cuda_kernel_axis(float* data, float* result_data, int* strides, int target_axis, int inner_size, int outer_size); __host__ void sum_tensor_cuda(Tensor* tensor, float* result_data, int axis); __global__ void sub_tensor_cuda_kernel(float* data1, float* data2, float* result_data, int size); diff --git a/norch/csrc/tensor.cpp b/norch/csrc/tensor.cpp index 8e5ffa6..221ec62 100644 --- a/norch/csrc/tensor.cpp +++ b/norch/csrc/tensor.cpp @@ -210,7 +210,15 @@ extern "C" { if (strcmp(tensor->device, "cuda") == 0) { float* result_data; - cudaMalloc((void**)&result_data, tensor->size * sizeof(float)); + if (axis == -1) { + cudaMalloc((void**)&result_data, tensor->size * sizeof(float)); + } else { + + int target_axis_stride = tensor->strides[axis]; + int outer_size = tensor->size / target_axis_stride; + + cudaMalloc((void**)&result_data, (outer_size) * sizeof(float)); + } sum_tensor_cuda(tensor, result_data, axis); if (keepdim) { diff --git a/norch/libtensor.so b/norch/libtensor.so index 6c1f5ae..eb4a5ac 100755 Binary files a/norch/libtensor.so and b/norch/libtensor.so differ