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 9b56631aee8..aea2798af47 100644 --- a/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala +++ b/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala @@ -62,6 +62,15 @@ import org.apache.comet.shims.{CometTypeShim, ShimCometStreaming, ShimCometWindo object CometExecRule { + private[rules] def removePlaceholders(plan: SparkPlan): SparkPlan = plan.transformUp { + // revertUnsafePartialAggregates re-runs transform over already wrapped query stages, which + // can produce CometSinkPlaceHolder(CometSinkPlaceHolder(stage)). Remove sinks bottom-up. + case CometSinkPlaceHolder(_, _, child) => child + // Scan wrappers are leaves. Recurse explicitly to preserve the old top-down cleanup's + // coverage of the wrapped plan's descendants as well as remove wrappers at its root. + case CometScanWrapper(_, wrapped) => removePlaceholders(wrapped) + } + /** * Tag applied to Partial-mode aggregate operators that must NOT be converted to Comet because a * corresponding buffer-consuming aggregate cannot be converted, and the aggregate functions @@ -794,11 +803,7 @@ case class CometExecRule(session: SparkSession) } } - // Remove placeholders - newPlan = newPlan.transform { - case CometSinkPlaceHolder(_, _, s) => s - case CometScanWrapper(_, s) => s - } + newPlan = CometExecRule.removePlaceholders(newPlan) // Revert CometColumnarShuffle to Spark's ShuffleExchangeExec when both its parent and child // are non-Comet HashAggregate/ObjectHashAggregate operators that remained JVM after the main diff --git a/spark/src/test/scala/org/apache/comet/rules/CometExecRuleSuite.scala b/spark/src/test/scala/org/apache/comet/rules/CometExecRuleSuite.scala index d291e7ab1b7..a48ba2e5e1f 100644 --- a/spark/src/test/scala/org/apache/comet/rules/CometExecRuleSuite.scala +++ b/spark/src/test/scala/org/apache/comet/rules/CometExecRuleSuite.scala @@ -33,11 +33,12 @@ import org.apache.spark.sql.catalyst.expressions.{Attribute, AttributeReference, import org.apache.spark.sql.catalyst.expressions.aggregate.{AggregateExpression, BloomFilterAggregate, Final, Min, Partial, PartialMerge} import org.apache.spark.sql.catalyst.optimizer.NormalizeNaNAndZero import org.apache.spark.sql.catalyst.plans.logical.{LocalRelation, LogicalPlan} +import org.apache.spark.sql.catalyst.plans.physical.{HashPartitioning, RoundRobinPartitioning} import org.apache.spark.sql.catalyst.rules.Rule import org.apache.spark.sql.comet._ -import org.apache.spark.sql.comet.execution.shuffle.CometShuffleExchangeExec +import org.apache.spark.sql.comet.execution.shuffle.{CometNativeShuffle, CometShuffleExchangeExec} import org.apache.spark.sql.execution._ -import org.apache.spark.sql.execution.adaptive.{AdaptiveSparkPlanExec, BroadcastQueryStageExec, LogicalQueryStage, QueryStageExec, ShuffleQueryStageExec, SimpleCost, SimpleCostEvaluator} +import org.apache.spark.sql.execution.adaptive.{AdaptiveSparkPlanExec, AQEShuffleReadExec, BroadcastQueryStageExec, LogicalQueryStage, QueryStageExec, ShuffleQueryStageExec, SimpleCost, SimpleCostEvaluator} import org.apache.spark.sql.execution.aggregate.{HashAggregateExec, ObjectHashAggregateExec} import org.apache.spark.sql.execution.datasources.v2.BatchScanExec import org.apache.spark.sql.execution.exchange.{BroadcastExchangeExec, ShuffleExchangeExec} @@ -46,7 +47,7 @@ import org.apache.spark.sql.types.{DataTypes, DoubleType, FloatType, StructField import org.apache.comet.{CometConf, CometCoverageStats, CometExplainInfo, CometSparkSessionExtensions, ExtendedExplainInfo} import org.apache.comet.CometSparkSessionExtensions.{isSpark35Plus, isSpark40Plus, isSpark42Plus, withFallbackReason} -import org.apache.comet.serde.{CometAggregateExpressionSerde, Compatible, ExprOuterClass, QueryPlanSerde, Unsupported} +import org.apache.comet.serde.{CometAggregateExpressionSerde, Compatible, ExprOuterClass, OperatorOuterClass, QueryPlanSerde, Unsupported} import org.apache.comet.testing.{DataGenOptions, FuzzDataGenerator} /** @@ -159,6 +160,114 @@ class CometExecRuleSuite extends CometTestBase { .asInstanceOf[CometHashAggregateExec] } + test("placeholder cleanup preserves AQE stages behind nested and leaf wrappers") { + val input = createSparkPlan(spark, "SELECT * FROM VALUES (1), (2) AS t(id)") + val exchange = ShuffleExchangeExec(HashPartitioning(input.output, 2), input) + val stage = ShuffleQueryStageExec(0, exchange, exchange.canonicalized) + val nativeOp = OperatorOuterClass.Operator.getDefaultInstance + def sink(child: SparkPlan): SparkPlan = CometSinkPlaceHolder(nativeOp, child, child) + def scan(child: SparkPlan): SparkPlan = CometScanWrapper(nativeOp, child) + val specs = Seq(CoalescedPartitionSpec(0, 2, Some(16L))) + + for { + child <- Seq(sink(sink(stage)), sink(scan(stage)), scan(sink(stage)), scan(scan(stage))) + wrapRead <- Seq(false, true) + } { + val read = AQEShuffleReadExec(child, specs) + val input = if (wrapRead) scan(read) else read + val cleaned = CometExecRule.removePlaceholders(ProjectExec(read.output, input)) + val cleanedRead = cleaned.children.head.asInstanceOf[AQEShuffleReadExec] + assert(cleanedRead.child eq stage) + assert(cleanedRead.partitionSpecs eq specs) + assert(!cleaned.exists { + case _: CometSinkPlaceHolder | _: CometScanWrapper => true + case _ => false + }) + assert(CometExecRule.removePlaceholders(cleaned) eq cleaned) + } + } + + test("aggregate fallback removes nested placeholders around an existing AQE stage") { + withSQLConf( + SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false", + SQLConf.WHOLESTAGE_CODEGEN_ENABLED.key -> "false", + CometConf.COMET_EXEC_LOCAL_TABLE_SCAN_ENABLED.key -> "true", + CometConf.COMET_SHUFFLE_ENABLED.key -> "true", + CometConf.COMET_SHUFFLE_MODE.key -> "native", + CometConf.COMET_SHUFFLE_NATIVE_HASH_PARTITIONING_ENABLED.key -> "false", + CometConf.COMET_SHUFFLE_NATIVE_ROUND_ROBIN_PARTITIONING_ENABLED.key -> "true") { + withTempView("test_data") { + createTestDataFrame.createOrReplaceTempView("test_data") + val original = + createSparkPlan(spark, "SELECT AVG(id) FROM test_data GROUP BY (id % 3)") + val partial = original.collectFirst { + case agg: HashAggregateExec if agg.aggregateExpressions.forall(_.mode == Partial) => agg + }.get + // Keep a previously converted input stage below the aggregate being planned. Its + // output attributes must match the Partial's original input, including grouping keys. + val exchange = + applyCometExecRule(ShuffleExchangeExec(RoundRobinPartitioning(4), partial.child)) + .asInstanceOf[CometShuffleExchangeExec] + val stage = ShuffleQueryStageExec(0, exchange, exchange.canonicalized) + val staged = original.transformUp { + case agg: HashAggregateExec if agg eq partial => agg.copy(child = stage) + } + + // The Partial initially converts, but its upper hash shuffle cannot. Repair restores + // Spark's AVG buffer producer and revisits the already wrapped round-robin input stage. + val result = applyCometExecRule(staged) + assert(result.collect { case agg: HashAggregateExec => agg }.size == 2) + assert(!result.exists(_.isInstanceOf[CometHashAggregateExec])) + val restoredPartial = result.collectFirst { + case agg: HashAggregateExec if agg.aggregateExpressions.forall(_.mode == Partial) => + agg + }.get + assert(restoredPartial.getTagValue(CometExecRule.COMET_UNSAFE_PARTIAL).isDefined) + assert(!result.exists { + case _: CometSinkPlaceHolder | _: CometScanWrapper => true + case _ => false + }) + assert(restoredPartial.child eq stage) + assert(result.collect { case s: ShuffleQueryStageExec => s } == Seq(stage)) + } + } + } + + test("aggregate fallback executes after an AQE round-robin repartition") { + withSQLConf( + SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "true", + SQLConf.SHUFFLE_PARTITIONS.key -> "4", + CometConf.COMET_SHUFFLE_ENABLED.key -> "true", + CometConf.COMET_SHUFFLE_MODE.key -> "native", + CometConf.COMET_SHUFFLE_NATIVE_HASH_PARTITIONING_ENABLED.key -> "false", + CometConf.COMET_SHUFFLE_NATIVE_ROUND_ROBIN_PARTITIONING_ENABLED.key -> "true") { + val data = Seq((1, Some(10)), (1, Some(20)), (2, None), (2, Some(40)), (3, None)) + withParquetTable(data, "test_data") { + val df = sql(""" + |SELECT k, AVG(v) FROM + | (SELECT /*+ REPARTITION(4) */ _1 AS k, _2 AS v FROM test_data) + |GROUP BY k + |""".stripMargin) + QueryTest.checkAnswer(df, Seq(Row(1, 15.0), Row(2, 40.0), Row(3, null))) + val plan = df.queryExecution.executedPlan + assert(plan.asInstanceOf[AdaptiveSparkPlanExec].isFinalPlan) + val nativeShuffles = collect(plan) { case s: CometShuffleExchangeExec => s } + assert(nativeShuffles.size == 1) + assert(nativeShuffles.head.shuffleType == CometNativeShuffle) + assert(nativeShuffles.head.outputPartitioning == RoundRobinPartitioning(4)) + val aggregates = collect(plan) { case a: HashAggregateExec => a } + assert(aggregates.size == 2) + assert(collect(plan) { case a: CometHashAggregateExec => a }.isEmpty) + val partial = aggregates.find(_.aggregateExpressions.forall(_.mode == Partial)).get + assert(partial.getTagValue(CometExecRule.COMET_UNSAFE_PARTIAL).isDefined) + assert(collect(plan) { + case p: CometSinkPlaceHolder => p + case p: CometScanWrapper => p + }.isEmpty) + } + } + } + test("CometExecRule preserves the current direct AQE logical link") { withSQLConf( SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false",