| 129 | |
| 130 | |
| 131 | std::vector<torch::Tensor> BuildNearestNeighbor( |
| 132 | torch::Tensor database, // database points: concat_Np * 3 |
| 133 | torch::Tensor query, // query points: concat_Mp * 3 |
| 134 | torch::Tensor nv_database, // batch: each element is the vertex number of a database sample |
| 135 | torch::Tensor nv_query) // batch: each element is the vertex number of a query sample |
| 136 | { |
| 137 | CHECK_INPUT(database,2); |
| 138 | CHECK_INPUT(query,2); |
| 139 | CHECK_INPUT(nv_database,1); |
| 140 | CHECK_INPUT(nv_query,1); |
| 141 | |
| 142 | // get the dims required by computations |
| 143 | int Np = database.size(0); // number of database points |
| 144 | int Mp = query.size(0); // number of query points |
| 145 | int B = nv_database.size(0); // batch size |
| 146 | |
| 147 | TORCH_CHECK(database.dim()==2 && database.size(1)==3, |
| 148 | "Shape of database points requires to be (Np, 3)"); |
| 149 | TORCH_CHECK(query.dim()==2 && query.size(1)==3, |
| 150 | "Shape of query points requires to be (Mp, 3)"); |
| 151 | TORCH_CHECK(nv_database.dim()==1 && nv_query.dim()==1 && nv_database.size(0)==nv_query.size(0), |
| 152 | "Shape of nv_database and nv_query should be identical"); |
| 153 | |
| 154 | const int nn_out = 3; // find 3 nearest neighbors |
| 155 | |
| 156 | // get pointers the input tensors |
| 157 | const float* database_ptr = database.data_ptr<float>(); |
| 158 | const float* query_ptr = query.data_ptr<float>(); |
| 159 | const int* nvDatabase_ptr = nv_database.data_ptr<int32_t>(); |
| 160 | const int* nvQuery_ptr = nv_query.data_ptr<int32_t>(); |
| 161 | |
| 162 | // create output tensors |
| 163 | auto nn_index = torch::zeros({Mp,nn_out}, database.options().dtype(torch::kInt32)); |
| 164 | auto nn_dist = torch::zeros({Mp,nn_out}, database.options()); |
| 165 | int* nnIndex_ptr = nn_index.data_ptr<int32_t>(); |
| 166 | float* nnDist_ptr = nn_dist.data_ptr<float>(); |
| 167 | |
| 168 | buildNearestNeighborLauncher(B, Np, Mp, database_ptr, query_ptr, |
| 169 | nvDatabase_ptr, nvQuery_ptr, nnIndex_ptr, nnDist_ptr); |
| 170 | return {nn_index, nn_dist}; |
| 171 | } |
| 172 | |
| 173 | |
| 174 | PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) |
nothing calls this directly
no outgoing calls
no test coverage detected