Aluode/PerceptionLabPortable
0
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 