diff --git a/src/mlpack/methods/kmeans/CMakeLists.txt b/src/mlpack/methods/kmeans/CMakeLists.txt index 1cc3d9433b..5d488aabef 100644 --- a/src/mlpack/methods/kmeans/CMakeLists.txt +++ b/src/mlpack/methods/kmeans/CMakeLists.txt @@ -2,11 +2,11 @@ # Anything not in this list will not be compiled into MLPACK. set(SOURCES allow_empty_clusters.hpp - dual_tree_kmeans.hpp - dual_tree_kmeans_impl.hpp - dual_tree_kmeans_rules.hpp - dual_tree_kmeans_rules_impl.hpp - dual_tree_kmeans_statistic.hpp + dtnn_kmeans.hpp + dtnn_kmeans_impl.hpp + dtnn_rules.hpp + dtnn_rules_impl.hpp + dtnn_statistic.hpp elkan_kmeans.hpp elkan_kmeans_impl.hpp hamerly_kmeans.hpp diff --git a/src/mlpack/methods/kmeans/dual_tree_kmeans.hpp b/src/mlpack/methods/kmeans/dual_tree_kmeans.hpp deleted file mode 100644 index 914586f24f..0000000000 --- a/src/mlpack/methods/kmeans/dual_tree_kmeans.hpp +++ /dev/null @@ -1,92 +0,0 @@ -/** - * @file dual_tree_kmeans.hpp - * @author Ryan Curtin - * - * A dual-tree algorithm for a single k-means iteration. - */ -#ifndef __MLPACK_METHODS_KMEANS_DUAL_TREE_KMEANS_HPP -#define __MLPACK_METHODS_KMEANS_DUAL_TREE_KMEANS_HPP - -#include "dual_tree_kmeans_statistic.hpp" - -namespace mlpack { -namespace kmeans { - -template< - typename MetricType, - typename MatType, - typename TreeType = tree::BinarySpaceTree, - DualTreeKMeansStatistic> -> -class DualTreeKMeans -{ - public: - DualTreeKMeans(const MatType& dataset, MetricType& metric); - - ~DualTreeKMeans(); - - double Iterate(const arma::mat& centroids, - arma::mat& newCentroids, - arma::Col& counts); - - //! Return the number of distance calculations. - size_t DistanceCalculations() const { return distanceCalculations; } - //! Modify the number of distance calculations. - size_t& DistanceCalculations() { return distanceCalculations; } - - private: - //! The original dataset reference. - const MatType& datasetOrig; - //! The dataset we are using. - const MatType& dataset; - //! A copy of the dataset, if necessary. - MatType datasetCopy; - //! The metric. - MetricType metric; - - //! The tree built on the points. - TreeType* tree; - - arma::vec clusterDistances; - arma::Col assignments; - arma::vec distances; - arma::Col visited; - arma::Col distanceIteration; - arma::vec hamerlyBounds; - - //! The current iteration. - size_t iteration; - - //! Track distance calculations. - size_t distanceCalculations; - - void ClusterTreeUpdate(TreeType* node, - const arma::mat& distances); - - void UpdateOwner(TreeType* node, - const size_t clusters, - const arma::Col& assignments) const; - - void TreeUpdate(TreeType* node, - const size_t clusters, - const arma::vec& clusterDistances, - const arma::Col& assignments, - const arma::mat& oldCentroids, - const arma::mat& dataset, - const std::vector& oldFromNew, - size_t& hamerlyPruned, - size_t& hamerlyPrunedNodes, - size_t& totalNodes, - const arma::mat& interclusterDistances); -}; - -template -using DefaultDualTreeKMeans = DualTreeKMeans; - -} // namespace kmeans -} // namespace mlpack - -// Include implementation. -#include "dual_tree_kmeans_impl.hpp" - -#endif diff --git a/src/mlpack/methods/kmeans/dual_tree_kmeans_impl.hpp b/src/mlpack/methods/kmeans/dual_tree_kmeans_impl.hpp deleted file mode 100644 index be7bfac817..0000000000 --- a/src/mlpack/methods/kmeans/dual_tree_kmeans_impl.hpp +++ /dev/null @@ -1,433 +0,0 @@ -/** - * @file dual_tree_kmeans_impl.hpp - * @author Ryan Curtin - * - * A dual-tree algorithm for a single k-means iteration. - */ -#ifndef __MLPACK_METHODS_KMEANS_DUAL_TREE_KMEANS_IMPL_HPP -#define __MLPACK_METHODS_KMEANS_DUAL_TREE_KMEANS_IMPL_HPP - -// In case it hasn't been included yet. -#include "dual_tree_kmeans.hpp" -#include "dual_tree_kmeans_rules.hpp" - -#include - -namespace mlpack { -namespace kmeans { - -template -DualTreeKMeans::DualTreeKMeans( - const MatType& dataset, - MetricType& metric) : - datasetOrig(dataset), - dataset(tree::TreeTraits::RearrangesDataset ? datasetCopy : - datasetOrig), - metric(metric), - iteration(0), - distanceCalculations(0) -{ - distances.set_size(dataset.n_cols); - distances.fill(DBL_MAX); - assignments.zeros(dataset.n_cols); - visited.zeros(dataset.n_cols); - distanceIteration.zeros(dataset.n_cols); - hamerlyBounds.set_size(dataset.n_cols); - hamerlyBounds.fill(DBL_MAX); - - Timer::Start("tree_building"); - - // Copy the dataset, if necessary. - if (tree::TreeTraits::RearrangesDataset) - datasetCopy = datasetOrig; - - // Now build the tree. We don't need any mappings. - tree = new TreeType(const_cast(this->dataset), 1); - - Timer::Stop("tree_building"); -} - -template -DualTreeKMeans::~DualTreeKMeans() -{ - if (tree) - delete tree; -} - -template -double DualTreeKMeans::Iterate( - const arma::mat& centroids, - arma::mat& newCentroids, - arma::Col& counts) -{ - newCentroids.zeros(centroids.n_rows, centroids.n_cols); - counts.zeros(centroids.n_cols); - if (clusterDistances.n_elem != centroids.n_cols + 1) - { - clusterDistances.set_size(centroids.n_cols + 1); - clusterDistances.fill(DBL_MAX / 2.0); // To prevent overflow. - } - - // Build a tree on the centroids. - arma::mat oldCentroids(centroids); - std::vector oldFromNewCentroids; - TreeType* centroidTree = BuildTree( - const_cast(centroids), oldFromNewCentroids); - - // Now calculate distances between centroids. - neighbor::NeighborSearch - nns(centroidTree, centroids); - arma::mat interclusterDistances; - arma::Mat closestClusters; // We don't actually care about these. - nns.Search(1, closestClusters, interclusterDistances); - - distanceCalculations += nns.BaseCases(); - distanceCalculations += nns.Scores(); - - // Update FirstBound(). - ClusterTreeUpdate(centroidTree, interclusterDistances); - - // Now run the dual-tree algorithm. - typedef DualTreeKMeansRules RulesType; - visited.zeros(dataset.n_cols); - RulesType rules(dataset, centroids, newCentroids, counts, oldFromNewCentroids, - iteration, clusterDistances, distances, assignments, visited, - distanceIteration, hamerlyBounds, interclusterDistances, metric); - - // Use the dual-tree traverser. -//typename TreeType::template DualTreeTraverser traverser(rules); - typename TreeType::template BreadthFirstDualTreeTraverser - traverser(rules); - - tree->Stat().ClustersPruned() = 0; // The constructor sets this to -1. - traverser.Traverse(*centroidTree, *tree); - - distanceCalculations += rules.DistanceCalculations(); - - // Now, calculate how far the clusters moved, after normalizing them. - double residual = 0.0; - clusterDistances.zeros(); - for (size_t c = 0; c < centroids.n_cols; ++c) - { - if (counts[c] == 0) - { - newCentroids.col(c).fill(DBL_MAX); // Should have happened anyway I think. - } - else - { - const size_t oldCluster = oldFromNewCentroids[c]; - newCentroids.col(oldCluster) /= counts(oldCluster); - const double dist = metric.Evaluate(centroids.col(c), - newCentroids.col(oldCluster)); - if (dist > clusterDistances[centroids.n_cols]) - clusterDistances[centroids.n_cols] = dist; - clusterDistances[oldCluster] = dist; - residual += std::pow(dist, 2.0); - } - } - - // Update the tree with the centroid movement information. - size_t hamerlyPruned = 0; - size_t hamerlyPrunedNodes = 0; - size_t totalNodes = 0; - UpdateOwner(tree, centroids.n_cols, assignments); - TreeUpdate(tree, centroids.n_cols, clusterDistances, assignments, - oldCentroids, dataset, oldFromNewCentroids, hamerlyPruned, - hamerlyPrunedNodes, totalNodes, interclusterDistances); - - delete centroidTree; - - ++iteration; - return std::sqrt(residual); -} - -template -void DualTreeKMeans::ClusterTreeUpdate( - TreeType* node, - const arma::mat& distances) -{ - // Just update the first bound, after recursing to the bottom. - double firstBound = 0.0; - for (size_t i = 0; i < node->NumChildren(); ++i) - { - ClusterTreeUpdate(&node->Child(i), distances); - if (node->Child(i).Stat().FirstBound() >= firstBound) - firstBound = node->Child(i).Stat().FirstBound(); - } - for (size_t i = 0; i < node->NumPoints(); ++i) - { - if (distances(0, node->Point(i)) > firstBound) - firstBound = distances(0, node->Point(i)); - } - - node->Stat().FirstBound() = firstBound; -} - -template -bool IsDescendantOf( - const TreeType& potentialParent, - const TreeType& potentialChild) -{ - if (potentialChild.Parent() == &potentialParent) - return true; - else if (&potentialChild == &potentialParent) - return true; - else if (potentialChild.Parent() == NULL) - return false; - else - return IsDescendantOf(potentialParent, *potentialChild.Parent()); -} - -template -void DualTreeKMeans::UpdateOwner( - TreeType* node, - const size_t clusters, - const arma::Col& assignments) const -{ - size_t owner = clusters + 1; - bool same = true; - for (size_t i = 0; i < node->NumChildren(); ++i) - { - UpdateOwner(&node->Child(i), clusters, assignments); - if (owner == clusters + 1) - owner = node->Child(i).Stat().Owner(); - else if (owner != node->Child(i).Stat().Owner()) - { - same = false; - owner = clusters; - break; - } - } - - if (same) - { - for (size_t i = 0; i < node->NumPoints(); ++i) - { - if (owner == clusters + 1) - owner = assignments[node->Point(i)]; - else if (owner != assignments[node->Point(i)]) - { - same = false; - break; - } - } - } - - if (same) - node->Stat().Owner() = owner; - else - node->Stat().Owner() = clusters; -} - -template -void DualTreeKMeans::TreeUpdate( - TreeType* node, - const size_t clusters, - const arma::vec& clusterDistances, - const arma::Col& assignments, - const arma::mat& centroids, - const arma::mat& dataset, - const std::vector& oldFromNew, - size_t& hamerlyPruned, - size_t& hamerlyPrunedNodes, - size_t& totalNodes, - const arma::mat& interclusterDistances) -{ - // This is basically IterationUpdate(), but pulled out to be separate from the - // actual dual-tree algorithm. - const bool prunedLastIteration = node->Stat().HamerlyPruned(); - node->Stat().HamerlyPruned() = false; - ++totalNodes; - -/* - for (size_t i = 0; i < node->NumPoints(); ++i) - { - if (!prunedLastIteration && - distanceIteration[node->Point(i)] < iteration) - Log::Warn << "Point " << node->Point(i) << " was never visited!" -<< " (" << distanceIteration[node->Point(i)] << ", " << prunedLastIteration -<< ")\n"; - if (!prunedLastIteration && - node->Stat().ClustersPruned() + visited[node->Point(i)] < clusters) - Log::Fatal << "Point " << node->Point(i) << " was only visited " << -node->Stat().ClustersPruned() << " + " << visited[node->Point(i)] << -" times!\n"; - } -*/ - - // The easy case: this node had an owner. - if (node->Stat().Owner() < clusters) - { -/* - if (prunedLastIteration && node->Stat().MaxQueryNodeDistance() == DBL_MAX) - Log::Fatal << "r" << node->Begin() << "c" << node->Count() << " was " - << "Hamerly pruned but was not visited!\n"; - // Verify correctness... - for (size_t i = 0; i < node->NumDescendants(); ++i) - { - size_t closest = clusters; - double closestDistance = DBL_MAX; - arma::vec distances(centroids.n_cols); - for (size_t j = 0; j < centroids.n_cols; ++j) - { - const double distance = metric.Evaluate(centroids.col(j), - dataset.col(node->Descendant(i))); - if (distance < closestDistance) - { - closest = j; - closestDistance = distance; - } - distances(j) = distance; - } - - if (closest != node->Stat().Owner()) - { - Log::Warn << distances.t(); - Log::Fatal << "Point " << node->Descendant(i) << " mistakenly assigned " - << "to cluster " << node->Stat().Owner() << ", but should be " << -closest << "! It's part of node r" << node->Begin() << "c" << node->Count() << -".\n"; - } - } -*/ - - // During the last iteration, this node was pruned. - const size_t owner = node->Stat().Owner(); - if (node->Stat().MaxQueryNodeDistance() != DBL_MAX) - node->Stat().MaxQueryNodeDistance() += clusterDistances[owner]; - if (node->Stat().MinQueryNodeDistance() != DBL_MAX) - node->Stat().MinQueryNodeDistance() += clusterDistances[owner]; - - if (prunedLastIteration) - { - // Can we continue being Hamerly pruned? If not, we'll have to update the - // bound next iteration. - if (node->Stat().MaxQueryNodeDistance() < - node->Stat().LastSecondClosestBound() - clusterDistances[clusters]) - { - node->Stat().HamerlyPruned() = true; - if (!node->Parent()->Stat().HamerlyPruned()) - hamerlyPruned += node->NumDescendants(); - } - else if (node->Stat().MaxQueryNodeDistance() < 0.5 * - interclusterDistances(0, owner)) - { -// Log::Warn << "Secondary Elkan prune! r" << node->Begin() << "c" << -//node->Count() << ".\n"; - node->Stat().HamerlyPruned() = true; - if (!node->Parent()->Stat().HamerlyPruned()) - hamerlyPruned += node->NumDescendants(); - } - } - else - { - // Now we check for a Hamerly prune. We know that we have an accurate - // second bound since nothing can be pruned. - if (node->Stat().MaxQueryNodeDistance() /* already adjusted */ < - node->Stat().SecondMinQueryNodeDistance() - clusterDistances[clusters]) - { - node->Stat().HamerlyPruned() = true; - if (!node->Parent()->Stat().HamerlyPruned()) - hamerlyPruned += node->NumDescendants(); - ++hamerlyPrunedNodes; - } - } - - if (!node->Stat().HamerlyPruned()) - { - if (node->Parent() != NULL && node->Parent()->Stat().HamerlyPruned()) - { - node->Stat().HamerlyPruned() = true; - node->Stat().MinQueryNodeDistance() = DBL_MAX; - } - else - { - if (node->Stat().SecondMaxQueryNodeDistance() != DBL_MAX) - node->Stat().SecondMaxQueryNodeDistance() += clusterDistances[clusters]; - if (node->Stat().SecondMinQueryNodeDistance() != DBL_MAX) - node->Stat().SecondMinQueryNodeDistance() += clusterDistances[clusters]; - } - } - else - node->Stat().MinQueryNodeDistance() = DBL_MAX; - } - else - { - // This node did not have a single owner, but did have a closest query - // node. So we will simply loosen that bound. The loosening here is too - // loose; TODO: tighten to the max cluster movement in the closest query - // node. - if (node->Stat().MaxQueryNodeDistance() != DBL_MAX) - node->Stat().MaxQueryNodeDistance() += clusterDistances[clusters]; - if (node->Stat().MinQueryNodeDistance() != DBL_MAX) - node->Stat().MinQueryNodeDistance() += clusterDistances[clusters]; - if (node->Stat().SecondMaxQueryNodeDistance() != DBL_MAX) - node->Stat().SecondMaxQueryNodeDistance() += clusterDistances[clusters]; - if (node->Stat().SecondMinQueryNodeDistance() != DBL_MAX) - node->Stat().SecondMinQueryNodeDistance() += clusterDistances[clusters]; - - // Since the node didn't have an owner, it can't be Hamerly pruned. - node->Stat().HamerlyPruned() = false; - node->Stat().Owner() = centroids.n_cols; - } - - bool allPruned = true; - size_t owner = clusters; - for (size_t i = 0; i < node->NumChildren(); ++i) - { - TreeUpdate(&node->Child(i), clusters, clusterDistances, assignments, - centroids, dataset, oldFromNew, hamerlyPruned, hamerlyPrunedNodes, - totalNodes, interclusterDistances); - if (!node->Child(i).Stat().HamerlyPruned()) - allPruned = false; - else if (owner == clusters) - owner = node->Child(i).Stat().Owner(); - else if (owner < clusters && owner != node->Child(i).Stat().Owner()) - owner = clusters + 1; - } - - if (node->NumChildren() == 0 && !node->Stat().HamerlyPruned()) - allPruned = false; - - if (allPruned && owner < clusters && !node->Stat().HamerlyPruned()) - { - node->Stat().MinQueryNodeDistance() = DBL_MAX; - node->Stat().HamerlyPruned() = true; - hamerlyPrunedNodes++; - } - - node->Stat().Iteration() = iteration; - node->Stat().ClustersPruned() = (node->Parent() == NULL) ? 0 : -1; - // We have to set the closest query node to NULL because the cluster tree will - // be rebuilt. -// node->Stat().ClosestQueryNode() = NULL; - - if (prunedLastIteration) - node->Stat().LastSecondClosestBound() -= clusterDistances[clusters]; - else - node->Stat().LastSecondClosestBound() = - node->Stat().SecondMinQueryNodeDistance() - clusterDistances[clusters]; -// node->Stat().MinQueryNodeDistance() = DBL_MAX; - node->Stat().MinQueryNodeDistance() = DBL_MAX; - node->Stat().SecondMinQueryNodeDistance() = DBL_MAX; - if (prunedLastIteration && !node->Stat().HamerlyPruned()) - { - node->Stat().MaxQueryNodeDistance() = DBL_MAX; - node->Stat().SecondMaxQueryNodeDistance() = DBL_MAX; - } - // This should change later, but I'm not yet sure how to do it. -// node->Stat().SecondClosestBound() = DBL_MAX; -// node->Stat().SecondClosestQueryNode() = NULL; - - if (node->Parent() == NULL) - { - Log::Info << "Total Hamerly pruned points: " << hamerlyPruned << ".\n"; - Log::Info << "Total pruned Hamerly nodes: " << hamerlyPrunedNodes << ".\n"; - Log::Info << "Total nodes in tree: " << totalNodes << ".\n"; - } -} - -} // namespace kmeans -} // namespace mlpack - -#endif diff --git a/src/mlpack/methods/kmeans/dual_tree_kmeans_rules.hpp b/src/mlpack/methods/kmeans/dual_tree_kmeans_rules.hpp deleted file mode 100644 index 4c960bea01..0000000000 --- a/src/mlpack/methods/kmeans/dual_tree_kmeans_rules.hpp +++ /dev/null @@ -1,83 +0,0 @@ -/** - * @file dual_tree_kmeans_rules.hpp - * @author Ryan Curtin - * - * A set of tree traversal rules for dual-tree k-means clustering. - */ -#ifndef __MLPACK_METHODS_KMEANS_DUAL_TREE_KMEANS_RULES_HPP -#define __MLPACK_METHODS_KMEANS_DUAL_TREE_KMEANS_RULES_HPP - -namespace mlpack { -namespace kmeans { - -template -class DualTreeKMeansRules -{ - public: - DualTreeKMeansRules(const typename TreeType::Mat& dataset, - const arma::mat& centroids, - arma::mat& newCentroids, - arma::Col& counts, - const std::vector& mappings, - const size_t iteration, - const arma::vec& clusterDistances, - arma::vec& distances, - arma::Col& assignments, - arma::Col& visited, - arma::Col& distanceIteration, - arma::vec& hamerlyBounds, - const arma::mat& interclusterDistances, - 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) const; - - double Rescore(TreeType& queryNode, - TreeType& referenceNode, - const double oldScore) const; - - size_t DistanceCalculations() const { return distanceCalculations; } - size_t& DistanceCalculations() { return distanceCalculations; } - - typedef neighbor::NeighborSearchTraversalInfo TraversalInfoType; - - const TraversalInfoType& TraversalInfo() const { return traversalInfo; } - TraversalInfoType& TraversalInfo() { return traversalInfo; } - - private: - const typename TreeType::Mat& dataset; - const arma::mat& centroids; - arma::mat& newCentroids; - arma::Col& counts; - const std::vector& mappings; - const size_t iteration; - const arma::vec& clusterDistances; - arma::vec& distances; - arma::Col& assignments; - arma::Col& visited; - arma::Col& distanceIteration; - arma::vec& hamerlyBounds; - const arma::mat& interclusterDistances; - MetricType& metric; - - size_t distanceCalculations; - - TraversalInfoType traversalInfo; - - bool IsDescendantOf(const TreeType& potentialParent, const TreeType& - potentialChild) const; -}; - -} // namespace kmeans -} // namespace mlpack - -#include "dual_tree_kmeans_rules_impl.hpp" - -#endif diff --git a/src/mlpack/methods/kmeans/dual_tree_kmeans_rules_impl.hpp b/src/mlpack/methods/kmeans/dual_tree_kmeans_rules_impl.hpp deleted file mode 100644 index bb2003837c..0000000000 --- a/src/mlpack/methods/kmeans/dual_tree_kmeans_rules_impl.hpp +++ /dev/null @@ -1,233 +0,0 @@ -/** - * @file dual_tree_kmeans_rules_impl.hpp - * @author Ryan Curtin - * - * A set of tree traversal rules for dual-tree k-means clustering. - */ -#ifndef __MLPACK_METHODS_KMEANS_DUAL_TREE_KMEANS_RULES_IMPL_HPP -#define __MLPACK_METHODS_KMEANS_DUAL_TREE_KMEANS_RULES_IMPL_HPP - -// In case it hasn't been included yet. -#include "dual_tree_kmeans_rules.hpp" - -namespace mlpack { -namespace kmeans { - -template -DualTreeKMeansRules::DualTreeKMeansRules( - const typename TreeType::Mat& dataset, - const arma::mat& centroids, - arma::mat& newCentroids, - arma::Col& counts, - const std::vector& mappings, - const size_t iteration, - const arma::vec& clusterDistances, - arma::vec& distances, - arma::Col& assignments, - arma::Col& visited, - arma::Col& distanceIteration, - arma::vec& hamerlyBounds, - const arma::mat& interclusterDistances, - MetricType& metric) : - dataset(dataset), - centroids(centroids), - newCentroids(newCentroids), - counts(counts), - mappings(mappings), - iteration(iteration), - clusterDistances(clusterDistances), - distances(distances), - assignments(assignments), - visited(visited), - distanceIteration(distanceIteration), - hamerlyBounds(hamerlyBounds), - interclusterDistances(interclusterDistances), - metric(metric), - distanceCalculations(0) -{ } - -template -inline force_inline double DualTreeKMeansRules::BaseCase( - const size_t queryIndex, - const size_t referenceIndex) -{ - // Collect the number of clusters that have been pruned during the traversal. - // The ternary operator may not be necessary. - const size_t traversalPruned = (traversalInfo.LastReferenceNode() != NULL) ? - traversalInfo.LastReferenceNode()->Stat().ClustersPruned() : 0; - - // It's possible that the reference node has been pruned before we got to the - // base case. In that case, don't do the base case, and just return. - if (traversalInfo.LastReferenceNode()->Stat().ClustersPruned() + - visited[referenceIndex] == centroids.n_cols) - return 0.0; - - ++distanceCalculations; - - const double distance = metric.Evaluate(centroids.col(queryIndex), - dataset.col(referenceIndex)); - - // Iteration change check. - if (distanceIteration[referenceIndex] < iteration) - { - distanceIteration[referenceIndex] = iteration; - distances[referenceIndex] = distance; - assignments[referenceIndex] = mappings[queryIndex]; - hamerlyBounds[referenceIndex] = DBL_MAX; // Not sure about this one. - } - else if (distance < distances[referenceIndex]) - { - distances[referenceIndex] = distance; - assignments[referenceIndex] = mappings[queryIndex]; - } - else if (distance < hamerlyBounds[referenceIndex]) - { - hamerlyBounds[referenceIndex] = distance; // Not yet done. - } - - ++visited[referenceIndex]; - - if (visited[referenceIndex] + traversalPruned == centroids.n_cols) - { - newCentroids.col(assignments[referenceIndex]) += - dataset.col(referenceIndex); - ++counts(assignments[referenceIndex]); -// Log::Warn << "Commit base case " << referenceIndex << ".\n"; - } - - return distance; -} - -template -double DualTreeKMeansRules::Score( - const size_t /* queryIndex */, - TreeType& /* referenceNode */) -{ - // No pruning here, for now. - return 0.0; -} - -template -double DualTreeKMeansRules::Score( - TreeType& queryNode, - TreeType& referenceNode) -{ -// if (referenceNode.Begin() == 33313 || referenceNode.Begin() == 37121 || -// if (referenceNode.Begin() == 37447) -// Log::Warn << "Visit r" << referenceNode.Begin() << "c" << -//referenceNode.Count() << ", q" << queryNode.Begin() << "c" << queryNode.Count() -//<< ":\n" << referenceNode.Stat(); - - // This won't happen with the root since it is explicitly set to 0. - if (referenceNode.Stat().ClustersPruned() == size_t(-1)) - referenceNode.Stat().ClustersPruned() = - referenceNode.Parent()->Stat().ClustersPruned(); - - if (referenceNode.Stat().HamerlyPruned()) - { - // Add to centroids if necessary. - if (referenceNode.Stat().MinQueryNodeDistance() == DBL_MAX /* hack */) - { - newCentroids.col(referenceNode.Stat().Owner()) += - referenceNode.NumDescendants() * referenceNode.Stat().Centroid(); - counts(referenceNode.Stat().Owner()) += referenceNode.NumDescendants(); - referenceNode.Stat().MinQueryNodeDistance() = 0.0; - } - return DBL_MAX; // No need to go further. - } - - traversalInfo.LastReferenceNode() = &referenceNode; - - // Calculate distance to node. - // This costs about the same (in terms of runtime) as a single MinDistance() - // call, so there only need to add one distance computation. - const math::Range distances = referenceNode.RangeDistance(&queryNode); - ++distanceCalculations; - - // Is this closer than the current best query node? - if (distances.Lo() < referenceNode.Stat().MinQueryNodeDistance()) - { - // This is the new closest node. - if (queryNode.NumDescendants() >= 2) - { - referenceNode.Stat().SecondMinQueryNodeDistance() = distances.Lo(); - referenceNode.Stat().SecondMaxQueryNodeDistance() = distances.Hi(); - } - else - { - referenceNode.Stat().SecondMinQueryNodeDistance() = - referenceNode.Stat().MinQueryNodeDistance(); - referenceNode.Stat().SecondMaxQueryNodeDistance() = - referenceNode.Stat().MaxQueryNodeDistance(); - } - referenceNode.Stat().MinQueryNodeDistance() = distances.Lo(); - referenceNode.Stat().MaxQueryNodeDistance() = distances.Hi(); - } - else if (distances.Lo() < referenceNode.Stat().SecondMinQueryNodeDistance()) - { - // This is the new second closest node. - referenceNode.Stat().SecondMinQueryNodeDistance() = distances.Lo(); - referenceNode.Stat().SecondMaxQueryNodeDistance() = distances.Hi(); - } - else if (distances.Lo() > referenceNode.Stat().SecondMaxQueryNodeDistance()) - { -// if (referenceNode.Begin() == 37447) -// Log::Warn << "Pelleg-Moore pruned.\n"; - referenceNode.Stat().ClustersPruned() += queryNode.NumDescendants(); - - // Is everything pruned? Then commit the points. - if (referenceNode.Stat().ClustersPruned() + - visited[referenceNode.Descendant(0)] == centroids.n_cols) - { -// Log::Warn << "Commit points in r" << referenceNode.Begin() << "c" << -//referenceNode.Count() << ".\n"; - for (size_t i = 0; i < referenceNode.NumDescendants(); ++i) - { - const size_t index = referenceNode.Descendant(i); - const size_t cluster = assignments[index]; - referenceNode.Stat().Owner() = cluster; - newCentroids.col(cluster) += dataset.col(index); - ++counts(cluster); - } - } - return DBL_MAX; - } - - return distances.Lo(); // No pruning allowed at this time. -} - -template -double DualTreeKMeansRules::Rescore( - const size_t /* queryIndex */, - TreeType& /* referenceNode */, - const double oldScore) const -{ - return oldScore; -} - -template -double DualTreeKMeansRules::Rescore( - TreeType& /* queryNode */, - TreeType& /* referenceNode */, - const double oldScore) const -{ - return oldScore; -} - -template -bool DualTreeKMeansRules::IsDescendantOf( - const TreeType& potentialParent, - const TreeType& potentialChild) const -{ - if (potentialChild.Parent() == &potentialParent) - return true; - else if (potentialChild.Parent() == NULL) - return false; - else - return IsDescendantOf(potentialParent, *potentialChild.Parent()); -} - -} // namespace kmeans -} // namespace mlpack - -#endif diff --git a/src/mlpack/methods/kmeans/dual_tree_kmeans_statistic.hpp b/src/mlpack/methods/kmeans/dual_tree_kmeans_statistic.hpp deleted file mode 100644 index 01fa6bc33c..0000000000 --- a/src/mlpack/methods/kmeans/dual_tree_kmeans_statistic.hpp +++ /dev/null @@ -1,190 +0,0 @@ -/** - * @file dual_tree_kmeans_statistic.hpp - * @author Ryan Curtin - * - * Statistic for dual-tree k-means traversal. - */ -#ifndef __MLPACK_METHODS_KMEANS_DUAL_TREE_KMEANS_STATISTIC_HPP -#define __MLPACK_METHODS_KMEANS_DUAL_TREE_KMEANS_STATISTIC_HPP - -namespace mlpack { -namespace kmeans { - -class DualTreeKMeansStatistic -{ - public: - DualTreeKMeansStatistic() { /* Nothing to do. */ } - - template - DualTreeKMeansStatistic(TreeType& node) : -// closestQueryNode(NULL), -// secondClosestQueryNode(NULL), - minQueryNodeDistance(DBL_MAX), - maxQueryNodeDistance(DBL_MAX), - secondMinQueryNodeDistance(DBL_MAX), - secondMaxQueryNodeDistance(DBL_MAX), - lastSecondClosestBound(DBL_MAX), - hamerlyPruned(false), - clustersPruned(size_t(-1)), - iteration(size_t() - 1), - owner(size_t(-1)), - firstBound(DBL_MAX), - secondBound(DBL_MAX), - bound(DBL_MAX), - lastDistanceNode(NULL), - lastDistance(0.0) - { - // Empirically calculate the centroid. - centroid.zeros(node.Dataset().n_rows); - for (size_t i = 0; i < node.NumPoints(); ++i) - centroid += node.Dataset().col(node.Point(i)); - - for (size_t i = 0; i < node.NumChildren(); ++i) - centroid += node.Child(i).NumDescendants() * - node.Child(i).Stat().Centroid(); - - centroid /= node.NumDescendants(); - } - - //! Return the centroid. - const arma::vec& Centroid() const { return centroid; } - //! Modify the centroid. - arma::vec& Centroid() { return centroid; } - - //! Get the current closest query node. -// void* ClosestQueryNode() const { return closestQueryNode; } - //! Modify the current closest query node. -// void*& ClosestQueryNode() { return closestQueryNode; } - - //! Get the second closest query node. -// void* SecondClosestQueryNode() const { return secondClosestQueryNode; } - //! Modify the second closest query node. -// void*& SecondClosestQueryNode() { return secondClosestQueryNode; } - - //! Get the minimum distance to the closest query node. - double MinQueryNodeDistance() const { return minQueryNodeDistance; } - //! Modify the minimum distance to the closest query node. - double& MinQueryNodeDistance() { return minQueryNodeDistance; } - - //! Get the maximum distance to the closest query node. - double MaxQueryNodeDistance() const { return maxQueryNodeDistance; } - //! Modify the maximum distance to the closest query node. - double& MaxQueryNodeDistance() { return maxQueryNodeDistance; } - - //! Get the minimum distance to the second closest query node. - double SecondMinQueryNodeDistance() const - { return secondMinQueryNodeDistance; } - //! Modify the minimum distance to the second closest query node. - double& SecondMinQueryNodeDistance() { return secondMinQueryNodeDistance; } - - //! Get the maximum distance to the second closest query node. - double SecondMaxQueryNodeDistance() const - { return secondMaxQueryNodeDistance; } - //! Modify the maximum distance to the second closest query node. - double& SecondMaxQueryNodeDistance() { return secondMaxQueryNodeDistance; } - - //! Get last iteration's second closest bound. - double LastSecondClosestBound() const { return lastSecondClosestBound; } - //! Modify last iteration's second closest bound. - double& LastSecondClosestBound() { return lastSecondClosestBound; } - - //! Get whether or not this node is Hamerly pruned this iteration. - bool HamerlyPruned() const { return hamerlyPruned; } - //! Modify whether or not this node is Hamerly pruned this iteration. - bool& HamerlyPruned() { return hamerlyPruned; } - - //! Get the number of clusters that have been pruned during this iteration. - size_t ClustersPruned() const { return clustersPruned; } - //! Modify the number of clusters that have been pruned during this iteration. - size_t& ClustersPruned() { return clustersPruned; } - - //! Get the current iteration. - size_t Iteration() const { return iteration; } - //! Modify the current iteration. - size_t& Iteration() { return iteration; } - - //! Get the current owner (if any) of these reference points. - size_t Owner() const { return owner; } - //! Modify the current owner (if any) of these reference points. - size_t& Owner() { return owner; } - - // For nearest neighbor search. - - //! Get the first bound. - double FirstBound() const { return firstBound; } - //! Modify the first bound. - double& FirstBound() { return firstBound; } - //! Get the second bound. - double SecondBound() const { return secondBound; } - //! Modify the second bound. - double& SecondBound() { return secondBound; } - //! Get the overall bound. - double Bound() const { return bound; } - //! Modify the overall bound. - double& Bound() { return bound; } - //! Get the last distance evaluation node. - void* LastDistanceNode() const { return lastDistanceNode; } - //! Modify the last distance evaluation node. - void*& LastDistanceNode() { return lastDistanceNode; } - //! Get the last distance calculation. - double LastDistance() const { return lastDistance; } - //! Modify the last distance calculation. - double& LastDistance() { return lastDistance; } - - std::string ToString() const - { - std::ostringstream convert; - convert << "DualTreeKMeansStatistic [" << this << "]" << std::endl; - convert << " minQueryNodeDistance: " << minQueryNodeDistance << ".\n"; - convert << " maxQueryNodeDistance: " << maxQueryNodeDistance << ".\n"; - convert << " secondMinQueryNodeDistance: " << secondMinQueryNodeDistance << ".\n"; - convert << " secondMaxQueryNodeDistance: " << secondMaxQueryNodeDistance << ".\n"; - convert << " hamerlyPruned: " << hamerlyPruned << ".\n"; - convert << " lastSecondClosestBound: " << lastSecondClosestBound << ".\n"; - convert << " clustersPruned: " << clustersPruned << ".\n"; - return convert.str(); - } - - private: - //! The empirically calculated centroid of the node. - arma::vec centroid; - - //! The current closest query node to this reference node. -// void* closestQueryNode; - //! The second closest query node. -// void* secondClosestQueryNode; - //! The minimum distance to the closest query node. - double minQueryNodeDistance; - //! The maximum distance to the closest query node. - double maxQueryNodeDistance; - //! The minimum distance to the second closest query node. - double secondMinQueryNodeDistance; - //! The maximum distance to the second closest query node. - double secondMaxQueryNodeDistance; - - //! The second closest lower bound, on the previous iteration. - double lastSecondClosestBound; - //! Whether or not this node is pruned for the next iteration. - bool hamerlyPruned; - - //! The number of clusters that have been pruned. - size_t clustersPruned; - //! The current iteration. - size_t iteration; - //! The owner of these reference nodes (centroids.n_cols if there is no - //! owner). - size_t owner; - - // For nearest neighbor search. - - double firstBound; - double secondBound; - double bound; - void* lastDistanceNode; - double lastDistance; -}; - -} // namespace kmeans -} // namespace mlpack - -#endif diff --git a/src/mlpack/methods/kmeans/kmeans_main.cpp b/src/mlpack/methods/kmeans/kmeans_main.cpp index f14aaa4f14..2c5561bccc 100644 --- a/src/mlpack/methods/kmeans/kmeans_main.cpp +++ b/src/mlpack/methods/kmeans/kmeans_main.cpp @@ -13,7 +13,6 @@ #include "hamerly_kmeans.hpp" #include "pelleg_moore_kmeans.hpp" #include "dtnn_kmeans.hpp" -#include "dual_tree_kmeans.hpp" using namespace mlpack; using namespace mlpack::kmeans; @@ -153,15 +152,12 @@ void FindLloydStepType(const InitialPartitionPolicy& ipp) else if (algorithm == "pelleg-moore") RunKMeans(ipp); - else if (algorithm == "dtnn") - RunKMeans(ipp); - else if (algorithm == "dtnn-covertree") - RunKMeans(ipp); else if (algorithm == "dualtree") RunKMeans(ipp); + DefaultDTNNKMeans>(ipp); + else if (algorithm == "dualtree-covertree") + RunKMeans(ipp); else if (algorithm == "naive") RunKMeans(ipp); else diff --git a/src/mlpack/tests/kmeans_test.cpp b/src/mlpack/tests/kmeans_test.cpp index 6d61993320..54ca95cd0c 100644 --- a/src/mlpack/tests/kmeans_test.cpp +++ b/src/mlpack/tests/kmeans_test.cpp @@ -663,416 +663,4 @@ BOOST_AUTO_TEST_CASE(DTNNCoverTreeTest) */ } -/* -BOOST_AUTO_TEST_CASE(DualTreeKMeansTest) -{ - const size_t trials = 5; - - for (size_t t = 0; t < trials; ++t) - { - arma::mat dataset(10, 1000); - dataset.randu(); - - const size_t k = 5 * (t + 1); - arma::mat centroids(10, k); - centroids.randu(); - - arma::mat naiveCentroids(centroids); - KMeans<> km; - arma::Col assignments; - km.Cluster(dataset, k, assignments, naiveCentroids, false, true); - - KMeans dtnn; - arma::Col dtnnAssignments; - arma::mat dtnnCentroids(centroids); - dtnn.Cluster(dataset, k, dtnnAssignments, dtnnCentroids, false, true); - - for (size_t i = 0; i < dataset.n_cols; ++i) - BOOST_REQUIRE_EQUAL(assignments[i], dtnnAssignments[i]); - - for (size_t i = 0; i < centroids.n_elem; ++i) - BOOST_REQUIRE_CLOSE(naiveCentroids[i], dtnnCentroids[i], 1e-5); - } -} - -BOOST_AUTO_TEST_CASE(DualTreeKMeansBaseCaseTest) -{ - // If we run BaseCase() on all the points, do we get valid results? - const size_t points = 1000; - const size_t clusters = 5; - arma::mat dataset(5, points); - dataset.randu(); - arma::mat centroids(5, clusters); - centroids.randu(); - - // Create the Rules object. - arma::Col assignments(points); - arma::vec upperBounds(points); - arma::vec lowerBounds(points); - upperBounds.fill(DBL_MAX); - lowerBounds.fill(DBL_MAX); - std::vector visited(points, false); // Fill with false. - std::vector oldFromNewCentroids(clusters); - for (size_t i = 0; i < clusters; ++i) - oldFromNewCentroids[i] = i; - std::vector prunedPoints(points, false); // Fill with false. - - EuclideanDistance e; - DTNNKMeansRules, - EuclideanDistance, DTNNStatistic> > rules(centroids, dataset, - assignments, upperBounds, lowerBounds, e, prunedPoints, - oldFromNewCentroids, visited); - - for (size_t i = 0; i < points; ++i) - { - for (size_t j = 0; j < clusters; ++j) - { - rules.BaseCase(i, j); - } - } - - // Now, run nearest neighbors to establish true bounds. - neighbor::AllkNN allknn(centroids, dataset); - arma::Mat trueAssignments; - arma::mat trueDistances; - allknn.Search(2, trueAssignments, trueDistances); - - for (size_t i = 0; i < points; ++i) - { - BOOST_REQUIRE_GE(upperBounds[i], trueDistances(0, i)); - BOOST_REQUIRE_LE(lowerBounds[i], trueDistances(1, i)); - BOOST_REQUIRE_EQUAL(assignments[i], trueAssignments(0, i)); - } -} - -BOOST_AUTO_TEST_CASE(DualTreeKMeansScoreKDTreeOneLeafTest) -{ - // If we run a dual-tree algorithm, do we get valid results for each point - // and/or node when we use the kd-tree with a leaf size of one? - const size_t points = 5000; - const size_t clusters = 100; - arma::mat dataset(5, points); - dataset.randu(); - arma::mat centroids(5, clusters); - centroids.randu(); - - arma::mat datasetCopy(dataset); - arma::mat centroidsCopy(centroids); - - // Create the trees. - typedef BinarySpaceTree, DTNNStatistic> TreeType; - TreeType pointTree(dataset, 1); - TreeType centroidTree(centroids, 1); - - // Create the Rules object. - arma::Col assignments(points); - arma::vec upperBounds(points); - arma::vec lowerBounds(points); - upperBounds.fill(DBL_MAX); - lowerBounds.fill(DBL_MAX); - std::vector visited(points, false); // Fill with false. - std::vector oldFromNewCentroids(clusters); - for (size_t i = 0; i < clusters; ++i) - oldFromNewCentroids[i] = i; - std::vector prunedPoints(points, false); // Fill with false. - - EuclideanDistance e; - typedef DTNNKMeansRules RuleType; - RuleType rules(centroids, dataset, assignments, upperBounds, lowerBounds, e, - prunedPoints, oldFromNewCentroids, visited); - - // Now create the traverser. - typename TreeType::template BreadthFirstDualTreeTraverser - traverser(rules); - - pointTree.Stat().Pruned() = 0; - traverser.Traverse(pointTree, centroidTree); - - // Get true bounds. - AllkNN allknn(centroids, dataset); - - arma::Mat trueAssignments; - arma::mat trueDistances; - allknn.Search(2, trueAssignments, trueDistances); - - // Check the points first. Lots of weird mappings have to go on in this stage - // because the tree building procedure changed all the points around. - for (size_t i = 0; i < dataset.n_cols; ++i) - { - if (visited[i]) - { - BOOST_REQUIRE_GE(upperBounds[i], trueDistances(0, i)); - BOOST_REQUIRE_EQUAL(assignments[i], trueAssignments(0, i)); - } - } - - // Now traverse the tree to see if it is correct. - std::queue nodeQueue; - nodeQueue.push(&pointTree); - while (!nodeQueue.empty()) - { - // This is an expensive operation. We must ensure that the upper bound is - // valid and that the lower bound is valid. Both can be needlessly loose, - // but that will simply affect the pruning of the method, not the - // correctness. Here we care about correctness. - TreeType* node = nodeQueue.front(); - nodeQueue.pop(); - - // We must make sure the upper bound and lower bound are both valid for all - // descendant points. The lower bound only matters if the node was pruned. - // So, we must calculate the upper and lower bounds manually for the - // descendants. - double exactUpperBound = 0.0; - double exactLowerBound = DBL_MAX; - for (size_t i = 0; i < node->NumDescendants(); ++i) - { - if (trueDistances(0, node->Descendant(i)) > exactUpperBound) - exactUpperBound = trueDistances(0, node->Descendant(i)); - if (trueDistances(1, node->Descendant(i)) < exactLowerBound) - exactLowerBound = trueDistances(1, node->Descendant(i)); - } - - // Multiplication is to add some tolerance for floating point discrepancies. - BOOST_REQUIRE_GE(node->Stat().UpperBound() * 1.000001, exactUpperBound); - - if (node->Stat().Pruned() == centroids.n_cols) - { - BOOST_REQUIRE_LE(node->Stat().LowerBound() * 0.99999, exactLowerBound); - } - else - { - for (size_t i = 0; i < node->NumPoints(); ++i) - { - const double bestLower = std::min(node->Stat().LowerBound(), - lowerBounds[node->Point(i)]); - BOOST_REQUIRE_LE(bestLower * 0.99999, trueDistances(1, node->Point(i))); - } - - // Recurse. - for (size_t i = 0; i < node->NumChildren(); ++i) - nodeQueue.push(&node->Child(i)); - } - } -} - -BOOST_AUTO_TEST_CASE(DualTreeKMeansScoreKDTreeTest) -{ - // If we run a dual-tree algorithm, do we get valid results for each point - // and/or node when we use the kd-tree with the default leaf size? - const size_t points = 5000; - const size_t clusters = 100; - arma::mat dataset(5, points); - dataset.randu(); - arma::mat centroids(5, clusters); - centroids.randu(); - - arma::mat datasetCopy(dataset); - arma::mat centroidsCopy(centroids); - - // Create the trees. - typedef BinarySpaceTree, DTNNStatistic> TreeType; - TreeType pointTree(dataset); - TreeType centroidTree(centroids); - - // Create the Rules object. - arma::Col assignments(points); - arma::vec upperBounds(points); - arma::vec lowerBounds(points); - upperBounds.fill(DBL_MAX); - lowerBounds.fill(DBL_MAX); - std::vector visited(points, false); // Fill with false. - std::vector oldFromNewCentroids(clusters); - for (size_t i = 0; i < clusters; ++i) - oldFromNewCentroids[i] = i; - std::vector prunedPoints(points, false); // Fill with false. - - EuclideanDistance e; - typedef DTNNKMeansRules RuleType; - RuleType rules(centroids, dataset, assignments, upperBounds, lowerBounds, e, - prunedPoints, oldFromNewCentroids, visited); - - // Now create the traverser. - typename TreeType::template BreadthFirstDualTreeTraverser - traverser(rules); - - pointTree.Stat().Pruned() = 0; - traverser.Traverse(pointTree, centroidTree); - - // Get true bounds. - AllkNN allknn(centroids, dataset); - - arma::Mat trueAssignments; - arma::mat trueDistances; - allknn.Search(2, trueAssignments, trueDistances); - - // Check the points first. Lots of weird mappings have to go on in this stage - // because the tree building procedure changed all the points around. - for (size_t i = 0; i < dataset.n_cols; ++i) - { - if (visited[i]) - { - BOOST_REQUIRE_GE(upperBounds[i], trueDistances(0, i)); - BOOST_REQUIRE_EQUAL(assignments[i], trueAssignments(0, i)); - } - } - - // Now traverse the tree to see if it is correct. - std::queue nodeQueue; - nodeQueue.push(&pointTree); - while (!nodeQueue.empty()) - { - // This is an expensive operation. We must ensure that the upper bound is - // valid and that the lower bound is valid. Both can be needlessly loose, - // but that will simply affect the pruning of the method, not the - // correctness. Here we care about correctness. - TreeType* node = nodeQueue.front(); - nodeQueue.pop(); - - // We must make sure the upper bound and lower bound are both valid for all - // descendant points. The lower bound only matters if the node was pruned. - // So, we must calculate the upper and lower bounds manually for the - // descendants. - double exactUpperBound = 0.0; - double exactLowerBound = DBL_MAX; - for (size_t i = 0; i < node->NumDescendants(); ++i) - { - if (trueDistances(0, node->Descendant(i)) > exactUpperBound) - exactUpperBound = trueDistances(0, node->Descendant(i)); - if (trueDistances(1, node->Descendant(i)) < exactLowerBound) - exactLowerBound = trueDistances(1, node->Descendant(i)); - } - - // Multiplication is to add some tolerance for floating point discrepancies. - BOOST_REQUIRE_GE(node->Stat().UpperBound() * 1.000001, exactUpperBound); - - if (node->Stat().Pruned() == centroids.n_cols) - { - BOOST_REQUIRE_LE(node->Stat().LowerBound() * 0.99999, exactLowerBound); - } - else - { - for (size_t i = 0; i < node->NumPoints(); ++i) - { - const double bestLower = std::min(node->Stat().LowerBound(), - lowerBounds[node->Point(i)]); - BOOST_REQUIRE_LE(bestLower * 0.99999, trueDistances(1, node->Point(i))); - } - - // Recurse. - for (size_t i = 0; i < node->NumChildren(); ++i) - nodeQueue.push(&node->Child(i)); - } - } -} - -BOOST_AUTO_TEST_CASE(DualTreeKMeansScoreCoverTreeTest) -{ - // If we run a dual-tree algorithm, do we get valid results for each point - // and/or node when we use the kd-tree with the default leaf size? - const size_t points = 5000; - const size_t clusters = 100; - arma::mat dataset(5, points); - dataset.randu(); - arma::mat centroids(5, clusters); - centroids.randu(); - - arma::mat datasetCopy(dataset); - arma::mat centroidsCopy(centroids); - - // Create the trees. - typedef CoverTree - TreeType; - TreeType pointTree(dataset); - TreeType centroidTree(centroids); - - // Create the Rules object. - arma::Col assignments(points); - arma::vec upperBounds(points); - arma::vec lowerBounds(points); - upperBounds.fill(DBL_MAX); - lowerBounds.fill(DBL_MAX); - std::vector visited(points, false); // Fill with false. - std::vector oldFromNewCentroids; // Not used. - std::vector prunedPoints(points, false); // Fill with false. - - EuclideanDistance e; - typedef DTNNKMeansRules RuleType; - RuleType rules(centroids, dataset, assignments, upperBounds, lowerBounds, e, - prunedPoints, oldFromNewCentroids, visited); - - // Now create the traverser. - typename TreeType::template DualTreeTraverser traverser(rules); - - pointTree.Stat().Pruned() = 0; - traverser.Traverse(pointTree, centroidTree); - - // Get true bounds. - AllkNN allknn(centroids, dataset); - - arma::Mat trueAssignments; - arma::mat trueDistances; - allknn.Search(2, trueAssignments, trueDistances); - - // Check the points first. Lots of weird mappings have to go on in this stage - // because the tree building procedure changed all the points around. - for (size_t i = 0; i < dataset.n_cols; ++i) - { - if (visited[i]) - { - BOOST_REQUIRE_GE(upperBounds[i], trueDistances(0, i)); - BOOST_REQUIRE_EQUAL(assignments[i], trueAssignments(0, i)); - } - } - - // Now traverse the tree to see if it is correct. - std::queue nodeQueue; - nodeQueue.push(&pointTree); - while (!nodeQueue.empty()) - { - // This is an expensive operation. We must ensure that the upper bound is - // valid and that the lower bound is valid. Both can be needlessly loose, - // but that will simply affect the pruning of the method, not the - // correctness. Here we care about correctness. - TreeType* node = nodeQueue.front(); - nodeQueue.pop(); - - // We must make sure the upper bound and lower bound are both valid for all - // descendant points. The lower bound only matters if the node was pruned. - // So, we must calculate the upper and lower bounds manually for the - // descendants. - double exactUpperBound = 0.0; - double exactLowerBound = DBL_MAX; - for (size_t i = 0; i < node->NumDescendants(); ++i) - { - if (trueDistances(0, node->Descendant(i)) > exactUpperBound) - exactUpperBound = trueDistances(0, node->Descendant(i)); - if (trueDistances(1, node->Descendant(i)) < exactLowerBound) - exactLowerBound = trueDistances(1, node->Descendant(i)); - } - - // Multiplication is to add some tolerance for floating point discrepancies. - BOOST_REQUIRE_GE(node->Stat().UpperBound() * 1.000001, exactUpperBound); - - if (node->Stat().Pruned() == centroids.n_cols) - { - BOOST_REQUIRE_LE(node->Stat().LowerBound() * 0.99999, exactLowerBound); - } - else if (node->NumChildren() == 0) - { - // The node has one point. - const double bestLower = std::min(node->Stat().LowerBound(), - lowerBounds[node->Point(0)]); - BOOST_REQUIRE_LE(bestLower * 0.99999, trueDistances(1, node->Point(0))); - } - else - { - // Recurse. - for (size_t i = 0; i < node->NumChildren(); ++i) - nodeQueue.push(&node->Child(i)); - } - } -} -*/ - BOOST_AUTO_TEST_SUITE_END();