From 6ae5f6656c3600f8f5c5f66318d19d7a8c59b2bc Mon Sep 17 00:00:00 2001 From: Dustin Smith Date: Wed, 30 Sep 2026 00:28:30 +0700 Subject: [PATCH] fix: add native metrics from every plan a task runs instead of keeping the last Native plans publish absolute metric values, and CometMetricNode.set wrote them into the SQL metric with set. When one task runs several native plans over the same metric tree, as a coalesce without a shuffle does over a Comet RDD, each plan overwrote the one before it, so the task reported only the last plan: a COALESCE(1) over four files showed 2500 of 10000 scanned rows and one file's bytes, and the task input metrics derived from them were short too. Each native plan now gets its own copy of the metric tree, sharing the same SQL metrics, which remembers the last value that plan reported and adds only the increase, so periodic updates are not counted twice and several plans add up. The peak memory gauges keep the largest value instead. --- docs/source/user-guide/latest/metrics.md | 6 +- .../org/apache/comet/CometExecIterator.scala | 4 +- .../spark/sql/comet/CometMetricNode.scala | 38 +++- .../CometExecIteratorLifecycleSuite.scala | 2 + .../sql/comet/CometTaskMetricsSuite.scala | 207 +++++++++++++++++- 5 files changed, 250 insertions(+), 7 deletions(-) diff --git a/docs/source/user-guide/latest/metrics.md b/docs/source/user-guide/latest/metrics.md index ba198b97d60..960bfb47e1a 100644 --- a/docs/source/user-guide/latest/metrics.md +++ b/docs/source/user-guide/latest/metrics.md @@ -144,7 +144,11 @@ so this measures memory released from shuffle buffering, not necessarily a drop ## Native Metrics Setting `spark.comet.explain.native.enabled=true` will cause native plans to be logged in each executor. Metrics are -logged for each native plan (and there is one plan per task, so this is very verbose). +logged for each native plan (and there is usually one plan per task, so this is very verbose). + +The SQL metrics that native plans report accumulate per task. When one task runs several native plans, such as a +coalesce without a shuffle that reads several partitions, each metric reports their sum, while the two memory +high-water marks, `peak_mem_used` and `build_mem_used`, report the largest value among them. Here is a guide to some of the native metrics. diff --git a/spark/src/main/scala/org/apache/comet/CometExecIterator.scala b/spark/src/main/scala/org/apache/comet/CometExecIterator.scala index 498086e3ec9..74a54427d1b 100644 --- a/spark/src/main/scala/org/apache/comet/CometExecIterator.scala +++ b/spark/src/main/scala/org/apache/comet/CometExecIterator.scala @@ -128,7 +128,9 @@ class CometExecIterator( protobufQueryPlan, protobufSparkConfigs, numParts, - nativeMetrics, + // A copy per native plan, so plans that share the tree within one task add up their + // metrics instead of overwriting each other (see CometMetricNode.set). + nativeMetrics.newInstance(), metricsUpdateInterval = COMET_METRICS_UPDATE_INTERVAL.get(), cometTaskMemoryManager, localDiskDirs, diff --git a/spark/src/main/scala/org/apache/spark/sql/comet/CometMetricNode.scala b/spark/src/main/scala/org/apache/spark/sql/comet/CometMetricNode.scala index 36e679be5aa..ad137722637 100644 --- a/spark/src/main/scala/org/apache/spark/sql/comet/CometMetricNode.scala +++ b/spark/src/main/scala/org/apache/spark/sql/comet/CometMetricNode.scala @@ -22,6 +22,7 @@ package org.apache.spark.sql.comet import java.util.IdentityHashMap import java.util.concurrent.ConcurrentHashMap +import scala.collection.mutable import scala.jdk.CollectionConverters._ import org.apache.spark.{SparkContext, TaskContext} @@ -44,6 +45,10 @@ import org.apache.comet.serde.Metric case class CometMetricNode(metrics: Map[String, SQLMetric], children: Seq[CometMetricNode]) extends Logging { + // The highest value each metric has reported through this node (see [[set]]). A body field + // rather than a constructor parameter, so it takes no part in equality. + private val lastReported = mutable.HashMap.empty[String, Long] + /** * Returns the leaf node (deepest single-child descendant). For a native scan plan like * FilterExec -> DataSourceExec, this returns the DataSourceExec node which has the @@ -201,17 +206,37 @@ case class CometMetricNode(metrics: Map[String, SQLMetric], children: Seq[CometM } /** - * Update the value of a metric. This method will typically be called multiple times for the - * same metric during multiple calls to executePlan. + * Returns a copy of this tree for one native plan. The copy updates the same `SQLMetric`s but + * tracks its own last reported values, so [[set]] can add what this plan reports to what + * earlier plans in the same task reported. + */ + def newInstance(): CometMetricNode = CometMetricNode(metrics, children.map(_.newInstance())) + + /** + * Folds a value reported by native code into a metric. Called from native, typically many times + * for the same metric while the plan runs. + * + * Native code reports the absolute value for its own plan, while the `SQLMetric` is shared by + * every native plan that a task runs over this tree, such as one plan per parent partition + * under a coalesce. Each plan works on its own copy of the tree (see [[newInstance]]), and the + * metric grows by the increase since that copy's last report, so plans that run one after + * another in a task add up. A value below an earlier report adds nothing and the last reported + * value keeps the higher one; recording the lower value would count the recovery a second time. + * Gauges such as peak memory keep the maximum across the task's plans instead. * * @param metricName * the name of the metric at native operator. * @param v - * the value to set. + * the absolute value native code reports for this plan. */ def set(metricName: String, v: Long): Unit = { metrics.get(metricName) match { - case Some(metric) => metric.set(v) + case Some(metric) if CometMetricNode.gaugeMetricNames.contains(metricName) => + metric.set(math.max(metric.value, v)) + case Some(metric) => + val last = lastReported.getOrElse(metricName, 0L) + metric.add(math.max(0L, v - last)) + lastReported.update(metricName, math.max(last, v)) case None => // no-op logDebug(s"Non-existing metric: $metricName. Ignored") @@ -239,6 +264,11 @@ object CometMetricNode { private val aggregateMetricNames = Set("spill_count", "spilled_bytes", "spilled_rows", "peak_mem_used") + // High-water marks that native code reports, which [[CometMetricNode.set]] keeps the maximum + // of rather than adding up. A gauge missing here would be summed across the native plans that + // one task runs. + private val gaugeMetricNames = Set("peak_mem_used", "build_mem_used") + private type SeenMetricSet = IdentityHashMap[SQLMetric, java.lang.Boolean] private case class SeenMetrics( diff --git a/spark/src/test/scala/org/apache/spark/CometExecIteratorLifecycleSuite.scala b/spark/src/test/scala/org/apache/spark/CometExecIteratorLifecycleSuite.scala index 43e5bb61759..202f6f17272 100644 --- a/spark/src/test/scala/org/apache/spark/CometExecIteratorLifecycleSuite.scala +++ b/spark/src/test/scala/org/apache/spark/CometExecIteratorLifecycleSuite.scala @@ -184,6 +184,8 @@ class CometExecIteratorLifecycleSuite extends CometTestBase { withTaskContext(4400000L) { val failMetrics = new AtomicBoolean(false) class ThrowingMetricNode extends CometMetricNode(Map.empty, Nil) { + // The iterator hands native a copy of its tree; keep this node so the failure fires. + override def newInstance(): CometMetricNode = this override def set_all_from_bytes(bytes: Array[Byte]): Unit = { if (failMetrics.get()) { throw new IllegalStateException("injected metrics update failure") diff --git a/spark/src/test/scala/org/apache/spark/sql/comet/CometTaskMetricsSuite.scala b/spark/src/test/scala/org/apache/spark/sql/comet/CometTaskMetricsSuite.scala index 153a7b422f9..783abb914fa 100644 --- a/spark/src/test/scala/org/apache/spark/sql/comet/CometTaskMetricsSuite.scala +++ b/spark/src/test/scala/org/apache/spark/sql/comet/CometTaskMetricsSuite.scala @@ -44,7 +44,7 @@ import org.apache.spark.unsafe.Platform import org.apache.comet.CometConf import org.apache.comet.CometSparkSessionExtensions.{isSpark40Plus, isSpark41Plus} -import org.apache.comet.serde.OperatorOuterClass +import org.apache.comet.serde.{Metric, OperatorOuterClass} class CometTaskMetricsSuite extends CometTestBase with AdaptiveSparkPlanHelper { @@ -127,6 +127,58 @@ class CometTaskMetricsSuite extends CometTestBase with AdaptiveSparkPlanHelper { } } + test("native metric updates from plan instances run in one task add up") { + val bytes = new SQLMetric("size", -1L) + val rows = new SQLMetric("sum") + val peakMemory = new SQLMetric("size", -1L) + val childRows = new SQLMetric("sum") + // Every instance of the tree updates these accumulators, like the parent partitions a + // coalesce computes one after another in a single task. + val tree = CometMetricNode( + Map("bytes_scanned" -> bytes, "output_rows" -> rows, "peak_mem_used" -> peakMemory), + Seq(CometMetricNode(Map("output_rows" -> childRows)))) + + // Native code reports absolute values for its own plan, several times while it runs. + def report(instance: CometMetricNode, b: Long, r: Long, peak: Long, child: Long): Unit = + instance.set_all_from_bytes( + nativeMetricUpdate( + Map("bytes_scanned" -> b, "output_rows" -> r, "peak_mem_used" -> peak), + nativeMetricUpdate(Map("output_rows" -> child))).toByteArray) + def values: (Long, Long, Long, Long) = + (bytes.value, rows.value, peakMemory.value, childRows.value) + + val first = tree.newInstance() + report(first, 100L, 10L, 5L, 7L) + report(first, 250L, 20L, 8L, 9L) + assert(values == ((250L, 20L, 8L, 9L))) + + // Counters add to what the first instance reported, and peak memory keeps the maximum. + val second = tree.newInstance() + report(second, 100L, 10L, 3L, 4L) + report(second, 300L, 40L, 3L, 6L) + assert(values == ((550L, 60L, 8L, 15L))) + } + + test("native metric updates keep a first zero and count a dip and recovery once") { + val bytes = new SQLMetric("size", -1L) + val rows = new SQLMetric("sum") + val instance = + CometMetricNode(Map("bytes_scanned" -> bytes, "output_rows" -> rows)).newInstance() + def report(b: Long, r: Long): Unit = + instance.set_all_from_bytes( + nativeMetricUpdate(Map("bytes_scanned" -> b, "output_rows" -> r)).toByteArray) + + // A size metric starts at -1 (unset); a reported zero marks it as set. + report(0L, 0L) + assert(bytes.value == 0L && !bytes.isZero) + + report(100L, 50L) + report(80L, 40L) + report(120L, 60L) + assert(bytes.value == 120L) + assert(rows.value == 60L) + } + // The unit tests below build their contexts with TaskContext.empty(), which always carries // task attempt id 0, so they share one per-task registry entry. Each test must reach // markTaskCompleted so the cleanup listener discards that entry before the next test claims. @@ -1287,6 +1339,149 @@ class CometTaskMetricsSuite extends CometTestBase with AdaptiveSparkPlanHelper { } } + test("coalesced native scan reports rows and bytes from every partition it reads") { + val totalRows = 10000 + withTempPath { dir => + spark + .createDataFrame((0 until totalRows).map(i => (i, s"elem_$i"))) + .repartition(4) + .write + .parquet(dir.getAbsolutePath) + spark.read.parquet(dir.getAbsolutePath).createOrReplaceTempView("coalesce_scan_tbl") + + // One scan partition per file: the plain query runs four native scan plans in four + // tasks, and the coalesced query runs the same four plans one after another in one task. + val fileSize = dir.listFiles().filter(_.getName.endsWith(".parquet")).map(_.length()).max + val onePartitionPerFile = Seq( + SQLConf.FILES_MAX_PARTITION_BYTES.key -> fileSize.toString, + SQLConf.FILES_OPEN_COST_IN_BYTES.key -> "0", + SQLConf.FILES_MIN_PARTITION_NUM.key -> "1") + val plainQuery = "SELECT * FROM coalesce_scan_tbl" + val coalescedQuery = "SELECT /*+ COALESCE(1) */ * FROM coalesce_scan_tbl" + + // Returns the task input bytes and records, the number of tasks, the scan's SQL + // output_rows and bytes_scanned, and the executed plan. + def run(query: String, confs: Seq[(String, String)]) = { + val store = spark.sparkContext.statusStore + val stagesBefore = store.stageList(null).map(_.stageId).toSet + val (bytesRead, recordsRead, plan) = + collectInputMetrics(query, (CometConf.COMET_ENABLED.key -> "true") +: confs: _*) + val numTasks = + store.stageList(null).filterNot(s => stagesBefore.contains(s.stageId)).map(_.numTasks) + val scan = collectFirst(plan) { case s: CometNativeScanExec => s } + assert(scan.isDefined, s"Expected CometNativeScanExec in plan:\n${plan.treeString}") + val scanRows = scan.get.metrics("output_rows").value + val scanBytes = scan.get.metrics("bytes_scanned").value + (bytesRead, recordsRead, numTasks.sum, scanRows, scanBytes, plan) + } + + Seq("-1", CometConf.COMET_METRICS_UPDATE_INTERVAL.defaultValueString).foreach { interval => + val confs = + onePartitionPerFile :+ (CometConf.COMET_METRICS_UPDATE_INTERVAL.key -> interval) + + val (plainBytesRead, plainRecordsRead, plainTasks, plainRows, plainScanBytes, plainPlan) = + run(plainQuery, confs) + val plainScan = collectFirst(plainPlan) { case s: CometNativeScanExec => s }.get + assert( + plainScan.outputPartitioning.numPartitions == 4, + s"Expected one scan partition per file:\n${plainPlan.treeString}") + assert(plainTasks == 4, s"Expected 4 tasks at interval $interval, got $plainTasks") + assert(plainRows == totalRows, s"output_rows at interval $interval: $plainRows") + assert(plainScanBytes > 0, s"bytes_scanned at interval $interval: $plainScanBytes") + assert(plainRecordsRead == totalRows, s"recordsRead at interval $interval") + + val (bytesRead, recordsRead, tasks, rows, scanBytes, plan) = run(coalescedQuery, confs) + assert( + find(plan)(_.isInstanceOf[CometCoalesceExec]).isDefined, + s"Expected CometCoalesceExec in plan:\n${plan.treeString}") + assert(tasks == 1, s"Expected 1 task at interval $interval, got $tasks") + assert(rows == totalRows, s"output_rows at interval $interval: $rows") + assert( + scanBytes == plainScanBytes, + s"bytes_scanned at interval $interval: coalesced=$scanBytes, plain=$plainScanBytes") + assert(recordsRead == totalRows, s"recordsRead at interval $interval: $recordsRead") + assert( + bytesRead == plainBytesRead, + s"bytesRead at interval $interval: coalesced=$bytesRead, plain=$plainBytesRead") + } + + val (_, sparkRecords, _) = collectInputMetrics( + coalescedQuery, + (CometConf.COMET_ENABLED.key -> "false") +: onePartitionPerFile: _*) + assert(sparkRecords == totalRows, s"Spark recordsRead: $sparkRecords") + } + } + + test("native blocks that feed each other through a JVM input count each operator's rows once") { + val totalRows = 10000 + withTempPath { dir => + spark + .createDataFrame((0 until totalRows).map(i => (i, s"elem_$i"))) + .repartition(1) + .write + .parquet(dir.getAbsolutePath) + spark.read.parquet(dir.getAbsolutePath).createOrReplaceTempView("nested_blocks_tbl") + val evenRows = totalRows / 2 + + def outputRows(plans: Seq[SparkPlan]): Seq[Long] = plans.map(_.metrics("output_rows").value) + // Rows the scan read, whether it returned them or pruned them with a pushed-down filter. + def scanRows(plan: SparkPlan): Seq[Long] = collect(plan) { case s: CometNativeScanExec => + s.metrics("output_rows").value + s.metrics.get("pushdown_rows_pruned").fold(0L)(_.value) + } + def filters(plan: SparkPlan): Seq[SparkPlan] = + collect(plan) { case f: CometFilterExec => f } + def aggregates(plan: SparkPlan): Seq[SparkPlan] = + collect(plan) { case a: CometHashAggregateExec => a } + + // A native project above a coalesce reads the scan and filter block through a JVM input. + val coalesced = spark + .table("nested_blocks_tbl") + .filter($"_1" % 2 === 0) + .coalesce(1) + .selectExpr("_1 + 1 AS x") + assert(coalesced.collect().length == evenRows) + val coalescedPlan = stripAQEPlan(coalesced.queryExecution.executedPlan) + val project = collectFirst(coalescedPlan) { + case p: CometProjectExec if p.child.isInstanceOf[CometCoalesceExec] => p + } + assert( + project.isDefined, + s"Expected a project over a coalesce:\n${coalescedPlan.treeString}") + assert(scanRows(coalescedPlan) == Seq(totalRows.toLong), coalescedPlan.treeString) + assert(outputRows(filters(coalescedPlan)) == Seq(evenRows.toLong), coalescedPlan.treeString) + assert(outputRows(project.toSeq) == Seq(evenRows.toLong), coalescedPlan.treeString) + + // A partial aggregate above a union reads two scan and filter blocks through a JVM input, + // one per task. + val unioned = sql( + "SELECT count(*), sum(_1) FROM (" + + "SELECT _1 FROM nested_blocks_tbl WHERE _1 % 2 = 0 UNION ALL " + + "SELECT _1 FROM nested_blocks_tbl WHERE _1 % 2 = 1)") + assert(unioned.collect().head.getLong(0) == totalRows) + val unionedPlan = stripAQEPlan(unioned.queryExecution.executedPlan) + assert( + find(unionedPlan)(_.isInstanceOf[CometUnionExec]).isDefined, + s"Expected CometUnionExec in plan:\n${unionedPlan.treeString}") + assert(scanRows(unionedPlan) == Seq(totalRows.toLong, totalRows.toLong)) + assert(outputRows(filters(unionedPlan)) == Seq(evenRows.toLong, evenRows.toLong)) + // One partial row from each of the union's two partitions, then the final row. + assert(outputRows(aggregates(unionedPlan)).sorted == Seq(1L, 2L), unionedPlan.treeString) + + // The final aggregate reads the partial aggregate's output through a shuffle. + val grouped = sql( + "SELECT _1 % 10 AS k, count(*) FROM nested_blocks_tbl WHERE _1 % 2 = 0 GROUP BY _1 % 10") + assert(grouped.collect().length == 5) + val groupedPlan = stripAQEPlan(grouped.queryExecution.executedPlan) + assert( + find(groupedPlan)(_.isInstanceOf[CometShuffleExchangeExec]).isDefined, + s"Expected CometShuffleExchangeExec in plan:\n${groupedPlan.treeString}") + assert(scanRows(groupedPlan) == Seq(totalRows.toLong), groupedPlan.treeString) + assert(outputRows(filters(groupedPlan)) == Seq(evenRows.toLong), groupedPlan.treeString) + // A single map task, so the partial and final aggregates both produce the five groups. + assert(outputRows(aggregates(groupedPlan)) == Seq(5L, 5L), groupedPlan.treeString) + } + } + test("task input metrics keep rows read from a cached JVM input in the same task") { withTempPath { dir => spark @@ -1416,6 +1611,16 @@ class CometTaskMetricsSuite extends CometTestBase with AdaptiveSparkPlanHelper { } } + /** Builds one metrics update in the form native code sends through JNI. */ + private def nativeMetricUpdate( + metrics: Map[String, Long], + children: Metric.NativeMetricNode*): Metric.NativeMetricNode = { + val builder = Metric.NativeMetricNode.newBuilder() + metrics.foreach { case (name, value) => builder.putMetrics(name, value) } + children.foreach(child => builder.addChildren(child)) + builder.build() + } + /** The task input bytes a query reports on its own with Comet enabled. */ private def cometBytesRead(query: String, confs: (String, String)*): Long = collectInputMetrics(query, (CometConf.COMET_ENABLED.key -> "true") +: confs: _*)._1