Skip to main content

vortex_array/aggregate_fn/fns/mean/
mod.rs

1// SPDX-License-Identifier: Apache-2.0
2// SPDX-FileCopyrightText: Copyright the Vortex contributors
3
4use vortex_error::VortexResult;
5use vortex_error::vortex_bail;
6use vortex_session::registry::CachedId;
7
8use crate::ArrayRef;
9use crate::ExecutionCtx;
10use crate::IntoArray;
11use crate::aggregate_fn::Accumulator;
12use crate::aggregate_fn::AggregateFnId;
13use crate::aggregate_fn::DynAccumulator;
14use crate::aggregate_fn::NumericalAggregateOpts;
15use crate::aggregate_fn::combined::BinaryCombined;
16use crate::aggregate_fn::combined::Combined;
17use crate::aggregate_fn::combined::CombinedOptions;
18use crate::aggregate_fn::combined::PairOptions;
19use crate::aggregate_fn::fns::count::Count;
20use crate::aggregate_fn::fns::sum::Sum;
21use crate::aggregate_fn::fns::sum::sum_decimal_dtype;
22use crate::arrays::ConstantArray;
23use crate::builtins::ArrayBuiltins;
24use crate::dtype::DType;
25use crate::dtype::DecimalDType;
26use crate::dtype::MAX_PRECISION;
27use crate::dtype::MAX_SCALE;
28use crate::dtype::Nullability;
29use crate::dtype::PType;
30use crate::dtype::i256;
31use crate::scalar::DecimalValue;
32use crate::scalar::Scalar;
33use crate::scalar_fn::fns::operators::Operator;
34
35/// Compute the arithmetic mean of an array.
36///
37/// See [`Mean`] for details.
38pub fn mean(array: &ArrayRef, ctx: &mut ExecutionCtx) -> VortexResult<Scalar> {
39    let mut acc = Accumulator::try_new(
40        Mean::combined(),
41        PairOptions(
42            NumericalAggregateOpts::default(),
43            NumericalAggregateOpts::default(),
44        ),
45        array.dtype().clone(),
46    )?;
47    acc.accumulate(array, ctx)?;
48    acc.finish()
49}
50
51/// Compute the arithmetic mean of an array.
52///
53/// Implemented as `Sum / Count` via [`BinaryCombined`].
54///
55/// Booleans and primitive numeric types produce nullable `f64` results.
56/// Decimals produce a nullable decimal result.
57#[derive(Clone, Debug)]
58pub struct Mean;
59
60impl Mean {
61    pub fn combined() -> Combined<Self> {
62        Combined(Mean)
63    }
64}
65
66impl BinaryCombined for Mean {
67    type Left = Sum;
68    type Right = Count;
69
70    fn id(&self) -> AggregateFnId {
71        static ID: CachedId = CachedId::new("vortex.mean");
72        *ID
73    }
74
75    fn left(&self) -> Sum {
76        Sum
77    }
78
79    fn right(&self) -> Count {
80        Count
81    }
82
83    fn left_name(&self) -> &'static str {
84        "sum"
85    }
86
87    fn right_name(&self) -> &'static str {
88        "count"
89    }
90
91    fn return_dtype(&self, input_dtype: &DType) -> Option<DType> {
92        Some(mean_output_dtype(input_dtype)?.with_nullability(Nullability::Nullable))
93    }
94
95    fn finalize(&self, sum: ArrayRef, count: ArrayRef) -> VortexResult<ArrayRef> {
96        if let DType::Decimal(..) = sum.dtype() {
97            vortex_bail!("grouped mean over decimals is not yet supported");
98        }
99        let target = DType::Primitive(PType::F64, Nullability::Nullable);
100        let sum = sum.cast(target.clone())?;
101        let count = count.cast(target.clone())?;
102
103        let non_zero = count
104            .binary(
105                ConstantArray::new(Scalar::zero_value(&target), count.len()).into_array(),
106                Operator::NotEq,
107            )?
108            .fill_null(false)?;
109        // if count is 0, dividing by 0 below produces NaN, and we need Null.
110        // mask values to skip 0 so on 0 count turns into Null, dividing by
111        // Null is always Null
112        let count = count.mask(non_zero)?;
113
114        sum.binary(count, Operator::Div)
115    }
116
117    fn finalize_scalar(&self, left_scalar: Scalar, right_scalar: Scalar) -> VortexResult<Scalar> {
118        if let DType::Decimal(decimal_dtype, _) = *left_scalar.dtype() {
119            return finalize_decimal_scalar(&left_scalar, &right_scalar, decimal_dtype);
120        }
121
122        let target = DType::Primitive(PType::F64, Nullability::Nullable);
123        let sum_cast = left_scalar.cast(&target)?;
124        let count_cast = right_scalar.cast(&target)?;
125
126        let sum = sum_cast.as_primitive().typed_value::<f64>();
127        let count = count_cast.as_primitive().typed_value::<f64>();
128        let value = match (sum, count) {
129            // None sum means sum overflowed, 0 count means empty input
130            (None, _) | (_, None) | (_, Some(0.0)) => return Ok(Scalar::null(target)),
131            (Some(s), Some(c)) => s / c,
132        };
133        Ok(Scalar::primitive(value, Nullability::Nullable))
134    }
135
136    fn serialize(&self, _options: &CombinedOptions<Self>) -> VortexResult<Option<Vec<u8>>> {
137        unimplemented!("mean is not yet serializable");
138    }
139}
140
141fn mean_output_dtype(input_dtype: &DType) -> Option<DType> {
142    match input_dtype {
143        DType::Bool(_) | DType::Primitive(..) => {
144            Some(DType::Primitive(PType::F64, Nullability::Nullable))
145        }
146        DType::Decimal(decimal_dtype, _) => Some(DType::Decimal(
147            mean_decimal_dtype(&sum_decimal_dtype(decimal_dtype)),
148            Nullability::Nullable,
149        )),
150        _ => None,
151    }
152}
153
154/// mean() output decimal type mimicking Spark/DataFusion/MySQL: decimal(p+4, s+4)
155fn mean_decimal_dtype(sum: &DecimalDType) -> DecimalDType {
156    DecimalDType::new(
157        u8::min(MAX_PRECISION, sum.precision().saturating_sub(6)),
158        i8::min(MAX_SCALE, sum.scale() + 4),
159    )
160}
161
162fn finalize_decimal_scalar(
163    sum: &Scalar,
164    count: &Scalar,
165    sum_decimal: DecimalDType,
166) -> VortexResult<Scalar> {
167    let target_decimal_dtype = mean_decimal_dtype(&sum_decimal);
168    let target_dtype = DType::Decimal(target_decimal_dtype, Nullability::Nullable);
169
170    // overflow
171    let Some(sum_value) = sum.as_decimal().decimal_value() else {
172        return Ok(Scalar::null(target_dtype));
173    };
174    // empty input
175    let count = count.as_primitive().typed_value::<u64>().unwrap_or(0);
176    if count == 0 {
177        return Ok(Scalar::null(target_dtype));
178    }
179
180    let Ok(sum) = DecimalValue::rescale_i256(
181        sum_value.as_i256(),
182        sum_decimal.scale(),
183        target_decimal_dtype.scale(),
184    ) else {
185        return Ok(Scalar::null(target_dtype));
186    };
187    let mean = sum / i256::from_i128(i128::from(count));
188
189    let Ok(mean) = DecimalValue::try_from_i256(mean, target_decimal_dtype) else {
190        return Ok(Scalar::null(target_dtype));
191    };
192    Ok(Scalar::decimal(
193        mean,
194        target_decimal_dtype,
195        Nullability::Nullable,
196    ))
197}
198
199#[cfg(test)]
200mod tests {
201    use vortex_buffer::buffer;
202    use vortex_error::VortexResult;
203
204    use super::*;
205    use crate::VortexSessionExecute;
206    use crate::aggregate_fn::DynGroupedAccumulator;
207    use crate::aggregate_fn::GroupedAccumulator;
208    use crate::array_session;
209    use crate::arrays::BoolArray;
210    use crate::arrays::ChunkedArray;
211    use crate::arrays::DecimalArray;
212    use crate::arrays::FixedSizeListArray;
213    use crate::arrays::PrimitiveArray;
214    use crate::dtype::DecimalDType;
215    use crate::validity::Validity;
216
217    #[test]
218    fn mean_all_valid() -> VortexResult<()> {
219        let array = PrimitiveArray::new(buffer![1.0f64, 2.0, 3.0, 4.0, 5.0], Validity::NonNullable)
220            .into_array();
221        let mut ctx = array_session().create_execution_ctx();
222        let result = mean(&array, &mut ctx)?;
223        assert_eq!(result.as_primitive().as_::<f64>(), Some(3.0));
224        Ok(())
225    }
226
227    #[test]
228    fn mean_with_nulls() -> VortexResult<()> {
229        let array = PrimitiveArray::from_option_iter([Some(2.0f64), None, Some(4.0)]).into_array();
230        let mut ctx = array_session().create_execution_ctx();
231        let result = mean(&array, &mut ctx)?;
232        assert_eq!(result.as_primitive().as_::<f64>(), Some(3.0));
233        Ok(())
234    }
235
236    #[test]
237    fn mean_integers() -> VortexResult<()> {
238        let array = PrimitiveArray::new(buffer![10i32, 20, 30], Validity::NonNullable).into_array();
239        let mut ctx = array_session().create_execution_ctx();
240        let result = mean(&array, &mut ctx)?;
241        assert_eq!(result.as_primitive().as_::<f64>(), Some(20.0));
242        Ok(())
243    }
244
245    #[test]
246    fn mean_bool() -> VortexResult<()> {
247        let array: BoolArray = [true, false, true, true].into_iter().collect();
248        let mut ctx = array_session().create_execution_ctx();
249        let result = mean(&array.into_array(), &mut ctx)?;
250        assert_eq!(result.as_primitive().as_::<f64>(), Some(0.75));
251        Ok(())
252    }
253
254    #[test]
255    fn mean_constant_non_null() -> VortexResult<()> {
256        let array = ConstantArray::new(5.0f64, 4);
257        let mut ctx = array_session().create_execution_ctx();
258        let result = mean(&array.into_array(), &mut ctx)?;
259        assert_eq!(result.as_primitive().as_::<f64>(), Some(5.0));
260        Ok(())
261    }
262
263    #[test]
264    fn mean_chunked() -> VortexResult<()> {
265        let chunk1 = PrimitiveArray::from_option_iter([Some(1.0f64), None, Some(3.0)]);
266        let chunk2 = PrimitiveArray::from_option_iter([Some(5.0f64), None]);
267        let dtype = chunk1.dtype().clone();
268        let chunked = ChunkedArray::try_new(vec![chunk1.into_array(), chunk2.into_array()], dtype)?;
269        let mut ctx = array_session().create_execution_ctx();
270        let result = mean(&chunked.into_array(), &mut ctx)?;
271        assert_eq!(result.as_primitive().as_::<f64>(), Some(3.0));
272        Ok(())
273    }
274
275    #[test]
276    fn mean_skips_nans_by_default() -> VortexResult<()> {
277        // NaNs are excluded from both the sum and the count.
278        let array =
279            PrimitiveArray::new(buffer![1.0f64, f64::NAN, 3.0], Validity::NonNullable).into_array();
280        let mut ctx = array_session().create_execution_ctx();
281        let result = mean(&array, &mut ctx)?;
282        assert_eq!(result.as_primitive().as_::<f64>(), Some(2.0));
283        Ok(())
284    }
285
286    #[test]
287    fn mean_with_nan_not_skipping() -> VortexResult<()> {
288        let array =
289            PrimitiveArray::new(buffer![1.0f64, f64::NAN, 3.0], Validity::NonNullable).into_array();
290        let mut ctx = array_session().create_execution_ctx();
291        let keep_nans = NumericalAggregateOpts::include_nans();
292        let mut acc = Accumulator::try_new(
293            Mean::combined(),
294            PairOptions(keep_nans, keep_nans),
295            array.dtype().clone(),
296        )?;
297        acc.accumulate(&array, &mut ctx)?;
298        let result = acc.finish()?;
299        assert!(result.as_primitive().as_::<f64>().is_some_and(f64::is_nan));
300        Ok(())
301    }
302
303    #[test]
304    fn mean_all_null_returns_null() -> VortexResult<()> {
305        let array = PrimitiveArray::from_option_iter::<f64, _>([None, None, None]).into_array();
306        let mut ctx = array_session().create_execution_ctx();
307        let result = mean(&array, &mut ctx)?;
308        assert_eq!(result.as_primitive().as_::<f64>(), None);
309        Ok(())
310    }
311
312    #[test]
313    fn mean_decimal() -> VortexResult<()> {
314        let dtype = DecimalDType::new(6, 2);
315        let array =
316            DecimalArray::new(buffer![100i32, 200, 300], dtype, Validity::NonNullable).into_array();
317        let mut ctx = array_session().create_execution_ctx();
318        let result = mean(&array, &mut ctx)?;
319        assert_eq!(
320            result.dtype(),
321            &DType::Decimal(DecimalDType::new(10, 6), Nullability::Nullable)
322        );
323        // mean(1.00, 2.00, 3.00) = 2.000000
324        assert_eq!(
325            result.as_decimal().decimal_value(),
326            Some(DecimalValue::I256(i256::from_i128(2_000_000)))
327        );
328        Ok(())
329    }
330
331    #[test]
332    fn mean_decimal_null() -> VortexResult<()> {
333        let dtype = DecimalDType::new(6, 2);
334        let validity = Validity::from_iter([true, false, true]);
335        let array = DecimalArray::new(buffer![150i32, 0, 450], dtype, validity).into_array();
336        let mut ctx = array_session().create_execution_ctx();
337        let result = mean(&array, &mut ctx)?;
338        // mean(1.50, 4.50) = 3.000000
339        assert_eq!(
340            result.as_decimal().decimal_value(),
341            Some(DecimalValue::I256(i256::from_i128(3_000_000)))
342        );
343        Ok(())
344    }
345
346    #[test]
347    fn mean_decimal_chunked() -> VortexResult<()> {
348        let dtype = DecimalDType::new(6, 2);
349        let validity = Validity::NonNullable;
350        let chunk1 = DecimalArray::new(buffer![100i32, 200], dtype, validity.clone()).into_array();
351        let chunk2 = DecimalArray::new(buffer![300i32, 400, 500], dtype, validity).into_array();
352        let dtype = chunk1.dtype().clone();
353        let chunked = ChunkedArray::try_new(vec![chunk1, chunk2], dtype)?;
354        let mut ctx = array_session().create_execution_ctx();
355        let result = mean(&chunked.into_array(), &mut ctx)?;
356        // mean(1.00, 2.00, 3.00, 4.00, 5.00) = 3.000000
357        assert_eq!(
358            result.as_decimal().decimal_value(),
359            Some(DecimalValue::I256(i256::from_i128(3_000_000)))
360        );
361        Ok(())
362    }
363
364    #[test]
365    fn mean_decimal_33() -> VortexResult<()> {
366        let dtype = DecimalDType::new(6, 2);
367        let buf = buffer![100i32, 0, 0];
368        let array = DecimalArray::new(buf, dtype, Validity::NonNullable).into_array();
369        let mut ctx = array_session().create_execution_ctx();
370        let result = mean(&array, &mut ctx)?;
371        // mean(1.00, 0.00, 0.00) = 1/3 => 0.333333
372        assert_eq!(
373            result.as_decimal().decimal_value(),
374            Some(DecimalValue::I256(i256::from_i128(333_333)))
375        );
376        Ok(())
377    }
378
379    #[test]
380    fn mean_multi_batch() -> VortexResult<()> {
381        let mut ctx = array_session().create_execution_ctx();
382        let dtype = DType::Primitive(PType::F64, Nullability::NonNullable);
383        let mut acc = Accumulator::try_new(
384            Mean::combined(),
385            PairOptions(
386                NumericalAggregateOpts::default(),
387                NumericalAggregateOpts::default(),
388            ),
389            dtype,
390        )?;
391
392        let batch1 =
393            PrimitiveArray::new(buffer![1.0f64, 2.0, 3.0], Validity::NonNullable).into_array();
394        acc.accumulate(&batch1, &mut ctx)?;
395
396        let batch2 = PrimitiveArray::new(buffer![4.0f64, 5.0], Validity::NonNullable).into_array();
397        acc.accumulate(&batch2, &mut ctx)?;
398
399        let result = acc.finish()?;
400        assert_eq!(result.as_primitive().as_::<f64>(), Some(3.0));
401        Ok(())
402    }
403
404    fn mean_nan_null() -> Vec<(Vec<Option<f64>>, Option<f64>)> {
405        vec![
406            (vec![Some(f64::NAN), Some(1.0), None], Some(1.0)),
407            (vec![Some(f64::NAN), Some(1.0), Some(3.0)], Some(2.0)),
408            (vec![None, None, Some(f64::NAN)], None),
409            (vec![None, None, None], None),
410            (vec![Some(1.0), Some(2.0), Some(3.0)], Some(2.0)),
411        ]
412    }
413
414    #[test]
415    fn mean_combined_partials() -> VortexResult<()> {
416        let mut ctx = array_session().create_execution_ctx();
417        for (case, (group, expected)) in mean_nan_null().into_iter().enumerate() {
418            let mut acc = Accumulator::try_new(
419                Mean::combined(),
420                PairOptions(
421                    NumericalAggregateOpts::default(),
422                    NumericalAggregateOpts::default(),
423                ),
424                DType::Primitive(PType::F64, Nullability::Nullable),
425            )?;
426            let (head, tail) = group.split_at(2);
427            let head = PrimitiveArray::from_option_iter(head.iter().copied()).into_array();
428            let tail = PrimitiveArray::from_option_iter(tail.iter().copied()).into_array();
429            acc.accumulate(&head, &mut ctx)?;
430            acc.accumulate(&tail, &mut ctx)?;
431            let result = acc.finish()?;
432            assert_eq!(result.as_primitive().as_::<f64>(), expected, "case {case}");
433        }
434        Ok(())
435    }
436
437    #[test]
438    fn mean_grouped_finalize() -> VortexResult<()> {
439        let cases = mean_nan_null();
440        let elements = PrimitiveArray::from_option_iter(
441            cases.iter().flat_map(|(group, _)| group.iter().copied()),
442        )
443        .into_array();
444        let groups = FixedSizeListArray::try_new(elements, 3, Validity::NonNullable, cases.len())?;
445
446        let mut acc = GroupedAccumulator::try_new(
447            Mean::combined(),
448            PairOptions(
449                NumericalAggregateOpts::default(),
450                NumericalAggregateOpts::default(),
451            ),
452            DType::Primitive(PType::F64, Nullability::Nullable),
453        )?;
454        let mut ctx = array_session().create_execution_ctx();
455        acc.accumulate_list(&groups.into_array(), &mut ctx)?;
456        let result = acc.finish()?;
457
458        for (case, (_, expected)) in cases.into_iter().enumerate() {
459            let actual = result.execute_scalar(case, &mut ctx)?;
460            assert_eq!(actual.as_primitive().as_::<f64>(), expected, "case {case}");
461        }
462        Ok(())
463    }
464}