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..7b33aad7b0b 100644 --- a/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala +++ b/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala @@ -22,8 +22,8 @@ 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.aggregate.{AggregateMode, Final, Partial, PartialMerge} +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 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,17 @@ 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 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") + /** * 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 +598,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 +658,7 @@ case class CometExecRule(session: SparkSession, wholePlan: Boolean = false) } plan.transformUp { case op => - val converted = convertNode(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 @@ -928,6 +942,105 @@ case class CometExecRule(session: SparkSession, wholePlan: Boolean = false) } } + 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 + } + } + if (changed) op.withNewChildren(children) else op + } + + 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. 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) + 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 (!agg.aggregateExpressions.forall(e => orderInsensitive(e.aggregateFunction))) { + withFallbackReason( + agg, + "SortAggregate with an aggregate function that depends on the input order stays in Spark") + 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) + converted.foreach(_.setTagValue(CometExecRule.CONVERTED_SORT_AGGREGATE_TAG, ())) + 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..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 expect_fallback(SortAggregate is not supported) +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 expect_fallback(SortAggregate is not supported) +query expect_fallback(depends on the input order) 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..298d2a73eb5 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,133 @@ 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("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")) + 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") + 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") + } + } + } + + 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")