| // /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); |
| } |
|
|