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
28 changes: 14 additions & 14 deletions native/spark-expr/src/agg_funcs/regr.rs
Original file line number Diff line number Diff line change
Expand Up @@ -72,13 +72,13 @@ pub struct Regr {
name: String,
signature: Signature,
regr_type: RegrType,
/// Only consulted for `Slope` / `Intercept`. When `true` (Spark 3.5+),
/// `VariancePop(x)` counts only rows where both `y` and `x` are non-null.
/// When `false` (Spark 3.4), it counts every row where `x` is non-null.
/// Only consulted for `Slope` / `Intercept`. When `true` (Spark 3.5.2 and
/// 4.0.0 onward), `VariancePop(x)` counts only rows where both `y` and `x`
/// are non-null. When `false`, it counts every row where `x` is non-null.
filter_var_by_pair_nulls: bool,
/// Only consulted for `R2`. When `true` (Spark 3.5+), a constant dependent
/// variable evaluates to `1.0` and a constant independent variable to `null`.
/// When `false` (Spark 3.4), those two cases are reversed.
/// Only consulted for `R2`. When `true` (Spark 3.5.9, 4.0.3, 4.1.2 and 4.2.0
/// onward), a constant dependent variable evaluates to `1.0` and a constant
/// independent variable to `null`. When `false`, those two cases are reversed.
r2_constant_dependent_is_perfect_fit: bool,
}

Expand Down Expand Up @@ -282,13 +282,13 @@ impl Accumulator for RegrCovAccumulator {
/// count, mean1, mean2, algo_const, m2(y), m2(x).
///
/// Spark's degenerate-case handling (see `RegrR2.evaluateExpression`) was
/// swapped by SPARK-55969, which shipped in 3.5.9, 4.0.3 and 4.1.0. In both
/// eras one degenerate case returns `null` and the other returns `1.0` (a
/// swapped by SPARK-55969, which shipped in 3.5.9, 4.0.3, 4.1.2 and 4.2.0. In
/// both eras one degenerate case returns `null` and the other returns `1.0` (a
/// perfect fit), but which is which differs:
/// - before SPARK-55969 (Spark 3.4): constant dependent `y` (`m2(y) == 0`) ->
/// `null`; constant independent `x` (`m2(x) == 0`) -> `1.0`.
/// - after (the Spark 3.5+ versions Comet builds against): constant dependent
/// `y` -> `1.0`; constant independent `x` -> `null`.
/// - before SPARK-55969 (Spark 3.4 and earlier patches of 3.5, 4.0 and 4.1):
/// constant dependent `y` (`m2(y) == 0`) -> `null`; constant independent `x`
/// (`m2(x) == 0`) -> `1.0`.
/// - after: constant dependent `y` -> `1.0`; constant independent `x` -> `null`.
///
/// `m2(y) == 0` also covers fewer than two rows. DataFusion returns `null` in
/// both degenerate cases.
Expand Down Expand Up @@ -368,10 +368,10 @@ impl Accumulator for RegrR2Accumulator {
// The two degenerate cases (constant dependent y, constant independent x)
// return null and 1.0 respectively, but SPARK-55969 swapped which is which.
let (null_case, perfect_fit_case) = if self.constant_dependent_is_perfect_fit {
// Spark 3.5+: constant x -> null, constant y -> 1.0.
// After SPARK-55969: constant x -> null, constant y -> 1.0.
(m2_x == 0.0, m2_y == 0.0)
} else {
// Spark 3.4: constant y -> null, constant x -> 1.0.
// Before SPARK-55969: constant y -> null, constant x -> 1.0.
(m2_y == 0.0, m2_x == 0.0)
};
if null_case {
Expand Down
63 changes: 52 additions & 11 deletions spark/src/main/scala/org/apache/comet/serde/aggregates.scala
Original file line number Diff line number Diff line change
Expand Up @@ -29,7 +29,7 @@ import org.apache.spark.sql.internal.SQLConf
import org.apache.spark.sql.types.{BinaryType, BooleanType, ByteType, DataType, DateType, DecimalType, DoubleType, FloatType, IntegerType, LongType, NumericType, ShortType, StringType, TimestampNTZType, TimestampType}

import org.apache.comet.CometConf.COMET_EXEC_STRICT_FLOATING_POINT
import org.apache.comet.CometSparkSessionExtensions.{isSpark35Plus, isSpark41Plus, isSpark42Plus, withFallbackReason}
import org.apache.comet.CometSparkSessionExtensions.{isSpark41Plus, isSpark42Plus, withFallbackReason}
import org.apache.comet.expressions.CometEvalMode
import org.apache.comet.serde.QueryPlanSerde.{evalModeToProto, exprToProto, serializeDataType}
import org.apache.comet.shims.{CometCollectShim, CometEvalModeUtil, CometTypeShim}
Expand Down Expand Up @@ -847,6 +847,51 @@ object CometCorr extends CometAggregateExpressionSerde[Corr] {
}
}

/**
* The Spark releases that changed the regression aggregates. Both changes shipped in patch
* releases, so the minor version alone places them on the wrong side for earlier patches.
*/
private[comet] object RegrSparkVersions {

/**
* SPARK-48719 made regr_slope and regr_intercept count VariancePop(x) only over rows where both
* y and x are non-null. Present from 3.5.2 and 4.0.0; earlier 3.5 patches and 3.4 count every
* row where x is non-null.
*/
def slopeFiltersVarByPairNulls(sparkVersion: String): Boolean =
majorMinorPatch(sparkVersion) match {
case Some((3, 5, patch)) => patch >= 2
case Some((major, minor, _)) => isAfterMinor(major, minor, 3, 5)
case None => sparkVersion >= "3.5"
}

/**
* SPARK-55969 swapped regr_r2's degenerate cases: a constant dependent variable yields 1.0 (was
* null) and a constant independent variable yields null (was 1.0). Present from 3.5.9, 4.0.3,
* 4.1.2 and 4.2.0; earlier patches of those lines and 3.4 keep the old cases.
*/
def r2DegenerateCasesSwapped(sparkVersion: String): Boolean =
majorMinorPatch(sparkVersion) match {
case Some((3, 5, patch)) => patch >= 9
case Some((4, 0, patch)) => patch >= 3
case Some((4, 1, patch)) => patch >= 2
case Some((major, minor, _)) => isAfterMinor(major, minor, 3, 5)
case None => sparkVersion >= "3.5"
}

private def isAfterMinor(major: Int, minor: Int, thanMajor: Int, thanMinor: Int): Boolean =
major > thanMajor || (major == thanMajor && minor > thanMinor)

// Spark's VersionUtils is private to its packages, so parse the leading major.minor.patch
// here; a missing patch reads as 0 and a suffix such as -SNAPSHOT is ignored.
private val versionPattern = """^(\d+)\.(\d+)(?:\.(\d+))?""".r

private def majorMinorPatch(sparkVersion: String): Option[(Int, Int, Int)] =
versionPattern.findFirstMatchIn(sparkVersion).map { m =>
(m.group(1).toInt, m.group(2).toInt, Option(m.group(3)).map(_.toInt).getOrElse(0))
}
}

/**
* Shared serialization for the simple linear regression aggregates. `child1` is the dependent
* variable (y) and `child2` is the independent variable (x), matching the native accumulator's
Expand All @@ -870,16 +915,12 @@ trait CometRegrBase {
builder.setChild2(child2Expr.get)
builder.setRegrType(regrType)
builder.setDatatype(dataType.get)
// Spark 3.5 fixed regr_slope/regr_intercept so VariancePop(x) only counts
// rows where both y and x are non-null. Spark 3.4 counts every row where x
// is non-null. The native accumulator only consults this for slope/intercept.
builder.setFilterVarByPairNulls(isSpark35Plus)
// Spark swapped regr_r2's degenerate-case handling: a constant dependent
// variable now yields 1.0 (was null) and a constant independent variable
// yields null (was 1.0). The swap is present in the Spark versions Comet
// builds against for 3.5 and later (3.5.9, 4.0.3+, 4.1, 4.2) but not in 3.4.
// The native accumulator only consults this for R2.
builder.setR2ConstantDependentIsPerfectFit(isSpark35Plus)
// Both regression fixes shipped in patch releases, so the running Spark's exact version
// decides which behaviour the native accumulator mirrors.
val sparkVersion = org.apache.spark.SPARK_VERSION
builder.setFilterVarByPairNulls(RegrSparkVersions.slopeFiltersVarByPairNulls(sparkVersion))
builder.setR2ConstantDependentIsPerfectFit(
RegrSparkVersions.r2DegenerateCasesSwapped(sparkVersion))

Some(
ExprOuterClass.AggExpr
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -144,8 +144,8 @@ SELECT regr_slope(y, x), regr_intercept(y, x), regr_r2(y, x) FROM test_regr_sing

-- edge case: independent variable (x) is constant but y varies.
-- var_pop(x) = 0, so slope/intercept are NULL and regr_sxx = 0. regr_r2 is a
-- degenerate case whose value depends on the Spark version (1.0 on 3.4, NULL on
-- 3.5+ after SPARK-55969), which Comet matches per version.
-- degenerate case whose value depends on the Spark patch release (1.0 before
-- SPARK-55969, NULL from 3.5.9, 4.0.3, 4.1.2 and 4.2.0), which Comet matches.
statement
CREATE TABLE test_regr_const_x(y double, x double) USING parquet

Expand All @@ -160,8 +160,8 @@ SELECT regr_sxx(y, x), regr_syy(y, x), regr_sxy(y, x) FROM test_regr_const_x

-- edge case: dependent variable (y) is constant but x varies.
-- The slope is 0 and the intercept equals the constant y. regr_r2 is a degenerate
-- case whose value depends on the Spark version (NULL on 3.4, 1.0 on 3.5+ after
-- SPARK-55969), which Comet matches per version.
-- case whose value depends on the Spark patch release (NULL before SPARK-55969,
-- 1.0 from 3.5.9, 4.0.3, 4.1.2 and 4.2.0), which Comet matches.
statement
CREATE TABLE test_regr_const_y(y double, x double) USING parquet

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -43,6 +43,7 @@ import org.apache.spark.sql.types.{DataTypes, StructField, StructType}
import org.apache.comet.CometConf
import org.apache.comet.CometConf.COMET_EXEC_STRICT_FLOATING_POINT
import org.apache.comet.CometSparkSessionExtensions.isSpark41Plus
import org.apache.comet.serde.RegrSparkVersions
import org.apache.comet.testing.{DataGenOptions, FuzzDataGenerator, ParquetGenerator, SchemaGenOptions}

/**
Expand Down Expand Up @@ -1404,6 +1405,57 @@ class CometAggregateSuite extends CometTestBase with AdaptiveSparkPlanHelper {
}
}

test("regression aggregate flags follow the Spark patch release that changed them") {
// Both fixes shipped mid-line, so the flag must flip at the exact patch and stay off for
// earlier patches of the same minor version. A vendor suffix parses, a missing patch reads
// as 0, and an unparsable string falls back to the old minor-version rule.
val slopeExpectations = Seq(
"3.4.3" -> false,
"3.5.0" -> false,
"3.5.1" -> false,
"3.5.2" -> true,
"3.5.9" -> true,
"3.5.10" -> true,
"4.0.0" -> true,
"4.1.3" -> true,
"4.2.0-SNAPSHOT" -> true,
"4.1.3-amzn-1" -> true,
"3.5" -> false,
"v3.5.9" -> true)
slopeExpectations.foreach { case (version, expected) =>
assert(
RegrSparkVersions.slopeFiltersVarByPairNulls(version) == expected,
s"slope pair-null filtering for Spark $version should be $expected")
}

val r2Expectations = Seq(
"3.4.3" -> false,
"3.5.3" -> false,
"3.5.8" -> false,
"3.5.9" -> true,
"3.5.10" -> true,
"4.0.1" -> false,
"4.0.2" -> false,
"4.0.3" -> true,
"4.0.4" -> true,
"4.1.0" -> false,
"4.1.1" -> false,
"4.1.2" -> true,
"4.1.3" -> true,
"4.2.0" -> true,
"4.2.0-SNAPSHOT" -> true,
"3.6.0" -> true,
"5.0.0" -> true,
"4.1.3-amzn-1" -> true,
"3.5" -> false,
"v3.5.9" -> true)
r2Expectations.foreach { case (version, expected) =>
assert(
RegrSparkVersions.r2DegenerateCasesSwapped(version) == expected,
s"regr_r2 degenerate-case swap for Spark $version should be $expected")
}
}

test("avg/sum overflow on decimal(38, _)") {
val table = "overflow_decimal_38"
withTable(table) {
Expand Down
Loading