comphead commented on code in PR #6810:
URL: https://github.com/apache/datafusion-comet/pull/6810#discussion_r4232154480
##########
native/core/src/execution/planner.rs:
##########
@@ -1235,6 +1235,27 @@ impl PhysicalPlanner {
}
}
+ /// Returns a sort of `plan` by the SortOrder exprs in `ordering`, or
`None` when `ordering`
+ /// is empty or the plan's equivalence properties already satisfy it.
+ fn sort_unless_ordered(
Review Comment:
This is the same check and wrap that the sort aggregate already does for
`ordered_by_grouping_keys` (the `ordering_satisfy` block in the `HashAgg` arm,
around line 1712). Could `sort_unless_ordered` take the `PhysicalSortExpr`s, or
a `LexOrdering`, instead of Spark `Expr`s, with the join mapping
`create_sort_expr` at its call site? Then the aggregate arm could call it too,
and the planner would keep one place that decides when an output sort is needed.
##########
spark/src/test/scala/org/apache/comet/exec/CometJoinSuite.scala:
##########
@@ -1219,6 +1219,298 @@ class CometJoinSuite extends CometTestBase {
}
}
+ // A hash join reports its streamed side's ordering, so Spark drops a sort
above it when the
+ // streamed side already arrives partitioned and sorted. Here that side is a
cached table that
+ // reaches the native join through Comet's JVM to native conversion.
+ private def withSortedCache(data: DataFrame, order: Column*)(f: DataFrame =>
Unit): Unit = {
+ withSQLConf(
+ CometConf.COMET_CONVERT_FROM_IN_MEMORY_CACHE_ENABLED.key -> "true",
+ SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false",
+ SQLConf.AUTO_BROADCASTJOIN_THRESHOLD.key -> "-1",
+ SQLConf.SHUFFLE_PARTITIONS.key -> "2") {
+ val cached = data.repartition(2, data("_1")).sortWithinPartitions(order:
_*).cache()
+ try {
+ cached.count()
+ f(cached)
+ } finally {
+ cached.unpersist()
+ }
+ }
+ }
+
+ /**
+ * Checks the query's rows against Spark and that every partition is sorted
by `order`, which is
+ * bound to the query's output. Returns the plan's single native hash join
and its build side.
+ */
+ private def checkJoinKeepsOrder(
+ query: => DataFrame,
+ order: SortOrder*): (CometBinaryExec, BuildSide) = {
+ // withSQLConf returns Unit on Spark 3.x, so the expected rows are
assigned inside it.
+ var expected = Seq.empty[String]
+ withSQLConf(CometConf.COMET_ENABLED.key -> "false") {
+ expected = query.collect().map(_.toString).sorted.toSeq
+ }
+ val df = query
+ val rows = df.collect()
+ assert(rows.map(_.toString).sorted.toSeq == expected)
+ val plan = df.queryExecution.executedPlan
+ val joins = collect(plan) {
+ case j: CometHashJoinExec => (j, j.buildSide)
+ case j: CometBroadcastHashJoinExec => (j, j.buildSide)
+ }
+ assert(joins.length == 1, plan)
+ // Read before the plan runs again below, so the count covers the one run
above.
+ assert(joins.head._1.metrics("output_rows").value == rows.length, plan)
+ assert(joins.head._1.metrics("build_input_rows").value > 0, plan)
+ // Spark dropped the sort above the join, so the order comes from the join
alone.
+ val sortsAboveJoin = collect(plan) {
+ case s: SortExec => s
+ case s: CometSortExec => s
+ }.filter(s => find(s)(_ eq joins.head._1).isDefined)
+ assert(sortsAboveJoin.isEmpty, plan)
+ val ordering = new InterpretedOrdering(order)
+ // Propagates the session's confs, such as the batch size, to the tasks as
collect() does.
+ val parts = SQLExecution.withSQLConfPropagated(spark) {
+ df.queryExecution.toRdd.map(_.copy()).glom().collect()
+ }
+ parts.zipWithIndex.foreach { case (part, i) =>
+ val outOfOrder = part.zip(part.drop(1)).indexWhere { case (a, b) =>
ordering.gt(a, b) }
+ assert(outOfOrder < 0, s"partition $i is out of order at row
$outOfOrder\n$plan")
+ }
+ joins.head
+ }
+
+ private def keyOrder(
+ dataType: DataType,
+ direction: SortDirection,
+ ordinal: Int = 0): SortOrder =
+ SortOrder(BoundReference(ordinal, dataType, nullable = true), direction)
+
+ // (hint, join type, extra condition, AQE, native batch size)
+ private val streamedOrderCases =
+ (for (hint <- Seq("shuffle_hash", "broadcast"); joinType <-
Seq("left_outer", "right_outer"))
+ yield (hint, joinType, false, false, 8192)) ++ Seq(
+ ("shuffle_hash", "left_outer", true, false, 8192),
+ ("shuffle_hash", "right_outer", false, true, 8192),
+ ("shuffle_hash", "left_outer", false, false, 100))
+
+ for ((hint, joinType, withCondition, adaptive, batchSize) <-
streamedOrderCases) {
+ test(
+ s"$hint $joinType join keeps a cached streamed side's order, " +
+ s"condition=$withCondition, AQE=$adaptive, batchSize=$batchSize") {
+ withSQLConf(CometConf.COMET_BATCH_SIZE.key -> batchSize.toString) {
+ withParquetTable((0 until 10000).map(i => (i % 100, i)), "big") {
+ withParquetTable((0 until 10).map(i => (i * 10, i)), "small") {
+ withSortedCache(spark.table("big"), $"_1") { streamed =>
+ val build = spark.table("small").hint(hint)
+ val keys = streamed("_1") === build("_1")
+ val on = if (withCondition) keys && build("_2") > 3 else keys
+ def query: DataFrame = {
+ val joined =
+ if (joinType == "left_outer") streamed.join(build, on,
joinType)
+ else build.join(streamed, on, joinType)
+ joined
+ .select(streamed("_1").as("k"), streamed("_2").as("v"),
build("_2").as("w"))
+ .sortWithinPartitions("k")
+ }
+ withSQLConf(SQLConf.ADAPTIVE_EXECUTION_ENABLED.key ->
adaptive.toString) {
+ val (join, buildSide) =
+ checkJoinKeepsOrder(query, keyOrder(IntegerType, Ascending))
+ assert(buildSide == (if (joinType == "left_outer") BuildRight
else BuildLeft))
+ hint match {
+ case "broadcast" =>
assert(join.isInstanceOf[CometBroadcastHashJoinExec])
+ case _ => assert(join.isInstanceOf[CometHashJoinExec])
+ }
+ if (batchSize < 1000) {
+ // Several batches per partition come back to the JVM from
the native join.
+ val partitions = 2
+ val plan = query.queryExecution.executedPlan
+ val batches = collectFirst(plan) { case c:
ColumnarToRowTransition =>
+ SQLExecution.withSQLConfPropagated(spark) {
+ c.child.executeColumnar().count()
+ }
+ }
+ assert(batches.exists(_ > partitions), plan)
+ }
+ }
+ }
+ }
+ }
+ }
+ }
+ }
+
+ test("shuffle_hash join keeps a descending, nulls last streamed order with
null keys") {
+ val data = (0 until 10000).map(i => (if (i % 7 == 0) None else Some(i %
100), i))
+ withParquetTable(data, "big") {
+ withParquetTable((0 until 10).map(i => (i * 10, i)), "small") {
+ withSortedCache(spark.table("big"), $"_1".desc_nulls_last) { streamed
=>
+ val build = spark.table("small").hint("shuffle_hash")
+ checkJoinKeepsOrder(
+ streamed
+ .join(build, streamed("_1") === build("_1"), "left_outer")
+ .select(streamed("_1").as("k"), streamed("_2").as("v"),
build("_2").as("w"))
+ .sortWithinPartitions($"k".desc_nulls_last),
+ keyOrder(IntegerType, Descending))
+ }
+ }
+ }
+ }
+
+ test("shuffle_hash join keeps a descending streamed order on a double with
NaN and -0.0") {
+ def d(i: Int): Option[Double] = i % 9 match {
+ case 0 => None
+ case 1 => Some(Double.NaN)
+ case 2 => Some(-0.0)
+ case 3 => Some(0.0)
+ case _ => Some(i % 50 - 25.5)
+ }
+ withParquetTable((0 until 10000).map(i => (i % 100, i, d(i))), "big") {
+ withParquetTable((0 until 10).map(i => (i * 10, i)), "small") {
+ withSortedCache(spark.table("big"), $"_3".desc_nulls_last) { streamed
=>
+ val build = spark.table("small").hint("shuffle_hash")
+ checkJoinKeepsOrder(
+ streamed
+ .join(build, streamed("_1") === build("_1"), "left_outer")
+ .select(streamed("_3").as("d"), streamed("_2").as("v"),
build("_2").as("w"))
+ .sortWithinPartitions($"d".desc_nulls_last),
+ keyOrder(DoubleType, Descending))
+ }
+ }
+ }
+ }
+
+ test("shuffle_hash join keeps a two-key streamed order") {
+ withParquetTable((0 until 10000).map(i => (i % 100, i)), "big") {
+ withParquetTable((0 until 10).map(i => (i * 10, i)), "small") {
+ withSortedCache(spark.table("big"), $"_1", $"_2".desc) { streamed =>
+ val build = spark.table("small").hint("shuffle_hash")
+ checkJoinKeepsOrder(
+ streamed
+ .join(build, streamed("_1") === build("_1"), "left_outer")
+ .select(streamed("_1").as("k"), streamed("_2").as("v"),
build("_2").as("w"))
+ .sortWithinPartitions($"k", $"v".desc),
+ keyOrder(IntegerType, Ascending),
+ keyOrder(IntegerType, Descending, ordinal = 1))
+ }
+ }
+ }
+ }
+
+ // The streamed side is sorted inside the native plan, and the join's output
keeps that order in
+ // every partition. EliminateSorts would drop that sort as one below a join.
+ test("shuffle_hash join keeps a natively sorted streamed side's order") {
+ withSQLConf(
+ SQLConf.OPTIMIZER_EXCLUDED_RULES.key -> EliminateSorts.ruleName,
+ SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false",
+ SQLConf.AUTO_BROADCASTJOIN_THRESHOLD.key -> "-1",
+ SQLConf.SHUFFLE_PARTITIONS.key -> "2") {
+ withParquetTable((0 until 10000).map(i => (i % 100, i)), "big") {
+ withParquetTable((0 until 10).map(i => (i * 10, i)), "small") {
+ val streamed = spark.table("big").repartition(2,
$"_1").sortWithinPartitions("_1")
+ val build = spark.table("small").hint("shuffle_hash")
+ val (join, buildSide) = checkJoinKeepsOrder(
+ streamed
+ .join(build, streamed("_1") === build("_1"), "left_outer")
+ .select(streamed("_1").as("k"), streamed("_2").as("v"),
build("_2").as("w"))
+ .sortWithinPartitions("k"),
+ keyOrder(IntegerType, Ascending))
+ assert(join.isInstanceOf[CometHashJoinExec])
+ assert(buildSide == BuildRight)
+ }
+ }
+ }
+ }
+
+ // Spark plans NOT IN as a null-aware anti join that streams the left side.
Comet runs it with
+ // the left side as DataFusion's build side, which gives up that side's
batch order. Sparse keys
+ // keep DataFusion on its hash map build, which stores the batches last to
first.
+ test("null-aware anti join keeps a cached streamed side's order across
batches") {
+ withSQLConf(CometConf.COMET_BATCH_SIZE.key -> "100") {
+ withParquetTable((0 until 10000).map(i => (i % 100 * 1000, i)), "big") {
+ withParquetTable((0 until 10).map(i => (i * 10000, i)), "small") {
+ withSortedCache(spark.table("big"), $"_1") { streamed =>
+ withSQLConf(SQLConf.AUTO_BROADCASTJOIN_THRESHOLD.key ->
"10485760") {
+ withTempView("streamed") {
+ streamed.createOrReplaceTempView("streamed")
+ def query: DataFrame =
+ sql("SELECT _1 AS k, _2 AS v FROM streamed WHERE _1 NOT IN "
+
+ "(SELECT _1 FROM small)").sortWithinPartitions("k")
+ val (join, _) = checkJoinKeepsOrder(query,
keyOrder(IntegerType, Ascending))
+ join match {
+ case j: CometBroadcastHashJoinExec =>
assert(j.isNullAwareAntiJoin, j)
+ case j => fail(s"expected a null-aware broadcast hash join:
$j")
+ }
+ // The build side spans several batches per partition. Read on
a fresh plan that
+ // runs once, so the count covers that run alone.
+ val partitions = 2
+ val df = query
+ assert(df.collect().nonEmpty)
+ val plan = df.queryExecution.executedPlan
+ val buildBatches = collectFirst(plan) { case j:
CometBroadcastHashJoinExec =>
+ j.metrics("build_input_batches").value
+ }
+ assert(buildBatches.exists(_ > partitions), plan)
+ }
+ }
+ }
+ }
+ }
+ }
+ }
+
+ test("only outer hash joins that can move unmatched streamed rows send their
ordering") {
+ withParquetTable((0 until 10000).map(i => (i % 100, i)), "big") {
+ withParquetTable((0 until 10).map(i => (i * 10, i)), "small") {
+ withSortedCache(spark.table("big"), $"_1") { streamed =>
+ val build = spark.table("small").hint("shuffle_hash")
+ def query(joinType: String): DataFrame = streamed
+ .join(build, streamed("_1") === build("_1"), joinType)
+ .select(streamed("_1").as("k"), streamed("_2").as("v"),
build("_2").as("w"))
+ .sortWithinPartitions("k")
+ val (inner, _) = checkJoinKeepsOrder(query("inner"),
keyOrder(IntegerType, Ascending))
+ assert(inner.nativeOp.getHashJoin.getOutputOrderingCount == 0)
+ val (leftOuter, buildSide) =
+ checkJoinKeepsOrder(query("left_outer"), keyOrder(IntegerType,
Ascending))
+ assert(buildSide == BuildRight)
+ assert(leftOuter.nativeOp.getHashJoin.getOutputOrderingCount > 0)
Review Comment:
The `left_outer` half of this test runs the same tables, cache, query and
confs as the matrix case `shuffle_hash left_outer ... condition=false,
AQE=false, batchSize=8192` above, so it repeats that run. Could the
`getOutputOrderingCount > 0` assertion move into the matrix loop instead? That
would check the field for every affected shape, including `broadcast` and
`right_outer`, and this test could keep only the `inner` check.
--
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]