Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
10 changes: 10 additions & 0 deletions spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala
Original file line number Diff line number Diff line change
Expand Up @@ -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)

Expand Down
28 changes: 27 additions & 1 deletion spark/src/test/scala/org/apache/comet/exec/CometExecSuite.scala
Original file line number Diff line number Diff line change
Expand Up @@ -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}
Expand Down Expand Up @@ -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",
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -31,15 +31,15 @@ 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}
import org.apache.spark.sql.comet.execution.arrow.{ArrowCachedBatchSerializer, CometCachedBatchHelper}
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
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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",
Expand Down
21 changes: 20 additions & 1 deletion spark/src/test/scala/org/apache/spark/sql/CometTestBase.scala
Original file line number Diff line number Diff line change
Expand Up @@ -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._
Expand Down Expand Up @@ -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
Expand Down
Loading