#include <algorithm>
#include <cstdint>
#include <cstring>
#include <limits>
#include <vector>
#include "indices/nonlearned/skdtree/skdtree.hpp"
#include "utils/type.hpp"
namespace {
constexpr std::size_t kDim = 3;
using SkdTree = bench::index::SKDTREE<kDim>;
using Point = point_t<kDim>;
using Points = std::vector<Point>;
struct Handle {
Points points;
std::vector<Point> results_scratch;
SkdTree* index = nullptr;
~Handle() { delete index; }
};
bool global_tree_in_use = false;
}
extern "C" {
void* skdtree_build_f64(const double* points, uint64_t n, uint32_t dim) {
if (dim != kDim || global_tree_in_use) {
return nullptr;
}
auto* handle = new Handle();
handle->points.resize(static_cast<std::size_t>(n));
for (std::size_t i = 0; i < static_cast<std::size_t>(n); ++i) {
for (std::size_t d = 0; d < kDim; ++d) {
const double coordinate = points[i * kDim + d];
if (!(coordinate >= 0.0) || coordinate >= 2.0) {
delete handle;
return nullptr;
}
handle->points[i][d] = coordinate;
}
}
handle->index = new SkdTree(handle->points);
global_tree_in_use = true;
return handle;
}
void skdtree_free_f64(void* raw_handle) {
if (raw_handle == nullptr) {
return;
}
delete static_cast<Handle*>(raw_handle);
global_tree_in_use = false;
}
uint64_t skdtree_nearest_n_f64(void* raw_handle, const double* q, uint64_t k,
double* out_dist2) {
auto* handle = static_cast<Handle*>(raw_handle);
Point query;
for (std::size_t d = 0; d < kDim; ++d) {
query[d] = q[d];
}
Points results = knnQuery<kDim>[treeType](query, static_cast<uint32_t>(k));
const std::size_t produced = std::min<std::size_t>(results.size(), static_cast<std::size_t>(k));
for (std::size_t i = 0; i < produced; ++i) {
double sum = 0.0;
for (std::size_t d = 0; d < kDim; ++d) {
const double delta = results[i][d] - query[d];
sum += delta * delta;
}
out_dist2[i] = sum;
}
std::sort(out_dist2, out_dist2 + produced);
return produced;
}
void skdtree_nearest_one_f64(void* raw_handle, const double* q, double* out_dist2) {
double best = std::numeric_limits<double>::infinity();
if (skdtree_nearest_n_f64(raw_handle, q, 1, &best) == 0) {
best = std::numeric_limits<double>::infinity();
}
*out_dist2 = best;
}
}