29 #include <torch/script.h>
33 template <
class TFeat,
class TOut,
class TIndex,
class TKernelIndex>
35 const torch::Tensor& filters,
36 const torch::Tensor& out_importance,
37 const torch::Tensor& inp_features,
38 const torch::Tensor& inp_neighbors_importance_sum,
39 const torch::Tensor& inp_neighbors_row_splits,
40 const torch::Tensor& neighbors_index,
41 const torch::Tensor& neighbors_kernel_index,
42 const torch::Tensor& neighbors_importance,
43 const torch::Tensor& neighbors_row_splits,
44 const torch::Tensor& out_features_gradient,
46 const int64_t max_temp_mem_MB,
47 torch::Tensor& filter_backprop);
49 #ifdef BUILD_CUDA_MODULE
50 template <
class TFeat,
class TOut,
class TIndex,
class TKernelIndex>
51 void SparseConvTransposeBackpropFilterCUDA(
52 const torch::Tensor& filters,
53 const torch::Tensor& out_importance,
54 const torch::Tensor& inp_features,
55 const torch::Tensor& inp_neighbors_importance_sum,
56 const torch::Tensor& inp_neighbors_row_splits,
57 const torch::Tensor& neighbors_index,
58 const torch::Tensor& neighbors_kernel_index,
59 const torch::Tensor& neighbors_importance,
60 const torch::Tensor& neighbors_row_splits,
61 const torch::Tensor& out_features_gradient,
63 const int64_t max_temp_mem_MB,
64 torch::Tensor& filter_backprop);
void SparseConvTransposeBackpropFilterCPU(const torch::Tensor &filters, const torch::Tensor &out_importance, const torch::Tensor &inp_features, const torch::Tensor &inp_neighbors_importance_sum, const torch::Tensor &inp_neighbors_row_splits, const torch::Tensor &neighbors_index, const torch::Tensor &neighbors_kernel_index, const torch::Tensor &neighbors_importance, const torch::Tensor &neighbors_row_splits, const torch::Tensor &out_features_gradient, const bool normalize, const int64_t max_temp_mem_MB, torch::Tensor &filter_backprop)
Definition: SparseConvTransposeBackpropFilterOpKernel.cpp:37