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..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 @@ -363,28 +363,32 @@ class StringIndexerModel ( .where(conditions.reduce(_ and _)) } - private def getIndexer(labels: Seq[String], labelToIndex: OpenHashMap[String, Double]) = { - val keepInvalid = (getHandleInvalid == StringIndexer.KEEP_INVALID) - - udf { label: String => - if (label == null) { - if (keepInvalid) { - labels.length + private def getIndexer( + labels: Seq[String], + labelToIndex: OpenHashMap[String, Double], + keepInvalid: Boolean) = { + 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 { - 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).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") @@ -400,6 +404,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 +421,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)