|
10 | 10 | #include <limits> |
11 | 11 | #include <cmath> |
12 | 12 | #include <cstddef> |
| 13 | +#include <type_traits> |
13 | 14 |
|
14 | 15 | /** |
15 | | - * @file Kmknn.hpp |
16 | | - * |
| 16 | + * @file knncolle_kmknn.hpp |
17 | 17 | * @brief Implements the k-means with k-nearest neighbors (KMKNN) algorithm. |
18 | 18 | */ |
19 | 19 |
|
@@ -201,35 +201,21 @@ struct KmknnOptions { |
201 | 201 | std::shared_ptr<Refine<Index_, Data_, Distance_, KmeansMatrix_> > refine_algorithm; |
202 | 202 | }; |
203 | 203 |
|
| 204 | +/** |
| 205 | + * @cond |
| 206 | + */ |
204 | 207 | template<typename Index_, typename Data_, typename Distance_, class DistanceMetric_> |
205 | 208 | class KmknnPrebuilt; |
206 | 209 |
|
207 | | -/** |
208 | | - * @brief KMKNN searcher. |
209 | | - * |
210 | | - * Instances of this class are usually constructed using `KmknnPrebuilt::initialize()`. |
211 | | - * |
212 | | - * @tparam Index_ Integer type for the observation indices. |
213 | | - * @tparam Data_ Numeric type for the input and query data. |
214 | | - * @tparam Distance_ Floating-point type for the distances. |
215 | | - * @tparam DistanceMetric_ Class implementing the distance metric calculation. |
216 | | - * This should satisfy the `knncolle::DistanceMetric` interface. |
217 | | - */ |
218 | 210 | template<typename Index_, typename Data_, typename Distance_, class DistanceMetric_ = DistanceMetric<Data_, Distance_> > |
219 | 211 | class KmknnSearcher final : public knncolle::Searcher<Index_, Data_, Distance_> { |
220 | 212 | public: |
221 | | - /** |
222 | | - * @cond |
223 | | - */ |
224 | 213 | KmknnSearcher(const KmknnPrebuilt<Index_, Data_, Distance_, DistanceMetric_>& parent) : my_parent(parent) { |
225 | 214 | my_center_order.reserve(my_parent.my_sizes.size()); |
226 | 215 | if constexpr(needs_conversion) { |
227 | 216 | my_conversion_buffer.resize(my_parent.my_dim); |
228 | 217 | } |
229 | 218 | } |
230 | | - /** |
231 | | - * @endcond |
232 | | - */ |
233 | 219 |
|
234 | 220 | private: |
235 | 221 | const KmknnPrebuilt<Index_, Data_, Distance_, DistanceMetric_>& my_parent; |
@@ -315,17 +301,6 @@ class KmknnSearcher final : public knncolle::Searcher<Index_, Data_, Distance_> |
315 | 301 | } |
316 | 302 | }; |
317 | 303 |
|
318 | | -/** |
319 | | - * @brief Index for a KMKNN search. |
320 | | - * |
321 | | - * Instances of this class are usually constructed using `KmknnBuilder::build_raw()`. |
322 | | - * |
323 | | - * @tparam Index_ Integer type for the indices. |
324 | | - * @tparam Data_ Numeric type for the input and query data. |
325 | | - * @tparam Distance_ Floating-point type for the distances. |
326 | | - * @tparam DistanceMetric_ Class implementing the distance metric calculation. |
327 | | - * This should satisfy the `knncolle::DistanceMetric` interface. |
328 | | - */ |
329 | 304 | template<typename Index_, typename Data_, typename Distance_, class DistanceMetric_ = DistanceMetric<Data_, Distance_> > |
330 | 305 | class KmknnPrebuilt final : public knncolle::Prebuilt<Index_, Data_, Distance_> { |
331 | 306 | private: |
@@ -354,9 +329,6 @@ class KmknnPrebuilt final : public knncolle::Prebuilt<Index_, Data_, Distance_> |
354 | 329 | std::vector<Distance_> my_dist_to_centroid; |
355 | 330 |
|
356 | 331 | public: |
357 | | - /** |
358 | | - * @cond |
359 | | - */ |
360 | 332 | template<class KmeansMatrix_ = kmeans::SimpleMatrix<Index_, Data_> > |
361 | 333 | KmknnPrebuilt( |
362 | 334 | std::size_t num_dim, |
@@ -490,9 +462,6 @@ class KmknnPrebuilt final : public knncolle::Prebuilt<Index_, Data_, Distance_> |
490 | 462 |
|
491 | 463 | return; |
492 | 464 | } |
493 | | - /** |
494 | | - * @endcond |
495 | | - */ |
496 | 465 |
|
497 | 466 | private: |
498 | 467 | void search_nn(const Common<Data_, Distance_>* target, knncolle::NeighborQueue<Index_, Distance_>& nearest, std::vector<std::pair<Distance_, Index_> >& center_order) const { |
@@ -640,13 +609,17 @@ class KmknnPrebuilt final : public knncolle::Prebuilt<Index_, Data_, Distance_> |
640 | 609 | friend class KmknnSearcher<Index_, Data_, Distance_, DistanceMetric_>; |
641 | 610 |
|
642 | 611 | public: |
643 | | - /** |
644 | | - * Creates a `KmknnSearcher` instance. |
645 | | - */ |
646 | 612 | std::unique_ptr<knncolle::Searcher<Index_, Data_, Distance_> > initialize() const { |
| 613 | + return initialize_known(); |
| 614 | + } |
| 615 | + |
| 616 | + auto initialize_known() const { |
647 | 617 | return std::make_unique<KmknnSearcher<Index_, Data_, Distance_, DistanceMetric_> >(*this); |
648 | 618 | } |
649 | 619 | }; |
| 620 | +/** |
| 621 | + * @endcond |
| 622 | + */ |
650 | 623 |
|
651 | 624 | /** |
652 | 625 | * @brief Perform a nearest neighbor search based on k-means clustering. |
@@ -722,23 +695,42 @@ class KmknnBuilder final : public knncolle::Builder<Index_, Data_, Distance_, Ma |
722 | 695 | return my_options; |
723 | 696 | } |
724 | 697 |
|
| 698 | +public: |
| 699 | + knncolle::Prebuilt<Index_, Data_, Distance_>* build_raw(const Matrix_& data) const { |
| 700 | + return build_known_raw(data); |
| 701 | + } |
| 702 | + |
725 | 703 | public: |
726 | 704 | /** |
727 | | - * Creates a `KmknnPrebuilt` instance. |
| 705 | + * Override to assist devirtualization. |
728 | 706 | */ |
729 | | - knncolle::Prebuilt<Index_, Data_, Distance_>* build_raw(const Matrix_& data) const { |
| 707 | + auto build_known_raw(const Matrix_& data) const { |
730 | 708 | std::size_t ndim = data.num_dimensions(); |
731 | 709 | auto nobs = data.num_observations(); |
732 | 710 |
|
733 | 711 | std::vector<Common<Data_, Distance_> > store(ndim * static_cast<std::size_t>(nobs)); // cast to size_t to avoid overflow problems. |
734 | | - auto work = data.new_extractor(); |
| 712 | + auto work = data.new_known_extractor(); |
735 | 713 | for (Index_ o = 0; o < nobs; ++o) { |
736 | 714 | auto ptr = work->next(); |
737 | 715 | std::copy_n(ptr, ndim, store.begin() + static_cast<std::size_t>(o) * ndim); // cast to size_t to avoid overflow. |
738 | 716 | } |
739 | 717 |
|
740 | 718 | return new KmknnPrebuilt<Index_, Data_, Distance_, DistanceMetric_>(ndim, nobs, std::move(store), my_metric, my_options); |
741 | 719 | } |
| 720 | + |
| 721 | + /** |
| 722 | + * Override to assist devirtualization. |
| 723 | + */ |
| 724 | + auto build_known_unique(const Matrix_& data) const { |
| 725 | + return std::unique_ptr<std::remove_reference_t<decltype(*build_known_raw(data))> >(build_known_raw(data)); |
| 726 | + } |
| 727 | + |
| 728 | + /** |
| 729 | + * Override to assist devirtualization. |
| 730 | + */ |
| 731 | + auto build_known_shared(const Matrix_& data) const { |
| 732 | + return std::shared_ptr<std::remove_reference_t<decltype(*build_known_raw(data))> >(build_known_raw(data)); |
| 733 | + } |
742 | 734 | }; |
743 | 735 |
|
744 | 736 | } |
|
0 commit comments