Skip to content

Instantly share code, notes, and snippets.

@rajanim
Created October 8, 2017 15:35
Show Gist options
  • Select an option

  • Save rajanim/e01d1af591d697202c75ee0a785fa94f to your computer and use it in GitHub Desktop.

Select an option

Save rajanim/e01d1af591d697202c75ee0a785fa94f to your computer and use it in GitHub Desktop.
package org.sfsu.cs.selectivesearch.common.distance
import org.apache.spark.mllib.linalg.{DenseVector, SparseVector, Vector, Vectors}
/**
* Created by rajanishivarajmaski1 on 4/2/17.
* fork from spark-scala apis
* * The cosine distance between two points: cosineDistance(a,b) = (a dot b)/(norm(a) * norm(b))
*/
object CosineDistanceMeasure {
def distance(v1: org.apache.spark.mllib.linalg.Vector, v2: org.apache.spark.mllib.linalg.Vector): Double = {
// if (v1.size != v2.size) throw new CardinalityException(v1.size, v2.size)
val dotProduct: Double = dot(v1, v2)
val denom = Vectors.norm(v1, 2.0) * Vectors.norm(v2, 2.0)
if (denom == 0.0) {
0.0
} else {
(dotProduct / denom)
}
}
/**
* dot(x, y)
*/
def dot(x: Vector, y: Vector): Double = {
val x1 = x.toDense
val y1 = y.toDense
//require(x.size == y.size,
// "BLAS.dot(x: Vector, y:Vector) was given Vectors with non-matching sizes:" +
// " x.size = " + x.size + ", y.size = " + y.size)
(x1, y1) match {
case (dx: DenseVector, dy: DenseVector) =>
dot(dx, dy)
// case (sx: SparseVector, sy: SparseVector) =>
// dot(sx, sy)
case _ =>
throw new IllegalArgumentException(s"dot doesn't support (${x.getClass}, ${y.getClass}).")
}
}
/**
* L2Norm
*/
def getL2Norm(values: Array[Double]): Double = {
var sum = 0.0
var i = 0
val size = values.length
while (i < size) {
sum += values(i) * values(i)
i += 1
}
math.sqrt(sum)
}
/**
* dot(x, y) array of [double] of same size
*/
def dot(v1: Array[Double], v2: Array[Double]): Double = {
if (v1.size != v2.size) throw new CardinalityException(v1.size, v2.size)
var result: Double = 0.0
val max: Int = v1.size
var i: Int = 0
while (i < max) {
result += v1(i) * v2(i)
i += 1
}
result
}
/**
* dot(x, y) DenseVector
*/
def dot(v1: DenseVector, v2: DenseVector): Double = {
if (v1.size != v2.size) throw new CardinalityException(v1.size, v2.size)
var result: Double = 0.0
val max: Int = v1.size
var i: Int = 0
while (i < max) {
result += v1(i) * v2(i)
i += 1
}
result
}
/**
* dot(x, y) SparseVector
*/
def dot(x: SparseVector, y: SparseVector): Double = {
val xValues = x.values
val xIndices = x.indices
val yValues = y.values
val yIndices = y.indices
val nnzx = xIndices.length
val nnzy = yIndices.length
var kx = 0
var ky = 0
var sum = 0.0
// y catching x
while (kx < nnzx && ky < nnzy) {
val ix = xIndices(kx)
while (ky < nnzy && yIndices(ky) < ix) {
ky += 1
}
if (ky < nnzy && yIndices(ky) == ix) {
sum += xValues(kx) * yValues(ky)
ky += 1
}
kx += 1
}
sum
}
}
/**
* Forked from spark-scala apis
* Exception thrown when there is a cardinality mismatch in matrix or vector operations.
* For example, vectors of differing cardinality cannot be added.
*/
class CardinalityException(val expected: Int, val cardinality: Int) extends IllegalArgumentException("Required cardinality " + expected + " but got " + cardinality) {
}
class TestCosineDistanceMeasure extends TestSuiteBuilder {
test("cosineDistance") {
val vector1 = Vectors.sparse(5, Seq((1, 1.0), (2, 1.0)))
val vector2 = Vectors.sparse(5, Seq((1, 1.0), (2, 1.0)))
println("dot product", CosineDistanceMeasure.dot(vector1.toArray, vector2.toArray))
println("Vectors.norm", Vectors.norm(vector2, 2.0))
println("denom", Vectors.norm(vector1, 2.0) * Vectors.norm(vector2, 2.0))
println(s"distance measure of vector 1 and 2, same", CosineDistanceMeasure.distance(vector1, vector2))
val vector3 = Vectors.sparse(5, Seq((1, 1.0), (2, 1.0), (3, 1.0)))
println(s"distance measure of vector 1 and 3, nearly similar", CosineDistanceMeasure.distance(vector1, vector3) )
val vector4 = Vectors.sparse(5, Seq((0, 1.0), (1, 0.0), (2, 0.0), (3, 1.0)))
println(s"distance measure of vector 1 and 4, dissimilar ", CosineDistanceMeasure.distance(vector1, vector4) )
}
}
@nullm4ri

Copy link
Copy Markdown

Hello,
I saw your answer in here and I wonder if you can share your custom k-means implemented.

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment