diff --git a/mllib/src/main/scala/org/apache/spark/ml/regression/RandomForestRegressor.scala b/mllib/src/main/scala/org/apache/spark/ml/regression/RandomForestRegressor.scala index d0c342b450202..c642568b73f20 100644 --- a/mllib/src/main/scala/org/apache/spark/ml/regression/RandomForestRegressor.scala +++ b/mllib/src/main/scala/org/apache/spark/ml/regression/RandomForestRegressor.scala @@ -238,17 +238,22 @@ class RandomForestRegressionModel private[ml] ( if ($(predictionCol).nonEmpty || $(leafCol).nonEmpty) { var predColNames = Seq.empty[String] var predCols = Seq.empty[Column] - val bcModel = dataset.sparkSession.sparkContext.broadcast(this) + val bcRootNodes = dataset.sparkSession.sparkContext.broadcast(_trees.map(_.rootNode)) if ($(predictionCol).nonEmpty) { - val predUDF = udf { features: Vector => bcModel.value.predict(features) } + val predUDF = udf { features: Vector => + val rootNodes = bcRootNodes.value + TreeEnsembleModel.predictRaw(features, rootNodes) / rootNodes.length + } predColNames :+= $(predictionCol) predCols :+= predUDF(col($(featuresCol))) .as($(predictionCol), outputSchema($(predictionCol)).metadata) } if ($(leafCol).nonEmpty) { - val leafUDF = udf { features: Vector => bcModel.value.predictLeaf(features) } + val leafUDF = udf { features: Vector => + TreeEnsembleModel.predictLeaf(features, bcRootNodes.value) + } predColNames :+= $(leafCol) predCols :+= 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 f970726760b26..b750411d5edf0 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 @@ -207,6 +207,16 @@ private[ml] object TreeEnsembleModel { prediction } + private[ml] def predictRaw(features: Vector, rootNodes: Array[Node]): Double = { + var prediction = 0.0 + var i = 0 + while (i < rootNodes.length) { + prediction += rootNodes(i).predictImpl(features).prediction + i += 1 + } + prediction + } + private[ml] def predictLeaf(features: Vector, rootNodes: Array[Node]): Vector = { val indices = Array.ofDim[Double](rootNodes.length) var i = 0