From 3988fc6987178a006091c98bd046b097f8b412da Mon Sep 17 00:00:00 2001 From: Erich Schubert Date: Thu, 24 Mar 2016 22:02:27 +0100 Subject: [PATCH 1/2] Bug fix for Elkan? AFAICT, the old code always killed all centroids if one cluster is empty. Keeping the last centroids is much more stable. --- src/mlpack/methods/kmeans/elkan_kmeans_impl.hpp | 15 ++++++++------- 1 file changed, 8 insertions(+), 7 deletions(-) diff --git a/src/mlpack/methods/kmeans/elkan_kmeans_impl.hpp b/src/mlpack/methods/kmeans/elkan_kmeans_impl.hpp index 638e35c873..d4a1b1fd13 100644 --- a/src/mlpack/methods/kmeans/elkan_kmeans_impl.hpp +++ b/src/mlpack/methods/kmeans/elkan_kmeans_impl.hpp @@ -153,14 +153,15 @@ double ElkanKMeans::Iterate(const arma::mat& centroids, double cNorm = 0.0; // Cluster movement for residual. for (size_t c = 0; c < centroids.n_cols; ++c) { - if (counts[c] > 0) + if (counts[c] > 0) { newCentroids.col(c) /= counts[c]; - else - newCentroids.fill(DBL_MAX); // Fill with invalid value. - - moveDistances(c) = metric.Evaluate(newCentroids.col(c), centroids.col(c)); - cNorm += std::pow(moveDistances(c), 2.0); - distanceCalculations++; + moveDistances(c) = metric.Evaluate(newCentroids.col(c), centroids.col(c)); + cNorm += std::pow(moveDistances(c), 2.0); + distanceCalculations++; + } else { + newCentroids.col(c) = centroids.col(c); // Keep old centroid. + moveDistances(c) = 0.0; + } } for (size_t i = 0; i < dataset.n_cols; ++i) From 04b0c0867f13d88edff8419d7deaec2315b54258 Mon Sep 17 00:00:00 2001 From: Erich Schubert Date: Fri, 25 Mar 2016 22:50:01 +0100 Subject: [PATCH 2/2] Minimal fix for bug in Elkan implementation. Kill only one cluster, not all. --- src/mlpack/methods/kmeans/elkan_kmeans_impl.hpp | 15 +++++++-------- 1 file changed, 7 insertions(+), 8 deletions(-) diff --git a/src/mlpack/methods/kmeans/elkan_kmeans_impl.hpp b/src/mlpack/methods/kmeans/elkan_kmeans_impl.hpp index d4a1b1fd13..c2682f7382 100644 --- a/src/mlpack/methods/kmeans/elkan_kmeans_impl.hpp +++ b/src/mlpack/methods/kmeans/elkan_kmeans_impl.hpp @@ -153,15 +153,14 @@ double ElkanKMeans::Iterate(const arma::mat& centroids, double cNorm = 0.0; // Cluster movement for residual. for (size_t c = 0; c < centroids.n_cols; ++c) { - if (counts[c] > 0) { + if (counts[c] > 0) newCentroids.col(c) /= counts[c]; - moveDistances(c) = metric.Evaluate(newCentroids.col(c), centroids.col(c)); - cNorm += std::pow(moveDistances(c), 2.0); - distanceCalculations++; - } else { - newCentroids.col(c) = centroids.col(c); // Keep old centroid. - moveDistances(c) = 0.0; - } + else + newCentroids.col(c).fill(DBL_MAX); // Fill with invalid value. + + moveDistances(c) = metric.Evaluate(newCentroids.col(c), centroids.col(c)); + cNorm += std::pow(moveDistances(c), 2.0); + distanceCalculations++; } for (size_t i = 0; i < dataset.n_cols; ++i)