fix sum cuda axis atomicAdd

This commit is contained in:
lucasdelimanogueira 2024-05-23 00:30:48 -03:00
parent b8ad9ef282
commit 192a50f773
6 changed files with 34 additions and 35 deletions

Binary file not shown.

Binary file not shown.

View file

@ -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<<<num_blocks, THREADS_PER_BLOCK>>>(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<<<num_blocks_next, THREADS_PER_BLOCK>>>(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<<<num_blocks, THREADS_PER_BLOCK>>>(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);
}
}

View file

@ -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);

View file

@ -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) {

Binary file not shown.