From 4db93bfa4c4e67bad7a3f1778018ece2b6fb4fc0 Mon Sep 17 00:00:00 2001 From: Ryan Curtin Date: Thu, 29 Jan 2015 22:34:00 -0500 Subject: [PATCH] Refactor to use new set of rules. How many times will I restart writing this algorithm until I actually get it working well? --- .../methods/kmeans/dtnn_kmeans_impl.hpp | 33 +++++--- src/mlpack/methods/kmeans/dtnn_rules.hpp | 60 ++++++++++++++ src/mlpack/methods/kmeans/dtnn_rules_impl.hpp | 78 +++++++++++++++++++ 3 files changed, 159 insertions(+), 12 deletions(-) create mode 100644 src/mlpack/methods/kmeans/dtnn_rules.hpp create mode 100644 src/mlpack/methods/kmeans/dtnn_rules_impl.hpp diff --git a/src/mlpack/methods/kmeans/dtnn_kmeans_impl.hpp b/src/mlpack/methods/kmeans/dtnn_kmeans_impl.hpp index fde2b44ee6..6112ca72e6 100644 --- a/src/mlpack/methods/kmeans/dtnn_kmeans_impl.hpp +++ b/src/mlpack/methods/kmeans/dtnn_kmeans_impl.hpp @@ -13,6 +13,8 @@ // In case it hasn't been included yet. #include "dtnn_kmeans.hpp" +#include "dtnn_rules.hpp" + namespace mlpack { namespace kmeans { @@ -85,28 +87,35 @@ double DTNNKMeans::Iterate( TreeType* centroidTree = BuildTree( const_cast(centroids), oldFromNewCentroids); - typedef neighbor::NeighborSearch AllkNNType; - AllkNNType allknn(centroidTree, tree, centroids, dataset, false, metric); - + // We won't use the AllkNN class here because we have our own set of rules. // This is a lot of overhead. We don't need the distances. - arma::mat distances; - arma::Mat assignments; - allknn.Search(1, assignments, distances); - distanceCalculations += allknn.BaseCases() + allknn.Scores(); + arma::mat distances(5, dataset.n_cols); + arma::Mat assignments(5, dataset.n_cols); + distances.fill(DBL_MAX); + assignments.fill(size_t(-1)); + typedef DTNNKMeansRules RuleType; + RuleType rules(centroids, dataset, assignments, distances, metric); + + // Now construct the traverser ourselves. + typename TreeType::template DualTreeTraverser traverser(rules); + + traverser.Traverse(*tree, *centroidTree); + + distanceCalculations += rules.BaseCases() + rules.Scores(); // From the assignments, calculate the new centroids and counts. for (size_t i = 0; i < dataset.n_cols; ++i) { if (tree::TreeTraits::RearrangesDataset) { - newCentroids.col(oldFromNewCentroids[assignments[i]]) += dataset.col(i); - ++counts(oldFromNewCentroids[assignments[i]]); + newCentroids.col(oldFromNewCentroids[assignments(0, i)]) += + dataset.col(i); + ++counts(oldFromNewCentroids[assignments(0, i)]); } else { - newCentroids.col(assignments[i]) += dataset.col(i); - ++counts(assignments[i]); + newCentroids.col(assignments(0, i)) += dataset.col(i); + ++counts(assignments(0, i)); } } diff --git a/src/mlpack/methods/kmeans/dtnn_rules.hpp b/src/mlpack/methods/kmeans/dtnn_rules.hpp new file mode 100644 index 0000000000..44647ce158 --- /dev/null +++ b/src/mlpack/methods/kmeans/dtnn_rules.hpp @@ -0,0 +1,60 @@ +/** + * @file dtnn_rules.hpp + * @author Ryan Curtin + * + * A set of rules for the dual-tree k-means algorithm which uses dual-tree + * nearest neighbor search. For the most part we'll call out to + * NeighborSearchRules when we can. + */ +#ifndef __MLPACK_METHODS_KMEANS_DTNN_RULES_HPP +#define __MLPACK_METHODS_KMEANS_DTNN_RULES_HPP + +#include + +namespace mlpack { +namespace kmeans { + +template +class DTNNKMeansRules +{ + public: + DTNNKMeansRules(const arma::mat& centroids, + const arma::mat& dataset, + arma::Mat& neighbors, + arma::mat& distances, + MetricType& metric); + + double BaseCase(const size_t queryIndex, const size_t referenceIndex); + + double Score(const size_t queryIndex, TreeType& referenceNode); + double Score(TreeType& queryNode, TreeType& referenceNode); + double Rescore(const size_t queryIndex, + TreeType& referenceNode, + const double oldScore); + double Rescore(TreeType& queryNode, + TreeType& referenceNode, + const double oldScore); + + typedef neighbor::NeighborSearchTraversalInfo TraversalInfoType; + + size_t Scores() const { return rules.Scores(); } + size_t& Scores() { return rules.Scores(); } + size_t BaseCases() const { return rules.BaseCases(); } + size_t& BaseCases() { return rules.BaseCases(); } + + const TraversalInfoType& TraversalInfo() const + { return rules.TraversalInfo(); } + TraversalInfoType& TraversalInfo() { return rules.TraversalInfo(); } + + private: + + typename neighbor::NeighborSearchRules rules; +}; + +} // namespace kmeans +} // namespace mlpack + +#include "dtnn_rules_impl.hpp" + +#endif diff --git a/src/mlpack/methods/kmeans/dtnn_rules_impl.hpp b/src/mlpack/methods/kmeans/dtnn_rules_impl.hpp new file mode 100644 index 0000000000..c4492c1158 --- /dev/null +++ b/src/mlpack/methods/kmeans/dtnn_rules_impl.hpp @@ -0,0 +1,78 @@ +/** + * @file dtnn_rules_impl.hpp + * @author Ryan Curtin + * + * Implementation of DualTreeKMeansRules. + */ +#ifndef __MLPACK_METHODS_KMEANS_DTNN_RULES_IMPL_HPP +#define __MLPACK_METHODS_KMEANS_DTNN_RULES_IMPL_HPP + +#include "dtnn_rules.hpp" + +namespace mlpack { +namespace kmeans { + +template +DTNNKMeansRules::DTNNKMeansRules( + const arma::mat& centroids, + const arma::mat& dataset, + arma::Mat& neighbors, + arma::mat& distances, + MetricType& metric) : + rules(centroids, dataset, neighbors, distances, metric) +{ + // Nothing to do. +} + +template +inline force_inline double DTNNKMeansRules::BaseCase( + const size_t queryIndex, + const size_t referenceIndex) +{ + // We'll check if the query point has been Hamerly pruned. If so, don't + // continue. + + return rules.BaseCase(queryIndex, referenceIndex); +} + +template +inline double DTNNKMeansRules::Score( + const size_t queryIndex, + TreeType& referenceNode) +{ + return rules.Score(queryIndex, referenceNode); +} + +template +inline double DTNNKMeansRules::Score( + TreeType& queryNode, + TreeType& referenceNode) +{ + // Check if the query node is Hamerly pruned, and if not, then don't continue. + return rules.Score(queryNode, referenceNode); +} + +template +inline double DTNNKMeansRules::Rescore( + const size_t queryIndex, + TreeType& referenceNode, + const double oldScore) +{ + return rules.Rescore(queryIndex, referenceNode, oldScore); +} + +template +inline double DTNNKMeansRules::Rescore( + TreeType& queryNode, + TreeType& referenceNode, + const double oldScore) +{ + // No need to check for a Hamerly prune. Because we've already done that in + // Score(). + return rules.Rescore(queryNode, referenceNode, oldScore); +} + +} // namespace kmeans +} // namespace mlpack + +#endif