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 286f5adea00..6009906c843 100644 --- a/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala +++ b/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala @@ -408,6 +408,16 @@ case class CometExecRule(session: SparkSession) } } + // For AQE table-cache stage (Spark 3.5+) on a Comet cache scan. The operators above it are + // planned again once it materializes, and like a Comet shuffle stage it is a native input. + case s: QueryStageExec if s.plan.isInstanceOf[CometInMemoryTableScanExec] => + convertToComet(s, CometExchangeSink).getOrElse(s) + + // A CometSparkToColumnarExec from an earlier pass, which AQE reuses over a table-cache stage + // because it carries its scan's logical link. Wrap it again so re-planned parents convert. + case c: CometSparkToColumnarExec => + convertToComet(c, CometScanWrapper).getOrElse(c) + case op if shouldApplySparkToColumnar(conf, op) => convertToComet(op, CometSparkToColumnarExec).getOrElse(op) 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 973260b69f3..f730c1f43df 100644 --- a/spark/src/test/scala/org/apache/comet/exec/CometExecSuite.scala +++ b/spark/src/test/scala/org/apache/comet/exec/CometExecSuite.scala @@ -39,7 +39,7 @@ import org.apache.spark.sql.comet.execution.shuffle.{CometColumnarShuffle, Comet 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.columnar.CometInMemoryRelationHelper +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} import org.apache.spark.sql.execution.joins.{BroadcastHashJoinExec, BroadcastNestedLoopJoinExec, CartesianProductExec, SortMergeJoinExec} @@ -3949,6 +3949,32 @@ class CometExecSuite extends CometTestBase { } } + // https://github.com/apache/datafusion-comet/issues/6202 + test("SparkToColumnar over InMemoryTableScanExec survives the AQE re-plan") { + assume(isSpark35Plus, "Table-cache query stages require Spark 3.5+") + // CometInMemoryCacheSuite covers this too, but registers Comet's rules twice, through its + // plugin and CometTestBase. This suite registers them once, as production does. Reset the + // serializer so the relation is cached in Spark's format, as in the test above. + CometInMemoryRelationHelper.clearSerializer() + withSQLConf(SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "true") { + withTempView("table_cache_stage") { + spark + .range(0, 10000, 1, 4) + .selectExpr("id", "id % 10 AS k") + .createOrReplaceTempView("table_cache_stage") + spark.catalog.cacheTable("table_cache_stage") + val df = spark.sql("SELECT k, count(*) FROM table_cache_stage GROUP BY k") + checkAnswer(df, (0L until 10L).map(k => Row(k, 1000L))) + val plan = df.queryExecution.executedPlan + checkCometOperatorsInFinalPlan(plan, classOf[InMemoryTableScanExec]) + val conversions = collect(plan) { + case c: CometSparkToColumnarExec if isTableCacheStage(c.child) => c + } + assert(conversions.size == 1, plan) + } + } + } + test("SparkToColumnar eliminate redundant in AQE") { withSQLConf( SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "true", 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 bb9583e4a4d..b12ef02c5f7 100644 --- a/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheSuite.scala +++ b/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheSuite.scala @@ -31,7 +31,7 @@ import org.apache.arrow.vector.compression.{CompressionCodec, CompressionUtil, N import org.apache.arrow.vector.types.pojo.ArrowType import org.apache.spark.CometDriverPlugin import org.apache.spark.SparkConf -import org.apache.spark.sql.{CometTestBase, Row} +import org.apache.spark.sql.{CometTestBase, QueryTest, Row} import org.apache.spark.sql.catalyst.expressions.{And, Attribute, AttributeReference, Expression, GreaterThanOrEqual, LessThan, Literal} import org.apache.spark.sql.columnar.{CachedBatch, SimpleMetricsCachedBatch} import org.apache.spark.sql.comet.{CometBroadcastHashJoinExec, CometInMemoryTableScanExec, CometSortExec, CometSortMergeJoinExec} @@ -39,7 +39,7 @@ import org.apache.spark.sql.comet.execution.arrow.{ArrowCachedBatchSerializer, C import org.apache.spark.sql.comet.util.Utils import org.apache.spark.sql.execution.SortExec import org.apache.spark.sql.execution.adaptive.{AdaptiveSparkPlanExec, AQEShuffleReadExec, QueryStageExec, ShuffleQueryStageExec} -import org.apache.spark.sql.execution.columnar.{CometInMemoryRelationHelper, InMemoryRelation} +import org.apache.spark.sql.execution.columnar.{CometInMemoryRelationHelper, InMemoryRelation, InMemoryTableScanExec} import org.apache.spark.sql.execution.exchange.{Exchange, ReusedExchangeExec, ShuffleExchangeLike} import org.apache.spark.sql.execution.joins.{BroadcastHashJoinExec, SortMergeJoinExec} import org.apache.spark.sql.functions.max @@ -166,10 +166,7 @@ class CometInMemoryCacheSuite extends CometTestBase { checkAnswer(df, Seq(Row(1, 1), Row(2, 2))) val plan = df.queryExecution.executedPlan assert(plan.asInstanceOf[AdaptiveSparkPlanExec].isFinalPlan) - // Match by name because this suite must also compile on Spark 3.4. - assert(collect(plan) { - case s: QueryStageExec if s.getClass.getSimpleName == "TableCacheQueryStageExec" => s - }.size == 1) + assert(collect(plan) { case s: QueryStageExec if isTableCacheStage(s) => s }.size == 1) assert(collect(plan) { case s: CometInMemoryTableScanExec => s }.size == 1) assert(collect(plan) { case s: ShuffleQueryStageExec => s }.size == 1) assert(collect(plan) { case s: AQEShuffleReadExec => s }.isEmpty) @@ -230,6 +227,56 @@ class CometInMemoryCacheSuite extends CometTestBase { } } + // https://github.com/apache/datafusion-comet/issues/6202 + test("AQE keeps the operators above a table cache stage native once the stage materializes") { + assume(isSpark35Plus, "Table-cache query stages require Spark 3.5+") + // `t` holds 1000 ids for each k, and the ids for a given k sum to 1000 * k + 4995000. + val queries = Seq( + "SELECT k, count(*) FROM t GROUP BY k" -> (0L until 10L).map(k => Row(k, 1000L)), + "SELECT name, sum(id) FROM t JOIN d ON k = k2 GROUP BY name" -> + (0L until 10L).map(k => Row(s"n$k", 1000L * k + 4995000L)), + "SELECT sum(id) FROM t WHERE k = 3" -> Seq(Row(4998000L))) + withTempView("t", "d") { + spark + .range(0, 10000, 1, 4) + .selectExpr("id", "id % 10 AS k") + .createOrReplaceTempView("t") + spark + .range(10) + .selectExpr("id AS k2", "concat('n', id) AS name") + .createOrReplaceTempView("d") + for { + nativeCache <- Seq(true, false) + (query, expected) <- queries + } { + withAQECache { + withSQLConf(CometConf.COMET_EXEC_IN_MEMORY_CACHE_ENABLED.key -> nativeCache.toString) { + spark.catalog.cacheTable("t") + spark.catalog.cacheTable("d") + val scanClass: Class[_] = + if (nativeCache) classOf[CometInMemoryTableScanExec] + else classOf[InMemoryTableScanExec] + // The first run materializes the cache through the stage and the second reads it warm. + // checkToRDD = false keeps checkAnswer from warming the cache with a run of its own. + Seq("cold", "warm").foreach { cache => + val df = sql(query) + QueryTest.checkAnswer(df, expected, checkToRDD = false) + val plan = df.queryExecution.executedPlan + withClue(s"nativeCache=$nativeCache, $cache cache: $query\n$plan\n") { + val cacheScans = collect(plan) { + case s: QueryStageExec if isTableCacheStage(s) => s.plan + } + assert(cacheScans.nonEmpty) + assert(cacheScans.forall(scanClass.isInstance)) + checkCometOperatorsInFinalPlan(plan, classOf[InMemoryTableScanExec]) + } + } + } + } + } + } + } + test("CometInMemoryTableScan over CometCachedBatch") { withSQLConf( SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false", diff --git a/spark/src/test/scala/org/apache/spark/sql/CometTestBase.scala b/spark/src/test/scala/org/apache/spark/sql/CometTestBase.scala index 9740a4a5468..ff72204033f 100644 --- a/spark/src/test/scala/org/apache/spark/sql/CometTestBase.scala +++ b/spark/src/test/scala/org/apache/spark/sql/CometTestBase.scala @@ -42,7 +42,7 @@ import org.apache.spark.sql.catalyst.util.sideBySide import org.apache.spark.sql.comet.CometPlanChecker import org.apache.spark.sql.comet.execution.shuffle.{CometColumnarShuffle, CometNativeShuffle, CometShuffleExchangeExec} import org.apache.spark.sql.execution._ -import org.apache.spark.sql.execution.adaptive.AdaptiveSparkPlanHelper +import org.apache.spark.sql.execution.adaptive.{AdaptiveSparkPlanExec, AdaptiveSparkPlanHelper, AQEShuffleReadExec, QueryStageExec} import org.apache.spark.sql.execution.exchange.ReusedExchangeExec import org.apache.spark.sql.internal._ import org.apache.spark.sql.test._ @@ -629,6 +629,25 @@ abstract class CometTestBase checkPlanNotMissingInput(plan) } + /** + * [[checkCometOperators]] for an adaptive query that has run. `checkSparkAnswerAndOperator` + * inspects a DataFrame that has not run, which for an adaptive query is its initial plan, and + * `checkCometOperators` treats query stages as leaves, so this checks the final plan and the + * plan inside each of its stages. + */ + protected def checkCometOperatorsInFinalPlan( + plan: SparkPlan, + excludedClasses: Class[_]*): Unit = { + assert(plan.asInstanceOf[AdaptiveSparkPlanExec].isFinalPlan, s"The query has not run:\n$plan") + val stagePlans = collect(plan) { case s: QueryStageExec => s.plan } + val excluded = excludedClasses :+ classOf[QueryStageExec] :+ classOf[AQEShuffleReadExec] + (stripAQEPlan(plan) +: stagePlans).foreach(checkCometOperators(_, excluded: _*)) + } + + // Matched by name because Spark 3.4 has no TableCacheQueryStageExec. + protected def isTableCacheStage(plan: SparkPlan): Boolean = + plan.getClass.getSimpleName == "TableCacheQueryStageExec" + // checks the plan node has no missing inputs // such nodes represented in plan with exclamation mark ! // example: !CometWindowExec