Skip to content
3 changes: 2 additions & 1 deletion datafusion/expr-common/src/accumulator.rs
Original file line number Diff line number Diff line change
Expand Up @@ -97,7 +97,7 @@ impl Drop for AggregateMetricTimer<'_> {
/// `Accumulator`s are stateful objects that implement a single group. They
/// aggregate values from multiple rows together into a final output aggregate.
///
/// [`GroupsAccumulator]` is an additional more performant (but also complex) API
/// [`GroupsAccumulator`] is an additional more performant (but also complex) API
/// that manages state for multiple groups at once.
///
/// An accumulator knows how to:
Expand All @@ -118,6 +118,7 @@ impl Drop for AggregateMetricTimer<'_> {
/// [`state`]: Self::state
/// [`evaluate`]: Self::evaluate
/// [`merge_batch`]: Self::merge_batch
/// [`GroupsAccumulator`]: crate::groups_accumulator::GroupsAccumulator
/// [window function]: https://en.wikipedia.org/wiki/Window_function_(SQL)
pub trait Accumulator: Send + Sync + Debug + std::any::Any {
/// Supplies optional metrics owned by this aggregate expression.
Expand Down
118 changes: 70 additions & 48 deletions datafusion/functions-aggregate/src/approx_distinct.rs
Original file line number Diff line number Diff line change
Expand Up @@ -126,6 +126,65 @@ impl<A: Accumulator> Accumulator for ApproxDistinctBitmapWrapper<A> {
}
}

/// A validated, zero-copy view of a serialized `approx_distinct` partial state,
/// as produced by [`GroupHll::serialize`] or by the per-group [`Accumulator`]s.
enum SerializedHll<'a> {
/// The raw [`NUM_REGISTERS`] registers of a dense sketch.
Dense(&'a [u8; NUM_REGISTERS]),
/// Little-endian hashes of at most [`SPARSE_LIMIT`] distinct values. An
/// empty state decodes as an empty sparse state.
Sparse(&'a [[u8; size_of::<u64>()]]),
}

impl<'a> SerializedHll<'a> {
fn decode(bytes: &'a [u8]) -> Result<Self> {
if let Ok(registers) = <&[u8; NUM_REGISTERS]>::try_from(bytes) {
return Ok(Self::Dense(registers));
}
let (chunks, rest) = bytes.as_chunks::<{ size_of::<u64>() }>();
if !rest.is_empty() {
return internal_err!(
"approx_distinct: malformed sparse state: length {} is not a multiple of {}",
bytes.len(),
size_of::<u64>()
);
}
if chunks.len() > SPARSE_LIMIT {
return internal_err!(
"approx_distinct: malformed sparse state: length {} exceeds sparse limit {}",
bytes.len(),
SPARSE_LIMIT * size_of::<u64>()
);
}
Ok(Self::Sparse(chunks))
}
}

/// Merge the serialized partial states in `states` into `hll`.
fn merge_states<T: Hash + ?Sized>(
hll: &mut HyperLogLog<T>,
states: &[ArrayRef],
) -> Result<()> {
assert_eq!(1, states.len(), "expect only 1 element in the states");
let binary_array = downcast_value!(states[0], BinaryArray);
for v in binary_array.iter() {
let v = v.ok_or_else(|| {
internal_datafusion_err!("Impossibly got empty binary array from states")
})?;
match SerializedHll::decode(v)? {
SerializedHll::Dense(registers) => {
hll.merge(&HyperLogLog::new_with_registers(*registers));
}
SerializedHll::Sparse(chunks) => {
for chunk in chunks {
hll.add_hashed(u64::from_le_bytes(*chunk));
}
}
}
}
Ok(())
}

#[derive(Debug)]
struct HLLAccumulator {
hll: HyperLogLog<u8>,
Expand Down Expand Up @@ -166,16 +225,7 @@ impl Accumulator for HLLAccumulator {
}

fn merge_batch(&mut self, states: &[ArrayRef]) -> Result<()> {
assert_eq!(1, states.len(), "expect only 1 element in the states");
let binary_array = downcast_value!(states[0], BinaryArray);
for v in binary_array.iter() {
let v = v.ok_or_else(|| {
internal_datafusion_err!("Impossibly got empty binary array from states")
})?;
let other = v.try_into()?;
self.hll.merge(&other);
}
Ok(())
merge_states(&mut self.hll, states)
}

fn state(&mut self) -> Result<Vec<ScalarValue>> {
Expand Down Expand Up @@ -226,16 +276,7 @@ where
}

fn merge_batch(&mut self, states: &[ArrayRef]) -> Result<()> {
assert_eq!(1, states.len(), "expect only 1 element in the states");
let binary_array = downcast_value!(states[0], BinaryArray);
for v in binary_array.iter() {
let v = v.ok_or_else(|| {
internal_datafusion_err!("Impossibly got empty binary array from states")
})?;
let other = v.try_into()?;
self.hll.merge(&other);
}
Ok(())
merge_states(&mut self.hll, states)
}

fn state(&mut self) -> Result<Vec<ScalarValue>> {
Expand Down Expand Up @@ -335,36 +376,17 @@ impl GroupHll {
}
}

/// Merge a serialized state (produced by [`Self::serialize`] or by the
/// per-group [`Accumulator`]) into this sketch.
/// Merge a serialized state (see [`SerializedHll`]) into this sketch.
fn merge_serialized(&mut self, bytes: &[u8]) -> Result<isize> {
if bytes.is_empty() {
return Ok(0);
}
if bytes.len() == NUM_REGISTERS {
let other: HyperLogLog<u8> = bytes.try_into()?;
Ok(self.merge_dense(&other))
} else {
if !bytes.len().is_multiple_of(size_of::<u64>()) {
return internal_err!(
"approx_distinct: malformed sparse state: length {} is not a multiple of {}",
bytes.len(),
size_of::<u64>()
);
Ok(match SerializedHll::decode(bytes)? {
SerializedHll::Dense(registers) => {
self.merge_dense(&HyperLogLog::new_with_registers(*registers))
}
if bytes.len() > SPARSE_LIMIT * size_of::<u64>() {
return internal_err!(
"approx_distinct: malformed sparse state: length {} exceeds sparse limit {}",
bytes.len(),
SPARSE_LIMIT * size_of::<u64>()
);
}
let mut delta = 0;
for chunk in bytes.as_chunks::<{ size_of::<u64>() }>().0 {
delta += self.add_hash(u64::from_le_bytes(*chunk));
}
Ok(delta)
}
SerializedHll::Sparse(chunks) => chunks
.iter()
.map(|chunk| self.add_hash(u64::from_le_bytes(*chunk)))
.sum(),
})
}

/// Merge a dense sketch into this one, promoting to dense if necessary.
Expand Down
40 changes: 29 additions & 11 deletions datafusion/functions-aggregate/src/average.rs
Original file line number Diff line number Diff line change
Expand Up @@ -653,10 +653,14 @@ impl Accumulator for AvgAccumulator {
}

fn state(&mut self) -> Result<Vec<ScalarValue>> {
Ok(vec![
ScalarValue::from(self.count),
ScalarValue::Float64(self.sum),
])
// With no non-NULL values both fields are NULL, matching
// `AvgGroupsAccumulator::state`
let (count, sum) = if self.count == 0 {
(None, None)
} else {
(Some(self.count), self.sum)
};
Ok(vec![ScalarValue::UInt64(count), ScalarValue::Float64(sum)])
}

fn merge_batch(&mut self, states: &[ArrayRef]) -> Result<()> {
Expand Down Expand Up @@ -833,9 +837,16 @@ where
}

fn state(&mut self) -> Result<Vec<ScalarValue>> {
// With no non-NULL values both fields are NULL, matching
// `AvgGroupsAccumulator::state`
let (count, sum) = if self.count == 0 {
(None, None)
} else {
(Some(self.count), self.sum)
};
Ok(vec![
ScalarValue::from(self.count),
ScalarValue::new_primitive::<S>(self.sum, &self.sum_data_type)?,
ScalarValue::UInt64(count),
ScalarValue::new_primitive::<S>(sum, &self.sum_data_type)?,
])
}

Expand Down Expand Up @@ -916,14 +927,21 @@ impl Accumulator for DurationAvgAccumulator {
}

fn state(&mut self) -> Result<Vec<ScalarValue>> {
// With no non-NULL values both fields are NULL, matching
// `AvgGroupsAccumulator::state`
let (count, sum) = if self.count == 0 {
(None, None)
} else {
(Some(self.count), self.sum)
};
let duration_value = match self.time_unit {
TimeUnit::Second => ScalarValue::DurationSecond(self.sum),
TimeUnit::Millisecond => ScalarValue::DurationMillisecond(self.sum),
TimeUnit::Microsecond => ScalarValue::DurationMicrosecond(self.sum),
TimeUnit::Nanosecond => ScalarValue::DurationNanosecond(self.sum),
TimeUnit::Second => ScalarValue::DurationSecond(sum),
TimeUnit::Millisecond => ScalarValue::DurationMillisecond(sum),
TimeUnit::Microsecond => ScalarValue::DurationMicrosecond(sum),
TimeUnit::Nanosecond => ScalarValue::DurationNanosecond(sum),
};

Ok(vec![ScalarValue::from(self.count), duration_value])
Ok(vec![ScalarValue::UInt64(count), duration_value])
}

fn merge_batch(&mut self, states: &[ArrayRef]) -> Result<()> {
Expand Down
Loading
Loading