Skip to content
Closed
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
12 changes: 7 additions & 5 deletions spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
}

Expand All @@ -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)
Expand Down
27 changes: 14 additions & 13 deletions spark/src/main/scala/org/apache/comet/serde/aggregates.scala
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand Down
41 changes: 30 additions & 11 deletions spark/src/test/resources/sql-tests/expressions/aggregate/max_by.sql
Original file line number Diff line number Diff line change
Expand Up @@ -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).
Expand Down Expand Up @@ -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
Expand All @@ -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
Original file line number Diff line number Diff line change
Expand Up @@ -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).
Expand Down Expand Up @@ -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
Expand All @@ -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
Original file line number Diff line number Diff line change
Expand Up @@ -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")
}
}
}
Loading