diff --git a/mllib/src/main/scala/org/apache/spark/ml/classification/DecisionTreeClassifier.scala b/mllib/src/main/scala/org/apache/spark/ml/classification/DecisionTreeClassifier.scala index 11d43d1b31fb1..d83df697a04a6 100644 --- a/mllib/src/main/scala/org/apache/spark/ml/classification/DecisionTreeClassifier.scala +++ b/mllib/src/main/scala/org/apache/spark/ml/classification/DecisionTreeClassifier.scala @@ -209,6 +209,44 @@ class DecisionTreeClassificationModel private[ml] ( rootNode.predictImpl(features).prediction } + override protected def predictRawColumn(features: Column): Column = { + val localRootNode = rootNode + udf((features: Vector) => + DecisionTreeClassificationModel.predictRaw(features, localRootNode) + ).apply(features) + } + + override protected def raw2probabilityColumn(rawPrediction: Column): Column = { + udf((rawPrediction: Vector) => + DecisionTreeClassificationModel.raw2probability(rawPrediction) + ).apply(rawPrediction) + } + + override protected def predictProbabilityColumn(features: Column): Column = { + val localRootNode = rootNode + udf((features: Vector) => { + val rawPrediction = DecisionTreeClassificationModel.predictRaw(features, localRootNode) + DecisionTreeClassificationModel.raw2probability(rawPrediction) + }).apply(features) + } + + override protected def raw2predictionColumn(rawPrediction: Column): Column = { + if (isDefined(thresholds)) { + val localThresholds = getThresholds.clone() + udf((rawPrediction: Vector) => { + val probability = DecisionTreeClassificationModel.raw2probability(rawPrediction) + ProbabilisticClassificationModel.probability2prediction(probability, localThresholds) + }).apply(rawPrediction) + } else { + udf((rawPrediction: Vector) => rawPrediction.argmax.toDouble).apply(rawPrediction) + } + } + + override protected def predictionColumn(features: Column): Column = { + val localRootNode = rootNode + udf((features: Vector) => localRootNode.predictImpl(features).prediction).apply(features) + } + @Since("3.0.0") override def transformSchema(schema: StructType): StructType = { var outputSchema = super.transformSchema(schema) @@ -223,7 +261,10 @@ class DecisionTreeClassificationModel private[ml] ( val outputData = super.transform(dataset) if ($(leafCol).nonEmpty) { - val leafUDF = udf { features: Vector => predictLeaf(features) } + val localRootNode = rootNode + val leafUDF = udf { features: Vector => + DecisionTreeModel.predictLeaf(features, localRootNode) + } outputData.withColumn($(leafCol), leafUDF(col($(featuresCol))), outputSchema($(leafCol)).metadata) } else { @@ -233,7 +274,7 @@ class DecisionTreeClassificationModel private[ml] ( @Since("3.0.0") override def predictRaw(features: Vector): Vector = { - Vectors.dense(rootNode.predictImpl(features).impurityStats.stats.clone()) + DecisionTreeClassificationModel.predictRaw(features, rootNode) } override protected def raw2probabilityInPlace(rawPrediction: Vector): Vector = { @@ -291,6 +332,16 @@ class DecisionTreeClassificationModel private[ml] ( @Since("2.0.0") object DecisionTreeClassificationModel extends MLReadable[DecisionTreeClassificationModel] { + private def predictRaw(features: Vector, rootNode: Node): Vector = { + Vectors.dense(rootNode.predictImpl(features).impurityStats.stats.clone()) + } + + private def raw2probability(rawPrediction: Vector): Vector = { + val probability = rawPrediction.copy.toDense + ProbabilisticClassificationModel.normalizeToProbabilitiesInPlace(probability) + probability + } + @Since("2.0.0") override def read: MLReader[DecisionTreeClassificationModel] = new DecisionTreeClassificationModelReader diff --git a/mllib/src/main/scala/org/apache/spark/ml/classification/ProbabilisticClassifier.scala b/mllib/src/main/scala/org/apache/spark/ml/classification/ProbabilisticClassifier.scala index 2e3e681df9391..8be033d22db84 100644 --- a/mllib/src/main/scala/org/apache/spark/ml/classification/ProbabilisticClassifier.scala +++ b/mllib/src/main/scala/org/apache/spark/ml/classification/ProbabilisticClassifier.scala @@ -194,7 +194,14 @@ abstract class ProbabilisticClassificationModel[ * @note This method honors [[thresholds]] when they are set. */ protected def probability2predictionColumn(probability: Column): Column = { - udf(probability2prediction _).apply(probability) + if (isDefined(thresholds)) { + val localThresholds = getThresholds.clone() + udf((probability: Vector) => + ProbabilisticClassificationModel.probability2prediction(probability, localThresholds) + ).apply(probability) + } else { + udf((probability: Vector) => probability.argmax.toDouble).apply(probability) + } } /** @group setParam */ diff --git a/mllib/src/main/scala/org/apache/spark/ml/regression/DecisionTreeRegressor.scala b/mllib/src/main/scala/org/apache/spark/ml/regression/DecisionTreeRegressor.scala index 4a5c3b5790354..79339807a2b6b 100644 --- a/mllib/src/main/scala/org/apache/spark/ml/regression/DecisionTreeRegressor.scala +++ b/mllib/src/main/scala/org/apache/spark/ml/regression/DecisionTreeRegressor.scala @@ -218,26 +218,33 @@ class DecisionTreeRegressionModel private[ml] ( @Since("2.0.0") override def transform(dataset: Dataset[_]): DataFrame = { val outputSchema = transformSchema(dataset.schema, logging = true) + val localRootNode = rootNode var predictionColNames = Seq.empty[String] var predictionColumns = Seq.empty[Column] if ($(predictionCol).nonEmpty) { - val predictUDF = udf { features: Vector => predict(features) } + val predictUDF = udf { features: Vector => + localRootNode.predictImpl(features).prediction + } predictionColNames :+= $(predictionCol) predictionColumns :+= predictUDF(col($(featuresCol))) .as($(predictionCol), outputSchema($(predictionCol)).metadata) } if (isDefined(varianceCol) && $(varianceCol).nonEmpty) { - val predictVarianceUDF = udf { features: Vector => predictVariance(features) } + val predictVarianceUDF = udf { features: Vector => + localRootNode.predictImpl(features).impurityStats.calculate() + } predictionColNames :+= $(varianceCol) predictionColumns :+= predictVarianceUDF(col($(featuresCol))) .as($(varianceCol), outputSchema($(varianceCol)).metadata) } if ($(leafCol).nonEmpty) { - val leafUDF = udf { features: Vector => predictLeaf(features) } + val leafUDF = udf { features: Vector => + DecisionTreeModel.predictLeaf(features, localRootNode) + } predictionColNames :+= $(leafCol) predictionColumns :+= leafUDF(col($(featuresCol))) .as($(leafCol), outputSchema($(leafCol)).metadata) diff --git a/mllib/src/main/scala/org/apache/spark/ml/tree/treeModels.scala b/mllib/src/main/scala/org/apache/spark/ml/tree/treeModels.scala index 5f45cc3315621..70d5589ab7b6a 100644 --- a/mllib/src/main/scala/org/apache/spark/ml/tree/treeModels.scala +++ b/mllib/src/main/scala/org/apache/spark/ml/tree/treeModels.scala @@ -94,9 +94,7 @@ private[spark] trait DecisionTreeModel { * Leaves are indexed in pre-order from 0. */ def predictLeaf(features: Vector): Double = { - val leaf = rootNode.predictImpl(features) - assert(leaf.leafIndex >= 0, "Leaf indices are not assigned.") - leaf.leafIndex.toDouble + DecisionTreeModel.predictLeaf(features, rootNode) } def getEstimatedSize(): Long = { @@ -104,6 +102,15 @@ private[spark] trait DecisionTreeModel { } } +private[spark] object DecisionTreeModel { + + private[ml] def predictLeaf(features: Vector, rootNode: Node): Double = { + val leaf = rootNode.predictImpl(features) + assert(leaf.leafIndex >= 0, "Leaf indices are not assigned.") + leaf.leafIndex.toDouble + } +} + /** * Abstraction for models which are ensembles of decision trees * @tparam M Type of tree model in this ensemble