From 0fa7cd9d93c4b0e1061725c9f22230abefc65078 Mon Sep 17 00:00:00 2001 From: discord9 Date: Wed, 15 Jul 2026 11:47:18 +0800 Subject: [PATCH] fix(backport): account for grouped aggregate memory Backport the grouped AVG and MEDIAN accumulator memory accounting fix from upstream DataFusion commit 73d5d780. Signed-off-by: discord9 --- datafusion/functions-aggregate/src/average.rs | 39 ++++++++++++++++- datafusion/functions-aggregate/src/median.rs | 43 +++++++++++++++++-- 2 files changed, 77 insertions(+), 5 deletions(-) diff --git a/datafusion/functions-aggregate/src/average.rs b/datafusion/functions-aggregate/src/average.rs index 543116db1ddb6..6c220ae28f0f5 100644 --- a/datafusion/functions-aggregate/src/average.rs +++ b/datafusion/functions-aggregate/src/average.rs @@ -956,6 +956,43 @@ where } fn size(&self) -> usize { - self.counts.capacity() * size_of::() + self.sums.capacity() * size_of::() + size_of_val(&self.counts) + + self.counts.capacity() * size_of::() + + size_of_val(&self.sums) + + self.sums.capacity() * size_of::() + + self.null_state.size() + } +} + +#[cfg(test)] +mod tests { + use super::*; + use arrow::array::Float64Array; + + #[test] + fn test_groups_accumulator_size_accounts_for_owned_state() -> Result<()> { + let mut accumulator = AvgGroupsAccumulator::::new( + &DataType::Float64, + &DataType::Float64, + |sum, count| Ok(sum / count as f64), + ); + let values: ArrayRef = Arc::new(Float64Array::from_iter( + (0..64).map(|value| (value != 0).then_some(value as f64)), + )); + let input_size = values.get_array_memory_size(); + let group_indices = (0..64).collect::>(); + + accumulator.update_batch(&[values], &group_indices, None, 64)?; + + let expected_size = size_of_val(&accumulator.counts) + + accumulator.counts.capacity() * size_of::() + + size_of_val(&accumulator.sums) + + accumulator.sums.capacity() * size_of::() + + accumulator.null_state.size(); + + // The input Arrow array is borrowed during update and is not owned by the accumulator. + assert_eq!(accumulator.size(), expected_size); + assert_ne!(accumulator.size(), expected_size + input_size); + Ok(()) } } diff --git a/datafusion/functions-aggregate/src/median.rs b/datafusion/functions-aggregate/src/median.rs index db769918d1353..2eee5dcbe422e 100644 --- a/datafusion/functions-aggregate/src/median.rs +++ b/datafusion/functions-aggregate/src/median.rs @@ -528,12 +528,47 @@ impl GroupsAccumulator for MedianGroupsAccumulator usize { - self.group_values - .iter() - .map(|values| values.capacity() * size_of::()) + size_of_val(&self.group_values) + + self + .group_values + .iter() + .map(|values| values.capacity() * size_of::()) .sum::() // account for size of self.grou_values too - + self.group_values.capacity() * size_of::>() + + self.group_values.capacity() * size_of::>() + } +} + +#[cfg(test)] +mod tests { + use super::*; + use arrow::array::Int8Array; + + #[test] + fn test_groups_accumulator_size_accounts_for_owned_state() -> Result<()> { + let mut accumulator = + MedianGroupsAccumulator::::new(DataType::Int8); + let values: ArrayRef = Arc::new(Int8Array::from_iter_values( + (0..8).flat_map(|_| (0..128).map(|value| value as i8)), + )); + let group_indices = (0..8) + .flat_map(|group| std::iter::repeat_n(group, 128)) + .collect::>(); + + assert_eq!(accumulator.size(), size_of_val(&accumulator.group_values)); + + accumulator.update_batch(&[values], &group_indices, None, 8)?; + + let expected_size = size_of_val(&accumulator.group_values) + + accumulator + .group_values + .iter() + .map(|values| values.capacity() * size_of::()) + .sum::() + + accumulator.group_values.capacity() * size_of::>(); + + assert_eq!(accumulator.size(), expected_size); + Ok(()) } }