From 67e9f11f0975d598e1f8837d0402a53df8063100 Mon Sep 17 00:00:00 2001 From: Ruifeng Zheng Date: Mon, 3 Aug 2026 09:15:14 +0000 Subject: [PATCH] [SPARK-58509][ML][CONNECT] Exclude OneVsRest classifier from model size estimation --- .../org/apache/spark/ml/classification/OneVsRest.scala | 10 ++++++++-- .../spark/ml/classification/OneVsRestSuite.scala | 9 ++++++--- 2 files changed, 14 insertions(+), 5 deletions(-) diff --git a/mllib/src/main/scala/org/apache/spark/ml/classification/OneVsRest.scala b/mllib/src/main/scala/org/apache/spark/ml/classification/OneVsRest.scala index 38c68fd7000ca..42c6cf69600b5 100644 --- a/mllib/src/main/scala/org/apache/spark/ml/classification/OneVsRest.scala +++ b/mllib/src/main/scala/org/apache/spark/ml/classification/OneVsRest.scala @@ -150,8 +150,14 @@ final class OneVsRestModel private[ml] ( val numFeatures: Int = models.head.numFeatures private[spark] override def estimatedSize: Long = { - estimateMatadataSize + SizeEstimator.estimate(labelMetadata) + - models.iterator.map(_.estimatedSize).sum + var size = estimateMatadataSize(excluded = Seq( + // classifier: Param[ClassifierType] + classifier)) + // labelMetadata: Metadata + size += SizeEstimator.estimate(labelMetadata) + // models: Array[_ <: ClassificationModel[_, _]] + size += models.iterator.map(_.estimatedSize).sum + size } /** @group setParam */ diff --git a/mllib/src/test/scala/org/apache/spark/ml/classification/OneVsRestSuite.scala b/mllib/src/test/scala/org/apache/spark/ml/classification/OneVsRestSuite.scala index 70408731a20f6..75839a81c1f88 100644 --- a/mllib/src/test/scala/org/apache/spark/ml/classification/OneVsRestSuite.scala +++ b/mllib/src/test/scala/org/apache/spark/ml/classification/OneVsRestSuite.scala @@ -68,7 +68,7 @@ class OneVsRestSuite extends MLTest with DefaultReadWriteTest { ParamsSuite.checkParams(model) } - test("SPARK-58250: OneVsRestModel estimated size") { + test("SPARK-58250 and SPARK-58509: OneVsRestModel estimated size") { val trainingData = Seq( (0.0, Vectors.dense(0.0, 0.0)), (0.0, Vectors.dense(0.0, 1.0)), @@ -77,13 +77,16 @@ class OneVsRestSuite extends MLTest with DefaultReadWriteTest { (2.0, Vectors.dense(2.0, 0.0)), (2.0, Vectors.dense(2.0, 1.0))).toDF("label", "features") + val classifier = new LogisticRegression().setMaxIter(1) + // Initialize the classifier's logger before retaining it in the OneVsRestModel. + classifier.fit(trainingData) val model = new OneVsRest() - .setClassifier(new LogisticRegression().setMaxIter(1)) + .setClassifier(classifier) .fit(trainingData) val maxSize = 32 * 1024 assert(model.estimatedSize < maxSize, - s"Estimation (${model.estimatedSize}) should be less than $maxSize") + s"Estimation (${model.estimatedSize}) should not include shared runtime state") } test("one-vs-rest: default params") {