From bce6482563336fe3111e79ecb166f39e1fcc0d7d Mon Sep 17 00:00:00 2001 From: Ryan Curtin Date: Thu, 29 Jan 2015 19:35:12 -0500 Subject: [PATCH] Slightly tighter prune. --- .../methods/kmeans/dual_tree_kmeans_impl.hpp | 6 +++--- .../kmeans/dual_tree_kmeans_rules_impl.hpp | 16 ++++++++++++---- 2 files changed, 15 insertions(+), 7 deletions(-) diff --git a/src/mlpack/methods/kmeans/dual_tree_kmeans_impl.hpp b/src/mlpack/methods/kmeans/dual_tree_kmeans_impl.hpp index dbefdc9470..a2c91b55e2 100644 --- a/src/mlpack/methods/kmeans/dual_tree_kmeans_impl.hpp +++ b/src/mlpack/methods/kmeans/dual_tree_kmeans_impl.hpp @@ -42,7 +42,7 @@ DualTreeKMeans::DualTreeKMeans( datasetCopy = datasetOrig; // Now build the tree. We don't need any mappings. - tree = new TreeType(const_cast(this->dataset), 1); + tree = new TreeType(const_cast(this->dataset), 10); Timer::Stop("tree_building"); } @@ -312,8 +312,8 @@ closest << "! It's part of node r" << node->Begin() << "c" << node->Count() << else if (node->Stat().MaxQueryNodeDistance() < 0.5 * interclusterDistances(0, owner)) { - Log::Warn << "Secondary Elkan prune! r" << node->Begin() << "c" << -node->Count() << ".\n"; +// Log::Warn << "Secondary Elkan prune! r" << node->Begin() << "c" << +//node->Count() << ".\n"; node->Stat().HamerlyPruned() = true; if (!node->Parent()->Stat().HamerlyPruned()) hamerlyPruned += node->NumDescendants(); diff --git a/src/mlpack/methods/kmeans/dual_tree_kmeans_rules_impl.hpp b/src/mlpack/methods/kmeans/dual_tree_kmeans_rules_impl.hpp index 21d526d7f3..403ed699a2 100644 --- a/src/mlpack/methods/kmeans/dual_tree_kmeans_rules_impl.hpp +++ b/src/mlpack/methods/kmeans/dual_tree_kmeans_rules_impl.hpp @@ -148,10 +148,18 @@ double DualTreeKMeansRules::Score( if (distances.Lo() < referenceNode.Stat().MinQueryNodeDistance()) { // This is the new closest node. - referenceNode.Stat().SecondMinQueryNodeDistance() = - referenceNode.Stat().MinQueryNodeDistance(); - referenceNode.Stat().SecondMaxQueryNodeDistance() = - referenceNode.Stat().MaxQueryNodeDistance(); + 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(); }