From 0a4e6b5daa48194d16de7e5b12d2426123e3848a Mon Sep 17 00:00:00 2001 From: Ruifeng Zheng Date: Thu, 10 Sep 2026 10:02:27 +0000 Subject: [PATCH 1/7] [SPARK-59335][FOLLOWUP][ML] Reuse tree ensemble prediction helper --- .../ml/classification/GBTClassifier.scala | 11 ++------ .../spark/ml/regression/GBTRegressor.scala | 8 +----- .../ml/regression/RandomForestRegressor.scala | 2 +- .../org/apache/spark/ml/tree/treeModels.scala | 25 +++++++++++++++++++ 4 files changed, 29 insertions(+), 17 deletions(-) diff --git a/mllib/src/main/scala/org/apache/spark/ml/classification/GBTClassifier.scala b/mllib/src/main/scala/org/apache/spark/ml/classification/GBTClassifier.scala index ff7e449653c03..021fe76f711a9 100644 --- a/mllib/src/main/scala/org/apache/spark/ml/classification/GBTClassifier.scala +++ b/mllib/src/main/scala/org/apache/spark/ml/classification/GBTClassifier.scala @@ -429,15 +429,8 @@ class GBTClassificationModel private[ml]( TreeEnsembleModel.featureImportances(trees, numFeatures, perTreeNormalization = false) /** Raw prediction for the positive class. */ - private def margin(features: Vector): Double = { - var prediction = 0.0 - var i = 0 - while (i < _trees.length) { - prediction += _trees(i).rootNode.predictImpl(features).prediction * _treeWeights(i) - i += 1 - } - prediction - } + private def margin(features: Vector): Double = + TreeEnsembleModel.predict(features, _trees, _treeWeights) /** (private[ml]) Convert to a model in the old API */ private[ml] def toOld: OldGBTModel = { diff --git a/mllib/src/main/scala/org/apache/spark/ml/regression/GBTRegressor.scala b/mllib/src/main/scala/org/apache/spark/ml/regression/GBTRegressor.scala index 069cf39f19ada..a5e121c810369 100644 --- a/mllib/src/main/scala/org/apache/spark/ml/regression/GBTRegressor.scala +++ b/mllib/src/main/scala/org/apache/spark/ml/regression/GBTRegressor.scala @@ -307,13 +307,7 @@ class GBTRegressionModel private[ml]( override def predict(features: Vector): Double = { // TODO: When we add a generic Boosting class, handle transform there? SPARK-7129 // Classifies by thresholding sum of weighted tree predictions - var prediction = 0.0 - var i = 0 - while (i < _trees.length) { - prediction += _trees(i).rootNode.predictImpl(features).prediction * _treeWeights(i) - i += 1 - } - prediction + TreeEnsembleModel.predict(features, _trees, _treeWeights) } @Since("1.4.0") 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 4d9a039ecaca9..7f95d9ac6220e 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 @@ -266,7 +266,7 @@ class RandomForestRegressionModel private[ml] ( // TODO: When we add a generic Bagging class, handle transform there. SPARK-7128 // Predict average of tree predictions. // Ignore the weights since all are 1.0 for now. - _trees.map(_.rootNode.predictImpl(features).prediction).sum / getNumTrees + TreeEnsembleModel.predict(features, _trees) / getNumTrees } @Since("1.4.0") 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 d436fc896382b..3f05ef4a0f5be 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 @@ -169,6 +169,31 @@ private[spark] trait TreeEnsembleModel[M <: DecisionTreeModel] { private[ml] object TreeEnsembleModel { + private[ml] def predict[M <: DecisionTreeModel]( + features: Vector, + trees: Array[M], + treeWeights: Array[Double]): Double = { + var prediction = 0.0 + var i = 0 + while (i < trees.length) { + prediction += trees(i).rootNode.predictImpl(features).prediction * treeWeights(i) + i += 1 + } + prediction + } + + private[ml] def predict[M <: DecisionTreeModel]( + features: Vector, + trees: Array[M]): Double = { + var prediction = 0.0 + var i = 0 + while (i < trees.length) { + prediction += trees(i).rootNode.predictImpl(features).prediction + i += 1 + } + prediction + } + private[ml] def predict( features: Vector, rootNodes: Array[Node], From ccdfd053f6a7205e9bc08352856d86eed6765f1c Mon Sep 17 00:00:00 2001 From: Ruifeng Zheng Date: Thu, 10 Sep 2026 10:10:58 +0000 Subject: [PATCH 2/7] [SPARK-59400][ML] Inline GBT classification prediction helper --- .../apache/spark/ml/classification/GBTClassifier.scala | 8 ++------ 1 file changed, 2 insertions(+), 6 deletions(-) diff --git a/mllib/src/main/scala/org/apache/spark/ml/classification/GBTClassifier.scala b/mllib/src/main/scala/org/apache/spark/ml/classification/GBTClassifier.scala index 021fe76f711a9..d4e3bd458bccb 100644 --- a/mllib/src/main/scala/org/apache/spark/ml/classification/GBTClassifier.scala +++ b/mllib/src/main/scala/org/apache/spark/ml/classification/GBTClassifier.scala @@ -380,13 +380,13 @@ class GBTClassificationModel private[ml]( if (isDefined(thresholds)) { super.predict(features) } else { - if (margin(features) > 0.0) 1.0 else 0.0 + if (TreeEnsembleModel.predict(features, _trees, _treeWeights) > 0.0) 1.0 else 0.0 } } @Since("3.0.0") override def predictRaw(features: Vector): Vector = { - val prediction: Double = margin(features) + val prediction = TreeEnsembleModel.predict(features, _trees, _treeWeights) Vectors.dense(Array(-prediction, prediction)) } @@ -428,10 +428,6 @@ class GBTClassificationModel private[ml]( lazy val featureImportances: Vector = TreeEnsembleModel.featureImportances(trees, numFeatures, perTreeNormalization = false) - /** Raw prediction for the positive class. */ - private def margin(features: Vector): Double = - TreeEnsembleModel.predict(features, _trees, _treeWeights) - /** (private[ml]) Convert to a model in the old API */ private[ml] def toOld: OldGBTModel = { new OldGBTModel(OldAlgo.Classification, _trees.map(_.toOld), _treeWeights) From 6465c7a23d8f9898493cab7fe5bf602a261022f9 Mon Sep 17 00:00:00 2001 From: Ruifeng Zheng Date: Thu, 10 Sep 2026 10:14:51 +0000 Subject: [PATCH 3/7] [SPARK-59400][ML] Centralize random forest raw prediction --- .../RandomForestClassifier.scala | 42 +++++++++++-------- 1 file changed, 24 insertions(+), 18 deletions(-) diff --git a/mllib/src/main/scala/org/apache/spark/ml/classification/RandomForestClassifier.scala b/mllib/src/main/scala/org/apache/spark/ml/classification/RandomForestClassifier.scala index cf990923bc7fd..0385ff677674d 100644 --- a/mllib/src/main/scala/org/apache/spark/ml/classification/RandomForestClassifier.scala +++ b/mllib/src/main/scala/org/apache/spark/ml/classification/RandomForestClassifier.scala @@ -372,24 +372,8 @@ class RandomForestClassificationModel private[ml] ( } @Since("3.0.0") - override def predictRaw(features: Vector): Vector = { - // TODO: When we add a generic Bagging class, handle transform there: SPARK-7128 - // Classifies using majority votes. - // Ignore the tree weights since all are 1.0 for now. - val votes = Array.ofDim[Double](numClasses) - _trees.foreach { tree => - val classCounts = tree.rootNode.predictImpl(features).impurityStats.stats - val total = classCounts.sum - if (total != 0) { - var i = 0 - while (i < numClasses) { - votes(i) += classCounts(i) / total - i += 1 - } - } - } - Vectors.dense(votes) - } + override def predictRaw(features: Vector): Vector = + RandomForestClassificationModel.predictRaw(features, _trees, numClasses) override protected def raw2probabilityInPlace(rawPrediction: Vector): Vector = { rawPrediction match { @@ -469,6 +453,28 @@ class RandomForestClassificationModel private[ml] ( @Since("2.0.0") object RandomForestClassificationModel extends MLReadable[RandomForestClassificationModel] { + private def predictRaw( + features: Vector, + trees: Array[DecisionTreeClassificationModel], + numClasses: Int): Vector = { + // TODO: When we add a generic Bagging class, handle transform there: SPARK-7128 + // Classifies using majority votes. + // Ignore the tree weights since all are 1.0 for now. + val votes = Array.ofDim[Double](numClasses) + trees.foreach { tree => + val classCounts = tree.rootNode.predictImpl(features).impurityStats.stats + val total = classCounts.sum + if (total != 0) { + var i = 0 + while (i < numClasses) { + votes(i) += classCounts(i) / total + i += 1 + } + } + } + Vectors.dense(votes) + } + private def predictRaw( features: Vector, rootNodes: Array[Node], From eb12faa5da51663494a3fe809d4fe8bc2e7b0d06 Mon Sep 17 00:00:00 2001 From: Ruifeng Zheng Date: Thu, 10 Sep 2026 10:29:50 +0000 Subject: [PATCH 4/7] [SPARK-59400][ML] Remove stale random forest prediction comments --- .../spark/ml/classification/RandomForestClassifier.scala | 6 ------ 1 file changed, 6 deletions(-) diff --git a/mllib/src/main/scala/org/apache/spark/ml/classification/RandomForestClassifier.scala b/mllib/src/main/scala/org/apache/spark/ml/classification/RandomForestClassifier.scala index 0385ff677674d..26c5e2fc856dd 100644 --- a/mllib/src/main/scala/org/apache/spark/ml/classification/RandomForestClassifier.scala +++ b/mllib/src/main/scala/org/apache/spark/ml/classification/RandomForestClassifier.scala @@ -457,9 +457,6 @@ object RandomForestClassificationModel extends MLReadable[RandomForestClassifica features: Vector, trees: Array[DecisionTreeClassificationModel], numClasses: Int): Vector = { - // TODO: When we add a generic Bagging class, handle transform there: SPARK-7128 - // Classifies using majority votes. - // Ignore the tree weights since all are 1.0 for now. val votes = Array.ofDim[Double](numClasses) trees.foreach { tree => val classCounts = tree.rootNode.predictImpl(features).impurityStats.stats @@ -479,9 +476,6 @@ object RandomForestClassificationModel extends MLReadable[RandomForestClassifica features: Vector, rootNodes: Array[Node], numClasses: Int): Vector = { - // TODO: When we add a generic Bagging class, handle transform there: SPARK-7128 - // Classifies using majority votes. - // Ignore the tree weights since all are 1.0 for now. val votes = Array.ofDim[Double](numClasses) rootNodes.foreach { rootNode => val classCounts = rootNode.predictImpl(features).impurityStats.stats From 138483027caf264c33f3f89f6860a635f60d56f2 Mon Sep 17 00:00:00 2001 From: Ruifeng Zheng Date: Thu, 10 Sep 2026 10:31:33 +0000 Subject: [PATCH 5/7] [SPARK-59400][ML] Remove stale tree ensemble prediction comments --- .../scala/org/apache/spark/ml/regression/GBTRegressor.scala | 2 -- .../org/apache/spark/ml/regression/RandomForestRegressor.scala | 3 --- 2 files changed, 5 deletions(-) diff --git a/mllib/src/main/scala/org/apache/spark/ml/regression/GBTRegressor.scala b/mllib/src/main/scala/org/apache/spark/ml/regression/GBTRegressor.scala index a5e121c810369..d2ebbcc2563cc 100644 --- a/mllib/src/main/scala/org/apache/spark/ml/regression/GBTRegressor.scala +++ b/mllib/src/main/scala/org/apache/spark/ml/regression/GBTRegressor.scala @@ -305,8 +305,6 @@ class GBTRegressionModel private[ml]( } override def predict(features: Vector): Double = { - // TODO: When we add a generic Boosting class, handle transform there? SPARK-7129 - // Classifies by thresholding sum of weighted tree predictions TreeEnsembleModel.predict(features, _trees, _treeWeights) } 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 7f95d9ac6220e..e38af66e23b03 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 @@ -263,9 +263,6 @@ class RandomForestRegressionModel private[ml] ( } override def predict(features: Vector): Double = { - // TODO: When we add a generic Bagging class, handle transform there. SPARK-7128 - // Predict average of tree predictions. - // Ignore the weights since all are 1.0 for now. TreeEnsembleModel.predict(features, _trees) / getNumTrees } From 28aed82d60b20808afc6656bd68175232931f7a1 Mon Sep 17 00:00:00 2001 From: Ruifeng Zheng Date: Thu, 10 Sep 2026 10:49:18 +0000 Subject: [PATCH 6/7] [SPARK-59400][ML] Rename tree ensemble prediction helpers --- .../apache/spark/ml/classification/GBTClassifier.scala | 8 ++++---- .../org/apache/spark/ml/regression/GBTRegressor.scala | 4 ++-- .../spark/ml/regression/RandomForestRegressor.scala | 2 +- .../main/scala/org/apache/spark/ml/tree/treeModels.scala | 6 +++--- 4 files changed, 10 insertions(+), 10 deletions(-) diff --git a/mllib/src/main/scala/org/apache/spark/ml/classification/GBTClassifier.scala b/mllib/src/main/scala/org/apache/spark/ml/classification/GBTClassifier.scala index d4e3bd458bccb..05e3d0ab26657 100644 --- a/mllib/src/main/scala/org/apache/spark/ml/classification/GBTClassifier.scala +++ b/mllib/src/main/scala/org/apache/spark/ml/classification/GBTClassifier.scala @@ -369,7 +369,7 @@ class GBTClassificationModel private[ml]( }).apply(features) } else { udf((features: Vector) => { - val margin = TreeEnsembleModel.predict(features, localRootNodes, localTreeWeights) + val margin = TreeEnsembleModel.predictRaw(features, localRootNodes, localTreeWeights) if (margin > 0.0) 1.0 else 0.0 }).apply(features) } @@ -380,13 +380,13 @@ class GBTClassificationModel private[ml]( if (isDefined(thresholds)) { super.predict(features) } else { - if (TreeEnsembleModel.predict(features, _trees, _treeWeights) > 0.0) 1.0 else 0.0 + if (TreeEnsembleModel.predictRaw(features, _trees, _treeWeights) > 0.0) 1.0 else 0.0 } } @Since("3.0.0") override def predictRaw(features: Vector): Vector = { - val prediction = TreeEnsembleModel.predict(features, _trees, _treeWeights) + val prediction = TreeEnsembleModel.predictRaw(features, _trees, _treeWeights) Vectors.dense(Array(-prediction, prediction)) } @@ -459,7 +459,7 @@ object GBTClassificationModel extends MLReadable[GBTClassificationModel] { features: Vector, rootNodes: Array[Node], treeWeights: Array[Double]): Vector = { - val prediction = TreeEnsembleModel.predict(features, rootNodes, treeWeights) + val prediction = TreeEnsembleModel.predictRaw(features, rootNodes, treeWeights) Vectors.dense(-prediction, prediction) } diff --git a/mllib/src/main/scala/org/apache/spark/ml/regression/GBTRegressor.scala b/mllib/src/main/scala/org/apache/spark/ml/regression/GBTRegressor.scala index d2ebbcc2563cc..2d67c14ecbd27 100644 --- a/mllib/src/main/scala/org/apache/spark/ml/regression/GBTRegressor.scala +++ b/mllib/src/main/scala/org/apache/spark/ml/regression/GBTRegressor.scala @@ -280,7 +280,7 @@ class GBTRegressionModel private[ml]( if ($(predictionCol).nonEmpty) { val predUDF = udf { features: Vector => val (rootNodes, treeWeights) = bcTreeData.value - TreeEnsembleModel.predict(features, rootNodes, treeWeights) + TreeEnsembleModel.predictRaw(features, rootNodes, treeWeights) } predColNames :+= $(predictionCol) predCols :+= predUDF(col($(featuresCol))) @@ -305,7 +305,7 @@ class GBTRegressionModel private[ml]( } override def predict(features: Vector): Double = { - TreeEnsembleModel.predict(features, _trees, _treeWeights) + TreeEnsembleModel.predictRaw(features, _trees, _treeWeights) } @Since("1.4.0") 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 e38af66e23b03..e4f6d0a2d0882 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 @@ -263,7 +263,7 @@ class RandomForestRegressionModel private[ml] ( } override def predict(features: Vector): Double = { - TreeEnsembleModel.predict(features, _trees) / getNumTrees + TreeEnsembleModel.predictRaw(features, _trees) / getNumTrees } @Since("1.4.0") 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 3f05ef4a0f5be..f970726760b26 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 @@ -169,7 +169,7 @@ private[spark] trait TreeEnsembleModel[M <: DecisionTreeModel] { private[ml] object TreeEnsembleModel { - private[ml] def predict[M <: DecisionTreeModel]( + private[ml] def predictRaw[M <: DecisionTreeModel]( features: Vector, trees: Array[M], treeWeights: Array[Double]): Double = { @@ -182,7 +182,7 @@ private[ml] object TreeEnsembleModel { prediction } - private[ml] def predict[M <: DecisionTreeModel]( + private[ml] def predictRaw[M <: DecisionTreeModel]( features: Vector, trees: Array[M]): Double = { var prediction = 0.0 @@ -194,7 +194,7 @@ private[ml] object TreeEnsembleModel { prediction } - private[ml] def predict( + private[ml] def predictRaw( features: Vector, rootNodes: Array[Node], treeWeights: Array[Double]): Double = { From b61563e4974fe1f57a2b09798e40f79bdaf79418 Mon Sep 17 00:00:00 2001 From: Ruifeng Zheng Date: Thu, 10 Sep 2026 10:52:04 +0000 Subject: [PATCH 7/7] [SPARK-59400][ML] Simplify regression prediction overrides --- .../scala/org/apache/spark/ml/regression/GBTRegressor.scala | 3 +-- .../org/apache/spark/ml/regression/RandomForestRegressor.scala | 3 +-- 2 files changed, 2 insertions(+), 4 deletions(-) diff --git a/mllib/src/main/scala/org/apache/spark/ml/regression/GBTRegressor.scala b/mllib/src/main/scala/org/apache/spark/ml/regression/GBTRegressor.scala index 2d67c14ecbd27..c33ce2552b09a 100644 --- a/mllib/src/main/scala/org/apache/spark/ml/regression/GBTRegressor.scala +++ b/mllib/src/main/scala/org/apache/spark/ml/regression/GBTRegressor.scala @@ -304,9 +304,8 @@ class GBTRegressionModel private[ml]( } } - override def predict(features: Vector): Double = { + override def predict(features: Vector): Double = TreeEnsembleModel.predictRaw(features, _trees, _treeWeights) - } @Since("1.4.0") override def copy(extra: ParamMap): GBTRegressionModel = { 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 e4f6d0a2d0882..d0c342b450202 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 @@ -262,9 +262,8 @@ class RandomForestRegressionModel private[ml] ( } } - override def predict(features: Vector): Double = { + override def predict(features: Vector): Double = TreeEnsembleModel.predictRaw(features, _trees) / getNumTrees - } @Since("1.4.0") override def copy(extra: ParamMap): RandomForestRegressionModel = {