From faa46184da51817488da7e210edc081ed8ed2fd4 Mon Sep 17 00:00:00 2001 From: Ruifeng Zheng Date: Thu, 10 Sep 2026 10:13:13 +0000 Subject: [PATCH 01/12] [SPARK-59398][ML][SQL] Add a SQL expression for ML vector affine transformations --- .../catalyst/analysis/FunctionRegistry.scala | 3 +- .../ml/VectorAffineTransform.scala | 240 ++++++++++++++++++ .../ml/VectorAffineTransformSuite.scala | 132 ++++++++++ 3 files changed, 374 insertions(+), 1 deletion(-) create mode 100644 sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorAffineTransform.scala create mode 100644 sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorAffineTransformSuite.scala diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/FunctionRegistry.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/FunctionRegistry.scala index a31e7ac4d74c7..8c67d994c7c87 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/FunctionRegistry.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/FunctionRegistry.scala @@ -1222,8 +1222,9 @@ object FunctionRegistry { registerInternalExpression[NullIndex]("null_index") registerInternalExpression[CastTimestampNTZToLong]("timestamp_ntz_to_long") registerInternalExpression[ArrayBinarySearch]("array_binary_search") - registerInternalExpression[VectorPosExplode]("ml_vector_posexplode") + registerInternalExpression[VectorAffineTransform]("ml_vector_affine_transform") registerInternalExpression[VectorDotProduct]("ml_vector_dot_product") + registerInternalExpression[VectorPosExplode]("ml_vector_posexplode") } registerInternalExpressions() diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorAffineTransform.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorAffineTransform.scala new file mode 100644 index 0000000000000..5d14576f2c182 --- /dev/null +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorAffineTransform.scala @@ -0,0 +1,240 @@ +/* + * 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.{ExpectsInputTypes, Expression, GenericInternalRow, TernaryExpression, UnsafeArrayData} +import org.apache.spark.sql.catalyst.expressions.codegen.{CodegenContext, ExprCode} +import org.apache.spark.sql.catalyst.util.ArrayData +import org.apache.spark.sql.types._ + +/** + * Applies an element-wise affine transformation to SQL struct representations of MLlib vectors: + * `vector(i) * scale(i) + shift(i)`. This expression is dedicated only for Spark ML and should be + * used together with `unwrap_udt` and `wrap_udt`. + */ +case class VectorAffineTransform( + vector: Expression, + scale: Expression, + shift: Expression) + extends TernaryExpression with ExpectsInputTypes { + + override def nullIntolerant: Boolean = true + + override def first: Expression = vector + override def second: Expression = scale + override def third: Expression = shift + + override def prettyName: String = "ml_vector_affine_transform" + + override def inputTypes: Seq[AbstractDataType] = Seq.fill(3)(VectorAffineTransform.vectorSqlType) + + override def dataType: DataType = VectorAffineTransform.vectorSqlType + + override protected def nullSafeEval( + vectorInput: Any, + scaleInput: Any, + shiftInput: Any): Any = { + VectorAffineTransform.transform( + vectorInput.asInstanceOf[InternalRow], + scaleInput.asInstanceOf[InternalRow], + shiftInput.asInstanceOf[InternalRow]) + } + + override protected def doGenCode(ctx: CodegenContext, ev: ExprCode): ExprCode = { + val cls = VectorAffineTransform.getClass.getName + defineCodeGen(ctx, ev, (vectorInput, scaleInput, shiftInput) => + s"$cls.MODULE$$.transform($vectorInput, $scaleInput, $shiftInput)") + } + + override protected def withNewChildrenInternal( + newVector: Expression, + newScale: Expression, + newShift: Expression): VectorAffineTransform = { + copy(vector = newVector, scale = newScale, shift = newShift) + } +} + +object VectorAffineTransform { + private val SparseVectorType: Byte = 0 + private val DenseVectorType: Byte = 1 + + private[ml] val vectorSqlType = StructType(Array( + StructField("type", ByteType, nullable = false), + StructField("size", IntegerType, nullable = true), + StructField("indices", ArrayType(IntegerType, containsNull = false), nullable = true), + StructField("values", ArrayType(DoubleType, containsNull = false), nullable = true))) + + private def vectorSize(vector: InternalRow, vectorType: Byte, values: ArrayData): Int = { + vectorType match { + case SparseVectorType => vector.getInt(1) + case DenseVectorType => values.numElements() + case _ => throw new IllegalArgumentException(s"Unknown vector type $vectorType.") + } + } + + private def isZeroVector(vectorType: Byte, values: ArrayData): Boolean = { + var index = 0 + while (index < values.numElements()) { + if (values.getDouble(index) != 0.0) return false + index += 1 + } + vectorType match { + case SparseVectorType | DenseVectorType => true + case _ => throw new IllegalArgumentException(s"Unknown vector type $vectorType.") + } + } + + private def sparseResult( + vector: InternalRow, + scale: InternalRow, + size: Int, + vectorValues: ArrayData, + scaleType: Byte, + scaleValues: ArrayData): InternalRow = { + val vectorIndices = vector.getArray(2) + val scaleIndices = if (scaleType == SparseVectorType) scale.getArray(2) else null + val resultValues = new Array[Double](vectorValues.numElements()) + var vectorIndex = 0 + var scaleIndex = 0 + while (vectorIndex < resultValues.length) { + val featureIndex = vectorIndices.getInt(vectorIndex) + val scaleValue = if (scaleType == DenseVectorType) { + scaleValues.getDouble(featureIndex) + } else { + while (scaleIndex < scaleValues.numElements() && + scaleIndices.getInt(scaleIndex) < featureIndex) { + scaleIndex += 1 + } + if (scaleIndex < scaleValues.numElements() && + scaleIndices.getInt(scaleIndex) == featureIndex) { + scaleValues.getDouble(scaleIndex) + } else { + 0.0 + } + } + resultValues(vectorIndex) = vectorValues.getDouble(vectorIndex) * scaleValue + vectorIndex += 1 + } + new GenericInternalRow(Array[Any]( + SparseVectorType, + size, + vectorIndices, + UnsafeArrayData.fromPrimitiveArray(resultValues))) + } + + private def denseResult( + vector: InternalRow, + scale: InternalRow, + shift: InternalRow, + size: Int, + vectorType: Byte, + scaleType: Byte, + shiftType: Byte, + vectorValues: ArrayData, + scaleValues: ArrayData, + shiftValues: ArrayData): InternalRow = { + val vectorIndices = if (vectorType == SparseVectorType) vector.getArray(2) else null + val scaleIndices = if (scaleType == SparseVectorType) scale.getArray(2) else null + val shiftIndices = if (shiftType == SparseVectorType) shift.getArray(2) else null + val resultValues = new Array[Double](size) + var vectorIndex = 0 + var scaleIndex = 0 + var shiftIndex = 0 + var featureIndex = 0 + while (featureIndex < size) { + val vectorIsActive = vectorType == DenseVectorType || + (vectorIndex < vectorValues.numElements() && + vectorIndices.getInt(vectorIndex) == featureIndex) + val vectorValue = if (vectorType == DenseVectorType) { + vectorValues.getDouble(featureIndex) + } else if (vectorIsActive) { + val value = vectorValues.getDouble(vectorIndex) + vectorIndex += 1 + value + } else { + 0.0 + } + + val scaleValue = if (scaleType == DenseVectorType) { + scaleValues.getDouble(featureIndex) + } else if (scaleIndex < scaleValues.numElements() && + scaleIndices.getInt(scaleIndex) == featureIndex) { + val value = scaleValues.getDouble(scaleIndex) + scaleIndex += 1 + value + } else { + 0.0 + } + + val shiftValue = if (shiftType == DenseVectorType) { + shiftValues.getDouble(featureIndex) + } else if (shiftIndex < shiftValues.numElements() && + shiftIndices.getInt(shiftIndex) == featureIndex) { + val value = shiftValues.getDouble(shiftIndex) + shiftIndex += 1 + value + } else { + 0.0 + } + + resultValues(featureIndex) = + (if (vectorIsActive) vectorValue * scaleValue else 0.0) + shiftValue + featureIndex += 1 + } + new GenericInternalRow(Array[Any]( + DenseVectorType, + null, + null, + UnsafeArrayData.fromPrimitiveArray(resultValues))) + } + + private[ml] def transform( + vector: InternalRow, + scale: InternalRow, + shift: InternalRow): InternalRow = { + val vectorType = vector.getByte(0) + val scaleType = scale.getByte(0) + val shiftType = shift.getByte(0) + val vectorValues = vector.getArray(3) + val scaleValues = scale.getArray(3) + val shiftValues = shift.getArray(3) + val size = vectorSize(vector, vectorType, vectorValues) + val scaleSize = vectorSize(scale, scaleType, scaleValues) + val shiftSize = vectorSize(shift, shiftType, shiftValues) + require(size == scaleSize && size == shiftSize, + "VectorAffineTransform was given vectors with non-matching sizes:" + + s" vector.size = $size, scale.size = $scaleSize, shift.size = $shiftSize") + + if (vectorType == SparseVectorType && isZeroVector(shiftType, shiftValues)) { + sparseResult(vector, scale, size, vectorValues, scaleType, scaleValues) + } else { + denseResult( + vector, + scale, + shift, + size, + vectorType, + scaleType, + shiftType, + vectorValues, + scaleValues, + shiftValues) + } + } +} diff --git a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorAffineTransformSuite.scala b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorAffineTransformSuite.scala new file mode 100644 index 0000000000000..0703f401b235b --- /dev/null +++ b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorAffineTransformSuite.scala @@ -0,0 +1,132 @@ +/* + * 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.SparkFunSuite +import org.apache.spark.sql.catalyst.InternalRow +import org.apache.spark.sql.catalyst.expressions.{ExpressionEvalHelper, GenericInternalRow, Literal, UnsafeArrayData} + +class VectorAffineTransformSuite extends SparkFunSuite with ExpressionEvalHelper { + private val vectorSqlType = VectorAffineTransform.vectorSqlType + + private def denseRow(values: Double*): InternalRow = { + new GenericInternalRow(Array[Any]( + 1.toByte, + null, + null, + UnsafeArrayData.fromPrimitiveArray(values.toArray))) + } + + private def dense(values: Double*): Literal = Literal(denseRow(values: _*), vectorSqlType) + + private def sparseRow(size: Int, indices: Array[Int], values: Array[Double]): InternalRow = { + new GenericInternalRow(Array[Any]( + 0.toByte, + size, + UnsafeArrayData.fromPrimitiveArray(indices), + UnsafeArrayData.fromPrimitiveArray(values))) + } + + private def sparse(size: Int, indices: Array[Int], values: Array[Double]): Literal = { + Literal(sparseRow(size, indices, values), vectorSqlType) + } + + test("vector affine transform interpreted and code-generated evaluation") { + val expression = VectorAffineTransform( + dense(1.0, 2.0, 3.0), + dense(2.0, 3.0, 4.0), + dense(5.0, 6.0, 7.0)) + assert(expression.prettyName === "ml_vector_affine_transform") + checkEvaluation(expression, denseRow(7.0, 12.0, 19.0)) + + checkEvaluation( + VectorAffineTransform( + dense(1.0, 2.0, 3.0), + sparse(3, Array(0, 2), Array(2.0, 4.0)), + sparse(3, Array(1), Array(1.0))), + denseRow(2.0, 1.0, 12.0)) + } + + test("vector affine transform preserves sparse vectors for a zero shift") { + val vector = sparse(3, Array(0, 2), Array(1.0, 3.0)) + val expected = sparseRow(3, Array(0, 2), Array(2.0, 12.0)) + + checkEvaluation( + VectorAffineTransform(vector, dense(2.0, 3.0, 4.0), dense(0.0, 0.0, 0.0)), + expected) + checkEvaluation( + VectorAffineTransform( + vector, + sparse(3, Array(0, 2), Array(2.0, 4.0)), + sparse(3, Array.emptyIntArray, Array.emptyDoubleArray)), + expected) + } + + test("vector affine transform produces a dense vector for a nonzero shift") { + checkEvaluation( + VectorAffineTransform( + sparse(3, Array(0, 2), Array(1.0, 3.0)), + dense(2.0, 3.0, 4.0), + sparse(3, Array(1), Array(1.0))), + denseRow(2.0, 1.0, 12.0)) + } + + test("vector affine transform with null vectors") { + val nullVector = Literal(null, vectorSqlType) + val vector = dense(1.0) + + checkEvaluation(VectorAffineTransform(nullVector, vector, vector), null) + checkEvaluation(VectorAffineTransform(vector, nullVector, vector), null) + checkEvaluation(VectorAffineTransform(vector, vector, nullVector), null) + } + + test("vector affine transform with empty vectors") { + val emptyDense = dense() + val emptySparse = sparse(0, Array.emptyIntArray, Array.emptyDoubleArray) + + checkEvaluation( + VectorAffineTransform(emptyDense, emptyDense, emptyDense), + denseRow()) + checkEvaluation( + VectorAffineTransform(emptySparse, emptySparse, emptySparse), + sparseRow(0, Array.emptyIntArray, Array.emptyDoubleArray)) + } + + test("vector affine transform with infinite and NaN values") { + Seq(Double.PositiveInfinity, Double.NegativeInfinity, Double.NaN).foreach { value => + checkEvaluation( + VectorAffineTransform(dense(value), dense(1.0), dense(0.0)), + denseRow(value)) + checkEvaluation( + VectorAffineTransform(dense(1.0), dense(value), dense(0.0)), + denseRow(value)) + checkEvaluation( + VectorAffineTransform(dense(1.0), dense(1.0), dense(value)), + denseRow(value)) + } + } + + test("vector affine transform rejects vectors with different sizes") { + checkExceptionInExpression[IllegalArgumentException]( + VectorAffineTransform(dense(1.0), dense(1.0, 2.0), dense(1.0)), + "vectors with non-matching sizes") + checkExceptionInExpression[IllegalArgumentException]( + VectorAffineTransform(dense(1.0), dense(1.0), dense(1.0, 2.0)), + "vectors with non-matching sizes") + } +} From 94d60e13bdaa2d54e840db217df0caddeb254f8e Mon Sep 17 00:00:00 2001 From: Ruifeng Zheng Date: Thu, 10 Sep 2026 10:29:15 +0000 Subject: [PATCH 02/12] [SPARK-59398][ML][SQL] Support optional scale and shift --- .../ml/VectorAffineTransform.scala | 121 ++++++++++++------ .../ml/VectorAffineTransformSuite.scala | 43 ++++++- 2 files changed, 125 insertions(+), 39 deletions(-) diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorAffineTransform.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorAffineTransform.scala index 5d14576f2c182..4f1fa63ab3d53 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorAffineTransform.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorAffineTransform.scala @@ -18,55 +18,93 @@ package org.apache.spark.sql.catalyst.expressions.ml import org.apache.spark.sql.catalyst.InternalRow -import org.apache.spark.sql.catalyst.expressions.{ExpectsInputTypes, Expression, GenericInternalRow, TernaryExpression, UnsafeArrayData} -import org.apache.spark.sql.catalyst.expressions.codegen.{CodegenContext, ExprCode} +import org.apache.spark.sql.catalyst.expressions.{ExpectsInputTypes, Expression, GenericInternalRow, UnsafeArrayData} +import org.apache.spark.sql.catalyst.expressions.codegen.{CodegenContext, CodeGenerator, ExprCode} +import org.apache.spark.sql.catalyst.expressions.codegen.Block._ import org.apache.spark.sql.catalyst.util.ArrayData import org.apache.spark.sql.types._ /** * Applies an element-wise affine transformation to SQL struct representations of MLlib vectors: * `vector(i) * scale(i) + shift(i)`. This expression is dedicated only for Spark ML and should be - * used together with `unwrap_udt` and `wrap_udt`. + * used together with `unwrap_udt` and `wrap_udt`. A null scale is treated as an identity scale, and + * a null shift is treated as a zero shift. The scale and shift cannot both be null. */ case class VectorAffineTransform( vector: Expression, scale: Expression, shift: Expression) - extends TernaryExpression with ExpectsInputTypes { + extends Expression with ExpectsInputTypes { - override def nullIntolerant: Boolean = true + require(scale != null || shift != null, "The scale and shift cannot both be null.") - override def first: Expression = vector - override def second: Expression = scale - override def third: Expression = shift + override def children: Seq[Expression] = + Seq(vector) ++ Option(scale) ++ Option(shift) override def prettyName: String = "ml_vector_affine_transform" - override def inputTypes: Seq[AbstractDataType] = Seq.fill(3)(VectorAffineTransform.vectorSqlType) + override def inputTypes: Seq[AbstractDataType] = + Seq.fill(children.length)(VectorAffineTransform.vectorSqlType) override def dataType: DataType = VectorAffineTransform.vectorSqlType - override protected def nullSafeEval( - vectorInput: Any, - scaleInput: Any, - shiftInput: Any): Any = { - VectorAffineTransform.transform( - vectorInput.asInstanceOf[InternalRow], - scaleInput.asInstanceOf[InternalRow], - shiftInput.asInstanceOf[InternalRow]) + override def nullable: Boolean = vector.nullable + + override def eval(input: InternalRow): Any = { + val vectorInput = vector.eval(input) + if (vectorInput == null) { + null + } else { + VectorAffineTransform.transform( + vectorInput.asInstanceOf[InternalRow], + if (scale == null) null else scale.eval(input).asInstanceOf[InternalRow], + if (shift == null) null else shift.eval(input).asInstanceOf[InternalRow]) + } } override protected def doGenCode(ctx: CodegenContext, ev: ExprCode): ExprCode = { val cls = VectorAffineTransform.getClass.getName - defineCodeGen(ctx, ev, (vectorInput, scaleInput, shiftInput) => - s"$cls.MODULE$$.transform($vectorInput, $scaleInput, $shiftInput)") + val javaType = CodeGenerator.javaType(dataType) + val vectorGen = vector.genCode(ctx) + val scaleInput = ctx.freshName("scaleInput") + val shiftInput = ctx.freshName("shiftInput") + val scaleGen = Option(scale).map(_.genCode(ctx)) + val shiftGen = Option(shift).map(_.genCode(ctx)) + val scaleEval = scaleGen.map { gen => + code""" + ${gen.code} + $javaType $scaleInput = ${gen.isNull} ? null : ${gen.value}; + """ + }.getOrElse(code"$javaType $scaleInput = null;") + val shiftEval = shiftGen.map { gen => + code""" + ${gen.code} + $javaType $shiftInput = ${gen.isNull} ? null : ${gen.value}; + """ + }.getOrElse(code"$javaType $shiftInput = null;") + + ev.copy(code = code""" + ${vectorGen.code} + boolean ${ev.isNull} = ${vectorGen.isNull}; + $javaType ${ev.value} = null; + if (!${ev.isNull}) { + $scaleEval + $shiftEval + ${ev.value} = $cls.MODULE$$.transform(${vectorGen.value}, $scaleInput, $shiftInput); + } + """) } override protected def withNewChildrenInternal( - newVector: Expression, - newScale: Expression, - newShift: Expression): VectorAffineTransform = { - copy(vector = newVector, scale = newScale, shift = newShift) + newChildren: IndexedSeq[Expression]): VectorAffineTransform = { + var index = 1 + val newScale = if (scale == null) null else { + val child = newChildren(index) + index += 1 + child + } + val newShift = if (shift == null) null else newChildren(index) + copy(vector = newChildren.head, scale = newScale, shift = newShift) } } @@ -108,13 +146,16 @@ object VectorAffineTransform { scaleType: Byte, scaleValues: ArrayData): InternalRow = { val vectorIndices = vector.getArray(2) - val scaleIndices = if (scaleType == SparseVectorType) scale.getArray(2) else null + val scaleIndices = + if (scale != null && scaleType == SparseVectorType) scale.getArray(2) else null val resultValues = new Array[Double](vectorValues.numElements()) var vectorIndex = 0 var scaleIndex = 0 while (vectorIndex < resultValues.length) { val featureIndex = vectorIndices.getInt(vectorIndex) - val scaleValue = if (scaleType == DenseVectorType) { + val scaleValue = if (scale == null) { + 1.0 + } else if (scaleType == DenseVectorType) { scaleValues.getDouble(featureIndex) } else { while (scaleIndex < scaleValues.numElements() && @@ -150,8 +191,10 @@ object VectorAffineTransform { scaleValues: ArrayData, shiftValues: ArrayData): InternalRow = { val vectorIndices = if (vectorType == SparseVectorType) vector.getArray(2) else null - val scaleIndices = if (scaleType == SparseVectorType) scale.getArray(2) else null - val shiftIndices = if (shiftType == SparseVectorType) shift.getArray(2) else null + val scaleIndices = + if (scale != null && scaleType == SparseVectorType) scale.getArray(2) else null + val shiftIndices = + if (shift != null && shiftType == SparseVectorType) shift.getArray(2) else null val resultValues = new Array[Double](size) var vectorIndex = 0 var scaleIndex = 0 @@ -171,7 +214,9 @@ object VectorAffineTransform { 0.0 } - val scaleValue = if (scaleType == DenseVectorType) { + val scaleValue = if (scale == null) { + 1.0 + } else if (scaleType == DenseVectorType) { scaleValues.getDouble(featureIndex) } else if (scaleIndex < scaleValues.numElements() && scaleIndices.getInt(scaleIndex) == featureIndex) { @@ -182,7 +227,9 @@ object VectorAffineTransform { 0.0 } - val shiftValue = if (shiftType == DenseVectorType) { + val shiftValue = if (shift == null) { + 0.0 + } else if (shiftType == DenseVectorType) { shiftValues.getDouble(featureIndex) } else if (shiftIndex < shiftValues.numElements() && shiftIndices.getInt(shiftIndex) == featureIndex) { @@ -208,20 +255,22 @@ object VectorAffineTransform { vector: InternalRow, scale: InternalRow, shift: InternalRow): InternalRow = { + require(scale != null || shift != null, "The scale and shift cannot both be null.") val vectorType = vector.getByte(0) - val scaleType = scale.getByte(0) - val shiftType = shift.getByte(0) + val scaleType = if (scale == null) DenseVectorType else scale.getByte(0) + val shiftType = if (shift == null) DenseVectorType else shift.getByte(0) val vectorValues = vector.getArray(3) - val scaleValues = scale.getArray(3) - val shiftValues = shift.getArray(3) + val scaleValues = if (scale == null) null else scale.getArray(3) + val shiftValues = if (shift == null) null else shift.getArray(3) val size = vectorSize(vector, vectorType, vectorValues) - val scaleSize = vectorSize(scale, scaleType, scaleValues) - val shiftSize = vectorSize(shift, shiftType, shiftValues) + val scaleSize = if (scale == null) size else vectorSize(scale, scaleType, scaleValues) + val shiftSize = if (shift == null) size else vectorSize(shift, shiftType, shiftValues) require(size == scaleSize && size == shiftSize, "VectorAffineTransform was given vectors with non-matching sizes:" + s" vector.size = $size, scale.size = $scaleSize, shift.size = $shiftSize") - if (vectorType == SparseVectorType && isZeroVector(shiftType, shiftValues)) { + if (vectorType == SparseVectorType && + (shift == null || isZeroVector(shiftType, shiftValues))) { sparseResult(vector, scale, size, vectorValues, scaleType, scaleValues) } else { denseResult( diff --git a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorAffineTransformSuite.scala b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorAffineTransformSuite.scala index 0703f401b235b..a20db82116047 100644 --- a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorAffineTransformSuite.scala +++ b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorAffineTransformSuite.scala @@ -86,13 +86,50 @@ class VectorAffineTransformSuite extends SparkFunSuite with ExpressionEvalHelper denseRow(2.0, 1.0, 12.0)) } - test("vector affine transform with null vectors") { + test("vector affine transform with a null vector") { val nullVector = Literal(null, vectorSqlType) val vector = dense(1.0) checkEvaluation(VectorAffineTransform(nullVector, vector, vector), null) - checkEvaluation(VectorAffineTransform(vector, nullVector, vector), null) - checkEvaluation(VectorAffineTransform(vector, vector, nullVector), null) + } + + test("vector affine transform with a null scale") { + val nullVector = Literal(null, vectorSqlType) + + checkEvaluation( + VectorAffineTransform(dense(1.0, 2.0), nullVector, dense(3.0, 4.0)), + denseRow(4.0, 6.0)) + checkEvaluation( + VectorAffineTransform( + sparse(3, Array(0, 2), Array(1.0, 3.0)), + null, + sparse(3, Array(1), Array(2.0))), + denseRow(1.0, 2.0, 3.0)) + } + + test("vector affine transform with a null shift") { + val nullVector = Literal(null, vectorSqlType) + + checkEvaluation( + VectorAffineTransform(dense(1.0, 2.0), dense(3.0, 4.0), nullVector), + denseRow(3.0, 8.0)) + checkEvaluation( + VectorAffineTransform( + sparse(3, Array(0, 2), Array(1.0, 3.0)), + dense(2.0, 3.0, 4.0), + null), + sparseRow(3, Array(0, 2), Array(2.0, 12.0))) + } + + test("vector affine transform rejects a null scale and shift") { + val nullVector = Literal(null, vectorSqlType) + checkExceptionInExpression[IllegalArgumentException]( + VectorAffineTransform(dense(1.0), nullVector, nullVector), + "scale and shift cannot both be null") + val error = intercept[IllegalArgumentException] { + VectorAffineTransform(dense(1.0), null, null) + } + assert(error.getMessage.contains("scale and shift cannot both be null")) } test("vector affine transform with empty vectors") { From 55b8a61b3e378aa97e151b2cc9ec9c3f0531946d Mon Sep 17 00:00:00 2001 From: Ruifeng Zheng Date: Thu, 10 Sep 2026 10:37:34 +0000 Subject: [PATCH 03/12] [SPARK-59398][ML][SQL] Treat null scale and shift values as optional --- .../ml/VectorAffineTransform.scala | 53 +++++++------------ .../ml/VectorAffineTransformSuite.scala | 8 +-- 2 files changed, 20 insertions(+), 41 deletions(-) diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorAffineTransform.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorAffineTransform.scala index 4f1fa63ab3d53..2c5fc58d74de4 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorAffineTransform.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorAffineTransform.scala @@ -18,7 +18,7 @@ package org.apache.spark.sql.catalyst.expressions.ml import org.apache.spark.sql.catalyst.InternalRow -import org.apache.spark.sql.catalyst.expressions.{ExpectsInputTypes, Expression, GenericInternalRow, UnsafeArrayData} +import org.apache.spark.sql.catalyst.expressions.{ExpectsInputTypes, Expression, GenericInternalRow, TernaryExpression, UnsafeArrayData} import org.apache.spark.sql.catalyst.expressions.codegen.{CodegenContext, CodeGenerator, ExprCode} import org.apache.spark.sql.catalyst.expressions.codegen.Block._ import org.apache.spark.sql.catalyst.util.ArrayData @@ -34,17 +34,15 @@ case class VectorAffineTransform( vector: Expression, scale: Expression, shift: Expression) - extends Expression with ExpectsInputTypes { + extends TernaryExpression with ExpectsInputTypes { - require(scale != null || shift != null, "The scale and shift cannot both be null.") - - override def children: Seq[Expression] = - Seq(vector) ++ Option(scale) ++ Option(shift) + override def first: Expression = vector + override def second: Expression = scale + override def third: Expression = shift override def prettyName: String = "ml_vector_affine_transform" - override def inputTypes: Seq[AbstractDataType] = - Seq.fill(children.length)(VectorAffineTransform.vectorSqlType) + override def inputTypes: Seq[AbstractDataType] = Seq.fill(3)(VectorAffineTransform.vectorSqlType) override def dataType: DataType = VectorAffineTransform.vectorSqlType @@ -57,8 +55,8 @@ case class VectorAffineTransform( } else { VectorAffineTransform.transform( vectorInput.asInstanceOf[InternalRow], - if (scale == null) null else scale.eval(input).asInstanceOf[InternalRow], - if (shift == null) null else shift.eval(input).asInstanceOf[InternalRow]) + scale.eval(input).asInstanceOf[InternalRow], + shift.eval(input).asInstanceOf[InternalRow]) } } @@ -68,43 +66,28 @@ case class VectorAffineTransform( val vectorGen = vector.genCode(ctx) val scaleInput = ctx.freshName("scaleInput") val shiftInput = ctx.freshName("shiftInput") - val scaleGen = Option(scale).map(_.genCode(ctx)) - val shiftGen = Option(shift).map(_.genCode(ctx)) - val scaleEval = scaleGen.map { gen => - code""" - ${gen.code} - $javaType $scaleInput = ${gen.isNull} ? null : ${gen.value}; - """ - }.getOrElse(code"$javaType $scaleInput = null;") - val shiftEval = shiftGen.map { gen => - code""" - ${gen.code} - $javaType $shiftInput = ${gen.isNull} ? null : ${gen.value}; - """ - }.getOrElse(code"$javaType $shiftInput = null;") + val scaleGen = scale.genCode(ctx) + val shiftGen = shift.genCode(ctx) ev.copy(code = code""" ${vectorGen.code} boolean ${ev.isNull} = ${vectorGen.isNull}; $javaType ${ev.value} = null; if (!${ev.isNull}) { - $scaleEval - $shiftEval + ${scaleGen.code} + ${shiftGen.code} + $javaType $scaleInput = ${scaleGen.isNull} ? null : ${scaleGen.value}; + $javaType $shiftInput = ${shiftGen.isNull} ? null : ${shiftGen.value}; ${ev.value} = $cls.MODULE$$.transform(${vectorGen.value}, $scaleInput, $shiftInput); } """) } override protected def withNewChildrenInternal( - newChildren: IndexedSeq[Expression]): VectorAffineTransform = { - var index = 1 - val newScale = if (scale == null) null else { - val child = newChildren(index) - index += 1 - child - } - val newShift = if (shift == null) null else newChildren(index) - copy(vector = newChildren.head, scale = newScale, shift = newShift) + newVector: Expression, + newScale: Expression, + newShift: Expression): VectorAffineTransform = { + copy(vector = newVector, scale = newScale, shift = newShift) } } diff --git a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorAffineTransformSuite.scala b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorAffineTransformSuite.scala index a20db82116047..815783815dd8b 100644 --- a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorAffineTransformSuite.scala +++ b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorAffineTransformSuite.scala @@ -102,7 +102,7 @@ class VectorAffineTransformSuite extends SparkFunSuite with ExpressionEvalHelper checkEvaluation( VectorAffineTransform( sparse(3, Array(0, 2), Array(1.0, 3.0)), - null, + nullVector, sparse(3, Array(1), Array(2.0))), denseRow(1.0, 2.0, 3.0)) } @@ -117,7 +117,7 @@ class VectorAffineTransformSuite extends SparkFunSuite with ExpressionEvalHelper VectorAffineTransform( sparse(3, Array(0, 2), Array(1.0, 3.0)), dense(2.0, 3.0, 4.0), - null), + nullVector), sparseRow(3, Array(0, 2), Array(2.0, 12.0))) } @@ -126,10 +126,6 @@ class VectorAffineTransformSuite extends SparkFunSuite with ExpressionEvalHelper checkExceptionInExpression[IllegalArgumentException]( VectorAffineTransform(dense(1.0), nullVector, nullVector), "scale and shift cannot both be null") - val error = intercept[IllegalArgumentException] { - VectorAffineTransform(dense(1.0), null, null) - } - assert(error.getMessage.contains("scale and shift cannot both be null")) } test("vector affine transform with empty vectors") { From 6fb9ea168a11202874ec14aec7e85a3c50729cad Mon Sep 17 00:00:00 2001 From: Ruifeng Zheng Date: Thu, 10 Sep 2026 10:48:49 +0000 Subject: [PATCH 04/12] [SPARK-59398][ML][SQL] Add ML helpers for vector affine transform --- .../scala/org/apache/spark/ml/functions.scala | 58 ++++++++++++++----- .../org/apache/spark/ml/FunctionsSuite.scala | 43 +++++++++++++- .../ml/VectorAffineTransform.scala | 5 +- .../ml/VectorAffineTransformSuite.scala | 15 +++-- 4 files changed, 98 insertions(+), 23 deletions(-) diff --git a/mllib/src/main/scala/org/apache/spark/ml/functions.scala b/mllib/src/main/scala/org/apache/spark/ml/functions.scala index 87a8a7d98ea0c..0ca8f12b661f3 100644 --- a/mllib/src/main/scala/org/apache/spark/ml/functions.scala +++ b/mllib/src/main/scala/org/apache/spark/ml/functions.scala @@ -18,7 +18,7 @@ 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} @@ -65,24 +65,50 @@ 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), + sf.unwrap_udt(scale), + sf.unwrap_udt(shift)) + sf.wrap_udt(transformed, new VectorUDT) + } + + private[ml] def vector_affine_transform( + vector: Column, + scale: Vector, + shift: Vector): Column = { + val transformed = Column.internalFn( + "ml_vector_affine_transform", + sf.unwrap_udt(vector), + vectorToStruct(scale), + vectorToStruct(shift)) + sf.wrap_udt(transformed, new VectorUDT) + } + + 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 = diff --git a/mllib/src/test/scala/org/apache/spark/ml/FunctionsSuite.scala b/mllib/src/test/scala/org/apache/spark/ml/FunctionsSuite.scala index ddc70e532b828..d868c11581c38 100644 --- a/mllib/src/test/scala/org/apache/spark/ml/FunctionsSuite.scala +++ b/mllib/src/test/scala/org/apache/spark/ml/FunctionsSuite.scala @@ -24,7 +24,7 @@ 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.catalyst.expressions.ml.{VectorAffineTransform, VectorPosExplode} import org.apache.spark.sql.functions.{col, unwrap_udt, wrap_udt} import org.apache.spark.sql.types.{StructField, StructType, UserDefinedType} @@ -256,6 +256,47 @@ 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), Vectors.dense(2.0, 3.0), Vectors.dense(4.0, 5.0)), + (Vectors.sparse(2, Seq((0, 1.0))), Vectors.dense(2.0, 3.0), null), + (Vectors.dense(1.0, 2.0), null, Vectors.dense(4.0, 5.0)), + (Vectors.sparse(2, Seq((0, 1.0))), null, null), + (null, Vectors.dense(2.0, 3.0), Vectors.dense(4.0, 5.0))) + .toDF("vector", "scale", "shift") + + 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.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", + Vectors.dense(2.0, 3.0), + Vectors.dense(4.0, 5.0))) + .first() + .getAs[Vector](0) + assert(constantResult === Vectors.dense(6.0, 11.0)) + + val nullConstantsResult = df.limit(1) + .select(vector_affine_transform( + $"vector", + null.asInstanceOf[Vector], + null.asInstanceOf[Vector])) + .first() + .getAs[Vector](0) + assert(nullConstantsResult === Vectors.dense(1.0, 2.0)) + } + test("test get_vector") { val df = Seq( (Vectors.dense(1.0, 2.0, 3.0), 0), diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorAffineTransform.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorAffineTransform.scala index 2c5fc58d74de4..d2f8adb346df2 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorAffineTransform.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorAffineTransform.scala @@ -28,7 +28,8 @@ import org.apache.spark.sql.types._ * Applies an element-wise affine transformation to SQL struct representations of MLlib vectors: * `vector(i) * scale(i) + shift(i)`. This expression is dedicated only for Spark ML and should be * used together with `unwrap_udt` and `wrap_udt`. A null scale is treated as an identity scale, and - * a null shift is treated as a zero shift. The scale and shift cannot both be null. + * a null shift is treated as a zero shift. If both are null, the input vector is returned + * unchanged. */ case class VectorAffineTransform( vector: Expression, @@ -238,7 +239,7 @@ object VectorAffineTransform { vector: InternalRow, scale: InternalRow, shift: InternalRow): InternalRow = { - require(scale != null || shift != null, "The scale and shift cannot both be null.") + if (scale == null && shift == null) return vector val vectorType = vector.getByte(0) val scaleType = if (scale == null) DenseVectorType else scale.getByte(0) val shiftType = if (shift == null) DenseVectorType else shift.getByte(0) diff --git a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorAffineTransformSuite.scala b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorAffineTransformSuite.scala index 815783815dd8b..702be124ba3cb 100644 --- a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorAffineTransformSuite.scala +++ b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorAffineTransformSuite.scala @@ -121,11 +121,18 @@ class VectorAffineTransformSuite extends SparkFunSuite with ExpressionEvalHelper sparseRow(3, Array(0, 2), Array(2.0, 12.0))) } - test("vector affine transform rejects a null scale and shift") { + test("vector affine transform with a null scale and shift") { val nullVector = Literal(null, vectorSqlType) - checkExceptionInExpression[IllegalArgumentException]( - VectorAffineTransform(dense(1.0), nullVector, nullVector), - "scale and shift cannot both be null") + + checkEvaluation( + VectorAffineTransform(dense(1.0, 2.0), nullVector, nullVector), + denseRow(1.0, 2.0)) + checkEvaluation( + VectorAffineTransform( + sparse(3, Array(0, 2), Array(1.0, 3.0)), + nullVector, + nullVector), + sparseRow(3, Array(0, 2), Array(1.0, 3.0))) } test("vector affine transform with empty vectors") { From 4263e3ae4d841c903cd58ca753cb58ead55b4cac Mon Sep 17 00:00:00 2001 From: Ruifeng Zheng Date: Thu, 10 Sep 2026 11:20:02 +0000 Subject: [PATCH 05/12] [SPARK-59398][ML][SQL] Use arrays for affine scale and shift --- .../scala/org/apache/spark/ml/functions.scala | 22 ++-- .../org/apache/spark/ml/FunctionsSuite.scala | 27 ++-- .../ml/VectorAffineTransform.scala | 118 ++++++------------ .../ml/VectorAffineTransformSuite.scala | 85 ++++++++----- 4 files changed, 121 insertions(+), 131 deletions(-) diff --git a/mllib/src/main/scala/org/apache/spark/ml/functions.scala b/mllib/src/main/scala/org/apache/spark/ml/functions.scala index 0ca8f12b661f3..869ba15e6ab2b 100644 --- a/mllib/src/main/scala/org/apache/spark/ml/functions.scala +++ b/mllib/src/main/scala/org/apache/spark/ml/functions.scala @@ -21,7 +21,7 @@ import org.apache.spark.annotation.Since 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") @@ -78,23 +78,31 @@ object functions { val transformed = Column.internalFn( "ml_vector_affine_transform", sf.unwrap_udt(vector), - sf.unwrap_udt(scale), - sf.unwrap_udt(shift)) + scale, + shift) sf.wrap_udt(transformed, new VectorUDT) } private[ml] def vector_affine_transform( vector: Column, - scale: Vector, - shift: Vector): Column = { + scale: Array[Double], + shift: Array[Double]): Column = { val transformed = Column.internalFn( "ml_vector_affine_transform", sf.unwrap_udt(vector), - vectorToStruct(scale), - vectorToStruct(shift)) + 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 => diff --git a/mllib/src/test/scala/org/apache/spark/ml/FunctionsSuite.scala b/mllib/src/test/scala/org/apache/spark/ml/FunctionsSuite.scala index d868c11581c38..905405ac0c703 100644 --- a/mllib/src/test/scala/org/apache/spark/ml/FunctionsSuite.scala +++ b/mllib/src/test/scala/org/apache/spark/ml/FunctionsSuite.scala @@ -26,7 +26,7 @@ import org.apache.spark.mllib.linalg.{Matrices => OldMatrices, MatrixUDT => OldM import org.apache.spark.sql.{AnalysisException, DataFrame, Row} import org.apache.spark.sql.catalyst.expressions.ml.{VectorAffineTransform, 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.types.{ArrayType, DoubleType, StructField, StructType, UserDefinedType} class FunctionsSuite extends MLTest { @@ -258,13 +258,15 @@ class FunctionsSuite extends MLTest { test("test vector_affine_transform") { val df = Seq( - (Vectors.dense(1.0, 2.0), Vectors.dense(2.0, 3.0), Vectors.dense(4.0, 5.0)), - (Vectors.sparse(2, Seq((0, 1.0))), Vectors.dense(2.0, 3.0), null), - (Vectors.dense(1.0, 2.0), null, Vectors.dense(4.0, 5.0)), + (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))), null, null), - (null, Vectors.dense(2.0, 3.0), Vectors.dense(4.0, 5.0))) + (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( @@ -281,8 +283,8 @@ class FunctionsSuite extends MLTest { val constantResult = df.limit(1) .select(vector_affine_transform( $"vector", - Vectors.dense(2.0, 3.0), - Vectors.dense(4.0, 5.0))) + Array(2.0, 3.0), + Array(4.0, 5.0))) .first() .getAs[Vector](0) assert(constantResult === Vectors.dense(6.0, 11.0)) @@ -290,11 +292,18 @@ class FunctionsSuite extends MLTest { val nullConstantsResult = df.limit(1) .select(vector_affine_transform( $"vector", - null.asInstanceOf[Vector], - null.asInstanceOf[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)) } test("test get_vector") { diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorAffineTransform.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorAffineTransform.scala index d2f8adb346df2..45618bced8e4b 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorAffineTransform.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorAffineTransform.scala @@ -43,7 +43,10 @@ case class VectorAffineTransform( override def prettyName: String = "ml_vector_affine_transform" - override def inputTypes: Seq[AbstractDataType] = Seq.fill(3)(VectorAffineTransform.vectorSqlType) + override def inputTypes: Seq[AbstractDataType] = Seq( + VectorAffineTransform.vectorSqlType, + VectorAffineTransform.NonNullableDoubleArrayType, + VectorAffineTransform.NonNullableDoubleArrayType) override def dataType: DataType = VectorAffineTransform.vectorSqlType @@ -56,14 +59,15 @@ case class VectorAffineTransform( } else { VectorAffineTransform.transform( vectorInput.asInstanceOf[InternalRow], - scale.eval(input).asInstanceOf[InternalRow], - shift.eval(input).asInstanceOf[InternalRow]) + scale.eval(input).asInstanceOf[ArrayData], + shift.eval(input).asInstanceOf[ArrayData]) } } override protected def doGenCode(ctx: CodegenContext, ev: ExprCode): ExprCode = { val cls = VectorAffineTransform.getClass.getName - val javaType = CodeGenerator.javaType(dataType) + val vectorJavaType = CodeGenerator.javaType(dataType) + val arrayJavaType = CodeGenerator.javaType(VectorAffineTransform.doubleArraySqlType) val vectorGen = vector.genCode(ctx) val scaleInput = ctx.freshName("scaleInput") val shiftInput = ctx.freshName("shiftInput") @@ -73,12 +77,12 @@ case class VectorAffineTransform( ev.copy(code = code""" ${vectorGen.code} boolean ${ev.isNull} = ${vectorGen.isNull}; - $javaType ${ev.value} = null; + $vectorJavaType ${ev.value} = null; if (!${ev.isNull}) { ${scaleGen.code} ${shiftGen.code} - $javaType $scaleInput = ${scaleGen.isNull} ? null : ${scaleGen.value}; - $javaType $shiftInput = ${shiftGen.isNull} ? null : ${shiftGen.value}; + $arrayJavaType $scaleInput = ${scaleGen.isNull} ? null : ${scaleGen.value}; + $arrayJavaType $shiftInput = ${shiftGen.isNull} ? null : ${shiftGen.value}; ${ev.value} = $cls.MODULE$$.transform(${vectorGen.value}, $scaleInput, $shiftInput); } """) @@ -102,6 +106,16 @@ object VectorAffineTransform { StructField("indices", ArrayType(IntegerType, containsNull = false), nullable = true), StructField("values", ArrayType(DoubleType, containsNull = false), nullable = true))) + private[ml] val doubleArraySqlType = ArrayType(DoubleType, containsNull = false) + + private object NonNullableDoubleArrayType extends AbstractDataType { + override private[sql] def defaultConcreteType: DataType = doubleArraySqlType + + override private[sql] def acceptsType(other: DataType): Boolean = other == doubleArraySqlType + + override private[spark] def simpleString: String = doubleArraySqlType.simpleString + } + private def vectorSize(vector: InternalRow, vectorType: Byte, values: ArrayData): Int = { vectorType match { case SparseVectorType => vector.getInt(1) @@ -110,48 +124,29 @@ object VectorAffineTransform { } } - private def isZeroVector(vectorType: Byte, values: ArrayData): Boolean = { + private def isZeroArray(values: ArrayData): Boolean = { var index = 0 while (index < values.numElements()) { if (values.getDouble(index) != 0.0) return false index += 1 } - vectorType match { - case SparseVectorType | DenseVectorType => true - case _ => throw new IllegalArgumentException(s"Unknown vector type $vectorType.") - } + true } private def sparseResult( vector: InternalRow, - scale: InternalRow, size: Int, vectorValues: ArrayData, - scaleType: Byte, - scaleValues: ArrayData): InternalRow = { + scale: ArrayData): InternalRow = { val vectorIndices = vector.getArray(2) - val scaleIndices = - if (scale != null && scaleType == SparseVectorType) scale.getArray(2) else null val resultValues = new Array[Double](vectorValues.numElements()) var vectorIndex = 0 - var scaleIndex = 0 while (vectorIndex < resultValues.length) { val featureIndex = vectorIndices.getInt(vectorIndex) val scaleValue = if (scale == null) { 1.0 - } else if (scaleType == DenseVectorType) { - scaleValues.getDouble(featureIndex) } else { - while (scaleIndex < scaleValues.numElements() && - scaleIndices.getInt(scaleIndex) < featureIndex) { - scaleIndex += 1 - } - if (scaleIndex < scaleValues.numElements() && - scaleIndices.getInt(scaleIndex) == featureIndex) { - scaleValues.getDouble(scaleIndex) - } else { - 0.0 - } + scale.getDouble(featureIndex) } resultValues(vectorIndex) = vectorValues.getDouble(vectorIndex) * scaleValue vectorIndex += 1 @@ -165,24 +160,14 @@ object VectorAffineTransform { private def denseResult( vector: InternalRow, - scale: InternalRow, - shift: InternalRow, size: Int, vectorType: Byte, - scaleType: Byte, - shiftType: Byte, vectorValues: ArrayData, - scaleValues: ArrayData, - shiftValues: ArrayData): InternalRow = { + scale: ArrayData, + shift: ArrayData): InternalRow = { val vectorIndices = if (vectorType == SparseVectorType) vector.getArray(2) else null - val scaleIndices = - if (scale != null && scaleType == SparseVectorType) scale.getArray(2) else null - val shiftIndices = - if (shift != null && shiftType == SparseVectorType) shift.getArray(2) else null val resultValues = new Array[Double](size) var vectorIndex = 0 - var scaleIndex = 0 - var shiftIndex = 0 var featureIndex = 0 while (featureIndex < size) { val vectorIsActive = vectorType == DenseVectorType || @@ -200,28 +185,14 @@ object VectorAffineTransform { val scaleValue = if (scale == null) { 1.0 - } else if (scaleType == DenseVectorType) { - scaleValues.getDouble(featureIndex) - } else if (scaleIndex < scaleValues.numElements() && - scaleIndices.getInt(scaleIndex) == featureIndex) { - val value = scaleValues.getDouble(scaleIndex) - scaleIndex += 1 - value } else { - 0.0 + scale.getDouble(featureIndex) } val shiftValue = if (shift == null) { 0.0 - } else if (shiftType == DenseVectorType) { - shiftValues.getDouble(featureIndex) - } else if (shiftIndex < shiftValues.numElements() && - shiftIndices.getInt(shiftIndex) == featureIndex) { - val value = shiftValues.getDouble(shiftIndex) - shiftIndex += 1 - value } else { - 0.0 + shift.getDouble(featureIndex) } resultValues(featureIndex) = @@ -237,37 +208,22 @@ object VectorAffineTransform { private[ml] def transform( vector: InternalRow, - scale: InternalRow, - shift: InternalRow): InternalRow = { + scale: ArrayData, + shift: ArrayData): InternalRow = { if (scale == null && shift == null) return vector val vectorType = vector.getByte(0) - val scaleType = if (scale == null) DenseVectorType else scale.getByte(0) - val shiftType = if (shift == null) DenseVectorType else shift.getByte(0) val vectorValues = vector.getArray(3) - val scaleValues = if (scale == null) null else scale.getArray(3) - val shiftValues = if (shift == null) null else shift.getArray(3) val size = vectorSize(vector, vectorType, vectorValues) - val scaleSize = if (scale == null) size else vectorSize(scale, scaleType, scaleValues) - val shiftSize = if (shift == null) size else vectorSize(shift, shiftType, shiftValues) + val scaleSize = if (scale == null) size else scale.numElements() + val shiftSize = if (shift == null) size else shift.numElements() require(size == scaleSize && size == shiftSize, - "VectorAffineTransform was given vectors with non-matching sizes:" + + "VectorAffineTransform was given inputs with non-matching sizes:" + s" vector.size = $size, scale.size = $scaleSize, shift.size = $shiftSize") - if (vectorType == SparseVectorType && - (shift == null || isZeroVector(shiftType, shiftValues))) { - sparseResult(vector, scale, size, vectorValues, scaleType, scaleValues) + if (vectorType == SparseVectorType && (shift == null || isZeroArray(shift))) { + sparseResult(vector, size, vectorValues, scale) } else { - denseResult( - vector, - scale, - shift, - size, - vectorType, - scaleType, - shiftType, - vectorValues, - scaleValues, - shiftValues) + denseResult(vector, size, vectorType, vectorValues, scale, shift) } } } diff --git a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorAffineTransformSuite.scala b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorAffineTransformSuite.scala index 702be124ba3cb..03e0943354270 100644 --- a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorAffineTransformSuite.scala +++ b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorAffineTransformSuite.scala @@ -20,9 +20,11 @@ package org.apache.spark.sql.catalyst.expressions.ml import org.apache.spark.SparkFunSuite import org.apache.spark.sql.catalyst.InternalRow import org.apache.spark.sql.catalyst.expressions.{ExpressionEvalHelper, GenericInternalRow, Literal, UnsafeArrayData} +import org.apache.spark.sql.types.{ArrayType, DoubleType} class VectorAffineTransformSuite extends SparkFunSuite with ExpressionEvalHelper { private val vectorSqlType = VectorAffineTransform.vectorSqlType + private val doubleArraySqlType = VectorAffineTransform.doubleArraySqlType private def denseRow(values: Double*): InternalRow = { new GenericInternalRow(Array[Any]( @@ -46,19 +48,23 @@ class VectorAffineTransformSuite extends SparkFunSuite with ExpressionEvalHelper Literal(sparseRow(size, indices, values), vectorSqlType) } + private def array(values: Double*): Literal = { + Literal(UnsafeArrayData.fromPrimitiveArray(values.toArray), doubleArraySqlType) + } + test("vector affine transform interpreted and code-generated evaluation") { val expression = VectorAffineTransform( dense(1.0, 2.0, 3.0), - dense(2.0, 3.0, 4.0), - dense(5.0, 6.0, 7.0)) + array(2.0, 3.0, 4.0), + array(5.0, 6.0, 7.0)) assert(expression.prettyName === "ml_vector_affine_transform") checkEvaluation(expression, denseRow(7.0, 12.0, 19.0)) checkEvaluation( VectorAffineTransform( dense(1.0, 2.0, 3.0), - sparse(3, Array(0, 2), Array(2.0, 4.0)), - sparse(3, Array(1), Array(1.0))), + array(2.0, 0.0, 4.0), + array(0.0, 1.0, 0.0)), denseRow(2.0, 1.0, 12.0)) } @@ -67,13 +73,13 @@ class VectorAffineTransformSuite extends SparkFunSuite with ExpressionEvalHelper val expected = sparseRow(3, Array(0, 2), Array(2.0, 12.0)) checkEvaluation( - VectorAffineTransform(vector, dense(2.0, 3.0, 4.0), dense(0.0, 0.0, 0.0)), + VectorAffineTransform(vector, array(2.0, 3.0, 4.0), array(0.0, 0.0, 0.0)), expected) checkEvaluation( VectorAffineTransform( vector, - sparse(3, Array(0, 2), Array(2.0, 4.0)), - sparse(3, Array.emptyIntArray, Array.emptyDoubleArray)), + array(2.0, 3.0, 4.0), + array(0.0, 0.0, 0.0)), expected) } @@ -81,92 +87,103 @@ class VectorAffineTransformSuite extends SparkFunSuite with ExpressionEvalHelper checkEvaluation( VectorAffineTransform( sparse(3, Array(0, 2), Array(1.0, 3.0)), - dense(2.0, 3.0, 4.0), - sparse(3, Array(1), Array(1.0))), + array(2.0, 3.0, 4.0), + array(0.0, 1.0, 0.0)), denseRow(2.0, 1.0, 12.0)) } test("vector affine transform with a null vector") { val nullVector = Literal(null, vectorSqlType) - val vector = dense(1.0) + val values = array(1.0) - checkEvaluation(VectorAffineTransform(nullVector, vector, vector), null) + checkEvaluation(VectorAffineTransform(nullVector, values, values), null) } test("vector affine transform with a null scale") { - val nullVector = Literal(null, vectorSqlType) + val nullArray = Literal(null, doubleArraySqlType) checkEvaluation( - VectorAffineTransform(dense(1.0, 2.0), nullVector, dense(3.0, 4.0)), + VectorAffineTransform(dense(1.0, 2.0), nullArray, array(3.0, 4.0)), denseRow(4.0, 6.0)) checkEvaluation( VectorAffineTransform( sparse(3, Array(0, 2), Array(1.0, 3.0)), - nullVector, - sparse(3, Array(1), Array(2.0))), + nullArray, + array(0.0, 2.0, 0.0)), denseRow(1.0, 2.0, 3.0)) } test("vector affine transform with a null shift") { - val nullVector = Literal(null, vectorSqlType) + val nullArray = Literal(null, doubleArraySqlType) checkEvaluation( - VectorAffineTransform(dense(1.0, 2.0), dense(3.0, 4.0), nullVector), + VectorAffineTransform(dense(1.0, 2.0), array(3.0, 4.0), nullArray), denseRow(3.0, 8.0)) checkEvaluation( VectorAffineTransform( sparse(3, Array(0, 2), Array(1.0, 3.0)), - dense(2.0, 3.0, 4.0), - nullVector), + array(2.0, 3.0, 4.0), + nullArray), sparseRow(3, Array(0, 2), Array(2.0, 12.0))) } test("vector affine transform with a null scale and shift") { - val nullVector = Literal(null, vectorSqlType) + val nullArray = Literal(null, doubleArraySqlType) checkEvaluation( - VectorAffineTransform(dense(1.0, 2.0), nullVector, nullVector), + VectorAffineTransform(dense(1.0, 2.0), nullArray, nullArray), denseRow(1.0, 2.0)) checkEvaluation( VectorAffineTransform( sparse(3, Array(0, 2), Array(1.0, 3.0)), - nullVector, - nullVector), + nullArray, + nullArray), sparseRow(3, Array(0, 2), Array(1.0, 3.0))) } test("vector affine transform with empty vectors") { - val emptyDense = dense() val emptySparse = sparse(0, Array.emptyIntArray, Array.emptyDoubleArray) + val emptyArray = array() checkEvaluation( - VectorAffineTransform(emptyDense, emptyDense, emptyDense), + VectorAffineTransform(dense(), emptyArray, emptyArray), denseRow()) checkEvaluation( - VectorAffineTransform(emptySparse, emptySparse, emptySparse), + VectorAffineTransform(emptySparse, emptyArray, emptyArray), sparseRow(0, Array.emptyIntArray, Array.emptyDoubleArray)) } test("vector affine transform with infinite and NaN values") { Seq(Double.PositiveInfinity, Double.NegativeInfinity, Double.NaN).foreach { value => checkEvaluation( - VectorAffineTransform(dense(value), dense(1.0), dense(0.0)), + VectorAffineTransform(dense(value), array(1.0), array(0.0)), denseRow(value)) checkEvaluation( - VectorAffineTransform(dense(1.0), dense(value), dense(0.0)), + VectorAffineTransform(dense(1.0), array(value), array(0.0)), denseRow(value)) checkEvaluation( - VectorAffineTransform(dense(1.0), dense(1.0), dense(value)), + VectorAffineTransform(dense(1.0), array(1.0), array(value)), denseRow(value)) } } - test("vector affine transform rejects vectors with different sizes") { + test("vector affine transform rejects inputs with different sizes") { checkExceptionInExpression[IllegalArgumentException]( - VectorAffineTransform(dense(1.0), dense(1.0, 2.0), dense(1.0)), - "vectors with non-matching sizes") + VectorAffineTransform(dense(1.0), array(1.0, 2.0), array(1.0)), + "inputs with non-matching sizes") checkExceptionInExpression[IllegalArgumentException]( - VectorAffineTransform(dense(1.0), dense(1.0), dense(1.0, 2.0)), - "vectors with non-matching sizes") + VectorAffineTransform(dense(1.0), array(1.0), array(1.0, 2.0)), + "inputs with non-matching sizes") + } + + test("vector affine transform requires arrays without null elements") { + val nullableArray = Literal( + UnsafeArrayData.fromPrimitiveArray(Array(1.0)), + ArrayType(DoubleType, containsNull = true)) + + assert(VectorAffineTransform(dense(1.0), nullableArray, array(0.0)) + .checkInputDataTypes().isFailure) + assert(VectorAffineTransform(dense(1.0), array(1.0), nullableArray) + .checkInputDataTypes().isFailure) } } From a8009e2d50f4db54fcb83acdfdc6ad543968f62c Mon Sep 17 00:00:00 2001 From: Ruifeng Zheng Date: Thu, 10 Sep 2026 13:10:42 +0000 Subject: [PATCH 06/12] [SPARK-59398][ML][SQL] Define affine transform sparsity behavior --- .../org/apache/spark/ml/FunctionsSuite.scala | 2 ++ .../expressions/ml/VectorAffineTransform.scala | 11 +---------- .../ml/VectorAffineTransformSuite.scala | 16 +++------------- 3 files changed, 6 insertions(+), 23 deletions(-) diff --git a/mllib/src/test/scala/org/apache/spark/ml/FunctionsSuite.scala b/mllib/src/test/scala/org/apache/spark/ml/FunctionsSuite.scala index 905405ac0c703..e81a8dc34ad4c 100644 --- a/mllib/src/test/scala/org/apache/spark/ml/FunctionsSuite.scala +++ b/mllib/src/test/scala/org/apache/spark/ml/FunctionsSuite.scala @@ -261,6 +261,7 @@ class FunctionsSuite extends MLTest { (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") @@ -273,6 +274,7 @@ class FunctionsSuite extends MLTest { 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)) diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorAffineTransform.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorAffineTransform.scala index 45618bced8e4b..cbde2d741d49b 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorAffineTransform.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorAffineTransform.scala @@ -124,15 +124,6 @@ object VectorAffineTransform { } } - private def isZeroArray(values: ArrayData): Boolean = { - var index = 0 - while (index < values.numElements()) { - if (values.getDouble(index) != 0.0) return false - index += 1 - } - true - } - private def sparseResult( vector: InternalRow, size: Int, @@ -220,7 +211,7 @@ object VectorAffineTransform { "VectorAffineTransform was given inputs with non-matching sizes:" + s" vector.size = $size, scale.size = $scaleSize, shift.size = $shiftSize") - if (vectorType == SparseVectorType && (shift == null || isZeroArray(shift))) { + if (vectorType == SparseVectorType && shift == null) { sparseResult(vector, size, vectorValues, scale) } else { denseResult(vector, size, vectorType, vectorValues, scale, shift) diff --git a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorAffineTransformSuite.scala b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorAffineTransformSuite.scala index 03e0943354270..e87fa67ce3ac6 100644 --- a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorAffineTransformSuite.scala +++ b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorAffineTransformSuite.scala @@ -68,26 +68,16 @@ class VectorAffineTransformSuite extends SparkFunSuite with ExpressionEvalHelper denseRow(2.0, 1.0, 12.0)) } - test("vector affine transform preserves sparse vectors for a zero shift") { + test("vector affine transform produces a dense vector for a non-null shift") { val vector = sparse(3, Array(0, 2), Array(1.0, 3.0)) - val expected = sparseRow(3, Array(0, 2), Array(2.0, 12.0)) checkEvaluation( VectorAffineTransform(vector, array(2.0, 3.0, 4.0), array(0.0, 0.0, 0.0)), - expected) + denseRow(2.0, 0.0, 12.0)) checkEvaluation( VectorAffineTransform( vector, array(2.0, 3.0, 4.0), - array(0.0, 0.0, 0.0)), - expected) - } - - test("vector affine transform produces a dense vector for a nonzero shift") { - checkEvaluation( - VectorAffineTransform( - sparse(3, Array(0, 2), Array(1.0, 3.0)), - array(2.0, 3.0, 4.0), array(0.0, 1.0, 0.0)), denseRow(2.0, 1.0, 12.0)) } @@ -150,7 +140,7 @@ class VectorAffineTransformSuite extends SparkFunSuite with ExpressionEvalHelper denseRow()) checkEvaluation( VectorAffineTransform(emptySparse, emptyArray, emptyArray), - sparseRow(0, Array.emptyIntArray, Array.emptyDoubleArray)) + denseRow()) } test("vector affine transform with infinite and NaN values") { From f0422973c463ff290eab9486b96c725d4304327c Mon Sep 17 00:00:00 2001 From: Ruifeng Zheng Date: Thu, 10 Sep 2026 13:34:47 +0000 Subject: [PATCH 07/12] [SPARK-59398][ML][SQL] Hoist affine transform null checks --- .../ml/VectorAffineTransform.scala | 72 ++++++++++--------- 1 file changed, 40 insertions(+), 32 deletions(-) diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorAffineTransform.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorAffineTransform.scala index cbde2d741d49b..12c34ebd64c1b 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorAffineTransform.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorAffineTransform.scala @@ -134,12 +134,8 @@ object VectorAffineTransform { var vectorIndex = 0 while (vectorIndex < resultValues.length) { val featureIndex = vectorIndices.getInt(vectorIndex) - val scaleValue = if (scale == null) { - 1.0 - } else { - scale.getDouble(featureIndex) - } - resultValues(vectorIndex) = vectorValues.getDouble(vectorIndex) * scaleValue + resultValues(vectorIndex) = + vectorValues.getDouble(vectorIndex) * scale.getDouble(featureIndex) vectorIndex += 1 } new GenericInternalRow(Array[Any]( @@ -156,39 +152,51 @@ object VectorAffineTransform { vectorValues: ArrayData, scale: ArrayData, shift: ArrayData): InternalRow = { - val vectorIndices = if (vectorType == SparseVectorType) vector.getArray(2) else null val resultValues = new Array[Double](size) - var vectorIndex = 0 - var featureIndex = 0 - while (featureIndex < size) { - val vectorIsActive = vectorType == DenseVectorType || - (vectorIndex < vectorValues.numElements() && - vectorIndices.getInt(vectorIndex) == featureIndex) - val vectorValue = if (vectorType == DenseVectorType) { - vectorValues.getDouble(featureIndex) - } else if (vectorIsActive) { - val value = vectorValues.getDouble(vectorIndex) - vectorIndex += 1 - value + if (vectorType == DenseVectorType) { + var featureIndex = 0 + if (scale == null) { + while (featureIndex < size) { + resultValues(featureIndex) = + vectorValues.getDouble(featureIndex) + shift.getDouble(featureIndex) + featureIndex += 1 + } + } else if (shift == null) { + while (featureIndex < size) { + resultValues(featureIndex) = + vectorValues.getDouble(featureIndex) * scale.getDouble(featureIndex) + featureIndex += 1 + } } else { - 0.0 + while (featureIndex < size) { + resultValues(featureIndex) = vectorValues.getDouble(featureIndex) * + scale.getDouble(featureIndex) + shift.getDouble(featureIndex) + featureIndex += 1 + } } - - val scaleValue = if (scale == null) { - 1.0 - } else { - scale.getDouble(featureIndex) + } else { + var featureIndex = 0 + while (featureIndex < size) { + resultValues(featureIndex) = 0.0 + shift.getDouble(featureIndex) + featureIndex += 1 } - val shiftValue = if (shift == null) { - 0.0 + val vectorIndices = vector.getArray(2) + var vectorIndex = 0 + if (scale == null) { + while (vectorIndex < vectorValues.numElements()) { + val featureIndex = vectorIndices.getInt(vectorIndex) + resultValues(featureIndex) += vectorValues.getDouble(vectorIndex) + vectorIndex += 1 + } } else { - shift.getDouble(featureIndex) + while (vectorIndex < vectorValues.numElements()) { + val featureIndex = vectorIndices.getInt(vectorIndex) + resultValues(featureIndex) += + vectorValues.getDouble(vectorIndex) * scale.getDouble(featureIndex) + vectorIndex += 1 + } } - - resultValues(featureIndex) = - (if (vectorIsActive) vectorValue * scaleValue else 0.0) + shiftValue - featureIndex += 1 } new GenericInternalRow(Array[Any]( DenseVectorType, From 126df4a3b2a9d79fedb1a3c228a0c96ceb953659 Mon Sep 17 00:00:00 2001 From: Ruifeng Zheng Date: Thu, 10 Sep 2026 13:50:33 +0000 Subject: [PATCH 08/12] [SPARK-59398][ML][SQL] Expand affine transform function tests --- .../org/apache/spark/ml/FunctionsSuite.scala | 29 +++++++++++++++++++ 1 file changed, 29 insertions(+) diff --git a/mllib/src/test/scala/org/apache/spark/ml/FunctionsSuite.scala b/mllib/src/test/scala/org/apache/spark/ml/FunctionsSuite.scala index e81a8dc34ad4c..7e6233440f616 100644 --- a/mllib/src/test/scala/org/apache/spark/ml/FunctionsSuite.scala +++ b/mllib/src/test/scala/org/apache/spark/ml/FunctionsSuite.scala @@ -306,6 +306,35 @@ class FunctionsSuite extends MLTest { .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") { From f5590131a210204adfdbeb030190a21c74720c76 Mon Sep 17 00:00:00 2001 From: Ruifeng Zheng Date: Thu, 10 Sep 2026 13:59:14 +0000 Subject: [PATCH 09/12] [SPARK-59398][ML][SQL] Generate vector affine transform evaluation --- .../ml/VectorAffineTransform.scala | 102 +++++++++++++++++- 1 file changed, 100 insertions(+), 2 deletions(-) diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorAffineTransform.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorAffineTransform.scala index 12c34ebd64c1b..2533f2cf7fdb4 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorAffineTransform.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorAffineTransform.scala @@ -65,7 +65,6 @@ case class VectorAffineTransform( } override protected def doGenCode(ctx: CodegenContext, ev: ExprCode): ExprCode = { - val cls = VectorAffineTransform.getClass.getName val vectorJavaType = CodeGenerator.javaType(dataType) val arrayJavaType = CodeGenerator.javaType(VectorAffineTransform.doubleArraySqlType) val vectorGen = vector.genCode(ctx) @@ -73,6 +72,15 @@ case class VectorAffineTransform( val shiftInput = ctx.freshName("shiftInput") val scaleGen = scale.genCode(ctx) val shiftGen = shift.genCode(ctx) + val vectorType = ctx.freshName("vectorType") + val vectorValues = ctx.freshName("vectorValues") + val size = ctx.freshName("size") + val scaleSize = ctx.freshName("scaleSize") + val shiftSize = ctx.freshName("shiftSize") + val vectorIndices = ctx.freshName("vectorIndices") + val resultValues = ctx.freshName("resultValues") + val featureIndex = ctx.freshName("featureIndex") + val vectorIndex = ctx.freshName("vectorIndex") ev.copy(code = code""" ${vectorGen.code} @@ -83,7 +91,97 @@ case class VectorAffineTransform( ${shiftGen.code} $arrayJavaType $scaleInput = ${scaleGen.isNull} ? null : ${scaleGen.value}; $arrayJavaType $shiftInput = ${shiftGen.isNull} ? null : ${shiftGen.value}; - ${ev.value} = $cls.MODULE$$.transform(${vectorGen.value}, $scaleInput, $shiftInput); + if ($scaleInput == null && $shiftInput == null) { + ${ev.value} = ${vectorGen.value}; + } else { + final byte $vectorType = ${vectorGen.value}.getByte(0); + final ArrayData $vectorValues = ${vectorGen.value}.getArray(3); + int $size = -1; + if ($vectorType == ${VectorAffineTransform.SparseVectorType}) { + $size = ${vectorGen.value}.getInt(1); + } else if ($vectorType == ${VectorAffineTransform.DenseVectorType}) { + $size = $vectorValues.numElements(); + } else { + throw new IllegalArgumentException("Unknown vector type " + $vectorType + "."); + } + + final int $scaleSize = $scaleInput == null ? $size : $scaleInput.numElements(); + final int $shiftSize = $shiftInput == null ? $size : $shiftInput.numElements(); + 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 == ${VectorAffineTransform.SparseVectorType} && + $shiftInput == null) { + final ArrayData $vectorIndices = ${vectorGen.value}.getArray(2); + final double[] $resultValues = new double[$vectorValues.numElements()]; + for (int $vectorIndex = 0; + $vectorIndex < $resultValues.length; + $vectorIndex++) { + final int $featureIndex = $vectorIndices.getInt($vectorIndex); + $resultValues[$vectorIndex] = $vectorValues.getDouble($vectorIndex) * + $scaleInput.getDouble($featureIndex); + } + ${ev.value} = new GenericInternalRow(new Object[] { + (byte) ${VectorAffineTransform.SparseVectorType}, + $size, + $vectorIndices, + UnsafeArrayData.fromPrimitiveArray($resultValues) + }); + } else { + final double[] $resultValues = new double[$size]; + if ($vectorType == ${VectorAffineTransform.DenseVectorType}) { + if ($scaleInput == null) { + for (int $featureIndex = 0; $featureIndex < $size; $featureIndex++) { + $resultValues[$featureIndex] = $vectorValues.getDouble($featureIndex) + + $shiftInput.getDouble($featureIndex); + } + } else if ($shiftInput == null) { + for (int $featureIndex = 0; $featureIndex < $size; $featureIndex++) { + $resultValues[$featureIndex] = $vectorValues.getDouble($featureIndex) * + $scaleInput.getDouble($featureIndex); + } + } else { + for (int $featureIndex = 0; $featureIndex < $size; $featureIndex++) { + $resultValues[$featureIndex] = $vectorValues.getDouble($featureIndex) * + $scaleInput.getDouble($featureIndex) + + $shiftInput.getDouble($featureIndex); + } + } + } else { + for (int $featureIndex = 0; $featureIndex < $size; $featureIndex++) { + $resultValues[$featureIndex] = 0.0D + $shiftInput.getDouble($featureIndex); + } + + final ArrayData $vectorIndices = ${vectorGen.value}.getArray(2); + if ($scaleInput == null) { + for (int $vectorIndex = 0; + $vectorIndex < $vectorValues.numElements(); + $vectorIndex++) { + final int $featureIndex = $vectorIndices.getInt($vectorIndex); + $resultValues[$featureIndex] += $vectorValues.getDouble($vectorIndex); + } + } else { + for (int $vectorIndex = 0; + $vectorIndex < $vectorValues.numElements(); + $vectorIndex++) { + final int $featureIndex = $vectorIndices.getInt($vectorIndex); + $resultValues[$featureIndex] += $vectorValues.getDouble($vectorIndex) * + $scaleInput.getDouble($featureIndex); + } + } + } + ${ev.value} = new GenericInternalRow(new Object[] { + (byte) ${VectorAffineTransform.DenseVectorType}, + null, + null, + UnsafeArrayData.fromPrimitiveArray($resultValues) + }); + } + } } """) } From d5f8e31897ff4fd28394facb8c53ceb54f3ba834 Mon Sep 17 00:00:00 2001 From: Ruifeng Zheng Date: Thu, 10 Sep 2026 15:12:10 +0000 Subject: [PATCH 10/12] [SPARK-59398][ML][SQL] Share vector affine transform evaluation --- .../org/apache/spark/ml/FunctionsSuite.scala | 32 ++- .../expressions/ml/MLExpressionUtils.java | 181 +++++++++++++ .../ml/VectorAffineTransform.scala | 250 ++---------------- 3 files changed, 241 insertions(+), 222 deletions(-) create mode 100644 sql/catalyst/src/main/java/org/apache/spark/sql/catalyst/expressions/ml/MLExpressionUtils.java diff --git a/mllib/src/test/scala/org/apache/spark/ml/FunctionsSuite.scala b/mllib/src/test/scala/org/apache/spark/ml/FunctionsSuite.scala index 7e6233440f616..04601f88e2c8f 100644 --- a/mllib/src/test/scala/org/apache/spark/ml/FunctionsSuite.scala +++ b/mllib/src/test/scala/org/apache/spark/ml/FunctionsSuite.scala @@ -25,7 +25,7 @@ import org.apache.spark.mllib.linalg.{Matrices => OldMatrices, MatrixUDT => OldM Vector => OldVector, Vectors => OldVectors, VectorUDT => OldVectorUDT} import org.apache.spark.sql.{AnalysisException, DataFrame, Row} import org.apache.spark.sql.catalyst.expressions.ml.{VectorAffineTransform, VectorPosExplode} -import org.apache.spark.sql.functions.{col, unwrap_udt, wrap_udt} +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 { @@ -291,6 +291,36 @@ class FunctionsSuite extends MLTest { .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", diff --git a/sql/catalyst/src/main/java/org/apache/spark/sql/catalyst/expressions/ml/MLExpressionUtils.java b/sql/catalyst/src/main/java/org/apache/spark/sql/catalyst/expressions/ml/MLExpressionUtils.java new file mode 100644 index 0000000000000..3874ecb0b5a75 --- /dev/null +++ b/sql/catalyst/src/main/java/org/apache/spark/sql/catalyst/expressions/ml/MLExpressionUtils.java @@ -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) + }); + } +} diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorAffineTransform.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorAffineTransform.scala index 2533f2cf7fdb4..ce646cdd0bed7 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorAffineTransform.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorAffineTransform.scala @@ -18,7 +18,7 @@ package org.apache.spark.sql.catalyst.expressions.ml import org.apache.spark.sql.catalyst.InternalRow -import org.apache.spark.sql.catalyst.expressions.{ExpectsInputTypes, Expression, GenericInternalRow, TernaryExpression, UnsafeArrayData} +import org.apache.spark.sql.catalyst.expressions.{ExpectsInputTypes, Expression, Literal, TernaryExpression} import org.apache.spark.sql.catalyst.expressions.codegen.{CodegenContext, CodeGenerator, ExprCode} import org.apache.spark.sql.catalyst.expressions.codegen.Block._ import org.apache.spark.sql.catalyst.util.ArrayData @@ -57,7 +57,7 @@ case class VectorAffineTransform( if (vectorInput == null) { null } else { - VectorAffineTransform.transform( + MLExpressionUtils.affineTransform( vectorInput.asInstanceOf[InternalRow], scale.eval(input).asInstanceOf[ArrayData], shift.eval(input).asInstanceOf[ArrayData]) @@ -65,123 +65,44 @@ case class VectorAffineTransform( } override protected def doGenCode(ctx: CodegenContext, ev: ExprCode): ExprCode = { + val utils = classOf[MLExpressionUtils].getName val vectorJavaType = CodeGenerator.javaType(dataType) val arrayJavaType = CodeGenerator.javaType(VectorAffineTransform.doubleArraySqlType) val vectorGen = vector.genCode(ctx) val scaleInput = ctx.freshName("scaleInput") val shiftInput = ctx.freshName("shiftInput") - val scaleGen = scale.genCode(ctx) - val shiftGen = shift.genCode(ctx) - val vectorType = ctx.freshName("vectorType") - val vectorValues = ctx.freshName("vectorValues") - val size = ctx.freshName("size") - val scaleSize = ctx.freshName("scaleSize") - val shiftSize = ctx.freshName("shiftSize") - val vectorIndices = ctx.freshName("vectorIndices") - val resultValues = ctx.freshName("resultValues") - val featureIndex = ctx.freshName("featureIndex") - val vectorIndex = ctx.freshName("vectorIndex") + val (scaleCode, cachedScale) = scale match { + case Literal(value: ArrayData, _) => + (code"$arrayJavaType $scaleInput = null;", + ctx.addReferenceObj("cachedScale", value.toDoubleArray(), "double[]")) + case _ => + val scaleGen = scale.genCode(ctx) + (code""" + ${scaleGen.code} + $arrayJavaType $scaleInput = ${scaleGen.isNull} ? null : ${scaleGen.value}; + """, "null") + } + val (shiftCode, cachedShift) = shift match { + case Literal(value: ArrayData, _) => + (code"$arrayJavaType $shiftInput = null;", + ctx.addReferenceObj("cachedShift", value.toDoubleArray(), "double[]")) + case _ => + val shiftGen = shift.genCode(ctx) + (code""" + ${shiftGen.code} + $arrayJavaType $shiftInput = ${shiftGen.isNull} ? null : ${shiftGen.value}; + """, "null") + } ev.copy(code = code""" ${vectorGen.code} boolean ${ev.isNull} = ${vectorGen.isNull}; $vectorJavaType ${ev.value} = null; if (!${ev.isNull}) { - ${scaleGen.code} - ${shiftGen.code} - $arrayJavaType $scaleInput = ${scaleGen.isNull} ? null : ${scaleGen.value}; - $arrayJavaType $shiftInput = ${shiftGen.isNull} ? null : ${shiftGen.value}; - if ($scaleInput == null && $shiftInput == null) { - ${ev.value} = ${vectorGen.value}; - } else { - final byte $vectorType = ${vectorGen.value}.getByte(0); - final ArrayData $vectorValues = ${vectorGen.value}.getArray(3); - int $size = -1; - if ($vectorType == ${VectorAffineTransform.SparseVectorType}) { - $size = ${vectorGen.value}.getInt(1); - } else if ($vectorType == ${VectorAffineTransform.DenseVectorType}) { - $size = $vectorValues.numElements(); - } else { - throw new IllegalArgumentException("Unknown vector type " + $vectorType + "."); - } - - final int $scaleSize = $scaleInput == null ? $size : $scaleInput.numElements(); - final int $shiftSize = $shiftInput == null ? $size : $shiftInput.numElements(); - 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 == ${VectorAffineTransform.SparseVectorType} && - $shiftInput == null) { - final ArrayData $vectorIndices = ${vectorGen.value}.getArray(2); - final double[] $resultValues = new double[$vectorValues.numElements()]; - for (int $vectorIndex = 0; - $vectorIndex < $resultValues.length; - $vectorIndex++) { - final int $featureIndex = $vectorIndices.getInt($vectorIndex); - $resultValues[$vectorIndex] = $vectorValues.getDouble($vectorIndex) * - $scaleInput.getDouble($featureIndex); - } - ${ev.value} = new GenericInternalRow(new Object[] { - (byte) ${VectorAffineTransform.SparseVectorType}, - $size, - $vectorIndices, - UnsafeArrayData.fromPrimitiveArray($resultValues) - }); - } else { - final double[] $resultValues = new double[$size]; - if ($vectorType == ${VectorAffineTransform.DenseVectorType}) { - if ($scaleInput == null) { - for (int $featureIndex = 0; $featureIndex < $size; $featureIndex++) { - $resultValues[$featureIndex] = $vectorValues.getDouble($featureIndex) + - $shiftInput.getDouble($featureIndex); - } - } else if ($shiftInput == null) { - for (int $featureIndex = 0; $featureIndex < $size; $featureIndex++) { - $resultValues[$featureIndex] = $vectorValues.getDouble($featureIndex) * - $scaleInput.getDouble($featureIndex); - } - } else { - for (int $featureIndex = 0; $featureIndex < $size; $featureIndex++) { - $resultValues[$featureIndex] = $vectorValues.getDouble($featureIndex) * - $scaleInput.getDouble($featureIndex) + - $shiftInput.getDouble($featureIndex); - } - } - } else { - for (int $featureIndex = 0; $featureIndex < $size; $featureIndex++) { - $resultValues[$featureIndex] = 0.0D + $shiftInput.getDouble($featureIndex); - } - - final ArrayData $vectorIndices = ${vectorGen.value}.getArray(2); - if ($scaleInput == null) { - for (int $vectorIndex = 0; - $vectorIndex < $vectorValues.numElements(); - $vectorIndex++) { - final int $featureIndex = $vectorIndices.getInt($vectorIndex); - $resultValues[$featureIndex] += $vectorValues.getDouble($vectorIndex); - } - } else { - for (int $vectorIndex = 0; - $vectorIndex < $vectorValues.numElements(); - $vectorIndex++) { - final int $featureIndex = $vectorIndices.getInt($vectorIndex); - $resultValues[$featureIndex] += $vectorValues.getDouble($vectorIndex) * - $scaleInput.getDouble($featureIndex); - } - } - } - ${ev.value} = new GenericInternalRow(new Object[] { - (byte) ${VectorAffineTransform.DenseVectorType}, - null, - null, - UnsafeArrayData.fromPrimitiveArray($resultValues) - }); - } - } + $scaleCode + $shiftCode + ${ev.value} = $utils.affineTransform( + ${vectorGen.value}, $scaleInput, $shiftInput, $cachedScale, $cachedShift); } """) } @@ -195,9 +116,6 @@ case class VectorAffineTransform( } object VectorAffineTransform { - private val SparseVectorType: Byte = 0 - private val DenseVectorType: Byte = 1 - private[ml] val vectorSqlType = StructType(Array( StructField("type", ByteType, nullable = false), StructField("size", IntegerType, nullable = true), @@ -213,114 +131,4 @@ object VectorAffineTransform { override private[spark] def simpleString: String = doubleArraySqlType.simpleString } - - private def vectorSize(vector: InternalRow, vectorType: Byte, values: ArrayData): Int = { - vectorType match { - case SparseVectorType => vector.getInt(1) - case DenseVectorType => values.numElements() - case _ => throw new IllegalArgumentException(s"Unknown vector type $vectorType.") - } - } - - private def sparseResult( - vector: InternalRow, - size: Int, - vectorValues: ArrayData, - scale: ArrayData): InternalRow = { - val vectorIndices = vector.getArray(2) - val resultValues = new Array[Double](vectorValues.numElements()) - var vectorIndex = 0 - while (vectorIndex < resultValues.length) { - val featureIndex = vectorIndices.getInt(vectorIndex) - resultValues(vectorIndex) = - vectorValues.getDouble(vectorIndex) * scale.getDouble(featureIndex) - vectorIndex += 1 - } - new GenericInternalRow(Array[Any]( - SparseVectorType, - size, - vectorIndices, - UnsafeArrayData.fromPrimitiveArray(resultValues))) - } - - private def denseResult( - vector: InternalRow, - size: Int, - vectorType: Byte, - vectorValues: ArrayData, - scale: ArrayData, - shift: ArrayData): InternalRow = { - val resultValues = new Array[Double](size) - if (vectorType == DenseVectorType) { - var featureIndex = 0 - if (scale == null) { - while (featureIndex < size) { - resultValues(featureIndex) = - vectorValues.getDouble(featureIndex) + shift.getDouble(featureIndex) - featureIndex += 1 - } - } else if (shift == null) { - while (featureIndex < size) { - resultValues(featureIndex) = - vectorValues.getDouble(featureIndex) * scale.getDouble(featureIndex) - featureIndex += 1 - } - } else { - while (featureIndex < size) { - resultValues(featureIndex) = vectorValues.getDouble(featureIndex) * - scale.getDouble(featureIndex) + shift.getDouble(featureIndex) - featureIndex += 1 - } - } - } else { - var featureIndex = 0 - while (featureIndex < size) { - resultValues(featureIndex) = 0.0 + shift.getDouble(featureIndex) - featureIndex += 1 - } - - val vectorIndices = vector.getArray(2) - var vectorIndex = 0 - if (scale == null) { - while (vectorIndex < vectorValues.numElements()) { - val featureIndex = vectorIndices.getInt(vectorIndex) - resultValues(featureIndex) += vectorValues.getDouble(vectorIndex) - vectorIndex += 1 - } - } else { - while (vectorIndex < vectorValues.numElements()) { - val featureIndex = vectorIndices.getInt(vectorIndex) - resultValues(featureIndex) += - vectorValues.getDouble(vectorIndex) * scale.getDouble(featureIndex) - vectorIndex += 1 - } - } - } - new GenericInternalRow(Array[Any]( - DenseVectorType, - null, - null, - UnsafeArrayData.fromPrimitiveArray(resultValues))) - } - - private[ml] def transform( - vector: InternalRow, - scale: ArrayData, - shift: ArrayData): InternalRow = { - if (scale == null && shift == null) return vector - val vectorType = vector.getByte(0) - val vectorValues = vector.getArray(3) - val size = vectorSize(vector, vectorType, vectorValues) - val scaleSize = if (scale == null) size else scale.numElements() - val shiftSize = if (shift == null) size else shift.numElements() - require(size == scaleSize && size == shiftSize, - "VectorAffineTransform was given inputs with non-matching sizes:" + - s" vector.size = $size, scale.size = $scaleSize, shift.size = $shiftSize") - - if (vectorType == SparseVectorType && shift == null) { - sparseResult(vector, size, vectorValues, scale) - } else { - denseResult(vector, size, vectorType, vectorValues, scale, shift) - } - } } From 552f8c3fb013250c8ae57f600580b8231caccd78 Mon Sep 17 00:00:00 2001 From: Ruifeng Zheng Date: Fri, 11 Sep 2026 03:33:01 +0000 Subject: [PATCH 11/12] [SPARK-59398][ML][SQL] Cover cached coefficient null combinations --- .../scala/org/apache/spark/ml/FunctionsSuite.scala | 14 ++++++++++++++ 1 file changed, 14 insertions(+) diff --git a/mllib/src/test/scala/org/apache/spark/ml/FunctionsSuite.scala b/mllib/src/test/scala/org/apache/spark/ml/FunctionsSuite.scala index 04601f88e2c8f..b2c4acb0ab35d 100644 --- a/mllib/src/test/scala/org/apache/spark/ml/FunctionsSuite.scala +++ b/mllib/src/test/scala/org/apache/spark/ml/FunctionsSuite.scala @@ -297,12 +297,26 @@ class FunctionsSuite extends MLTest { .getAs[Vector](0) assert(cachedScaleResult === Vectors.dense(6.0, 11.0)) + val cachedScaleWithNullShiftResult = df + .where($"scale".isNotNull && $"shift".isNull) + .select(vector_affine_transform($"vector", typedLit(Array(2.0, 3.0)), $"shift")) + .first() + .getAs[Vector](0) + assert(cachedScaleWithNullShiftResult === Vectors.sparse(2, Seq((0, 2.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 nullScaleWithCachedShiftResult = df + .where($"scale".isNull && $"shift".isNotNull) + .select(vector_affine_transform($"vector", $"scale", typedLit(Array(4.0, 5.0)))) + .first() + .getAs[Vector](0) + assert(nullScaleWithCachedShiftResult === Vectors.dense(5.0, 7.0)) + val scaleOnlyConstantResult = df.limit(1) .select(vector_affine_transform( $"vector", From 2a6fa230ce4c2493d8ede16070c6661980bfd9b9 Mon Sep 17 00:00:00 2001 From: Ruifeng Zheng Date: Fri, 11 Sep 2026 03:53:21 +0000 Subject: [PATCH 12/12] [SPARK-59398][ML][SQL] Rename vector affine transform to scale shift --- .../scala/org/apache/spark/ml/functions.scala | 8 +-- .../org/apache/spark/ml/FunctionsSuite.scala | 32 ++++----- .../expressions/ml/MLExpressionUtils.java | 8 +-- .../catalyst/analysis/FunctionRegistry.scala | 2 +- ...Transform.scala => VectorScaleShift.scala} | 24 +++---- ...uite.scala => VectorScaleShiftSuite.scala} | 68 +++++++++---------- 6 files changed, 71 insertions(+), 71 deletions(-) rename sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/ml/{VectorAffineTransform.scala => VectorScaleShift.scala} (87%) rename sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/ml/{VectorAffineTransformSuite.scala => VectorScaleShiftSuite.scala} (66%) diff --git a/mllib/src/main/scala/org/apache/spark/ml/functions.scala b/mllib/src/main/scala/org/apache/spark/ml/functions.scala index 869ba15e6ab2b..12c073d8ae6ef 100644 --- a/mllib/src/main/scala/org/apache/spark/ml/functions.scala +++ b/mllib/src/main/scala/org/apache/spark/ml/functions.scala @@ -71,24 +71,24 @@ object functions { vectorToStruct(right)) } - private[ml] def vector_affine_transform( + private[ml] def vector_scale_shift( vector: Column, scale: Column, shift: Column): Column = { val transformed = Column.internalFn( - "ml_vector_affine_transform", + "ml_vector_scale_shift", sf.unwrap_udt(vector), scale, shift) sf.wrap_udt(transformed, new VectorUDT) } - private[ml] def vector_affine_transform( + private[ml] def vector_scale_shift( vector: Column, scale: Array[Double], shift: Array[Double]): Column = { val transformed = Column.internalFn( - "ml_vector_affine_transform", + "ml_vector_scale_shift", sf.unwrap_udt(vector), doubleArrayLiteral(scale), doubleArrayLiteral(shift)) diff --git a/mllib/src/test/scala/org/apache/spark/ml/FunctionsSuite.scala b/mllib/src/test/scala/org/apache/spark/ml/FunctionsSuite.scala index b2c4acb0ab35d..3ca352df81710 100644 --- a/mllib/src/test/scala/org/apache/spark/ml/FunctionsSuite.scala +++ b/mllib/src/test/scala/org/apache/spark/ml/FunctionsSuite.scala @@ -24,7 +24,7 @@ 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.{VectorAffineTransform, VectorPosExplode} +import org.apache.spark.sql.catalyst.expressions.ml.{VectorPosExplode, VectorScaleShift} import org.apache.spark.sql.functions.{col, typedLit, unwrap_udt, wrap_udt} import org.apache.spark.sql.types.{ArrayType, DoubleType, StructField, StructType, UserDefinedType} @@ -256,7 +256,7 @@ class FunctionsSuite extends MLTest { assert(error.getMessage.contains("vectors with non-matching sizes")) } - test("test vector_affine_transform") { + test("test vector_scale_shift") { 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), @@ -268,7 +268,7 @@ class FunctionsSuite extends MLTest { 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")) + val transformed = df.select(vector_scale_shift($"vector", $"scale", $"shift")) assert(transformed.schema.head.dataType === new VectorUDT) assert(transformed.collect().map(_.get(0)).toSeq === Seq( Vectors.dense(6.0, 11.0), @@ -279,11 +279,11 @@ class FunctionsSuite extends MLTest { null)) val expressions = transformed.queryExecution.analyzed - .flatMap(_.expressions.flatMap(_.collect { case v: VectorAffineTransform => v })) - assert(expressions.map(_.prettyName).distinct === Seq("ml_vector_affine_transform")) + .flatMap(_.expressions.flatMap(_.collect { case v: VectorScaleShift => v })) + assert(expressions.map(_.prettyName).distinct === Seq("ml_vector_scale_shift")) val constantResult = df.limit(1) - .select(vector_affine_transform( + .select(vector_scale_shift( $"vector", Array(2.0, 3.0), Array(4.0, 5.0))) @@ -292,33 +292,33 @@ class FunctionsSuite extends MLTest { 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")) + .select(vector_scale_shift($"vector", typedLit(Array(2.0, 3.0)), $"shift")) .first() .getAs[Vector](0) assert(cachedScaleResult === Vectors.dense(6.0, 11.0)) val cachedScaleWithNullShiftResult = df .where($"scale".isNotNull && $"shift".isNull) - .select(vector_affine_transform($"vector", typedLit(Array(2.0, 3.0)), $"shift")) + .select(vector_scale_shift($"vector", typedLit(Array(2.0, 3.0)), $"shift")) .first() .getAs[Vector](0) assert(cachedScaleWithNullShiftResult === Vectors.sparse(2, Seq((0, 2.0)))) val cachedShiftResult = df.limit(1) - .select(vector_affine_transform($"vector", $"scale", typedLit(Array(4.0, 5.0)))) + .select(vector_scale_shift($"vector", $"scale", typedLit(Array(4.0, 5.0)))) .first() .getAs[Vector](0) assert(cachedShiftResult === Vectors.dense(6.0, 11.0)) val nullScaleWithCachedShiftResult = df .where($"scale".isNull && $"shift".isNotNull) - .select(vector_affine_transform($"vector", $"scale", typedLit(Array(4.0, 5.0)))) + .select(vector_scale_shift($"vector", $"scale", typedLit(Array(4.0, 5.0)))) .first() .getAs[Vector](0) assert(nullScaleWithCachedShiftResult === Vectors.dense(5.0, 7.0)) val scaleOnlyConstantResult = df.limit(1) - .select(vector_affine_transform( + .select(vector_scale_shift( $"vector", Array(2.0, 3.0), null.asInstanceOf[Array[Double]])) @@ -327,7 +327,7 @@ class FunctionsSuite extends MLTest { assert(scaleOnlyConstantResult === Vectors.dense(2.0, 6.0)) val shiftOnlyConstantResult = df.limit(1) - .select(vector_affine_transform( + .select(vector_scale_shift( $"vector", null.asInstanceOf[Array[Double]], Array(4.0, 5.0))) @@ -336,7 +336,7 @@ class FunctionsSuite extends MLTest { assert(shiftOnlyConstantResult === Vectors.dense(5.0, 7.0)) val nullConstantsResult = df.limit(1) - .select(vector_affine_transform( + .select(vector_scale_shift( $"vector", null.asInstanceOf[Array[Double]], null.asInstanceOf[Array[Double]])) @@ -346,7 +346,7 @@ class FunctionsSuite extends MLTest { val emptyConstantsResult = Seq(Tuple1(Vectors.dense(Array.emptyDoubleArray))) .toDF("vector") - .select(vector_affine_transform($"vector", Array.emptyDoubleArray, Array.emptyDoubleArray)) + .select(vector_scale_shift($"vector", Array.emptyDoubleArray, Array.emptyDoubleArray)) .first() .getAs[Vector](0) assert(emptyConstantsResult === Vectors.dense(Array.emptyDoubleArray)) @@ -360,7 +360,7 @@ class FunctionsSuite extends MLTest { } val specialValueResults = specialValueRows .toDF("vector", "scale", "shift") - .select(vector_affine_transform($"vector", $"scale", $"shift")) + .select(vector_scale_shift($"vector", $"scale", $"shift")) .collect() .map(_.getAs[Vector](0)(0)) specialValueResults.zip(specialValues.flatMap(value => Seq.fill(3)(value))) @@ -374,7 +374,7 @@ class FunctionsSuite extends MLTest { val error = intercept[IllegalArgumentException] { Seq((Vectors.dense(1.0), scale, shift)) .toDF("vector", "scale", "shift") - .select(vector_affine_transform($"vector", $"scale", $"shift")) + .select(vector_scale_shift($"vector", $"scale", $"shift")) .collect() } assert(error.getMessage.contains("inputs with non-matching sizes")) diff --git a/sql/catalyst/src/main/java/org/apache/spark/sql/catalyst/expressions/ml/MLExpressionUtils.java b/sql/catalyst/src/main/java/org/apache/spark/sql/catalyst/expressions/ml/MLExpressionUtils.java index 3874ecb0b5a75..1427775d28556 100644 --- a/sql/catalyst/src/main/java/org/apache/spark/sql/catalyst/expressions/ml/MLExpressionUtils.java +++ b/sql/catalyst/src/main/java/org/apache/spark/sql/catalyst/expressions/ml/MLExpressionUtils.java @@ -28,14 +28,14 @@ public class MLExpressionUtils { private MLExpressionUtils() {} - public static InternalRow affineTransform( + public static InternalRow scaleShift( InternalRow vector, ArrayData scale, ArrayData shift) { - return affineTransform(vector, scale, shift, null, null); + return scaleShift(vector, scale, shift, null, null); } - public static InternalRow affineTransform( + public static InternalRow scaleShift( InternalRow vector, ArrayData scale, ArrayData shift, @@ -64,7 +64,7 @@ public static InternalRow affineTransform( (cachedShift == null ? shift.numElements() : cachedShift.length); if (size != scaleSize || size != shiftSize) { throw new IllegalArgumentException( - "requirement failed: VectorAffineTransform was given inputs with non-matching sizes: " + + "requirement failed: VectorScaleShift was given inputs with non-matching sizes: " + "vector.size = " + size + ", scale.size = " + scaleSize + ", shift.size = " + shiftSize); } diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/FunctionRegistry.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/FunctionRegistry.scala index 8c67d994c7c87..26287782cb29a 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/FunctionRegistry.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/FunctionRegistry.scala @@ -1222,7 +1222,7 @@ object FunctionRegistry { registerInternalExpression[NullIndex]("null_index") registerInternalExpression[CastTimestampNTZToLong]("timestamp_ntz_to_long") registerInternalExpression[ArrayBinarySearch]("array_binary_search") - registerInternalExpression[VectorAffineTransform]("ml_vector_affine_transform") + registerInternalExpression[VectorScaleShift]("ml_vector_scale_shift") registerInternalExpression[VectorDotProduct]("ml_vector_dot_product") registerInternalExpression[VectorPosExplode]("ml_vector_posexplode") } diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorAffineTransform.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorScaleShift.scala similarity index 87% rename from sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorAffineTransform.scala rename to sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorScaleShift.scala index ce646cdd0bed7..8d9f758d898da 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorAffineTransform.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorScaleShift.scala @@ -25,13 +25,13 @@ import org.apache.spark.sql.catalyst.util.ArrayData import org.apache.spark.sql.types._ /** - * Applies an element-wise affine transformation to SQL struct representations of MLlib vectors: + * Applies element-wise scaling and shifting to SQL struct representations of MLlib vectors: * `vector(i) * scale(i) + shift(i)`. This expression is dedicated only for Spark ML and should be * used together with `unwrap_udt` and `wrap_udt`. A null scale is treated as an identity scale, and * a null shift is treated as a zero shift. If both are null, the input vector is returned * unchanged. */ -case class VectorAffineTransform( +case class VectorScaleShift( vector: Expression, scale: Expression, shift: Expression) @@ -41,14 +41,14 @@ case class VectorAffineTransform( override def second: Expression = scale override def third: Expression = shift - override def prettyName: String = "ml_vector_affine_transform" + override def prettyName: String = "ml_vector_scale_shift" override def inputTypes: Seq[AbstractDataType] = Seq( - VectorAffineTransform.vectorSqlType, - VectorAffineTransform.NonNullableDoubleArrayType, - VectorAffineTransform.NonNullableDoubleArrayType) + VectorScaleShift.vectorSqlType, + VectorScaleShift.NonNullableDoubleArrayType, + VectorScaleShift.NonNullableDoubleArrayType) - override def dataType: DataType = VectorAffineTransform.vectorSqlType + override def dataType: DataType = VectorScaleShift.vectorSqlType override def nullable: Boolean = vector.nullable @@ -57,7 +57,7 @@ case class VectorAffineTransform( if (vectorInput == null) { null } else { - MLExpressionUtils.affineTransform( + MLExpressionUtils.scaleShift( vectorInput.asInstanceOf[InternalRow], scale.eval(input).asInstanceOf[ArrayData], shift.eval(input).asInstanceOf[ArrayData]) @@ -67,7 +67,7 @@ case class VectorAffineTransform( override protected def doGenCode(ctx: CodegenContext, ev: ExprCode): ExprCode = { val utils = classOf[MLExpressionUtils].getName val vectorJavaType = CodeGenerator.javaType(dataType) - val arrayJavaType = CodeGenerator.javaType(VectorAffineTransform.doubleArraySqlType) + val arrayJavaType = CodeGenerator.javaType(VectorScaleShift.doubleArraySqlType) val vectorGen = vector.genCode(ctx) val scaleInput = ctx.freshName("scaleInput") val shiftInput = ctx.freshName("shiftInput") @@ -101,7 +101,7 @@ case class VectorAffineTransform( if (!${ev.isNull}) { $scaleCode $shiftCode - ${ev.value} = $utils.affineTransform( + ${ev.value} = $utils.scaleShift( ${vectorGen.value}, $scaleInput, $shiftInput, $cachedScale, $cachedShift); } """) @@ -110,12 +110,12 @@ case class VectorAffineTransform( override protected def withNewChildrenInternal( newVector: Expression, newScale: Expression, - newShift: Expression): VectorAffineTransform = { + newShift: Expression): VectorScaleShift = { copy(vector = newVector, scale = newScale, shift = newShift) } } -object VectorAffineTransform { +object VectorScaleShift { private[ml] val vectorSqlType = StructType(Array( StructField("type", ByteType, nullable = false), StructField("size", IntegerType, nullable = true), diff --git a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorAffineTransformSuite.scala b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorScaleShiftSuite.scala similarity index 66% rename from sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorAffineTransformSuite.scala rename to sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorScaleShiftSuite.scala index e87fa67ce3ac6..3eedefa74e73c 100644 --- a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorAffineTransformSuite.scala +++ b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/ml/VectorScaleShiftSuite.scala @@ -22,9 +22,9 @@ import org.apache.spark.sql.catalyst.InternalRow import org.apache.spark.sql.catalyst.expressions.{ExpressionEvalHelper, GenericInternalRow, Literal, UnsafeArrayData} import org.apache.spark.sql.types.{ArrayType, DoubleType} -class VectorAffineTransformSuite extends SparkFunSuite with ExpressionEvalHelper { - private val vectorSqlType = VectorAffineTransform.vectorSqlType - private val doubleArraySqlType = VectorAffineTransform.doubleArraySqlType +class VectorScaleShiftSuite extends SparkFunSuite with ExpressionEvalHelper { + private val vectorSqlType = VectorScaleShift.vectorSqlType + private val doubleArraySqlType = VectorScaleShift.doubleArraySqlType private def denseRow(values: Double*): InternalRow = { new GenericInternalRow(Array[Any]( @@ -52,128 +52,128 @@ class VectorAffineTransformSuite extends SparkFunSuite with ExpressionEvalHelper Literal(UnsafeArrayData.fromPrimitiveArray(values.toArray), doubleArraySqlType) } - test("vector affine transform interpreted and code-generated evaluation") { - val expression = VectorAffineTransform( + test("vector scale shift interpreted and code-generated evaluation") { + val expression = VectorScaleShift( dense(1.0, 2.0, 3.0), array(2.0, 3.0, 4.0), array(5.0, 6.0, 7.0)) - assert(expression.prettyName === "ml_vector_affine_transform") + assert(expression.prettyName === "ml_vector_scale_shift") checkEvaluation(expression, denseRow(7.0, 12.0, 19.0)) checkEvaluation( - VectorAffineTransform( + VectorScaleShift( dense(1.0, 2.0, 3.0), array(2.0, 0.0, 4.0), array(0.0, 1.0, 0.0)), denseRow(2.0, 1.0, 12.0)) } - test("vector affine transform produces a dense vector for a non-null shift") { + test("vector scale shift produces a dense vector for a non-null shift") { val vector = sparse(3, Array(0, 2), Array(1.0, 3.0)) checkEvaluation( - VectorAffineTransform(vector, array(2.0, 3.0, 4.0), array(0.0, 0.0, 0.0)), + VectorScaleShift(vector, array(2.0, 3.0, 4.0), array(0.0, 0.0, 0.0)), denseRow(2.0, 0.0, 12.0)) checkEvaluation( - VectorAffineTransform( + VectorScaleShift( vector, array(2.0, 3.0, 4.0), array(0.0, 1.0, 0.0)), denseRow(2.0, 1.0, 12.0)) } - test("vector affine transform with a null vector") { + test("vector scale shift with a null vector") { val nullVector = Literal(null, vectorSqlType) val values = array(1.0) - checkEvaluation(VectorAffineTransform(nullVector, values, values), null) + checkEvaluation(VectorScaleShift(nullVector, values, values), null) } - test("vector affine transform with a null scale") { + test("vector scale shift with a null scale") { val nullArray = Literal(null, doubleArraySqlType) checkEvaluation( - VectorAffineTransform(dense(1.0, 2.0), nullArray, array(3.0, 4.0)), + VectorScaleShift(dense(1.0, 2.0), nullArray, array(3.0, 4.0)), denseRow(4.0, 6.0)) checkEvaluation( - VectorAffineTransform( + VectorScaleShift( sparse(3, Array(0, 2), Array(1.0, 3.0)), nullArray, array(0.0, 2.0, 0.0)), denseRow(1.0, 2.0, 3.0)) } - test("vector affine transform with a null shift") { + test("vector scale shift with a null shift") { val nullArray = Literal(null, doubleArraySqlType) checkEvaluation( - VectorAffineTransform(dense(1.0, 2.0), array(3.0, 4.0), nullArray), + VectorScaleShift(dense(1.0, 2.0), array(3.0, 4.0), nullArray), denseRow(3.0, 8.0)) checkEvaluation( - VectorAffineTransform( + VectorScaleShift( sparse(3, Array(0, 2), Array(1.0, 3.0)), array(2.0, 3.0, 4.0), nullArray), sparseRow(3, Array(0, 2), Array(2.0, 12.0))) } - test("vector affine transform with a null scale and shift") { + test("vector scale shift with a null scale and shift") { val nullArray = Literal(null, doubleArraySqlType) checkEvaluation( - VectorAffineTransform(dense(1.0, 2.0), nullArray, nullArray), + VectorScaleShift(dense(1.0, 2.0), nullArray, nullArray), denseRow(1.0, 2.0)) checkEvaluation( - VectorAffineTransform( + VectorScaleShift( sparse(3, Array(0, 2), Array(1.0, 3.0)), nullArray, nullArray), sparseRow(3, Array(0, 2), Array(1.0, 3.0))) } - test("vector affine transform with empty vectors") { + test("vector scale shift with empty vectors") { val emptySparse = sparse(0, Array.emptyIntArray, Array.emptyDoubleArray) val emptyArray = array() checkEvaluation( - VectorAffineTransform(dense(), emptyArray, emptyArray), + VectorScaleShift(dense(), emptyArray, emptyArray), denseRow()) checkEvaluation( - VectorAffineTransform(emptySparse, emptyArray, emptyArray), + VectorScaleShift(emptySparse, emptyArray, emptyArray), denseRow()) } - test("vector affine transform with infinite and NaN values") { + test("vector scale shift with infinite and NaN values") { Seq(Double.PositiveInfinity, Double.NegativeInfinity, Double.NaN).foreach { value => checkEvaluation( - VectorAffineTransform(dense(value), array(1.0), array(0.0)), + VectorScaleShift(dense(value), array(1.0), array(0.0)), denseRow(value)) checkEvaluation( - VectorAffineTransform(dense(1.0), array(value), array(0.0)), + VectorScaleShift(dense(1.0), array(value), array(0.0)), denseRow(value)) checkEvaluation( - VectorAffineTransform(dense(1.0), array(1.0), array(value)), + VectorScaleShift(dense(1.0), array(1.0), array(value)), denseRow(value)) } } - test("vector affine transform rejects inputs with different sizes") { + test("vector scale shift rejects inputs with different sizes") { checkExceptionInExpression[IllegalArgumentException]( - VectorAffineTransform(dense(1.0), array(1.0, 2.0), array(1.0)), + VectorScaleShift(dense(1.0), array(1.0, 2.0), array(1.0)), "inputs with non-matching sizes") checkExceptionInExpression[IllegalArgumentException]( - VectorAffineTransform(dense(1.0), array(1.0), array(1.0, 2.0)), + VectorScaleShift(dense(1.0), array(1.0), array(1.0, 2.0)), "inputs with non-matching sizes") } - test("vector affine transform requires arrays without null elements") { + test("vector scale shift requires arrays without null elements") { val nullableArray = Literal( UnsafeArrayData.fromPrimitiveArray(Array(1.0)), ArrayType(DoubleType, containsNull = true)) - assert(VectorAffineTransform(dense(1.0), nullableArray, array(0.0)) + assert(VectorScaleShift(dense(1.0), nullableArray, array(0.0)) .checkInputDataTypes().isFailure) - assert(VectorAffineTransform(dense(1.0), array(1.0), nullableArray) + assert(VectorScaleShift(dense(1.0), array(1.0), nullableArray) .checkInputDataTypes().isFailure) } }