diff --git a/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala b/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala index 7b33aad7b0b..15e0737df44 100644 --- a/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala +++ b/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala @@ -23,7 +23,7 @@ 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, 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.expressions.aggregate.{AggregateFunction, AggregateMode, Average, Count, Final, Max, MaxMinBy, 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 @@ -978,7 +978,7 @@ case class CometExecRule(session: SparkSession, wholePlan: Boolean = false) } private def orderInsensitive(fn: AggregateFunction): Boolean = fn match { - case _: Min | _: Max | _: Count | _: Sum | _: Average => true + case _: Min | _: Max | _: Count | _: Sum | _: Average | _: MaxMinBy => true case _ => false } @@ -989,9 +989,11 @@ case class CometExecRule(session: SparkSession, wholePlan: Boolean = false) * 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]]. + * their input order, which the native hash aggregate does not guarantee. MAX_BY and MIN_BY + * depend on it only among rows tied on the ordering, where they are non-deterministic in Spark + * too and already differ from it on the native hash aggregate path. 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) diff --git a/spark/src/main/scala/org/apache/comet/serde/aggregates.scala b/spark/src/main/scala/org/apache/comet/serde/aggregates.scala index bc45c605025..4583c2bce00 100644 --- a/spark/src/main/scala/org/apache/comet/serde/aggregates.scala +++ b/spark/src/main/scala/org/apache/comet/serde/aggregates.scala @@ -129,27 +129,28 @@ abstract class CometMaxMinBy[T <: MaxMinBy] extends CometAggregateExpressionSerd " Results may differ from Spark in that case.") override def getUnsupportedReasons(): Seq[String] = Seq( - "The value and ordering must both be fixed-length types (boolean, integral, floating-point," + - " decimal, date, or timestamp). A variable-length or nested type such as string, binary, or" + - " struct falls back to Spark.") + "The ordering must be a fixed-length type (boolean, integral, floating-point, decimal, date," + + " or timestamp). The value may also be a string with the default UTF8_BINARY collation." + + " Other variable-length or nested types such as binary or struct fall back to Spark.") override def getSupportLevel(expr: T): SupportLevel = { - // Both the value and ordering must be fixed-length types. + // The ordering must be a fixed-length type. The value may also be a UTF8_BINARY string: the + // native side only stores it as Arrow row bytes and never compares it. // - // On its own a variable-length type never reaches here: Spark only uses HashAggregate (the - // aggregate operator Comet accelerates) when the aggregation buffer is mutable, and the buffer - // holds both the running value and the running ordering, so a StringType in either position - // forces SortAggregate, which Comet does not convert. + // A string in the buffer makes Spark plan a SortAggregate, which Comet converts to a native + // hash aggregate (see `CometExecRule.convertSortAggregate`), or an ObjectHashAggregate when a + // TypedImperativeAggregate sits in the same aggregate. // - // The check is still load-bearing, because a TypedImperativeAggregate elsewhere in the same - // aggregate switches Spark to ObjectHashAggregate, which Comet does convert. In that shape a - // string ordering would otherwise be compared by Arrow's row format as raw UTF-8 bytes, while + // A string ordering stays in Spark: Arrow's row format would compare raw UTF-8 bytes, while // Spark compares collation sort keys. See the fallback cases in max_by.sql. // // The native side compares the ordering column via Arrow's row format, which supports all of // the fixed-length orderable types allowed below. - if (!AggSerde.minMaxDataTypeSupported(expr.valueExpr.dataType)) { - Unsupported(Some(s"Unsupported value data type: ${expr.valueExpr.dataType}")) + val valueType = expr.valueExpr.dataType + val stringValue = + valueType.isInstanceOf[StringType] && !AggSerde.isStringCollationType(valueType) + if (!stringValue && !AggSerde.minMaxDataTypeSupported(valueType)) { + Unsupported(Some(s"Unsupported value data type: $valueType")) } else if (!AggSerde.minMaxDataTypeSupported(expr.orderingExpr.dataType)) { Unsupported(Some(s"Unsupported ordering data type: ${expr.orderingExpr.dataType}")) } else { diff --git a/spark/src/test/resources/sql-tests/expressions/aggregate/max_by.sql b/spark/src/test/resources/sql-tests/expressions/aggregate/max_by.sql index 9c8354f06ab..b8270db9e1a 100644 --- a/spark/src/test/resources/sql-tests/expressions/aggregate/max_by.sql +++ b/spark/src/test/resources/sql-tests/expressions/aggregate/max_by.sql @@ -17,9 +17,8 @@ -- max_by(x, y) returns the value of x associated with the maximum value of y. -- --- The value (x) must be a fixed-length type: Spark only uses HashAggregate (the aggregate --- operator Comet accelerates) when the aggregation buffer is mutable, so variable-length --- value types such as string force SortAggregate and fall back to Spark. +-- The value (x) must be a fixed-length type or a string. A string value makes Spark plan a +-- SortAggregate, which Comet converts to a native hash aggregate. -- -- Ordering values are kept unique within each group so results are deterministic (max_by is -- non-deterministic when several rows tie on the maximum ordering). @@ -260,14 +259,11 @@ query SELECT grp, max_by(v, ord) FROM mb_signed_zero GROUP BY grp ORDER BY grp -- ============================================================ --- Variable-length value or ordering falls back to Spark +-- Variable-length ordering falls back to Spark -- --- A plain max_by over a string is planned as SortAggregate, which Comet never converts, so the --- serde's type check is not what stops it. Pairing it with a TypedImperativeAggregate switches --- Spark to ObjectHashAggregate, which Comet does convert, and then getSupportLevel is the only --- thing that keeps the aggregate off the native path. That matters most for a string *ordering*: --- Arrow's row format compares raw UTF-8 bytes, while Spark compares collation sort keys, so this --- check is load-bearing for correctness and not just an optimisation. +-- A string value runs natively: as a plain aggregate through the converted SortAggregate, and +-- next to a TypedImperativeAggregate through ObjectHashAggregate. A string ordering stays in +-- Spark: Arrow's row format compares raw UTF-8 bytes, while Spark compares collation sort keys. -- ============================================================ statement @@ -277,8 +273,31 @@ statement INSERT INTO mb_varlen VALUES (1, 'a', 'g1'), (2, 'b', 'g1'), (3, 'c', 'g2') -query expect_fallback(Unsupported value data type) +query SELECT grp, max_by(s, v), percentile(v, 0.5) FROM mb_varlen GROUP BY grp ORDER BY grp query expect_fallback(Unsupported ordering data type) SELECT grp, max_by(v, s), percentile(v, 0.5) FROM mb_varlen GROUP BY grp ORDER BY grp + +-- ============================================================ +-- String value with an integer ordering +-- ============================================================ + +statement +CREATE TABLE mb_str(s string, ord int, grp string) USING parquet + +statement +INSERT INTO mb_str VALUES + ('b', 1, 'g1'), ('', 5, 'g1'), (NULL, 3, 'g1'), + ('é日本', 2, 'g2'), (NULL, 9, 'g2'), ('😀', 4, 'g2'), + ('z', NULL, 'g3'), ('y', NULL, 'g3'), + ('x', 6, 'g4') + +query +SELECT grp, max_by(s, ord) FROM mb_str GROUP BY grp ORDER BY grp + +query +SELECT max_by(s, ord) FROM mb_str + +query +SELECT grp, max_by(s, ord), max(s), count(*) FROM mb_str GROUP BY grp ORDER BY grp diff --git a/spark/src/test/resources/sql-tests/expressions/aggregate/min_by.sql b/spark/src/test/resources/sql-tests/expressions/aggregate/min_by.sql index d015acc35d3..55828e5c65c 100644 --- a/spark/src/test/resources/sql-tests/expressions/aggregate/min_by.sql +++ b/spark/src/test/resources/sql-tests/expressions/aggregate/min_by.sql @@ -17,9 +17,8 @@ -- min_by(x, y) returns the value of x associated with the minimum value of y. -- --- The value (x) must be a fixed-length type: Spark only uses HashAggregate (the aggregate --- operator Comet accelerates) when the aggregation buffer is mutable, so variable-length --- value types such as string force SortAggregate and fall back to Spark. +-- The value (x) must be a fixed-length type or a string. A string value makes Spark plan a +-- SortAggregate, which Comet converts to a native hash aggregate. -- -- Ordering values are kept unique within each group so results are deterministic (min_by is -- non-deterministic when several rows tie on the minimum ordering). @@ -245,11 +244,11 @@ query SELECT grp, min_by(v, ord) FROM mnb_signed_zero GROUP BY grp ORDER BY grp -- ============================================================ --- Variable-length value or ordering falls back to Spark +-- Variable-length ordering falls back to Spark -- --- See the equivalent section in max_by.sql for why this needs a TypedImperativeAggregate --- alongside it: without one, Spark plans SortAggregate and Comet never sees the aggregate at all, --- so the serde's type check would go untested. +-- A string value runs natively: as a plain aggregate through the converted SortAggregate, and +-- next to a TypedImperativeAggregate through ObjectHashAggregate. A string ordering stays in +-- Spark: Arrow's row format compares raw UTF-8 bytes, while Spark compares collation sort keys. -- ============================================================ statement @@ -259,8 +258,31 @@ statement INSERT INTO mnb_varlen VALUES (1, 'a', 'g1'), (2, 'b', 'g1'), (3, 'c', 'g2') -query expect_fallback(Unsupported value data type) +query SELECT grp, min_by(s, v), percentile(v, 0.5) FROM mnb_varlen GROUP BY grp ORDER BY grp query expect_fallback(Unsupported ordering data type) SELECT grp, min_by(v, s), percentile(v, 0.5) FROM mnb_varlen GROUP BY grp ORDER BY grp + +-- ============================================================ +-- String value with an integer ordering +-- ============================================================ + +statement +CREATE TABLE mnb_str(s string, ord int, grp string) USING parquet + +statement +INSERT INTO mnb_str VALUES + ('b', 1, 'g1'), ('', 5, 'g1'), (NULL, 3, 'g1'), + ('é日本', 2, 'g2'), (NULL, 9, 'g2'), ('😀', 4, 'g2'), + ('z', NULL, 'g3'), ('y', NULL, 'g3'), + ('x', 6, 'g4') + +query +SELECT grp, min_by(s, ord) FROM mnb_str GROUP BY grp ORDER BY grp + +query +SELECT min_by(s, ord) FROM mnb_str + +query +SELECT grp, min_by(s, ord), max(s), count(*) FROM mnb_str GROUP BY grp ORDER BY grp diff --git a/spark/src/test/scala/org/apache/comet/exec/CometAggregateSuite.scala b/spark/src/test/scala/org/apache/comet/exec/CometAggregateSuite.scala index 298d2a73eb5..04b3746a5d1 100644 --- a/spark/src/test/scala/org/apache/comet/exec/CometAggregateSuite.scala +++ b/spark/src/test/scala/org/apache/comet/exec/CometAggregateSuite.scala @@ -3577,4 +3577,69 @@ class CometAggregateSuite extends CometTestBase with AdaptiveSparkPlanHelper { } } } + + private val byRows: Seq[(Integer, String, Integer)] = Seq( + (1, "b", 1), + (1, "", 5), + (1, null, 3), + (2, "\u00e9\u65e5\u672c", 2), + (2, null, 9), + (2, "\ud83d\ude00", 4), + (3, "z", null), + (3, "y", null), + (4, "x", 6), + (null, "w", 7)) + + for (aqe <- Seq("false", "true"); fn <- Seq("max_by", "min_by")) { + test(s"$fn over a string value runs natively, grouped and ungrouped (AQE=$aqe)") { + withSQLConf(SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> aqe) { + withParquetTable(byRows, "by_tbl") { + for (query <- Seq( + s"SELECT _1, $fn(_2, _3) FROM by_tbl GROUP BY _1", + s"SELECT $fn(_2, _3) FROM by_tbl", + s"SELECT _1, $fn(_2, _3), max(_2), count(*) FROM by_tbl GROUP BY _1", + s"SELECT $fn(_2, _3) FROM by_tbl WHERE _1 = 3")) { + 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("max_by and min_by over a string value pick one of the values tied on the ordering") { + val rows = (0 until 400).map(i => (i % 4, s"v${i % 7}", if (i % 3 == 0) 10 else i % 3)) + withSQLConf(SQLConf.SHUFFLE_PARTITIONS.key -> "3") { + withParquetTable(rows, "by_ties") { + for ((fn, tiedOrd) <- Seq("max_by" -> 10, "min_by" -> 1)) { + val df = sql(s"SELECT _1, $fn(_2, _3) FROM by_ties GROUP BY _1") + val result = df.collect().map(r => r.getInt(0) -> r.getString(1)).toMap + assert(nativeAggregates(df.queryExecution.executedPlan).nonEmpty) + val allowed = rows.filter(_._3 == tiedOrd).groupBy(_._1).map { case (k, v) => + k -> v.map(_._2).toSet + } + assert(result.keySet == allowed.keySet, s"$fn: $result") + result.foreach { case (k, v) => assert(allowed(k).contains(v), s"$fn group $k: $v") } + } + } + } + } + + test("max_by over a string value keeps a native partial with a Spark final") { + withSQLConf(CometConf.COMET_ENABLE_FINAL_HASH_AGGREGATE.key -> "false") { + withParquetTable(byRows, "by_tbl") { + checkSparkAnswer(sql("SELECT _1, max_by(_2, _3), min_by(_2, _3) FROM by_tbl GROUP BY _1")) + } + } + } + + test("an ungrouped order-sensitive sort aggregate over an ordered input stays in Spark") { + withParquetTable(byRows, "by_tbl") { + val (_, cometPlan) = checkSparkAnswer( + sql("SELECT max_by(_2, _1), first(_2) FROM (SELECT * FROM by_tbl ORDER BY _3, _2)")) + assert(sortAggregates(cometPlan).nonEmpty, s"plan:\n$cometPlan") + } + } }