1use datafusion::common::utils::SingleRowListArrayBuilder;
19use datafusion::physical_expr::aggregate::utils::Hashable;
20use datafusion::{arrow, common, error, logical_expr, scalar};
21use std::{cmp, collections, fmt, hash, mem, sync};
22
23#[derive(fmt::Debug)]
24pub struct PrimitiveModeAccumulator<T>
25where
26 T: arrow::array::ArrowPrimitiveType + Send,
27 T::Native: Eq + hash::Hash,
28{
29 value_counts: collections::HashMap<T::Native, i64>,
30 data_type: arrow::datatypes::DataType,
31}
32
33impl<T> PrimitiveModeAccumulator<T>
34where
35 T: arrow::array::ArrowPrimitiveType + Send,
36 T::Native: Eq + hash::Hash + Clone,
37{
38 pub fn new(data_type: &arrow::datatypes::DataType) -> Self {
39 Self {
40 value_counts: collections::HashMap::default(),
41 data_type: data_type.clone(),
42 }
43 }
44}
45
46impl<T> logical_expr::Accumulator for PrimitiveModeAccumulator<T>
47where
48 T: arrow::array::ArrowPrimitiveType + Send + fmt::Debug,
49 T::Native: Eq + hash::Hash + Clone + PartialOrd + fmt::Debug,
50{
51 fn update_batch(&mut self, values: &[arrow::array::ArrayRef]) -> error::Result<()> {
52 if values.is_empty() {
53 return Ok(());
54 }
55 let arr = common::cast::as_primitive_array::<T>(&values[0])?;
56
57 for value in arr.iter().flatten() {
58 let counter = self.value_counts.entry(value).or_insert(0);
59 *counter += 1;
60 }
61
62 Ok(())
63 }
64
65 fn state(&mut self) -> error::Result<Vec<scalar::ScalarValue>> {
66 let values =
67 arrow::array::PrimitiveArray::<T>::from_iter_values(self.value_counts.keys().copied())
68 .with_data_type(self.data_type.clone());
69 let counts =
70 arrow::array::Int64Array::from_iter_values(self.value_counts.values().copied());
71
72 Ok(vec![
73 SingleRowListArrayBuilder::new(sync::Arc::new(values)).build_list_scalar(),
74 SingleRowListArrayBuilder::new(sync::Arc::new(counts)).build_list_scalar(),
75 ])
76 }
77
78 fn merge_batch(&mut self, states: &[arrow::array::ArrayRef]) -> error::Result<()> {
79 super::for_each_state_row(states, |values, counts| {
80 let values = common::cast::as_primitive_array::<T>(values)?;
81 for (value, count) in values.iter().zip(counts.values()) {
82 if let Some(value) = value {
83 *self.value_counts.entry(value).or_insert(0) += *count;
84 }
85 }
86 Ok(())
87 })
88 }
89
90 fn evaluate(&mut self) -> error::Result<scalar::ScalarValue> {
91 let mut max_value: Option<T::Native> = None;
92 let mut max_count: i64 = 0;
93
94 self.value_counts.iter().for_each(|(value, &count)| {
95 match count.cmp(&max_count) {
96 cmp::Ordering::Greater => {
97 max_value = Some(*value);
98 max_count = count;
99 }
100 cmp::Ordering::Equal => {
101 max_value = match max_value {
103 Some(ref current_max_value) if value < current_max_value => Some(*value),
104 Some(ref current_max_value) => Some(*current_max_value),
105 None => Some(*value),
106 };
107 }
108 _ => {} }
110 });
111
112 scalar::ScalarValue::new_primitive::<T>(max_value, &self.data_type)
113 }
114
115 fn size(&self) -> usize {
116 mem::size_of_val(&self.value_counts)
117 + self.value_counts.len() * mem::size_of::<(T::Native, i64)>()
118 }
119}
120
121#[derive(Debug)]
122pub struct FloatModeAccumulator<T>
123where
124 T: arrow::array::ArrowPrimitiveType,
125{
126 value_counts: collections::HashMap<Hashable<T::Native>, i64>,
127 data_type: arrow::datatypes::DataType,
128}
129
130impl<T> FloatModeAccumulator<T>
131where
132 T: arrow::array::ArrowPrimitiveType,
133{
134 pub fn new(data_type: &arrow::datatypes::DataType) -> Self {
135 Self {
136 value_counts: collections::HashMap::default(),
137 data_type: data_type.clone(),
138 }
139 }
140}
141
142impl<T> logical_expr::Accumulator for FloatModeAccumulator<T>
143where
144 T: arrow::array::ArrowPrimitiveType + Send + fmt::Debug,
145 T::Native: PartialOrd + fmt::Debug + Clone,
146{
147 fn update_batch(&mut self, values: &[arrow::array::ArrayRef]) -> error::Result<()> {
148 if values.is_empty() {
149 return Ok(());
150 }
151
152 let arr = common::cast::as_primitive_array::<T>(&values[0])?;
153
154 for value in arr.iter().flatten() {
155 let counter = self.value_counts.entry(Hashable(value)).or_insert(0);
156 *counter += 1;
157 }
158
159 Ok(())
160 }
161
162 fn state(&mut self) -> error::Result<Vec<scalar::ScalarValue>> {
163 let values = arrow::array::PrimitiveArray::<T>::from_iter_values(
164 self.value_counts.keys().map(|key| key.0),
165 )
166 .with_data_type(self.data_type.clone());
167 let counts =
168 arrow::array::Int64Array::from_iter_values(self.value_counts.values().copied());
169
170 Ok(vec![
171 SingleRowListArrayBuilder::new(sync::Arc::new(values)).build_list_scalar(),
172 SingleRowListArrayBuilder::new(sync::Arc::new(counts)).build_list_scalar(),
173 ])
174 }
175
176 fn merge_batch(&mut self, states: &[arrow::array::ArrayRef]) -> error::Result<()> {
177 super::for_each_state_row(states, |values, counts| {
178 let values = common::cast::as_primitive_array::<T>(values)?;
179 for (value, count) in values.iter().zip(counts.values()) {
180 if let Some(value) = value {
181 *self.value_counts.entry(Hashable(value)).or_insert(0) += *count;
182 }
183 }
184 Ok(())
185 })
186 }
187
188 fn evaluate(&mut self) -> error::Result<scalar::ScalarValue> {
189 let mut max_value: Option<T::Native> = None;
190 let mut max_count: i64 = 0;
191
192 self.value_counts.iter().for_each(|(value, &count)| {
193 match count.cmp(&max_count) {
194 cmp::Ordering::Greater => {
195 max_value = Some(value.0);
196 max_count = count;
197 }
198 cmp::Ordering::Equal => {
199 max_value = match max_value {
201 Some(current_max_value) if value.0 < current_max_value => Some(value.0),
202 Some(current_max_value) => Some(current_max_value),
203 None => Some(value.0),
204 };
205 }
206 _ => {} }
208 });
209
210 scalar::ScalarValue::new_primitive::<T>(max_value, &self.data_type)
211 }
212
213 fn size(&self) -> usize {
214 mem::size_of_val(&self.value_counts)
215 + self.value_counts.len() * mem::size_of::<(Hashable<T::Native>, i64)>()
216 }
217}
218
219#[cfg(test)]
220mod tests {
221
222 use super::*;
223
224 use datafusion::logical_expr::Accumulator;
225 use std::sync;
226
227 fn merge_from(dest: &mut impl Accumulator, src: &mut impl Accumulator) -> error::Result<()> {
228 let arrays = src
229 .state()?
230 .iter()
231 .map(|value| value.to_array())
232 .collect::<error::Result<Vec<_>>>()?;
233 dest.merge_batch(&arrays)
234 }
235
236 #[test]
237 fn test_mode_accumulator_single_mode_int64() -> error::Result<()> {
238 let mut acc = PrimitiveModeAccumulator::<arrow::datatypes::Int64Type>::new(
239 &arrow::datatypes::DataType::Int64,
240 );
241 let values: arrow::array::ArrayRef =
242 sync::Arc::new(arrow::array::Int64Array::from(vec![1, 2, 2, 3, 3, 3]));
243 acc.update_batch(&[values])?;
244 let result = acc.evaluate()?;
245 assert_eq!(
246 result,
247 scalar::ScalarValue::new_primitive::<arrow::datatypes::Int64Type>(
248 Some(3),
249 &arrow::datatypes::DataType::Int64
250 )?
251 );
252 Ok(())
253 }
254
255 #[test]
256 fn test_mode_accumulator_with_nulls_int64() -> error::Result<()> {
257 let mut acc = PrimitiveModeAccumulator::<arrow::datatypes::Int64Type>::new(
258 &arrow::datatypes::DataType::Int64,
259 );
260 let values: arrow::array::ArrayRef = sync::Arc::new(arrow::array::Int64Array::from(vec![
261 None,
262 Some(1),
263 Some(2),
264 Some(2),
265 Some(3),
266 Some(3),
267 Some(3),
268 ]));
269 acc.update_batch(&[values])?;
270 let result = acc.evaluate()?;
271 assert_eq!(
272 result,
273 scalar::ScalarValue::new_primitive::<arrow::datatypes::Int64Type>(
274 Some(3),
275 &arrow::datatypes::DataType::Int64
276 )?
277 );
278 Ok(())
279 }
280
281 #[test]
282 fn test_mode_accumulator_tie_case_int64() -> error::Result<()> {
283 let mut acc = PrimitiveModeAccumulator::<arrow::datatypes::Int64Type>::new(
284 &arrow::datatypes::DataType::Int64,
285 );
286 let values: arrow::array::ArrayRef =
287 sync::Arc::new(arrow::array::Int64Array::from(vec![1, 2, 2, 3, 3]));
288 acc.update_batch(&[values])?;
289 let result = acc.evaluate()?;
290 assert_eq!(
291 result,
292 scalar::ScalarValue::new_primitive::<arrow::datatypes::Int64Type>(
293 Some(2),
294 &arrow::datatypes::DataType::Int64
295 )?
296 );
297 Ok(())
298 }
299
300 #[test]
301 fn test_mode_accumulator_only_nulls_int64() -> error::Result<()> {
302 let mut acc = PrimitiveModeAccumulator::<arrow::datatypes::Int64Type>::new(
303 &arrow::datatypes::DataType::Int64,
304 );
305 let values: arrow::array::ArrayRef =
306 sync::Arc::new(arrow::array::Int64Array::from(vec![None, None, None, None]));
307 acc.update_batch(&[values])?;
308 let result = acc.evaluate()?;
309 assert_eq!(
310 result,
311 scalar::ScalarValue::new_primitive::<arrow::datatypes::Int64Type>(
312 None,
313 &arrow::datatypes::DataType::Int64
314 )?
315 );
316 Ok(())
317 }
318
319 #[test]
320 fn test_mode_accumulator_merge_overlapping_keys_int64() -> error::Result<()> {
321 let mut left = PrimitiveModeAccumulator::<arrow::datatypes::Int64Type>::new(
322 &arrow::datatypes::DataType::Int64,
323 );
324 let left_values: arrow::array::ArrayRef =
325 sync::Arc::new(arrow::array::Int64Array::from(vec![2, 2, 2]));
326 left.update_batch(&[left_values])?;
327
328 let mut right = PrimitiveModeAccumulator::<arrow::datatypes::Int64Type>::new(
329 &arrow::datatypes::DataType::Int64,
330 );
331 let right_values: arrow::array::ArrayRef =
333 sync::Arc::new(arrow::array::Int64Array::from(vec![1, 1, 1, 1, 2, 2]));
334 right.update_batch(&[right_values])?;
335
336 merge_from(&mut right, &mut left)?;
337 let result = right.evaluate()?;
338 assert_eq!(
339 result,
340 scalar::ScalarValue::new_primitive::<arrow::datatypes::Int64Type>(
341 Some(2),
342 &arrow::datatypes::DataType::Int64
343 )?
344 );
345 Ok(())
346 }
347
348 #[test]
349 fn test_mode_accumulator_single_mode_float64() -> error::Result<()> {
350 let mut acc = FloatModeAccumulator::<arrow::datatypes::Float64Type>::new(
351 &arrow::datatypes::DataType::Float64,
352 );
353 let values: arrow::array::ArrayRef =
354 sync::Arc::new(arrow::array::Float64Array::from(vec![
355 1.0, 2.0, 2.0, 3.0, 3.0, 3.0,
356 ]));
357 acc.update_batch(&[values])?;
358 let result = acc.evaluate()?;
359 assert_eq!(
360 result,
361 scalar::ScalarValue::new_primitive::<arrow::datatypes::Float64Type>(
362 Some(3.0),
363 &arrow::datatypes::DataType::Float64
364 )?
365 );
366 Ok(())
367 }
368
369 #[test]
370 fn test_mode_accumulator_with_nulls_float64() -> error::Result<()> {
371 let mut acc = FloatModeAccumulator::<arrow::datatypes::Float64Type>::new(
372 &arrow::datatypes::DataType::Float64,
373 );
374 let values: arrow::array::ArrayRef =
375 sync::Arc::new(arrow::array::Float64Array::from(vec![
376 None,
377 Some(1.0),
378 Some(2.0),
379 Some(2.0),
380 Some(3.0),
381 Some(3.0),
382 Some(3.0),
383 ]));
384 acc.update_batch(&[values])?;
385 let result = acc.evaluate()?;
386 assert_eq!(
387 result,
388 scalar::ScalarValue::new_primitive::<arrow::datatypes::Float64Type>(
389 Some(3.0),
390 &arrow::datatypes::DataType::Float64
391 )?
392 );
393 Ok(())
394 }
395
396 #[test]
397 fn test_mode_accumulator_tie_case_float64() -> error::Result<()> {
398 let mut acc = FloatModeAccumulator::<arrow::datatypes::Float64Type>::new(
399 &arrow::datatypes::DataType::Float64,
400 );
401 let values: arrow::array::ArrayRef =
402 sync::Arc::new(arrow::array::Float64Array::from(vec![
403 1.0, 2.0, 2.0, 3.0, 3.0,
404 ]));
405 acc.update_batch(&[values])?;
406 let result = acc.evaluate()?;
407 assert_eq!(
408 result,
409 scalar::ScalarValue::new_primitive::<arrow::datatypes::Float64Type>(
410 Some(2.0),
411 &arrow::datatypes::DataType::Float64
412 )?
413 );
414 Ok(())
415 }
416
417 #[test]
418 fn test_mode_accumulator_only_nulls_float64() -> error::Result<()> {
419 let mut acc = FloatModeAccumulator::<arrow::datatypes::Float64Type>::new(
420 &arrow::datatypes::DataType::Float64,
421 );
422 let values: arrow::array::ArrayRef =
423 sync::Arc::new(arrow::array::Float64Array::from(vec![
424 None, None, None, None,
425 ]));
426 acc.update_batch(&[values])?;
427 let result = acc.evaluate()?;
428 assert_eq!(
429 result,
430 scalar::ScalarValue::new_primitive::<arrow::datatypes::Float64Type>(
431 None,
432 &arrow::datatypes::DataType::Float64
433 )?
434 );
435 Ok(())
436 }
437
438 #[test]
439 fn test_mode_accumulator_merge_overlapping_keys_float64() -> error::Result<()> {
440 let mut left = FloatModeAccumulator::<arrow::datatypes::Float64Type>::new(
441 &arrow::datatypes::DataType::Float64,
442 );
443 let left_values: arrow::array::ArrayRef =
444 sync::Arc::new(arrow::array::Float64Array::from(vec![2.0, 2.0, 2.0]));
445 left.update_batch(&[left_values])?;
446
447 let mut right = FloatModeAccumulator::<arrow::datatypes::Float64Type>::new(
448 &arrow::datatypes::DataType::Float64,
449 );
450 let right_values: arrow::array::ArrayRef =
451 sync::Arc::new(arrow::array::Float64Array::from(vec![
452 1.0, 1.0, 1.0, 1.0, 2.0, 2.0,
453 ]));
454 right.update_batch(&[right_values])?;
455
456 merge_from(&mut right, &mut left)?;
457 let result = right.evaluate()?;
458 assert_eq!(
459 result,
460 scalar::ScalarValue::new_primitive::<arrow::datatypes::Float64Type>(
461 Some(2.0),
462 &arrow::datatypes::DataType::Float64
463 )?
464 );
465 Ok(())
466 }
467
468 #[test]
469 fn test_mode_accumulator_single_mode_date64() -> error::Result<()> {
470 let mut acc = PrimitiveModeAccumulator::<arrow::datatypes::Date64Type>::new(
471 &arrow::datatypes::DataType::Date64,
472 );
473 let values: arrow::array::ArrayRef = sync::Arc::new(arrow::array::Date64Array::from(vec![
474 1609459200000,
475 1609545600000,
476 1609545600000,
477 1609632000000,
478 1609632000000,
479 1609632000000,
480 ]));
481 acc.update_batch(&[values])?;
482 let result = acc.evaluate()?;
483 assert_eq!(
484 result,
485 scalar::ScalarValue::new_primitive::<arrow::datatypes::Date64Type>(
486 Some(1609632000000),
487 &arrow::datatypes::DataType::Date64
488 )?
489 );
490 Ok(())
491 }
492
493 #[test]
494 fn test_mode_accumulator_with_nulls_date64() -> error::Result<()> {
495 let mut acc = PrimitiveModeAccumulator::<arrow::datatypes::Date64Type>::new(
496 &arrow::datatypes::DataType::Date64,
497 );
498 let values: arrow::array::ArrayRef = sync::Arc::new(arrow::array::Date64Array::from(vec![
499 None,
500 Some(1609459200000),
501 Some(1609545600000),
502 Some(1609545600000),
503 Some(1609632000000),
504 Some(1609632000000),
505 Some(1609632000000),
506 ]));
507 acc.update_batch(&[values])?;
508 let result = acc.evaluate()?;
509 assert_eq!(
510 result,
511 scalar::ScalarValue::new_primitive::<arrow::datatypes::Date64Type>(
512 Some(1609632000000),
513 &arrow::datatypes::DataType::Date64
514 )?
515 );
516 Ok(())
517 }
518
519 #[test]
520 fn test_mode_accumulator_tie_case_date64() -> error::Result<()> {
521 let mut acc = PrimitiveModeAccumulator::<arrow::datatypes::Date64Type>::new(
522 &arrow::datatypes::DataType::Date64,
523 );
524 let values: arrow::array::ArrayRef = sync::Arc::new(arrow::array::Date64Array::from(vec![
525 1609459200000,
526 1609545600000,
527 1609545600000,
528 1609632000000,
529 1609632000000,
530 ]));
531 acc.update_batch(&[values])?;
532 let result = acc.evaluate()?;
533 assert_eq!(
534 result,
535 scalar::ScalarValue::new_primitive::<arrow::datatypes::Date64Type>(
536 Some(1609545600000),
537 &arrow::datatypes::DataType::Date64
538 )?
539 );
540 Ok(())
541 }
542
543 #[test]
544 fn test_mode_accumulator_only_nulls_date64() -> error::Result<()> {
545 let mut acc = PrimitiveModeAccumulator::<arrow::datatypes::Date64Type>::new(
546 &arrow::datatypes::DataType::Date64,
547 );
548 let values: arrow::array::ArrayRef = sync::Arc::new(arrow::array::Date64Array::from(vec![
549 None, None, None, None,
550 ]));
551 acc.update_batch(&[values])?;
552 let result = acc.evaluate()?;
553 assert_eq!(
554 result,
555 scalar::ScalarValue::new_primitive::<arrow::datatypes::Date64Type>(
556 None,
557 &arrow::datatypes::DataType::Date64
558 )?
559 );
560 Ok(())
561 }
562
563 #[test]
564 fn test_mode_accumulator_single_mode_time64() -> error::Result<()> {
565 let mut acc = PrimitiveModeAccumulator::<arrow::datatypes::Time64MicrosecondType>::new(
566 &arrow::datatypes::DataType::Time64(arrow::datatypes::TimeUnit::Microsecond),
567 );
568 let values: arrow::array::ArrayRef =
569 sync::Arc::new(arrow::array::Time64MicrosecondArray::from(vec![
570 3600000000,
571 7200000000,
572 7200000000,
573 10800000000,
574 10800000000,
575 10800000000,
576 ]));
577 acc.update_batch(&[values])?;
578 let result = acc.evaluate()?;
579 assert_eq!(
580 result,
581 scalar::ScalarValue::new_primitive::<arrow::datatypes::Time64MicrosecondType>(
582 Some(10800000000),
583 &arrow::datatypes::DataType::Time64(arrow::datatypes::TimeUnit::Microsecond)
584 )?
585 );
586 Ok(())
587 }
588
589 #[test]
590 fn test_mode_accumulator_with_nulls_time64() -> error::Result<()> {
591 let mut acc = PrimitiveModeAccumulator::<arrow::datatypes::Time64MicrosecondType>::new(
592 &arrow::datatypes::DataType::Time64(arrow::datatypes::TimeUnit::Microsecond),
593 );
594 let values: arrow::array::ArrayRef =
595 sync::Arc::new(arrow::array::Time64MicrosecondArray::from(vec![
596 None,
597 Some(3600000000),
598 Some(7200000000),
599 Some(7200000000),
600 Some(10800000000),
601 Some(10800000000),
602 Some(10800000000),
603 ]));
604 acc.update_batch(&[values])?;
605 let result = acc.evaluate()?;
606 assert_eq!(
607 result,
608 scalar::ScalarValue::new_primitive::<arrow::datatypes::Time64MicrosecondType>(
609 Some(10800000000),
610 &arrow::datatypes::DataType::Time64(arrow::datatypes::TimeUnit::Microsecond)
611 )?
612 );
613 Ok(())
614 }
615
616 #[test]
617 fn test_mode_accumulator_tie_case_time64() -> error::Result<()> {
618 let mut acc = PrimitiveModeAccumulator::<arrow::datatypes::Time64MicrosecondType>::new(
619 &arrow::datatypes::DataType::Time64(arrow::datatypes::TimeUnit::Microsecond),
620 );
621 let values: arrow::array::ArrayRef =
622 sync::Arc::new(arrow::array::Time64MicrosecondArray::from(vec![
623 3600000000,
624 7200000000,
625 7200000000,
626 10800000000,
627 10800000000,
628 ]));
629 acc.update_batch(&[values])?;
630 let result = acc.evaluate()?;
631 assert_eq!(
632 result,
633 scalar::ScalarValue::new_primitive::<arrow::datatypes::Time64MicrosecondType>(
634 Some(7200000000),
635 &arrow::datatypes::DataType::Time64(arrow::datatypes::TimeUnit::Microsecond)
636 )?
637 );
638 Ok(())
639 }
640
641 #[test]
642 fn test_mode_accumulator_only_nulls_time64() -> error::Result<()> {
643 let mut acc = PrimitiveModeAccumulator::<arrow::datatypes::Time64MicrosecondType>::new(
644 &arrow::datatypes::DataType::Time64(arrow::datatypes::TimeUnit::Microsecond),
645 );
646 let values: arrow::array::ArrayRef =
647 sync::Arc::new(arrow::array::Time64MicrosecondArray::from(vec![
648 None, None, None, None,
649 ]));
650 acc.update_batch(&[values])?;
651 let result = acc.evaluate()?;
652 assert_eq!(
653 result,
654 scalar::ScalarValue::new_primitive::<arrow::datatypes::Time64MicrosecondType>(
655 None,
656 &arrow::datatypes::DataType::Time64(arrow::datatypes::TimeUnit::Microsecond)
657 )?
658 );
659 Ok(())
660 }
661}