This is an automated email from the ASF dual-hosted git repository.
andygrove pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/datafusion-comet.git
The following commit(s) were added to refs/heads/main by this push:
new 6456f335c5 feat: enable mixed partial/final execution for sum and
non-decimal avg (#4861)
6456f335c5 is described below
commit 6456f335c5b913313865d30acd755cb1cb19d69c
Author: Andy Grove <[email protected]>
AuthorDate: Wed Jul 8 16:20:18 2026 -0600
feat: enable mixed partial/final execution for sum and non-decimal avg
(#4861)
---
dev/diffs/3.4.3.diff | 33 ++++-
dev/diffs/3.5.8.diff | 35 +++++-
dev/diffs/4.0.2.diff | 35 +++++-
dev/diffs/4.1.2.diff | 35 +++++-
.../serde/CometAggregateExpressionSerde.scala | 13 +-
.../org/apache/comet/serde/QueryPlanSerde.scala | 2 +-
.../scala/org/apache/comet/serde/aggregates.scala | 25 +++-
.../approved-plans-v1_4-spark3_5/q70/extended.txt | 8 +-
.../approved-plans-v1_4/q35/extended.txt | 2 +-
.../approved-plans-v1_4/q45/extended.txt | 89 ++++++-------
.../approved-plans-v2_7-spark3_5/q70a/extended.txt | 32 ++---
.../approved-plans-v2_7/q35/extended.txt | 2 +-
.../apache/comet/exec/CometAggregateSuite.scala | 26 ++++
.../apache/comet/rules/CometExecRuleSuite.scala | 140 ++++++++++++++++++++-
14 files changed, 371 insertions(+), 106 deletions(-)
diff --git a/dev/diffs/3.4.3.diff b/dev/diffs/3.4.3.diff
index 6faf9dab78..487c5d1482 100644
--- a/dev/diffs/3.4.3.diff
+++ b/dev/diffs/3.4.3.diff
@@ -1,5 +1,5 @@
diff --git a/pom.xml b/pom.xml
-index d3544881af1..aae0ae3b27b 100644
+index d3544881af1..ff963395ec3 100644
--- a/pom.xml
+++ b/pom.xml
@@ -148,6 +148,8 @@
@@ -260,19 +260,32 @@ index cf40e944c09..bdd5be4f462 100644
test("A cached table preserves the partitioning and ordering of its cached
SparkPlan") {
diff --git
a/sql/core/src/test/scala/org/apache/spark/sql/DataFrameAggregateSuite.scala
b/sql/core/src/test/scala/org/apache/spark/sql/DataFrameAggregateSuite.scala
-index 1cc09c3d7fc..f031fa45c33 100644
+index 1cc09c3d7fc..a84939c045b 100644
---
a/sql/core/src/test/scala/org/apache/spark/sql/DataFrameAggregateSuite.scala
+++
b/sql/core/src/test/scala/org/apache/spark/sql/DataFrameAggregateSuite.scala
-@@ -27,7 +27,7 @@ import org.apache.spark.SparkException
+@@ -26,8 +26,9 @@ import org.scalatest.matchers.must.Matchers.the
+ import org.apache.spark.SparkException
import org.apache.spark.sql.execution.WholeStageCodegenExec
import org.apache.spark.sql.execution.adaptive.AdaptiveSparkPlanHelper
++import org.apache.spark.sql.comet.CometHashAggregateExec
import org.apache.spark.sql.execution.aggregate.{HashAggregateExec,
ObjectHashAggregateExec, SortAggregateExec}
-import org.apache.spark.sql.execution.exchange.ShuffleExchangeExec
+import org.apache.spark.sql.execution.exchange.ShuffleExchangeLike
import org.apache.spark.sql.expressions.Window
import org.apache.spark.sql.functions._
import org.apache.spark.sql.internal.SQLConf
-@@ -755,7 +755,7 @@ class DataFrameAggregateSuite extends QueryTest
+@@ -691,7 +692,9 @@ class DataFrameAggregateSuite extends QueryTest
+ case _ => false
+ }.isDefined)
+ } else {
+- assert(stripAQEPlan(hashAggPlan).isInstanceOf[HashAggregateExec])
++ val strippedPlan = stripAQEPlan(hashAggPlan)
++ assert(strippedPlan.isInstanceOf[HashAggregateExec] ||
++ strippedPlan.exists(_.isInstanceOf[CometHashAggregateExec]))
+ }
+
+ // test case for ObjectHashAggregate and SortAggregate
+@@ -755,7 +758,7 @@ class DataFrameAggregateSuite extends QueryTest
assert(objHashAggPlans.nonEmpty)
val exchangePlans = collect(aggPlan) {
@@ -1532,7 +1545,7 @@ index ac710c32296..2854b433dd3 100644
import testImplicits._
diff --git
a/sql/core/src/test/scala/org/apache/spark/sql/execution/adaptive/AdaptiveQueryExecSuite.scala
b/sql/core/src/test/scala/org/apache/spark/sql/execution/adaptive/AdaptiveQueryExecSuite.scala
-index 593bd7bb4ba..32af28b0238 100644
+index 593bd7bb4ba..b327d84d5cc 100644
---
a/sql/core/src/test/scala/org/apache/spark/sql/execution/adaptive/AdaptiveQueryExecSuite.scala
+++
b/sql/core/src/test/scala/org/apache/spark/sql/execution/adaptive/AdaptiveQueryExecSuite.scala
@@ -26,9 +26,11 @@ import org.scalatest.time.SpanSugar._
@@ -1923,6 +1936,16 @@ index 593bd7bb4ba..32af28b0238 100644
}
assert(shuffles2.size == 4)
val smj2 = findTopLevelSortMergeJoin(adaptive2)
+@@ -2703,7 +2744,8 @@ class AdaptiveQueryExecSuite
+ }
+ }
+
+-test("SPARK-44040: Fix compute stats when AggregateExec nodes above
QueryStageExec") {
++test("SPARK-44040: Fix compute stats when AggregateExec nodes above
QueryStageExec",
++ IgnoreComet("https://github.com/apache/datafusion-comet/issues/4412")) {
+ val emptyDf = spark.range(1).where("false")
+ val aggDf1 = emptyDf.agg(sum("id").as("id")).withColumn("name",
lit("df1"))
+ val aggDf2 = emptyDf.agg(sum("id").as("id")).withColumn("name",
lit("df2"))
diff --git
a/sql/core/src/test/scala/org/apache/spark/sql/execution/datasources/SchemaPruningSuite.scala
b/sql/core/src/test/scala/org/apache/spark/sql/execution/datasources/SchemaPruningSuite.scala
index bd9c79e5b96..2ada8c28842 100644
---
a/sql/core/src/test/scala/org/apache/spark/sql/execution/datasources/SchemaPruningSuite.scala
diff --git a/dev/diffs/3.5.8.diff b/dev/diffs/3.5.8.diff
index 7155e9a9b4..fe87568b96 100644
--- a/dev/diffs/3.5.8.diff
+++ b/dev/diffs/3.5.8.diff
@@ -1,5 +1,5 @@
diff --git a/pom.xml b/pom.xml
-index edd2ad57880..a47b7dec672 100644
+index edd2ad57880..45f8fd01538 100644
--- a/pom.xml
+++ b/pom.xml
@@ -152,6 +152,8 @@
@@ -241,19 +241,32 @@ index e5494726695..00937f025c2 100644
test("A cached table preserves the partitioning and ordering of its cached
SparkPlan") {
diff --git
a/sql/core/src/test/scala/org/apache/spark/sql/DataFrameAggregateSuite.scala
b/sql/core/src/test/scala/org/apache/spark/sql/DataFrameAggregateSuite.scala
-index 6f3090d8908..c08a60fb0c2 100644
+index 6f3090d8908..a0e9309888d 100644
---
a/sql/core/src/test/scala/org/apache/spark/sql/DataFrameAggregateSuite.scala
+++
b/sql/core/src/test/scala/org/apache/spark/sql/DataFrameAggregateSuite.scala
-@@ -28,7 +28,7 @@ import org.apache.spark.sql.catalyst.plans.logical.Expand
+@@ -27,8 +27,9 @@ import org.apache.spark.{SparkException, SparkThrowable}
+ import org.apache.spark.sql.catalyst.plans.logical.Expand
import org.apache.spark.sql.execution.WholeStageCodegenExec
import org.apache.spark.sql.execution.adaptive.AdaptiveSparkPlanHelper
++import org.apache.spark.sql.comet.CometHashAggregateExec
import org.apache.spark.sql.execution.aggregate.{HashAggregateExec,
ObjectHashAggregateExec, SortAggregateExec}
-import org.apache.spark.sql.execution.exchange.ShuffleExchangeExec
+import org.apache.spark.sql.execution.exchange.ShuffleExchangeLike
import org.apache.spark.sql.expressions.Window
import org.apache.spark.sql.functions._
import org.apache.spark.sql.internal.SQLConf
-@@ -793,7 +793,7 @@ class DataFrameAggregateSuite extends QueryTest
+@@ -729,7 +730,9 @@ class DataFrameAggregateSuite extends QueryTest
+ case _ => false
+ }.isDefined)
+ } else {
+- assert(stripAQEPlan(hashAggPlan).isInstanceOf[HashAggregateExec])
++ val strippedPlan = stripAQEPlan(hashAggPlan)
++ assert(strippedPlan.isInstanceOf[HashAggregateExec] ||
++ strippedPlan.exists(_.isInstanceOf[CometHashAggregateExec]))
+ }
+
+ // test case for ObjectHashAggregate and SortAggregate
+@@ -793,7 +796,7 @@ class DataFrameAggregateSuite extends QueryTest
assert(objHashAggPlans.nonEmpty)
val exchangePlans = collect(aggPlan) {
@@ -1498,7 +1511,7 @@ index 5a413c77754..207b66e1d7b 100644
import testImplicits._
diff --git
a/sql/core/src/test/scala/org/apache/spark/sql/execution/adaptive/AdaptiveQueryExecSuite.scala
b/sql/core/src/test/scala/org/apache/spark/sql/execution/adaptive/AdaptiveQueryExecSuite.scala
-index 2f8e401e743..dbcf3171946 100644
+index 2f8e401e743..7849c685b19 100644
---
a/sql/core/src/test/scala/org/apache/spark/sql/execution/adaptive/AdaptiveQueryExecSuite.scala
+++
b/sql/core/src/test/scala/org/apache/spark/sql/execution/adaptive/AdaptiveQueryExecSuite.scala
@@ -27,9 +27,11 @@ import org.scalatest.time.SpanSugar._
@@ -1904,7 +1917,17 @@ index 2f8e401e743..dbcf3171946 100644
}.size == (if (firstAccess) 2 else 0))
assert(collect(initialExecutedPlan) {
case i: InMemoryTableScanLike => i
-@@ -2980,7 +3023,9 @@ class AdaptiveQueryExecSuite
+@@ -2938,7 +2981,8 @@ class AdaptiveQueryExecSuite
+ }
+ }
+
+- test("SPARK-44040: Fix compute stats when AggregateExec nodes above
QueryStageExec") {
++ test("SPARK-44040: Fix compute stats when AggregateExec nodes above
QueryStageExec",
++ IgnoreComet("https://github.com/apache/datafusion-comet/issues/4412")) {
+ val emptyDf = spark.range(1).where("false")
+ val aggDf1 = emptyDf.agg(sum("id").as("id")).withColumn("name",
lit("df1"))
+ val aggDf2 = emptyDf.agg(sum("id").as("id")).withColumn("name",
lit("df2"))
+@@ -2980,7 +3024,9 @@ class AdaptiveQueryExecSuite
val plan =
df.queryExecution.executedPlan.asInstanceOf[AdaptiveSparkPlanExec]
assert(plan.inputPlan.isInstanceOf[TakeOrderedAndProjectExec])
diff --git a/dev/diffs/4.0.2.diff b/dev/diffs/4.0.2.diff
index 28cac668c8..d8e506b631 100644
--- a/dev/diffs/4.0.2.diff
+++ b/dev/diffs/4.0.2.diff
@@ -39,7 +39,7 @@ index 6c51bd4ff2e..e72ec1d26e2 100644
withSpark(sc) { sc =>
TestUtils.waitUntilExecutorsUp(sc, 2, 60000)
diff --git a/pom.xml b/pom.xml
-index 252cfdf9073..60cb9dcb7cf 100644
+index 252cfdf9073..50ec9d6314e 100644
--- a/pom.xml
+++ b/pom.xml
@@ -148,6 +148,8 @@
@@ -378,19 +378,32 @@ index 0f42502f1d9..e9ff802141f 100644
withTempView("t0", "t1", "t2") {
diff --git
a/sql/core/src/test/scala/org/apache/spark/sql/DataFrameAggregateSuite.scala
b/sql/core/src/test/scala/org/apache/spark/sql/DataFrameAggregateSuite.scala
-index 9db406ff12f..245e4caa319 100644
+index 9db406ff12f..19b4bb8e39d 100644
---
a/sql/core/src/test/scala/org/apache/spark/sql/DataFrameAggregateSuite.scala
+++
b/sql/core/src/test/scala/org/apache/spark/sql/DataFrameAggregateSuite.scala
-@@ -30,7 +30,7 @@ import org.apache.spark.sql.errors.DataTypeErrors.toSQLId
+@@ -29,8 +29,9 @@ import
org.apache.spark.sql.catalyst.util.AUTO_GENERATED_ALIAS
+ import org.apache.spark.sql.errors.DataTypeErrors.toSQLId
import org.apache.spark.sql.execution.WholeStageCodegenExec
import org.apache.spark.sql.execution.adaptive.AdaptiveSparkPlanHelper
++import org.apache.spark.sql.comet.CometHashAggregateExec
import org.apache.spark.sql.execution.aggregate.{HashAggregateExec,
ObjectHashAggregateExec, SortAggregateExec}
-import org.apache.spark.sql.execution.exchange.ShuffleExchangeExec
+import org.apache.spark.sql.execution.exchange.ShuffleExchangeLike
import org.apache.spark.sql.expressions.Window
import org.apache.spark.sql.functions._
import org.apache.spark.sql.internal.SQLConf
-@@ -855,7 +855,7 @@ class DataFrameAggregateSuite extends QueryTest
+@@ -791,7 +792,9 @@ class DataFrameAggregateSuite extends QueryTest
+ case _ => false
+ }.isDefined)
+ } else {
+- assert(stripAQEPlan(hashAggPlan).isInstanceOf[HashAggregateExec])
++ val strippedPlan = stripAQEPlan(hashAggPlan)
++ assert(strippedPlan.isInstanceOf[HashAggregateExec] ||
++ strippedPlan.exists(_.isInstanceOf[CometHashAggregateExec]))
+ }
+
+ // test case for ObjectHashAggregate and SortAggregate
+@@ -855,7 +858,7 @@ class DataFrameAggregateSuite extends QueryTest
assert(objHashAggPlans.nonEmpty)
val exchangePlans = collect(aggPlan) {
@@ -2122,7 +2135,7 @@ index a3cfdc5a240..3793b6191bf 100644
})
checkAnswer(distinctWithId, Seq(Row(1, 0), Row(1, 0)))
diff --git
a/sql/core/src/test/scala/org/apache/spark/sql/execution/adaptive/AdaptiveQueryExecSuite.scala
b/sql/core/src/test/scala/org/apache/spark/sql/execution/adaptive/AdaptiveQueryExecSuite.scala
-index 272be70f9fe..12daa1f5932 100644
+index 272be70f9fe..4a175083adf 100644
---
a/sql/core/src/test/scala/org/apache/spark/sql/execution/adaptive/AdaptiveQueryExecSuite.scala
+++
b/sql/core/src/test/scala/org/apache/spark/sql/execution/adaptive/AdaptiveQueryExecSuite.scala
@@ -28,12 +28,14 @@ import org.apache.spark.SparkException
@@ -2539,7 +2552,17 @@ index 272be70f9fe..12daa1f5932 100644
}.isEmpty)
assert(collect(initialExecutedPlan) {
case i: InMemoryTableScanLike => i
-@@ -3129,7 +3173,8 @@ class AdaptiveQueryExecSuite
+@@ -3039,7 +3083,8 @@ class AdaptiveQueryExecSuite
+ }
+ }
+
+- test("SPARK-44040: Fix compute stats when AggregateExec nodes above
QueryStageExec") {
++ test("SPARK-44040: Fix compute stats when AggregateExec nodes above
QueryStageExec",
++ IgnoreComet("https://github.com/apache/datafusion-comet/issues/4412")) {
+ val emptyDf = spark.range(1).where("false")
+ val aggDf1 = emptyDf.agg(sum("id").as("id")).withColumn("name",
lit("df1"))
+ val aggDf2 = emptyDf.agg(sum("id").as("id")).withColumn("name",
lit("df2"))
+@@ -3129,7 +3174,8 @@ class AdaptiveQueryExecSuite
val plan =
df.queryExecution.executedPlan.asInstanceOf[AdaptiveSparkPlanExec]
assert(plan.inputPlan.isInstanceOf[TakeOrderedAndProjectExec])
diff --git a/dev/diffs/4.1.2.diff b/dev/diffs/4.1.2.diff
index 946dae3949..9a1b30b010 100644
--- a/dev/diffs/4.1.2.diff
+++ b/dev/diffs/4.1.2.diff
@@ -39,7 +39,7 @@ index 6df8bc85b51..dabb75e2b75 100644
withSpark(sc) { sc =>
TestUtils.waitUntilExecutorsUp(sc, 2, 60000)
diff --git a/pom.xml b/pom.xml
-index dc201151999..d5c08f11ded 100644
+index dc201151999..20ee0e7482a 100644
--- a/pom.xml
+++ b/pom.xml
@@ -152,6 +152,8 @@
@@ -392,19 +392,32 @@ index 0d807aeae4d..6d7744e771b 100644
withTempView("t0", "t1", "t2") {
diff --git
a/sql/core/src/test/scala/org/apache/spark/sql/DataFrameAggregateSuite.scala
b/sql/core/src/test/scala/org/apache/spark/sql/DataFrameAggregateSuite.scala
-index bfe15b33768..55c23a38ccc 100644
+index bfe15b33768..31f8ba4b876 100644
---
a/sql/core/src/test/scala/org/apache/spark/sql/DataFrameAggregateSuite.scala
+++
b/sql/core/src/test/scala/org/apache/spark/sql/DataFrameAggregateSuite.scala
-@@ -31,7 +31,7 @@ import org.apache.spark.sql.errors.DataTypeErrors.toSQLId
+@@ -30,8 +30,9 @@ import
org.apache.spark.sql.catalyst.util.AUTO_GENERATED_ALIAS
+ import org.apache.spark.sql.errors.DataTypeErrors.toSQLId
import org.apache.spark.sql.execution.WholeStageCodegenExec
import org.apache.spark.sql.execution.adaptive.AdaptiveSparkPlanHelper
++import org.apache.spark.sql.comet.CometHashAggregateExec
import org.apache.spark.sql.execution.aggregate.{HashAggregateExec,
ObjectHashAggregateExec, SortAggregateExec}
-import org.apache.spark.sql.execution.exchange.ShuffleExchangeExec
+import org.apache.spark.sql.execution.exchange.ShuffleExchangeLike
import org.apache.spark.sql.expressions.Window
import org.apache.spark.sql.functions._
import org.apache.spark.sql.internal.SQLConf
-@@ -856,7 +856,7 @@ class DataFrameAggregateSuite extends QueryTest
+@@ -792,7 +793,9 @@ class DataFrameAggregateSuite extends QueryTest
+ case _ => false
+ }.isDefined)
+ } else {
+- assert(stripAQEPlan(hashAggPlan).isInstanceOf[HashAggregateExec])
++ val strippedPlan = stripAQEPlan(hashAggPlan)
++ assert(strippedPlan.isInstanceOf[HashAggregateExec] ||
++ strippedPlan.exists(_.isInstanceOf[CometHashAggregateExec]))
+ }
+
+ // test case for ObjectHashAggregate and SortAggregate
+@@ -856,7 +859,7 @@ class DataFrameAggregateSuite extends QueryTest
assert(objHashAggPlans.nonEmpty)
val exchangePlans = collect(aggPlan) {
@@ -2245,7 +2258,7 @@ index a3cfdc5a240..3793b6191bf 100644
})
checkAnswer(distinctWithId, Seq(Row(1, 0), Row(1, 0)))
diff --git
a/sql/core/src/test/scala/org/apache/spark/sql/execution/adaptive/AdaptiveQueryExecSuite.scala
b/sql/core/src/test/scala/org/apache/spark/sql/execution/adaptive/AdaptiveQueryExecSuite.scala
-index 3e7d26f74bd..7e70e72fa3e 100644
+index 3e7d26f74bd..79232dc3664 100644
---
a/sql/core/src/test/scala/org/apache/spark/sql/execution/adaptive/AdaptiveQueryExecSuite.scala
+++
b/sql/core/src/test/scala/org/apache/spark/sql/execution/adaptive/AdaptiveQueryExecSuite.scala
@@ -27,12 +27,14 @@ import org.apache.spark.SparkException
@@ -2662,7 +2675,17 @@ index 3e7d26f74bd..7e70e72fa3e 100644
}.isEmpty)
assert(collect(initialExecutedPlan) {
case i: InMemoryTableScanLike => i
-@@ -3310,7 +3354,8 @@ class AdaptiveQueryExecSuite
+@@ -3220,7 +3264,8 @@ class AdaptiveQueryExecSuite
+ }
+ }
+
+- test("SPARK-44040: Fix compute stats when AggregateExec nodes above
QueryStageExec") {
++ test("SPARK-44040: Fix compute stats when AggregateExec nodes above
QueryStageExec",
++ IgnoreComet("https://github.com/apache/datafusion-comet/issues/4412")) {
+ val emptyDf = spark.range(1).where("false")
+ val aggDf1 = emptyDf.agg(sum("id").as("id")).withColumn("name",
lit("df1"))
+ val aggDf2 = emptyDf.agg(sum("id").as("id")).withColumn("name",
lit("df2"))
+@@ -3310,7 +3355,8 @@ class AdaptiveQueryExecSuite
val plan =
df.queryExecution.executedPlan.asInstanceOf[AdaptiveSparkPlanExec]
assert(plan.inputPlan.isInstanceOf[TakeOrderedAndProjectExec])
diff --git
a/spark/src/main/scala/org/apache/comet/serde/CometAggregateExpressionSerde.scala
b/spark/src/main/scala/org/apache/comet/serde/CometAggregateExpressionSerde.scala
index 9a83152168..a52d600821 100644
---
a/spark/src/main/scala/org/apache/comet/serde/CometAggregateExpressionSerde.scala
+++
b/spark/src/main/scala/org/apache/comet/serde/CometAggregateExpressionSerde.scala
@@ -82,15 +82,16 @@ trait CometAggregateExpressionSerde[T <: AggregateFunction]
{
def getSupportLevel(expr: T): SupportLevel = Compatible(None)
/**
- * Whether this aggregate's intermediate buffer format is compatible between
Spark and Comet,
- * making it safe to run the Partial in one engine and the Final in the
other. Aggregates with
- * simple single-value buffers (MIN, MAX, bitwise) are safe; those with
complex or
- * differently-encoded buffers (AVG, SUM with decimals, CollectSet,
Variance) are not. COUNT is
- * intentionally excluded: mixed COUNT partial/final regressed AQE's
+ * Whether this aggregate's intermediate buffer format is compatible between
Spark and Comet for
+ * the given function instance, making it safe to run the Partial in one
engine and the Final in
+ * the other. Aggregates with simple single-value buffers (MIN, MAX,
bitwise) are always safe;
+ * SUM and non-decimal AVG match Spark's buffer and are safe except where
noted per instance
+ * (e.g. TRY-mode SUM uses a Comet-internal flag column). COUNT is
intentionally excluded
+ * despite a matching buffer: mixed COUNT partial/final regressed AQE's
* PropagateEmptyRelationAfterAQE pattern (which matches BaseAggregateExec
only) and the Spark
* 4.0 count-bug decorrelation for correlated IN subqueries.
*/
- def supportsMixedPartialFinal: Boolean = false
+ def supportsMixedPartialFinal(fn: T): Boolean = false
/**
* Convert a Spark expression into a protocol buffer representation that can
be passed into
diff --git a/spark/src/main/scala/org/apache/comet/serde/QueryPlanSerde.scala
b/spark/src/main/scala/org/apache/comet/serde/QueryPlanSerde.scala
index 1e7aacd402..82874a7f19 100644
--- a/spark/src/main/scala/org/apache/comet/serde/QueryPlanSerde.scala
+++ b/spark/src/main/scala/org/apache/comet/serde/QueryPlanSerde.scala
@@ -433,7 +433,7 @@ object QueryPlanSerde extends Logging with CometExprShim
with CometTypeShim {
case Some(handler) =>
handler
.asInstanceOf[CometAggregateExpressionSerde[AggregateFunction]]
- .supportsMixedPartialFinal
+ .supportsMixedPartialFinal(fn)
case None => false
}
}
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 5710232cb4..510efd3c9f 100644
--- a/spark/src/main/scala/org/apache/comet/serde/aggregates.scala
+++ b/spark/src/main/scala/org/apache/comet/serde/aggregates.scala
@@ -28,12 +28,13 @@ import org.apache.spark.sql.types.{ByteType, DecimalType,
DoubleType, IntegerTyp
import org.apache.comet.CometConf.COMET_EXEC_STRICT_FLOATING_POINT
import org.apache.comet.CometSparkSessionExtensions.{isSpark41Plus,
withFallbackReason}
+import org.apache.comet.expressions.CometEvalMode
import org.apache.comet.serde.QueryPlanSerde.{evalModeToProto, exprToProto,
serializeDataType}
import org.apache.comet.shims.CometEvalModeUtil
object CometMin extends CometAggregateExpressionSerde[Min] {
- override def supportsMixedPartialFinal: Boolean = true
+ override def supportsMixedPartialFinal(fn: Min): Boolean = true
override def getSupportLevel(expr: Min): SupportLevel =
AggSerde.minMaxSupportLevel(expr.dataType)
@@ -70,7 +71,7 @@ object CometMin extends CometAggregateExpressionSerde[Min] {
object CometMax extends CometAggregateExpressionSerde[Max] {
- override def supportsMixedPartialFinal: Boolean = true
+ override def supportsMixedPartialFinal(fn: Max): Boolean = true
override def getSupportLevel(expr: Max): SupportLevel =
AggSerde.minMaxSupportLevel(expr.dataType)
@@ -130,6 +131,11 @@ object CometCount extends
CometAggregateExpressionSerde[Count] {
object CometAverage extends CometAggregateExpressionSerde[Average] {
+ override def supportsMixedPartialFinal(fn: Average): Boolean =
+ // Non-decimal AVG has a (sum: double, count: long) buffer matching Spark.
Decimal AVG is
+ // deferred (overflow nulls count differently) and stays unsafe for mixed
execution.
+ !fn.child.dataType.isInstanceOf[DecimalType]
+
override def getUnsupportedReasons(): Seq[String] = Seq(
"YearMonthIntervalType and DayTimeIntervalType inputs are not supported")
@@ -186,6 +192,13 @@ object CometAverage extends
CometAggregateExpressionSerde[Average] {
object CometSum extends CometAggregateExpressionSerde[Sum] {
+ override def supportsMixedPartialFinal(fn: Sum): Boolean =
+ // Decimal SUM is excluded: overflow detection (ANSI throw / Legacy null)
does not survive a
+ // Spark-partial / Comet-final split, so the required ArithmeticException
is never raised.
+ // TRY-mode integer SUM carries a Comet-internal has_all_nulls column that
Spark cannot read.
+ !fn.child.dataType.isInstanceOf[DecimalType] &&
+ CometEvalModeUtil.fromSparkEvalMode(CometEvalModeUtil.sumEvalMode(fn))
!= CometEvalMode.TRY
+
override def getSupportLevel(expr: Sum): SupportLevel =
if (AggSerde.sumDataTypeSupported(expr.dataType)) {
Compatible()
@@ -300,7 +313,7 @@ object CometLast extends
CometAggregateExpressionSerde[Last] {
}
object CometBitAndAgg extends CometAggregateExpressionSerde[BitAndAgg] {
- override def supportsMixedPartialFinal: Boolean = true
+ override def supportsMixedPartialFinal(fn: BitAndAgg): Boolean = true
override def getSupportLevel(expr: BitAndAgg): SupportLevel =
if (AggSerde.bitwiseAggTypeSupported(expr.dataType)) {
@@ -339,7 +352,7 @@ object CometBitAndAgg extends
CometAggregateExpressionSerde[BitAndAgg] {
}
object CometBitOrAgg extends CometAggregateExpressionSerde[BitOrAgg] {
- override def supportsMixedPartialFinal: Boolean = true
+ override def supportsMixedPartialFinal(fn: BitOrAgg): Boolean = true
override def getSupportLevel(expr: BitOrAgg): SupportLevel =
if (AggSerde.bitwiseAggTypeSupported(expr.dataType)) {
@@ -378,7 +391,7 @@ object CometBitOrAgg extends
CometAggregateExpressionSerde[BitOrAgg] {
}
object CometBitXOrAgg extends CometAggregateExpressionSerde[BitXorAgg] {
- override def supportsMixedPartialFinal: Boolean = true
+ override def supportsMixedPartialFinal(fn: BitXorAgg): Boolean = true
override def getSupportLevel(expr: BitXorAgg): SupportLevel =
if (AggSerde.bitwiseAggTypeSupported(expr.dataType)) {
@@ -707,7 +720,7 @@ object CometCorr extends
CometAggregateExpressionSerde[Corr] {
object CometBloomFilterAggregate extends
CometAggregateExpressionSerde[BloomFilterAggregate] {
- override def supportsMixedPartialFinal: Boolean = true
+ override def supportsMixedPartialFinal(fn: BloomFilterAggregate): Boolean =
true
override def getSupportLevel(expr: BloomFilterAggregate): SupportLevel =
expr.child.dataType match {
diff --git
a/spark/src/test/resources/tpcds-plan-stability/approved-plans-v1_4-spark3_5/q70/extended.txt
b/spark/src/test/resources/tpcds-plan-stability/approved-plans-v1_4-spark3_5/q70/extended.txt
index 08df1c04a0..f9fb64d9e4 100644
---
a/spark/src/test/resources/tpcds-plan-stability/approved-plans-v1_4-spark3_5/q70/extended.txt
+++
b/spark/src/test/resources/tpcds-plan-stability/approved-plans-v1_4-spark3_5/q70/extended.txt
@@ -3,9 +3,9 @@ CometNativeColumnarToRow
+- CometProject
+- CometWindowExec
+- CometSort
- +- CometColumnarExchange
- +- HashAggregate [COMET: Spark Final aggregate without Comet
Partial requires compatible intermediate buffer formats, but the following
aggregate function(s) have incompatible buffers: sum]
- +- Exchange
+ +- CometExchange
+ +- CometHashAggregate
+ +- CometColumnarExchange
+- HashAggregate
+- Expand
+- Project
@@ -56,4 +56,4 @@ CometNativeColumnarToRow
+- CometFilter
+- CometNativeScan parquet spark_catalog.default.date_dim
-Comet accelerated 37 out of 53 eligible operators (69%). Final plan contains 4
transitions between Spark and Comet.
\ No newline at end of file
+Comet accelerated 39 out of 53 eligible operators (73%). Final plan contains 4
transitions between Spark and Comet.
\ No newline at end of file
diff --git
a/spark/src/test/resources/tpcds-plan-stability/approved-plans-v1_4/q35/extended.txt
b/spark/src/test/resources/tpcds-plan-stability/approved-plans-v1_4/q35/extended.txt
index 50e9ebb965..fcb5e0003d 100644
---
a/spark/src/test/resources/tpcds-plan-stability/approved-plans-v1_4/q35/extended.txt
+++
b/spark/src/test/resources/tpcds-plan-stability/approved-plans-v1_4/q35/extended.txt
@@ -1,5 +1,5 @@
TakeOrderedAndProject
-+- HashAggregate [COMET: Spark Final aggregate without Comet Partial requires
compatible intermediate buffer formats, but the following aggregate function(s)
have incompatible buffers: avg, count]
++- HashAggregate [COMET: Spark Final aggregate without Comet Partial requires
compatible intermediate buffer formats, but the following aggregate function(s)
have incompatible buffers: count]
+- Exchange
+- HashAggregate
+- Project
diff --git
a/spark/src/test/resources/tpcds-plan-stability/approved-plans-v1_4/q45/extended.txt
b/spark/src/test/resources/tpcds-plan-stability/approved-plans-v1_4/q45/extended.txt
index 2d463ef986..6f73e28782 100644
---
a/spark/src/test/resources/tpcds-plan-stability/approved-plans-v1_4/q45/extended.txt
+++
b/spark/src/test/resources/tpcds-plan-stability/approved-plans-v1_4/q45/extended.txt
@@ -1,45 +1,46 @@
-TakeOrderedAndProject
-+- HashAggregate [COMET: Spark Final aggregate without Comet Partial requires
compatible intermediate buffer formats, but the following aggregate function(s)
have incompatible buffers: sum]
- +- Exchange
- +- HashAggregate
- +- Project
- +- Filter
- +- BroadcastHashJoin [COMET: Unsupported join type
ExistenceJoin(exists#1)]
- :- CometNativeColumnarToRow
- : +- CometProject
- : +- CometBroadcastHashJoin
- : :- CometProject
- : : +- CometBroadcastHashJoin
- : : :- CometProject
- : : : +- CometBroadcastHashJoin
- : : : :- CometProject
- : : : : +- CometBroadcastHashJoin
- : : : : :- CometFilter
- : : : : : +- CometNativeScan parquet
spark_catalog.default.web_sales
- : : : : : +- CometSubqueryBroadcast
- : : : : : +-
CometBroadcastExchange
- : : : : : +- CometProject
- : : : : : +- CometFilter
- : : : : : +-
CometNativeScan parquet spark_catalog.default.date_dim
- : : : : +- CometBroadcastExchange
- : : : : +- CometFilter
- : : : : +- CometNativeScan parquet
spark_catalog.default.customer
- : : : +- CometBroadcastExchange
- : : : +- CometProject
- : : : +- CometFilter
- : : : +- CometNativeScan parquet
spark_catalog.default.customer_address
- : : +- CometBroadcastExchange
- : : +- CometProject
- : : +- CometFilter
- : : +- CometNativeScan parquet
spark_catalog.default.date_dim
- : +- CometBroadcastExchange
- : +- CometProject
- : +- CometFilter
- : +- CometNativeScan parquet
spark_catalog.default.item
- +- BroadcastExchange
- +- CometNativeColumnarToRow
- +- CometProject
- +- CometFilter
- +- CometNativeScan parquet
spark_catalog.default.item
+CometNativeColumnarToRow
++- CometTakeOrderedAndProject
+ +- CometHashAggregate
+ +- CometColumnarExchange
+ +- HashAggregate
+ +- Project
+ +- Filter
+ +- BroadcastHashJoin [COMET: Unsupported join type
ExistenceJoin(exists#1)]
+ :- CometNativeColumnarToRow
+ : +- CometProject
+ : +- CometBroadcastHashJoin
+ : :- CometProject
+ : : +- CometBroadcastHashJoin
+ : : :- CometProject
+ : : : +- CometBroadcastHashJoin
+ : : : :- CometProject
+ : : : : +- CometBroadcastHashJoin
+ : : : : :- CometFilter
+ : : : : : +- CometNativeScan parquet
spark_catalog.default.web_sales
+ : : : : : +-
CometSubqueryBroadcast
+ : : : : : +-
CometBroadcastExchange
+ : : : : : +- CometProject
+ : : : : : +-
CometFilter
+ : : : : : +-
CometNativeScan parquet spark_catalog.default.date_dim
+ : : : : +- CometBroadcastExchange
+ : : : : +- CometFilter
+ : : : : +- CometNativeScan
parquet spark_catalog.default.customer
+ : : : +- CometBroadcastExchange
+ : : : +- CometProject
+ : : : +- CometFilter
+ : : : +- CometNativeScan parquet
spark_catalog.default.customer_address
+ : : +- CometBroadcastExchange
+ : : +- CometProject
+ : : +- CometFilter
+ : : +- CometNativeScan parquet
spark_catalog.default.date_dim
+ : +- CometBroadcastExchange
+ : +- CometProject
+ : +- CometFilter
+ : +- CometNativeScan parquet
spark_catalog.default.item
+ +- BroadcastExchange
+ +- CometNativeColumnarToRow
+ +- CometProject
+ +- CometFilter
+ +- CometNativeScan parquet
spark_catalog.default.item
-Comet accelerated 32 out of 41 eligible operators (78%). Final plan contains 2
transitions between Spark and Comet.
\ No newline at end of file
+Comet accelerated 35 out of 41 eligible operators (85%). Final plan contains 3
transitions between Spark and Comet.
\ No newline at end of file
diff --git
a/spark/src/test/resources/tpcds-plan-stability/approved-plans-v2_7-spark3_5/q70a/extended.txt
b/spark/src/test/resources/tpcds-plan-stability/approved-plans-v2_7-spark3_5/q70a/extended.txt
index e527d95f46..2aac2972c3 100644
---
a/spark/src/test/resources/tpcds-plan-stability/approved-plans-v2_7-spark3_5/q70a/extended.txt
+++
b/spark/src/test/resources/tpcds-plan-stability/approved-plans-v2_7-spark3_5/q70a/extended.txt
@@ -5,11 +5,11 @@ CometNativeColumnarToRow
+- CometSort
+- CometExchange
+- CometHashAggregate
- +- CometColumnarExchange
- +- HashAggregate
- +- Union
- :- HashAggregate [COMET: Spark Final aggregate
without Comet Partial requires compatible intermediate buffer formats, but the
following aggregate function(s) have incompatible buffers: sum]
- : +- Exchange
+ +- CometExchange
+ +- CometHashAggregate
+ +- CometUnion
+ :- CometHashAggregate
+ : +- CometColumnarExchange
: +- HashAggregate
: +- Project
: +- BroadcastHashJoin
@@ -58,11 +58,11 @@ CometNativeColumnarToRow
:
+- CometProject
:
+- CometFilter
:
+- CometNativeScan parquet spark_catalog.default.date_dim
- :- HashAggregate [COMET: Spark Final aggregate
without Comet Partial requires compatible intermediate buffer formats, but the
following aggregate function(s) have incompatible buffers: sum]
- : +- Exchange
- : +- HashAggregate
- : +- HashAggregate [COMET: Spark Final
aggregate without Comet Partial requires compatible intermediate buffer
formats, but the following aggregate function(s) have incompatible buffers: sum]
- : +- Exchange
+ :- CometHashAggregate
+ : +- CometExchange
+ : +- CometHashAggregate
+ : +- CometHashAggregate
+ : +- CometColumnarExchange
: +- HashAggregate
: +- Project
: +- BroadcastHashJoin
@@ -111,11 +111,11 @@ CometNativeColumnarToRow
:
+- CometProject
:
+- CometFilter
:
+- CometNativeScan parquet
spark_catalog.default.date_dim
- +- HashAggregate [COMET: Spark Final aggregate
without Comet Partial requires compatible intermediate buffer formats, but the
following aggregate function(s) have incompatible buffers: sum]
- +- Exchange
- +- HashAggregate
- +- HashAggregate [COMET: Spark Final
aggregate without Comet Partial requires compatible intermediate buffer
formats, but the following aggregate function(s) have incompatible buffers: sum]
- +- Exchange
+ +- CometHashAggregate
+ +- CometExchange
+ +- CometHashAggregate
+ +- CometHashAggregate
+ +- CometColumnarExchange
+- HashAggregate
+- Project
+- BroadcastHashJoin
@@ -165,4 +165,4 @@ CometNativeColumnarToRow
+- CometFilter
+- CometNativeScan parquet
spark_catalog.default.date_dim
-Comet accelerated 103 out of 156 eligible operators (66%). Final plan contains
10 transitions between Spark and Comet.
\ No newline at end of file
+Comet accelerated 117 out of 156 eligible operators (75%). Final plan contains
10 transitions between Spark and Comet.
\ No newline at end of file
diff --git
a/spark/src/test/resources/tpcds-plan-stability/approved-plans-v2_7/q35/extended.txt
b/spark/src/test/resources/tpcds-plan-stability/approved-plans-v2_7/q35/extended.txt
index c340d4a42c..fcb5e0003d 100644
---
a/spark/src/test/resources/tpcds-plan-stability/approved-plans-v2_7/q35/extended.txt
+++
b/spark/src/test/resources/tpcds-plan-stability/approved-plans-v2_7/q35/extended.txt
@@ -1,5 +1,5 @@
TakeOrderedAndProject
-+- HashAggregate [COMET: Spark Final aggregate without Comet Partial requires
compatible intermediate buffer formats, but the following aggregate function(s)
have incompatible buffers: avg, count, sum]
++- HashAggregate [COMET: Spark Final aggregate without Comet Partial requires
compatible intermediate buffer formats, but the following aggregate function(s)
have incompatible buffers: count]
+- Exchange
+- HashAggregate
+- Project
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 ae14c68207..e597e6e6e4 100644
--- a/spark/src/test/scala/org/apache/comet/exec/CometAggregateSuite.scala
+++ b/spark/src/test/scala/org/apache/comet/exec/CometAggregateSuite.scala
@@ -183,6 +183,32 @@ class CometAggregateSuite extends CometTestBase with
AdaptiveSparkPlanHelper {
}
}
+ test("mixed engine sum/avg: Comet partial + Spark final matches Spark") {
+ val data = (0 until 100).map(i => (i, i.toLong, i.toDouble, i % 7))
+ withParquetTable(data, "tbl") {
+ withSQLConf(
+ CometConf.COMET_ENABLE_FINAL_HASH_AGGREGATE.key -> "false",
+ CometConf.COMET_EXEC_SHUFFLE_ENABLED.key -> "true",
+ CometConf.COMET_SHUFFLE_MODE.key -> "jvm") {
+ checkSparkAnswer(
+ "SELECT _4, SUM(_1), SUM(_2), SUM(_3), AVG(_1), AVG(_2), AVG(_3)
FROM tbl GROUP BY _4")
+ }
+ }
+ }
+
+ test("mixed engine sum/avg: Spark partial + Comet final matches Spark") {
+ val data = (0 until 100).map(i => (i, i.toLong, i.toDouble, i % 7))
+ withParquetTable(data, "tbl") {
+ withSQLConf(
+ CometConf.COMET_ENABLE_PARTIAL_HASH_AGGREGATE.key -> "false",
+ CometConf.COMET_EXEC_SHUFFLE_ENABLED.key -> "true",
+ CometConf.COMET_SHUFFLE_MODE.key -> "jvm") {
+ checkSparkAnswer(
+ "SELECT _4, SUM(_1), SUM(_2), SUM(_3), AVG(_1), AVG(_2), AVG(_3)
FROM tbl GROUP BY _4")
+ }
+ }
+ }
+
test("Aggregation without aggregate expressions should use correct result
expressions") {
withSQLConf(
CometConf.COMET_ENABLED.key -> "true",
diff --git
a/spark/src/test/scala/org/apache/comet/rules/CometExecRuleSuite.scala
b/spark/src/test/scala/org/apache/comet/rules/CometExecRuleSuite.scala
index 7fa06a26cc..5135d8483c 100644
--- a/spark/src/test/scala/org/apache/comet/rules/CometExecRuleSuite.scala
+++ b/spark/src/test/scala/org/apache/comet/rules/CometExecRuleSuite.scala
@@ -34,7 +34,7 @@ import
org.apache.spark.sql.execution.exchange.{BroadcastExchangeExec, ShuffleEx
import org.apache.spark.sql.types.{DataTypes, StructField, StructType}
import org.apache.comet.CometConf
-import org.apache.comet.CometSparkSessionExtensions.{isSpark40Plus,
isSpark42Plus}
+import org.apache.comet.CometSparkSessionExtensions.{isSpark35Plus,
isSpark40Plus, isSpark42Plus}
import org.apache.comet.testing.{DataGenOptions, FuzzDataGenerator}
/**
@@ -152,8 +152,8 @@ class CometExecRuleSuite extends CometTestBase {
CometConf.COMET_EXEC_LOCAL_TABLE_SCAN_ENABLED.key -> "true") {
val transformedPlan = applyCometExecRule(sparkPlan)
- // SUM has incompatible intermediate buffers, so if the final
aggregate cannot
- // be converted to Comet, neither should be
+ // COUNT is intentionally excluded from mixed execution (AQE /
count-bug reasons), so if
+ // the final aggregate cannot be converted to Comet, neither should
the partial.
assert(
countOperators(transformedPlan, classOf[HashAggregateExec]) ==
originalHashAggCount)
assert(countOperators(transformedPlan,
classOf[CometHashAggregateExec]) == 0)
@@ -177,7 +177,8 @@ class CometExecRuleSuite extends CometTestBase {
CometConf.COMET_EXEC_LOCAL_TABLE_SCAN_ENABLED.key -> "true") {
val transformedPlan = applyCometExecRule(sparkPlan)
- // if the partial aggregate cannot be converted to Comet, then neither
should be
+ // COUNT blocks mixed execution, so if the partial cannot be
converted, neither should
+ // the final.
assert(
countOperators(transformedPlan, classOf[HashAggregateExec]) ==
originalHashAggCount)
assert(countOperators(transformedPlan,
classOf[CometHashAggregateExec]) == 0)
@@ -231,6 +232,137 @@ class CometExecRuleSuite extends CometTestBase {
}
}
+ test("CometExecRule should allow SUM mixed Comet partial and Spark final") {
+ withTempView("test_data") {
+ createTestDataFrame.createOrReplaceTempView("test_data")
+ val sparkPlan =
+ createSparkPlan(spark, "SELECT SUM(id) FROM test_data GROUP BY (id %
3)")
+ assert(countOperators(sparkPlan, classOf[HashAggregateExec]) == 2)
+ withSQLConf(
+ CometConf.COMET_ENABLE_FINAL_HASH_AGGREGATE.key -> "false",
+ CometConf.COMET_EXEC_LOCAL_TABLE_SCAN_ENABLED.key -> "true") {
+ val transformedPlan = applyCometExecRule(sparkPlan)
+ // SUM buffer matches Spark: partial converts to Comet, final stays
Spark.
+ assert(countOperators(transformedPlan, classOf[HashAggregateExec]) ==
1) // final
+ assert(countOperators(transformedPlan,
classOf[CometHashAggregateExec]) == 1) // partial
+ }
+ }
+ }
+
+ test("CometExecRule should allow SUM mixed Spark partial and Comet final") {
+ withTempView("test_data") {
+ createTestDataFrame.createOrReplaceTempView("test_data")
+ val sparkPlan =
+ createSparkPlan(spark, "SELECT SUM(id) FROM test_data GROUP BY (id %
3)")
+ assert(countOperators(sparkPlan, classOf[HashAggregateExec]) == 2)
+ withSQLConf(
+ CometConf.COMET_ENABLE_PARTIAL_HASH_AGGREGATE.key -> "false",
+ CometConf.COMET_EXEC_LOCAL_TABLE_SCAN_ENABLED.key -> "true") {
+ val transformedPlan = applyCometExecRule(sparkPlan)
+ assert(countOperators(transformedPlan, classOf[HashAggregateExec]) ==
1) // partial
+ assert(countOperators(transformedPlan,
classOf[CometHashAggregateExec]) == 1) // final
+ }
+ }
+ }
+
+ test("CometExecRule should allow AVG mixed Comet partial and Spark final") {
+ withTempView("test_data") {
+ createTestDataFrame.createOrReplaceTempView("test_data")
+ val sparkPlan =
+ createSparkPlan(spark, "SELECT AVG(id) FROM test_data GROUP BY (id %
3)")
+ assert(countOperators(sparkPlan, classOf[HashAggregateExec]) == 2)
+ withSQLConf(
+ CometConf.COMET_ENABLE_FINAL_HASH_AGGREGATE.key -> "false",
+ CometConf.COMET_EXEC_LOCAL_TABLE_SCAN_ENABLED.key -> "true") {
+ val transformedPlan = applyCometExecRule(sparkPlan)
+ assert(countOperators(transformedPlan, classOf[HashAggregateExec]) ==
1) // final
+ assert(countOperators(transformedPlan,
classOf[CometHashAggregateExec]) == 1) // partial
+ }
+ }
+ }
+
+ test("CometExecRule should not allow try_sum mixed execution") {
+ assume(isSpark35Plus, "try_sum was added in Spark 3.5")
+ withTempView("test_data") {
+ createTestDataFrame.createOrReplaceTempView("test_data")
+ val sparkPlan =
+ createSparkPlan(spark, "SELECT try_sum(id) FROM test_data GROUP BY (id
% 3)")
+ assert(countOperators(sparkPlan, classOf[HashAggregateExec]) == 2)
+ withSQLConf(
+ CometConf.COMET_ENABLE_FINAL_HASH_AGGREGATE.key -> "false",
+ CometConf.COMET_EXEC_LOCAL_TABLE_SCAN_ENABLED.key -> "true") {
+ val transformedPlan = applyCometExecRule(sparkPlan)
+ // TRY-mode SUM uses a Comet-internal buffer column, so mixing is
unsafe:
+ // the partial must also fall back to Spark.
+ assert(countOperators(transformedPlan, classOf[HashAggregateExec]) ==
2)
+ assert(countOperators(transformedPlan,
classOf[CometHashAggregateExec]) == 0)
+ }
+ }
+ }
+
+ test("CometExecRule should not allow decimal AVG mixed execution") {
+ withTempView("test_data") {
+ createTestDataFrame.createOrReplaceTempView("test_data")
+ // Precision must be large enough (prec + 4 > 15) that Spark's own
DecimalAggregates
+ // optimizer rule does not rewrite AVG to operate on the unscaled Long
value, which would
+ // sidestep the decimal buffer path this test is meant to exercise.
+ val sparkPlan =
+ createSparkPlan(
+ spark,
+ "SELECT AVG(CAST(id AS DECIMAL(20, 2))) FROM test_data GROUP BY (id
% 3)")
+ assert(countOperators(sparkPlan, classOf[HashAggregateExec]) == 2)
+ withSQLConf(
+ CometConf.COMET_ENABLE_FINAL_HASH_AGGREGATE.key -> "false",
+ CometConf.COMET_EXEC_LOCAL_TABLE_SCAN_ENABLED.key -> "true") {
+ val transformedPlan = applyCometExecRule(sparkPlan)
+ // Decimal AVG is deferred (its overflow path nulls count differently
from Spark), so
+ // mixed execution is unsafe and the partial must also fall back to
Spark.
+ assert(countOperators(transformedPlan, classOf[HashAggregateExec]) ==
2)
+ assert(countOperators(transformedPlan,
classOf[CometHashAggregateExec]) == 0)
+ }
+ }
+ }
+
+ test("CometExecRule should not allow decimal SUM mixed execution") {
+ withTempView("test_data") {
+ createTestDataFrame.createOrReplaceTempView("test_data")
+ // Precision must be large enough (prec + 4 > 15) that Spark's own
DecimalAggregates
+ // optimizer rule does not rewrite SUM to operate on the unscaled Long
value, which would
+ // sidestep the decimal buffer path this test is meant to exercise.
+ val sparkPlan =
+ createSparkPlan(
+ spark,
+ "SELECT SUM(CAST(id AS DECIMAL(20, 2))) FROM test_data GROUP BY (id
% 3)")
+ assert(countOperators(sparkPlan, classOf[HashAggregateExec]) == 2)
+ withSQLConf(
+ CometConf.COMET_ENABLE_FINAL_HASH_AGGREGATE.key -> "false",
+ CometConf.COMET_EXEC_LOCAL_TABLE_SCAN_ENABLED.key -> "true") {
+ val transformedPlan = applyCometExecRule(sparkPlan)
+ // Decimal SUM overflow detection (ANSI throw / Legacy null) does not
survive a
+ // Spark-partial / Comet-final split, so mixed execution is unsafe and
the partial
+ // must also fall back to Spark.
+ assert(countOperators(transformedPlan, classOf[HashAggregateExec]) ==
2)
+ assert(countOperators(transformedPlan,
classOf[CometHashAggregateExec]) == 0)
+ }
+ }
+ }
+
+ test("CometExecRule should allow AVG mixed Spark partial and Comet final") {
+ withTempView("test_data") {
+ createTestDataFrame.createOrReplaceTempView("test_data")
+ val sparkPlan =
+ createSparkPlan(spark, "SELECT AVG(id) FROM test_data GROUP BY (id %
3)")
+ assert(countOperators(sparkPlan, classOf[HashAggregateExec]) == 2)
+ withSQLConf(
+ CometConf.COMET_ENABLE_PARTIAL_HASH_AGGREGATE.key -> "false",
+ CometConf.COMET_EXEC_LOCAL_TABLE_SCAN_ENABLED.key -> "true") {
+ val transformedPlan = applyCometExecRule(sparkPlan)
+ assert(countOperators(transformedPlan, classOf[HashAggregateExec]) ==
1) // partial
+ assert(countOperators(transformedPlan,
classOf[CometHashAggregateExec]) == 1) // final
+ }
+ }
+ }
+
test("CometExecRule should allow BloomFilter mixed Comet partial and Spark
final") {
assume(!isSpark42Plus,
"https://github.com/apache/datafusion-comet/issues/4142")
val funcId = new FunctionIdentifier("bloom_filter_agg")
---------------------------------------------------------------------
To unsubscribe, e-mail: [email protected]
For additional commands, e-mail: [email protected]