From 62b7d1ca974f63d3b49e8f9e112f97ea1516671d Mon Sep 17 00:00:00 2001 From: Ruifeng Zheng Date: Tue, 4 Aug 2026 07:17:09 +0000 Subject: [PATCH 1/2] [SPARK-58546][ML] Avoid redundant StringIndexer label lookups --- .../spark/ml/feature/StringIndexer.scala | 28 +++++++++---------- 1 file changed, 13 insertions(+), 15 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 c0f7e312dcbb..efafa03155ba 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 @@ -363,9 +363,10 @@ class StringIndexerModel ( .where(conditions.reduce(_ and _)) } - private def getIndexer(labels: Seq[String], labelToIndex: OpenHashMap[String, Double]) = { - val keepInvalid = (getHandleInvalid == StringIndexer.KEEP_INVALID) - + private def getIndexer( + labels: Seq[String], + labelToIndex: OpenHashMap[String, Double], + keepInvalid: Boolean) = { udf { label: String => if (label == null) { if (keepInvalid) { @@ -375,13 +376,12 @@ class StringIndexerModel ( "NULLS, try setting StringIndexer.handleInvalid.") } } else { - if (labelToIndex.contains(label)) { - labelToIndex(label) - } else if (keepInvalid) { - labels.length - } else { - throw new SparkException(s"Unseen label: $label. To handle unseen labels, " + - s"set Param handleInvalid to ${StringIndexer.KEEP_INVALID}.") + labelToIndex.get(label) match { + case Some(index) => index + case None if keepInvalid => labels.length + case None => + throw new SparkException(s"Unseen label: $label. To handle unseen labels, " + + s"set Param handleInvalid to ${StringIndexer.KEEP_INVALID}.") } } }.asNondeterministic() @@ -400,6 +400,7 @@ class StringIndexerModel ( map } 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) { @@ -416,16 +417,13 @@ class StringIndexerModel ( try { dataset.col(inputColName) - val filteredLabels = getHandleInvalid match { - case StringIndexer.KEEP_INVALID => labels :+ "__unknown" - case _ => labels - } + val filteredLabels = if (keepInvalid) labels :+ "__unknown" else labels val metadata = NominalAttribute.defaultAttr .withName(outputColName) .withValues(filteredLabels) .toMetadata() - val indexer = getIndexer(labels.toImmutableArraySeq, labelToIndex) + val indexer = getIndexer(labels.toImmutableArraySeq, labelToIndex, keepInvalid) outputColumns(i) = indexer(dataset(inputColName).cast(StringType)) .as(outputColName, metadata) From 5848de685a681b023643c7ae083fda6df164e951 Mon Sep 17 00:00:00 2001 From: Ruifeng Zheng Date: Tue, 4 Aug 2026 07:25:50 +0000 Subject: [PATCH 2/2] [SPARK-58546][ML] Move StringIndexer invalid handling outside UDF --- .../spark/ml/feature/StringIndexer.scala | 28 +++++++++++-------- 1 file changed, 16 insertions(+), 12 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 efafa03155ba..bfb78adfd42d 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 @@ -367,24 +367,28 @@ class StringIndexerModel ( labels: Seq[String], labelToIndex: OpenHashMap[String, Double], keepInvalid: Boolean) = { - udf { label: String => - if (label == null) { - if (keepInvalid) { - labels.length + val unknownIndex = labels.length.toDouble + if (keepInvalid) { + udf { label: String => + if (label == null) { + unknownIndex } else { + labelToIndex.get(label).getOrElse(unknownIndex) + } + }.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) match { - case Some(index) => index - case None if keepInvalid => labels.length - case None => + } 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() + }.asNondeterministic() + } } @Since("2.0.0")