diff --git a/norch/csrc/cuda.cu b/norch/csrc/cuda.cu index 460baec..e79955a 100644 --- a/norch/csrc/cuda.cu +++ b/norch/csrc/cuda.cu @@ -379,6 +379,42 @@ __host__ void matmul_tensor_cuda(Tensor* tensor1, Tensor* tensor2, float* result cudaDeviceSynchronize(); } +__global__ void batched_matmul_tensor_cuda_kernel(float* data1, float* data2, float* result_data, int batch_size, int rows1, int cols1, int cols2) { + int batch = blockIdx.z; + int row = blockIdx.y * blockDim.y + threadIdx.y; + int col = blockIdx.x * blockDim.x + threadIdx.x; + + if (row < rows1 && col < cols2) { + float sum = 0.0f; + for (int k = 0; k < cols1; ++k) { + sum += data1[batch * rows1 * cols1 + row * cols1 + k] * + data2[batch * cols1 * cols2 + k * cols2 + col]; + } + result_data[batch * rows1 * cols2 + row * cols2 + col] = sum; + } +} + +__host__ void batched_matmul_tensor_cuda(Tensor* tensor1, Tensor* tensor2, float* result_data) { + + int batch_size = tensor2->shape[0]; + int rows1 = tensor1->shape[1]; + int cols1 = tensor1->shape[2]; + int cols2 = tensor2->shape[2]; + + dim3 threadsPerBlock(16, 16); + dim3 number_of_blocks((cols2 + threadsPerBlock.x - 1) / threadsPerBlock.x, (rows1 + threadsPerBlock.y - 1) / threadsPerBlock.y); + batched_matmul_tensor_cuda_kernel<<>>(tensor1->data, tensor2->data, result_data, batch_size, rows1, cols1, cols2); + + + cudaError_t error = cudaGetLastError(); + if (error != cudaSuccess) { + printf("CUDA error: %s\n", cudaGetErrorString(error)); + exit(-1); + } + + cudaDeviceSynchronize(); +} + __global__ void tensor_pow_scalar_cuda_kernel(float* data, float exponent, float* result_data, int size) { int i = blockIdx.x * blockDim.x + threadIdx.x; diff --git a/norch/csrc/cuda.h b/norch/csrc/cuda.h index e333ea7..5d5e353 100644 --- a/norch/csrc/cuda.h +++ b/norch/csrc/cuda.h @@ -34,6 +34,9 @@ __global__ void matmul_tensor_cuda_kernel(float* data1, float* data2, float* result_data, int rows1, int cols1, int cols2); __host__ void matmul_tensor_cuda(Tensor* tensor1, Tensor* tensor2, float* result_data); + __global__ void batched_matmul_tensor_cuda_kernel(float* data1, float* data2, float* result_data, int batch_size, int rows1, int cols1, int cols2); + __host__ void batched_matmul_tensor_cuda(Tensor* tensor1, Tensor* tensor2, float* result_data); + __global__ void tensor_pow_scalar_cuda_kernel(float* data, float exponent, float* result_data, int size); __host__ void tensor_pow_scalar_cuda(Tensor* tensor, float exponent, float* result_data); diff --git a/norch/csrc/tensor.cpp b/norch/csrc/tensor.cpp index bdb432a..104994e 100644 --- a/norch/csrc/tensor.cpp +++ b/norch/csrc/tensor.cpp @@ -695,7 +695,7 @@ extern "C" { float* result_data; cudaMalloc((void **)&result_data, size * sizeof(float)); - //batched_matmul_tensor_cuda(tensor1, tensor2, result_data); + batched_matmul_tensor_cuda(tensor1, tensor2, result_data); return create_tensor(result_data, shape, ndim, device); } else { diff --git a/norch/libtensor.so b/norch/libtensor.so index c0f88c2..06f1149 100755 Binary files a/norch/libtensor.so and b/norch/libtensor.so differ