Allanatrix's picture
Publish PyC CUDA kernels as source artifacts
bf4d3fe verified
Raw
History Blame Contribute Delete
1.13 kB
// /kernel/matrix_mult.cu
#include <cuda_runtime.h>
__global__ void matrix_mult_kernel(float* a, float* b, float* c, int m, int n, int k) {
int row = blockIdx.y * blockDim.y + threadIdx.y;
int col = blockIdx.x * blockDim.x + threadIdx.x;
if (row < m && col < n) {
float sum = 0.0f;
for (int i = 0; i < k; i++) {
sum += a[row * k + i] * b[i * n + col];
}
c[row * n + col] = sum;
}
}
void matrix_multiply(float* a, float* b, float* c, int m, int n, int k) {
float *d_a, *d_b, *d_c;
cudaMalloc(&d_a, m * k * sizeof(float));
cudaMalloc(&d_b, k * n * sizeof(float));
cudaMalloc(&d_c, m * n * sizeof(float));
cudaMemcpy(d_a, a, m * k * sizeof(float), cudaMemcpyHostToDevice);
cudaMemcpy(d_b, b, k * n * sizeof(float), cudaMemcpyHostToDevice);
dim3 threads(16, 16);
dim3 blocks((n + threads.x - 1) / threads.x, (m + threads.y - 1) / threads.y);
matrix_mult_kernel<<<blocks, threads>>>(d_a, d_b, d_c, m, n, k);
cudaMemcpy(c, d_c, m * n * sizeof(float), cudaMemcpyDeviceToHost);
cudaFree(d_a); cudaFree(d_b); cudaFree(d_c);
}