Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
95 changes: 85 additions & 10 deletions vortex-array/src/aggregate_fn/fns/mean/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand All @@ -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;
Expand Down Expand Up @@ -95,9 +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)?;
sum_cast.binary(count_cast, Operator::Div)
let sum = sum.cast(target.clone())?;
let count = count.cast(target.clone())?;

let non_zero = count
.binary(
ConstantArray::new(Scalar::zero_value(&target), count.len()).into_array(),
Operator::NotEq,
)?
.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 turns 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<Scalar> {
Expand All @@ -112,9 +126,8 @@ impl BinaryCombined for Mean {
let sum = sum_cast.as_primitive().typed_value::<f64>();
let count = count_cast.as_primitive().typed_value::<f64>();
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))
Expand Down Expand Up @@ -217,13 +230,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;
Expand Down Expand Up @@ -315,11 +329,11 @@ mod tests {
}

#[test]
fn mean_all_null_returns_nan() -> VortexResult<()> {
fn mean_all_null_returns_null() -> VortexResult<()> {
let array = PrimitiveArray::from_option_iter::<f64, _>([None, None, None]).into_array();
let mut ctx = array_session().create_execution_ctx();
let result = mean(&array, &mut ctx)?;
assert!(result.as_primitive().as_::<f64>().is_some_and(f64::is_nan));
assert_eq!(result.as_primitive().as_::<f64>(), None);
Ok(())
}

Expand Down Expand Up @@ -414,4 +428,65 @@ mod tests {
assert_eq!(result.as_primitive().as_::<f64>(), Some(3.0));
Ok(())
}

fn mean_nan_null() -> Vec<(Vec<Option<f64>>, Option<f64>)> {
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_::<f64>(), 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_::<f64>(), expected, "case {case}");
}
Ok(())
}
}
30 changes: 30 additions & 0 deletions vortex-sqllogictest/slt/datafusion/mean_all_null.slt
Original file line number Diff line number Diff line change
@@ -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
26 changes: 26 additions & 0 deletions vortex-sqllogictest/slt/duckdb/mean_all_null.slt
Original file line number Diff line number Diff line change
@@ -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';
----
<!REGEX>:.*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
44 changes: 44 additions & 0 deletions vortex-sqllogictest/slt/mean_all_null.slt
Original file line number Diff line number Diff line change
@@ -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
Loading