From 1299da8b9d9134a9eb5637f15d7759f773bc36c5 Mon Sep 17 00:00:00 2001 From: rich7420 Date: Mon, 21 Sep 2026 13:48:04 +0800 Subject: [PATCH 1/9] fix: match Spark variance updates for large nearby values --- .../spark-expr/src/agg_funcs/correlation.rs | 12 +++-- native/spark-expr/src/agg_funcs/regr.rs | 6 ++- native/spark-expr/src/agg_funcs/stddev.rs | 5 ++ native/spark-expr/src/agg_funcs/variance.rs | 49 ++++++++++++++++++- native/spark-expr/src/agg_funcs/welford.rs | 24 +++++++-- .../comet/exec/CometAggregateSuite.scala | 40 +++++++++++++++ 6 files changed, 126 insertions(+), 10 deletions(-) diff --git a/native/spark-expr/src/agg_funcs/correlation.rs b/native/spark-expr/src/agg_funcs/correlation.rs index d69ac61def7..4dbdef6b51b 100644 --- a/native/spark-expr/src/agg_funcs/correlation.rs +++ b/native/spark-expr/src/agg_funcs/correlation.rs @@ -144,8 +144,10 @@ impl CorrelationAccumulator { pub fn try_new(null_on_divide_by_zero: bool) -> Result { Ok(Self { covar: CovarianceAccumulator::try_new(StatsType::Population, null_on_divide_by_zero)?, - stddev1: StddevAccumulator::try_new(StatsType::Population, null_on_divide_by_zero)?, - stddev2: StddevAccumulator::try_new(StatsType::Population, null_on_divide_by_zero)?, + stddev1: StddevAccumulator::try_new(StatsType::Population, null_on_divide_by_zero)? + .with_pearson_update(), + stddev2: StddevAccumulator::try_new(StatsType::Population, null_on_divide_by_zero)? + .with_pearson_update(), null_on_divide_by_zero, }) } @@ -280,8 +282,10 @@ impl CorrelationGroupsAccumulator { // that intent explicit. Self { covar: CovarianceGroupsAccumulator::new(StatsType::Population, false), - var1: VarianceGroupsAccumulator::new(StatsType::Population, false), - var2: VarianceGroupsAccumulator::new(StatsType::Population, false), + var1: VarianceGroupsAccumulator::new(StatsType::Population, false) + .with_pearson_update(), + var2: VarianceGroupsAccumulator::new(StatsType::Population, false) + .with_pearson_update(), null_on_divide_by_zero, } } diff --git a/native/spark-expr/src/agg_funcs/regr.rs b/native/spark-expr/src/agg_funcs/regr.rs index 72511fc9874..7c73181708d 100644 --- a/native/spark-expr/src/agg_funcs/regr.rs +++ b/native/spark-expr/src/agg_funcs/regr.rs @@ -306,8 +306,10 @@ impl RegrR2Accumulator { fn try_new(constant_dependent_is_perfect_fit: bool) -> Result { Ok(Self { covar: CovarianceAccumulator::try_new(StatsType::Population, false)?, - var_y: VarianceAccumulator::try_new(StatsType::Population, false)?, - var_x: VarianceAccumulator::try_new(StatsType::Population, false)?, + var_y: VarianceAccumulator::try_new(StatsType::Population, false)? + .with_pearson_update(), + var_x: VarianceAccumulator::try_new(StatsType::Population, false)? + .with_pearson_update(), constant_dependent_is_perfect_fit, }) } diff --git a/native/spark-expr/src/agg_funcs/stddev.rs b/native/spark-expr/src/agg_funcs/stddev.rs index bbceaa72dcd..1ef31f7a10e 100644 --- a/native/spark-expr/src/agg_funcs/stddev.rs +++ b/native/spark-expr/src/agg_funcs/stddev.rs @@ -163,6 +163,11 @@ impl StddevAccumulator { pub fn get_m2(&self) -> f64 { self.variance.get_m2() } + + pub(super) fn with_pearson_update(mut self) -> Self { + self.variance = self.variance.with_pearson_update(); + self + } } impl Accumulator for StddevAccumulator { diff --git a/native/spark-expr/src/agg_funcs/variance.rs b/native/spark-expr/src/agg_funcs/variance.rs index 57a8f6da501..146a6dd1c39 100644 --- a/native/spark-expr/src/agg_funcs/variance.rs +++ b/native/spark-expr/src/agg_funcs/variance.rs @@ -29,6 +29,8 @@ use datafusion::physical_expr::expressions::StatsType; use std::mem::size_of; use std::sync::Arc; +use super::welford::VarianceUpdate; + /// VAR_SAMP and VAR_POP aggregate expression /// The implementation mostly is the same as the DataFusion's implementation. The reason /// we have our own implementation is that DataFusion has UInt64 for state_field `count`, @@ -144,6 +146,7 @@ pub struct VarianceAccumulator { count: f64, stats_type: StatsType, null_on_divide_by_zero: bool, + update: VarianceUpdate, } impl VarianceAccumulator { @@ -155,9 +158,15 @@ impl VarianceAccumulator { count: 0_f64, stats_type: s_type, null_on_divide_by_zero, + update: VarianceUpdate::CentralMoment, }) } + pub(super) fn with_pearson_update(mut self) -> Self { + self.update = VarianceUpdate::Pearson; + self + } + pub fn get_count(&self) -> f64 { self.count } @@ -184,7 +193,8 @@ impl Accumulator for VarianceAccumulator { let arr = downcast_value!(&values[0], Float64Array).iter().flatten(); for value in arr { - let (c, m, m2) = super::welford::variance_update(self.count, self.mean, self.m2, value); + let (c, m, m2) = + super::welford::variance_update(self.count, self.mean, self.m2, value, self.update); self.count = c; self.mean = m; self.m2 = m2; @@ -271,6 +281,7 @@ pub(crate) struct VarianceGroupsAccumulator { pub(super) m2s: Vec, stats_type: StatsType, null_on_divide_by_zero: bool, + update: VarianceUpdate, } impl VarianceGroupsAccumulator { @@ -281,9 +292,15 @@ impl VarianceGroupsAccumulator { m2s: Vec::new(), stats_type, null_on_divide_by_zero, + update: VarianceUpdate::CentralMoment, } } + pub(super) fn with_pearson_update(mut self) -> Self { + self.update = VarianceUpdate::Pearson; + self + } + fn resize(&mut self, total_num_groups: usize) { self.counts.resize(total_num_groups, 0.0); self.means.resize(total_num_groups, 0.0); @@ -327,6 +344,7 @@ impl GroupsAccumulator for VarianceGroupsAccumulator { self.means[group_index], self.m2s[group_index], value, + self.update, ); self.counts[group_index] = c; self.means[group_index] = m; @@ -421,6 +439,35 @@ mod groups_tests { .collect() } + #[test] + fn large_offset_variance() { + for pair in [[1e16, 1e16 + 2.0], [1e16 + 2.0, 1e16], [-1e16, -1e16 - 2.0]] { + for (stats, expected) in [(StatsType::Population, 1.0), (StatsType::Sample, 2.0)] { + let values: ArrayRef = + Arc::new(Float64Array::from(vec![Some(pair[0]), None, Some(pair[1])])); + for batch_size in [1, 3] { + let mut scalar = VarianceAccumulator::try_new(stats, true).unwrap(); + let mut grouped = VarianceGroupsAccumulator::new(stats, true); + for offset in (0..3).step_by(batch_size) { + let batch = [values.slice(offset, batch_size)]; + scalar.update_batch(&batch).unwrap(); + grouped + .update_batch(&batch, &vec![0; batch_size], None, 2) + .unwrap(); + } + assert_eq!( + scalar.evaluate().unwrap(), + ScalarValue::Float64(Some(expected)) + ); + let states = grouped.state(EmitTo::All).unwrap(); + let mut merged = VarianceGroupsAccumulator::new(stats, true); + merged.merge_batch(&states, &[0, 1], 2).unwrap(); + assert_eq!(evaluate(&mut merged), vec![Some(expected), None]); + } + } + } + } + #[test] fn pop_variance_single_group() { let mut acc = pop_acc(); diff --git a/native/spark-expr/src/agg_funcs/welford.rs b/native/spark-expr/src/agg_funcs/welford.rs index bcc44b29887..04fbec7d72e 100644 --- a/native/spark-expr/src/agg_funcs/welford.rs +++ b/native/spark-expr/src/agg_funcs/welford.rs @@ -23,12 +23,30 @@ use arrow::buffer::NullBuffer; use datafusion::physical_expr::expressions::StatsType; +#[derive(Debug, Clone, Copy)] +pub(super) enum VarianceUpdate { + CentralMoment, + Pearson, +} + #[inline] -pub(crate) fn variance_update(count: f64, mean: f64, m2: f64, value: f64) -> (f64, f64, f64) { +pub(super) fn variance_update( + count: f64, + mean: f64, + m2: f64, + value: f64, + update: VarianceUpdate, +) -> (f64, f64, f64) { let new_count = count + 1.0; let delta1 = value - mean; - let new_mean = delta1 / new_count + mean; - let delta2 = value - new_mean; + let delta_n = delta1 / new_count; + let new_mean = mean + delta_n; + // Match Spark's CentralMomentAgg without subtracting the rounded new mean. + // PearsonCorrelation (also used by regr_r2) deliberately uses that subtraction. + let delta2 = match update { + VarianceUpdate::CentralMoment => delta1 - delta_n, + VarianceUpdate::Pearson => value - new_mean, + }; let new_m2 = m2 + delta1 * delta2; (new_count, new_mean, new_m2) } 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 55c2ebdb578..6b66e0e5b16 100644 --- a/spark/src/test/scala/org/apache/comet/exec/CometAggregateSuite.scala +++ b/spark/src/test/scala/org/apache/comet/exec/CometAggregateSuite.scala @@ -2343,6 +2343,46 @@ class CometAggregateSuite extends CometTestBase with AdaptiveSparkPlanHelper { } } + test("statistical aggregates with large nearby values") { + withSQLConf( + SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false", + SQLConf.SHUFFLE_PARTITIONS.key -> "1", + "spark.sql.files.minPartitionNum" -> "1", + CometConf.COMET_SHUFFLE_ENABLED.key -> "true", + CometConf.COMET_SHUFFLE_MODE.key -> "native") { + for (values <- Seq(Seq(1e16, 1e16 + 2), Seq(1e16 + 2, 1e16), Seq(-1e16, -1e16 - 2))) { + // One ordered file keeps both values in the same partial accumulator. Splitting + // them across files would only exercise merging two single-row states. + withTempPath { path => + (Seq(Some(values.head), None, Some(values.last))) + .map(v => (0, v)) + .toDF("g", "v") + .coalesce(1) + .write + .parquet(path.getCanonicalPath) + withParquetTable(path.getCanonicalPath, "large_moments") { + for (groupBy <- Seq("", " GROUP BY g")) { + val query = "SELECT var_pop(v), var_samp(v), stddev_pop(v), stddev_samp(v) " + + "FROM large_moments" + groupBy + val (_, cometPlan) = checkSparkAnswerAndOperator(query) + val aggregates = cometPlan.collect { case a: CometHashAggregateExec => a } + assert(aggregates.exists(_.modes.contains(Partial))) + assert(aggregates.exists(_.modes.contains(Final))) + checkAnswer(sql(query), Seq(Row(1.0, 2.0, 1.0, math.sqrt(2.0)))) + + // CORR and REGR_R2 use PearsonCorrelation's update, while REGR_SXX/SYY + // and the variance used by slope/intercept follow CentralMomentAgg. + checkSparkAnswerWithTolAndNumOfAggregates( + "SELECT corr(v, v), regr_r2(v, v), regr_sxx(v, v), regr_syy(v, v), " + + "regr_slope(v, v), regr_intercept(v, v) FROM large_moments" + groupBy, + 2) + } + } + } + } + } + } + test("var_pop and var_samp") { withSQLConf(CometConf.COMET_SHUFFLE_ENABLED.key -> "true") { Seq("native", "jvm").foreach { cometShuffleMode => From 72f12f2cb71c7130a3e842b302a369b256a0ba15 Mon Sep 17 00:00:00 2001 From: rich7420 Date: Wed, 23 Sep 2026 18:15:06 +0800 Subject: [PATCH 2/9] fix: match Spark statistical aggregate merges --- native/spark-expr/src/agg_funcs/covariance.rs | 28 +++++++++ native/spark-expr/src/agg_funcs/variance.rs | 58 +++++++++++++++++++ native/spark-expr/src/agg_funcs/welford.rs | 34 ++++++++--- .../comet/exec/CometAggregateSuite.scala | 42 ++++++++++++++ 4 files changed, 154 insertions(+), 8 deletions(-) diff --git a/native/spark-expr/src/agg_funcs/covariance.rs b/native/spark-expr/src/agg_funcs/covariance.rs index 548c1187208..058c2576f5b 100644 --- a/native/spark-expr/src/agg_funcs/covariance.rs +++ b/native/spark-expr/src/agg_funcs/covariance.rs @@ -513,6 +513,34 @@ mod groups_tests { assert!((evaluate(&mut acc)[0].unwrap() - 4.0).abs() < 1e-12); } + #[test] + fn large_offset_covariance_merge() { + let mut scalar = CovarianceAccumulator::try_new(StatsType::Population, true).unwrap(); + let mut grouped = pop(); + for (value, count) in [(0.0, 0), (1e17 - 32.0, 3), (1e17 - 16.0, 2), (0.0, 0)] { + let mut xs = vec![Some(value); count]; + xs.push(None); + let ys = xs.iter().map(|v| v.map(|v| -v)).collect::>(); + let values: Vec = vec![ + Arc::new(Float64Array::from(xs)), + Arc::new(Float64Array::from(ys)), + ]; + let mut partial = pop(); + partial + .update_batch(&values, &vec![0; count + 1], None, 2) + .unwrap(); + let state = partial.state(EmitTo::All).unwrap(); + scalar.merge_batch(&state).unwrap(); + grouped.merge_batch(&state, &[0, 1], 2).unwrap(); + } + let expected = -61.44000000000001; + assert_eq!( + scalar.evaluate().unwrap(), + ScalarValue::Float64(Some(expected)) + ); + assert_eq!(evaluate(&mut grouped), vec![Some(expected), None]); + } + #[test] fn null_in_either_column_skipped() { let mut acc = pop(); diff --git a/native/spark-expr/src/agg_funcs/variance.rs b/native/spark-expr/src/agg_funcs/variance.rs index 146a6dd1c39..0c722d71de0 100644 --- a/native/spark-expr/src/agg_funcs/variance.rs +++ b/native/spark-expr/src/agg_funcs/variance.rs @@ -468,6 +468,64 @@ mod groups_tests { } } + #[test] + fn large_offset_variance_merge() { + for (partitions, population, sample) in [ + ( + [(1e17 - 32.0, 3), (1e17 - 16.0, 2)], + 61.44000000000001, + 76.80000000000001, + ), + ([(1e17 - 96.0, 3), (1e17 - 32.0, 3)], 1024.0, 1228.8), + ] { + for sign in [1.0, -1.0] { + for reverse in [false, true] { + let mut partitions = partitions; + if reverse { + partitions.reverse(); + } + for (stats, expected) in [ + (StatsType::Population, population), + (StatsType::Sample, sample), + ] { + let mut scalar = VarianceAccumulator::try_new(stats, true).unwrap(); + let mut grouped = VarianceGroupsAccumulator::new(stats, true); + // Include empty partials before and after the nonempty states. + for (value, count) in + [(0.0, 0)].into_iter().chain(partitions).chain([(0.0, 0)]) + { + let mut values = vec![Some(sign * value); count]; + values.push(None); + let values: ArrayRef = Arc::new(Float64Array::from(values)); + let mut partial = VarianceAccumulator::try_new(stats, true).unwrap(); + partial.update_batch(&[Arc::clone(&values)]).unwrap(); + let state = partial + .state() + .unwrap() + .iter() + .map(|v| v.to_array_of_size(1).unwrap()) + .collect::>(); + scalar.merge_batch(&state).unwrap(); + + let mut partial = VarianceGroupsAccumulator::new(stats, true); + partial + .update_batch(&[values], &vec![0; count + 1], None, 2) + .unwrap(); + grouped + .merge_batch(&partial.state(EmitTo::All).unwrap(), &[0, 1], 2) + .unwrap(); + } + assert_eq!( + scalar.evaluate().unwrap(), + ScalarValue::Float64(Some(expected)) + ); + assert_eq!(evaluate(&mut grouped), vec![Some(expected), None]); + } + } + } + } + } + #[test] fn pop_variance_single_group() { let mut acc = pop_acc(); diff --git a/native/spark-expr/src/agg_funcs/welford.rs b/native/spark-expr/src/agg_funcs/welford.rs index 04fbec7d72e..dc3f8551e30 100644 --- a/native/spark-expr/src/agg_funcs/welford.rs +++ b/native/spark-expr/src/agg_funcs/welford.rs @@ -71,9 +71,16 @@ pub(crate) fn variance_merge( m2_b: f64, ) -> (f64, f64, f64) { let new_count = count_a + count_b; - let new_mean = mean_a * count_a / new_count + mean_b * count_b / new_count; - let delta = mean_a - mean_b; - let new_m2 = m2_a + m2_b + delta * delta * count_a * count_b / new_count; + // CentralMomentAgg and PearsonCorrelation use the same merge expressions. + // Preserve Spark's operation order to avoid rounding large means differently. + let delta = mean_b - mean_a; + let delta_n = if new_count == 0.0 { + 0.0 + } else { + delta / new_count + }; + let new_mean = mean_a + delta_n * count_b; + let new_m2 = m2_a + m2_b + delta * delta_n * count_a * count_b; (new_count, new_mean, new_m2) } @@ -167,10 +174,21 @@ pub(crate) fn covariance_merge( c_b: f64, ) -> (f64, f64, f64, f64) { let new_count = count_a + count_b; - let new_mean1 = mean1_a * count_a / new_count + mean1_b * count_b / new_count; - let new_mean2 = mean2_a * count_a / new_count + mean2_b * count_b / new_count; - let delta1 = mean1_a - mean1_b; - let delta2 = mean2_a - mean2_b; - let new_c = c_a + c_b + delta1 * delta2 * count_a * count_b / new_count; + // Keep covariance aligned with variance and Spark's Covariance/PearsonCorrelation. + let delta1 = mean1_b - mean1_a; + let delta2 = mean2_b - mean2_a; + let delta1_n = if new_count == 0.0 { + 0.0 + } else { + delta1 / new_count + }; + let delta2_n = if new_count == 0.0 { + 0.0 + } else { + delta2 / new_count + }; + let new_mean1 = mean1_a + delta1_n * count_b; + let new_mean2 = mean2_a + delta2_n * count_b; + let new_c = c_a + c_b + delta1 * delta2_n * count_a * count_b; (new_count, new_mean1, new_mean2, new_c) } 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 6b66e0e5b16..5526acfe18c 100644 --- a/spark/src/test/scala/org/apache/comet/exec/CometAggregateSuite.scala +++ b/spark/src/test/scala/org/apache/comet/exec/CometAggregateSuite.scala @@ -2383,6 +2383,48 @@ class CometAggregateSuite extends CometTestBase with AdaptiveSparkPlanHelper { } } + test("statistical aggregates merge large nearby values across partitions") { + withSQLConf( + SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false", + SQLConf.SHUFFLE_PARTITIONS.key -> "1", + SQLConf.FILES_MAX_PARTITION_BYTES.key -> "1048576", + SQLConf.FILES_OPEN_COST_IN_BYTES.key -> "1048576", + CometConf.COMET_SHUFFLE_ENABLED.key -> "true", + CometConf.COMET_SHUFFLE_MODE.key -> "native") { + withTempPath { path => + // Two constant-valued files produce separate partials with zero M2. The + // old merge returns 576 instead of 1024, regardless of which partial arrives first. + for (value <- Seq(1e17 - 96, 1e17 - 32)) { + (Seq.fill(3)((0, Option(value))) ++ Seq((0, None), (1, None))) + .toDF("g", "v") + .coalesce(1) + .write + .mode("append") + .parquet(path.getCanonicalPath) + } + withParquetTable(path.getCanonicalPath, "merged_moments") { + assert(spark.table("merged_moments").rdd.getNumPartitions == 2) + for (groupBy <- Seq("", " GROUP BY g")) { + val query = "SELECT var_pop(v), var_samp(v), stddev_pop(v), stddev_samp(v) " + + "FROM merged_moments" + groupBy + val (_, cometPlan) = checkSparkAnswerAndOperator(query) + val aggregates = cometPlan.collect { case a: CometHashAggregateExec => a } + assert(aggregates.exists(_.modes.contains(Partial))) + assert(aggregates.exists(_.modes.contains(Final))) + val expected = Seq(Row(1024.0, 1228.8, 32.0, math.sqrt(1228.8))) ++ + (if (groupBy.isEmpty) Seq.empty else Seq(Row(null, null, null, null))) + checkAnswer(sql(query), expected) + checkSparkAnswerWithTolAndNumOfAggregates( + "SELECT covar_pop(v, -v), covar_samp(v, -v), corr(v, -v), regr_r2(v, -v), " + + "regr_sxx(v, -v), regr_syy(v, -v), regr_sxy(v, -v), " + + "regr_slope(v, -v), regr_intercept(v, -v) FROM merged_moments" + groupBy, + 2) + } + } + } + } + } + test("var_pop and var_samp") { withSQLConf(CometConf.COMET_SHUFFLE_ENABLED.key -> "true") { Seq("native", "jvm").foreach { cometShuffleMode => From a16b7e13bf038827d73be1ebdbb95f9907c1ad12 Mon Sep 17 00:00:00 2001 From: rich7420 Date: Mon, 28 Sep 2026 22:26:15 +0800 Subject: [PATCH 3/9] fix: preserve Spark semantics when merging empty partials --- .../spark-expr/src/agg_funcs/correlation.rs | 32 ++++++++ native/spark-expr/src/agg_funcs/covariance.rs | 64 ++++++++++++++-- native/spark-expr/src/agg_funcs/regr.rs | 39 ++++++++++ native/spark-expr/src/agg_funcs/stddev.rs | 34 +++++++++ native/spark-expr/src/agg_funcs/variance.rs | 73 +++++++++++++++++-- 5 files changed, 230 insertions(+), 12 deletions(-) diff --git a/native/spark-expr/src/agg_funcs/correlation.rs b/native/spark-expr/src/agg_funcs/correlation.rs index 4dbdef6b51b..8ea10a4c0fb 100644 --- a/native/spark-expr/src/agg_funcs/correlation.rs +++ b/native/spark-expr/src/agg_funcs/correlation.rs @@ -477,6 +477,38 @@ mod groups_tests { .collect() } + #[test] + fn correlation_merge_empty_partials() { + for empty_first in [false, true] { + let mut scalar = CorrelationAccumulator::try_new(true).unwrap(); + let mut grouped = CorrelationGroupsAccumulator::new(true); + let counts = if empty_first { + [0.0, 100.0] + } else { + [100.0, 0.0] + }; + for count in counts { + let mean = if count == 0.0 { 0.0 } else { 1e155 }; + let state: Vec = [count, mean, mean, 0.0, 0.0, 0.0] + .into_iter() + .map(|v| Arc::new(Float64Array::from(vec![v])) as ArrayRef) + .collect(); + scalar.merge_batch(&state).unwrap(); + grouped.merge_batch(&state, &[0], 1).unwrap(); + } + let ScalarValue::Float64(scalar) = scalar.evaluate().unwrap() else { + panic!("expected a double correlation"); + }; + for result in [scalar, evaluate(&mut grouped)[0]] { + if empty_first { + assert_eq!(result, None); + } else { + assert!(result.unwrap().is_nan()); + } + } + } + } + #[test] fn perfectly_correlated_single_group() { let mut a = acc(true); diff --git a/native/spark-expr/src/agg_funcs/covariance.rs b/native/spark-expr/src/agg_funcs/covariance.rs index 058c2576f5b..28e09574a8e 100644 --- a/native/spark-expr/src/agg_funcs/covariance.rs +++ b/native/spark-expr/src/agg_funcs/covariance.rs @@ -283,9 +283,7 @@ impl Accumulator for CovarianceAccumulator { for i in 0..counts.len() { let c = counts.value(i); - if c == 0.0 { - continue; - } + // Even empty partials affect Spark's NaN propagation during merge. let (new_count, new_mean1, new_mean2, new_c) = super::welford::covariance_merge( self.count, self.mean1, @@ -428,9 +426,7 @@ impl GroupsAccumulator for CovarianceGroupsAccumulator { for (i, &group_index) in group_indices.iter().enumerate() { let partial_count = counts.value(i); - if partial_count == 0.0 { - continue; - } + // Even empty partials affect Spark's NaN propagation during merge. let (new_count, new_m1, new_m2, new_c) = super::welford::covariance_merge( self.counts[group_index], self.mean1s[group_index], @@ -541,6 +537,62 @@ mod groups_tests { assert_eq!(evaluate(&mut grouped), vec![Some(expected), None]); } + #[test] + fn covariance_merge_empty_partials() { + for value in [1e154, -1e154, 1e155, -1e155] { + for empty_first in [false, true] { + for stats in [StatsType::Population, StatsType::Sample] { + let mut scalar = CovarianceAccumulator::try_new(stats, true).unwrap(); + let mut grouped = CovarianceGroupsAccumulator::new(stats, true); + let counts = if empty_first { [0, 100] } else { [100, 0] }; + for count in counts { + let xs: ArrayRef = Arc::new(Float64Array::from( + (0..101) + .map(|i| (i < count).then_some(value)) + .collect::>(), + )); + let ys: ArrayRef = Arc::new(Float64Array::from(vec![Some(-value); 101])); + // A null in just one input still produces a zero-count partial. + let values = [xs, ys]; + let mut partial = CovarianceAccumulator::try_new(stats, true).unwrap(); + partial.update_batch(&values).unwrap(); + let state = partial + .state() + .unwrap() + .iter() + .map(|v| v.to_array_of_size(1).unwrap()) + .collect::>(); + scalar.merge_batch(&state).unwrap(); + + let mut partial = CovarianceGroupsAccumulator::new(stats, true); + let mut group_indices = vec![0; 101]; + // The second group has no valid input pairs in either partial. + group_indices[100] = 1; + partial + .update_batch(&values, &group_indices, None, 2) + .unwrap(); + grouped + .merge_batch(&partial.state(EmitTo::All).unwrap(), &[0, 1], 2) + .unwrap(); + } + let expected_nan = value.abs() == 1e155 && !empty_first; + let ScalarValue::Float64(Some(actual)) = scalar.evaluate().unwrap() else { + panic!("expected a non-null covariance"); + }; + let grouped = evaluate(&mut grouped); + for result in [actual, grouped[0].unwrap()] { + if expected_nan { + assert!(result.is_nan(), "value={value}, empty_first={empty_first}"); + } else { + assert_eq!(result, 0.0); + } + } + assert_eq!(grouped[1], None); + } + } + } + } + #[test] fn null_in_either_column_skipped() { let mut acc = pop(); diff --git a/native/spark-expr/src/agg_funcs/regr.rs b/native/spark-expr/src/agg_funcs/regr.rs index 7c73181708d..012775f9414 100644 --- a/native/spark-expr/src/agg_funcs/regr.rs +++ b/native/spark-expr/src/agg_funcs/regr.rs @@ -535,6 +535,45 @@ mod tests { ); } + #[test] + fn regr_merge_empty_partials() { + for kind in [ + RegrType::SXX, + RegrType::SYY, + RegrType::SXY, + RegrType::R2, + RegrType::Slope, + RegrType::Intercept, + ] { + for empty_first in [false, true] { + let mut merged = acc(kind); + let counts = if empty_first { [0, 100] } else { [100, 0] }; + for count in counts { + let mut partial = acc(kind); + let values = vec![Some(1e155); count]; + partial.update_batch(&cols(values.clone(), values)).unwrap(); + let state = partial + .state() + .unwrap() + .iter() + .map(|v| v.to_array_of_size(1).unwrap()) + .collect::>(); + merged.merge_batch(&state).unwrap(); + } + let ScalarValue::Float64(result) = merged.evaluate().unwrap() else { + panic!("expected a double regression result"); + }; + if !empty_first { + assert!(result.unwrap().is_nan(), "{kind:?}"); + } else if matches!(kind, RegrType::SXX | RegrType::SYY | RegrType::SXY) { + assert_eq!(result, Some(0.0)); + } else { + assert_eq!(result, None); + } + } + } + } + fn perfect_line() -> (Vec>, Vec>) { // y = 2x + 1 let x = vec![1.0, 2.0, 3.0, 4.0, 5.0]; diff --git a/native/spark-expr/src/agg_funcs/stddev.rs b/native/spark-expr/src/agg_funcs/stddev.rs index 1ef31f7a10e..8ef9411119e 100644 --- a/native/spark-expr/src/agg_funcs/stddev.rs +++ b/native/spark-expr/src/agg_funcs/stddev.rs @@ -273,6 +273,40 @@ mod groups_tests { use arrow::array::AsArray; use arrow::datatypes::Float64Type; + #[test] + fn stddev_merge_empty_partials() { + for empty_first in [false, true] { + for stats in [StatsType::Population, StatsType::Sample] { + let mut scalar = StddevAccumulator::try_new(stats, true).unwrap(); + let mut grouped = StddevGroupsAccumulator::new(stats, true); + let counts = if empty_first { + [0.0, 100.0] + } else { + [100.0, 0.0] + }; + for count in counts { + let state: Vec = [count, if count == 0.0 { 0.0 } else { 1e155 }, 0.0] + .into_iter() + .map(|v| Arc::new(Float64Array::from(vec![v])) as ArrayRef) + .collect(); + scalar.merge_batch(&state).unwrap(); + grouped.merge_batch(&state, &[0], 1).unwrap(); + } + let ScalarValue::Float64(Some(scalar)) = scalar.evaluate().unwrap() else { + panic!("expected a non-null standard deviation"); + }; + let grouped = grouped.evaluate(EmitTo::All).unwrap(); + for result in [scalar, grouped.as_primitive::().value(0)] { + if empty_first { + assert_eq!(result, 0.0); + } else { + assert!(result.is_nan()); + } + } + } + } + } + #[test] fn pop_stddev_single_group() { let mut acc = StddevGroupsAccumulator::new(StatsType::Population, false); diff --git a/native/spark-expr/src/agg_funcs/variance.rs b/native/spark-expr/src/agg_funcs/variance.rs index 0c722d71de0..4d89bff5696 100644 --- a/native/spark-expr/src/agg_funcs/variance.rs +++ b/native/spark-expr/src/agg_funcs/variance.rs @@ -224,9 +224,7 @@ impl Accumulator for VarianceAccumulator { for i in 0..counts.len() { let c = counts.value(i); - if c == 0_f64 { - continue; - } + // Even empty partials affect Spark's NaN propagation during merge. let (new_count, new_mean, new_m2) = super::welford::variance_merge( self.count, self.mean, @@ -369,9 +367,7 @@ impl GroupsAccumulator for VarianceGroupsAccumulator { for (i, &group_index) in group_indices.iter().enumerate() { let partial_count = partial_counts.value(i); - if partial_count == 0.0 { - continue; - } + // Even empty partials affect Spark's NaN propagation during merge. let (new_count, new_mean, new_m2) = super::welford::variance_merge( self.counts[group_index], self.means[group_index], @@ -526,6 +522,71 @@ mod groups_tests { } } + #[test] + fn variance_merge_empty_partials() { + for value in [1e154, -1e154, 1e155, -1e155] { + for empty_first in [false, true] { + for stats in [StatsType::Population, StatsType::Sample] { + for update in [VarianceUpdate::CentralMoment, VarianceUpdate::Pearson] { + let mut scalar = VarianceAccumulator::try_new(stats, true).unwrap(); + let mut grouped = VarianceGroupsAccumulator::new(stats, true); + scalar.update = update; + grouped.update = update; + let counts = if empty_first { [0, 100] } else { [100, 0] }; + for count in counts { + let values: ArrayRef = Arc::new(Float64Array::from( + (0..101) + .map(|i| (i < count).then_some(value)) + .collect::>(), + )); + let mut partial = VarianceAccumulator::try_new(stats, true).unwrap(); + partial.update = update; + partial.update_batch(&[Arc::clone(&values)]).unwrap(); + let state = partial + .state() + .unwrap() + .iter() + .map(|v| v.to_array_of_size(1).unwrap()) + .collect::>(); + scalar.merge_batch(&state).unwrap(); + + let mut partial = VarianceGroupsAccumulator::new(stats, true); + partial.update = update; + let mut group_indices = vec![0; 101]; + // Keep a separate all-null group alongside the non-empty group. + group_indices[100] = 1; + partial + .update_batch(&[values], &group_indices, None, 2) + .unwrap(); + grouped + .merge_batch(&partial.state(EmitTo::All).unwrap(), &[0, 1], 2) + .unwrap(); + } + // Spark evaluates delta * deltaN * n1 * n2 even when n2 is zero. + // For 1e155, the third multiplication overflows before multiplying by + // zero. Empty-first and the smaller-magnitude control remain zero. + let expected_nan = value.abs() == 1e155 && !empty_first; + let ScalarValue::Float64(Some(actual)) = scalar.evaluate().unwrap() else { + panic!("expected a non-null variance"); + }; + let grouped = evaluate(&mut grouped); + for result in [actual, grouped[0].unwrap()] { + if expected_nan { + assert!( + result.is_nan(), + "value={value}, empty_first={empty_first}" + ); + } else { + assert_eq!(result, 0.0); + } + } + assert_eq!(grouped[1], None); + } + } + } + } + } + #[test] fn pop_variance_single_group() { let mut acc = pop_acc(); From d89c385e674026fcae52ab5ae256c77855bdf2d6 Mon Sep 17 00:00:00 2001 From: rich7420 Date: Wed, 30 Sep 2026 21:23:51 +0800 Subject: [PATCH 4/9] test: cover regression aggregates merging fractional constants --- native/spark-expr/src/agg_funcs/regr.rs | 45 +++++++++++++++++++ .../comet/exec/CometAggregateSuite.scala | 43 ++++++++++++++++++ 2 files changed, 88 insertions(+) diff --git a/native/spark-expr/src/agg_funcs/regr.rs b/native/spark-expr/src/agg_funcs/regr.rs index 012775f9414..0b8838c985b 100644 --- a/native/spark-expr/src/agg_funcs/regr.rs +++ b/native/spark-expr/src/agg_funcs/regr.rs @@ -721,6 +721,51 @@ mod tests { } } + #[test] + fn merge_fractional_constant_partials() { + // #6423: the first merge into an empty final buffer used to round the + // constant mean away from 0.1. Later merges then produced nonzero M2. + for (kind, constant_dependent, expected) in [ + (RegrType::Slope, false, None), + (RegrType::Intercept, false, None), + (RegrType::R2, false, None), + (RegrType::R2, true, Some(1.0)), + (RegrType::SXX, false, Some(0.0)), + (RegrType::SYY, true, Some(0.0)), + (RegrType::SXY, false, Some(0.0)), + ] { + for reverse in [false, true] { + let mut merged = acc(kind); + let starts = if reverse { [3, 0] } else { [0, 3] }; + for start in starts { + let varying = (start..start + 3).map(|v| Some(v as f64)).collect(); + let constant = vec![Some(0.1); 3]; + let values = match kind { + // Spark's RegrReplacement duplicates the selected column. + RegrType::SXX | RegrType::SYY => cols(constant.clone(), constant), + _ if constant_dependent => cols(constant, varying), + _ => cols(varying, constant), + }; + let mut partial = acc(kind); + partial.update_batch(&values).unwrap(); + let state = partial + .state() + .unwrap() + .iter() + .map(|v| v.to_array_of_size(1).unwrap()) + .collect::>(); + merged.merge_batch(&state).unwrap(); + } + // A tolerance would hide the tiny nonzero moments behind the bug. + assert_eq!( + merged.evaluate().unwrap(), + ScalarValue::Float64(expected), + "{kind:?}, constant_dependent={constant_dependent}, reverse={reverse}" + ); + } + } + } + #[test] fn merge_matches_single_batch() { let (y, x) = perfect_line(); 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 5526acfe18c..504407659f6 100644 --- a/spark/src/test/scala/org/apache/comet/exec/CometAggregateSuite.scala +++ b/spark/src/test/scala/org/apache/comet/exec/CometAggregateSuite.scala @@ -2425,6 +2425,49 @@ class CometAggregateSuite extends CometTestBase with AdaptiveSparkPlanHelper { } } + test("statistical aggregates merge fractional constants across partitions") { + // https://github.com/apache/datafusion-comet/issues/6423 + withSQLConf( + SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false", + SQLConf.SHUFFLE_PARTITIONS.key -> "1", + SQLConf.FILES_MAX_PARTITION_BYTES.key -> "1048576", + SQLConf.FILES_OPEN_COST_IN_BYTES.key -> "1048576", + CometConf.COMET_SHUFFLE_ENABLED.key -> "true", + CometConf.COMET_SHUFFLE_MODE.key -> "native") { + val swapped = RegrSparkVersions.r2DegenerateCasesSwapped(spark.version) + val constantX: java.lang.Double = if (swapped) null else 1.0 + val constantY: java.lang.Double = if (swapped) 1.0 else null + for (reverse <- Seq(false, true)) { + withTempPath { path => + val starts = if (reverse) Seq(3, 0) else Seq(0, 3) + for (start <- starts) { + (start until start + 3) + .map(y => (0, y.toDouble, 0.1)) + .toDF("g", "y", "x") + .coalesce(1) + .write + .mode("append") + .parquet(path.getCanonicalPath) + } + withParquetTable(path.getCanonicalPath, "fractional_constants") { + assert(spark.table("fractional_constants").rdd.getNumPartitions == 2) + for (groupBy <- Seq("", " GROUP BY g")) { + val query = "SELECT regr_slope(y, x), regr_intercept(y, x), regr_r2(y, x), " + + "regr_sxx(y, x), regr_sxy(y, x), regr_r2(x, y), regr_syy(x, y) " + + "FROM fractional_constants" + groupBy + val (_, cometPlan) = checkSparkAnswerAndOperator(query) + val aggregates = cometPlan.collect { case a: CometHashAggregateExec => a } + assert(aggregates.exists(_.modes.contains(Partial))) + assert(aggregates.exists(_.modes.contains(Final))) + // Keep the degenerate-case results exact. A tolerance can hide nonzero M2. + checkAnswer(sql(query), Seq(Row(null, null, constantX, 0.0, 0.0, constantY, 0.0))) + } + } + } + } + } + } + test("var_pop and var_samp") { withSQLConf(CometConf.COMET_SHUFFLE_ENABLED.key -> "true") { Seq("native", "jvm").foreach { cometShuffleMode => From d23925bb8281adc167320680e6b993ddb87455b5 Mon Sep 17 00:00:00 2001 From: rich7420 Date: Thu, 1 Oct 2026 14:37:59 +0800 Subject: [PATCH 5/9] fix: align scalar correlation evaluation with Spark --- .../spark-expr/src/agg_funcs/correlation.rs | 75 ++++++++++++++++--- .../comet/exec/CometAggregateSuite.scala | 51 ++++++++++++- 2 files changed, 113 insertions(+), 13 deletions(-) diff --git a/native/spark-expr/src/agg_funcs/correlation.rs b/native/spark-expr/src/agg_funcs/correlation.rs index 8ea10a4c0fb..de71b144434 100644 --- a/native/spark-expr/src/agg_funcs/correlation.rs +++ b/native/spark-expr/src/agg_funcs/correlation.rs @@ -229,10 +229,6 @@ impl Accumulator for CorrelationAccumulator { } fn evaluate(&mut self) -> Result { - let covar = self.covar.evaluate()?; - let stddev1 = self.stddev1.evaluate()?; - let stddev2 = self.stddev2.evaluate()?; - if self.covar.get_count() == 0.0 { return Ok(ScalarValue::Float64(None)); } else if self.covar.get_count() == 1.0 { @@ -242,14 +238,16 @@ impl Accumulator for CorrelationAccumulator { return Ok(ScalarValue::Float64(Some(f64::NAN))); } } - match (covar, stddev1, stddev2) { - ( - ScalarValue::Float64(Some(c)), - ScalarValue::Float64(Some(s1)), - ScalarValue::Float64(Some(s2)), - ) if s1 != 0.0 && s2 != 0.0 => Ok(ScalarValue::Float64(Some(c / (s1 * s2)))), - _ => Ok(ScalarValue::Float64(None)), + let m2_1 = self.stddev1.get_m2(); + let m2_2 = self.stddev2.get_m2(); + if m2_1 == 0.0 || m2_2 == 0.0 { + return Ok(ScalarValue::Float64(None)); } + // Match Spark and the grouped path's raw-moment evaluation. Normalizing + // first changes rounding and can avoid overflow in m2_1 * m2_2. + Ok(ScalarValue::Float64(Some( + self.covar.get_algo_const() / (m2_1 * m2_2).sqrt(), + ))) } fn size(&self) -> usize { @@ -477,6 +475,61 @@ mod groups_tests { .collect() } + #[test] + fn correlation_evaluates_raw_moments_exactly() { + // Spark divides ck by sqrt(m2_1 * m2_2). Normalizing the moments + // first changes rounding, and avoids overflow that Spark preserves. + for (values, expected) in [ + ([1e16, 1e16 + 2.0], Some(1.0)), + ([1e16 + 2.0, 1e16], None), + ([-1e16, -1e16 - 2.0], Some(1.0)), + ([1e100, 2e100], Some(0.0)), + ] { + for sign in [-1.0, 1.0] { + for null_on_divide_by_zero in [false, true] { + let input: Vec = vec![ + Arc::new(Float64Array::from(vec![ + Some(values[0]), + None, + Some(values[1]), + ])), + Arc::new(Float64Array::from(vec![ + Some(sign * values[0]), + Some(0.0), + Some(sign * values[1]), + ])), + ]; + let mut scalar = + CorrelationAccumulator::try_new(null_on_divide_by_zero).unwrap(); + let mut grouped = CorrelationGroupsAccumulator::new(null_on_divide_by_zero); + scalar.update_batch(&input).unwrap(); + grouped.update_batch(&input, &[0, 0, 0], None, 1).unwrap(); + let state = scalar + .state() + .unwrap() + .iter() + .map(|s| s.to_array_of_size(1).unwrap()) + .collect::>(); + let mut merged_scalar = + CorrelationAccumulator::try_new(null_on_divide_by_zero).unwrap(); + let mut merged_grouped = + CorrelationGroupsAccumulator::new(null_on_divide_by_zero); + merged_scalar.merge_batch(&state).unwrap(); + merged_grouped.merge_batch(&state, &[0], 1).unwrap(); + for result in [ + scalar.evaluate().unwrap(), + merged_scalar.evaluate().unwrap(), + ] { + assert_eq!(result, ScalarValue::Float64(expected.map(|v| sign * v))); + } + for result in [evaluate(&mut grouped)[0], evaluate(&mut merged_grouped)[0]] { + assert_eq!(result, expected.map(|v| sign * v)); + } + } + } + } + } + #[test] fn correlation_merge_empty_partials() { for empty_first in [false, true] { 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 504407659f6..2c68f37bb60 100644 --- a/spark/src/test/scala/org/apache/comet/exec/CometAggregateSuite.scala +++ b/spark/src/test/scala/org/apache/comet/exec/CometAggregateSuite.scala @@ -2350,7 +2350,11 @@ class CometAggregateSuite extends CometTestBase with AdaptiveSparkPlanHelper { "spark.sql.files.minPartitionNum" -> "1", CometConf.COMET_SHUFFLE_ENABLED.key -> "true", CometConf.COMET_SHUFFLE_MODE.key -> "native") { - for (values <- Seq(Seq(1e16, 1e16 + 2), Seq(1e16 + 2, 1e16), Seq(-1e16, -1e16 - 2))) { + val cases = Seq( + (Seq(1e16, 1e16 + 2), Some(1.0)), + (Seq(1e16 + 2, 1e16), None), + (Seq(-1e16, -1e16 - 2), Some(1.0))) + for ((values, expectedCorr) <- cases) { // One ordered file keeps both values in the same partial accumulator. Splitting // them across files would only exercise merging two single-row states. withTempPath { path => @@ -2372,8 +2376,13 @@ class CometAggregateSuite extends CometTestBase with AdaptiveSparkPlanHelper { // CORR and REGR_R2 use PearsonCorrelation's update, while REGR_SXX/SYY // and the variance used by slope/intercept follow CentralMomentAgg. + val corrQuery = "SELECT corr(v, v) FROM large_moments" + groupBy + checkSparkAnswerAndNumOfAggregates(corrQuery, 2) + // Reversing the positive pair rounds the Pearson mean to the second + // value, leaving zero M2 and a NULL correlation. + checkAnswer(sql(corrQuery), Seq(Row(expectedCorr.map(Double.box).orNull))) checkSparkAnswerWithTolAndNumOfAggregates( - "SELECT corr(v, v), regr_r2(v, v), regr_sxx(v, v), regr_syy(v, v), " + + "SELECT regr_r2(v, v), regr_sxx(v, v), regr_syy(v, v), " + "regr_slope(v, v), regr_intercept(v, v) FROM large_moments" + groupBy, 2) } @@ -2383,6 +2392,33 @@ class CometAggregateSuite extends CometTestBase with AdaptiveSparkPlanHelper { } } + test("statistical aggregates correlation preserves raw moment overflow") { + withSQLConf( + SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false", + SQLConf.SHUFFLE_PARTITIONS.key -> "1", + "spark.sql.files.minPartitionNum" -> "1", + CometConf.COMET_SHUFFLE_ENABLED.key -> "true", + CometConf.COMET_SHUFFLE_MODE.key -> "native") { + withTempPath { path => + Seq(Some(1e100), None, Some(2e100)) + .map(v => (0, v)) + .toDF("g", "v") + .coalesce(1) + .write + .parquet(path.getCanonicalPath) + withParquetTable(path.getCanonicalPath, "correlation_overflow") { + assert(spark.table("correlation_overflow").rdd.getNumPartitions == 1) + for (groupBy <- Seq("", " GROUP BY g")) { + val query = "SELECT corr(v, v), corr(v, -v) FROM correlation_overflow" + groupBy + checkSparkAnswerAndNumOfAggregates(query, 2) + // Spark's sqrt(m2_1 * m2_2) overflows to infinity, so corr is exactly zero. + checkAnswer(sql(query), Seq(Row(0.0, -0.0))) + } + } + } + } + } + test("statistical aggregates merge large nearby values across partitions") { withSQLConf( SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false", @@ -2427,7 +2463,9 @@ class CometAggregateSuite extends CometTestBase with AdaptiveSparkPlanHelper { test("statistical aggregates merge fractional constants across partitions") { // https://github.com/apache/datafusion-comet/issues/6423 + // https://github.com/apache/datafusion-comet/issues/6481 withSQLConf( + SQLConf.ANSI_ENABLED.key -> "false", SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false", SQLConf.SHUFFLE_PARTITIONS.key -> "1", SQLConf.FILES_MAX_PARTITION_BYTES.key -> "1048576", @@ -2461,6 +2499,15 @@ class CometAggregateSuite extends CometTestBase with AdaptiveSparkPlanHelper { assert(aggregates.exists(_.modes.contains(Final))) // Keep the degenerate-case results exact. A tolerance can hide nonzero M2. checkAnswer(sql(query), Seq(Row(null, null, constantX, 0.0, 0.0, constantY, 0.0))) + + val statsQuery = "SELECT corr(y, x), covar_pop(y, x), covar_samp(y, x), " + + "var_pop(x), var_samp(x), stddev_pop(x), stddev_samp(x) " + + "FROM fractional_constants" + groupBy + val (_, statsPlan) = checkSparkAnswerAndOperator(statsQuery) + val statsAggregates = statsPlan.collect { case a: CometHashAggregateExec => a } + assert(statsAggregates.exists(_.modes.contains(Partial))) + assert(statsAggregates.exists(_.modes.contains(Final))) + checkAnswer(sql(statsQuery), Seq(Row(null, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0))) } } } From c4baa2611fe62b48be23fb1ba2808d7d90f32c45 Mon Sep 17 00:00:00 2001 From: rich7420 Date: Thu, 1 Oct 2026 15:01:12 +0800 Subject: [PATCH 6/9] fix: preserve correlation underflow and NaN semantics --- .../spark-expr/src/agg_funcs/correlation.rs | 52 +++++++++++++++---- .../comet/exec/CometAggregateSuite.scala | 45 ++++++++++------ 2 files changed, 72 insertions(+), 25 deletions(-) diff --git a/native/spark-expr/src/agg_funcs/correlation.rs b/native/spark-expr/src/agg_funcs/correlation.rs index de71b144434..c75d61ea1e3 100644 --- a/native/spark-expr/src/agg_funcs/correlation.rs +++ b/native/spark-expr/src/agg_funcs/correlation.rs @@ -238,15 +238,16 @@ impl Accumulator for CorrelationAccumulator { return Ok(ScalarValue::Float64(Some(f64::NAN))); } } - let m2_1 = self.stddev1.get_m2(); - let m2_2 = self.stddev2.get_m2(); - if m2_1 == 0.0 || m2_2 == 0.0 { + let m2_product = self.stddev1.get_m2() * self.stddev2.get_m2(); + // The product can underflow even when both moments are nonzero. + // A zero moment paired with infinity produces NaN, not a zero denominator. + if m2_product == 0.0 { return Ok(ScalarValue::Float64(None)); } // Match Spark and the grouped path's raw-moment evaluation. Normalizing // first changes rounding and can avoid overflow in m2_1 * m2_2. Ok(ScalarValue::Float64(Some( - self.covar.get_algo_const() / (m2_1 * m2_2).sqrt(), + self.covar.get_algo_const() / m2_product.sqrt(), ))) } @@ -405,16 +406,15 @@ impl GroupsAccumulator for CorrelationGroupsAccumulator { } continue; } - // Population stats: divide m2 / count, c / count. The 1/count - // factors cancel in c / (s1 * s2), so we work with raw moments. - let s1_sq = m2_1s[i]; - let s2_sq = m2_2s[i]; - if s1_sq == 0.0 || s2_sq == 0.0 { + // Match Spark's raw-moment product, including overflow, underflow + // and NaN from a zero moment paired with infinity. + let m2_product = m2_1s[i] * m2_2s[i]; + if m2_product == 0.0 { values.push(0.0); validity.push(false); continue; } - values.push(algo_consts[i] / (s1_sq * s2_sq).sqrt()); + values.push(algo_consts[i] / m2_product.sqrt()); validity.push(true); } @@ -484,6 +484,7 @@ mod groups_tests { ([1e16 + 2.0, 1e16], None), ([-1e16, -1e16 - 2.0], Some(1.0)), ([1e100, 2e100], Some(0.0)), + ([1e-100, 2e-100], None), ] { for sign in [-1.0, 1.0] { for null_on_divide_by_zero in [false, true] { @@ -530,6 +531,37 @@ mod groups_tests { } } + #[test] + fn correlation_zero_and_infinite_moments_yield_nan() { + let input: Vec = vec![ + Arc::new(Float64Array::from(vec![1e200, -1e200])), + Arc::new(Float64Array::from(vec![0.1, 0.1])), + ]; + let mut scalar = CorrelationAccumulator::try_new(true).unwrap(); + let mut grouped = CorrelationGroupsAccumulator::new(true); + scalar.update_batch(&input).unwrap(); + grouped.update_batch(&input, &[0, 0], None, 1).unwrap(); + let state = scalar + .state() + .unwrap() + .iter() + .map(|s| s.to_array_of_size(1).unwrap()) + .collect::>(); + let mut merged_scalar = CorrelationAccumulator::try_new(true).unwrap(); + let mut merged_grouped = CorrelationGroupsAccumulator::new(true); + merged_scalar.merge_batch(&state).unwrap(); + merged_grouped.merge_batch(&state, &[0], 1).unwrap(); + for result in [ + scalar.evaluate().unwrap(), + merged_scalar.evaluate().unwrap(), + ] { + assert!(matches!(result, ScalarValue::Float64(Some(v)) if v.is_nan())); + } + for result in [evaluate(&mut grouped)[0], evaluate(&mut merged_grouped)[0]] { + assert!(result.unwrap().is_nan()); + } + } + #[test] fn correlation_merge_empty_partials() { for empty_first in [false, true] { 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 2c68f37bb60..ff76e293362 100644 --- a/spark/src/test/scala/org/apache/comet/exec/CometAggregateSuite.scala +++ b/spark/src/test/scala/org/apache/comet/exec/CometAggregateSuite.scala @@ -2392,27 +2392,42 @@ class CometAggregateSuite extends CometTestBase with AdaptiveSparkPlanHelper { } } - test("statistical aggregates correlation preserves raw moment overflow") { + test("statistical aggregates correlation uses raw moments at extreme magnitudes") { withSQLConf( + SQLConf.ANSI_ENABLED.key -> "false", SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false", SQLConf.SHUFFLE_PARTITIONS.key -> "1", "spark.sql.files.minPartitionNum" -> "1", CometConf.COMET_SHUFFLE_ENABLED.key -> "true", CometConf.COMET_SHUFFLE_MODE.key -> "native") { - withTempPath { path => - Seq(Some(1e100), None, Some(2e100)) - .map(v => (0, v)) - .toDF("g", "v") - .coalesce(1) - .write - .parquet(path.getCanonicalPath) - withParquetTable(path.getCanonicalPath, "correlation_overflow") { - assert(spark.table("correlation_overflow").rdd.getNumPartitions == 1) - for (groupBy <- Seq("", " GROUP BY g")) { - val query = "SELECT corr(v, v), corr(v, -v) FROM correlation_overflow" + groupBy - checkSparkAnswerAndNumOfAggregates(query, 2) - // Spark's sqrt(m2_1 * m2_2) overflows to infinity, so corr is exactly zero. - checkAnswer(sql(query), Seq(Row(0.0, -0.0))) + val cases = Seq( + (Seq(1e100, 2e100), Some(0.0), None), + (Seq(1e-100, 2e-100), None, None), + (Seq(1e200, -1e200), Some(Double.NaN), Some(Double.NaN))) + for ((values, expectedCorr, expectedConstantCorr) <- cases) { + withTempPath { path => + Seq(Some(values.head), None, Some(values.last)) + .map(v => (0, v, 0.1)) + .toDF("g", "v", "x") + .coalesce(1) + .write + .parquet(path.getCanonicalPath) + withParquetTable(path.getCanonicalPath, "correlation_extremes") { + assert(spark.table("correlation_extremes").rdd.getNumPartitions == 1) + for (groupBy <- Seq("", " GROUP BY g")) { + val query = "SELECT corr(v, v), corr(v, -v), corr(v, x) " + + "FROM correlation_extremes" + groupBy + checkSparkAnswerAndNumOfAggregates(query, 2) + // The raw-moment product can overflow, underflow, or become NaN (0 * Inf). + // ANSI-off division by a zero denominator returns NULL. + checkAnswer( + sql(query), + Seq( + Row( + expectedCorr.map(Double.box).orNull, + expectedCorr.map(v => Double.box(-v)).orNull, + expectedConstantCorr.map(Double.box).orNull))) + } } } } From 3166408559c421f42f120cdf9604ed60989d21f5 Mon Sep 17 00:00:00 2001 From: rich7420 Date: Fri, 2 Oct 2026 17:19:32 +0800 Subject: [PATCH 7/9] test: normalize NaN in correlation extreme-value assertions --- .../test/scala/org/apache/comet/exec/CometAggregateSuite.scala | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) 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 ff76e293362..8a7de3e0e80 100644 --- a/spark/src/test/scala/org/apache/comet/exec/CometAggregateSuite.scala +++ b/spark/src/test/scala/org/apache/comet/exec/CometAggregateSuite.scala @@ -2420,7 +2420,7 @@ class CometAggregateSuite extends CometTestBase with AdaptiveSparkPlanHelper { checkSparkAnswerAndNumOfAggregates(query, 2) // The raw-moment product can overflow, underflow, or become NaN (0 * Inf). // ANSI-off division by a zero denominator returns NULL. - checkAnswer( + checkCometAnswer( sql(query), Seq( Row( From 9dd422b23367a3007c69a42c19d66ee0e66a2aa6 Mon Sep 17 00:00:00 2001 From: rich7420 Date: Fri, 2 Oct 2026 23:10:33 +0800 Subject: [PATCH 8/9] fix: restore native regression aggregates after merge correction --- docs/source/user-guide/latest/expressions.md | 12 ++-- .../org/apache/comet/serde/aggregates.scala | 48 ++++------------ .../sql-tests/expressions/aggregate/regr.sql | 10 ---- .../expressions/aggregate/regr_fallback.sql | 56 ------------------- 4 files changed, 17 insertions(+), 109 deletions(-) delete mode 100644 spark/src/test/resources/sql-tests/expressions/aggregate/regr_fallback.sql diff --git a/docs/source/user-guide/latest/expressions.md b/docs/source/user-guide/latest/expressions.md index a809adedc7e..efdf1c4aad3 100644 --- a/docs/source/user-guide/latest/expressions.md +++ b/docs/source/user-guide/latest/expressions.md @@ -126,12 +126,12 @@ The tables below list every Spark built-in expression with its current status. | `regr_avgx` | ✅ | — | Native: Spark rewrites to `Average` (tests in [#4551](https://github.com/apache/datafusion-comet/pull/4551)) | | `regr_avgy` | ✅ | — | Native: Spark rewrites to `Average` (tests in [#4551](https://github.com/apache/datafusion-comet/pull/4551)) | | `regr_count` | ✅ | — | Native: Spark rewrites to `Count` (tests in [#4551](https://github.com/apache/datafusion-comet/pull/4551)) | -| `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` | -| `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` | -| `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` | -| `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`) | -| `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` | -| `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`) | +| `regr_intercept` | ✅ | Native | | +| `regr_r2` | ✅ | Native | | +| `regr_slope` | ✅ | Native | | +| `regr_sxx` | ✅ | Native | | +| `regr_sxy` | ✅ | Native | | +| `regr_syy` | ✅ | Native | | | `skewness` | 🔜 | — | Not yet implemented natively | | `some` | ✅ | — | | | `std` | ✅ | Native | | 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 8928e40cbeb..ad53cceb31b 100644 --- a/spark/src/main/scala/org/apache/comet/serde/aggregates.scala +++ b/spark/src/main/scala/org/apache/comet/serde/aggregates.scala @@ -22,7 +22,7 @@ package org.apache.comet.serde import scala.jdk.CollectionConverters._ import org.apache.spark.sql.catalyst.expressions.{Attribute, Cast, Expression, Literal} -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} +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} import org.apache.spark.sql.catalyst.util.ArrayData import org.apache.spark.sql.comet.CometExecUtils import org.apache.spark.sql.internal.SQLConf @@ -929,27 +929,7 @@ private[comet] object RegrSparkVersions { * variable (y) and `child2` is the independent variable (x), matching the native accumulator's * `regr_*(y, x)` convention. */ -trait CometRegrBase[T <: AggregateFunction] extends CometAggregateExpressionSerde[T] { - - /** The SQL function or functions this serde implements, named in the incompatibility note. */ - protected def sqlFunctions: String - - // The native merge (`variance_merge` and `covariance_merge` in welford.rs) orders its - // floating-point operations differently from Spark's CentralMomentAgg and Covariance. Merging - // the first partial buffer into the zero-initialized final buffer can leave a one-ULP error on - // the mean, so a constant variable ends up with a tiny non-zero m2 and Spark's exact `m2 == 0` - // degenerate-case checks never fire. Porting Spark's merge order would make these Compatible. - private def mergeOrderReason: String = - s"Comet merges the partial aggregates of $sqlFunctions in a different floating-point " + - "operation order from Spark. When a group's rows come from more than one partial " + - "aggregate and a variable is constant at a value that binary floating point cannot " + - "represent exactly, such as 0.1, Comet returns a wrong value where Spark returns NULL, " + - "0.0 or 1.0 (https://github.com/apache/datafusion-comet/issues/6423)" - - override def getIncompatibleReasons(): Seq[String] = Seq(mergeOrderReason) - - override def getSupportLevel(expr: T): SupportLevel = Incompatible(Some(mergeOrderReason)) - +trait CometRegrBase { def convertRegr( aggExpr: AggregateExpression, regrType: ExprOuterClass.Regr.RegrType, @@ -986,9 +966,7 @@ trait CometRegrBase[T <: AggregateFunction] extends CometAggregateExpressionSerd } } -object CometRegrSlope extends CometRegrBase[RegrSlope] { - override protected def sqlFunctions: String = "`regr_slope`" - +object CometRegrSlope extends CometAggregateExpressionSerde[RegrSlope] with CometRegrBase { override def convert( aggExpr: AggregateExpression, expr: RegrSlope, @@ -1004,9 +982,9 @@ object CometRegrSlope extends CometRegrBase[RegrSlope] { binding) } -object CometRegrIntercept extends CometRegrBase[RegrIntercept] { - override protected def sqlFunctions: String = "`regr_intercept`" - +object CometRegrIntercept + extends CometAggregateExpressionSerde[RegrIntercept] + with CometRegrBase { override def convert( aggExpr: AggregateExpression, expr: RegrIntercept, @@ -1022,9 +1000,7 @@ object CometRegrIntercept extends CometRegrBase[RegrIntercept] { binding) } -object CometRegrR2 extends CometRegrBase[RegrR2] { - override protected def sqlFunctions: String = "`regr_r2`" - +object CometRegrR2 extends CometAggregateExpressionSerde[RegrR2] with CometRegrBase { override def convert( aggExpr: AggregateExpression, expr: RegrR2, @@ -1034,9 +1010,7 @@ object CometRegrR2 extends CometRegrBase[RegrR2] { convertRegr(aggExpr, ExprOuterClass.Regr.RegrType.R2, expr.y, expr.x, inputs, binding) } -object CometRegrSXY extends CometRegrBase[RegrSXY] { - override protected def sqlFunctions: String = "`regr_sxy`" - +object CometRegrSXY extends CometAggregateExpressionSerde[RegrSXY] with CometRegrBase { override def convert( aggExpr: AggregateExpression, expr: RegrSXY, @@ -1053,9 +1027,9 @@ object CometRegrSXY extends CometRegrBase[RegrSXY] { * deviations) of its single child. We serialize it as the `SXX` regression statistic with the * child duplicated, since `regr_sxx(c, c) = m2(c)`. */ -object CometRegrReplacement extends CometRegrBase[RegrReplacement] { - override protected def sqlFunctions: String = "`regr_sxx` and `regr_syy`" - +object CometRegrReplacement + extends CometAggregateExpressionSerde[RegrReplacement] + with CometRegrBase { override def convert( aggExpr: AggregateExpression, expr: RegrReplacement, diff --git a/spark/src/test/resources/sql-tests/expressions/aggregate/regr.sql b/spark/src/test/resources/sql-tests/expressions/aggregate/regr.sql index 3618039e296..90c3017978f 100644 --- a/spark/src/test/resources/sql-tests/expressions/aggregate/regr.sql +++ b/spark/src/test/resources/sql-tests/expressions/aggregate/regr.sql @@ -19,16 +19,6 @@ -- regr_avgy, regr_sxx, regr_syy, regr_sxy, regr_slope, regr_intercept, regr_r2. -- All functions take (y, x) and operate only on rows where BOTH y and x are non-null. --- regr_slope, regr_intercept, regr_r2, regr_sxx, regr_syy and regr_sxy fall back to Spark by --- default because their native merge of partial aggregates differs from Spark's --- (https://github.com/apache/datafusion-comet/issues/6423). Opt in so the queries below cover --- the native path. Spark plans regr_sxx and regr_syy as RegrReplacement. --- Config: spark.comet.expression.RegrSlope.allowIncompatible=true --- Config: spark.comet.expression.RegrIntercept.allowIncompatible=true --- Config: spark.comet.expression.RegrR2.allowIncompatible=true --- Config: spark.comet.expression.RegrSXY.allowIncompatible=true --- Config: spark.comet.expression.RegrReplacement.allowIncompatible=true - statement CREATE TABLE test_regr(y double, x double, grp string) USING parquet diff --git a/spark/src/test/resources/sql-tests/expressions/aggregate/regr_fallback.sql b/spark/src/test/resources/sql-tests/expressions/aggregate/regr_fallback.sql deleted file mode 100644 index fa12cbf2976..00000000000 --- a/spark/src/test/resources/sql-tests/expressions/aggregate/regr_fallback.sql +++ /dev/null @@ -1,56 +0,0 @@ --- Licensed to the Apache Software Foundation (ASF) under one --- or more contributor license agreements. See the NOTICE file --- distributed with this work for additional information --- regarding copyright ownership. The ASF licenses this file --- to you under the Apache License, Version 2.0 (the --- "License"); you may not use this file except in compliance --- with the License. You may obtain a copy of the License at --- --- http://www.apache.org/licenses/LICENSE-2.0 --- --- Unless required by applicable law or agreed to in writing, --- software distributed under the License is distributed on an --- "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY --- KIND, either express or implied. See the License for the --- specific language governing permissions and limitations --- under the License. - --- regr_slope, regr_intercept, regr_r2, regr_sxx, regr_syy and regr_sxy fall back to Spark by --- default, because their native merge of partial aggregates orders its floating-point operations --- differently from Spark's (https://github.com/apache/datafusion-comet/issues/6423). Each query --- below holds a single regr function, so it checks that function's own fallback. regr.sql opts in --- and covers the native path. - --- The data from #6423: x is constant at 0.1, and each INSERT writes one file of three rows, so the --- rows are merged from two partial aggregates. The native merge leaves x with a tiny non-zero --- variance there, and returns wrong values where Spark returns NULL, 0.0 or 1.0. -statement -CREATE TABLE test_regr_fallback(y double, x double) USING parquet - -statement -INSERT INTO test_regr_fallback SELECT CAST(id AS DOUBLE), 0.1D FROM range(0, 3, 1, 1) - -statement -INSERT INTO test_regr_fallback SELECT CAST(id AS DOUBLE), 0.1D FROM range(3, 6, 1, 1) - -query expect_fallback(issues/6423) -SELECT regr_slope(y, x) FROM test_regr_fallback - -query expect_fallback(issues/6423) -SELECT regr_intercept(y, x) FROM test_regr_fallback - -query expect_fallback(issues/6423) -SELECT regr_r2(y, x) FROM test_regr_fallback - --- The constant as the dependent variable -query expect_fallback(issues/6423) -SELECT regr_r2(x, y) FROM test_regr_fallback - -query expect_fallback(issues/6423) -SELECT regr_sxx(y, x) FROM test_regr_fallback - -query expect_fallback(issues/6423) -SELECT regr_sxy(y, x) FROM test_regr_fallback - -query expect_fallback(issues/6423) -SELECT regr_syy(x, y) FROM test_regr_fallback From 418accc5238d6bab0300fae4c9655bfe4526b6ee Mon Sep 17 00:00:00 2001 From: rich7420 Date: Sat, 3 Oct 2026 02:08:41 +0800 Subject: [PATCH 9/9] fix: preserve Spark regr_r2 divide-by-zero semantics --- native/core/src/execution/planner.rs | 1 + native/proto/src/proto/expr.proto | 3 + native/spark-expr/src/agg_funcs/regr.rs | 159 +++++++++++++++--- .../org/apache/comet/serde/aggregates.scala | 29 +++- .../sql-tests/expressions/aggregate/regr.sql | 45 ++++- .../comet/exec/CometAggregateSuite.scala | 95 ++++++++++- 6 files changed, 294 insertions(+), 38 deletions(-) diff --git a/native/core/src/execution/planner.rs b/native/core/src/execution/planner.rs index 42571d649a0..a8d8c944753 100644 --- a/native/core/src/execution/planner.rs +++ b/native/core/src/execution/planner.rs @@ -3103,6 +3103,7 @@ impl PhysicalPlanner { name, expr.filter_var_by_pair_nulls, expr.r2_constant_dependent_is_perfect_fit, + from_protobuf_eval_mode(expr.eval_mode)?, )); Self::create_aggr_func_expr(name, schema, vec![child1, child2], func) } diff --git a/native/proto/src/proto/expr.proto b/native/proto/src/proto/expr.proto index a53a97f6c03..1ec4d5bbe4b 100644 --- a/native/proto/src/proto/expr.proto +++ b/native/proto/src/proto/expr.proto @@ -298,6 +298,9 @@ message Regr { // Comet builds against for 3.5 and later (3.5.9, 4.0.3+, 4.1, 4.2) but not // in 3.4. bool r2_constant_dependent_is_perfect_fit = 6; + // Only consulted for R2. Division by an underflowed moment product returns + // null in LEGACY mode and raises DIVIDE_BY_ZERO in ANSI mode. + EvalMode eval_mode = 7; } message Percentile { diff --git a/native/spark-expr/src/agg_funcs/regr.rs b/native/spark-expr/src/agg_funcs/regr.rs index 0b8838c985b..15c95cbcdb8 100644 --- a/native/spark-expr/src/agg_funcs/regr.rs +++ b/native/spark-expr/src/agg_funcs/regr.rs @@ -45,6 +45,7 @@ use std::sync::Arc; use crate::agg_funcs::covariance::CovarianceAccumulator; use crate::agg_funcs::variance::VarianceAccumulator; +use crate::{divide_by_zero_error, EvalMode}; /// The kind of linear-regression statistic to compute. /// @@ -80,6 +81,8 @@ pub struct Regr { /// 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, + /// Only consulted for `R2` when the raw moment product underflows to zero. + eval_mode: EvalMode, } impl Regr { @@ -88,6 +91,7 @@ impl Regr { name: impl Into, filter_var_by_pair_nulls: bool, r2_constant_dependent_is_perfect_fit: bool, + eval_mode: EvalMode, ) -> Self { Self { name: name.into(), @@ -98,6 +102,7 @@ impl Regr { regr_type, filter_var_by_pair_nulls, r2_constant_dependent_is_perfect_fit, + eval_mode, } } @@ -140,6 +145,7 @@ impl AggregateUDFImpl for Regr { RegrType::SXY => Box::new(RegrCovAccumulator::try_new()?), RegrType::R2 => Box::new(RegrR2Accumulator::try_new( self.r2_constant_dependent_is_perfect_fit, + self.eval_mode, )?), RegrType::Slope => Box::new(RegrLineAccumulator::try_new( false, @@ -297,13 +303,14 @@ struct RegrR2Accumulator { covar: CovarianceAccumulator, var_y: VarianceAccumulator, var_x: VarianceAccumulator, - /// When `true` (Spark 3.5+), a constant dependent variable yields `1.0` and a - /// constant independent variable yields `null`; reversed when `false`. + /// When `true` (after SPARK-55969), a constant dependent variable yields + /// `1.0` and a constant independent variable yields `null`; reversed otherwise. constant_dependent_is_perfect_fit: bool, + eval_mode: EvalMode, } impl RegrR2Accumulator { - fn try_new(constant_dependent_is_perfect_fit: bool) -> Result { + fn try_new(constant_dependent_is_perfect_fit: bool, eval_mode: EvalMode) -> Result { Ok(Self { covar: CovarianceAccumulator::try_new(StatsType::Population, false)?, var_y: VarianceAccumulator::try_new(StatsType::Population, false)? @@ -311,6 +318,7 @@ impl RegrR2Accumulator { var_x: VarianceAccumulator::try_new(StatsType::Population, false)? .with_pearson_update(), constant_dependent_is_perfect_fit, + eval_mode, }) } } @@ -381,12 +389,22 @@ impl Accumulator for RegrR2Accumulator { } else if perfect_fit_case { Ok(ScalarValue::Float64(Some(1.0))) } else { + // Nonzero moments can still have a product that underflows to zero. + // Spark's Divide handles this only after the constant-input guards. + let product = m2_y * m2_x; + if product == 0.0 { + return if self.eval_mode == EvalMode::Ansi { + Err(divide_by_zero_error().into()) + } else { + Ok(ScalarValue::Float64(None)) + }; + } // Mirror Spark's exact evaluation order (corr = ck / sqrt(m2_y * m2_x); // corr * corr) so the last-ULP rounding matches bit-for-bit. Writing // it as (ck * ck) / (m2_x * m2_y) is mathematically equal but rounds // differently. let ck = self.covar.get_algo_const(); - let corr = ck / (m2_y * m2_x).sqrt(); + let corr = ck / product.sqrt(); Ok(ScalarValue::Float64(Some(corr * corr))) } } @@ -499,13 +517,15 @@ impl Accumulator for RegrLineAccumulator { mod tests { use super::*; use arrow::array::Float64Array; + use datafusion::common::DataFusionError; + use datafusion_comet_common::SparkError; fn acc(regr_type: RegrType) -> Box { match regr_type { RegrType::SXX | RegrType::SYY => Box::new(RegrMomentAccumulator::try_new().unwrap()), RegrType::SXY => Box::new(RegrCovAccumulator::try_new().unwrap()), // Default to the post-SPARK-55969 degenerate-case semantics. - RegrType::R2 => Box::new(RegrR2Accumulator::try_new(true).unwrap()), + RegrType::R2 => Box::new(RegrR2Accumulator::try_new(true, EvalMode::Legacy).unwrap()), // Existing tests exercise the Spark 3.5+ both-non-null semantics. RegrType::Slope => Box::new(RegrLineAccumulator::try_new(false, true).unwrap()), RegrType::Intercept => Box::new(RegrLineAccumulator::try_new(true, true).unwrap()), @@ -597,6 +617,106 @@ mod tests { approx(eval(RegrType::R2, y, x), 1.0); } + fn assert_r2_zero_divisor(a: &mut RegrR2Accumulator, eval_mode: EvalMode) { + // Both moments are nonzero. It is their product, not a constant input, + // that makes Spark's final division have a zero denominator. + assert_ne!(a.var_y.get_m2(), 0.0); + assert_ne!(a.var_x.get_m2(), 0.0); + assert_eq!(a.var_y.get_m2() * a.var_x.get_m2(), 0.0); + if eval_mode == EvalMode::Ansi { + let DataFusionError::External(error) = a.evaluate().unwrap_err() else { + panic!("expected a structured Spark divide-by-zero error"); + }; + assert!(matches!( + error.downcast_ref::(), + Some(SparkError::DivideByZero) + )); + } else { + assert_eq!(a.evaluate().unwrap(), ScalarValue::Float64(None)); + } + } + + #[test] + fn r2_moment_product_underflow() { + for eval_mode in [EvalMode::Legacy, EvalMode::Ansi] { + for perfect_dep in [false, true] { + for sign in [-1.0, 1.0] { + for merge in [false, true] { + let mut a = RegrR2Accumulator::try_new(perfect_dep, eval_mode).unwrap(); + let y = vec![Some(1e-100), None, Some(2e-100)]; + let x = vec![Some(sign * 1e-100), Some(1e200), Some(sign * 2e-100)]; + if merge { + // Merge two single-row partials, including an unpaired null. + for range in [0..2, 2..3] { + let mut partial = + RegrR2Accumulator::try_new(perfect_dep, eval_mode).unwrap(); + partial + .update_batch(&cols( + y[range.clone()].to_vec(), + x[range].to_vec(), + )) + .unwrap(); + let state = partial + .state() + .unwrap() + .iter() + .map(|v| v.to_array_of_size(1).unwrap()) + .collect::>(); + a.merge_batch(&state).unwrap(); + } + } else { + a.update_batch(&cols(y, x)).unwrap(); + } + assert_r2_zero_divisor(&mut a, eval_mode); + } + } + } + } + } + + #[test] + fn r2_zero_covariance_with_underflowed_product() { + for eval_mode in [EvalMode::Legacy, EvalMode::Ansi] { + for perfect_dep in [false, true] { + let mut a = RegrR2Accumulator::try_new(perfect_dep, eval_mode).unwrap(); + // A zero co-moment still divides by zero, rather than returning NaN. + let state = [4.0, 0.0, 0.0, 0.0, 4e-200, 4e-200] + .into_iter() + .map(|v| Arc::new(Float64Array::from(vec![v])) as ArrayRef) + .collect::>(); + a.merge_batch(&state).unwrap(); + assert_r2_zero_divisor(&mut a, eval_mode); + } + } + } + + #[test] + fn r2_raw_moment_overflow() { + for eval_mode in [EvalMode::Legacy, EvalMode::Ansi] { + for perfect_dep in [false, true] { + for sign in [-1.0, 1.0] { + for (values, nan) in [(vec![1e100, 2e100], false), (vec![1e200, -1e200], true)] + { + let mut a = RegrR2Accumulator::try_new(perfect_dep, eval_mode).unwrap(); + a.update_batch(&cols( + values.iter().copied().map(Some).collect(), + values.iter().map(|v| Some(sign * v)).collect(), + )) + .unwrap(); + let ScalarValue::Float64(Some(result)) = a.evaluate().unwrap() else { + panic!("expected a non-null double result"); + }; + if nan { + assert!(result.is_nan()); + } else { + assert_eq!(result.to_bits(), 0.0_f64.to_bits()); + } + } + } + } + } + } + #[test] fn r2_constant_y_is_perfect_fit() { // Dependent variable constant: post-SPARK-55969 regr_r2 returns 1.0. @@ -617,21 +737,22 @@ mod tests { vec![Some(5.0), Some(5.0), Some(5.0), Some(5.0)], ); - let r2 = |perfect_dep: bool, y: Vec>, x: Vec>| { - let mut a = RegrR2Accumulator::try_new(perfect_dep).unwrap(); - a.update_batch(&cols(y, x)).unwrap(); - match a.evaluate().unwrap() { - ScalarValue::Float64(v) => v, - other => panic!("unexpected {other:?}"), - } - }; + for eval_mode in [EvalMode::Legacy, EvalMode::Ansi] { + let r2 = |perfect_dep: bool, y: Vec>, x: Vec>| { + let mut a = RegrR2Accumulator::try_new(perfect_dep, eval_mode).unwrap(); + a.update_batch(&cols(y, x)).unwrap(); + match a.evaluate().unwrap() { + ScalarValue::Float64(v) => v, + other => panic!("unexpected {other:?}"), + } + }; - // Spark 3.5+: constant dependent -> 1.0, constant independent -> null. - approx(r2(true, const_y.0.clone(), const_y.1.clone()), 1.0); - assert_eq!(r2(true, const_x.0.clone(), const_x.1.clone()), None); - // Spark 3.4: constant dependent -> null, constant independent -> 1.0. - assert_eq!(r2(false, const_y.0, const_y.1), None); - approx(r2(false, const_x.0, const_x.1), 1.0); + // These guards must run before the zero-product check, even in ANSI mode. + assert_eq!(r2(true, const_y.0.clone(), const_y.1.clone()), Some(1.0)); + assert_eq!(r2(true, const_x.0.clone(), const_x.1.clone()), None); + assert_eq!(r2(false, const_y.0.clone(), const_y.1.clone()), None); + assert_eq!(r2(false, const_x.0.clone(), const_x.1.clone()), Some(1.0)); + } } #[test] 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 ad53cceb31b..7eeb51e38de 100644 --- a/spark/src/main/scala/org/apache/comet/serde/aggregates.scala +++ b/spark/src/main/scala/org/apache/comet/serde/aggregates.scala @@ -21,7 +21,7 @@ package org.apache.comet.serde import scala.jdk.CollectionConverters._ -import org.apache.spark.sql.catalyst.expressions.{Attribute, Cast, Expression, Literal} +import org.apache.spark.sql.catalyst.expressions.{Attribute, Cast, Divide, Expression, Literal} 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} import org.apache.spark.sql.catalyst.util.ArrayData import org.apache.spark.sql.comet.CometExecUtils @@ -936,7 +936,8 @@ trait CometRegrBase { y: Expression, x: Expression, inputs: Seq[Attribute], - binding: Boolean): Option[ExprOuterClass.AggExpr] = { + binding: Boolean, + evalMode: CometEvalMode.Value = CometEvalMode.LEGACY): Option[ExprOuterClass.AggExpr] = { val child1Expr = exprToProto(y, inputs, binding) val child2Expr = exprToProto(x, inputs, binding) val dataType = serializeDataType(DoubleType) @@ -947,6 +948,7 @@ trait CometRegrBase { builder.setChild2(child2Expr.get) builder.setRegrType(regrType) builder.setDatatype(dataType.get) + builder.setEvalMode(evalModeToProto(evalMode)) // 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 @@ -1006,8 +1008,27 @@ object CometRegrR2 extends CometAggregateExpressionSerde[RegrR2] with CometRegrB expr: RegrR2, inputs: Seq[Attribute], binding: Boolean, - conf: SQLConf): Option[ExprOuterClass.AggExpr] = - convertRegr(aggExpr, ExprOuterClass.Regr.RegrType.R2, expr.y, expr.x, inputs, binding) + conf: SQLConf): Option[ExprOuterClass.AggExpr] = { + // Divide captures its mode when Spark constructs the expression. Reading the + // current SQLConf could change semantics if ANSI mode has changed since then. + val evalMode = expr.evaluateExpression.collectFirst { case divide: Divide => + CometEvalModeUtil.fromSparkEvalMode(divide.evalMode) + } + evalMode match { + case Some(mode) => + convertRegr( + aggExpr, + ExprOuterClass.Regr.RegrType.R2, + expr.y, + expr.x, + inputs, + binding, + mode) + case None => + withFallbackReason(aggExpr, "REGR_R2 division evaluation mode not supported") + None + } + } } object CometRegrSXY extends CometAggregateExpressionSerde[RegrSXY] with CometRegrBase { diff --git a/spark/src/test/resources/sql-tests/expressions/aggregate/regr.sql b/spark/src/test/resources/sql-tests/expressions/aggregate/regr.sql index 90c3017978f..e17268eccbd 100644 --- a/spark/src/test/resources/sql-tests/expressions/aggregate/regr.sql +++ b/spark/src/test/resources/sql-tests/expressions/aggregate/regr.sql @@ -118,11 +118,11 @@ query SELECT regr_count(y, x), regr_avgx(y, x), regr_avgy(y, x) FROM test_regr_all_null -- regr_sxx/syy/sxy return NULL when there are no non-null pairs -query tolerance=1e-6 +query SELECT regr_sxx(y, x), regr_syy(y, x), regr_sxy(y, x) FROM test_regr_all_null -- regr_slope/intercept/r2 return NULL when there are no non-null pairs -query tolerance=1e-6 +query SELECT regr_slope(y, x), regr_intercept(y, x), regr_r2(y, x) FROM test_regr_all_null -- edge case: single non-null pair (slope/intercept/r2 require >= 2 rows) @@ -136,10 +136,10 @@ query SELECT regr_count(y, x), regr_avgx(y, x), regr_avgy(y, x) FROM test_regr_single -- sxx/syy/sxy are 0 for a single pair; slope/intercept/r2 are NULL -query tolerance=1e-6 +query SELECT regr_sxx(y, x), regr_syy(y, x), regr_sxy(y, x) FROM test_regr_single -query tolerance=1e-6 +query SELECT regr_slope(y, x), regr_intercept(y, x), regr_r2(y, x) FROM test_regr_single -- edge case: independent variable (x) is constant but y varies. @@ -152,10 +152,10 @@ CREATE TABLE test_regr_const_x(y double, x double) USING parquet statement INSERT INTO test_regr_const_x VALUES (1.0, 5.0), (2.0, 5.0), (3.0, 5.0), (4.0, 5.0) -query tolerance=1e-6 +query SELECT regr_slope(y, x), regr_intercept(y, x), regr_r2(y, x) FROM test_regr_const_x -query tolerance=1e-6 +query 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. @@ -168,8 +168,37 @@ CREATE TABLE test_regr_const_y(y double, x double) USING parquet statement INSERT INTO test_regr_const_y VALUES (7.0, 1.0), (7.0, 2.0), (7.0, 3.0), (7.0, 4.0) -query tolerance=1e-6 +query SELECT regr_slope(y, x), regr_intercept(y, x), regr_r2(y, x) FROM test_regr_const_y -query tolerance=1e-6 +query SELECT regr_sxx(y, x), regr_syy(y, x), regr_sxy(y, x) FROM test_regr_const_y + +-- Both moments are nonzero, but their product underflows to zero. +statement +CREATE TABLE test_regr_underflow(v double, grp int) USING parquet + +statement +INSERT INTO test_regr_underflow VALUES (1e-100, 0), (NULL, 0), (2e-100, 0) + +query +SELECT regr_r2(v, v), regr_r2(v, -v) FROM test_regr_underflow + +query +SELECT grp, regr_r2(v, v), regr_r2(v, -v) FROM test_regr_underflow GROUP BY grp + +statement +SET spark.sql.ansi.enabled=true + +query expect_error(DIVIDE_BY_ZERO) +SELECT regr_r2(v, v) FROM test_regr_underflow + +query expect_error(DIVIDE_BY_ZERO) +SELECT grp, regr_r2(v, -v) FROM test_regr_underflow GROUP BY grp + +-- The existing constant-input guards precede division, even in ANSI mode. +query +SELECT regr_r2(y, x) FROM test_regr_const_x + +query +SELECT regr_r2(y, x) FROM test_regr_const_y 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 47c2b4321d5..44cce97e005 100644 --- a/spark/src/test/scala/org/apache/comet/exec/CometAggregateSuite.scala +++ b/spark/src/test/scala/org/apache/comet/exec/CometAggregateSuite.scala @@ -27,8 +27,8 @@ import org.apache.hadoop.fs.Path import org.apache.spark.{CometListenerBusUtils, SparkConf} import org.apache.spark.scheduler.{SparkListener, SparkListenerTaskEnd} import org.apache.spark.sql.{CometTestBase, DataFrame, Row} -import org.apache.spark.sql.catalyst.expressions.Cast -import org.apache.spark.sql.catalyst.expressions.aggregate.{Final, Partial, PartialMerge} +import org.apache.spark.sql.catalyst.expressions.{Cast, Literal} +import org.apache.spark.sql.catalyst.expressions.aggregate.{Final, Partial, PartialMerge, RegrR2} import org.apache.spark.sql.catalyst.optimizer.EliminateSorts import org.apache.spark.sql.catalyst.plans.physical.{HashPartitioning, RangePartitioning} import org.apache.spark.sql.comet.{CometFilterExec, CometHashAggregateExec, CometNativeExec, CometProjectExec} @@ -45,7 +45,7 @@ import org.apache.comet.CometConf import org.apache.comet.CometConf.COMET_EXEC_STRICT_FLOATING_POINT import org.apache.comet.CometSparkSessionExtensions.{isSpark35Plus, isSpark40Plus, isSpark41Plus} import org.apache.comet.rules.CometExecRule -import org.apache.comet.serde.RegrSparkVersions +import org.apache.comet.serde.{CometRegrR2, ExprOuterClass, RegrSparkVersions} import org.apache.comet.testing.{DataGenOptions, FuzzDataGenerator, ParquetGenerator, SchemaGenOptions} /** @@ -2797,7 +2797,7 @@ class CometAggregateSuite extends CometTestBase with AdaptiveSparkPlanHelper { // Reversing the positive pair rounds the Pearson mean to the second // value, leaving zero M2 and a NULL correlation. checkAnswer(sql(corrQuery), Seq(Row(expectedCorr.map(Double.box).orNull))) - checkSparkAnswerWithTolAndNumOfAggregates( + checkSparkAnswerAndNumOfAggregates( "SELECT regr_r2(v, v), regr_sxx(v, v), regr_syy(v, v), " + "regr_slope(v, v), regr_intercept(v, v) FROM large_moments" + groupBy, 2) @@ -2808,7 +2808,7 @@ class CometAggregateSuite extends CometTestBase with AdaptiveSparkPlanHelper { } } - test("statistical aggregates correlation uses raw moments at extreme magnitudes") { + test("statistical aggregates use raw moments at extreme magnitudes") { withSQLConf( SQLConf.ANSI_ENABLED.key -> "false", SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false", @@ -2831,7 +2831,8 @@ class CometAggregateSuite extends CometTestBase with AdaptiveSparkPlanHelper { withParquetTable(path.getCanonicalPath, "correlation_extremes") { assert(spark.table("correlation_extremes").rdd.getNumPartitions == 1) for (groupBy <- Seq("", " GROUP BY g")) { - val query = "SELECT corr(v, v), corr(v, -v), corr(v, x) " + + val query = "SELECT corr(v, v), corr(v, -v), corr(v, x), " + + "regr_r2(v, v), regr_r2(v, -v) " + "FROM correlation_extremes" + groupBy checkSparkAnswerAndNumOfAggregates(query, 2) // The raw-moment product can overflow, underflow, or become NaN (0 * Inf). @@ -2842,7 +2843,87 @@ class CometAggregateSuite extends CometTestBase with AdaptiveSparkPlanHelper { Row( expectedCorr.map(Double.box).orNull, expectedCorr.map(v => Double.box(-v)).orNull, - expectedConstantCorr.map(Double.box).orNull))) + expectedConstantCorr.map(Double.box).orNull, + expectedCorr.map(v => Double.box(v * v)).orNull, + expectedCorr.map(v => Double.box(v * v)).orNull))) + } + } + } + } + } + } + + test("statistical aggregates regr_r2 preserves its captured evaluation mode") { + for (ansi <- Seq(false, true)) { + withSQLConf(SQLConf.ANSI_ENABLED.key -> ansi.toString) { + val expr = RegrR2(Literal(1.0), Literal(2.0)) + withSQLConf(SQLConf.ANSI_ENABLED.key -> (!ansi).toString) { + val serialized = CometRegrR2.convert( + expr.toAggregateExpression(), + expr, + Seq.empty, + binding = false, + conf = SQLConf.get) + assert(serialized.isDefined) + val expected = + if (ansi) ExprOuterClass.EvalMode.ANSI else ExprOuterClass.EvalMode.LEGACY + assert(serialized.get.getRegr.getEvalMode == expected) + } + } + } + } + + test("statistical aggregates regr_r2 underflow respects ANSI mode") { + withSQLConf( + SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false", + SQLConf.SHUFFLE_PARTITIONS.key -> "1", + SQLConf.FILES_MAX_PARTITION_BYTES.key -> "1048576", + SQLConf.FILES_OPEN_COST_IN_BYTES.key -> "1048576", + "spark.sql.files.minPartitionNum" -> "1", + CometConf.COMET_SHUFFLE_ENABLED.key -> "true", + CometConf.COMET_SHUFFLE_MODE.key -> "native") { + for (numPartitions <- Seq(1, 2)) { + withTempPath { path => + // One file exercises updating a partial with both values. Two files also + // exercise merging independent single-row partials in the final aggregate. + val files = if (numPartitions == 1) { + Seq(Seq(Some(1e-100), None, Some(2e-100))) + } else { + Seq(Seq(Some(1e-100), None), Seq(Some(2e-100))) + } + files.foreach { values => + values + .map(v => (0, v)) + .toDF("g", "v") + .coalesce(1) + .write + .mode("append") + .parquet(path.getCanonicalPath) + } + withParquetTable(path.getCanonicalPath, "r2_underflow") { + assert(spark.table("r2_underflow").rdd.getNumPartitions == numPartitions) + for { + ansi <- Seq(false, true) + groupBy <- Seq("", " GROUP BY g") + x <- Seq("v", "-v") + } { + withSQLConf(SQLConf.ANSI_ENABLED.key -> ansi.toString) { + val query = s"SELECT regr_r2(v, $x) FROM r2_underflow" + groupBy + val df = sql(query) + val aggregates = stripAQEPlan(df.queryExecution.executedPlan).collect { + case a: CometHashAggregateExec => a + } + assert(aggregates.size == 2) + assert(aggregates.exists(_.modes.contains(Partial))) + assert(aggregates.exists(_.modes.contains(Final))) + if (ansi) { + val error = checkSparkError(df, "DIVIDE_BY_ZERO") + assert(error.getSqlState == "22012") + } else { + checkSparkAnswer(df) + checkCometAnswer(df, Seq(Row(null))) + } + } } } }