diff --git a/spark/src/main/scala/org/apache/comet/CometSparkSessionExtensions.scala b/spark/src/main/scala/org/apache/comet/CometSparkSessionExtensions.scala index 782fe2beab9..bc77c102dc2 100644 --- a/spark/src/main/scala/org/apache/comet/CometSparkSessionExtensions.scala +++ b/spark/src/main/scala/org/apache/comet/CometSparkSessionExtensions.scala @@ -33,7 +33,7 @@ import org.apache.spark.sql.internal.SQLConf import org.apache.comet.CometConf._ import org.apache.comet.iceberg.IcebergWriteStrategy -import org.apache.comet.rules.{CometPlanAdaptiveDynamicPruningFilters, CometReuseSubquery, CometRule, CometSpark34AqeDppFallbackRule} +import org.apache.comet.rules.{CometCoalesceShufflePartitions, CometPlanAdaptiveDynamicPruningFilters, CometReuseSubquery, CometRule, CometSpark34AqeDppFallbackRule} import org.apache.comet.shims.ShimCometSparkSessionExtensions /** @@ -107,6 +107,7 @@ class CometSparkSessionExtensions } injectQueryStageOptimizerRuleShim(extensions, CometPlanAdaptiveDynamicPruningFilters) injectQueryStageOptimizerRuleShim(extensions, CometReuseSubquery) + injectQueryStageOptimizerRuleShim(extensions, CometCoalesceShufflePartitions) extensions.injectPlannerStrategy { session => IcebergWriteStrategy(session) } } diff --git a/spark/src/main/scala/org/apache/comet/rules/CometCoalesceShufflePartitions.scala b/spark/src/main/scala/org/apache/comet/rules/CometCoalesceShufflePartitions.scala new file mode 100644 index 00000000000..f970ae50b0c --- /dev/null +++ b/spark/src/main/scala/org/apache/comet/rules/CometCoalesceShufflePartitions.scala @@ -0,0 +1,128 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.comet.rules + +import org.apache.spark.sql.SparkSession +import org.apache.spark.sql.catalyst.trees.TreeNodeTag +import org.apache.spark.sql.comet.CometExec +import org.apache.spark.sql.execution.{SparkPlan, UnionExec} +import org.apache.spark.sql.execution.adaptive.{AQEShuffleReadExec, AQEShuffleReadRule, CoalesceShufflePartitions, ShuffleQueryStageExec} +import org.apache.spark.sql.execution.exchange.ShuffleOrigin +import org.apache.spark.sql.execution.joins.{BroadcastHashJoinExec, BroadcastNestedLoopJoinExec, CartesianProductExec} + +/** + * Coalesces the shuffle partitions below a Comet operator that Spark's CoalesceShufflePartitions + * coalesces child by child but does not recognize. + * + * Spark coalesces each child of a `UnionExec` as a group of its own, and from Spark 4.0 each + * child of a `CartesianProductExec`, `BroadcastHashJoinExec` or `BroadcastNestedLoopJoinExec` + * too. It matches those classes, and the Comet operators that replace them are other classes, so + * it falls through to the case that coalesces only when every leaf below the operator is an + * exchange stage. A union with a scan or a table-cache stage in one branch then keeps every + * partition of the shuffles in the others: `spark.sql.shuffle.partitions` tasks for a query that + * needs a few. + * + * Comet replaces these operators while AQE prepares a stage, before its optimizer rules run, and + * plans the operators above them against the Comet versions. So this runs after Spark's rule + * instead, on each such operator whose shuffle stages that rule left untouched. It rebuilds the + * Spark operator each Comet one replaced over the Comet children, has Spark's own rule coalesce + * that, and swaps the Comet operators back in. The partitions come out as Spark would have + * coalesced them, down to which operators count, since it is Spark's code deciding. The one + * difference is that Spark divides its minimum partition count among the coalesce groups of the + * whole plan, and this among those below the Comet operator, which are usually all of them. + * + * When every leaf below such an operator is an exchange stage, Spark's rule already coalesces its + * shuffles, together rather than child by child, and this leaves them as they are. + * + * Extending `AQEShuffleReadRule` gets this the same treatment from AQE as Spark's rule: it is + * skipped for the final stage when that stage's shuffle optimizations are off, and its result is + * discarded if it breaks a distribution required above it. + */ +case object CometCoalesceShufflePartitions extends AQEShuffleReadRule { + + // The Comet operator that a stand-in Spark operator was rebuilt from. + private val COMET_OPERATOR = TreeNodeTag[SparkPlan]("cometCoalesceShufflePartitions") + + // Required by the trait. Which shuffles are coalesced is decided by Spark's rule, which applies + // its own list. + override protected def supportedShuffleOrigins: Seq[ShuffleOrigin] = + CoalesceShufflePartitions(SparkSession.active).supportedShuffleOrigins + + override def apply(plan: SparkPlan): SparkPlan = { + if (!conf.coalesceShufflePartitionsEnabled || !plan.exists(replaced(_).isDefined)) { + return plan + } + plan.transformDown { + case p if replaced(p).isDefined && untouched(p) => coalesceBelow(p) + } + } + + // The Spark operator a Comet operator replaced, if Spark's rule coalesces its children one by + // one. The class match mirrors Spark's, and Spark's rule decides, for its version, which of + // these it actually treats that way. + private def replaced(plan: SparkPlan): Option[SparkPlan] = plan match { + case comet: CometExec => + comet.originalPlan match { + case original @ (_: UnionExec | _: CartesianProductExec | _: BroadcastHashJoinExec | + _: BroadcastNestedLoopJoinExec) + if original.children.length == comet.children.length => + Some(original) + case _ => None + } + case _ => None + } + + // No AQE rule has put a read over any shuffle stage below `plan`: Spark's rule coalesced none + // of them, and none is a skew-split or local read that coalescing now could disturb. + private def untouched(plan: SparkPlan): Boolean = + plan.exists(_.isInstanceOf[ShuffleQueryStageExec]) && + !plan.exists(_.isInstanceOf[AQEShuffleReadExec]) + + private def coalesceBelow(plan: SparkPlan): SparkPlan = { + val asSpark = plan.transformUp { case p => + replaced(p) match { + case Some(original) => + val standIn = original.withNewChildren(p.children) + // `withNewChildren` hands back the original itself when the children are the same ones, + // and the tag must not land on the operator that the Comet one keeps. + if (standIn eq original) { + p + } else { + standIn.setTagValue(COMET_OPERATOR, p) + standIn + } + case None => p + } + } + val coalesced = CoalesceShufflePartitions(SparkSession.active).apply(asSpark) + if (coalesced eq asSpark) plan else restore(coalesced) + } + + // Put each Comet operator back over the children of its stand-in. Rebuilt by hand rather than + // with transformUp, which copies a replaced node's tags onto a replacement that has none, and so + // could leave the stand-in's tag on the Comet operator. + private def restore(plan: SparkPlan): SparkPlan = { + val children = plan.children.map(restore) + plan.getTagValue(COMET_OPERATOR) match { + case Some(comet) => comet.withNewChildren(children) + case None => plan.withNewChildren(children) + } + } +} 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 9d1bf44d033..7c0f2e78999 100644 --- a/spark/src/test/scala/org/apache/comet/exec/CometExecSuite.scala +++ b/spark/src/test/scala/org/apache/comet/exec/CometExecSuite.scala @@ -38,7 +38,7 @@ 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, LogicalQueryStage} +import org.apache.spark.sql.execution.adaptive.{AdaptiveSparkPlanExec, AQEShuffleReadExec, 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} @@ -3573,6 +3573,44 @@ class CometExecSuite extends CometTestBase { } } + // https://github.com/apache/datafusion-comet/issues/6454 + test("AQE coalesces the shuffle partitions of a union whose other branch is a scan") { + // Spark coalesces each child of a union as its own group, but its rule did not recognize + // Comet's union, so the shuffled branch of a union with a scan kept every shuffle partition. + // Comet's rule defers to Spark's, so the query should come out partitioned as it is on Spark. + assume(isSpark35Plus, "Comet's query-stage optimizer rules need Spark 3.5+") + withTempPath { dir => + spark.range(0, 100, 1, 1).toDF("c").write.parquet(dir.getCanonicalPath) + withSQLConf( + SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "true", + SQLConf.COALESCE_PARTITIONS_MIN_PARTITION_NUM.key -> "1", + SQLConf.SHUFFLE_PARTITIONS.key -> "200") { + def query() = spark + .range(0, 10, 1, 2) + .toDF("c") + .repartition($"c") + .union(spark.read.parquet(dir.getCanonicalPath)) + var sparkPartitions = 0 + withSQLConf(CometConf.COMET_ENABLED.key -> "false") { + val df = query() + df.collect() + sparkPartitions = df.rdd.getNumPartitions + } + assert(sparkPartitions < 200, "Spark should have coalesced the shuffled branch") + + val df = query() + checkSparkAnswer(df) + // checkSparkAnswer runs copies of the query, so run this one to finalize its own plan. + df.collect() + val plan = df.queryExecution.executedPlan + assert(plan.asInstanceOf[AdaptiveSparkPlanExec].isFinalPlan) + assert(collect(plan) { case u: CometUnionExec => u }.size == 1) + assert(collect(plan) { case r: AQEShuffleReadExec if r.isCoalescedRead => r }.size == 1) + assert(df.rdd.getNumPartitions == sparkPartitions) + } + } + } + test("native execution after coalesce") { withTable("t1") { (0 until 5) diff --git a/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheSuite.scala b/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheSuite.scala index ed3d518d3ce..921dd2c5207 100644 --- a/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheSuite.scala +++ b/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheSuite.scala @@ -188,6 +188,26 @@ class CometInMemoryCacheSuite extends CometTestBase { } } + // https://github.com/apache/spark/blob/v4.1.2/sql/core/src/test/scala/org/apache/spark/sql/execution/adaptive/AdaptiveQueryExecSuite.scala#L3178-L3191 + test("AQE SPARK-42101: coalesce the shuffle partitions of a union with a table cache stage") { + assume(isSpark35Plus, "Table-cache query stages require Spark 3.5+") + withAQECache { + withSQLConf(SQLConf.COALESCE_PARTITIONS_MIN_PARTITION_NUM.key -> "1") { + val cached = Seq(1).toDF("c").cache() + val df = Seq(2).toDF("c").repartition($"c").union(cached) + checkAnswer(df, Seq(Row(1), Row(2))) + val plan = df.queryExecution.executedPlan + assert(plan.asInstanceOf[AdaptiveSparkPlanExec].isFinalPlan) + assert(collect(plan) { case u: org.apache.spark.sql.comet.CometUnionExec => u }.size == 1) + assert(collect(plan) { case r @ AQEShuffleReadExec(_: ShuffleQueryStageExec, _) => + r + }.size == 1) + assert(collect(plan) { case s: QueryStageExec if isTableCacheStage(s) => s }.size == 1) + assert(collect(plan) { case s: CometInMemoryTableScanExec => s }.size == 1) + } + } + } + // https://github.com/apache/spark/blob/v4.1.2/sql/core/src/test/scala/org/apache/spark/sql/execution/adaptive/AdaptiveQueryExecSuite.scala#L2780-L2832 test("AQE SPARK-37742: use valid Comet cache statistics for join selection") { withAQECache {