CoolFace
Apppublic

Aluode/PerceptionLabPortable

sourceHugging Faceupdated 9mo agoView on Hugging Face
0likes
cuda_launch.cu155 linesDownload Raw Back to mra
1#include <torch/extension.h>2#include <ATen/ATen.h>3#include "cuda_launch.h"4#include "cuda_kernel.h"5#include <vector>6 7//////////////////////////////////////////////////////////////////////////////////////////////////8//////////////////////////////////////////////////////////////////////////////////////////////////9 10std::vector<at::Tensor> index_max_kernel(11  at::Tensor index_vals,  // [batch_size, 32, num_block]12  at::Tensor indices,     // [batch_size, num_block],13  int A_num_block,14  int B_num_block15) {16  int batch_size = indices.size(0);17  int num_block = indices.size(1);18 19  at::Tensor max_vals = at::zeros({batch_size, A_num_block * 32}, index_vals.options());20  at::Tensor max_vals_scatter = at::zeros({batch_size, 32, num_block}, index_vals.options());21 22  dim3 threads(256);23  dim3 blocks(batch_size);24  int shared_mem = A_num_block * 32 * sizeof(float);25 26  index_max_cuda_kernel<<<blocks, threads, shared_mem>>>(27    index_vals.data_ptr<float>(),28    indices.data_ptr<int>(),29    max_vals.data_ptr<float>(),30    max_vals_scatter.data_ptr<float>(),31    batch_size,32    A_num_block,33    B_num_block,34    num_block35  );36 37  return {max_vals, max_vals_scatter};38}39 40at::Tensor mm_to_sparse_kernel(41  at::Tensor dense_A,  // [batch_size, A_num_block, dim, 32]42  at::Tensor dense_B,  // [batch_size, B_num_block, dim, 32]43  at::Tensor indices   // [batch_size, num_block]44) {45  int batch_size = dense_A.size(0);46  int A_num_block = dense_A.size(1);47  int B_num_block = dense_B.size(1);48  int dim = dense_A.size(2);49  int num_block = indices.size(1);50 51  at::Tensor sparse_C = at::zeros({batch_size, num_block, 32, 32}, dense_A.options());52 53  dim3 threads(64, 4);54  dim3 blocks(num_block / 4, batch_size);55 56  mm_to_sparse_cuda_kernel<<<blocks, threads>>>(57    dense_A.data_ptr<float>(),58    dense_B.data_ptr<float>(),59    indices.data_ptr<int>(),60    sparse_C.data_ptr<float>(),61    batch_size,62    A_num_block,63    B_num_block,64    dim,65    num_block66  );67 68  return sparse_C;69}70 71at::Tensor sparse_dense_mm_kernel(72  at::Tensor sparse_A,  // [batch_size, num_block, 32, 32]73  at::Tensor indices,   // [batch_size, num_block]74  at::Tensor dense_B,   // [batch_size, B_num_block, dim, 32]75  int A_num_block76) {77  int batch_size = sparse_A.size(0);78  int num_block = sparse_A.size(1);79  int B_num_block = dense_B.size(1);80  int dim = dense_B.size(2);81 82  at::Tensor dense_C = at::zeros({batch_size, A_num_block, dim, 32}, dense_B.options());83 84  dim3 threads(128, 2);85  dim3 blocks(num_block / 2, batch_size);86 87  sparse_dense_mm_cuda_kernel<<<blocks, threads>>>(88    sparse_A.data_ptr<float>(),89    indices.data_ptr<int>(),90    dense_B.data_ptr<float>(),91    dense_C.data_ptr<float>(),92    batch_size,93    A_num_block,94    B_num_block,95    dim,96    num_block97  );98 99  return dense_C;100}101 102at::Tensor reduce_sum_kernel(103  at::Tensor sparse_A,  // [batch_size, num_block, 32, 32]104  at::Tensor indices,   // [batch_size, num_block]105  int A_num_block,106  int B_num_block107) {108  int batch_size = sparse_A.size(0);109  int num_block = sparse_A.size(1);110 111  at::Tensor dense_C = at::zeros({batch_size, A_num_block, 32}, sparse_A.options());112 113  dim3 threads(32, 4);114  dim3 blocks(num_block / 4, batch_size);115 116  reduce_sum_cuda_kernel<<<blocks, threads>>>(117    sparse_A.data_ptr<float>(),118    indices.data_ptr<int>(),119    dense_C.data_ptr<float>(),120    batch_size,121    A_num_block,122    B_num_block,123    num_block124  );125 126  return dense_C;127}128 129at::Tensor scatter_kernel(130  at::Tensor dense_A,   // [batch_size, A_num_block, 32]131  at::Tensor indices,   // [batch_size, num_block]132  int B_num_block133) {134  int batch_size = dense_A.size(0);135  int A_num_block = dense_A.size(1);136  int num_block = indices.size(1);137 138  at::Tensor sparse_C = at::zeros({batch_size, num_block, 32, 32}, dense_A.options());139 140  dim3 threads(32, 4);141  dim3 blocks(num_block / 4, batch_size);142 143  scatter_cuda_kernel<<<blocks, threads>>>(144    dense_A.data_ptr<float>(),145    indices.data_ptr<int>(),146    sparse_C.data_ptr<float>(),147    batch_size,148    A_num_block,149    B_num_block,150    num_block151  );152 153  return sparse_C;154}155 
Aluode/PerceptionLabPortable · CoolFace