From 5e71d731247692356ad6aa0028b88188fe936204 Mon Sep 17 00:00:00 2001 From: Ruifeng Zheng Date: Fri, 31 Jul 2026 12:43:20 +0000 Subject: [PATCH] [ML] Improve TrainValidationSplitModel size estimation --- .../ml/tuning/TrainValidationSplit.scala | 25 ++++++++++++++++++- .../ml/tuning/TrainValidationSplitSuite.scala | 15 +++++++++++ 2 files changed, 39 insertions(+), 1 deletion(-) diff --git a/mllib/src/main/scala/org/apache/spark/ml/tuning/TrainValidationSplit.scala b/mllib/src/main/scala/org/apache/spark/ml/tuning/TrainValidationSplit.scala index 6ee64ef99a668..5b0b77abb2dd0 100644 --- a/mllib/src/main/scala/org/apache/spark/ml/tuning/TrainValidationSplit.scala +++ b/mllib/src/main/scala/org/apache/spark/ml/tuning/TrainValidationSplit.scala @@ -38,8 +38,8 @@ import org.apache.spark.ml.util._ import org.apache.spark.ml.util.Instrumentation.instrumented import org.apache.spark.sql.{DataFrame, Dataset} import org.apache.spark.sql.types.StructType +import org.apache.spark.util.{SizeEstimator, ThreadUtils} import org.apache.spark.util.ArrayImplicits._ -import org.apache.spark.util.ThreadUtils /** * Params for [[TrainValidationSplit]] and [[TrainValidationSplitModel]]. @@ -293,6 +293,29 @@ class TrainValidationSplitModel private[ml] ( @Since("2.3.0") def hasSubModels: Boolean = _subModels.isDefined + private[spark] override def estimatedSize: Long = { + var size = estimateMatadataSize(excluded = Seq( + // estimator: Param[Estimator[_]] + estimator, + // estimatorParamMaps: Param[Array[ParamMap]] + estimatorParamMaps, + // evaluator: Param[Evaluator] + evaluator)) + // bestModel: Model[_] + size += bestModel.estimatedSize + // validationMetrics: Array[Double] + size += SizeEstimator.estimate(validationMetrics) + // _subModels: Option[Array[Model[_]]] + _subModels.foreach { modelArray => + modelArray.foreach { model => + if (model != null) { + size += model.estimatedSize + } + } + } + size + } + @Since("2.0.0") override def transform(dataset: Dataset[_]): DataFrame = { transformSchema(dataset.schema, logging = true) diff --git a/mllib/src/test/scala/org/apache/spark/ml/tuning/TrainValidationSplitSuite.scala b/mllib/src/test/scala/org/apache/spark/ml/tuning/TrainValidationSplitSuite.scala index 20ba69a5adbb9..0b4b022f7e6c4 100644 --- a/mllib/src/test/scala/org/apache/spark/ml/tuning/TrainValidationSplitSuite.scala +++ b/mllib/src/test/scala/org/apache/spark/ml/tuning/TrainValidationSplitSuite.scala @@ -73,6 +73,21 @@ class TrainValidationSplitSuite } } + test("TrainValidationSplitModel estimated size") { + val estimator = new LogisticRegression().setMaxIter(1) + // Initialize the estimator's logger before retaining it in the TrainValidationSplitModel. + estimator.fit(dataset) + val model = new TrainValidationSplit() + .setEstimator(estimator) + .setEstimatorParamMaps(Array(ParamMap.empty)) + .setEvaluator(new BinaryClassificationEvaluator()) + .fit(dataset) + + val maxSize = 16 * 1024 + assert(model.estimatedSize < maxSize, + s"Estimation (${model.estimatedSize}) should not include shared runtime state") + } + test("train validation with linear regression") { val dataset = sc.parallelize( LinearDataGenerator.generateLinearInput(