Repository: spark
Updated Branches:
  refs/heads/master 3027f06b4 -> 246111d17


[SPARK-5365][MLlib] Refactor KMeans to reduce redundant data

If a point is selected as new centers for many runs, it would collect many 
redundant data. This pr refactors it.

Author: Liang-Chi Hsieh <[email protected]>

Closes #4159 from viirya/small_refactor_kmeans and squashes the following 
commits:

25487e6 [Liang-Chi Hsieh] Refactor codes to reduce redundant data.


Project: http://git-wip-us.apache.org/repos/asf/spark/repo
Commit: http://git-wip-us.apache.org/repos/asf/spark/commit/246111d1
Tree: http://git-wip-us.apache.org/repos/asf/spark/tree/246111d1
Diff: http://git-wip-us.apache.org/repos/asf/spark/diff/246111d1

Branch: refs/heads/master
Commit: 246111d179a2f3f6b97a5c2b121d8ddbfd1c9aad
Parents: 3027f06
Author: Liang-Chi Hsieh <[email protected]>
Authored: Thu Jan 22 08:16:35 2015 -0800
Committer: Xiangrui Meng <[email protected]>
Committed: Thu Jan 22 08:16:35 2015 -0800

----------------------------------------------------------------------
 .../scala/org/apache/spark/mllib/clustering/KMeans.scala    | 9 +++++----
 1 file changed, 5 insertions(+), 4 deletions(-)
----------------------------------------------------------------------


http://git-wip-us.apache.org/repos/asf/spark/blob/246111d1/mllib/src/main/scala/org/apache/spark/mllib/clustering/KMeans.scala
----------------------------------------------------------------------
diff --git 
a/mllib/src/main/scala/org/apache/spark/mllib/clustering/KMeans.scala 
b/mllib/src/main/scala/org/apache/spark/mllib/clustering/KMeans.scala
index fc46da3..11633e8 100644
--- a/mllib/src/main/scala/org/apache/spark/mllib/clustering/KMeans.scala
+++ b/mllib/src/main/scala/org/apache/spark/mllib/clustering/KMeans.scala
@@ -328,14 +328,15 @@ class KMeans private (
       val chosen = data.zip(costs).mapPartitionsWithIndex { (index, 
pointsWithCosts) =>
         val rand = new XORShiftRandom(seed ^ (step << 16) ^ index)
         pointsWithCosts.flatMap { case (p, c) =>
-          (0 until runs).filter { r =>
+          val rs = (0 until runs).filter { r =>
             rand.nextDouble() < 2.0 * c(r) * k / sumCosts(r)
-          }.map((_, p))
+          }
+          if (rs.length > 0) Some(p, rs) else None
         }
       }.collect()
       mergeNewCenters()
-      chosen.foreach { case (r, p) =>
-        newCenters(r) += p.toDense
+      chosen.foreach { case (p, rs) =>
+        rs.foreach(newCenters(_) += p.toDense)
       }
       step += 1
     }


---------------------------------------------------------------------
To unsubscribe, e-mail: [email protected]
For additional commands, e-mail: [email protected]

Reply via email to