36inline static constexpr const char* bruteforce_prebuilt_save_name =
"knncolle::Bruteforce";
41template<
typename Index_,
typename Data_,
typename Distance_,
typename DistanceMetric_>
42class BruteforcePrebuilt;
44template<
typename Index_,
typename Data_,
typename Distance_,
class DistanceMetric_>
45class BruteforceSearcher final :
public Searcher<Index_, Data_, Distance_> {
47 BruteforceSearcher(
const BruteforcePrebuilt<Index_, Data_, Distance_, DistanceMetric_>& parent) : my_parent(parent) {}
50 const BruteforcePrebuilt<Index_, Data_, Distance_, DistanceMetric_>& my_parent;
52 std::vector<std::pair<Distance_, Index_> > my_all_neighbors;
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);
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);
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);
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());
77 output_indices->clear();
79 if (output_distances) {
80 output_distances->clear();
84 my_parent.search(query, my_nearest);
85 my_nearest.
report(output_indices, output_distances);
86 normalize(output_distances);
90 bool can_search_all()
const {
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);
97 if (!output_indices && !output_distances) {
99 my_parent.template search_all<true>(ptr, d, count);
103 my_all_neighbors.clear();
104 my_parent.template search_all<false>(ptr, d, my_all_neighbors);
106 normalize(output_distances);
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) {
114 my_parent.template search_all<true>(query, d, count);
118 my_all_neighbors.clear();
119 my_parent.template search_all<false>(query, d, my_all_neighbors);
121 normalize(output_distances);
122 return my_all_neighbors.size();
127template<
typename Index_,
typename Data_,
typename Distance_,
class DistanceMetric_>
128class BruteforcePrebuilt final :
public Prebuilt<Index_, Data_, Distance_> {
132 std::vector<Data_> my_data;
133 std::shared_ptr<const DistanceMetric_> my_metric;
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)) {}
140 std::size_t num_dimensions()
const {
144 Index_ num_observations()
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);
156 threshold_raw = nearest.
limit();
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_) {
171 all_neighbors.emplace_back(raw, x);
177 friend class BruteforceSearcher<Index_, Data_, Distance_, DistanceMetric_>;
180 std::unique_ptr<Searcher<Index_, Data_, Distance_> > initialize()
const {
181 return initialize_known();
184 auto initialize_known()
const {
185 return std::make_unique<BruteforceSearcher<Index_, Data_, Distance_, DistanceMetric_> >(*this);
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());
195 const auto distdir = dir /
"DISTANCE";
196 std::filesystem::create_directory(distdir);
197 my_metric->save(distdir);
200 BruteforcePrebuilt(
const std::filesystem::path& dir) {
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());
207 auto dptr = load_distance_metric_raw<Data_, Distance_>(dir /
"DISTANCE");
208 auto xptr =
dynamic_cast<DistanceMetric_*
>(dptr);
210 throw std::runtime_error(
"cannot cast the loaded distance metric to a DistanceMetric_");
212 my_metric.reset(xptr);
247 BruteforceBuilder(std::shared_ptr<const DistanceMetric_> metric) : my_metric(std::move(metric)) {}
250 std::shared_ptr<const DistanceMetric_> my_metric;
268 std::size_t ndim = data.num_dimensions();
269 const Index_ nobs = data.num_observations();
270 auto work = data.new_known_extractor();
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));
278 return new BruteforcePrebuilt<Index_, Data_, Distance_, DistanceMetric_>(ndim, nobs, std::move(store), my_metric);