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
7 changes: 7 additions & 0 deletions spark/src/main/scala/org/apache/comet/CometConf.scala
Original file line number Diff line number Diff line change
Expand Up @@ -237,6 +237,13 @@ object CometConf extends ShimCometConf {
createExecEnabledConfig("sortMergeJoin", defaultValue = true)
val COMET_EXEC_AGGREGATE_ENABLED: ConfigEntry[Boolean] =
createExecEnabledConfig("aggregate", defaultValue = true)
val COMET_EXEC_SORT_AGGREGATE_ENABLED: ConfigEntry[Boolean] =
createExecEnabledConfig(
"sortAggregate",
defaultValue = true,
notes = Some(
"When enabled, a SortAggregate that Comet can run is converted to a native hash " +
"aggregate, with a native sort restoring its output ordering"))
val COMET_EXEC_COLLECT_LIMIT_ENABLED: ConfigEntry[Boolean] =
createExecEnabledConfig("collectLimit", defaultValue = true)
val COMET_EXEC_COALESCE_ENABLED: ConfigEntry[Boolean] =
Expand Down
121 changes: 117 additions & 4 deletions spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala
Original file line number Diff line number Diff line change
Expand Up @@ -22,8 +22,8 @@ package org.apache.comet.rules
import scala.collection.mutable.ListBuffer

import org.apache.spark.sql.SparkSession
import org.apache.spark.sql.catalyst.expressions.{Divide, DoubleLiteral, EqualNullSafe, EqualTo, Expression, FloatLiteral, GreaterThan, GreaterThanOrEqual, KnownFloatingPointNormalized, LessThan, LessThanOrEqual, NamedExpression, Remainder}
import org.apache.spark.sql.catalyst.expressions.aggregate.{AggregateMode, Final, Partial, PartialMerge}
import org.apache.spark.sql.catalyst.expressions.{Divide, DoubleLiteral, EqualNullSafe, EqualTo, Expression, FloatLiteral, GreaterThan, GreaterThanOrEqual, KnownFloatingPointNormalized, LessThan, LessThanOrEqual, NamedExpression, Remainder, SortOrder}
import org.apache.spark.sql.catalyst.expressions.aggregate.{AggregateFunction, AggregateMode, Average, Count, Final, Max, Min, Partial, PartialMerge, Sum}
import org.apache.spark.sql.catalyst.optimizer.NormalizeNaNAndZero
import org.apache.spark.sql.catalyst.rules.Rule
import org.apache.spark.sql.catalyst.trees.TreeNodeTag
Expand All @@ -35,7 +35,7 @@ import org.apache.spark.sql.comet.shims.ShimCometEmptyRelation
import org.apache.spark.sql.comet.util.Utils
import org.apache.spark.sql.execution._
import org.apache.spark.sql.execution.adaptive.{AdaptiveSparkPlanExec, AQEShuffleReadExec, BroadcastQueryStageExec, LogicalQueryStage, QueryStageExec, ShuffleQueryStageExec}
import org.apache.spark.sql.execution.aggregate.{BaseAggregateExec, HashAggregateExec, ObjectHashAggregateExec}
import org.apache.spark.sql.execution.aggregate.{BaseAggregateExec, HashAggregateExec, ObjectHashAggregateExec, SortAggregateExec}
import org.apache.spark.sql.execution.columnar.InMemoryTableScanExec
import org.apache.spark.sql.execution.command.{DataWritingCommandExec, ExecutedCommandExec}
import org.apache.spark.sql.execution.datasources.{InsertIntoHadoopFsRelationCommand, WriteFilesExec}
Expand Down Expand Up @@ -150,6 +150,17 @@ object CometExecRule {
*/
val ENGINE_CHOICE_SPARK_TAG: TreeNodeTag[Unit] = TreeNodeTag[Unit]("comet.engineChoiceSpark")

/** Tag set on a native hash aggregate converted from a `SortAggregateExec`. */
val CONVERTED_SORT_AGGREGATE_TAG: TreeNodeTag[Unit] =
TreeNodeTag[Unit]("comet.convertedSortAggregate")

/**
* Tag set on the sort inserted above a converted `SortAggregateExec` for a consumer that
* requires the ordering the sort aggregate provided.
*/
val SORT_AGGREGATE_ORDER_TAG: TreeNodeTag[Unit] =
TreeNodeTag[Unit]("comet.sortAggregateOrder")

/**
* Serializes the native plan of each block of adjacent native operators into its topmost
* operator. Blocks that already hold a serialized plan are left as they are, so this can run
Expand Down Expand Up @@ -587,6 +598,9 @@ case class CometExecRule(session: SparkSession, wholePlan: Boolean = false)
convertToComet(s, CometShuffleExchangeExec)
.getOrElse(preserveSparkAggregateBuffers(s))

case agg: SortAggregateExec if CometConf.COMET_EXEC_SORT_AGGREGATE_ENABLED.get(conf) =>
convertSortAggregate(agg).getOrElse(agg)

case op =>
// if all children are native (or if this is a leaf node) then see if there is a
// registered handler for creating a fully native plan
Expand Down Expand Up @@ -644,7 +658,7 @@ case class CometExecRule(session: SparkSession, wholePlan: Boolean = false)
}

plan.transformUp { case op =>
val converted = convertNode(op)
val converted = convertNode(restoreSortAggregateOrdering(op))
// Replace SubqueryBroadcastExec with CometSubqueryBroadcastExec in DPP expressions
// when the broadcast child has a Comet plan underneath. This enables exchange reuse
// between the DPP subquery and the join's CometBroadcastExchangeExec because both
Expand Down Expand Up @@ -928,6 +942,105 @@ case class CometExecRule(session: SparkSession, wholePlan: Boolean = false)
}
}

private def leadsToConvertedSortAggregate(plan: SparkPlan): Boolean =
plan.getTagValue(CometExecRule.CONVERTED_SORT_AGGREGATE_TAG).isDefined ||
(plan.children.size == 1 && leadsToConvertedSortAggregate(plan.children.head))

/**
* Spark's EnsureRequirements left out the sort a consumer needs when a `SortAggregateExec`
* below it already provided the ordering. Once that aggregate is a native hash aggregate, the
* sort goes back on exactly the edges whose requirement is no longer met, before the consumer
* itself is converted.
*/
private def restoreSortAggregateOrdering(op: SparkPlan): SparkPlan = {
if (op.isInstanceOf[CometPlan]) return op
val required = op.requiredChildOrdering
if (required.length != op.children.length || required.forall(_.isEmpty)) return op
var changed = false
val children = op.children.zip(required).map { case (child, ordering) =>
if (ordering.nonEmpty && leadsToConvertedSortAggregate(child) &&
!SortOrder.orderingSatisfies(child.outputOrdering, ordering)) {
changed = true
val sort = SortExec(ordering, global = false, child)
child.logicalLink.foreach(sort.setLogicalLink)
val restored = if (child.isInstanceOf[CometNativeExec]) {
tryConvertToComet(sort, CometSortExec).getOrElse(sort)
} else {
sort
}
restored.setTagValue(CometExecRule.SORT_AGGREGATE_ORDER_TAG, ())
restored
} else {
child
}
}
if (changed) op.withNewChildren(children) else op
}

private def orderInsensitive(fn: AggregateFunction): Boolean = fn match {
case _: Min | _: Max | _: Count | _: Sum | _: Average => true
case _ => false
}

/**
* A `SortAggregateExec` as a native hash aggregate. Spark plans one when an aggregation buffer
* is not mutable, such as MIN or MAX over strings, and runs it over input sorted by the
* grouping keys. The hash aggregate starts from an `ObjectHashAggregateExec` with the same
* fields, which Spark runs correctly over any input order, so any operator reverted to Spark
* later stays correct. Only aggregates whose result does not depend on the input order are
* converted: Spark's sort is stable, so FIRST, LAST and the like see the rows of a group in
* their input order, which the native hash aggregate does not guarantee. The sort by the
* grouping keys directly below the aggregate is then dropped. A consumer that relied on the
* ordering of the sort aggregate gets a sort back, see [[restoreSortAggregateOrdering]].
*/
private def convertSortAggregate(agg: SortAggregateExec): Option[SparkPlan] = {
val required = agg.requiredChildOrdering.headOption.getOrElse(Nil)
val sortedByKeys = agg.child match {
case sort: CometSortExec =>
sort.sortOrder.length == required.length &&
sort.sortOrder.zip(required).forall { case (a, b) =>
a.semanticEquals(b)
} &&
(sort.originalPlan match {
case s: SortExec => !s.global
case _ => false
})
case _ => false
}
val input = if (sortedByKeys) agg.child.children.head else agg.child
if (!agg.aggregateExpressions.forall(e => orderInsensitive(e.aggregateFunction))) {
withFallbackReason(
agg,
"SortAggregate with an aggregate function that depends on the input order stays in Spark")
return None
}
if (!input.isInstanceOf[CometNativeExec]) {
return None
}
val hashAgg = ObjectHashAggregateExec(
agg.requiredChildDistributionExpressions,
agg.isStreaming,
agg.numShufflePartitions,
agg.groupingExpressions,
agg.aggregateExpressions,
agg.aggregateAttributes,
agg.initialInputBufferOffset,
agg.resultExpressions,
input)
hashAgg.copyTagsFrom(agg)
val converted = tryConvertToComet(hashAgg, CometObjectHashAggregateExec)
converted.foreach(_.setTagValue(CometExecRule.CONVERTED_SORT_AGGREGATE_TAG, ()))
if (converted.isEmpty) {
hashAgg
.getTagValue(CometExplainInfo.FALLBACK_REASONS)
.foreach(reasons => withFallbackReasons(agg, reasons))
if (!hasFallbackReason(agg)) {
withFallbackReason(agg, "SortAggregate could not be converted to a native hash aggregate")
}
}
converted
}

/** Convert a Spark plan to a Comet plan using the specified serde handler */
private def convertToComet(op: SparkPlan, handler: CometOperatorSerde[_]): Option[SparkPlan] = {
val converted = tryConvertToComet(op, handler)
Expand Down
10 changes: 8 additions & 2 deletions spark/src/main/scala/org/apache/comet/serde/aggregates.scala
Original file line number Diff line number Diff line change
Expand Up @@ -1395,7 +1395,7 @@ object CometMode extends CometAggregateExpressionSerde[Mode] with CometTypeShim
}
}

object AggSerde {
object AggSerde extends CometTypeShim {
import org.apache.spark.sql.types._

def minMaxDataTypeSupported(dt: DataType): Boolean = {
Expand Down Expand Up @@ -1436,7 +1436,13 @@ object AggSerde {

/** Shared support level for `Min` / `Max` based on the result data type. */
def minMaxSupportLevel(dt: DataType): SupportLevel = {
if (!minMaxDataTypeSupported(dt)) {
if (dt.isInstanceOf[StringType]) {
if (isStringCollationType(dt)) {
Unsupported(Some(s"Unsupported collated string type: $dt"))
} else {
Compatible()
}
} else if (!minMaxDataTypeSupported(dt)) {
Unsupported(Some(s"Unsupported data type: $dt"))
} else if ((dt == FloatType || dt == DoubleType) &&
COMET_EXEC_STRICT_FLOATING_POINT.get()) {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -183,7 +183,7 @@ SELECT first(val) IGNORE NULLS FROM test_single_null
-- first IGNORE NULLS: multiple data types
-- ============================================================

query expect_fallback(SortAggregate is not supported)
query expect_fallback(depends on the input order)
SELECT grp,
first(i_val) IGNORE NULLS,
first(l_val) IGNORE NULLS,
Expand Down Expand Up @@ -327,7 +327,7 @@ SELECT last(val) IGNORE NULLS FROM test_single_null
-- last IGNORE NULLS: multiple data types
-- ============================================================

query expect_fallback(SortAggregate is not supported)
query expect_fallback(depends on the input order)
SELECT grp,
last(i_val) IGNORE NULLS,
last(l_val) IGNORE NULLS,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -21,7 +21,7 @@ CREATE TABLE test_min_max(i int, d double, s string, grp string) USING parquet
statement
INSERT INTO test_min_max VALUES (1, 1.5, 'b', 'x'), (3, 3.5, 'a', 'x'), (2, 2.5, 'c', 'y'), (NULL, NULL, NULL, 'y'), (-1, -1.5, 'z', 'x')

query expect_fallback(SortAggregate is not supported)
query
SELECT min(i), max(i), min(d), max(d), min(s), max(s) FROM test_min_max

query
Expand Down
135 changes: 132 additions & 3 deletions spark/src/test/scala/org/apache/comet/exec/CometAggregateSuite.scala
Original file line number Diff line number Diff line change
Expand Up @@ -31,11 +31,11 @@ import org.apache.spark.sql.catalyst.expressions.Cast
import org.apache.spark.sql.catalyst.expressions.aggregate.{Final, Partial, PartialMerge}
import org.apache.spark.sql.catalyst.optimizer.EliminateSorts
import org.apache.spark.sql.catalyst.plans.physical.{HashPartitioning, RangePartitioning}
import org.apache.spark.sql.comet.{CometFilterExec, CometHashAggregateExec, CometNativeExec, CometProjectExec}
import org.apache.spark.sql.comet.{CometFilterExec, CometHashAggregateExec, CometNativeExec, CometProjectExec, CometSortExec, CometSortMergeJoinExec}
import org.apache.spark.sql.comet.execution.shuffle.CometShuffleExchangeExec
import org.apache.spark.sql.execution.SQLExecution
import org.apache.spark.sql.execution.{SparkPlan, SQLExecution}
import org.apache.spark.sql.execution.adaptive.{AdaptiveSparkPlanExec, AdaptiveSparkPlanHelper, ShuffleQueryStageExec}
import org.apache.spark.sql.execution.aggregate.{BaseAggregateExec, ObjectHashAggregateExec}
import org.apache.spark.sql.execution.aggregate.{BaseAggregateExec, ObjectHashAggregateExec, SortAggregateExec}
import org.apache.spark.sql.execution.exchange.{ReusedExchangeExec, ShuffleExchangeExec}
import org.apache.spark.sql.functions.{avg, col, count_distinct, expr, sum}
import org.apache.spark.sql.internal.SQLConf
Expand Down Expand Up @@ -3448,4 +3448,133 @@ class CometAggregateSuite extends CometTestBase with AdaptiveSparkPlanHelper {
}
}

private val stringRows: Seq[(Integer, String)] = Seq(
(1, "b"),
(1, ""),
(1, null),
(1, "a"),
(2, null),
(2, null),
(3, "\u00e9t\u00e9"),
(3, "ete"),
(3, "\u65e5\u672c"),
(4, "\uff61"),
(4, "\ud83d\ude00"),
(4, "z"),
(5, ""),
(null, "x"),
(null, "\u00ff"))

private def withStrings(f: => Unit): Unit =
withParquetTable(stringRows, "str_tbl")(f)

private def sortAggregates(plan: SparkPlan): Seq[SortAggregateExec] =
collect(plan) { case a: SortAggregateExec => a }

private def nativeAggregates(plan: SparkPlan): Seq[CometHashAggregateExec] =
collect(plan) { case a: CometHashAggregateExec => a }

for (aqe <- Seq("false", "true")) {
test(s"min and max over strings run natively, grouped and ungrouped (AQE=$aqe)") {
withSQLConf(SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> aqe) {
withStrings {
for (query <- Seq(
"SELECT min(_2), max(_2), count(_2) FROM str_tbl",
"SELECT _1, min(_2), max(_2) FROM str_tbl GROUP BY _1",
"SELECT _1, max(_2) FROM str_tbl WHERE _1 = 2 GROUP BY _1",
"SELECT max(_2) FROM str_tbl WHERE _1 > 100",
"SELECT _1, min(_2), sum(_1) FROM str_tbl GROUP BY _1")) {
val (sparkPlan, cometPlan) = checkSparkAnswer(sql(query))
assert(sortAggregates(sparkPlan).nonEmpty, s"$query:\n$sparkPlan")
assert(sortAggregates(cometPlan).isEmpty, s"$query:\n$cometPlan")
assert(nativeAggregates(cometPlan).nonEmpty, s"$query:\n$cometPlan")
}
}
}
}
}

test("min and max over strings compare UTF-8 bytes as Spark does") {
withStrings {
checkSparkAnswerAndOperator(sql("SELECT min(_2), max(_2) FROM str_tbl WHERE _1 = 4"))
checkSparkAnswerAndOperator(sql("SELECT _1, min(_2), max(_2) FROM str_tbl GROUP BY _1"))
}
}

test("a converted sort aggregate keeps the ordering a sort-merge join relies on") {
withSQLConf(
SQLConf.AUTO_BROADCASTJOIN_THRESHOLD.key -> "-1",
SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false") {
withStrings {
withParquetTable((0 until 50).map(i => (i % 6, s"v$i")), "str_right") {
val query =
"""SELECT a._1, a.m, r._2 FROM (SELECT _1, max(_2) AS m FROM str_tbl GROUP BY _1) a
|JOIN str_right r ON a._1 = r._1""".stripMargin
val (_, cometPlan) = checkSparkAnswer(sql(query))
assert(sortAggregates(cometPlan).isEmpty, s"plan:\n$cometPlan")
val join = collect(cometPlan) { case j: CometSortMergeJoinExec => j }
assert(join.size == 1, s"plan:\n$cometPlan")
val orderSorts = collect(cometPlan) {
case s: CometSortExec
if s.getTagValue(CometExecRule.SORT_AGGREGATE_ORDER_TAG).isDefined =>
s
}
assert(orderSorts.nonEmpty, s"plan:\n$cometPlan")
}
}
}
}

test("a converted sort aggregate gets no sort without a consumer that needs its ordering") {
withSQLConf(SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false") {
withStrings {
val (_, cometPlan) = checkSparkAnswer(sql("SELECT _1, max(_2) FROM str_tbl GROUP BY _1"))
val shuffles = collect(cometPlan) { case s: CometShuffleExchangeExec => s }
assert(shuffles.nonEmpty, s"plan:\n$cometPlan")
assert(shuffles.forall(s => !s.child.isInstanceOf[CometSortExec]), s"plan:\n$cometPlan")
assert(nativeAggregates(cometPlan).size == 2, s"plan:\n$cometPlan")
assert(
collect(cometPlan) {
case s: CometSortExec
if s.getTagValue(CometExecRule.SORT_AGGREGATE_ORDER_TAG).isDefined =>
s
}.isEmpty,
s"plan:\n$cometPlan")
assert(collect(cometPlan) { case s: CometSortExec => s }.isEmpty, s"plan:\n$cometPlan")
}
}
}

test("a sort aggregate Comet cannot run natively stays in Spark") {
withStrings {
val (_, cometPlan) =
checkSparkAnswer(sql("SELECT _1, max_by(_2, _2) FROM str_tbl GROUP BY _1"))
assert(sortAggregates(cometPlan).nonEmpty, s"plan:\n$cometPlan")
}
}

test("an order-sensitive sort aggregate over input sorted beyond its keys stays in Spark") {
withStrings {
val (_, cometPlan) = checkSparkAnswer(
sql("SELECT _1, first(_2) FROM (SELECT * FROM str_tbl DISTRIBUTE BY _1 SORT BY _1, _2) " +
"GROUP BY _1"))
assert(sortAggregates(cometPlan).nonEmpty, s"plan:\n$cometPlan")
}
}

test("a converted sort aggregate reverted by the cost-based choice runs as an object hash") {
withSQLConf(
CometConf.COMET_EXEC_COST_BASED_ENGINES_ENABLED.key -> "true",
CometConf.COMET_EXEC_COST_BASED_ENGINES_COST_TABLE.key -> "agg.flat.comet=1000000,0,0",
SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false") {
withStrings {
val (_, cometPlan) =
checkSparkAnswer(sql("SELECT _1, min(_2), max(_2) FROM str_tbl GROUP BY _1"))
assert(sortAggregates(cometPlan).isEmpty, s"plan:\n$cometPlan")
assert(
collect(cometPlan) { case a: ObjectHashAggregateExec => a }.nonEmpty,
s"plan:\n$cometPlan")
}
}
}
}
Loading
Loading