1use 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
35pub 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#[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 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, _) | (_, 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
154fn 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 let Some(sum_value) = sum.as_decimal().decimal_value() else {
172 return Ok(Scalar::null(target_dtype));
173 };
174 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 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 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 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 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 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}