From 6f7a970eb61fbbd872404a6fb32d3061e6498593 Mon Sep 17 00:00:00 2001 From: msaf Date: Sat, 10 Oct 2026 11:58:01 +0100 Subject: [PATCH 1/3] feat: native MIN/MAX over strings and SortAggregate as a native hash aggregate Spark plans MIN or MAX over a string as a SortAggregate, since the aggregation buffer is not mutable, and Comet did not convert SortAggregate, so such an aggregate and often the joins around it ran in Spark. MIN and MAX now accept StringType with the default UTF8_BINARY collation. The native min/max compare UTF-8 bytes, which is Spark's binary order; collated strings stay in Spark. A SortAggregateExec is converted to a CometHashAggregateExec built from an ObjectHashAggregateExec with the same fields. Spark runs that operator over any input order, so an aggregate reverted to Spark later (cost-based engine choice, unsafe partial aggregates) stays correct, and the cost-based choice prices it as an object hash aggregate. A sort by exactly the grouping keys directly below the aggregate is dropped; without one only order-insensitive aggregates (MIN, MAX, COUNT, SUM, AVG) are converted. A native sort above the aggregate restores the output ordering a consumer such as a sort-merge join may rely on, and is dropped again when a shuffle reads it directly. Aggregates Comet does not support keep the SortAggregate. spark.comet.exec.sortAggregate.enabled (default true) turns the conversion off. Co-Authored-By: Claude Opus 5.5 (1M context) --- .../scala/org/apache/comet/CometConf.scala | 7 + .../apache/comet/rules/CometExecRule.scala | 104 +++++++++++++- .../org/apache/comet/serde/aggregates.scala | 10 +- .../expressions/aggregate/first_last.sql | 4 +- .../expressions/aggregate/min_max.sql | 2 +- .../comet/exec/CometAggregateSuite.scala | 127 +++++++++++++++++- .../rules/CostBasedEngineChoiceSuite.scala | 24 +++- .../rules/WideRowSortFallbackSuite.scala | 1 + 8 files changed, 267 insertions(+), 12 deletions(-) diff --git a/spark/src/main/scala/org/apache/comet/CometConf.scala b/spark/src/main/scala/org/apache/comet/CometConf.scala index 03bd2701754..4c8ca3b402c 100644 --- a/spark/src/main/scala/org/apache/comet/CometConf.scala +++ b/spark/src/main/scala/org/apache/comet/CometConf.scala @@ -237,6 +237,13 @@ object CometConf extends ShimCometConf { createExecEnabledConfig("sortMergeJoin", defaultValue = true) val COMET_EXEC_AGGREGATE_ENABLED: ConfigEntry[Boolean] = createExecEnabledConfig("aggregate", defaultValue = true) + val COMET_EXEC_SORT_AGGREGATE_ENABLED: ConfigEntry[Boolean] = + createExecEnabledConfig( + "sortAggregate", + defaultValue = true, + notes = Some( + "When enabled, a SortAggregate that Comet can run is converted to a native hash " + + "aggregate, with a native sort restoring its output ordering")) val COMET_EXEC_COLLECT_LIMIT_ENABLED: ConfigEntry[Boolean] = createExecEnabledConfig("collectLimit", defaultValue = true) val COMET_EXEC_COALESCE_ENABLED: ConfigEntry[Boolean] = diff --git a/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala b/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala index 65b85540e41..2cc826e0eef 100644 --- a/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala +++ b/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala @@ -23,7 +23,7 @@ import scala.collection.mutable.ListBuffer import org.apache.spark.sql.SparkSession import org.apache.spark.sql.catalyst.expressions.{Divide, DoubleLiteral, EqualNullSafe, EqualTo, Expression, FloatLiteral, GreaterThan, GreaterThanOrEqual, KnownFloatingPointNormalized, LessThan, LessThanOrEqual, NamedExpression, Remainder} -import org.apache.spark.sql.catalyst.expressions.aggregate.{AggregateMode, Final, Partial, PartialMerge} +import org.apache.spark.sql.catalyst.expressions.aggregate.{AggregateFunction, AggregateMode, Average, Count, Final, Max, Min, Partial, PartialMerge, Sum} import org.apache.spark.sql.catalyst.optimizer.NormalizeNaNAndZero import org.apache.spark.sql.catalyst.rules.Rule import org.apache.spark.sql.catalyst.trees.TreeNodeTag @@ -35,7 +35,7 @@ import org.apache.spark.sql.comet.shims.ShimCometEmptyRelation import org.apache.spark.sql.comet.util.Utils import org.apache.spark.sql.execution._ import org.apache.spark.sql.execution.adaptive.{AdaptiveSparkPlanExec, AQEShuffleReadExec, BroadcastQueryStageExec, LogicalQueryStage, QueryStageExec, ShuffleQueryStageExec} -import org.apache.spark.sql.execution.aggregate.{BaseAggregateExec, HashAggregateExec, ObjectHashAggregateExec} +import org.apache.spark.sql.execution.aggregate.{BaseAggregateExec, HashAggregateExec, ObjectHashAggregateExec, SortAggregateExec} import org.apache.spark.sql.execution.columnar.InMemoryTableScanExec import org.apache.spark.sql.execution.command.{DataWritingCommandExec, ExecutedCommandExec} import org.apache.spark.sql.execution.datasources.{InsertIntoHadoopFsRelationCommand, WriteFilesExec} @@ -150,6 +150,14 @@ object CometExecRule { */ val ENGINE_CHOICE_SPARK_TAG: TreeNodeTag[Unit] = TreeNodeTag[Unit]("comet.engineChoiceSpark") + /** + * Tag set on the native sort that restores the output ordering of a `SortAggregateExec` + * converted to a native hash aggregate. A shuffle directly above it discards the ordering, so + * the sort is dropped there. + */ + val SORT_AGGREGATE_ORDER_TAG: TreeNodeTag[Unit] = + TreeNodeTag[Unit]("comet.sortAggregateOrder") + /** * Serializes the native plan of each block of adjacent native operators into its topmost * operator. Blocks that already hold a serialized plan are left as they are, so this can run @@ -587,6 +595,9 @@ case class CometExecRule(session: SparkSession, wholePlan: Boolean = false) convertToComet(s, CometShuffleExchangeExec) .getOrElse(preserveSparkAggregateBuffers(s)) + case agg: SortAggregateExec if CometConf.COMET_EXEC_SORT_AGGREGATE_ENABLED.get(conf) => + convertSortAggregate(agg).getOrElse(agg) + case op => // if all children are native (or if this is a leaf node) then see if there is a // registered handler for creating a fully native plan @@ -644,7 +655,7 @@ case class CometExecRule(session: SparkSession, wholePlan: Boolean = false) } plan.transformUp { case op => - val converted = convertNode(op) + val converted = convertNode(dropSortAggregateOrderUnderShuffle(op)) // Replace SubqueryBroadcastExec with CometSubqueryBroadcastExec in DPP expressions // when the broadcast child has a Comet plan underneath. This enables exchange reuse // between the DPP subquery and the join's CometBroadcastExchangeExec because both @@ -928,6 +939,93 @@ case class CometExecRule(session: SparkSession, wholePlan: Boolean = false) } } + private def dropSortAggregateOrderUnderShuffle(op: SparkPlan): SparkPlan = op match { + case shuffle: ShuffleExchangeLike => + shuffle.child match { + case sort: CometSortExec + if sort.getTagValue(CometExecRule.SORT_AGGREGATE_ORDER_TAG).isDefined => + shuffle.withNewChildren(Seq(sort.child)) + case _ => shuffle + } + case other => other + } + + private def orderInsensitive(fn: AggregateFunction): Boolean = fn match { + case _: Min | _: Max | _: Count | _: Sum | _: Average => true + case _ => false + } + + /** + * A `SortAggregateExec` as a native hash aggregate. Spark plans one when an aggregation buffer + * is not mutable, such as MIN or MAX over strings, and runs it over input sorted by the + * grouping keys. The hash aggregate starts from an `ObjectHashAggregateExec` with the same + * fields, which Spark runs correctly over any input order, so any operator reverted to Spark + * later stays correct. The sort by the grouping keys directly below the aggregate is dropped, + * as rows within a group were in no particular order anyway. Without such a sort the input + * order could carry meaning, so only aggregates that do not depend on it are converted. A + * native sort above restores the ordering consumers may rely on. + */ + private def convertSortAggregate(agg: SortAggregateExec): Option[SparkPlan] = { + val required = agg.requiredChildOrdering.headOption.getOrElse(Nil) + val sortedByKeys = agg.child match { + case sort: CometSortExec => + sort.sortOrder.length == required.length && + sort.sortOrder.zip(required).forall { case (a, b) => + a.semanticEquals(b) + } && + (sort.originalPlan match { + case s: SortExec => !s.global + case _ => false + }) + case _ => false + } + val input = if (sortedByKeys) agg.child.children.head else agg.child + if (!sortedByKeys && + !agg.aggregateExpressions.forall(e => orderInsensitive(e.aggregateFunction))) { + withFallbackReason( + agg, + "SortAggregate over input not sorted by its grouping keys alone may depend on the " + + "input order") + return None + } + if (!input.isInstanceOf[CometNativeExec]) { + return None + } + val hashAgg = ObjectHashAggregateExec( + agg.requiredChildDistributionExpressions, + agg.isStreaming, + agg.numShufflePartitions, + agg.groupingExpressions, + agg.aggregateExpressions, + agg.aggregateAttributes, + agg.initialInputBufferOffset, + agg.resultExpressions, + input) + hashAgg.copyTagsFrom(agg) + val converted = tryConvertToComet(hashAgg, CometObjectHashAggregateExec).flatMap { native => + val ordering = agg.outputOrdering + if (ordering.isEmpty) { + Some(native) + } else { + val sort = SortExec(ordering, global = false, native) + agg.logicalLink.foreach(sort.setLogicalLink) + tryConvertToComet(sort, CometSortExec).map { nativeSort => + nativeSort.setTagValue(CometExecRule.SORT_AGGREGATE_ORDER_TAG, ()) + nativeSort + } + } + } + if (converted.isEmpty) { + hashAgg + .getTagValue(CometExplainInfo.FALLBACK_REASONS) + .foreach(reasons => withFallbackReasons(agg, reasons)) + if (!hasFallbackReason(agg)) { + withFallbackReason(agg, "SortAggregate could not be converted to a native hash aggregate") + } + } + converted + } + /** Convert a Spark plan to a Comet plan using the specified serde handler */ private def convertToComet(op: SparkPlan, handler: CometOperatorSerde[_]): Option[SparkPlan] = { val converted = tryConvertToComet(op, handler) diff --git a/spark/src/main/scala/org/apache/comet/serde/aggregates.scala b/spark/src/main/scala/org/apache/comet/serde/aggregates.scala index ad53cceb31b..bc45c605025 100644 --- a/spark/src/main/scala/org/apache/comet/serde/aggregates.scala +++ b/spark/src/main/scala/org/apache/comet/serde/aggregates.scala @@ -1395,7 +1395,7 @@ object CometMode extends CometAggregateExpressionSerde[Mode] with CometTypeShim } } -object AggSerde { +object AggSerde extends CometTypeShim { import org.apache.spark.sql.types._ def minMaxDataTypeSupported(dt: DataType): Boolean = { @@ -1436,7 +1436,13 @@ object AggSerde { /** Shared support level for `Min` / `Max` based on the result data type. */ def minMaxSupportLevel(dt: DataType): SupportLevel = { - if (!minMaxDataTypeSupported(dt)) { + if (dt.isInstanceOf[StringType]) { + if (isStringCollationType(dt)) { + Unsupported(Some(s"Unsupported collated string type: $dt")) + } else { + Compatible() + } + } else if (!minMaxDataTypeSupported(dt)) { Unsupported(Some(s"Unsupported data type: $dt")) } else if ((dt == FloatType || dt == DoubleType) && COMET_EXEC_STRICT_FLOATING_POINT.get()) { diff --git a/spark/src/test/resources/sql-tests/expressions/aggregate/first_last.sql b/spark/src/test/resources/sql-tests/expressions/aggregate/first_last.sql index 8aedc4f5f17..7fd5dcde731 100644 --- a/spark/src/test/resources/sql-tests/expressions/aggregate/first_last.sql +++ b/spark/src/test/resources/sql-tests/expressions/aggregate/first_last.sql @@ -183,7 +183,7 @@ SELECT first(val) IGNORE NULLS FROM test_single_null -- first IGNORE NULLS: multiple data types -- ============================================================ -query expect_fallback(SortAggregate is not supported) +query SELECT grp, first(i_val) IGNORE NULLS, first(l_val) IGNORE NULLS, @@ -327,7 +327,7 @@ SELECT last(val) IGNORE NULLS FROM test_single_null -- last IGNORE NULLS: multiple data types -- ============================================================ -query expect_fallback(SortAggregate is not supported) +query SELECT grp, last(i_val) IGNORE NULLS, last(l_val) IGNORE NULLS, diff --git a/spark/src/test/resources/sql-tests/expressions/aggregate/min_max.sql b/spark/src/test/resources/sql-tests/expressions/aggregate/min_max.sql index dbec88e2f66..9fcde06186e 100644 --- a/spark/src/test/resources/sql-tests/expressions/aggregate/min_max.sql +++ b/spark/src/test/resources/sql-tests/expressions/aggregate/min_max.sql @@ -21,7 +21,7 @@ CREATE TABLE test_min_max(i int, d double, s string, grp string) USING parquet statement INSERT INTO test_min_max VALUES (1, 1.5, 'b', 'x'), (3, 3.5, 'a', 'x'), (2, 2.5, 'c', 'y'), (NULL, NULL, NULL, 'y'), (-1, -1.5, 'z', 'x') -query expect_fallback(SortAggregate is not supported) +query SELECT min(i), max(i), min(d), max(d), min(s), max(s) FROM test_min_max query diff --git a/spark/src/test/scala/org/apache/comet/exec/CometAggregateSuite.scala b/spark/src/test/scala/org/apache/comet/exec/CometAggregateSuite.scala index e82831ec426..20cbd42042b 100644 --- a/spark/src/test/scala/org/apache/comet/exec/CometAggregateSuite.scala +++ b/spark/src/test/scala/org/apache/comet/exec/CometAggregateSuite.scala @@ -31,11 +31,11 @@ import org.apache.spark.sql.catalyst.expressions.Cast import org.apache.spark.sql.catalyst.expressions.aggregate.{Final, Partial, PartialMerge} import org.apache.spark.sql.catalyst.optimizer.EliminateSorts import org.apache.spark.sql.catalyst.plans.physical.{HashPartitioning, RangePartitioning} -import org.apache.spark.sql.comet.{CometFilterExec, CometHashAggregateExec, CometNativeExec, CometProjectExec} +import org.apache.spark.sql.comet.{CometFilterExec, CometHashAggregateExec, CometNativeExec, CometProjectExec, CometSortExec, CometSortMergeJoinExec} import org.apache.spark.sql.comet.execution.shuffle.CometShuffleExchangeExec -import org.apache.spark.sql.execution.SQLExecution +import org.apache.spark.sql.execution.{SparkPlan, SQLExecution} import org.apache.spark.sql.execution.adaptive.{AdaptiveSparkPlanExec, AdaptiveSparkPlanHelper, ShuffleQueryStageExec} -import org.apache.spark.sql.execution.aggregate.{BaseAggregateExec, ObjectHashAggregateExec} +import org.apache.spark.sql.execution.aggregate.{BaseAggregateExec, ObjectHashAggregateExec, SortAggregateExec} import org.apache.spark.sql.execution.exchange.{ReusedExchangeExec, ShuffleExchangeExec} import org.apache.spark.sql.functions.{avg, col, count_distinct, expr, sum} import org.apache.spark.sql.internal.SQLConf @@ -3448,4 +3448,125 @@ class CometAggregateSuite extends CometTestBase with AdaptiveSparkPlanHelper { } } + private val stringRows: Seq[(Integer, String)] = Seq( + (1, "b"), + (1, ""), + (1, null), + (1, "a"), + (2, null), + (2, null), + (3, "\u00e9t\u00e9"), + (3, "ete"), + (3, "\u65e5\u672c"), + (4, "\uff61"), + (4, "\ud83d\ude00"), + (4, "z"), + (5, ""), + (null, "x"), + (null, "\u00ff")) + + private def withStrings(f: => Unit): Unit = + withParquetTable(stringRows, "str_tbl")(f) + + private def sortAggregates(plan: SparkPlan): Seq[SortAggregateExec] = + collect(plan) { case a: SortAggregateExec => a } + + private def nativeAggregates(plan: SparkPlan): Seq[CometHashAggregateExec] = + collect(plan) { case a: CometHashAggregateExec => a } + + for (aqe <- Seq("false", "true")) { + test(s"min and max over strings run natively, grouped and ungrouped (AQE=$aqe)") { + withSQLConf(SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> aqe) { + withStrings { + for (query <- Seq( + "SELECT min(_2), max(_2), count(_2) FROM str_tbl", + "SELECT _1, min(_2), max(_2) FROM str_tbl GROUP BY _1", + "SELECT _1, max(_2) FROM str_tbl WHERE _1 = 2 GROUP BY _1", + "SELECT max(_2) FROM str_tbl WHERE _1 > 100", + "SELECT _1, min(_2), sum(_1) FROM str_tbl GROUP BY _1")) { + val (sparkPlan, cometPlan) = checkSparkAnswer(sql(query)) + assert(sortAggregates(sparkPlan).nonEmpty, s"$query:\n$sparkPlan") + assert(sortAggregates(cometPlan).isEmpty, s"$query:\n$cometPlan") + assert(nativeAggregates(cometPlan).nonEmpty, s"$query:\n$cometPlan") + } + } + } + } + } + + test("min and max over strings compare UTF-8 bytes as Spark does") { + withStrings { + checkSparkAnswerAndOperator(sql("SELECT min(_2), max(_2) FROM str_tbl WHERE _1 = 4")) + checkSparkAnswerAndOperator(sql("SELECT _1, min(_2), max(_2) FROM str_tbl GROUP BY _1")) + } + } + + test("a converted sort aggregate keeps the ordering a sort-merge join relies on") { + withSQLConf( + SQLConf.AUTO_BROADCASTJOIN_THRESHOLD.key -> "-1", + SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false") { + withStrings { + withParquetTable((0 until 50).map(i => (i % 6, s"v$i")), "str_right") { + val query = + """SELECT a._1, a.m, r._2 FROM (SELECT _1, max(_2) AS m FROM str_tbl GROUP BY _1) a + |JOIN str_right r ON a._1 = r._1""".stripMargin + val (_, cometPlan) = checkSparkAnswer(sql(query)) + assert(sortAggregates(cometPlan).isEmpty, s"plan:\n$cometPlan") + val join = collect(cometPlan) { case j: CometSortMergeJoinExec => j } + assert(join.size == 1, s"plan:\n$cometPlan") + val orderSorts = collect(cometPlan) { + case s: CometSortExec + if s.getTagValue(CometExecRule.SORT_AGGREGATE_ORDER_TAG).isDefined => + s + } + assert(orderSorts.nonEmpty, s"plan:\n$cometPlan") + } + } + } + } + + test("the ordering sort of a converted partial sort aggregate is dropped under its shuffle") { + withSQLConf(SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false") { + withStrings { + val (_, cometPlan) = checkSparkAnswer(sql("SELECT _1, max(_2) FROM str_tbl GROUP BY _1")) + val shuffles = collect(cometPlan) { case s: CometShuffleExchangeExec => s } + assert(shuffles.nonEmpty, s"plan:\n$cometPlan") + assert(shuffles.forall(s => !s.child.isInstanceOf[CometSortExec]), s"plan:\n$cometPlan") + assert(nativeAggregates(cometPlan).size == 2, s"plan:\n$cometPlan") + } + } + } + + test("a sort aggregate Comet cannot run natively stays in Spark") { + withStrings { + val (_, cometPlan) = + checkSparkAnswer(sql("SELECT _1, max_by(_2, _2) FROM str_tbl GROUP BY _1")) + assert(sortAggregates(cometPlan).nonEmpty, s"plan:\n$cometPlan") + } + } + + test("an order-sensitive sort aggregate over input sorted beyond its keys stays in Spark") { + withStrings { + val (_, cometPlan) = checkSparkAnswer( + sql("SELECT _1, first(_2) FROM (SELECT * FROM str_tbl DISTRIBUTE BY _1 SORT BY _1, _2) " + + "GROUP BY _1")) + assert(sortAggregates(cometPlan).nonEmpty, s"plan:\n$cometPlan") + } + } + + test("a converted sort aggregate reverted by the cost-based choice runs as an object hash") { + withSQLConf( + CometConf.COMET_EXEC_COST_BASED_ENGINES_ENABLED.key -> "true", + CometConf.COMET_EXEC_COST_BASED_ENGINES_COST_TABLE.key -> "agg.flat.comet=1000000,0,0", + SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false") { + withStrings { + val (_, cometPlan) = + checkSparkAnswer(sql("SELECT _1, min(_2), max(_2) FROM str_tbl GROUP BY _1")) + assert(sortAggregates(cometPlan).isEmpty, s"plan:\n$cometPlan") + assert( + collect(cometPlan) { case a: ObjectHashAggregateExec => a }.nonEmpty, + s"plan:\n$cometPlan") + } + } + } } diff --git a/spark/src/test/scala/org/apache/comet/rules/CostBasedEngineChoiceSuite.scala b/spark/src/test/scala/org/apache/comet/rules/CostBasedEngineChoiceSuite.scala index 25ec4c2e0eb..2e2ff2673e6 100644 --- a/spark/src/test/scala/org/apache/comet/rules/CostBasedEngineChoiceSuite.scala +++ b/spark/src/test/scala/org/apache/comet/rules/CostBasedEngineChoiceSuite.scala @@ -19,6 +19,7 @@ package org.apache.comet.rules +import org.apache.spark.SparkConf import org.apache.spark.sql.{CometTestBase, DataFrame} import org.apache.spark.sql.catalyst.expressions.{Attribute, AttributeReference} import org.apache.spark.sql.catalyst.expressions.aggregate.Partial @@ -27,7 +28,7 @@ import org.apache.spark.sql.comet._ import org.apache.spark.sql.comet.execution.shuffle.CometShuffleExchangeExec import org.apache.spark.sql.execution.{ExpandExec, ProjectExec, SortExec, SparkPlan} import org.apache.spark.sql.execution.adaptive.{AdaptiveSparkPlanExec, QueryStageExec} -import org.apache.spark.sql.execution.aggregate.{HashAggregateExec, SortAggregateExec} +import org.apache.spark.sql.execution.aggregate.{HashAggregateExec, ObjectHashAggregateExec, SortAggregateExec} import org.apache.spark.sql.execution.exchange.ReusedExchangeExec import org.apache.spark.sql.execution.joins.{BroadcastHashJoinExec, SortMergeJoinExec} import org.apache.spark.sql.functions.{col, sum} @@ -43,6 +44,9 @@ import org.apache.comet.rules.EngineCostTable.CostClass._ class CostBasedEngineChoiceSuite extends CometTestBase { + override protected def sparkConf: SparkConf = + super.sparkConf.set(CometConf.COMET_EXEC_SORT_AGGREGATE_ENABLED.key, "false") + private val flag = CometConf.COMET_EXEC_COST_BASED_ENGINES_ENABLED.key private val costTable = CometConf.COMET_EXEC_COST_BASED_ENGINES_COST_TABLE.key private val expensiveNativeSort = costTable -> "sort.flat.comet=1000,0,0" @@ -1669,4 +1673,22 @@ class CostBasedEngineChoiceSuite extends CometTestBase { } } } + + test("a sort aggregate converted to a native hash aggregate is priced as an object hash") { + withTables { + withSQLConf( + CometConf.COMET_EXEC_SORT_AGGREGATE_ENABLED.key -> "true", + SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false", + flag -> "true") { + val plan = run("SELECT k, max(s) FROM t GROUP BY k") + assert(count(plan) { case a: SortAggregateExec => a } == 0, s"plan:\n$plan") + val aggs = nodes(plan).collect { case a: CometHashAggregateExec => a } + assert(aggs.size == 2, s"plan:\n$plan") + aggs.foreach { agg => + assert(agg.originalPlan.isInstanceOf[ObjectHashAggregateExec], s"plan:\n$plan") + assert(model.costClasses(agg) == Seq(Agg, AggObjectHash), s"plan:\n$plan") + } + } + } + } } diff --git a/spark/src/test/scala/org/apache/comet/rules/WideRowSortFallbackSuite.scala b/spark/src/test/scala/org/apache/comet/rules/WideRowSortFallbackSuite.scala index 37253b2485c..5de75ae186b 100644 --- a/spark/src/test/scala/org/apache/comet/rules/WideRowSortFallbackSuite.scala +++ b/spark/src/test/scala/org/apache/comet/rules/WideRowSortFallbackSuite.scala @@ -43,6 +43,7 @@ class WideRowSortFallbackSuite extends CometTestBase { override protected def sparkConf: SparkConf = super.sparkConf + .set(CometConf.COMET_EXEC_SORT_AGGREGATE_ENABLED.key, "false") .set(shuffleMinLeaves, "0") .set(CometConf.COMET_EXEC_COST_BASED_ENGINES_ENABLED.key, "false") From 4d7eec2d2965545fac326680da2f41b05ff7a99a Mon Sep 17 00:00:00 2001 From: msaf Date: Sat, 10 Oct 2026 13:17:25 +0100 Subject: [PATCH 2/3] fix: sort a converted SortAggregate's output only for a consumer that needs it The native sort above a SortAggregate converted to a hash aggregate was added unconditionally and dropped only under a shuffle. It is now added the way EnsureRequirements would: when a consumer is visited, before it is converted, each child whose required ordering is no longer met and that leads through unary operators to a converted SortAggregate gets a sort on that edge. A partial aggregate under a shuffle, a final aggregate read by a hash aggregate, a project or a write, and any consumer without an ordering requirement get none. Co-Authored-By: Claude Opus 5.5 (1M context) --- .../apache/comet/rules/CometExecRule.scala | 71 ++++++++++++------- .../comet/exec/CometAggregateSuite.scala | 10 ++- 2 files changed, 53 insertions(+), 28 deletions(-) diff --git a/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala b/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala index 2cc826e0eef..c461f43abfb 100644 --- a/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala +++ b/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala @@ -22,7 +22,7 @@ package org.apache.comet.rules import scala.collection.mutable.ListBuffer import org.apache.spark.sql.SparkSession -import org.apache.spark.sql.catalyst.expressions.{Divide, DoubleLiteral, EqualNullSafe, EqualTo, Expression, FloatLiteral, GreaterThan, GreaterThanOrEqual, KnownFloatingPointNormalized, LessThan, LessThanOrEqual, NamedExpression, Remainder} +import org.apache.spark.sql.catalyst.expressions.{Divide, DoubleLiteral, EqualNullSafe, EqualTo, Expression, FloatLiteral, GreaterThan, GreaterThanOrEqual, KnownFloatingPointNormalized, LessThan, LessThanOrEqual, NamedExpression, Remainder, SortOrder} import org.apache.spark.sql.catalyst.expressions.aggregate.{AggregateFunction, AggregateMode, Average, Count, Final, Max, Min, Partial, PartialMerge, Sum} import org.apache.spark.sql.catalyst.optimizer.NormalizeNaNAndZero import org.apache.spark.sql.catalyst.rules.Rule @@ -150,10 +150,13 @@ object CometExecRule { */ val ENGINE_CHOICE_SPARK_TAG: TreeNodeTag[Unit] = TreeNodeTag[Unit]("comet.engineChoiceSpark") + /** Tag set on a native hash aggregate converted from a `SortAggregateExec`. */ + val CONVERTED_SORT_AGGREGATE_TAG: TreeNodeTag[Unit] = + TreeNodeTag[Unit]("comet.convertedSortAggregate") + /** - * Tag set on the native sort that restores the output ordering of a `SortAggregateExec` - * converted to a native hash aggregate. A shuffle directly above it discards the ordering, so - * the sort is dropped there. + * Tag set on the sort inserted above a converted `SortAggregateExec` for a consumer that + * requires the ordering the sort aggregate provided. */ val SORT_AGGREGATE_ORDER_TAG: TreeNodeTag[Unit] = TreeNodeTag[Unit]("comet.sortAggregateOrder") @@ -655,7 +658,7 @@ case class CometExecRule(session: SparkSession, wholePlan: Boolean = false) } plan.transformUp { case op => - val converted = convertNode(dropSortAggregateOrderUnderShuffle(op)) + val converted = convertNode(restoreSortAggregateOrdering(op)) // Replace SubqueryBroadcastExec with CometSubqueryBroadcastExec in DPP expressions // when the broadcast child has a Comet plan underneath. This enables exchange reuse // between the DPP subquery and the join's CometBroadcastExchangeExec because both @@ -939,15 +942,39 @@ case class CometExecRule(session: SparkSession, wholePlan: Boolean = false) } } - private def dropSortAggregateOrderUnderShuffle(op: SparkPlan): SparkPlan = op match { - case shuffle: ShuffleExchangeLike => - shuffle.child match { - case sort: CometSortExec - if sort.getTagValue(CometExecRule.SORT_AGGREGATE_ORDER_TAG).isDefined => - shuffle.withNewChildren(Seq(sort.child)) - case _ => shuffle + private def leadsToConvertedSortAggregate(plan: SparkPlan): Boolean = + plan.getTagValue(CometExecRule.CONVERTED_SORT_AGGREGATE_TAG).isDefined || + (plan.children.size == 1 && leadsToConvertedSortAggregate(plan.children.head)) + + /** + * Spark's EnsureRequirements left out the sort a consumer needs when a `SortAggregateExec` + * below it already provided the ordering. Once that aggregate is a native hash aggregate, the + * sort goes back on exactly the edges whose requirement is no longer met, before the consumer + * itself is converted. + */ + private def restoreSortAggregateOrdering(op: SparkPlan): SparkPlan = { + if (op.isInstanceOf[CometPlan]) return op + val required = op.requiredChildOrdering + if (required.length != op.children.length || required.forall(_.isEmpty)) return op + var changed = false + val children = op.children.zip(required).map { case (child, ordering) => + if (ordering.nonEmpty && leadsToConvertedSortAggregate(child) && + !SortOrder.orderingSatisfies(child.outputOrdering, ordering)) { + changed = true + val sort = SortExec(ordering, global = false, child) + child.logicalLink.foreach(sort.setLogicalLink) + val restored = if (child.isInstanceOf[CometNativeExec]) { + tryConvertToComet(sort, CometSortExec).getOrElse(sort) + } else { + sort + } + restored.setTagValue(CometExecRule.SORT_AGGREGATE_ORDER_TAG, ()) + restored + } else { + child } - case other => other + } + if (changed) op.withNewChildren(children) else op } private def orderInsensitive(fn: AggregateFunction): Boolean = fn match { @@ -963,7 +990,8 @@ case class CometExecRule(session: SparkSession, wholePlan: Boolean = false) * later stays correct. The sort by the grouping keys directly below the aggregate is dropped, * as rows within a group were in no particular order anyway. Without such a sort the input * order could carry meaning, so only aggregates that do not depend on it are converted. A - * native sort above restores the ordering consumers may rely on. + * consumer that relied on the ordering of the sort aggregate gets a sort back, see + * [[restoreSortAggregateOrdering]]. */ private def convertSortAggregate(agg: SortAggregateExec): Option[SparkPlan] = { val required = agg.requiredChildOrdering.headOption.getOrElse(Nil) @@ -1002,19 +1030,8 @@ case class CometExecRule(session: SparkSession, wholePlan: Boolean = false) agg.resultExpressions, input) hashAgg.copyTagsFrom(agg) - val converted = tryConvertToComet(hashAgg, CometObjectHashAggregateExec).flatMap { native => - val ordering = agg.outputOrdering - if (ordering.isEmpty) { - Some(native) - } else { - val sort = SortExec(ordering, global = false, native) - agg.logicalLink.foreach(sort.setLogicalLink) - tryConvertToComet(sort, CometSortExec).map { nativeSort => - nativeSort.setTagValue(CometExecRule.SORT_AGGREGATE_ORDER_TAG, ()) - nativeSort - } - } - } + val converted = tryConvertToComet(hashAgg, CometObjectHashAggregateExec) + converted.foreach(_.setTagValue(CometExecRule.CONVERTED_SORT_AGGREGATE_TAG, ())) if (converted.isEmpty) { hashAgg .getTagValue(CometExplainInfo.FALLBACK_REASONS) diff --git a/spark/src/test/scala/org/apache/comet/exec/CometAggregateSuite.scala b/spark/src/test/scala/org/apache/comet/exec/CometAggregateSuite.scala index 20cbd42042b..298d2a73eb5 100644 --- a/spark/src/test/scala/org/apache/comet/exec/CometAggregateSuite.scala +++ b/spark/src/test/scala/org/apache/comet/exec/CometAggregateSuite.scala @@ -3525,7 +3525,7 @@ class CometAggregateSuite extends CometTestBase with AdaptiveSparkPlanHelper { } } - test("the ordering sort of a converted partial sort aggregate is dropped under its shuffle") { + test("a converted sort aggregate gets no sort without a consumer that needs its ordering") { withSQLConf(SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false") { withStrings { val (_, cometPlan) = checkSparkAnswer(sql("SELECT _1, max(_2) FROM str_tbl GROUP BY _1")) @@ -3533,6 +3533,14 @@ class CometAggregateSuite extends CometTestBase with AdaptiveSparkPlanHelper { assert(shuffles.nonEmpty, s"plan:\n$cometPlan") assert(shuffles.forall(s => !s.child.isInstanceOf[CometSortExec]), s"plan:\n$cometPlan") assert(nativeAggregates(cometPlan).size == 2, s"plan:\n$cometPlan") + assert( + collect(cometPlan) { + case s: CometSortExec + if s.getTagValue(CometExecRule.SORT_AGGREGATE_ORDER_TAG).isDefined => + s + }.isEmpty, + s"plan:\n$cometPlan") + assert(collect(cometPlan) { case s: CometSortExec => s }.isEmpty, s"plan:\n$cometPlan") } } } From e930327f406b071e102b25272a7ce7b81a5dda8f Mon Sep 17 00:00:00 2001 From: msaf Date: Sat, 10 Oct 2026 14:57:06 +0100 Subject: [PATCH 3/3] fix: keep a SortAggregate with an order-sensitive aggregate in Spark Spark's sort is stable, so a SortAggregate over input sorted by its grouping keys sees each group's rows in their input order, and FIRST, LAST, ANY_VALUE and FIRST_VALUE return the first row of a group as read. The native hash aggregate does not guarantee that order, so converting such a SortAggregate broke the exact first/last results the fork promises on an ordered single split (CometFirstLastBoolAggFuzzSuite, 240 mismatches on Spark 4.1 and 559 on 3.5 for first, any_value and first_value over strings). A SortAggregate is now converted only when every aggregate function is order-insensitive (MIN, MAX, COUNT, SUM, AVG), whatever its input. Co-Authored-By: Claude Opus 5.5 (1M context) --- .../org/apache/comet/rules/CometExecRule.scala | 16 +++++++--------- .../expressions/aggregate/first_last.sql | 4 ++-- 2 files changed, 9 insertions(+), 11 deletions(-) diff --git a/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala b/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala index c461f43abfb..7b33aad7b0b 100644 --- a/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala +++ b/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala @@ -987,11 +987,11 @@ case class CometExecRule(session: SparkSession, wholePlan: Boolean = false) * is not mutable, such as MIN or MAX over strings, and runs it over input sorted by the * grouping keys. The hash aggregate starts from an `ObjectHashAggregateExec` with the same * fields, which Spark runs correctly over any input order, so any operator reverted to Spark - * later stays correct. The sort by the grouping keys directly below the aggregate is dropped, - * as rows within a group were in no particular order anyway. Without such a sort the input - * order could carry meaning, so only aggregates that do not depend on it are converted. A - * consumer that relied on the ordering of the sort aggregate gets a sort back, see - * [[restoreSortAggregateOrdering]]. + * later stays correct. Only aggregates whose result does not depend on the input order are + * converted: Spark's sort is stable, so FIRST, LAST and the like see the rows of a group in + * their input order, which the native hash aggregate does not guarantee. The sort by the + * grouping keys directly below the aggregate is then dropped. A consumer that relied on the + * ordering of the sort aggregate gets a sort back, see [[restoreSortAggregateOrdering]]. */ private def convertSortAggregate(agg: SortAggregateExec): Option[SparkPlan] = { val required = agg.requiredChildOrdering.headOption.getOrElse(Nil) @@ -1008,12 +1008,10 @@ case class CometExecRule(session: SparkSession, wholePlan: Boolean = false) case _ => false } val input = if (sortedByKeys) agg.child.children.head else agg.child - if (!sortedByKeys && - !agg.aggregateExpressions.forall(e => orderInsensitive(e.aggregateFunction))) { + if (!agg.aggregateExpressions.forall(e => orderInsensitive(e.aggregateFunction))) { withFallbackReason( agg, - "SortAggregate over input not sorted by its grouping keys alone may depend on the " + - "input order") + "SortAggregate with an aggregate function that depends on the input order stays in Spark") return None } if (!input.isInstanceOf[CometNativeExec]) { diff --git a/spark/src/test/resources/sql-tests/expressions/aggregate/first_last.sql b/spark/src/test/resources/sql-tests/expressions/aggregate/first_last.sql index 7fd5dcde731..b67a4247eb8 100644 --- a/spark/src/test/resources/sql-tests/expressions/aggregate/first_last.sql +++ b/spark/src/test/resources/sql-tests/expressions/aggregate/first_last.sql @@ -183,7 +183,7 @@ SELECT first(val) IGNORE NULLS FROM test_single_null -- first IGNORE NULLS: multiple data types -- ============================================================ -query +query expect_fallback(depends on the input order) SELECT grp, first(i_val) IGNORE NULLS, first(l_val) IGNORE NULLS, @@ -327,7 +327,7 @@ SELECT last(val) IGNORE NULLS FROM test_single_null -- last IGNORE NULLS: multiple data types -- ============================================================ -query +query expect_fallback(depends on the input order) SELECT grp, last(i_val) IGNORE NULLS, last(l_val) IGNORE NULLS,