Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
102 changes: 45 additions & 57 deletions mllib/src/main/scala/org/apache/spark/ml/feature/StringIndexer.scala
Original file line number Diff line number Diff line change
Expand Up @@ -339,55 +339,40 @@ 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) = {
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).getOrElse {
throw new SparkException(s"Unseen label: $label. To handle unseen labels, " +
s"set Param handleInvalid to ${StringIndexer.KEEP_INVALID}.")
handleInvalid: String) = {
handleInvalid match {
case StringIndexer.KEEP_INVALID =>
val unknownIndex = labels.length.toDouble
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()
}
}

Expand All @@ -404,14 +389,6 @@ 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) {
filterInvalidData(dataset, inputColNames.toImmutableArraySeq, labelsToIndexArray)
} else {
dataset
}

for (i <- outputColNames.indices) {
val inputColName = inputColNames(i)
Expand All @@ -421,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, keepInvalid)
val indexer = getIndexer(labels.toImmutableArraySeq, labelToIndex, getHandleInvalid)

outputColumns(i) = indexer(dataset(inputColName).cast(StringType))
.as(outputColName, metadata)
Expand All @@ -443,10 +424,17 @@ class StringIndexerModel (

require(filteredOutputColNames.length == filteredOutputColumns.length)
if (filteredOutputColNames.length > 0) {
filteredDataset.withColumns(
val transformedDataset = dataset.withColumns(
filteredOutputColNames.toImmutableArraySeq, filteredOutputColumns.toImmutableArraySeq)
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)
} else {
transformedDataset
}
} else {
filteredDataset.toDF()
dataset.toDF()
}
}

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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")

Expand Down