Open3D (C++ API)  0.18.0
SparseConvTransposeBackpropFilterOpKernel.h
Go to the documentation of this file.
1 // ----------------------------------------------------------------------------
2 // - Open3D: www.open3d.org -
3 // ----------------------------------------------------------------------------
4 // Copyright (c) 2018-2023 www.open3d.org
5 // SPDX-License-Identifier: MIT
6 // ----------------------------------------------------------------------------
7 //
8 #pragma once
9 
10 #include <torch/script.h>
11 
12 #include <vector>
13 
14 template <class TFeat, class TOut, class TIndex, class TKernelIndex>
16  const torch::Tensor& filters,
17  const torch::Tensor& out_importance,
18  const torch::Tensor& inp_features,
19  const torch::Tensor& inp_neighbors_importance_sum,
20  const torch::Tensor& inp_neighbors_row_splits,
21  const torch::Tensor& neighbors_index,
22  const torch::Tensor& neighbors_kernel_index,
23  const torch::Tensor& neighbors_importance,
24  const torch::Tensor& neighbors_row_splits,
25  const torch::Tensor& out_features_gradient,
26  const bool normalize,
27  const int64_t max_temp_mem_MB,
28  torch::Tensor& filter_backprop);
29 
30 #ifdef BUILD_CUDA_MODULE
31 template <class TFeat, class TOut, class TIndex, class TKernelIndex>
32 void SparseConvTransposeBackpropFilterCUDA(
33  const torch::Tensor& filters,
34  const torch::Tensor& out_importance,
35  const torch::Tensor& inp_features,
36  const torch::Tensor& inp_neighbors_importance_sum,
37  const torch::Tensor& inp_neighbors_row_splits,
38  const torch::Tensor& neighbors_index,
39  const torch::Tensor& neighbors_kernel_index,
40  const torch::Tensor& neighbors_importance,
41  const torch::Tensor& neighbors_row_splits,
42  const torch::Tensor& out_features_gradient,
43  const bool normalize,
44  const int64_t max_temp_mem_MB,
45  torch::Tensor& filter_backprop);
46 #endif
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:18