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
19 changes: 18 additions & 1 deletion spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala
Original file line number Diff line number Diff line change
Expand Up @@ -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}
Expand Down Expand Up @@ -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.
Comment on lines +813 to +816

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

is it possible to link some spark code snippet here with version tag to show why this case is added? so readers can understand these comment with more context

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Updated in 92b0baa. Added version-pinned Spark 4.1.3 links beside the guard: LogicalQueryStageStrategy returns the existing physical root, then SparkStrategies.plan assigns its direct logical link. The comment also distinguishes the ordinary exchange path, where the exchange is behind a query-stage leaf, and keeps the direct/inherited distinction explicit.

// 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)
Expand Down
52 changes: 50 additions & 2 deletions spark/src/test/scala/org/apache/comet/exec/CometExecSuite.scala
Original file line number Diff line number Diff line change
Expand Up @@ -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}
Expand Down Expand Up @@ -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
Expand Down
231 changes: 229 additions & 2 deletions spark/src/test/scala/org/apache/comet/rules/CometExecRuleSuite.scala
Original file line number Diff line number Diff line change
Expand Up @@ -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}
Expand All @@ -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))
Expand Down Expand Up @@ -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
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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)
}
}
Expand Down
Loading