dwsmith1983 commented on code in PR #6547:
URL: https://github.com/apache/datafusion-comet/pull/6547#discussion_r4171286552
##########
spark/src/test/scala/org/apache/comet/rules/CometExecRuleSuite.scala:
##########
@@ -368,76 +439,148 @@ class CometExecRuleSuite extends CometTestBase {
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))
+ f
+ }
+ }
+ }
+ }
+
+ /**
+ * Runs the DPP broadcast link query, checks its answer and that the join
and the pruning
+ * broadcast are native, and returns the executed plan.
+ */
+ private def checkDppLinkQuery(): SparkPlan = {
+ 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)
+ plan
+ }
+
+ test("AQE DPP broadcast roots retain temporary logical links after an
unchanged replan") {
+ assume(isSpark35Plus, "Native AQE DPP requires Spark 3.5+")
+ // With direct read, the replan that first sees the dim side's shuffle
stage moves the
+ // aggregate under the broadcast onto a ShuffleScan, so that replan is not
unchanged.
+ withDppLinkViews(directRead = false) {
+ 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)
}
- }
- 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))
+ 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 plan = checkDppLinkQuery()
+ 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("AQE DPP broadcast roots keep their logical links through a ShuffleScan
replan") {
+ assume(isSpark35Plus, "Native AQE DPP requires Spark 3.5+")
+ withDppLinkViews(directRead = true) {
+ val pending = new ThreadLocal[List[Option[(SparkPlan, LogicalPlan)]]] {
+ override def initialValue(): List[Option[(SparkPlan, LogicalPlan)]] =
Nil
+ }
+ // For each replanned broadcast root: whether Comet changed it, whether
it kept the
+ // stage's direct and temporary links, and the leaf kinds of the native
plan under it.
+ val observed = new ConcurrentLinkedQueue[(Boolean, Boolean, Boolean,
Seq[String])]()
+ val tempTag = AdaptiveSparkPlanExec.TEMP_LOGICAL_PLAN_TAG
+ beforeCometPreparation = plan => {
+ val replan = plan match {
+ case broadcast: CometBroadcastExchangeExec =>
+ broadcast.getTagValue(SparkPlan.LOGICAL_PLAN_TAG).collect {
+ case stage: LogicalQueryStage => (broadcast.clone(),
stage.logicalPlan)
}
- } finally {
- beforeCometPreparation = _ => ()
- afterCometPreparation = _ => ()
+ 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) =>
Review Comment:
Added `withReplanObserver`, taking the captured value and the per-replan
check; both DPP tests use it with their assertions unchanged.
--
This is an automated message from the Apache Git Service.
To respond to the message, please log on to GitHub and use the
URL above to go to the specific comment.
To unsubscribe, e-mail: [email protected]
For queries about this service, please contact Infrastructure at:
[email protected]
---------------------------------------------------------------------
To unsubscribe, e-mail: [email protected]
For additional commands, e-mail: [email protected]