Skip to content

Commit feed6c1

Browse files
committed
[SPARK-58963][ML][SQL] Add internal ML-specific vector posexplode function
### What changes were proposed in this pull request? This PR adds an internal ML-specific `vector_posexplode` generator expression for MLlib vector columns. The Catalyst expression is registered as an internal expression and operates on the SQL struct representation produced by `unwrap_udt`: ```scala struct<type:tinyint,size:int,indices:array<int>,values:array<double>> ``` The ML-side helper is `private[ml]` and wraps vector UDT columns by calling `unwrap_udt` before invoking the internal expression, so it can be used by Spark ML internals with both `org.apache.spark.ml.linalg.VectorUDT` and `org.apache.spark.mllib.linalg.VectorUDT` columns. The generator emits `(index, value)` rows with two modes: * `sparse` emits nonzero vector entries. * `dense` emits every vector entry, including zeros. For each non-null vector, it first emits a marker row with index `-1 - vector.size` and value `Double.NaN`. Null vectors emit no rows. ### Why are the changes needed? Spark SQL generator functions such as `posexplode` only accept arrays and maps. Spark ML internals sometimes need to explode vector columns into index-value rows while preserving sparse vectors without densifying through `vector_to_array`. This adds a vector-aware internal generator that works from the shared vector SQL representation and keeps the API surface internal to ML. ### Does this PR introduce _any_ user-facing change? No. The Scala helper is `private[ml]`, and the Catalyst expression is registered with `registerInternalExpression`. This is intended only for Spark ML internals. ### How was this patch tested? Added coverage in `FunctionsSuite` for: * `spark.ml` and `spark.mllib` vector UDT columns * dense vectors * sparse vectors with explicit zero values * null vectors * empty sparse and dense vectors * `dense` and `sparse` modes * output schema Also ran: ```bash build/sbt "mllib/Test/compile" ``` ### Was this patch authored or co-authored using generative AI tooling? Generated-by: Codex (GPT-5) Closes #58194 from zhengruifeng/vector-posexplode-dev5. Authored-by: Ruifeng Zheng <ruifengz@apache.org> Signed-off-by: Ruifeng Zheng <ruifengz@foxmail.com>
1 parent 06017e3 commit feed6c1

4 files changed

Lines changed: 387 additions & 1 deletion

File tree

‎mllib/src/main/scala/org/apache/spark/ml/functions.scala‎

Lines changed: 15 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -44,6 +44,21 @@ object functions {
4444
*/
4545
def array_to_vector(v: Column): Column = Column.internalFn("array_to_vector", v)
4646

47+
/**
48+
* Creates a new row for each index-value pair in the given vector column. This expression is
49+
* dedicated only for Spark ML. It always emits a marker row with index `-1 - vector.size` and
50+
* value `Double.NaN` before each non-null vector.
51+
* @param v: the column of MLlib sparse/dense vectors
52+
* @param mode: `dense` emits all elements, and `sparse` emits nonzero elements
53+
* @return the index and value columns of the vector elements
54+
* @since 4.4.0
55+
*/
56+
private[ml] def vector_posexplode(
57+
v: Column,
58+
mode: String = "sparse"): Column = {
59+
Column.internalFn("vector_posexplode", sf.unwrap_udt(v), sf.lit(mode))
60+
}
61+
4762
private[ml] def array_binary_search(a: Column, v: Column): Column =
4863
Column.internalFn("array_binary_search", a, v)
4964

‎mllib/src/test/scala/org/apache/spark/ml/FunctionsSuite.scala‎

Lines changed: 79 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -22,7 +22,7 @@ import org.apache.spark.ml.functions._
2222
import org.apache.spark.ml.linalg.{Matrices, MatrixUDT, Vector, Vectors, VectorUDT}
2323
import org.apache.spark.ml.util.MLTest
2424
import org.apache.spark.mllib.linalg.{Matrices => OldMatrices, MatrixUDT => OldMatrixUDT,
25-
Vectors => OldVectors, VectorUDT => OldVectorUDT}
25+
Vector => OldVector, Vectors => OldVectors, VectorUDT => OldVectorUDT}
2626
import org.apache.spark.sql.{AnalysisException, DataFrame, Row}
2727
import org.apache.spark.sql.functions.{col, unwrap_udt, wrap_udt}
2828
import org.apache.spark.sql.types.{StructField, StructType, UserDefinedType}
@@ -63,6 +63,13 @@ class FunctionsSuite extends MLTest {
6363
assert(converted.collect().map(_.get(0)).toSeq === Seq(expected, null))
6464
}
6565

66+
private def normalizeNaN(rows: Seq[(Int, Int, Double)]): Seq[(Int, Int, String)] = {
67+
rows.map {
68+
case (id, index, value) if value.isNaN => (id, index, "NaN")
69+
case (id, index, value) => (id, index, value.toString)
70+
}
71+
}
72+
6673
test("test vector_to_array") {
6774
val df = Seq(
6875
(Vectors.dense(1.0, 2.0, 3.0), OldVectors.dense(10.0, 20.0, 30.0)),
@@ -137,6 +144,77 @@ class FunctionsSuite extends MLTest {
137144
assert(resultVec3 === Vectors.dense(Array(1.0, 2.0)))
138145
}
139146

147+
test("test vector_posexplode with vector UDT") {
148+
val df = Seq(
149+
(0, Vectors.dense(1.0, 0.0, 3.0), OldVectors.dense(10.0, 0.0, 30.0)),
150+
(1, Vectors.sparse(4, Seq((1, 2.0), (2, 0.0), (3, 4.0))),
151+
OldVectors.sparse(4, Seq((0, 20.0), (1, 0.0), (2, 30.0)))),
152+
(2, null.asInstanceOf[Vector], null.asInstanceOf[OldVector]),
153+
(3, Vectors.sparse(10, Array.emptyIntArray, Array.emptyDoubleArray),
154+
OldVectors.sparse(10, Array.emptyIntArray, Array.emptyDoubleArray)),
155+
(4, Vectors.dense(Array.emptyDoubleArray),
156+
OldVectors.dense(Array.emptyDoubleArray))
157+
).toDF("id", "vec", "oldVec")
158+
159+
val result = df.select($"id", vector_posexplode($"vec"))
160+
.as[(Int, Int, Double)]
161+
.collect()
162+
.toSeq
163+
assert(normalizeNaN(result) === Seq(
164+
(0, -4, "NaN"),
165+
(0, 0, "1.0"),
166+
(0, 2, "3.0"),
167+
(1, -5, "NaN"),
168+
(1, 1, "2.0"),
169+
(1, 3, "4.0"),
170+
(3, -11, "NaN"),
171+
(4, -1, "NaN")))
172+
173+
val oldResult = df.select($"id", vector_posexplode($"oldVec"))
174+
.as[(Int, Int, Double)]
175+
.collect()
176+
.toSeq
177+
assert(normalizeNaN(oldResult) === Seq(
178+
(0, -4, "NaN"),
179+
(0, 0, "10.0"),
180+
(0, 2, "30.0"),
181+
(1, -5, "NaN"),
182+
(1, 0, "20.0"),
183+
(1, 2, "30.0"),
184+
(3, -11, "NaN"),
185+
(4, -1, "NaN")))
186+
187+
val denseResult = df
188+
.where($"id" === 1)
189+
.select($"id", vector_posexplode($"vec", mode = "dense"))
190+
.as[(Int, Int, Double)]
191+
.collect()
192+
.toSeq
193+
assert(normalizeNaN(denseResult) === Seq(
194+
(1, -5, "NaN"),
195+
(1, 0, "0.0"),
196+
(1, 1, "2.0"),
197+
(1, 2, "0.0"),
198+
(1, 3, "4.0")))
199+
200+
val sparseResult = df.select($"id", vector_posexplode($"vec", mode = "sparse"))
201+
.as[(Int, Int, Double)]
202+
.collect()
203+
.toSeq
204+
assert(normalizeNaN(sparseResult) === Seq(
205+
(0, -4, "NaN"),
206+
(0, 0, "1.0"),
207+
(0, 2, "3.0"),
208+
(1, -5, "NaN"),
209+
(1, 1, "2.0"),
210+
(1, 3, "4.0"),
211+
(3, -11, "NaN"),
212+
(4, -1, "NaN")))
213+
214+
val schema = df.select(vector_posexplode($"vec")).schema
215+
assert(schema.simpleString === "struct<index:int,value:double>")
216+
}
217+
140218
test("test get_vector") {
141219
val df = Seq(
142220
(Vectors.dense(1.0, 2.0, 3.0), 0),

‎sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/analysis/FunctionRegistry.scala‎

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -30,6 +30,7 @@ import org.apache.spark.sql.AnalysisException
3030
import org.apache.spark.sql.catalyst.FunctionIdentifier
3131
import org.apache.spark.sql.catalyst.expressions._
3232
import org.apache.spark.sql.catalyst.expressions.aggregate._
33+
import org.apache.spark.sql.catalyst.expressions.ml._
3334
import org.apache.spark.sql.catalyst.expressions.st._
3435
import org.apache.spark.sql.catalyst.expressions.variant._
3536
import org.apache.spark.sql.catalyst.expressions.xml._
@@ -1137,6 +1138,7 @@ object FunctionRegistry {
11371138
registerInternalExpression[NullIndex]("null_index")
11381139
registerInternalExpression[CastTimestampNTZToLong]("timestamp_ntz_to_long")
11391140
registerInternalExpression[ArrayBinarySearch]("array_binary_search")
1141+
registerInternalExpression[VectorPosExplode]("vector_posexplode")
11401142

11411143
private def makeExprInfoForVirtualOperator(name: String, usage: String): ExpressionInfo = {
11421144
new ExpressionInfo(

0 commit comments

Comments
 (0)