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 69bd0c9165..19c88a6a9d 100644 --- a/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala +++ b/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala @@ -34,7 +34,7 @@ import org.apache.spark.sql.comet.execution.shuffle.{CometColumnarShuffle, Comet 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, QueryStageExec, ShuffleQueryStageExec} +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.columnar.InMemoryTableScanExec import org.apache.spark.sql.execution.command.{DataWritingCommandExec, ExecutedCommandExec} @@ -806,6 +806,23 @@ case class CometExecRule(session: SparkSession) // Set up logical links newPlan = newPlan.transform { + case op: CometExec + if op + .getTagValue(SparkPlan.LOGICAL_PLAN_TAG) + .exists(_.isInstanceOf[LogicalQueryStage]) => + // AQE replanning reuses this physical root and links it to the current logical stage. + // originalPlan can still point to a subtree hidden inside that logical leaf, which + // AQE cannot replace in the current logical plan. Only preserve a direct stage link, + // not a link inherited from an ancestor. + // On the ordinary exchange path, the exchange itself is behind a QueryStageExec + // leaf and is not visited by this transform. + // Spark 4.1.3 returns the existing root in LogicalQueryStageStrategy and then calls + // setLogicalLink from SparkStrategies.plan: + // scalastyle:off line.size.limit + // https://github.com/apache/spark/blob/v4.1.3/sql/core/src/main/scala/org/apache/spark/sql/execution/adaptive/LogicalQueryStageStrategy.scala#L64-L65 + // https://github.com/apache/spark/blob/v4.1.3/sql/core/src/main/scala/org/apache/spark/sql/execution/SparkStrategies.scala#L78-L87 + // scalastyle:on line.size.limit + op case op: CometExec => if (op.originalPlan.logicalLink.isEmpty) { op.unsetTagValue(SparkPlan.LOGICAL_PLAN_TAG) diff --git a/spark/src/test/scala/org/apache/comet/exec/CometExecSuite.scala b/spark/src/test/scala/org/apache/comet/exec/CometExecSuite.scala index f730c1f43d..89f2a6ee41 100644 --- a/spark/src/test/scala/org/apache/comet/exec/CometExecSuite.scala +++ b/spark/src/test/scala/org/apache/comet/exec/CometExecSuite.scala @@ -33,12 +33,12 @@ import org.apache.spark.sql._ import org.apache.spark.sql.catalyst.{FunctionIdentifier, TableIdentifier} import org.apache.spark.sql.catalyst.catalog.{BucketSpec, CatalogStatistics, CatalogTable} import org.apache.spark.sql.catalyst.expressions.{DynamicPruningExpression, Expression, ExpressionInfo, Hex, Literal} -import org.apache.spark.sql.catalyst.expressions.aggregate.{AggregateMode, BloomFilterAggregate} +import org.apache.spark.sql.catalyst.expressions.aggregate.{AggregateMode, BloomFilterAggregate, Final} import org.apache.spark.sql.comet._ import org.apache.spark.sql.comet.execution.shuffle.{CometColumnarShuffle, CometNativeShuffle, CometShuffleExchangeExec} import org.apache.spark.sql.connector.catalog.InMemoryTableCatalog import org.apache.spark.sql.execution._ -import org.apache.spark.sql.execution.adaptive.{AdaptiveSparkPlanExec, BroadcastQueryStageExec} +import org.apache.spark.sql.execution.adaptive.{AdaptiveSparkPlanExec, BroadcastQueryStageExec, LogicalQueryStage} import org.apache.spark.sql.execution.columnar.{CometInMemoryRelationHelper, InMemoryTableScanExec} import org.apache.spark.sql.execution.datasources.parquet.ParquetFileFormat import org.apache.spark.sql.execution.exchange.{BroadcastExchangeExec, BroadcastExchangeLike, ReusedExchangeExec, ShuffleExchangeExec} @@ -2132,6 +2132,54 @@ class CometExecSuite extends CometTestBase { } } + test("AQE broadcasts native aggregates after replanning") { + withSQLConf( + SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "true", + SQLConf.AUTO_BROADCASTJOIN_THRESHOLD.key -> "-1", + SQLConf.ADAPTIVE_AUTO_BROADCASTJOIN_THRESHOLD.key -> "10485760", + SQLConf.SHUFFLE_PARTITIONS.key -> "4", + CometConf.COMET_SHUFFLE_MODE.key -> "native", + CometConf.COMET_SPARK_TO_ARROW_SUPPORTED_OPERATOR_LIST.key -> "Range") { + val df = sql(""" + |WITH s AS ( + | SELECT id % 64 AS k, SUM(id) AS v FROM range(0, 4096, 1, 4) GROUP BY id % 64 + |), r1 AS ( + | SELECT id % 64 AS k, SUM(id + 1) AS v FROM range(0, 3072, 1, 4) GROUP BY id % 64 + |), r2 AS ( + | SELECT id % 64 AS k, SUM(id + 7) AS v FROM range(0, 2048, 1, 4) GROUP BY id % 64 + |), g AS ( + | SELECT SUM(id) AS v FROM range(0, 1024, 1, 4) + |) + |SELECT SUM(s.v + COALESCE(r1.v, 0) + COALESCE(r2.v, 0) + g.v) + |FROM s LEFT JOIN r1 ON s.k = r1.k LEFT JOIN r2 ON s.k = r2.k CROSS JOIN g + |""".stripMargin) + val adaptive = df.queryExecution.executedPlan.asInstanceOf[AdaptiveSparkPlanExec] + assert(collect(adaptive.executedPlan) { case b: CometBroadcastHashJoinExec => b }.isEmpty) + + checkAnswer(df, Seq(Row(48738816L))) + + val finalPlan = adaptive.executedPlan + assert(collect(finalPlan) { case b: CometBroadcastHashJoinExec => b }.size == 2) + val broadcasts = collect(finalPlan) { case b: CometBroadcastExchangeExec => b } + val aggregates = broadcasts.flatMap { broadcast => + collect(broadcast.child) { + case a: CometHashAggregateExec + if a.modes.contains(Final) && a.groupingExpressions.nonEmpty => + a + } + } + assert(aggregates.size == 2) + aggregates.foreach { aggregate => + assert(aggregate.longMetric("output_rows").value == 64) + assert(aggregate.longMetric("elapsed_compute").value > 0) + assert( + aggregate + .getTagValue(SparkPlan.LOGICAL_PLAN_TAG) + .exists(_.isInstanceOf[LogicalQueryStage])) + } + } + } + test("CometShuffleExchangeExec logical link should be correct") { withTempView("v") { spark.sparkContext 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 1769663dbc..9a475ae2bf 100644 --- a/spark/src/test/scala/org/apache/comet/rules/CometExecRuleSuite.scala +++ b/spark/src/test/scala/org/apache/comet/rules/CometExecRuleSuite.scala @@ -19,25 +19,32 @@ package org.apache.comet.rules +import java.util.concurrent.ConcurrentLinkedQueue + +import scala.jdk.CollectionConverters._ import scala.util.Random +import org.scalatest.PrivateMethodTester._ + import org.apache.logging.log4j.Level import org.apache.spark.sql._ import org.apache.spark.sql.catalyst.FunctionIdentifier import org.apache.spark.sql.catalyst.expressions.{Attribute, AttributeReference, Expression, ExpressionInfo, In, InSet, KnownFloatingPointNormalized, Literal, Not} 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.rules.Rule import org.apache.spark.sql.comet._ import org.apache.spark.sql.comet.execution.shuffle.CometShuffleExchangeExec import org.apache.spark.sql.execution._ -import org.apache.spark.sql.execution.adaptive.{QueryStageExec, ShuffleQueryStageExec} +import org.apache.spark.sql.execution.adaptive.{AdaptiveSparkPlanExec, 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} import org.apache.spark.sql.internal.SQLConf import org.apache.spark.sql.types.{DataTypes, DoubleType, FloatType, StructField, StructType} -import org.apache.comet.{CometConf, CometCoverageStats, CometExplainInfo, ExtendedExplainInfo} +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.testing.{DataGenOptions, FuzzDataGenerator} @@ -49,6 +56,40 @@ import org.apache.comet.testing.{DataGenOptions, FuzzDataGenerator} */ class CometExecRuleSuite extends CometTestBase { + // The observers are active only during the DPP lifecycle regression below. AQE can prepare + // subqueries on different threads, so publish the callbacks and pair each thread's invocations. + @volatile private var beforeCometPreparation: SparkPlan => Unit = (_: SparkPlan) => () + @volatile private var afterCometPreparation: SparkPlan => Unit = (_: SparkPlan) => () + + override protected def createSparkSession: SparkSessionType = { + SparkSession.clearActiveSession() + SparkSession.clearDefaultSession() + SparkSession + .builder() + .config(sparkContext.getConf) + .withExtensions { extensions => + extensions.injectQueryStagePrepRule { _ => + new Rule[SparkPlan] { + override def apply(plan: SparkPlan): SparkPlan = { + beforeCometPreparation(plan) + plan + } + } + } + new CometSparkSessionExtensions().apply(extensions) + extensions.injectQueryStagePrepRule { _ => + new Rule[SparkPlan] { + override def apply(plan: SparkPlan): SparkPlan = { + afterCometPreparation(plan) + plan + } + } + } + } + .getOrCreate() + .asInstanceOf[SparkSessionType] + } + /** Helper method to apply CometExecRule and return the transformed plan */ private def applyCometExecRule(plan: SparkPlan): SparkPlan = { CometExecRule(spark).apply(stripAQEPlan(plan)) @@ -106,6 +147,192 @@ class CometExecRuleSuite extends CometTestBase { case a: ObjectHashAggregateExec if a.aggregateExpressions.forall(_.mode == Partial) => a }.get + /** A native final aggregate over a shuffle stage, as reused by AQE replanning. */ + private def createAdaptiveAggregate(): CometHashAggregateExec = { + val plan = createSparkPlan( + spark, + "SELECT id % 3 AS k, SUM(id) AS total FROM range(0, 100, 1, 2) GROUP BY id % 3") + val aggregate = applyCometExecRule(plan).asInstanceOf[CometHashAggregateExec] + val shuffle = aggregate.child.asInstanceOf[CometShuffleExchangeExec] + aggregate + .withNewChildren(Seq(ShuffleQueryStageExec(0, shuffle, shuffle.canonicalized))) + .asInstanceOf[CometHashAggregateExec] + } + + test("CometExecRule preserves the current direct AQE logical link") { + withSQLConf( + SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false", + SQLConf.WHOLESTAGE_CODEGEN_ENABLED.key -> "false", + CometConf.COMET_SPARK_TO_ARROW_SUPPORTED_OPERATOR_LIST.key -> "Range") { + val originalTags = + Seq(Some(SparkPlan.LOGICAL_PLAN_TAG), Some(SparkPlan.LOGICAL_PLAN_INHERITED_TAG), None) + originalTags.foreach { originalTag => + withClue(s"original logical tag: $originalTag") { + val aggregate = createAdaptiveAggregate() + val original = aggregate.originalPlan + val originalLogicalPlan = original.logicalLink.get + original.unsetTagValue(SparkPlan.LOGICAL_PLAN_TAG) + original.unsetTagValue(SparkPlan.LOGICAL_PLAN_INHERITED_TAG) + originalTag.foreach(original.setTagValue(_, originalLogicalPlan)) + + var current: SparkPlan = aggregate + (1 to 2).foreach { _ => + val logicalStage = LogicalQueryStage(originalLogicalPlan, current) + val replanned = spark.sessionState.planner.plan(logicalStage).next() + assert(replanned eq current) + assert(replanned.getTagValue(SparkPlan.LOGICAL_PLAN_TAG).exists(_ eq logicalStage)) + + current = applyCometExecRule(replanned) + assert(current.getTagValue(SparkPlan.LOGICAL_PLAN_TAG).exists(_ eq logicalStage)) + } + } + } + } + } + + test("CometExecRule repairs ordinary and inherited logical links from the original plan") { + withSQLConf( + SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false", + SQLConf.WHOLESTAGE_CODEGEN_ENABLED.key -> "false", + CometConf.COMET_SPARK_TO_ARROW_SUPPORTED_OPERATOR_LIST.key -> "Range") { + val originalTags = + Seq(Some(SparkPlan.LOGICAL_PLAN_TAG), Some(SparkPlan.LOGICAL_PLAN_INHERITED_TAG), None) + for (originalTag <- originalTags; hasDirectLink <- Seq(false, true)) { + withClue(s"original logical tag: $originalTag, ordinary direct link: $hasDirectLink") { + val aggregate = createAdaptiveAggregate() + val original = aggregate.originalPlan + val originalLogicalPlan = original.logicalLink.get + original.unsetTagValue(SparkPlan.LOGICAL_PLAN_TAG) + original.unsetTagValue(SparkPlan.LOGICAL_PLAN_INHERITED_TAG) + originalTag.foreach(original.setTagValue(_, originalLogicalPlan)) + + aggregate.unsetTagValue(SparkPlan.LOGICAL_PLAN_TAG) + aggregate.setTagValue( + SparkPlan.LOGICAL_PLAN_INHERITED_TAG, + LogicalQueryStage(originalLogicalPlan, aggregate)) + if (hasDirectLink) { + aggregate.setTagValue(SparkPlan.LOGICAL_PLAN_TAG, LocalRelation(aggregate.output)) + } + + val transformed = applyCometExecRule(aggregate) + if (originalTag.isDefined) { + assert(transformed.logicalLink.exists(_ eq originalLogicalPlan)) + } else { + assert(transformed.getTagValue(SparkPlan.LOGICAL_PLAN_TAG).isEmpty) + assert(transformed.getTagValue(SparkPlan.LOGICAL_PLAN_INHERITED_TAG).isEmpty) + } + } + } + } + } + + test("AQE DPP broadcast roots retain temporary logical links after an unchanged replan") { + assume(isSpark35Plus, "Native AQE DPP requires Spark 3.5+") + withSQLConf( + SQLConf.USE_V1_SOURCE_LIST.key -> "parquet", + SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "true", + SQLConf.ADAPTIVE_FORCE_OPTIMIZE_SKEWED_JOIN.key -> "false", + SQLConf.DYNAMIC_PARTITION_PRUNING_ENABLED.key -> "true", + SQLConf.DYNAMIC_PARTITION_PRUNING_REUSE_BROADCAST_ONLY.key -> "true", + SQLConf.SHUFFLE_PARTITIONS.key -> "2", + CometConf.COMET_SHUFFLE_MODE.key -> "native") { + withTempDir { dir => + withTempView("dpp_link_fact", "dpp_link_dim") { + withSQLConf(CometConf.COMET_ENABLED.key -> "false") { + spark + .range(64) + .selectExpr("CAST(id % 8 AS INT) AS k", "id AS v") + .write + .partitionBy("k") + .parquet(s"$dir/fact") + spark + .range(32) + .selectExpr( + "CAST(id % 8 AS INT) AS k", + "id AS v", + "IF(id % 2 = 0, 'DE', 'US') AS country") + .write + .parquet(s"$dir/dim") + } + spark.read.parquet(s"$dir/fact").createOrReplaceTempView("dpp_link_fact") + spark.read.parquet(s"$dir/dim").createOrReplaceTempView("dpp_link_dim") + + assert( + spark.sessionState.conf.getConf(SQLConf.ADAPTIVE_CUSTOM_COST_EVALUATOR_CLASS).isEmpty) + type Replan = (SparkPlan, LogicalPlan) + val pending = new ThreadLocal[List[Option[Replan]]] { + override def initialValue(): List[Option[Replan]] = Nil + } + val observed = new ConcurrentLinkedQueue[(CometBroadcastExchangeExec, LogicalPlan)]() + val tempTag = AdaptiveSparkPlanExec.TEMP_LOGICAL_PLAN_TAG + val costEvaluator = SimpleCostEvaluator(forceOptimizeSkewedJoin = false) + beforeCometPreparation = plan => { + val replan = plan match { + case broadcast: CometBroadcastExchangeExec => + broadcast.getTagValue(SparkPlan.LOGICAL_PLAN_TAG).collect { + case stage: LogicalQueryStage => + assert(stage.physicalPlan eq broadcast) + assert(broadcast.getTagValue(tempTag).exists(_ eq stage.logicalPlan)) + (broadcast.clone(), stage.logicalPlan) + } + case _ => None + } + pending.set(replan :: pending.get()) + } + afterCometPreparation = plan => { + val replan = pending.get().head + val remaining = pending.get().tail + if (remaining.isEmpty) pending.remove() else pending.set(remaining) + replan.foreach { case (previous, logicalPlan) => + val broadcast = plan.asInstanceOf[CometBroadcastExchangeExec] + assert(broadcast.logicalLink.exists(_ eq logicalPlan)) + assert(broadcast.getTagValue(tempTag).exists(_ eq logicalPlan)) + // Spark rejects an equal-cost candidate when its physical tree is unchanged. + // Pin both inputs to that decision, including Comet's retained temporary link. + assert(previous == broadcast) + assert(costEvaluator.evaluateCost(previous) == SimpleCost(0)) + assert(costEvaluator.evaluateCost(broadcast) == SimpleCost(0)) + observed.add( + (broadcast.clone().asInstanceOf[CometBroadcastExchangeExec], logicalPlan)) + } + } + try { + val df = sql(""" + |SELECT /*+ BROADCAST(d) */ f.k, f.total, d.total + |FROM (SELECT k, SUM(v) AS total FROM dpp_link_fact GROUP BY k) f + |JOIN (SELECT k, SUM(v) AS total FROM dpp_link_dim + | WHERE country = 'DE' GROUP BY k) d ON f.k = d.k + |""".stripMargin) + QueryTest.checkAnswer( + df, + (0 until 8 by 2).map(k => Row(k, 224L + 8L * k, 48L + 4L * k)), + checkToRDD = false) + val plan = df.queryExecution.executedPlan + assert(collect(plan) { case b: CometBroadcastHashJoinExec => b }.nonEmpty) + assert(collectWithSubqueries(plan) { case s: CometSubqueryBroadcastExec => + s + }.nonEmpty) + assert(!observed.isEmpty, "Expected a DPP broadcast root with a direct logical stage") + observed.iterator().asScala.foreach { case (broadcast, logicalPlan) => + // Give the isolated snapshot conflicting links to pin Spark's TEMP-over-direct + // precedence, which would otherwise be invisible after Comet repairs both. + broadcast.setLogicalLink(LogicalQueryStage(logicalPlan, broadcast)) + val stage = BroadcastQueryStageExec(0, broadcast, broadcast.canonicalized) + val setStageLink = PrivateMethod[Unit](Symbol("setLogicalLinkForNewQueryStage")) + plan + .asInstanceOf[AdaptiveSparkPlanExec] + .invokePrivate(setStageLink(stage, broadcast)) + assert(stage.logicalLink.exists(_ eq logicalPlan)) + } + } finally { + beforeCometPreparation = _ => () + afterCometPreparation = _ => () + } + } + } + } + } + test("expression-level fallback reasons are rolled up onto the operator that falls back") { // Extended explain only walks plan nodes, so a reason recorded on a sub-expression is // invisible unless CometExecRule lifts it onto the enclosing operator. Disabling a single diff --git a/spark/src/test/scala/org/apache/spark/sql/comet/execution/shuffle/CometNativePositionalRoundRobinSuite.scala b/spark/src/test/scala/org/apache/spark/sql/comet/execution/shuffle/CometNativePositionalRoundRobinSuite.scala index 77f9cedffb..20a2e1592a 100644 --- a/spark/src/test/scala/org/apache/spark/sql/comet/execution/shuffle/CometNativePositionalRoundRobinSuite.scala +++ b/spark/src/test/scala/org/apache/spark/sql/comet/execution/shuffle/CometNativePositionalRoundRobinSuite.scala @@ -197,7 +197,7 @@ class CometNativePositionalRoundRobinSuite extends CometTestBase with AdaptiveSp // Drop the map output so the next job re-runs the map stage under the new batch size. SparkEnv.get.mapOutputTracker .asInstanceOf[MapOutputTrackerMaster] - .unregisterAllMapAndMergeOutput(exchange.shuffleId) + .unregisterAllMapAndMergeOutput(exchange.shuffleDependency.shuffleId) assert(placement() == first) } }