replicate/mra
0121
1 2#define WARP_SIZE 323#define FULL_MASK 0xffffffff4#define OPTIMAL_THREADS 2565 6__global__ void index_max_cuda_kernel(7 float *index_vals, // [batch_size, 32, num_block]8 int *indices, // [batch_size, num_block]9 float *max_vals, // [batch_size, A_num_block * 32]10 float *max_vals_scatter, // [batch_size, 32, num_block]11 long batch_size,12 long A_num_block,13 long B_num_block,14 long num_block15);16 17__global__ void mm_to_sparse_cuda_kernel(18 float *dense_A, // [batch_size, A_num_block, dim, 32]19 float *dense_B, // [batch_size, B_num_block, dim, 32]20 int *indices, // [batch_size, num_block]21 float *sparse_C, // [batch_size, num_block, 32, 32]22 long batch_size,23 long A_num_block,24 long B_num_block,25 long dim,26 long num_block27);28 29__global__ void sparse_dense_mm_cuda_kernel(30 float *sparse_A, // [batch_size, num_block, 32, 32]31 int *indices, // [batch_size, num_block]32 float *dense_B, // [batch_size, B_num_block, dim, 32]33 float *dense_C, // [batch_size, A_num_block, dim, 32]34 long batch_size,35 long A_num_block,36 long B_num_block,37 long dim,38 long num_block39);40 41__global__ void reduce_sum_cuda_kernel(42 float *sparse_A, // [batch_size, num_block, 32, 32]43 int *indices, // [batch_size, num_block]44 float *dense_C, // [batch_size, A_num_block, 32]45 long batch_size,46 long A_num_block,47 long B_num_block,48 long num_block49);50 51__global__ void scatter_cuda_kernel(52 float *dense_A, // [batch_size, A_num_block, 32]53 int *indices, // [batch_size, num_block]54 float *sparse_C, // [batch_size, num_block, 32, 32]55 long batch_size,56 long A_num_block,57 long B_num_block,58 long num_block59);60 