Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -65,6 +65,15 @@ case class CometColumnarToRowExec(child: SparkPlan)

override def outputOrdering: Seq[SortOrder] = child.outputOrdering

override def supportCodegen: Boolean = !readsBroadcast(child)

private def readsBroadcast(plan: SparkPlan): Boolean = plan match {
case _: CometBroadcastExchangeExec => true
case stage: BroadcastQueryStageExec => readsBroadcast(stage.plan)
case reused: ReusedExchangeExec => readsBroadcast(reused.child)
case _ => false
}

// `ColumnarToRowExec` processes the input RDD directly, which is kind of a leaf node in the
// codegen stage and needs to do the limit check.
protected override def canCheckLimitNotReached: Boolean = true
Expand Down
57 changes: 54 additions & 3 deletions spark/src/test/scala/org/apache/comet/exec/CometJoinSuite.scala
Original file line number Diff line number Diff line change
Expand Up @@ -31,11 +31,12 @@ import org.apache.spark.sql.catalyst.TableIdentifier
import org.apache.spark.sql.catalyst.analysis.UnresolvedRelation
import org.apache.spark.sql.catalyst.expressions.{And, AttributeReference, IsNotNull}
import org.apache.spark.sql.catalyst.optimizer.{BuildLeft, BuildRight}
import org.apache.spark.sql.comet.{CometBroadcastExchangeExec, CometBroadcastHashJoinExec, CometBroadcastNestedLoopJoinExec, CometFilterExec, CometHashJoinExec, CometNativeScanExec, CometSortMergeJoinExec, CometUnionExec}
import org.apache.spark.sql.execution.{LocalTableScanExec, SparkPlan}
import org.apache.spark.sql.execution.adaptive.{AdaptiveSparkPlanExec, AQEShuffleReadExec}
import org.apache.spark.sql.comet.{CometBroadcastExchangeExec, CometBroadcastHashJoinExec, CometBroadcastNestedLoopJoinExec, CometColumnarToRowExec, CometFilterExec, CometHashJoinExec, CometNativeScanExec, CometSortMergeJoinExec, CometUnionExec}
import org.apache.spark.sql.execution.{InputAdapter, LocalTableScanExec, SparkPlan, WholeStageCodegenExec}
import org.apache.spark.sql.execution.adaptive.{AdaptiveSparkPlanExec, AQEShuffleReadExec, BroadcastQueryStageExec}
import org.apache.spark.sql.execution.datasources.SchemaColumnConvertNotSupportedException
import org.apache.spark.sql.execution.exchange.ReusedExchangeExec
import org.apache.spark.sql.execution.joins.{BroadcastHashJoinExec, SortMergeJoinExec}
import org.apache.spark.sql.internal.SQLConf
import org.apache.spark.sql.types.{ArrayType, IntegerType, MetadataBuilder, StructField, StructType}

Expand Down Expand Up @@ -1479,6 +1480,56 @@ class CometJoinSuite extends CometTestBase {
}
}

for (costBased <- Seq("true", "false"); maxFields <- Seq("3", "100")) {
test(
"a Spark broadcast hash join over a skewed join reads a Comet broadcast stage " +
s"(costBasedEngines=$costBased, codegen.maxFields=$maxFields)") {
withSQLConf(
SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "true",
SQLConf.AUTO_BROADCASTJOIN_THRESHOLD.key -> "-1",
SQLConf.ADAPTIVE_AUTO_BROADCASTJOIN_THRESHOLD.key -> "-1",
SQLConf.COALESCE_PARTITIONS_ENABLED.key -> "false",
SQLConf.SHUFFLE_PARTITIONS.key -> "4",
SQLConf.SKEW_JOIN_ENABLED.key -> "true",
SQLConf.ADAPTIVE_FORCE_OPTIMIZE_SKEWED_JOIN.key -> "true",
SQLConf.SKEW_JOIN_SKEWED_PARTITION_FACTOR.key -> "2",
SQLConf.SKEW_JOIN_SKEWED_PARTITION_THRESHOLD.key -> "100",
SQLConf.ADVISORY_PARTITION_SIZE_IN_BYTES.key -> "100",
SQLConf.WHOLESTAGE_MAX_NUM_FIELDS.key -> maxFields,
CometConf.COMET_EXEC_COST_BASED_ENGINES_ENABLED.key -> costBased) {
withParquetTable(
(0 until 4000).map(i => (if (i % 10 < 8) 0 else i % 50, i % 10, i)),
"skew_f") {
withParquetTable((0 until 500).map(i => (i % 50, i)), "skew_g") {
withParquetTable((0 until 10).map(i => (i, s"name_$i")), "skew_d") {
val query =
"""SELECT /*+ BROADCAST(d), MERGE(f, g) */ f._1, f._2, f._3, g._2, d._2
|FROM skew_f f JOIN skew_g g ON f._1 = g._1
|LEFT JOIN skew_d d ON f._2 = d._1
|DISTRIBUTE BY f._3""".stripMargin
val (_, plan) = checkSparkAnswer(sql(query))
val skewed = collect(plan) { case j: SortMergeJoinExec if j.isSkewJoin => j }
assert(skewed.nonEmpty, s"expected a skewed Spark sort-merge join:\n$plan")
val sparkBroadcastJoins = collect(plan) { case j: BroadcastHashJoinExec => j }
assert(sparkBroadcastJoins.size == 1, s"plan:\n$plan")
val build = sparkBroadcastJoins.head.right
val readsCometBroadcast = build.collectFirst { case c: CometColumnarToRowExec =>
c.child match {
case InputAdapter(stage) => stage
case stage => stage
}
}
assert(
readsCometBroadcast.exists(_.isInstanceOf[BroadcastQueryStageExec]),
s"expected the join to read a Comet broadcast stage:\n$plan")
assert(!build.isInstanceOf[WholeStageCodegenExec], s"plan:\n$plan")
}
}
}
}
}
}

test("BroadcastNestedLoopJoin with unequal filter") {
withSQLConf(SQLConf.AUTO_BROADCASTJOIN_THRESHOLD.key -> "-1") {
// Include NULL keys: predicate `_1 > _1` returns NULL with a NULL operand, so
Expand Down
Loading