From 866c7bd43ddb88e76bf0b405342a4235f672ef61 Mon Sep 17 00:00:00 2001 From: Ruifeng Zheng Date: Tue, 4 Aug 2026 13:00:12 +0000 Subject: [PATCH 1/5] [WIP][ML] Avoid duplicate StringIndexer skip lookups --- .../spark/ml/feature/StringIndexer.scala | 57 +++++++------------ .../spark/ml/feature/StringIndexerSuite.scala | 20 +++++++ 2 files changed, 42 insertions(+), 35 deletions(-) diff --git a/mllib/src/main/scala/org/apache/spark/ml/feature/StringIndexer.scala b/mllib/src/main/scala/org/apache/spark/ml/feature/StringIndexer.scala index bfb78adfd42d0..f9783437dfbd3 100644 --- a/mllib/src/main/scala/org/apache/spark/ml/feature/StringIndexer.scala +++ b/mllib/src/main/scala/org/apache/spark/ml/feature/StringIndexer.scala @@ -339,34 +339,11 @@ class StringIndexerModel ( @Since("3.0.0") def setOutputCols(value: Array[String]): this.type = set(outputCols, value) - // This filters out any null values and also the input labels which are not in - // the dataset used for fitting. - private def filterInvalidData( - dataset: Dataset[_], - inputColNames: Seq[String], - labelsToIndexArray: Array[OpenHashMap[String, Double]]): Dataset[_] = { - val conditions: Seq[Column] = inputColNames.indices.map { i => - val inputColName = inputColNames(i) - val labelToIndex = labelsToIndexArray(i) - // We have this additional lookup at `labelToIndex` when `handleInvalid` is set to - // `StringIndexer.SKIP_INVALID`. Another idea is to do this lookup natively by SQL - // expression, however, lookup for a key in a map is not efficient in SparkSQL now. - // See `ElementAt` and `GetMapValue` expressions. If SQL's map lookup is improved, - // we can consider to change this. - val filter = udf { label: String => - labelToIndex.contains(label) - } - filter(dataset(inputColName)) - } - - dataset.na.drop(inputColNames.filter(dataset.schema.fieldNames.contains(_))) - .where(conditions.reduce(_ and _)) - } - private def getIndexer( labels: Seq[String], labelToIndex: OpenHashMap[String, Double], - keepInvalid: Boolean) = { + keepInvalid: Boolean, + skipInvalid: Boolean) = { val unknownIndex = labels.length.toDouble if (keepInvalid) { udf { label: String => @@ -376,6 +353,14 @@ class StringIndexerModel ( labelToIndex.get(label).getOrElse(unknownIndex) } }.asNondeterministic() + } else if (skipInvalid) { + udf { label: String => + if (label == null) { + null + } else { + labelToIndex.get(label).map(index => java.lang.Double.valueOf(index)).orNull + } + }.asNondeterministic() } else { udf { label: String => if (label == null) { @@ -405,13 +390,7 @@ class StringIndexerModel ( } val outputColumns = new Array[Column](outputColNames.length) val keepInvalid = getHandleInvalid == StringIndexer.KEEP_INVALID - - // Skips invalid rows if `handleInvalid` is set to `StringIndexer.SKIP_INVALID`. - val filteredDataset = if (getHandleInvalid == StringIndexer.SKIP_INVALID) { - filterInvalidData(dataset, inputColNames.toImmutableArraySeq, labelsToIndexArray) - } else { - dataset - } + val skipInvalid = getHandleInvalid == StringIndexer.SKIP_INVALID for (i <- outputColNames.indices) { val inputColName = inputColNames(i) @@ -427,7 +406,8 @@ class StringIndexerModel ( .withValues(filteredLabels) .toMetadata() - val indexer = getIndexer(labels.toImmutableArraySeq, labelToIndex, keepInvalid) + val indexer = + getIndexer(labels.toImmutableArraySeq, labelToIndex, keepInvalid, skipInvalid) outputColumns(i) = indexer(dataset(inputColName).cast(StringType)) .as(outputColName, metadata) @@ -443,10 +423,17 @@ class StringIndexerModel ( require(filteredOutputColNames.length == filteredOutputColumns.length) if (filteredOutputColNames.length > 0) { - filteredDataset.withColumns( + val transformedDataset = dataset.withColumns( filteredOutputColNames.toImmutableArraySeq, filteredOutputColumns.toImmutableArraySeq) + if (skipInvalid) { + // The skip indexers return null for invalid labels. Their nondeterminism keeps this filter + // above the projection, so each label is looked up only once. + transformedDataset.na.drop(filteredOutputColNames) + } else { + transformedDataset + } } else { - filteredDataset.toDF() + dataset.toDF() } } diff --git a/mllib/src/test/scala/org/apache/spark/ml/feature/StringIndexerSuite.scala b/mllib/src/test/scala/org/apache/spark/ml/feature/StringIndexerSuite.scala index 18005687f5656..01b53993e9f01 100644 --- a/mllib/src/test/scala/org/apache/spark/ml/feature/StringIndexerSuite.scala +++ b/mllib/src/test/scala/org/apache/spark/ml/feature/StringIndexerSuite.scala @@ -519,6 +519,26 @@ class StringIndexerSuite extends MLTest with DefaultReadWriteTest { } } + test("StringIndexer skips rows with invalid multiple input columns") { + val training = Seq(("a", "e"), ("b", "f")).toDF("label1", "label2") + val data = Seq( + (0, "a", "e"), + (1, "c", "e"), + (2, "a", "g"), + (3, null, "e"), + (4, "b", "f") + ).toDF("id", "label1", "label2") + + val model = new StringIndexer() + .setInputCols(Array("label1", "label2")) + .setOutputCols(Array("labelIndex1", "labelIndex2")) + .setHandleInvalid("skip") + .fit(training) + + val transformed = model.transform(data).select("id", "labelIndex1", "labelIndex2") + checkAnswer(transformed, Seq(Row(0, 0.0, 0.0), Row(4, 1.0, 1.0))) + } + test("Correctly skipping NULL and NaN values") { val df = Seq(("a", Double.NaN), (null, 1.0), ("b", 2.0), (null, 3.0)).toDF("str", "double") From ad931d15ee4c99ac0f8c05a3b1a4ccd277a48b09 Mon Sep 17 00:00:00 2001 From: Ruifeng Zheng Date: Tue, 4 Aug 2026 13:03:49 +0000 Subject: [PATCH 2/5] [WIP][ML] Pass StringIndexer invalid handling mode to indexer --- .../apache/spark/ml/feature/StringIndexer.scala | 15 +++++++-------- 1 file changed, 7 insertions(+), 8 deletions(-) diff --git a/mllib/src/main/scala/org/apache/spark/ml/feature/StringIndexer.scala b/mllib/src/main/scala/org/apache/spark/ml/feature/StringIndexer.scala index f9783437dfbd3..646d6fcaeac31 100644 --- a/mllib/src/main/scala/org/apache/spark/ml/feature/StringIndexer.scala +++ b/mllib/src/main/scala/org/apache/spark/ml/feature/StringIndexer.scala @@ -342,10 +342,9 @@ class StringIndexerModel ( private def getIndexer( labels: Seq[String], labelToIndex: OpenHashMap[String, Double], - keepInvalid: Boolean, - skipInvalid: Boolean) = { + handleInvalid: String) = { val unknownIndex = labels.length.toDouble - if (keepInvalid) { + if (handleInvalid == StringIndexer.KEEP_INVALID) { udf { label: String => if (label == null) { unknownIndex @@ -353,7 +352,7 @@ class StringIndexerModel ( labelToIndex.get(label).getOrElse(unknownIndex) } }.asNondeterministic() - } else if (skipInvalid) { + } else if (handleInvalid == StringIndexer.SKIP_INVALID) { udf { label: String => if (label == null) { null @@ -389,8 +388,9 @@ class StringIndexerModel ( map } val outputColumns = new Array[Column](outputColNames.length) - val keepInvalid = getHandleInvalid == StringIndexer.KEEP_INVALID - val skipInvalid = getHandleInvalid == StringIndexer.SKIP_INVALID + val handleInvalid = getHandleInvalid + val keepInvalid = handleInvalid == StringIndexer.KEEP_INVALID + val skipInvalid = handleInvalid == StringIndexer.SKIP_INVALID for (i <- outputColNames.indices) { val inputColName = inputColNames(i) @@ -406,8 +406,7 @@ class StringIndexerModel ( .withValues(filteredLabels) .toMetadata() - val indexer = - getIndexer(labels.toImmutableArraySeq, labelToIndex, keepInvalid, skipInvalid) + val indexer = getIndexer(labels.toImmutableArraySeq, labelToIndex, handleInvalid) outputColumns(i) = indexer(dataset(inputColName).cast(StringType)) .as(outputColName, metadata) From d02b01cf75c8ef6ec1bc687c4f8f01193354f01b Mon Sep 17 00:00:00 2001 From: Ruifeng Zheng Date: Tue, 4 Aug 2026 13:05:30 +0000 Subject: [PATCH 3/5] [WIP][ML] Match StringIndexer invalid handling mode --- .../spark/ml/feature/StringIndexer.scala | 55 ++++++++++--------- 1 file changed, 28 insertions(+), 27 deletions(-) diff --git a/mllib/src/main/scala/org/apache/spark/ml/feature/StringIndexer.scala b/mllib/src/main/scala/org/apache/spark/ml/feature/StringIndexer.scala index 646d6fcaeac31..4f2be5a73ea5e 100644 --- a/mllib/src/main/scala/org/apache/spark/ml/feature/StringIndexer.scala +++ b/mllib/src/main/scala/org/apache/spark/ml/feature/StringIndexer.scala @@ -344,34 +344,35 @@ class StringIndexerModel ( labelToIndex: OpenHashMap[String, Double], handleInvalid: String) = { val unknownIndex = labels.length.toDouble - if (handleInvalid == StringIndexer.KEEP_INVALID) { - udf { label: String => - if (label == null) { - unknownIndex - } else { - labelToIndex.get(label).getOrElse(unknownIndex) - } - }.asNondeterministic() - } else if (handleInvalid == StringIndexer.SKIP_INVALID) { - udf { label: String => - if (label == null) { - null - } else { - labelToIndex.get(label).map(index => java.lang.Double.valueOf(index)).orNull - } - }.asNondeterministic() - } else { - udf { label: String => - if (label == null) { - throw new SparkException("StringIndexer encountered NULL value. To handle or skip " + - "NULLS, try setting StringIndexer.handleInvalid.") - } else { - labelToIndex.get(label).getOrElse { - throw new SparkException(s"Unseen label: $label. To handle unseen labels, " + - s"set Param handleInvalid to ${StringIndexer.KEEP_INVALID}.") + handleInvalid match { + case StringIndexer.KEEP_INVALID => + udf { label: String => + if (label == null) { + unknownIndex + } else { + labelToIndex.get(label).getOrElse(unknownIndex) } - } - }.asNondeterministic() + }.asNondeterministic() + case StringIndexer.SKIP_INVALID => + udf { label: String => + if (label == null) { + null + } else { + labelToIndex.get(label).map(index => java.lang.Double.valueOf(index)).orNull + } + }.asNondeterministic() + case _ => + udf { label: String => + if (label == null) { + throw new SparkException("StringIndexer encountered NULL value. To handle or skip " + + "NULLS, try setting StringIndexer.handleInvalid.") + } else { + labelToIndex.get(label).getOrElse { + throw new SparkException(s"Unseen label: $label. To handle unseen labels, " + + s"set Param handleInvalid to ${StringIndexer.KEEP_INVALID}.") + } + } + }.asNondeterministic() } } From 3da44580a9781720010352a1a856b368d8d1c195 Mon Sep 17 00:00:00 2001 From: Ruifeng Zheng Date: Tue, 4 Aug 2026 13:06:45 +0000 Subject: [PATCH 4/5] [WIP][ML] Scope StringIndexer unknown index to keep handling --- .../main/scala/org/apache/spark/ml/feature/StringIndexer.scala | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/mllib/src/main/scala/org/apache/spark/ml/feature/StringIndexer.scala b/mllib/src/main/scala/org/apache/spark/ml/feature/StringIndexer.scala index 4f2be5a73ea5e..f9b5441220775 100644 --- a/mllib/src/main/scala/org/apache/spark/ml/feature/StringIndexer.scala +++ b/mllib/src/main/scala/org/apache/spark/ml/feature/StringIndexer.scala @@ -343,9 +343,9 @@ class StringIndexerModel ( labels: Seq[String], labelToIndex: OpenHashMap[String, Double], handleInvalid: String) = { - val unknownIndex = labels.length.toDouble handleInvalid match { case StringIndexer.KEEP_INVALID => + val unknownIndex = labels.length.toDouble udf { label: String => if (label == null) { unknownIndex From 199a0f774559900e307576dd9da8f0d9f1861e7f Mon Sep 17 00:00:00 2001 From: Ruifeng Zheng Date: Tue, 4 Aug 2026 13:11:10 +0000 Subject: [PATCH 5/5] [WIP][ML] Read StringIndexer invalid handling mode directly --- .../org/apache/spark/ml/feature/StringIndexer.scala | 13 +++++++------ 1 file changed, 7 insertions(+), 6 deletions(-) diff --git a/mllib/src/main/scala/org/apache/spark/ml/feature/StringIndexer.scala b/mllib/src/main/scala/org/apache/spark/ml/feature/StringIndexer.scala index f9b5441220775..d166b3d3ade4b 100644 --- a/mllib/src/main/scala/org/apache/spark/ml/feature/StringIndexer.scala +++ b/mllib/src/main/scala/org/apache/spark/ml/feature/StringIndexer.scala @@ -389,9 +389,6 @@ class StringIndexerModel ( map } val outputColumns = new Array[Column](outputColNames.length) - val handleInvalid = getHandleInvalid - val keepInvalid = handleInvalid == StringIndexer.KEEP_INVALID - val skipInvalid = handleInvalid == StringIndexer.SKIP_INVALID for (i <- outputColNames.indices) { val inputColName = inputColNames(i) @@ -401,13 +398,17 @@ class StringIndexerModel ( try { dataset.col(inputColName) - val filteredLabels = if (keepInvalid) labels :+ "__unknown" else labels + val filteredLabels = if (getHandleInvalid == StringIndexer.KEEP_INVALID) { + labels :+ "__unknown" + } else { + labels + } val metadata = NominalAttribute.defaultAttr .withName(outputColName) .withValues(filteredLabels) .toMetadata() - val indexer = getIndexer(labels.toImmutableArraySeq, labelToIndex, handleInvalid) + val indexer = getIndexer(labels.toImmutableArraySeq, labelToIndex, getHandleInvalid) outputColumns(i) = indexer(dataset(inputColName).cast(StringType)) .as(outputColName, metadata) @@ -425,7 +426,7 @@ class StringIndexerModel ( if (filteredOutputColNames.length > 0) { val transformedDataset = dataset.withColumns( filteredOutputColNames.toImmutableArraySeq, filteredOutputColumns.toImmutableArraySeq) - if (skipInvalid) { + if (getHandleInvalid == StringIndexer.SKIP_INVALID) { // The skip indexers return null for invalid labels. Their nondeterminism keeps this filter // above the projection, so each label is looked up only once. transformedDataset.na.drop(filteredOutputColNames)