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
15 changes: 10 additions & 5 deletions spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala
Original file line number Diff line number Diff line change
Expand Up @@ -62,6 +62,15 @@ import org.apache.comet.shims.{CometTypeShim, ShimCometStreaming, ShimCometWindo

object CometExecRule {

private[rules] def removePlaceholders(plan: SparkPlan): SparkPlan = plan.transformUp {
// revertUnsafePartialAggregates re-runs transform over already wrapped query stages, which
// can produce CometSinkPlaceHolder(CometSinkPlaceHolder(stage)). Remove sinks bottom-up.
case CometSinkPlaceHolder(_, _, child) => child
// Scan wrappers are leaves. Recurse explicitly to preserve the old top-down cleanup's
// coverage of the wrapped plan's descendants as well as remove wrappers at its root.
case CometScanWrapper(_, wrapped) => removePlaceholders(wrapped)
}

/**
* Tag applied to Partial-mode aggregate operators that must NOT be converted to Comet because a
* corresponding buffer-consuming aggregate cannot be converted, and the aggregate functions
Expand Down Expand Up @@ -794,11 +803,7 @@ case class CometExecRule(session: SparkSession)
}
}

// Remove placeholders
newPlan = newPlan.transform {
case CometSinkPlaceHolder(_, _, s) => s
case CometScanWrapper(_, s) => s
}
newPlan = CometExecRule.removePlaceholders(newPlan)

// Revert CometColumnarShuffle to Spark's ShuffleExchangeExec when both its parent and child
// are non-Comet HashAggregate/ObjectHashAggregate operators that remained JVM after the main
Expand Down
115 changes: 112 additions & 3 deletions spark/src/test/scala/org/apache/comet/rules/CometExecRuleSuite.scala
Original file line number Diff line number Diff line change
Expand Up @@ -33,11 +33,12 @@ import org.apache.spark.sql.catalyst.expressions.{Attribute, AttributeReference,
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.plans.physical.{HashPartitioning, RoundRobinPartitioning}
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.comet.execution.shuffle.{CometNativeShuffle, CometShuffleExchangeExec}
import org.apache.spark.sql.execution._
import org.apache.spark.sql.execution.adaptive.{AdaptiveSparkPlanExec, BroadcastQueryStageExec, LogicalQueryStage, QueryStageExec, ShuffleQueryStageExec, SimpleCost, SimpleCostEvaluator}
import org.apache.spark.sql.execution.adaptive.{AdaptiveSparkPlanExec, AQEShuffleReadExec, 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}
Expand All @@ -46,7 +47,7 @@ import org.apache.spark.sql.types.{DataTypes, DoubleType, FloatType, StructField

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.serde.{CometAggregateExpressionSerde, Compatible, ExprOuterClass, OperatorOuterClass, QueryPlanSerde, Unsupported}
import org.apache.comet.testing.{DataGenOptions, FuzzDataGenerator}

/**
Expand Down Expand Up @@ -159,6 +160,114 @@ class CometExecRuleSuite extends CometTestBase {
.asInstanceOf[CometHashAggregateExec]
}

test("placeholder cleanup preserves AQE stages behind nested and leaf wrappers") {
val input = createSparkPlan(spark, "SELECT * FROM VALUES (1), (2) AS t(id)")
val exchange = ShuffleExchangeExec(HashPartitioning(input.output, 2), input)
val stage = ShuffleQueryStageExec(0, exchange, exchange.canonicalized)
val nativeOp = OperatorOuterClass.Operator.getDefaultInstance
def sink(child: SparkPlan): SparkPlan = CometSinkPlaceHolder(nativeOp, child, child)
def scan(child: SparkPlan): SparkPlan = CometScanWrapper(nativeOp, child)
val specs = Seq(CoalescedPartitionSpec(0, 2, Some(16L)))

for {
child <- Seq(sink(sink(stage)), sink(scan(stage)), scan(sink(stage)), scan(scan(stage)))
wrapRead <- Seq(false, true)
} {
val read = AQEShuffleReadExec(child, specs)
val input = if (wrapRead) scan(read) else read
val cleaned = CometExecRule.removePlaceholders(ProjectExec(read.output, input))
val cleanedRead = cleaned.children.head.asInstanceOf[AQEShuffleReadExec]
assert(cleanedRead.child eq stage)
assert(cleanedRead.partitionSpecs eq specs)
assert(!cleaned.exists {
case _: CometSinkPlaceHolder | _: CometScanWrapper => true
case _ => false
})
assert(CometExecRule.removePlaceholders(cleaned) eq cleaned)
}
}

test("aggregate fallback removes nested placeholders around an existing AQE stage") {
withSQLConf(
SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false",
SQLConf.WHOLESTAGE_CODEGEN_ENABLED.key -> "false",
CometConf.COMET_EXEC_LOCAL_TABLE_SCAN_ENABLED.key -> "true",
CometConf.COMET_SHUFFLE_ENABLED.key -> "true",
CometConf.COMET_SHUFFLE_MODE.key -> "native",
CometConf.COMET_SHUFFLE_NATIVE_HASH_PARTITIONING_ENABLED.key -> "false",
CometConf.COMET_SHUFFLE_NATIVE_ROUND_ROBIN_PARTITIONING_ENABLED.key -> "true") {
withTempView("test_data") {
createTestDataFrame.createOrReplaceTempView("test_data")
val original =
createSparkPlan(spark, "SELECT AVG(id) FROM test_data GROUP BY (id % 3)")
val partial = original.collectFirst {
case agg: HashAggregateExec if agg.aggregateExpressions.forall(_.mode == Partial) => agg
}.get
// Keep a previously converted input stage below the aggregate being planned. Its
// output attributes must match the Partial's original input, including grouping keys.
val exchange =
applyCometExecRule(ShuffleExchangeExec(RoundRobinPartitioning(4), partial.child))
.asInstanceOf[CometShuffleExchangeExec]
val stage = ShuffleQueryStageExec(0, exchange, exchange.canonicalized)
val staged = original.transformUp {
case agg: HashAggregateExec if agg eq partial => agg.copy(child = stage)
}

// The Partial initially converts, but its upper hash shuffle cannot. Repair restores
// Spark's AVG buffer producer and revisits the already wrapped round-robin input stage.
val result = applyCometExecRule(staged)
assert(result.collect { case agg: HashAggregateExec => agg }.size == 2)
assert(!result.exists(_.isInstanceOf[CometHashAggregateExec]))
val restoredPartial = result.collectFirst {
case agg: HashAggregateExec if agg.aggregateExpressions.forall(_.mode == Partial) =>
agg
}.get
assert(restoredPartial.getTagValue(CometExecRule.COMET_UNSAFE_PARTIAL).isDefined)
assert(!result.exists {
case _: CometSinkPlaceHolder | _: CometScanWrapper => true
case _ => false
})
assert(restoredPartial.child eq stage)
assert(result.collect { case s: ShuffleQueryStageExec => s } == Seq(stage))
}
}
}

test("aggregate fallback executes after an AQE round-robin repartition") {
withSQLConf(
SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "true",
SQLConf.SHUFFLE_PARTITIONS.key -> "4",
CometConf.COMET_SHUFFLE_ENABLED.key -> "true",
CometConf.COMET_SHUFFLE_MODE.key -> "native",
CometConf.COMET_SHUFFLE_NATIVE_HASH_PARTITIONING_ENABLED.key -> "false",
CometConf.COMET_SHUFFLE_NATIVE_ROUND_ROBIN_PARTITIONING_ENABLED.key -> "true") {
val data = Seq((1, Some(10)), (1, Some(20)), (2, None), (2, Some(40)), (3, None))
withParquetTable(data, "test_data") {
val df = sql("""
|SELECT k, AVG(v) FROM
| (SELECT /*+ REPARTITION(4) */ _1 AS k, _2 AS v FROM test_data)
|GROUP BY k
|""".stripMargin)
QueryTest.checkAnswer(df, Seq(Row(1, 15.0), Row(2, 40.0), Row(3, null)))
val plan = df.queryExecution.executedPlan
assert(plan.asInstanceOf[AdaptiveSparkPlanExec].isFinalPlan)
val nativeShuffles = collect(plan) { case s: CometShuffleExchangeExec => s }
assert(nativeShuffles.size == 1)
assert(nativeShuffles.head.shuffleType == CometNativeShuffle)
assert(nativeShuffles.head.outputPartitioning == RoundRobinPartitioning(4))
val aggregates = collect(plan) { case a: HashAggregateExec => a }
assert(aggregates.size == 2)
assert(collect(plan) { case a: CometHashAggregateExec => a }.isEmpty)
val partial = aggregates.find(_.aggregateExpressions.forall(_.mode == Partial)).get
assert(partial.getTagValue(CometExecRule.COMET_UNSAFE_PARTIAL).isDefined)
assert(collect(plan) {
case p: CometSinkPlaceHolder => p
case p: CometScanWrapper => p
}.isEmpty)
}
}
}

test("CometExecRule preserves the current direct AQE logical link") {
withSQLConf(
SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false",
Expand Down
Loading