andygrove commented on code in PR #6607:
URL: https://github.com/apache/datafusion-comet/pull/6607#discussion_r4186482074


##########
spark/src/test/scala/org/apache/comet/exec/CometShuffleInputConversionSuite.scala:
##########
@@ -0,0 +1,310 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one
+ * or more contributor license agreements.  See the NOTICE file
+ * distributed with this work for additional information
+ * regarding copyright ownership.  The ASF licenses this file
+ * to you under the Apache License, Version 2.0 (the
+ * "License"); you may not use this file except in compliance
+ * with the License.  You may obtain a copy of the License at
+ *
+ *   http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing,
+ * software distributed under the License is distributed on an
+ * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+ * KIND, either express or implied.  See the License for the
+ * specific language governing permissions and limitations
+ * under the License.
+ */
+
+package org.apache.comet.exec
+
+import org.apache.spark.sql.{CometTestBase, DataFrame, Row}
+import org.apache.spark.sql.comet.CometSparkToColumnarExec
+import org.apache.spark.sql.comet.execution.shuffle.{CometColumnarShuffle, 
CometNativeShuffle, CometShuffleExchangeExec}
+import org.apache.spark.sql.execution.{ColumnarToRowExec, 
ColumnarToRowTransition, SparkPlan}
+import org.apache.spark.sql.functions.{array, avg, col, count, hash, length, 
max, min, size, sum}
+import org.apache.spark.sql.internal.SQLConf
+import org.apache.spark.sql.types.{ArrayType, BinaryType, 
CalendarIntervalType, DataTypes, DecimalType, IntegerType, LongType, 
StringType, StructType}
+import org.apache.spark.unsafe.types.CalendarInterval
+
+import org.apache.comet.{CometConf, ExtendedExplainInfo}
+
+// Top-level, so the encoder needs no outer pointer.
+case class ShuffleInputRec(a: Int, b: String)
+
+/** Tests for [[CometConf.COMET_CONVERT_FROM_SHUFFLE_INPUT_ENABLED]]. */
+class CometShuffleInputConversionSuite extends CometTestBase {
+
+  import testImplicits._
+
+  /**
+   * `CometTestBase` turns on the conversion of leaf operators such as RDD 
scans, which would
+   * convert the scans below the shuffles here before this conversion sees 
them. It is off by
+   * default.
+   */
+  private def withoutLeafConversion(f: => Unit): Unit =
+    withSQLConf(sparkToArrowConversionConfs(enabled = false): _*)(f)
+
+  /** Defines a test that runs with the shuffle input conversion enabled. */
+  private def convertTest(name: String)(f: => Unit): Unit =
+    test(name) {
+      withoutLeafConversion {
+        withSQLConf(CometConf.COMET_CONVERT_FROM_SHUFFLE_INPUT_ENABLED.key -> 
"true")(f)
+      }
+    }
+
+  private val rowSchema = new StructType()
+    .add("k", IntegerType)
+    .add("l", LongType)
+    .add("s", DataTypes.StringType)
+    .add("d", DataTypes.DoubleType)
+    .add("m", DecimalType(18, 2))
+
+  /**
+   * `rows` rows from an RDD, so the shuffle above them reads a Spark 
operator: the conversion of
+   * RDD scans, `spark.comet.convert.rdd.enabled`, is off by default.
+   */
+  private def rowsDf(rows: Int = 1000): DataFrame = {
+    val data = (0 until rows).map { i =>
+      Row(
+        i % 23,
+        i.toLong,
+        if (i % 11 == 0) null else s"value-$i",
+        i + 0.5d,
+        java.math.BigDecimal.valueOf(i * 7919L % 1000000L, 2))
+    }
+    spark.createDataFrame(spark.sparkContext.parallelize(data, 4), rowSchema)
+  }
+
+  /** Rows with an `array<int>` column, which `CometSparkToColumnarExec` does 
not convert. */
+  private def arraysDf(rows: Int = 100): DataFrame = {
+    val schema = new StructType().add("k", IntegerType).add("xs", 
ArrayType(IntegerType))
+    val data = (0 until rows).map(i => Row(i % 23, Seq(i, i + 1)))
+    spark.createDataFrame(spark.sparkContext.parallelize(data, 4), schema)
+  }
+
+  private def conversions(plan: SparkPlan): Seq[CometSparkToColumnarExec] =
+    collectWithSubqueries(plan) { case c: CometSparkToColumnarExec => c }
+
+  private def cometShuffles(plan: SparkPlan): Seq[CometShuffleExchangeExec] =
+    collectWithSubqueries(plan) { case s: CometShuffleExchangeExec => s }
+
+  /** The native shuffles that read rows `CometSparkToColumnarExec` converted. 
*/
+  private def convertedShuffles(plan: SparkPlan): 
Seq[CometShuffleExchangeExec] =
+    cometShuffles(plan).filter { s =>
+      s.shuffleType == CometNativeShuffle && 
s.child.isInstanceOf[CometSparkToColumnarExec]
+    }
+
+  /**
+   * Spark inserts no columnar transitions below a `CometSparkToColumnarExec`, 
so the rule adds
+   * them for the operators under the conversion. Without one, an operator 
reads its Comet child
+   * through `CometExec.doExecute`, Spark's interpreted columnar-to-row path, 
which gives the
+   * right answer slowly, so only the plan shows it.
+   */
+  private def assertRowOperatorsReadThroughTransitions(plan: SparkPlan): Unit 
= {
+    val converted = conversions(plan)
+    assert(converted.nonEmpty, plan)
+    converted.foreach { conversion =>
+      // `InputAdapter` and `WholeStageCodegenExec` report their child's 
`supportsColumnar`.
+      val bare = conversion.child.collect {
+        case p
+            if !p.supportsColumnar && !p.isInstanceOf[ColumnarToRowTransition] 
&&
+              p.children.exists(_.supportsColumnar) =>
+          p
+      }
+      assert(bare.isEmpty, s"row operators read a columnar child without a 
transition:\n$plan")
+      assert(
+        conversion.child.collect { case c: ColumnarToRowExec => c }.isEmpty,
+        s"expected Comet's columnar-to-row transitions, not Spark's:\n$plan")
+    }
+    // The transitions are added on the first of the rule's passes under AQE. 
A later pass must
+    // not report them as operators Comet failed to convert.
+    val reasons = new ExtendedExplainInfo().getFallbackReasons(plan)
+    assert(!reasons.exists(_.contains("ColumnarToRow")), reasons)
+  }
+
+  convertTest("the shuffle of a Spark operator's rows runs as native shuffle") 
{
+    Seq("true", "false").foreach { aqe =>
+      withSQLConf(SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> aqe) {
+        val (_, plan) = checkSparkAnswer(
+          rowsDf()
+            .repartition(10, col("k"))
+            .groupBy("k")
+            .agg(sum("l"), sum(length(col("s"))), sum("d"), sum("m")))
+        assert(convertedShuffles(plan).length == 1, s"AQE $aqe:\n$plan")
+        assert(
+          !cometShuffles(plan).exists(_.shuffleType == CometColumnarShuffle),
+          s"AQE $aqe:\n$plan")
+      }
+    }
+  }
+
+  convertTest("Spark operators over Comet operators read them through 
transitions") {
+    withParquetTable((0 until 200).map(i => (i, (i % 13).toString)), "tbl") {
+      Seq("true", "false").foreach { aqe =>
+        withSQLConf(SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> aqe) {
+          val ds = spark.sql("SELECT _1 AS a, _2 AS b FROM 
tbl").as[ShuffleInputRec]
+          // The shuffle reads the typed operation's SerializeFromObject.
+          val (_, repartitioned) = checkSparkAnswer(
+            ds.map(r => ShuffleInputRec(r.a % 13, r.b))
+              .repartition(7, col("a"))
+              .groupBy("a")
+              .agg(count("b")))
+          assert(convertedShuffles(repartitioned).length == 1, s"AQE 
$aqe:\n$repartitioned")
+          assertRowOperatorsReadThroughTransitions(repartitioned)
+          // The shuffle reads a Spark partial aggregate over the typed 
operation. The final
+          // aggregate stays on Spark too, so the shuffle would go back to 
Spark's own.
+          
withSQLConf(CometConf.COMET_SHUFFLE_REVERT_REDUNDANT_COLUMNAR_ENABLED.key -> 
"false") {
+            val (_, aggregated) =
+              checkSparkAnswer(ds.map(r => ShuffleInputRec(r.a % 7, 
r.b)).groupBy("a").count())
+            assert(convertedShuffles(aggregated).length == 1, s"AQE 
$aqe:\n$aggregated")
+            assertRowOperatorsReadThroughTransitions(aggregated)
+          }
+        }
+      }
+    }
+  }
+
+  convertTest("aggregates whose partial aggregate runs on Spark") {
+    checkSparkAnswer(
+      rowsDf()
+        .groupBy((col("k") % 7).as("g"))
+        .agg(sum("l"), count("s"), avg("d"), sum("m"), max("l"), min("s")))
+  }
+
+  convertTest("a shuffle between two Spark aggregates goes back to Spark's 
shuffle") {
+    withSQLConf(CometConf.COMET_EXEC_AGGREGATE_ENABLED.key -> "false") {
+      val (_, plan) = checkSparkAnswer(rowsDf().groupBy("k").agg(sum("l")))
+      assert(cometShuffles(plan).isEmpty, plan)
+      
withSQLConf(CometConf.COMET_SHUFFLE_REVERT_REDUNDANT_COLUMNAR_ENABLED.key -> 
"false") {
+        val (_, kept) = checkSparkAnswer(rowsDf().groupBy("k").agg(sum("l")))
+        assert(convertedShuffles(kept).length == 1, kept)
+      }
+    }
+  }
+
+  convertTest("a join with an input that stays on the JVM columnar shuffle") {

Review Comment:
   Done in ddb428bb3. There's now a test for each key type the gate admits: 
boolean, the four integer types, float and double with NaN and both zeros, 
decimal(18, 2), date, timestamp, timestamp_ntz and binary, each with boundary 
values and NULL keys. Each test first compares the `spark_partition_id()` of 
every row after a `repartition` with Spark's, which covers the NULL keys that 
the join filters out, and then runs the join with one input on each shuffle. 
They share a `withShuffledJoins` helper with the wide decimal and string tests. 
To check that the placement comparison would notice a difference, I took out 
the wide decimal guard and added a `decimal(38, 0)` key, and it failed on the 
partition ids.
   
   The Celeborn case is in `CometCelebornShufflePlanningSuite` (21651aa1d), and 
fails if the gate stops excluding Celeborn. The test over several batches now 
has `array<string>` and `map<string,string>` columns with NULL elements, values 
and rows next to the struct. Without the stream read the struct is what breaks: 
a map or an array column on its own passes over several batches, NULLs 
included, as #6685 says.
   



-- 
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]

Reply via email to