Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
68 changes: 51 additions & 17 deletions mllib/src/main/scala/org/apache/spark/ml/functions.scala
Original file line number Diff line number Diff line change
Expand Up @@ -18,10 +18,10 @@
package org.apache.spark.ml

import org.apache.spark.annotation.Since
import org.apache.spark.ml.linalg.{DenseVector, SparseVector, Vector}
import org.apache.spark.ml.linalg.{DenseVector, SparseVector, Vector, VectorUDT}
import org.apache.spark.sql.{functions => sf}
import org.apache.spark.sql.Column
import org.apache.spark.sql.types.{ArrayType, IntegerType}
import org.apache.spark.sql.types.{ArrayType, DoubleType, IntegerType}

// scalastyle:off
@Since("3.0.0")
Expand Down Expand Up @@ -65,24 +65,58 @@ object functions {
Column.internalFn("ml_vector_dot_product", sf.unwrap_udt(left), sf.unwrap_udt(right))

private[ml] def vector_dot_product(left: Column, right: Vector): Column = {
val rightStruct = right match {
case sparse: SparseVector =>
sf.struct(
sf.lit(0.toByte).alias("type"),
sf.lit(sparse.size).alias("size"),
sf.lit(sparse.indices).alias("indices"),
sf.lit(sparse.values).alias("values"))
case dense: DenseVector =>
sf.struct(
sf.lit(1.toByte).alias("type"),
sf.lit(null).cast(IntegerType).alias("size"),
sf.lit(null).cast(ArrayType(IntegerType)).alias("indices"),
sf.lit(dense.values).alias("values"))
}
Column.internalFn(
"ml_vector_dot_product",
sf.unwrap_udt(left),
rightStruct)
vectorToStruct(right))
}

private[ml] def vector_affine_transform(
vector: Column,
scale: Column,
shift: Column): Column = {
val transformed = Column.internalFn(
"ml_vector_affine_transform",
sf.unwrap_udt(vector),
scale,
shift)
sf.wrap_udt(transformed, new VectorUDT)
}

private[ml] def vector_affine_transform(
vector: Column,
scale: Array[Double],
shift: Array[Double]): Column = {
val transformed = Column.internalFn(
"ml_vector_affine_transform",
sf.unwrap_udt(vector),
doubleArrayLiteral(scale),
doubleArrayLiteral(shift))
sf.wrap_udt(transformed, new VectorUDT)
}

private def doubleArrayLiteral(values: Array[Double]): Column = {
if (values == null) {
sf.lit(null).cast(ArrayType(DoubleType, containsNull = false))
} else {
sf.typedLit(values)
}
}

private def vectorToStruct(vector: Vector): Column = vector match {
case null => sf.lit(null).cast(new VectorUDT().sqlType)
case sparse: SparseVector =>
sf.struct(
sf.lit(0.toByte).alias("type"),
sf.lit(sparse.size).alias("size"),
sf.lit(sparse.indices).alias("indices"),
sf.lit(sparse.values).alias("values"))
case dense: DenseVector =>
sf.struct(
sf.lit(1.toByte).alias("type"),
sf.lit(null).cast(IntegerType).alias("size"),
sf.lit(null).cast(ArrayType(IntegerType)).alias("indices"),
sf.lit(dense.values).alias("values"))
}

private[ml] def array_binary_search(a: Column, v: Column): Column =
Expand Down
117 changes: 114 additions & 3 deletions mllib/src/test/scala/org/apache/spark/ml/FunctionsSuite.scala
Original file line number Diff line number Diff line change
Expand Up @@ -24,9 +24,9 @@ import org.apache.spark.ml.util.MLTest
import org.apache.spark.mllib.linalg.{Matrices => OldMatrices, MatrixUDT => OldMatrixUDT,
Vector => OldVector, Vectors => OldVectors, VectorUDT => OldVectorUDT}
import org.apache.spark.sql.{AnalysisException, DataFrame, Row}
import org.apache.spark.sql.catalyst.expressions.ml.VectorPosExplode
import org.apache.spark.sql.functions.{col, unwrap_udt, wrap_udt}
import org.apache.spark.sql.types.{StructField, StructType, UserDefinedType}
import org.apache.spark.sql.catalyst.expressions.ml.{VectorAffineTransform, VectorPosExplode}
import org.apache.spark.sql.functions.{col, typedLit, unwrap_udt, wrap_udt}
import org.apache.spark.sql.types.{ArrayType, DoubleType, StructField, StructType, UserDefinedType}

class FunctionsSuite extends MLTest {

Expand Down Expand Up @@ -256,6 +256,117 @@ class FunctionsSuite extends MLTest {
assert(error.getMessage.contains("vectors with non-matching sizes"))
}

test("test vector_affine_transform") {
val df = Seq(
(Vectors.dense(1.0, 2.0), Array(2.0, 3.0), Array(4.0, 5.0)),
(Vectors.sparse(2, Seq((0, 1.0))), Array(2.0, 3.0), null),
(Vectors.dense(1.0, 2.0), null, Array(4.0, 5.0)),
(Vectors.sparse(2, Seq((0, 1.0))), Array(2.0, 3.0), Array(0.0, 0.0)),
(Vectors.sparse(2, Seq((0, 1.0))), null, null),
(null, Array(2.0, 3.0), Array(4.0, 5.0)))
.toDF("vector", "scale", "shift")

assert(df.schema("scale").dataType === ArrayType(DoubleType, containsNull = false))
assert(df.schema("shift").dataType === ArrayType(DoubleType, containsNull = false))
val transformed = df.select(vector_affine_transform($"vector", $"scale", $"shift"))
assert(transformed.schema.head.dataType === new VectorUDT)
assert(transformed.collect().map(_.get(0)).toSeq === Seq(
Vectors.dense(6.0, 11.0),
Vectors.sparse(2, Seq((0, 2.0))),
Vectors.dense(5.0, 7.0),
Vectors.dense(2.0, 0.0),
Vectors.sparse(2, Seq((0, 1.0))),
null))

val expressions = transformed.queryExecution.analyzed
.flatMap(_.expressions.flatMap(_.collect { case v: VectorAffineTransform => v }))
assert(expressions.map(_.prettyName).distinct === Seq("ml_vector_affine_transform"))

val constantResult = df.limit(1)
.select(vector_affine_transform(
$"vector",
Array(2.0, 3.0),
Array(4.0, 5.0)))
.first()
.getAs[Vector](0)
assert(constantResult === Vectors.dense(6.0, 11.0))

val cachedScaleResult = df.limit(1)
.select(vector_affine_transform($"vector", typedLit(Array(2.0, 3.0)), $"shift"))
.first()
.getAs[Vector](0)
assert(cachedScaleResult === Vectors.dense(6.0, 11.0))

val cachedShiftResult = df.limit(1)
.select(vector_affine_transform($"vector", $"scale", typedLit(Array(4.0, 5.0))))
.first()
.getAs[Vector](0)
assert(cachedShiftResult === Vectors.dense(6.0, 11.0))

val scaleOnlyConstantResult = df.limit(1)
.select(vector_affine_transform(
$"vector",
Array(2.0, 3.0),
null.asInstanceOf[Array[Double]]))
.first()
.getAs[Vector](0)
assert(scaleOnlyConstantResult === Vectors.dense(2.0, 6.0))

val shiftOnlyConstantResult = df.limit(1)
.select(vector_affine_transform(
$"vector",
null.asInstanceOf[Array[Double]],
Array(4.0, 5.0)))
.first()
.getAs[Vector](0)
assert(shiftOnlyConstantResult === Vectors.dense(5.0, 7.0))

val nullConstantsResult = df.limit(1)
.select(vector_affine_transform(
$"vector",
null.asInstanceOf[Array[Double]],
null.asInstanceOf[Array[Double]]))
.first()
.getAs[Vector](0)
assert(nullConstantsResult === Vectors.dense(1.0, 2.0))

val emptyConstantsResult = Seq(Tuple1(Vectors.dense(Array.emptyDoubleArray)))
.toDF("vector")
.select(vector_affine_transform($"vector", Array.emptyDoubleArray, Array.emptyDoubleArray))
.first()
.getAs[Vector](0)
assert(emptyConstantsResult === Vectors.dense(Array.emptyDoubleArray))

val specialValues = Seq(Double.NaN, Double.NegativeInfinity, Double.PositiveInfinity)
val specialValueRows = specialValues.flatMap { value =>
Seq(
(Vectors.dense(value), Array(1.0), Array(0.0)),
(Vectors.dense(1.0), Array(value), Array(0.0)),
(Vectors.dense(1.0), Array(1.0), Array(value)))
}
val specialValueResults = specialValueRows
.toDF("vector", "scale", "shift")
.select(vector_affine_transform($"vector", $"scale", $"shift"))
.collect()
.map(_.getAs[Vector](0)(0))
specialValueResults.zip(specialValues.flatMap(value => Seq.fill(3)(value)))
.foreach { case (actual, expected) =>
assert(java.lang.Double.compare(actual, expected) === 0)
}

Seq(
(Array(1.0, 2.0), Array(0.0)),
(Array(1.0), Array(0.0, 0.0))).foreach { case (scale, shift) =>
val error = intercept[IllegalArgumentException] {
Seq((Vectors.dense(1.0), scale, shift))
.toDF("vector", "scale", "shift")
.select(vector_affine_transform($"vector", $"scale", $"shift"))
.collect()
}
assert(error.getMessage.contains("inputs with non-matching sizes"))
}
}

test("test get_vector") {
val df = Seq(
(Vectors.dense(1.0, 2.0, 3.0), 0),
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,181 @@
/*
* Licensed to the Apache Software Foundation (ASF) under one or more
* contributor license agreements. See the NOTICE file distributed with
* this work for additional information regarding copyright ownership.
* The ASF licenses this file to You under the Apache License, Version 2.0
* (the "License"); you may not use this file except in compliance with
* the License. You may obtain a copy of the License at
*
* http://www.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*/

package org.apache.spark.sql.catalyst.expressions.ml;

import org.apache.spark.sql.catalyst.InternalRow;
import org.apache.spark.sql.catalyst.expressions.GenericInternalRow;
import org.apache.spark.sql.catalyst.expressions.UnsafeArrayData;
import org.apache.spark.sql.catalyst.util.ArrayData;

public class MLExpressionUtils {
private static final byte SPARSE_VECTOR_TYPE = 0;
private static final byte DENSE_VECTOR_TYPE = 1;

private MLExpressionUtils() {}

public static InternalRow affineTransform(
InternalRow vector,
ArrayData scale,
ArrayData shift) {
return affineTransform(vector, scale, shift, null, null);
}

public static InternalRow affineTransform(
InternalRow vector,
ArrayData scale,
ArrayData shift,
double[] cachedScale,
double[] cachedShift) {
boolean hasScale = scale != null || cachedScale != null;
boolean hasShift = shift != null || cachedShift != null;
if (!hasScale && !hasShift) {
return vector;
}

byte vectorType = vector.getByte(0);
ArrayData vectorValues = vector.getArray(3);
int size;
if (vectorType == SPARSE_VECTOR_TYPE) {
size = vector.getInt(1);
} else if (vectorType == DENSE_VECTOR_TYPE) {
size = vectorValues.numElements();
} else {
throw new IllegalArgumentException("Unknown vector type " + vectorType + ".");
}

int scaleSize = !hasScale ? size :
(cachedScale == null ? scale.numElements() : cachedScale.length);
int shiftSize = !hasShift ? size :
(cachedShift == null ? shift.numElements() : cachedShift.length);
if (size != scaleSize || size != shiftSize) {
throw new IllegalArgumentException(
"requirement failed: VectorAffineTransform was given inputs with non-matching sizes: " +
"vector.size = " + size + ", scale.size = " + scaleSize +
", shift.size = " + shiftSize);
}

if (vectorType == SPARSE_VECTOR_TYPE && !hasShift) {
ArrayData vectorIndices = vector.getArray(2);
double[] resultValues = new double[vectorValues.numElements()];
if (cachedScale == null) {
for (int vectorIndex = 0; vectorIndex < resultValues.length; vectorIndex++) {
int featureIndex = vectorIndices.getInt(vectorIndex);
resultValues[vectorIndex] =
vectorValues.getDouble(vectorIndex) * scale.getDouble(featureIndex);
}
} else {
for (int vectorIndex = 0; vectorIndex < resultValues.length; vectorIndex++) {
int featureIndex = vectorIndices.getInt(vectorIndex);
resultValues[vectorIndex] =
vectorValues.getDouble(vectorIndex) * cachedScale[featureIndex];
}
}
return new GenericInternalRow(new Object[] {
SPARSE_VECTOR_TYPE,
size,
vectorIndices,
UnsafeArrayData.fromPrimitiveArray(resultValues)
});
}

double[] resultValues = new double[size];
if (vectorType == DENSE_VECTOR_TYPE) {
if (!hasScale) {
if (cachedShift == null) {
for (int featureIndex = 0; featureIndex < size; featureIndex++) {
resultValues[featureIndex] =
vectorValues.getDouble(featureIndex) + shift.getDouble(featureIndex);
}
} else {
for (int featureIndex = 0; featureIndex < size; featureIndex++) {
resultValues[featureIndex] =
vectorValues.getDouble(featureIndex) + cachedShift[featureIndex];
}
}
} else if (!hasShift) {
if (cachedScale == null) {
for (int featureIndex = 0; featureIndex < size; featureIndex++) {
resultValues[featureIndex] =
vectorValues.getDouble(featureIndex) * scale.getDouble(featureIndex);
}
} else {
for (int featureIndex = 0; featureIndex < size; featureIndex++) {
resultValues[featureIndex] =
vectorValues.getDouble(featureIndex) * cachedScale[featureIndex];
}
}
} else if (cachedScale != null && cachedShift != null) {
for (int featureIndex = 0; featureIndex < size; featureIndex++) {
resultValues[featureIndex] = vectorValues.getDouble(featureIndex) *
cachedScale[featureIndex] + cachedShift[featureIndex];
}
} else if (cachedScale != null) {
for (int featureIndex = 0; featureIndex < size; featureIndex++) {
resultValues[featureIndex] = vectorValues.getDouble(featureIndex) *
cachedScale[featureIndex] + shift.getDouble(featureIndex);
}
} else if (cachedShift != null) {
for (int featureIndex = 0; featureIndex < size; featureIndex++) {
resultValues[featureIndex] = vectorValues.getDouble(featureIndex) *
scale.getDouble(featureIndex) + cachedShift[featureIndex];
}
} else {
for (int featureIndex = 0; featureIndex < size; featureIndex++) {
resultValues[featureIndex] = vectorValues.getDouble(featureIndex) *
scale.getDouble(featureIndex) + shift.getDouble(featureIndex);
}
}
} else {
if (cachedShift == null) {
for (int featureIndex = 0; featureIndex < size; featureIndex++) {
resultValues[featureIndex] = 0.0 + shift.getDouble(featureIndex);
}
} else {
for (int featureIndex = 0; featureIndex < size; featureIndex++) {
resultValues[featureIndex] = 0.0 + cachedShift[featureIndex];
}
}

ArrayData vectorIndices = vector.getArray(2);
if (!hasScale) {
for (int vectorIndex = 0; vectorIndex < vectorValues.numElements(); vectorIndex++) {
int featureIndex = vectorIndices.getInt(vectorIndex);
resultValues[featureIndex] += vectorValues.getDouble(vectorIndex);
}
} else if (cachedScale == null) {
for (int vectorIndex = 0; vectorIndex < vectorValues.numElements(); vectorIndex++) {
int featureIndex = vectorIndices.getInt(vectorIndex);
resultValues[featureIndex] +=
vectorValues.getDouble(vectorIndex) * scale.getDouble(featureIndex);
}
} else {
for (int vectorIndex = 0; vectorIndex < vectorValues.numElements(); vectorIndex++) {
int featureIndex = vectorIndices.getInt(vectorIndex);
resultValues[featureIndex] +=
vectorValues.getDouble(vectorIndex) * cachedScale[featureIndex];
}
}
}
return new GenericInternalRow(new Object[] {
DENSE_VECTOR_TYPE,
null,
null,
UnsafeArrayData.fromPrimitiveArray(resultValues)
});
}
}
Loading