From dea4b29b7185423243ee958d6ffb76c8ff9eaac7 Mon Sep 17 00:00:00 2001 From: Ryan Curtin Date: Wed, 14 Jan 2015 10:58:27 -0500 Subject: [PATCH] Update FirstBound correctly. Trivial speedup. --- src/mlpack/methods/kmeans/dual_tree_kmeans.hpp | 3 ++- src/mlpack/methods/kmeans/dual_tree_kmeans_impl.hpp | 12 +++++++++--- 2 files changed, 11 insertions(+), 4 deletions(-) diff --git a/src/mlpack/methods/kmeans/dual_tree_kmeans.hpp b/src/mlpack/methods/kmeans/dual_tree_kmeans.hpp index ebeb0edeb6..27bcf2554e 100644 --- a/src/mlpack/methods/kmeans/dual_tree_kmeans.hpp +++ b/src/mlpack/methods/kmeans/dual_tree_kmeans.hpp @@ -58,7 +58,8 @@ class DualTreeKMeans //! Track distance calculations. size_t distanceCalculations; - void ClusterTreeUpdate(TreeType* node); + void ClusterTreeUpdate(TreeType* node, + const arma::mat& distances); void TreeUpdate(TreeType* node, const size_t clusters, diff --git a/src/mlpack/methods/kmeans/dual_tree_kmeans_impl.hpp b/src/mlpack/methods/kmeans/dual_tree_kmeans_impl.hpp index 88c4e96f10..ed51137611 100644 --- a/src/mlpack/methods/kmeans/dual_tree_kmeans_impl.hpp +++ b/src/mlpack/methods/kmeans/dual_tree_kmeans_impl.hpp @@ -81,7 +81,7 @@ double DualTreeKMeans::Iterate( distanceCalculations += nns.Scores(); // Update FirstBound(). - ClusterTreeUpdate(centroidTree); + ClusterTreeUpdate(centroidTree, interclusterDistances); // Now run the dual-tree algorithm. typedef DualTreeKMeansRules RulesType; @@ -133,16 +133,22 @@ double DualTreeKMeans::Iterate( template void DualTreeKMeans::ClusterTreeUpdate( - TreeType* node) + 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)); + 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(1, node->Point(i)) > firstBound) + firstBound = distances(1, node->Point(i)); + } node->Stat().FirstBound() = firstBound; }