diff --git a/vortex-array/src/aggregate_fn/fns/mean/mod.rs b/vortex-array/src/aggregate_fn/fns/mean/mod.rs index 061ed92c4b1..41b542e6940 100644 --- a/vortex-array/src/aggregate_fn/fns/mean/mod.rs +++ b/vortex-array/src/aggregate_fn/fns/mean/mod.rs @@ -6,6 +6,7 @@ use vortex_error::vortex_bail; use crate::ArrayRef; use crate::ExecutionCtx; +use crate::IntoArray; use crate::aggregate_fn::Accumulator; use crate::aggregate_fn::AggregateFnId; use crate::aggregate_fn::AggregateFnVTable; @@ -17,6 +18,7 @@ use crate::aggregate_fn::combined::CombinedOptions; use crate::aggregate_fn::combined::PairOptions; use crate::aggregate_fn::fns::count::Count; use crate::aggregate_fn::fns::sum::Sum; +use crate::arrays::ConstantArray; use crate::builtins::ArrayBuiltins; use crate::dtype::DType; use crate::dtype::Nullability; @@ -91,8 +93,19 @@ impl BinaryCombined for Mean { _ => 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 are skipped during accumulation, so an all-null group has a count of zero and + // the division produces 0/0 = NaN. The mean of an empty group is null (as in SQL), so + // mask out zero-count entries. This matches `finalize_scalar`. + let non_empty = count_cast + .binary( + ConstantArray::new(Scalar::zero_value(&target), count.len()).into_array(), + Operator::NotEq, + )? + // A null count means a null group; keep it masked out. + .fill_null(false)?; + mean.mask(non_empty) } fn finalize_scalar(&self, left_scalar: Scalar, right_scalar: Scalar) -> VortexResult { @@ -107,9 +120,9 @@ 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. + // A `None` sum means the sum overflowed; a count of zero means an all-null (empty) + // input. Both are null, matching the array `finalize` path. + (None, _) | (_, None) | (_, Some(0.0)) => return Ok(Scalar::null(target)), (Some(s), Some(c)) => s / c, }; Ok(Scalar::primitive(value, Nullability::Nullable)) @@ -167,12 +180,13 @@ 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::FixedSizeListArray; use crate::arrays::PrimitiveArray; use crate::validity::Validity; @@ -263,11 +277,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(()) } @@ -295,4 +309,70 @@ mod tests { assert_eq!(result.as_primitive().as_::(), Some(3.0)); Ok(()) } + + /// Groups exercised by both finalize paths, under the default skip-NaN options. NaNs and nulls + /// are skipped during accumulation, so a group with no non-null, non-NaN values is empty and + /// its mean is null. + fn mean_cases() -> 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_via_combined_partials() -> VortexResult<()> { + let mut ctx = array_session().create_execution_ctx(); + for (case, (group, expected)) in mean_cases().into_iter().enumerate() { + let mut acc = Accumulator::try_new( + Mean::combined(), + PairOptions( + NumericalAggregateOpts::default(), + NumericalAggregateOpts::default(), + ), + DType::Primitive(PType::F64, Nullability::Nullable), + )?; + // Two batches per group so the result goes through partial combination and + // `finalize_scalar`. + 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_via_grouped_finalize() -> VortexResult<()> { + let cases = mean_cases(); + 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(()) + } }