-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathgemv_optimized.cu
More file actions
36 lines (32 loc) · 1.36 KB
/
Copy pathgemv_optimized.cu
File metadata and controls
36 lines (32 loc) · 1.36 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
#include "gemv.h"
#define FULL_MASK 0xffffffff
constexpr int WARP = 32;
__global__ void gemv_optimized_kernel(const float* d_mat, const float* d_vec, float* d_out_opt, int rows, int cols) {
int row = (blockIdx.x * blockDim.x + threadIdx.x) / WARP;
int lane_id = threadIdx.x % WARP;
const float4* d_mat4 = reinterpret_cast<const float4*>(d_mat);
const float4* d_vec4 = reinterpret_cast<const float4*>(d_vec);
if(cols % 4 == 0 && row < rows) {
float sum = 0.0f;
for(int i = lane_id; i < cols / 4; i += WARP) {
sum += d_mat4[row * (cols / 4) + i].x * d_vec4[i].x +
d_mat4[row * (cols / 4) + i].y * d_vec4[i].y +
d_mat4[row * (cols / 4) + i].z * d_vec4[i].z +
d_mat4[row * (cols / 4) + i].w * d_vec4[i].w;
}
int offset = WARP / 2;
for(int i = 0; i < 5; i++) {
sum += __shfl_down_sync(FULL_MASK, sum, offset);
offset /= 2;
}
if(lane_id == 0) {
d_out_opt[row] = sum;
}
}
}
void run_gemv_optimized(const float* d_mat, const float* d_vec, float* d_out_opt, int rows, int cols) {
int block_size = 256;
int grid_size = (rows * 32 + block_size - 1) / block_size;
gemv_optimized_kernel<<<grid_size, block_size>>>(d_mat, d_vec, d_out_opt, rows, cols);
cudaDeviceSynchronize();
}