dwsmith1983 commented on code in PR #6547:
URL: https://github.com/apache/datafusion-comet/pull/6547#discussion_r4171282389
##########
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") {
+ withSQLConf(aqeWithoutCoalescing: _*) {
+ val (_, plan) = checkSparkAnswerAndOperator(groupedSum)
+ checkCometOperatorsInFinalPlan(plan)
+ assert(shuffleReadingBlocks(plan) == Seq(shuffleStageRead), plan)
+ }
+ }
+
+ test("AQE final aggregate keeps a plain Scan with shuffle direct read
disabled") {
+ withSQLConf(
+ aqeWithoutCoalescing :+ (CometConf.COMET_SHUFFLE_DIRECT_READ_ENABLED.key
-> "false"): _*) {
+ val (_, plan) = checkSparkAnswerAndOperator(groupedSum)
+ assert(shuffleReadingBlocks(plan) == Seq(Seq("ShuffleQueryStageExec" ->
"SCAN")), plan)
+ }
+ }
+
+ test("AQE native aggregate chain over a shuffle stage reads it through
ShuffleScan") {
+ withSQLConf(aqeWithoutCoalescing: _*) {
+ // A distinct aggregate plans two shuffles, with two native aggregates
between them.
+ val df = spark
+ .range(0, 10000, 1, 20)
+ .groupBy((col("id") % 97).as("k"))
+ .agg(countDistinct(col("id") % 13), sum("id"))
+ val (_, plan) = checkSparkAnswerAndOperator(df)
+ checkCometOperatorsInFinalPlan(plan)
+ assert(shuffleReadingBlocks(plan) == Seq(shuffleStageRead,
shuffleStageRead), plan)
+ }
+ }
+
+ test("AQE join of final aggregates reads only its shuffle stages through
ShuffleScan") {
+ val agg = spark
+ .range(0, 10000, 1, 20)
+ .groupBy((col("id") % 97).as("k"))
+ .agg(sum("id").as("s"))
+ withSQLConf(
+ aqeWithoutCoalescing ++ Seq(
+ SQLConf.AUTO_BROADCASTJOIN_THRESHOLD.key -> "-1",
+ SQLConf.ADAPTIVE_AUTO_BROADCASTJOIN_THRESHOLD.key -> "-1"): _*) {
+ // The second side reuses the first side's shuffle.
+ val df = agg.join(agg.select(col("k"), (col("s") + 1).as("s2")), "k")
+ val (_, plan) = checkSparkAnswerAndOperator(df)
+ checkCometOperatorsInFinalPlan(plan, classOf[ReusedExchangeExec])
+ assert(collect(plan) { case r: ReusedExchangeExec => r }.nonEmpty, plan)
+ assert(shuffleReadingBlocks(plan) == Seq(shuffleStageRead ++
shuffleStageRead), plan)
+ }
+ withSQLConf(aqeWithoutCoalescing: _*) {
+ // The broadcast side stays a plain Scan.
+ val dim = spark.range(0, 50).select(col("id").as("k"), (col("id") *
2).as("w"))
+ val df = agg.join(broadcast(dim), "k")
+ val (_, plan) = checkSparkAnswerAndOperator(df)
+ checkCometOperatorsInFinalPlan(plan)
+ val joinBlock = shuffleStageRead :+ ("BroadcastQueryStageExec" -> "SCAN")
+ assert(shuffleReadingBlocks(plan) == Seq(joinBlock), plan)
+ }
+ }
+
+ test("AQE keeps exchange reuse above equivalent aggregates that read through
ShuffleScan") {
+ withSQLConf(aqeWithoutCoalescing: _*) {
+ // Each side's aggregate reads its own shuffle stage, so their
ShuffleScan sources differ.
+ val agg = spark
+ .range(0, 10000, 1, 20)
+ .groupBy((col("id") % 97).as("k"))
+ .agg(sum("id").as("s"))
+ .repartition(7, col("s"))
+ val (_, plan) = checkSparkAnswerAndOperator(agg.union(agg))
+ checkCometOperatorsInFinalPlan(plan, classOf[ReusedExchangeExec])
+ assertExchangeReuseOver(plan, "Expected the exchange above the aggregate
to be reused") {
+ case a: CometHashAggregateExec if a.modes.contains(Final) => a
+ }
+ val blocks = shuffleReadingBlocks(plan)
+ val readsShuffleScan = blocks.forall(_.forall { case (_, leaf) => leaf
== "SHUFFLE_SCAN" })
+ assert(blocks.nonEmpty && readsShuffleScan, plan)
+ }
+ }
+
+ test("AQE final aggregate over a coalesced shuffle read reads it through
ShuffleScan") {
+ // The aggregate takes the stage's ShuffleScan before AQE puts the
coalesced read between
+ // them, and the ShuffleScan then reads the partitions that the coalesced
read specifies.
+ withSQLConf(
+ SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "true",
+ SQLConf.COALESCE_PARTITIONS_ENABLED.key -> "true") {
+ val (_, plan) = checkSparkAnswerAndOperator(groupedSum)
+ checkCometOperatorsInFinalPlan(plan)
+ assert(shuffleReadingBlocks(plan) == Seq(Seq("AQEShuffleReadExec" ->
"SHUFFLE_SCAN")), plan)
+ }
+ }
+
+ test("final aggregate without AQE reads its shuffle exchange through a plain
Scan") {
+ withSQLConf(SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false") {
+ val (_, plan) = checkSparkAnswerAndOperator(groupedSum)
+ assert(shuffleReadingBlocks(plan) == Seq(Seq("CometShuffleExchangeExec"
-> "SCAN")), plan)
+ }
+ }
+
+ test("AQE DPP query with a broadcast side reads its shuffle stage through
ShuffleScan") {
+ assume(isSpark35Plus, "Native AQE DPP requires Spark 3.5+")
+ withTempDir { dir =>
+ withSQLConf(CometConf.COMET_ENABLED.key -> "false") {
+ spark
+ .range(0, 1000, 1, 10)
+ .selectExpr("id", "id % 10 AS p")
+ .write
+ .partitionBy("p")
+ .parquet(s"$dir/fact")
+ spark.range(0, 10).selectExpr("id AS p", "id % 3 AS
x").write.parquet(s"$dir/dim")
+ }
+ withTempView("fact", "dim") {
+ spark.read.parquet(s"$dir/fact").createOrReplaceTempView("fact")
+ spark.read.parquet(s"$dir/dim").createOrReplaceTempView("dim")
+ withSQLConf(
+ aqeWithoutCoalescing :+
+ (SQLConf.DYNAMIC_PARTITION_PRUNING_ENABLED.key -> "true"): _*) {
+ val df = sql("""SELECT f.p, sum(f.id), count(*) FROM fact f JOIN dim
d ON f.p = d.p
+ |WHERE d.x = 1 GROUP BY f.p""".stripMargin)
+ val (_, plan) = checkSparkAnswer(df)
+ assert(shuffleReadingBlocks(plan) == Seq(shuffleStageRead), plan)
+
+ // Aggregates on both sides, the broadcast one feeding the pruning
filter. This checks
+ // the answer only.
+ checkSparkAnswer(sql("""SELECT /*+ BROADCAST(d) */ f.p, f.total,
d.total
+ |FROM (SELECT p, sum(id) AS total FROM fact GROUP BY p) f
+ |JOIN (SELECT p, sum(x) AS total FROM dim WHERE x = 1 GROUP BY
p) d
+ |ON f.p = d.p""".stripMargin))
Review Comment:
Dropped.
--
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]