1#ifndef KNNCOLLE_VPTREE_HPP
2#define KNNCOLLE_VPTREE_HPP
24#include "sanisizer/sanisizer.hpp"
37inline static constexpr const char* vptree_prebuilt_save_name =
"knncolle::Vptree";
47 std::optional<typename std::mt19937_64::result_type>
seed;
53template<
typename Index_,
typename Data_,
typename Distance_,
class DistanceMetric_>
56template<
typename Index_>
57struct VptreeSearchHistory {
58 VptreeSearchHistory(
bool right, Index_ node) : node(node), right(right) {}
63template<
typename Index_,
typename Data_,
typename Distance_,
class DistanceMetric_>
64class VptreeSearcher final :
public Searcher<Index_, Data_, Distance_> {
66 VptreeSearcher(
const VptreePrebuilt<Index_, Data_, Distance_, DistanceMetric_>& parent) : my_parent(parent) {}
69 const VptreePrebuilt<Index_, Data_, Distance_, DistanceMetric_>& my_parent;
70 NeighborQueue<Index_, Distance_> my_nearest;
71 std::vector<VptreeSearchHistory<Index_> > my_history;
72 std::vector<std::pair<Distance_, Index_> > my_all_neighbors;
75 void search(Index_ i, Index_ k, std::vector<Index_>* output_indices, std::vector<Distance_>* output_distances) {
76 assert(k == 0 || k < my_parent.num_observations());
77 my_nearest.reset(k + 1);
78 auto iptr = my_parent.my_data.data() + sanisizer::product_unsafe<std::size_t>(my_parent.my_new_locations[i], my_parent.my_dim);
79 my_parent.search_nn(iptr, my_nearest, my_history);
80 my_nearest.report(output_indices, output_distances, i);
83 void search(
const Data_* query, Index_ k, std::vector<Index_>* output_indices, std::vector<Distance_>* output_distances) {
84 assert(k <= my_parent.num_observations());
87 if (k == 0 || my_parent.my_nodes.empty()) {
89 output_indices->clear();
91 if (output_distances) {
92 output_distances->clear();
97 my_parent.search_nn(query, my_nearest, my_history);
98 my_nearest.report(output_indices, output_distances);
102 bool can_search_all()
const {
106 Index_ search_all(Index_ i, Distance_ d, std::vector<Index_>* output_indices, std::vector<Distance_>* output_distances) {
107 auto iptr = my_parent.my_data.data() + sanisizer::product_unsafe<std::size_t>(my_parent.my_new_locations[i], my_parent.my_dim);
109 if (!output_indices && !output_distances) {
111 my_parent.template search_all<true>(iptr, d, count, my_history);
115 my_all_neighbors.clear();
116 my_parent.template search_all<false>(iptr, d, my_all_neighbors, my_history);
122 Index_ search_all(
const Data_* query, Distance_ d, std::vector<Index_>* output_indices, std::vector<Distance_>* output_distances) {
123 if (my_parent.my_nodes.empty()) {
124 my_all_neighbors.clear();
129 if (!output_indices && !output_distances) {
131 my_parent.template search_all<true>(query, d, count, my_history);
135 my_all_neighbors.clear();
136 my_parent.template search_all<false>(query, d, my_all_neighbors, my_history);
138 return my_all_neighbors.size();
143template<
typename Index_,
typename Data_,
typename Distance_,
class DistanceMetric_>
144class VptreePrebuilt final :
public Prebuilt<Index_, Data_, Distance_> {
148 std::vector<Data_> my_data;
149 std::shared_ptr<const DistanceMetric_> my_metric;
152 Index_ num_observations()
const {
156 std::size_t num_dimensions()
const {
166 static constexpr Index_ TERMINAL = 0;
170 Distance_ radius = 0;
176 Index_ left = TERMINAL;
179 Index_ right = TERMINAL;
182 std::vector<Node> my_nodes;
184 void build(
const VptreeOptions& options) {
185 typedef std::pair<Distance_, Index_> DataPoint;
186 std::vector<DataPoint> items;
187 items.reserve(my_obs);
188 for (Index_ i = 0; i < my_obs; ++i) {
189 items.emplace_back(0, i);
192 std::mt19937_64 rng([&]() {
193 if (options.seed.has_value()) {
194 return *(options.seed);
201 typedef typename std::mt19937_64::result_type SeedType;
202 const SeedType base = 1234567890, m1 = my_obs, m2 = my_dim;
203 return static_cast<SeedType
>(base * m1 + m2);
207 Index_ lower = 0, upper = my_obs;
210 my_nodes.reserve(my_obs);
211 const auto coords = my_data.data();
213 struct BuildHistory {
214 BuildHistory(Index_ lower, Index_ upper, Index_* right) : right(right), lower(lower), upper(upper) {}
218 std::vector<BuildHistory> history;
221 my_nodes.emplace_back();
222 Node& node = my_nodes.back();
224 const Index_ gap = upper - lower;
227 const auto& leaf = items[lower];
228 node.index = leaf.second;
231 if (history.empty()) {
234 *(history.back().right) = my_nodes.size();
235 lower = history.back().lower;
236 upper = history.back().upper;
244 const Index_ vp = (rng() % gap + lower);
245 std::swap(items[lower], items[vp]);
246 const auto& vantage = items[lower];
247 node.index = vantage.second;
248 const Data_* vantage_ptr = coords + sanisizer::product_unsafe<std::size_t>(vantage.second, my_dim);
252 const Index_ lower_p1 = lower + 1;
253 for (Index_ i = lower_p1 ; i < upper; ++i) {
254 const Data_* loc = coords + sanisizer::product_unsafe<std::size_t>(items[i].second, my_dim);
255 items[i].first = my_metric->raw(my_dim, vantage_ptr, loc);
260 const Index_ median = lower_p1 + (gap - 1)/2;
261 std::nth_element(items.begin() + lower_p1, items.begin() + median, items.begin() + upper);
264 node.radius = my_metric->normalize(items[median].first);
268 history.emplace_back(median, upper, &(node.right));
269 node.left = my_nodes.size();
276 const Index_ median = lower_p1;
277 node.radius = my_metric->normalize(items[median].first);
278 node.right = my_nodes.size();
290 std::vector<Index_> my_new_locations;
293 VptreePrebuilt(std::size_t num_dim, Index_ num_obs, std::vector<Data_> data, std::shared_ptr<const DistanceMetric_> metric,
const VptreeOptions& options) :
296 my_data(std::move(data)),
297 my_metric(std::move(metric))
303 auto used = sanisizer::create<std::vector<char> >(my_obs);
304 auto buffer = sanisizer::create<std::vector<Data_> >(my_dim);
305 sanisizer::resize(my_new_locations, my_obs);
306 auto host = my_data.data();
308 for (Index_ o = 0; o < num_obs; ++o) {
313 auto& current = my_nodes[o];
314 my_new_locations[current.index] = o;
315 if (current.index == o) {
319 auto optr = host + sanisizer::product_unsafe<std::size_t>(o, my_dim);
320 std::copy_n(optr, my_dim, buffer.begin());
321 Index_ replacement = current.index;
324 auto rptr = host + sanisizer::product_unsafe<std::size_t>(replacement, my_dim);
325 std::copy_n(rptr, my_dim, optr);
326 used[replacement] = 1;
328 const auto& next = my_nodes[replacement];
329 my_new_locations[next.index] = replacement;
332 replacement = next.index;
333 }
while (replacement != o);
335 std::copy(buffer.begin(), buffer.end(), optr);
341 static bool can_progress_left(
const Node& node,
const Distance_ dist_to_vp,
const Distance_ threshold) {
342 return node.left != TERMINAL && dist_to_vp - threshold <= node.radius;
345 static bool can_progress_right(
const Node& node,
const Distance_ dist_to_vp,
const Distance_ threshold) {
348 return node.right != TERMINAL && dist_to_vp + threshold >= node.radius;
351 void search_nn(
const Data_* target, NeighborQueue<Index_, Distance_>& nearest, std::vector<VptreeSearchHistory<Index_> >& history)
const {
353 Index_ curnode_offset = 0;
354 Distance_ max_dist = std::numeric_limits<Distance_>::max();
357 auto nptr = my_data.data() + sanisizer::product_unsafe<std::size_t>(curnode_offset, my_dim);
358 const Distance_ dist_to_vp = my_metric->normalize(my_metric->raw(my_dim, nptr, target));
360 const auto& curnode = my_nodes[curnode_offset];
361 if (dist_to_vp <= max_dist) {
362 nearest.add(curnode.index, dist_to_vp);
363 if (nearest.is_full()) {
364 max_dist = nearest.limit();
368 if (dist_to_vp < curnode.radius) {
374 const bool can_left = curnode.left != TERMINAL;
375 const bool can_right = can_progress_right(curnode, dist_to_vp, max_dist);
379 history.emplace_back(
false, curnode_offset);
381 curnode_offset = curnode.left;
383 }
else if (can_right) {
384 curnode_offset = curnode.right;
394 const bool can_right = curnode.right != TERMINAL;
395 const bool can_left = can_progress_left(curnode, dist_to_vp, max_dist);
399 history.emplace_back(
true, curnode_offset);
401 curnode_offset = curnode.right;
411 if (history.empty()) {
415 auto& histinfo = history.back();
416 if (!histinfo.right) {
417 curnode_offset = my_nodes[histinfo.node].right;
419 curnode_offset = my_nodes[histinfo.node].left;
425 template<
bool count_only_,
typename Output_>
426 void search_all(
const Data_* target,
const Distance_ threshold, Output_& all_neighbors, std::vector<VptreeSearchHistory<Index_> >& history)
const {
428 Index_ curnode_offset = 0;
431 auto nptr = my_data.data() + sanisizer::product_unsafe<std::size_t>(curnode_offset, my_dim);
432 const Distance_ dist_to_vp = my_metric->normalize(my_metric->raw(my_dim, nptr, target));
434 const auto& curnode = my_nodes[curnode_offset];
435 if (dist_to_vp <= threshold) {
436 if constexpr(count_only_) {
439 all_neighbors.emplace_back(dist_to_vp, curnode.index);
443 const bool can_left = can_progress_left(curnode, dist_to_vp, threshold);
444 const bool can_right = can_progress_right(curnode, dist_to_vp, threshold);
450 history.emplace_back(
false, curnode_offset);
452 curnode_offset = curnode.left;
454 }
else if (can_right) {
455 curnode_offset = curnode.right;
460 if (history.empty()) {
464 auto& histinfo = history.back();
465 curnode_offset = my_nodes[histinfo.node].right;
470 friend class VptreeSearcher<Index_, Data_, Distance_, DistanceMetric_>;
473 std::unique_ptr<Searcher<Index_, Data_, Distance_> > initialize()
const {
474 return initialize_known();
477 auto initialize_known()
const {
478 return std::make_unique<VptreeSearcher<Index_, Data_, Distance_, DistanceMetric_> >(*this);
482 void save(
const std::filesystem::path& dir)
const {
483 quick_save(dir /
"ALGORITHM", vptree_prebuilt_save_name, std::strlen(vptree_prebuilt_save_name));
484 quick_save(dir /
"DATA", my_data.data(), my_data.size());
487 quick_save(dir /
"NODES", my_nodes.data(), my_nodes.size());
488 quick_save(dir /
"NEW_LOCATIONS", my_new_locations.data(), my_new_locations.size());
490 const auto distdir = dir /
"DISTANCE";
491 std::filesystem::create_directory(distdir);
492 my_metric->save(distdir);
495 VptreePrebuilt(
const std::filesystem::path& dir) {
499 my_data.resize(sanisizer::product<I<
decltype(my_data.size())> >(my_obs, my_dim));
500 quick_load(dir /
"DATA", my_data.data(), my_data.size());
502 sanisizer::resize(my_nodes, my_obs);
503 quick_load(dir /
"NODES", my_nodes.data(), my_nodes.size());
505 sanisizer::resize(my_new_locations, my_obs);
506 quick_load(dir /
"NEW_LOCATIONS", my_new_locations.data(), my_new_locations.size());
508 auto dptr = load_distance_metric_raw<Data_, Distance_>(dir /
"DISTANCE");
509 auto xptr =
dynamic_cast<DistanceMetric_*
>(dptr);
511 throw std::runtime_error(
"cannot cast the loaded distance metric to a DistanceMetric_");
513 my_metric.reset(xptr);
565 class Matrix_ = Matrix<Index_, Data_>,
566 class DistanceMetric_ = DistanceMetric<Data_, Distance_>
574 VptreeBuilder(std::shared_ptr<const DistanceMetric_> metric,
VptreeOptions options) : my_metric(std::move(metric)), my_options(std::move(options)) {}
592 std::shared_ptr<const DistanceMetric_> my_metric;
611 std::size_t ndim = data.num_dimensions();
612 Index_ nobs = data.num_observations();
613 auto work = data.new_known_extractor();
616 std::vector<Data_> store(sanisizer::product<
typename std::vector<Data_>::size_type>(ndim, nobs));
617 for (Index_ o = 0; o < nobs; ++o) {
618 std::copy_n(work->next(), ndim, store.data() + sanisizer::product_unsafe<std::size_t>(o, ndim));
621 return new VptreePrebuilt<Index_, Data_, Distance_, DistanceMetric_>(ndim, nobs, std::move(store), my_metric, my_options);
Interface to build nearest-neighbor indices.
Interface for the input matrix.
Helper class to track nearest neighbors.
Interface for prebuilt nearest-neighbor indices.
Interface to build nearest-neighbor search indices.
Definition Builder.hpp:28
virtual Prebuilt< Index_, Data_, Distance_ > * build_raw(const Matrix_ &data) const =0
Interface for prebuilt nearest-neighbor search indices.
Definition Prebuilt.hpp:29
Perform a nearest neighbor search based on a vantage point (VP) tree.
Definition Vptree.hpp:568
VptreeBuilder(std::shared_ptr< const DistanceMetric_ > metric)
Definition Vptree.hpp:581
auto build_known_unique(const Matrix_ &data) const
Definition Vptree.hpp:627
VptreeOptions & get_options()
Definition Vptree.hpp:587
VptreeBuilder(std::shared_ptr< const DistanceMetric_ > metric, VptreeOptions options)
Definition Vptree.hpp:574
auto build_known_raw(const Matrix_ &data) const
Definition Vptree.hpp:610
auto build_known_shared(const Matrix_ &data) const
Definition Vptree.hpp:634
Classes for distance calculations.
Collection of KNN algorithms.
Definition Bruteforce.hpp:31
void quick_load(const std::filesystem::path &path, Input_ *const contents, const Length_ length)
Definition utils.hpp:57
Index_ count_all_neighbors_without_self(Index_ count)
Definition report_all_neighbors.hpp:23
void quick_save(const std::filesystem::path &path, const Input_ *const contents, const Length_ length)
Definition utils.hpp:33
void report_all_neighbors(std::vector< std::pair< Distance_, Index_ > > &all_neighbors, std::vector< Index_ > *output_indices, std::vector< Distance_ > *output_distances, Index_ self)
Definition report_all_neighbors.hpp:106
Format the output for Searcher::search_all().
Options for VptreeBuilder construction.
Definition Vptree.hpp:42
std::optional< typename std::mt19937_64::result_type > seed
Definition Vptree.hpp:47
Miscellaneous utilities for knncolle