dwsmith1983 commented on code in PR #6547:
URL: https://github.com/apache/datafusion-comet/pull/6547#discussion_r4171279527
##########
spark/src/test/scala/org/apache/comet/exec/CometNativeShuffleSuite.scala:
##########
@@ -1391,6 +1394,178 @@ class CometNativeShuffleSuite extends CometTestBase
with AdaptiveSparkPlanHelper
}
}
+ /**
+ * For each native block in `plan`, query stages included, that reads a
shuffle: each input of
+ * the block, in walk order, paired with the kind of the native leaf that
reads it.
+ */
+ private def shuffleReadingBlocks(plan: SparkPlan): Seq[Seq[(String,
String)]] = {
+ def inputs(op: SparkPlan): Seq[SparkPlan] = op.children.flatMap {
+ case native: CometNativeExec if native.children.nonEmpty =>
inputs(native)
+ case other => Seq(other)
+ }
+ def leaves(op: OperatorOuterClass.Operator):
Seq[OperatorOuterClass.Operator] =
+ if (op.getChildrenCount == 0) Seq(op)
+ else op.getChildrenList.asScala.toSeq.flatMap(leaves)
+ def readsShuffle(input: SparkPlan): Boolean = input match {
+ case _: ShuffleQueryStageExec | _: AQEShuffleReadExec | _:
CometShuffleExchangeExec => true
+ case _ => false
+ }
+ // A node AQE reuses can keep the serialized plan of a block it used to be
the root of.
+ val natives = collect(plan) { case native: CometNativeExec => native }
+ val nested = natives.flatMap(_.children.collect { case child:
CometNativeExec => child })
+ natives
+ .filter(root => root.serializedPlanOpt.isDefined && !nested.exists(_ eq
root))
+ .filter(inputs(_).exists(readsShuffle))
+ .map { root =>
+ val nativePlan =
OperatorOuterClass.Operator.parseFrom(root.serializedPlanOpt.plan.get)
+ val blockInputs = inputs(root).map(_.getClass.getSimpleName)
+ val leafKinds = leaves(nativePlan).map(_.getOpStructCase.name)
+ assert(blockInputs.length == leafKinds.length, s"$blockInputs vs
$leafKinds")
+ blockInputs.zip(leafKinds)
+ }
+ }
+
+ private val shuffleStageRead = Seq("ShuffleQueryStageExec" -> "SHUFFLE_SCAN")
+
+ private val aqeWithoutCoalescing = Seq(
+ SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "true",
+ SQLConf.COALESCE_PARTITIONS_ENABLED.key -> "false")
+
+ /** A two-phase aggregate over many map partitions. */
+ private def groupedSum: DataFrame =
+ spark
+ .range(0, 10000, 1, 50)
+ .groupBy((col("id") % 997).as("k"))
+ .agg(sum("id"), count("id"))
+
+ test("AQE final aggregate reads its shuffle stage through ShuffleScan") {
Review Comment:
Folded into one table-driven test over AQE, AQE with coalescing, and AQE
off. The direct-read-off case is dropped here and stays at the rule level.
--
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]