Aluode/PerceptionLabPortable
0
1#include <torch/extension.h>2#include <ATen/ATen.h>3#include "cuda_launch.h"4#include <vector>5 6std::vector<at::Tensor> index_max(7 at::Tensor index_vals,8 at::Tensor indices,9 int A_num_block,10 int B_num_block11) {12 return index_max_kernel(13 index_vals,14 indices,15 A_num_block,16 B_num_block17 );18}19 20at::Tensor mm_to_sparse(21 at::Tensor dense_A,22 at::Tensor dense_B,23 at::Tensor indices24) {25 return mm_to_sparse_kernel(26 dense_A,27 dense_B,28 indices29 );30}31 32at::Tensor sparse_dense_mm(33 at::Tensor sparse_A,34 at::Tensor indices,35 at::Tensor dense_B,36 int A_num_block37) {38 return sparse_dense_mm_kernel(39 sparse_A,40 indices,41 dense_B,42 A_num_block43 );44}45 46at::Tensor reduce_sum(47 at::Tensor sparse_A,48 at::Tensor indices,49 int A_num_block,50 int B_num_block51) {52 return reduce_sum_kernel(53 sparse_A,54 indices,55 A_num_block,56 B_num_block57 );58}59 60at::Tensor scatter(61 at::Tensor dense_A,62 at::Tensor indices,63 int B_num_block64) {65 return scatter_kernel(66 dense_A,67 indices,68 B_num_block69 );70}71 72PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {73 m.def("index_max", &index_max, "index_max (CUDA)");74 m.def("mm_to_sparse", &mm_to_sparse, "mm_to_sparse (CUDA)");75 m.def("sparse_dense_mm", &sparse_dense_mm, "sparse_dense_mm (CUDA)");76 m.def("reduce_sum", &reduce_sum, "reduce_sum (CUDA)");77 m.def("scatter", &scatter, "scatter (CUDA)");78}79 