From 3380d9315c131e259d9e0636ead0fc7cefb1dae6 Mon Sep 17 00:00:00 2001 From: Mikhail Kot Date: Thu, 30 Jul 2026 12:29:37 +0100 Subject: [PATCH 1/3] mean(all null) = NULL Signed-off-by: Mikhail Kot --- vortex-array/src/aggregate_fn/fns/mean/mod.rs | 91 +++++++++++++++++-- .../slt/datafusion/mean_all_null.slt | 30 ++++++ .../slt/duckdb/mean_all_null.slt | 26 ++++++ vortex-sqllogictest/slt/mean_all_null.slt | 44 +++++++++ 4 files changed, 182 insertions(+), 9 deletions(-) create mode 100644 vortex-sqllogictest/slt/datafusion/mean_all_null.slt create mode 100644 vortex-sqllogictest/slt/duckdb/mean_all_null.slt create mode 100644 vortex-sqllogictest/slt/mean_all_null.slt diff --git a/vortex-array/src/aggregate_fn/fns/mean/mod.rs b/vortex-array/src/aggregate_fn/fns/mean/mod.rs index 20b7f15834c..eaabda428b9 100644 --- a/vortex-array/src/aggregate_fn/fns/mean/mod.rs +++ b/vortex-array/src/aggregate_fn/fns/mean/mod.rs @@ -7,6 +7,7 @@ use vortex_session::registry::CachedId; use crate::ArrayRef; use crate::ExecutionCtx; +use crate::IntoArray; use crate::aggregate_fn::Accumulator; use crate::aggregate_fn::AggregateFnId; use crate::aggregate_fn::AggregateFnVTable; @@ -19,6 +20,7 @@ use crate::aggregate_fn::combined::PairOptions; use crate::aggregate_fn::fns::count::Count; use crate::aggregate_fn::fns::sum::Sum; use crate::aggregate_fn::fns::sum::sum_decimal_dtype; +use crate::arrays::ConstantArray; use crate::builtins::ArrayBuiltins; use crate::dtype::DType; use crate::dtype::DecimalDType; @@ -96,8 +98,18 @@ impl BinaryCombined for Mean { } let target = DType::Primitive(PType::F64, Nullability::Nullable); let sum_cast = sum.cast(target.clone())?; - let count_cast = count.cast(target)?; - sum_cast.binary(count_cast, Operator::Div) + let count_cast = count.cast(target.clone())?; + let mean = sum_cast.binary(count_cast.clone(), Operator::Div)?; + + // Nulls (and nans if skip_nan is enabled) are skipped so group with + // zero values has a count of 0, 0 / 0 = nan, and we need null. + let non_empty = count_cast + .binary( + ConstantArray::new(Scalar::zero_value(&target), count_cast.len()).into_array(), + Operator::NotEq, + )? + .fill_null(false)?; + mean.mask(non_empty) } fn finalize_scalar(&self, left_scalar: Scalar, right_scalar: Scalar) -> VortexResult { @@ -112,9 +124,8 @@ impl BinaryCombined for Mean { let sum = sum_cast.as_primitive().typed_value::(); let count = count_cast.as_primitive().typed_value::(); let value = match (sum, count) { - (None, _) | (_, None) => return Ok(Scalar::null(target)), // Sum overflowed - // A count of zero yields 0/0 = NaN, matching the array `finalize` path: nulls are - // skipped during accumulation, so an all-null input is an empty mean, not null. + // None sum means sum overflowed, 0 count means empty input + (None, _) | (_, None) | (_, Some(0.0)) => return Ok(Scalar::null(target)), (Some(s), Some(c)) => s / c, }; Ok(Scalar::primitive(value, Nullability::Nullable)) @@ -217,13 +228,14 @@ mod tests { use vortex_error::VortexResult; use super::*; - use crate::IntoArray; use crate::VortexSessionExecute; + use crate::aggregate_fn::DynGroupedAccumulator; + use crate::aggregate_fn::GroupedAccumulator; use crate::array_session; use crate::arrays::BoolArray; use crate::arrays::ChunkedArray; - use crate::arrays::ConstantArray; use crate::arrays::DecimalArray; + use crate::arrays::FixedSizeListArray; use crate::arrays::PrimitiveArray; use crate::dtype::DecimalDType; use crate::validity::Validity; @@ -315,11 +327,11 @@ mod tests { } #[test] - fn mean_all_null_returns_nan() -> VortexResult<()> { + fn mean_all_null_returns_null() -> VortexResult<()> { let array = PrimitiveArray::from_option_iter::([None, None, None]).into_array(); let mut ctx = array_session().create_execution_ctx(); let result = mean(&array, &mut ctx)?; - assert!(result.as_primitive().as_::().is_some_and(f64::is_nan)); + assert_eq!(result.as_primitive().as_::(), None); Ok(()) } @@ -414,4 +426,65 @@ mod tests { assert_eq!(result.as_primitive().as_::(), Some(3.0)); Ok(()) } + + fn mean_nan_null() -> Vec<(Vec>, Option)> { + vec![ + (vec![Some(f64::NAN), Some(1.0), None], Some(1.0)), + (vec![Some(f64::NAN), Some(1.0), Some(3.0)], Some(2.0)), + (vec![None, None, Some(f64::NAN)], None), + (vec![None, None, None], None), + (vec![Some(1.0), Some(2.0), Some(3.0)], Some(2.0)), + ] + } + + #[test] + fn mean_combined_partials() -> VortexResult<()> { + let mut ctx = array_session().create_execution_ctx(); + for (case, (group, expected)) in mean_nan_null().into_iter().enumerate() { + let mut acc = Accumulator::try_new( + Mean::combined(), + PairOptions( + NumericalAggregateOpts::default(), + NumericalAggregateOpts::default(), + ), + DType::Primitive(PType::F64, Nullability::Nullable), + )?; + let (head, tail) = group.split_at(2); + let head = PrimitiveArray::from_option_iter(head.iter().copied()).into_array(); + let tail = PrimitiveArray::from_option_iter(tail.iter().copied()).into_array(); + acc.accumulate(&head, &mut ctx)?; + acc.accumulate(&tail, &mut ctx)?; + let result = acc.finish()?; + assert_eq!(result.as_primitive().as_::(), expected, "case {case}"); + } + Ok(()) + } + + #[test] + fn mean_grouped_finalize() -> VortexResult<()> { + let cases = mean_nan_null(); + let elements = PrimitiveArray::from_option_iter( + cases.iter().flat_map(|(group, _)| group.iter().copied()), + ) + .into_array(); + let groups = FixedSizeListArray::try_new(elements, 3, Validity::NonNullable, cases.len())?; + + let mut acc = GroupedAccumulator::try_new( + Mean::combined(), + PairOptions( + NumericalAggregateOpts::default(), + NumericalAggregateOpts::default(), + ), + DType::Primitive(PType::F64, Nullability::Nullable), + )?; + let mut ctx = array_session().create_execution_ctx(); + acc.accumulate_list(&groups.into_array(), &mut ctx)?; + let result = acc.finish()?; + + for (case, (_, expected)) in cases.into_iter().enumerate() { + let actual = result.execute_scalar(case, &mut ctx)?; + assert_eq!(actual.as_primitive().as_::(), expected, "case {case}"); + } + Ok(()) + } } diff --git a/vortex-sqllogictest/slt/datafusion/mean_all_null.slt b/vortex-sqllogictest/slt/datafusion/mean_all_null.slt new file mode 100644 index 00000000000..b5cd82fa19a --- /dev/null +++ b/vortex-sqllogictest/slt/datafusion/mean_all_null.slt @@ -0,0 +1,30 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright the Vortex contributors + +include ../setup.slt.no + +statement ok +COPY (SELECT * FROM (VALUES (CAST(NULL AS BIGINT)),(CAST(NULL AS BIGINT))) AS t(x)) +TO '${WORK_DIR}/i-null.vortex'; + +statement ok +COPY (SELECT * FROM (VALUES (CAST(NULL AS DOUBLE)),(CAST(NULL AS DOUBLE))) AS t(x)) +TO '${WORK_DIR}/f-null.vortex'; + +statement ok +COPY (SELECT * FROM (VALUES ('NaN'::DOUBLE),(1.0),(2.0)) AS t(x)) TO '${WORK_DIR}/f-nan.vortex'; + +query IR +SELECT count(x), avg(x) FROM '${WORK_DIR}/i-null.vortex'; +---- +0 NULL + +query IR +SELECT count(x), avg(x) FROM '${WORK_DIR}/f-null.vortex'; +---- +0 NULL + +query IR +SELECT count(x), avg(x) FROM '${WORK_DIR}/f-nan.vortex'; +---- +3 NaN diff --git a/vortex-sqllogictest/slt/duckdb/mean_all_null.slt b/vortex-sqllogictest/slt/duckdb/mean_all_null.slt new file mode 100644 index 00000000000..78babdd7c39 --- /dev/null +++ b/vortex-sqllogictest/slt/duckdb/mean_all_null.slt @@ -0,0 +1,26 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright the Vortex contributors + +include ../setup.slt.no + +statement ok +COPY (SELECT * FROM (VALUES (CAST(NULL AS BIGINT)),(CAST(NULL AS BIGINT))) AS t(x)) +TO '${WORK_DIR}/i-null.vortex'; + +statement ok +COPY (SELECT * FROM (VALUES (1),(2),(CAST(NULL AS BIGINT))) AS t(x)) TO '${WORK_DIR}/i-some.vortex'; + +query TT +EXPLAIN SELECT avg(x) FROM '${WORK_DIR}/i-null.vortex'; +---- +:.*UNGROUPED_AGGREGATE.* + +query R +SELECT avg(x) FROM '${WORK_DIR}/i-null.vortex'; +---- +NULL + +query R +SELECT avg(x) FROM '${WORK_DIR}/i-some.vortex'; +---- +1.5 diff --git a/vortex-sqllogictest/slt/mean_all_null.slt b/vortex-sqllogictest/slt/mean_all_null.slt new file mode 100644 index 00000000000..557d730b5c9 --- /dev/null +++ b/vortex-sqllogictest/slt/mean_all_null.slt @@ -0,0 +1,44 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright the Vortex contributors + +include ./setup.slt.no + +statement ok +COPY (SELECT * FROM (VALUES (1),(2),(CAST(NULL AS BIGINT))) AS t(x)) TO '${WORK_DIR}/i-some.vortex'; + +statement ok +COPY (SELECT * FROM (VALUES (CAST(NULL AS BIGINT)),(CAST(NULL AS BIGINT))) AS t(x)) +TO '${WORK_DIR}/i-null.vortex'; + +query R +SELECT avg(x) FROM '${WORK_DIR}/i-some.vortex'; +---- +1.5 + +query IR +SELECT count(x), avg(x) FROM '${WORK_DIR}/i-null.vortex'; +---- +0 NULL + +statement ok +COPY (SELECT * FROM (VALUES (1.0),(3.0),(CAST(NULL AS DOUBLE))) AS t(x)) +TO '${WORK_DIR}/f-some.vortex'; + +statement ok +COPY (SELECT * FROM (VALUES (CAST(NULL AS DOUBLE)),(CAST(NULL AS DOUBLE))) AS t(x)) +TO '${WORK_DIR}/f-null.vortex'; + +query R +SELECT avg(x) FROM '${WORK_DIR}/f-some.vortex'; +---- +2 + +query IR +SELECT count(x), avg(x) FROM '${WORK_DIR}/f-null.vortex'; +---- +0 NULL + +query IR +SELECT count(x), avg(x) FROM '${WORK_DIR}/i-some.vortex' WHERE x > 100; +---- +0 NULL From ba617f1f38bdb367b54966358c2755d663c3c9ce Mon Sep 17 00:00:00 2001 From: Mikhail Kot Date: Thu, 30 Jul 2026 15:27:44 +0100 Subject: [PATCH 2/3] fix Signed-off-by: Mikhail Kot --- vortex-array/src/aggregate_fn/fns/mean/mod.rs | 18 ++++++++++-------- 1 file changed, 10 insertions(+), 8 deletions(-) diff --git a/vortex-array/src/aggregate_fn/fns/mean/mod.rs b/vortex-array/src/aggregate_fn/fns/mean/mod.rs index eaabda428b9..a8f103a621d 100644 --- a/vortex-array/src/aggregate_fn/fns/mean/mod.rs +++ b/vortex-array/src/aggregate_fn/fns/mean/mod.rs @@ -97,19 +97,21 @@ impl BinaryCombined for Mean { vortex_bail!("grouped mean over decimals is not yet supported"); } let target = DType::Primitive(PType::F64, Nullability::Nullable); - let sum_cast = sum.cast(target.clone())?; - let count_cast = count.cast(target.clone())?; - let mean = sum_cast.binary(count_cast.clone(), Operator::Div)?; + let sum = sum.cast(target.clone())?; + let count = count.cast(target.clone())?; - // Nulls (and nans if skip_nan is enabled) are skipped so group with - // zero values has a count of 0, 0 / 0 = nan, and we need null. - let non_empty = count_cast + let non_zero = count .binary( - ConstantArray::new(Scalar::zero_value(&target), count_cast.len()).into_array(), + ConstantArray::new(Scalar::zero_value(&target), count.len()).into_array(), Operator::NotEq, )? .fill_null(false)?; - mean.mask(non_empty) + // if count is 0, dividing by 0 below produces NaN, and we need Null. + // mask values to skip 0 so on 0 count turnes into Null, dividing by + // Null is always Null + let count = count.mask(non_zero)?; + + sum.binary(count, Operator::Div) } fn finalize_scalar(&self, left_scalar: Scalar, right_scalar: Scalar) -> VortexResult { From 1f560bfa5b741e4938a4b25bdbf18540bf19beaf Mon Sep 17 00:00:00 2001 From: Mikhail Kot Date: Thu, 30 Jul 2026 15:28:34 +0100 Subject: [PATCH 3/3] fix Signed-off-by: Mikhail Kot --- vortex-array/src/aggregate_fn/fns/mean/mod.rs | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/vortex-array/src/aggregate_fn/fns/mean/mod.rs b/vortex-array/src/aggregate_fn/fns/mean/mod.rs index a8f103a621d..da93d2352a3 100644 --- a/vortex-array/src/aggregate_fn/fns/mean/mod.rs +++ b/vortex-array/src/aggregate_fn/fns/mean/mod.rs @@ -107,7 +107,7 @@ impl BinaryCombined for Mean { )? .fill_null(false)?; // if count is 0, dividing by 0 below produces NaN, and we need Null. - // mask values to skip 0 so on 0 count turnes into Null, dividing by + // mask values to skip 0 so on 0 count turns into Null, dividing by // Null is always Null let count = count.mask(non_zero)?;