knncolle
Collection of KNN methods in C++
Loading...
Searching...
No Matches
Bruteforce.hpp
Go to the documentation of this file.
1#ifndef KNNCOLLE_BRUTEFORCE_HPP
2#define KNNCOLLE_BRUTEFORCE_HPP
3
4#include "distances.hpp"
5#include "NeighborQueue.hpp"
6#include "Searcher.hpp"
7#include "Builder.hpp"
8#include "Prebuilt.hpp"
9#include "Matrix.hpp"
11#include "utils.hpp"
12
13#include <vector>
14#include <limits>
15#include <memory>
16#include <cstddef>
17#include <string>
18#include <cstring>
19#include <filesystem>
20#include <cassert>
21#include <algorithm>
22
23#include "sanisizer/sanisizer.hpp"
24
31namespace knncolle {
32
36inline static constexpr const char* bruteforce_prebuilt_save_name = "knncolle::Bruteforce";
37
41template<typename Index_, typename Data_, typename Distance_, typename DistanceMetric_>
42class BruteforcePrebuilt;
43
44template<typename Index_, typename Data_, typename Distance_, class DistanceMetric_>
45class BruteforceSearcher final : public Searcher<Index_, Data_, Distance_> {
46public:
47 BruteforceSearcher(const BruteforcePrebuilt<Index_, Data_, Distance_, DistanceMetric_>& parent) : my_parent(parent) {}
48
49private:
50 const BruteforcePrebuilt<Index_, Data_, Distance_, DistanceMetric_>& my_parent;
52 std::vector<std::pair<Distance_, Index_> > my_all_neighbors;
53
54private:
55 void normalize(std::vector<Distance_>* output_distances) const {
56 if (output_distances) {
57 for (auto& d : *output_distances) {
58 d = my_parent.my_metric->normalize(d);
59 }
60 }
61 }
62
63public:
64 void search(Index_ i, Index_ k, std::vector<Index_>* output_indices, std::vector<Distance_>* output_distances) {
65 assert(k == 0 || k < my_parent.num_observations());
66 my_nearest.reset(k + 1); // +1 is safe as k < num_obs.
67 auto ptr = my_parent.my_data.data() + sanisizer::product_unsafe<std::size_t>(i, my_parent.my_dim);
68 my_parent.search(ptr, my_nearest);
69 my_nearest.report(output_indices, output_distances, i);
70 normalize(output_distances);
71 }
72
73 void search(const Data_* query, Index_ k, std::vector<Index_>* output_indices, std::vector<Distance_>* output_distances) {
74 assert(k <= my_parent.num_observations());
75 if (k == 0) { // protect the NeighborQueue from k = 0.
76 if (output_indices) {
77 output_indices->clear();
78 }
79 if (output_distances) {
80 output_distances->clear();
81 }
82 } else {
83 my_nearest.reset(k);
84 my_parent.search(query, my_nearest);
85 my_nearest.report(output_indices, output_distances);
86 normalize(output_distances);
87 }
88 }
89
90 bool can_search_all() const {
91 return true;
92 }
93
94 Index_ search_all(Index_ i, Distance_ d, std::vector<Index_>* output_indices, std::vector<Distance_>* output_distances) {
95 auto ptr = my_parent.my_data.data() + sanisizer::product_unsafe<std::size_t>(i, my_parent.my_dim);
96
97 if (!output_indices && !output_distances) {
98 Index_ count = 0;
99 my_parent.template search_all<true>(ptr, d, count);
101
102 } else {
103 my_all_neighbors.clear();
104 my_parent.template search_all<false>(ptr, d, my_all_neighbors);
105 report_all_neighbors(my_all_neighbors, output_indices, output_distances, i);
106 normalize(output_distances);
107 return count_all_neighbors_without_self(my_all_neighbors.size());
108 }
109 }
110
111 Index_ search_all(const Data_* query, Distance_ d, std::vector<Index_>* output_indices, std::vector<Distance_>* output_distances) {
112 if (!output_indices && !output_distances) {
113 Index_ count = 0;
114 my_parent.template search_all<true>(query, d, count);
115 return count;
116
117 } else {
118 my_all_neighbors.clear();
119 my_parent.template search_all<false>(query, d, my_all_neighbors);
120 report_all_neighbors(my_all_neighbors, output_indices, output_distances);
121 normalize(output_distances);
122 return my_all_neighbors.size();
123 }
124 }
125};
126
127template<typename Index_, typename Data_, typename Distance_, class DistanceMetric_>
128class BruteforcePrebuilt final : public Prebuilt<Index_, Data_, Distance_> {
129private:
130 std::size_t my_dim;
131 Index_ my_obs;
132 std::vector<Data_> my_data;
133 std::shared_ptr<const DistanceMetric_> my_metric;
134
135public:
136 BruteforcePrebuilt(std::size_t num_dim, Index_ num_obs, std::vector<Data_> data, std::shared_ptr<const DistanceMetric_> metric) :
137 my_dim(num_dim), my_obs(num_obs), my_data(std::move(data)), my_metric(std::move(metric)) {}
138
139public:
140 std::size_t num_dimensions() const {
141 return my_dim;
142 }
143
144 Index_ num_observations() const {
145 return my_obs;
146 }
147
148private:
149 void search(const Data_* query, NeighborQueue<Index_, Distance_>& nearest) const {
150 Distance_ threshold_raw = std::numeric_limits<Distance_>::infinity();
151 for (Index_ x = 0; x < my_obs; ++x) {
152 auto dist_raw = my_metric->raw(my_dim, query, my_data.data() + sanisizer::product_unsafe<std::size_t>(x, my_dim));
153 if (dist_raw <= threshold_raw) {
154 nearest.add(x, dist_raw);
155 if (nearest.is_full()) {
156 threshold_raw = nearest.limit();
157 }
158 }
159 }
160 }
161
162 template<bool count_only_, typename Output_>
163 void search_all(const Data_* query, Distance_ threshold, Output_& all_neighbors) const {
164 Distance_ threshold_raw = my_metric->denormalize(threshold);
165 for (Index_ x = 0; x < my_obs; ++x) {
166 Distance_ raw = my_metric->raw(my_dim, query, my_data.data() + sanisizer::product_unsafe<std::size_t>(x, my_dim));
167 if (threshold_raw >= raw) {
168 if constexpr(count_only_) {
169 ++all_neighbors; // expect this to be an integer.
170 } else {
171 all_neighbors.emplace_back(raw, x); // expect this to be a vector of (distance, index) pairs.
172 }
173 }
174 }
175 }
176
177 friend class BruteforceSearcher<Index_, Data_, Distance_, DistanceMetric_>;
178
179public:
180 std::unique_ptr<Searcher<Index_, Data_, Distance_> > initialize() const {
181 return initialize_known();
182 }
183
184 auto initialize_known() const {
185 return std::make_unique<BruteforceSearcher<Index_, Data_, Distance_, DistanceMetric_> >(*this);
186 }
187
188public:
189 void save(const std::filesystem::path& dir) const {
190 quick_save(dir / "ALGORITHM", bruteforce_prebuilt_save_name, std::strlen(bruteforce_prebuilt_save_name));
191 quick_save(dir / "DATA", my_data.data(), my_data.size());
192 quick_save(dir / "NUM_OBS", &my_obs, 1);
193 quick_save(dir / "NUM_DIM", &my_dim, 1);
194
195 const auto distdir = dir / "DISTANCE";
196 std::filesystem::create_directory(distdir);
197 my_metric->save(distdir);
198 }
199
200 BruteforcePrebuilt(const std::filesystem::path& dir) {
201 quick_load(dir / "NUM_OBS", &my_obs, 1);
202 quick_load(dir / "NUM_DIM", &my_dim, 1);
203
204 my_data.resize(sanisizer::product<I<decltype(my_data.size())> >(my_obs, my_dim));
205 quick_load(dir / "DATA", my_data.data(), my_data.size());
206
207 auto dptr = load_distance_metric_raw<Data_, Distance_>(dir / "DISTANCE");
208 auto xptr = dynamic_cast<DistanceMetric_*>(dptr);
209 if (xptr == NULL) {
210 throw std::runtime_error("cannot cast the loaded distance metric to a DistanceMetric_");
211 }
212 my_metric.reset(xptr);
213 }
214};
235template<
236 typename Index_,
237 typename Data_,
238 typename Distance_,
239 class Matrix_ = Matrix<Index_, Data_>,
240 class DistanceMetric_ = DistanceMetric<Data_, Distance_>
241>
242class BruteforceBuilder final : public Builder<Index_, Data_, Distance_, Matrix_> {
243public:
247 BruteforceBuilder(std::shared_ptr<const DistanceMetric_> metric) : my_metric(std::move(metric)) {}
248
249private:
250 std::shared_ptr<const DistanceMetric_> my_metric;
251
252public:
256 Prebuilt<Index_, Data_, Distance_>* build_raw(const Matrix_& data) const {
257 return build_known_raw(data);
258 }
263public:
267 auto build_known_raw(const Matrix_& data) const {
268 std::size_t ndim = data.num_dimensions();
269 const Index_ nobs = data.num_observations();
270 auto work = data.new_known_extractor();
271
272 // We assume that that vector::size_type <= size_t, otherwise data() wouldn't be a contiguous array.
273 std::vector<Data_> store(sanisizer::product<typename std::vector<Data_>::size_type>(ndim, nobs));
274 for (Index_ o = 0; o < nobs; ++o) {
275 std::copy_n(work->next(), ndim, store.data() + sanisizer::product_unsafe<std::size_t>(o, ndim));
276 }
277
278 return new BruteforcePrebuilt<Index_, Data_, Distance_, DistanceMetric_>(ndim, nobs, std::move(store), my_metric);
279 }
280
284 auto build_known_unique(const Matrix_& data) const {
285 return std::unique_ptr<I<decltype(*build_known_raw(data))> >(build_known_raw(data));
286 }
287
291 auto build_known_shared(const Matrix_& data) const {
292 return std::shared_ptr<I<decltype(*build_known_raw(data))> >(build_known_raw(data));
293 }
294};
295
296}
297
298#endif
Interface to build nearest-neighbor indices.
Interface for the input matrix.
Helper class to track nearest neighbors.
Interface for prebuilt nearest-neighbor indices.
Interface for searching nearest-neighbor indices.
Perform a brute-force nearest neighbor search.
Definition Bruteforce.hpp:242
BruteforceBuilder(std::shared_ptr< const DistanceMetric_ > metric)
Definition Bruteforce.hpp:247
auto build_known_shared(const Matrix_ &data) const
Definition Bruteforce.hpp:291
auto build_known_unique(const Matrix_ &data) const
Definition Bruteforce.hpp:284
auto build_known_raw(const Matrix_ &data) const
Definition Bruteforce.hpp:267
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 a distance metric.
Definition distances.hpp:30
Interface for matrix data.
Definition Matrix.hpp:59
Helper class to track nearest neighbors.
Definition NeighborQueue.hpp:36
void report(std::vector< Index_ > *output_indices, std::vector< Distance_ > *output_distances, Index_ self)
Definition NeighborQueue.hpp:178
void add(Index_ i, Distance_ d)
Definition NeighborQueue.hpp:100
Distance_ limit() const
Definition NeighborQueue.hpp:80
bool is_full() const
Definition NeighborQueue.hpp:72
void reset(Index_ k)
Definition NeighborQueue.hpp:52
Interface for prebuilt nearest-neighbor search indices.
Definition Prebuilt.hpp:29
Interface for searching nearest-neighbor search indices.
Definition Searcher.hpp:28
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().
Miscellaneous utilities for knncolle