MCPcopy Create free account
hub / github.com/EnyaHermite/PicassoPlus / BuildNearestNeighbor

Function BuildNearestNeighbor

picasso/point/pi_modules/source/nnquery.cpp:131–171  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

129
130
131std::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
174PYBIND11_MODULE(TORCH_EXTENSION_NAME, m)

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected