Skip to content

Commit a49fd8f

Browse files
authored
fix: fall back to Spark for regr_* aggregates until their merge matches Spark (#6451)
* fix: fall back to Spark for regr_* aggregates until their merge matches Spark The native merge of partial aggregates (variance_merge and covariance_merge in welford.rs) orders its floating-point operations differently from Spark's CentralMomentAgg and Covariance. For a variable that is constant at a value binary floating point cannot represent exactly, merged from two or more partial aggregates, m2 ends up around 1e-34 instead of 0, so Spark's exact m2 == 0 degenerate-case checks never fire and regr_slope, regr_intercept, regr_r2, regr_sxx, regr_syy and regr_sxy return wrong values. Mark the five serdes behind these six functions Incompatible, so they fall back to Spark by default as they did in 1.0.0 and run natively only with the per-expression allowIncompatible config. The regr SQL-file tests opt in to keep covering the native path, and a new test from the issue's reproducer checks that the default plan falls back and matches Spark. Closes #6423. * test: move the regr_* fallback check into a SQL file test One query per function, each expecting the #6423 fallback reason, over the reproducer's data: x constant at 0.1 in two three-row files, which the native merge gets wrong. regr.sql still opts in and covers the native path.
1 parent 02df813 commit a49fd8f

4 files changed

Lines changed: 109 additions & 17 deletions

File tree

‎docs/source/user-guide/latest/expressions.md‎

Lines changed: 6 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -126,12 +126,12 @@ The tables below list every Spark built-in expression with its current status.
126126
| `regr_avgx` | ✅ | — | Native: Spark rewrites to `Average` (tests in [#4551](https://github.com/apache/datafusion-comet/pull/4551)) |
127127
| `regr_avgy` | ✅ | — | Native: Spark rewrites to `Average` (tests in [#4551](https://github.com/apache/datafusion-comet/pull/4551)) |
128128
| `regr_count` | ✅ | — | Native: Spark rewrites to `Count` (tests in [#4551](https://github.com/apache/datafusion-comet/pull/4551)) |
129-
| `regr_intercept` | ✅ | Native | |
130-
| `regr_r2` | ✅ | Native | |
131-
| `regr_slope` | ✅ | Native | |
132-
| `regr_sxx` | ✅ | Native | |
133-
| `regr_sxy` | ✅ | Native | |
134-
| `regr_syy` | ✅ | Native | |
129+
| `regr_intercept` | ✅ | Native | Falls back by default because the native merge of partial aggregates differs from Spark ([#6423](https://github.com/apache/datafusion-comet/issues/6423)); the native path is opt-in via `spark.comet.expression.RegrIntercept.allowIncompatible=true` |
130+
| `regr_r2` | ✅ | Native | Falls back by default because the native merge of partial aggregates differs from Spark ([#6423](https://github.com/apache/datafusion-comet/issues/6423)); the native path is opt-in via `spark.comet.expression.RegrR2.allowIncompatible=true` |
131+
| `regr_slope` | ✅ | Native | Falls back by default because the native merge of partial aggregates differs from Spark ([#6423](https://github.com/apache/datafusion-comet/issues/6423)); the native path is opt-in via `spark.comet.expression.RegrSlope.allowIncompatible=true` |
132+
| `regr_sxx` | ✅ | Native | Falls back by default because the native merge of partial aggregates differs from Spark ([#6423](https://github.com/apache/datafusion-comet/issues/6423)); the native path is opt-in via `spark.comet.expression.RegrReplacement.allowIncompatible=true` (Spark plans `regr_sxx` as `RegrReplacement`) |
133+
| `regr_sxy` | ✅ | Native | Falls back by default because the native merge of partial aggregates differs from Spark ([#6423](https://github.com/apache/datafusion-comet/issues/6423)); the native path is opt-in via `spark.comet.expression.RegrSXY.allowIncompatible=true` |
134+
| `regr_syy` | ✅ | Native | Falls back by default because the native merge of partial aggregates differs from Spark ([#6423](https://github.com/apache/datafusion-comet/issues/6423)); the native path is opt-in via `spark.comet.expression.RegrReplacement.allowIncompatible=true` (Spark plans `regr_syy` as `RegrReplacement`) |
135135
| `skewness` | 🔜 | — | Not yet implemented natively |
136136
| `some` | ✅ | — | |
137137
| `std` | ✅ | Native | |

‎spark/src/main/scala/org/apache/comet/serde/aggregates.scala‎

Lines changed: 37 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -22,7 +22,7 @@ package org.apache.comet.serde
2222
import scala.jdk.CollectionConverters._
2323

2424
import org.apache.spark.sql.catalyst.expressions.{Attribute, Cast, Expression, Literal}
25-
import org.apache.spark.sql.catalyst.expressions.aggregate.{AggregateExpression, ApproximatePercentile, Average, BitAndAgg, BitOrAgg, BitXorAgg, BloomFilterAggregate, CentralMomentAgg, CollectList, CollectSet, Complete, Corr, Count, Covariance, CovPopulation, CovSample, First, HyperLogLogPlusPlus, Last, Max, MaxBy, MaxMinBy, Min, MinBy, Mode, Partial, Percentile, RegrIntercept, RegrR2, RegrReplacement, RegrSlope, RegrSXY, StddevPop, StddevSamp, Sum, VariancePop, VarianceSamp}
25+
import org.apache.spark.sql.catalyst.expressions.aggregate.{AggregateExpression, AggregateFunction, ApproximatePercentile, Average, BitAndAgg, BitOrAgg, BitXorAgg, BloomFilterAggregate, CentralMomentAgg, CollectList, CollectSet, Complete, Corr, Count, Covariance, CovPopulation, CovSample, First, HyperLogLogPlusPlus, Last, Max, MaxBy, MaxMinBy, Min, MinBy, Mode, Partial, Percentile, RegrIntercept, RegrR2, RegrReplacement, RegrSlope, RegrSXY, StddevPop, StddevSamp, Sum, VariancePop, VarianceSamp}
2626
import org.apache.spark.sql.catalyst.util.ArrayData
2727
import org.apache.spark.sql.comet.CometExecUtils
2828
import org.apache.spark.sql.internal.SQLConf
@@ -929,7 +929,27 @@ private[comet] object RegrSparkVersions {
929929
* variable (y) and `child2` is the independent variable (x), matching the native accumulator's
930930
* `regr_*(y, x)` convention.
931931
*/
932-
trait CometRegrBase {
932+
trait CometRegrBase[T <: AggregateFunction] extends CometAggregateExpressionSerde[T] {
933+
934+
/** The SQL function or functions this serde implements, named in the incompatibility note. */
935+
protected def sqlFunctions: String
936+
937+
// The native merge (`variance_merge` and `covariance_merge` in welford.rs) orders its
938+
// floating-point operations differently from Spark's CentralMomentAgg and Covariance. Merging
939+
// the first partial buffer into the zero-initialized final buffer can leave a one-ULP error on
940+
// the mean, so a constant variable ends up with a tiny non-zero m2 and Spark's exact `m2 == 0`
941+
// degenerate-case checks never fire. Porting Spark's merge order would make these Compatible.
942+
private def mergeOrderReason: String =
943+
s"Comet merges the partial aggregates of $sqlFunctions in a different floating-point " +
944+
"operation order from Spark. When a group's rows come from more than one partial " +
945+
"aggregate and a variable is constant at a value that binary floating point cannot " +
946+
"represent exactly, such as 0.1, Comet returns a wrong value where Spark returns NULL, " +
947+
"0.0 or 1.0 (https://github.com/apache/datafusion-comet/issues/6423)"
948+
949+
override def getIncompatibleReasons(): Seq[String] = Seq(mergeOrderReason)
950+
951+
override def getSupportLevel(expr: T): SupportLevel = Incompatible(Some(mergeOrderReason))
952+
933953
def convertRegr(
934954
aggExpr: AggregateExpression,
935955
regrType: ExprOuterClass.Regr.RegrType,
@@ -966,7 +986,9 @@ trait CometRegrBase {
966986
}
967987
}
968988

969-
object CometRegrSlope extends CometAggregateExpressionSerde[RegrSlope] with CometRegrBase {
989+
object CometRegrSlope extends CometRegrBase[RegrSlope] {
990+
override protected def sqlFunctions: String = "`regr_slope`"
991+
970992
override def convert(
971993
aggExpr: AggregateExpression,
972994
expr: RegrSlope,
@@ -982,9 +1004,9 @@ object CometRegrSlope extends CometAggregateExpressionSerde[RegrSlope] with Come
9821004
binding)
9831005
}
9841006

985-
object CometRegrIntercept
986-
extends CometAggregateExpressionSerde[RegrIntercept]
987-
with CometRegrBase {
1007+
object CometRegrIntercept extends CometRegrBase[RegrIntercept] {
1008+
override protected def sqlFunctions: String = "`regr_intercept`"
1009+
9881010
override def convert(
9891011
aggExpr: AggregateExpression,
9901012
expr: RegrIntercept,
@@ -1000,7 +1022,9 @@ object CometRegrIntercept
10001022
binding)
10011023
}
10021024

1003-
object CometRegrR2 extends CometAggregateExpressionSerde[RegrR2] with CometRegrBase {
1025+
object CometRegrR2 extends CometRegrBase[RegrR2] {
1026+
override protected def sqlFunctions: String = "`regr_r2`"
1027+
10041028
override def convert(
10051029
aggExpr: AggregateExpression,
10061030
expr: RegrR2,
@@ -1010,7 +1034,9 @@ object CometRegrR2 extends CometAggregateExpressionSerde[RegrR2] with CometRegrB
10101034
convertRegr(aggExpr, ExprOuterClass.Regr.RegrType.R2, expr.y, expr.x, inputs, binding)
10111035
}
10121036

1013-
object CometRegrSXY extends CometAggregateExpressionSerde[RegrSXY] with CometRegrBase {
1037+
object CometRegrSXY extends CometRegrBase[RegrSXY] {
1038+
override protected def sqlFunctions: String = "`regr_sxy`"
1039+
10141040
override def convert(
10151041
aggExpr: AggregateExpression,
10161042
expr: RegrSXY,
@@ -1027,9 +1053,9 @@ object CometRegrSXY extends CometAggregateExpressionSerde[RegrSXY] with CometReg
10271053
* deviations) of its single child. We serialize it as the `SXX` regression statistic with the
10281054
* child duplicated, since `regr_sxx(c, c) = m2(c)`.
10291055
*/
1030-
object CometRegrReplacement
1031-
extends CometAggregateExpressionSerde[RegrReplacement]
1032-
with CometRegrBase {
1056+
object CometRegrReplacement extends CometRegrBase[RegrReplacement] {
1057+
override protected def sqlFunctions: String = "`regr_sxx` and `regr_syy`"
1058+
10331059
override def convert(
10341060
aggExpr: AggregateExpression,
10351061
expr: RegrReplacement,

‎spark/src/test/resources/sql-tests/expressions/aggregate/regr.sql‎

Lines changed: 10 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -19,6 +19,16 @@
1919
-- regr_avgy, regr_sxx, regr_syy, regr_sxy, regr_slope, regr_intercept, regr_r2.
2020
-- All functions take (y, x) and operate only on rows where BOTH y and x are non-null.
2121

22+
-- regr_slope, regr_intercept, regr_r2, regr_sxx, regr_syy and regr_sxy fall back to Spark by
23+
-- default because their native merge of partial aggregates differs from Spark's
24+
-- (https://github.com/apache/datafusion-comet/issues/6423). Opt in so the queries below cover
25+
-- the native path. Spark plans regr_sxx and regr_syy as RegrReplacement.
26+
-- Config: spark.comet.expression.RegrSlope.allowIncompatible=true
27+
-- Config: spark.comet.expression.RegrIntercept.allowIncompatible=true
28+
-- Config: spark.comet.expression.RegrR2.allowIncompatible=true
29+
-- Config: spark.comet.expression.RegrSXY.allowIncompatible=true
30+
-- Config: spark.comet.expression.RegrReplacement.allowIncompatible=true
31+
2232
statement
2333
CREATE TABLE test_regr(y double, x double, grp string) USING parquet
2434

Lines changed: 56 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,56 @@
1+
-- Licensed to the Apache Software Foundation (ASF) under one
2+
-- or more contributor license agreements. See the NOTICE file
3+
-- distributed with this work for additional information
4+
-- regarding copyright ownership. The ASF licenses this file
5+
-- to you under the Apache License, Version 2.0 (the
6+
-- "License"); you may not use this file except in compliance
7+
-- with the License. You may obtain a copy of the License at
8+
--
9+
-- http://www.apache.org/licenses/LICENSE-2.0
10+
--
11+
-- Unless required by applicable law or agreed to in writing,
12+
-- software distributed under the License is distributed on an
13+
-- "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
14+
-- KIND, either express or implied. See the License for the
15+
-- specific language governing permissions and limitations
16+
-- under the License.
17+
18+
-- regr_slope, regr_intercept, regr_r2, regr_sxx, regr_syy and regr_sxy fall back to Spark by
19+
-- default, because their native merge of partial aggregates orders its floating-point operations
20+
-- differently from Spark's (https://github.com/apache/datafusion-comet/issues/6423). Each query
21+
-- below holds a single regr function, so it checks that function's own fallback. regr.sql opts in
22+
-- and covers the native path.
23+
24+
-- The data from #6423: x is constant at 0.1, and each INSERT writes one file of three rows, so the
25+
-- rows are merged from two partial aggregates. The native merge leaves x with a tiny non-zero
26+
-- variance there, and returns wrong values where Spark returns NULL, 0.0 or 1.0.
27+
statement
28+
CREATE TABLE test_regr_fallback(y double, x double) USING parquet
29+
30+
statement
31+
INSERT INTO test_regr_fallback SELECT CAST(id AS DOUBLE), 0.1D FROM range(0, 3, 1, 1)
32+
33+
statement
34+
INSERT INTO test_regr_fallback SELECT CAST(id AS DOUBLE), 0.1D FROM range(3, 6, 1, 1)
35+
36+
query expect_fallback(issues/6423)
37+
SELECT regr_slope(y, x) FROM test_regr_fallback
38+
39+
query expect_fallback(issues/6423)
40+
SELECT regr_intercept(y, x) FROM test_regr_fallback
41+
42+
query expect_fallback(issues/6423)
43+
SELECT regr_r2(y, x) FROM test_regr_fallback
44+
45+
-- The constant as the dependent variable
46+
query expect_fallback(issues/6423)
47+
SELECT regr_r2(x, y) FROM test_regr_fallback
48+
49+
query expect_fallback(issues/6423)
50+
SELECT regr_sxx(y, x) FROM test_regr_fallback
51+
52+
query expect_fallback(issues/6423)
53+
SELECT regr_sxy(y, x) FROM test_regr_fallback
54+
55+
query expect_fallback(issues/6423)
56+
SELECT regr_syy(x, y) FROM test_regr_fallback

0 commit comments

Comments
 (0)