Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
14 changes: 14 additions & 0 deletions vortex-array/src/aggregate_fn/accumulator.rs
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,7 @@ use crate::ExecutionCtx;
use crate::aggregate_fn::AggregateDTypes;
use crate::aggregate_fn::AggregateFn;
use crate::aggregate_fn::AggregateFnRef;
use crate::aggregate_fn::AggregateFnSatisfaction;
use crate::aggregate_fn::AggregateFnVTable;
use crate::aggregate_fn::session::AggregateFnSessionExt;
use crate::columnar::AnyColumnar;
Expand Down Expand Up @@ -192,6 +193,10 @@ pub trait DynAccumulator: 'static + Send {
/// Whether the accumulator's result is fully determined.
fn is_saturated(&self) -> bool;

/// Whether the current result can satisfy `requested`, taking the accumulated state into
/// account. See [`AggregateFnVTable::partial_can_satisfy`].
fn can_satisfy(&self, requested: &AggregateFnRef) -> AggregateFnSatisfaction;

/// Reset the accumulator's state to the empty group.
fn reset(&mut self);

Expand Down Expand Up @@ -318,6 +323,15 @@ impl<V: AggregateFnVTable> DynAccumulator for Accumulator<V> {
})
}

fn can_satisfy(&self, requested: &AggregateFnRef) -> AggregateFnSatisfaction {
let options = self.aggregate_fn.as_::<V>();
// Without a partial nothing was accumulated yet, so only the options can be consulted.
match &self.partial {
Some(partial) => self.vtable.partial_can_satisfy(options, partial, requested),
None => self.vtable.can_satisfy(options, requested),
}
}

fn reset(&mut self) {
self.partial = None;
}
Expand Down
177 changes: 147 additions & 30 deletions vortex-array/src/aggregate_fn/fns/bounded_max/mod.rs
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
// SPDX-License-Identifier: Apache-2.0
// SPDX-FileCopyrightText: Copyright the Vortex contributors

use std::cmp::Ordering;
use std::fmt::Display;
use std::fmt::Formatter;
use std::num::NonZeroUsize;
Expand Down Expand Up @@ -33,7 +34,6 @@ use crate::dtype::DType;
use crate::dtype::FieldNames;
use crate::dtype::Nullability;
use crate::dtype::StructFields;
use crate::partial_ord::partial_max;
use crate::scalar::Scalar;
use crate::scalar::ScalarTruncation;
use crate::scalar::upper_bound;
Expand Down Expand Up @@ -65,27 +65,59 @@ pub struct BoundedMax;

enum BoundedMaxState {
Empty,
Value(Scalar),
/// The true maximum: it fit within the byte bound, so it was never truncated.
Exact(Scalar),
/// An upper bound of the true maximum, either because the maximum was truncated or because
/// it was combined from a serialized partial, which doesn't record exactness.
Inexact(Scalar),
Comment on lines +69 to +72

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

can we pull this out into a different pr

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Do you want to base this on a new exactness tracking PR, or do you want to drop this for now and always read back string min/max as approximate?

Unknown,
}

impl BoundedMaxState {
fn bound(bound: Scalar, exact: bool) -> Self {
if exact {
Self::Exact(bound)
} else {
Self::Inexact(bound)
}
}
}

/// Partial accumulator state for the bounded maximum aggregate.
pub struct BoundedMaxPartial {
state: BoundedMaxState,
}

impl BoundedMaxPartial {
fn merge_bound(&mut self, max: Scalar) {
if max.is_null() {
fn merge_bound(&mut self, bound: Scalar, exact: bool) {
if bound.is_null() {
return;
}

self.state = match std::mem::replace(&mut self.state, BoundedMaxState::Empty) {
BoundedMaxState::Empty => BoundedMaxState::Value(max),
BoundedMaxState::Value(current) => BoundedMaxState::Value(
partial_max(max, current).vortex_expect("incomparable bounded max scalars"),
),
BoundedMaxState::Unknown => BoundedMaxState::Unknown,
let (current, current_exact) =
match std::mem::replace(&mut self.state, BoundedMaxState::Empty) {
BoundedMaxState::Empty => {
self.state = BoundedMaxState::bound(bound, exact);
return;
}
BoundedMaxState::Unknown => {
self.state = BoundedMaxState::Unknown;
return;
}
BoundedMaxState::Exact(current) => (current, true),
BoundedMaxState::Inexact(current) => (current, false),
};

// The larger bound wins. It is the true maximum iff it was not truncated: every other
// value is at most its own bound, which is below the winner. On a tie, either side being
// exact means the true maximum reaches the bound.
self.state = match bound
.partial_cmp(&current)
.vortex_expect("incomparable bounded max scalars")
{
Ordering::Greater => BoundedMaxState::bound(bound, exact),
Ordering::Less => BoundedMaxState::bound(current, current_exact),
Ordering::Equal => BoundedMaxState::bound(current, exact || current_exact),
};
}

Expand All @@ -96,7 +128,7 @@ impl BoundedMaxPartial {
fn final_scalar(&self, args: AggregateArgs<'_, BoundedMaxOptions>) -> VortexResult<Scalar> {
let dtype = args.return_dtype.clone();
match &self.state {
BoundedMaxState::Value(max) => max.cast(&dtype),
BoundedMaxState::Exact(bound) | BoundedMaxState::Inexact(bound) => bound.cast(&dtype),
BoundedMaxState::Empty | BoundedMaxState::Unknown => Ok(Scalar::null(dtype)),
}
}
Expand Down Expand Up @@ -183,6 +215,27 @@ impl AggregateFnVTable for BoundedMax {
}
}

fn partial_can_satisfy(
&self,
options: &Self::Options,
partial: &Self::Partial,
requested: &AggregateFnRef,
) -> AggregateFnSatisfaction {
let satisfaction = self.can_satisfy(options, requested);
let untruncated = matches!(
partial.state,
BoundedMaxState::Empty | BoundedMaxState::Exact(_)
);
if satisfaction == AggregateFnSatisfaction::Approximate
&& untruncated
&& requested.is::<Max>()
{
AggregateFnSatisfaction::Exact
} else {
satisfaction
}
}

fn partial_dtype(&self, options: &Self::Options, input_dtype: &DType) -> Option<DType> {
supported_dtype(options, input_dtype).map(make_bounded_max_partial_dtype)
}
Expand Down Expand Up @@ -226,7 +279,8 @@ impl AggregateFnVTable for BoundedMax {
} else if bound.is_null() {
BoundedMaxState::Empty
} else {
BoundedMaxState::Value(bound)
// Serialized partials don't record exactness.
BoundedMaxState::Inexact(bound)
}
};
Ok(BoundedMaxPartial { state })
Expand All @@ -240,7 +294,8 @@ impl AggregateFnVTable for BoundedMax {
) -> VortexResult<Self::Partial> {
match second.state {
BoundedMaxState::Empty => {}
BoundedMaxState::Value(max) => first.merge_bound(max),
BoundedMaxState::Exact(bound) => first.merge_bound(bound, true),
BoundedMaxState::Inexact(bound) => first.merge_bound(bound, false),
BoundedMaxState::Unknown => first.unknown(),
}
Ok(first)
Expand All @@ -255,10 +310,10 @@ impl AggregateFnVTable for BoundedMax {
let bound_dtype = args.dtype.as_nullable();
match &partial.state {
BoundedMaxState::Empty => Ok(Scalar::null(dtype)),
BoundedMaxState::Value(max) => Ok(Scalar::struct_(
BoundedMaxState::Exact(bound) | BoundedMaxState::Inexact(bound) => Ok(Scalar::struct_(
dtype,
vec![
max.cast(&bound_dtype)?,
bound.cast(&bound_dtype)?,
Scalar::bool(false, Nullability::NonNullable),
],
)),
Expand Down Expand Up @@ -293,7 +348,7 @@ impl AggregateFnVTable for BoundedMax {
return Ok(());
};
match truncate_max(result.max, args.options.max_bytes.get())? {
Some(bound) => partial.merge_bound(bound),
Some((bound, truncated)) => partial.merge_bound(bound, !truncated),
None => partial.unknown(),
}
Ok(())
Expand Down Expand Up @@ -322,22 +377,21 @@ fn supported_dtype<'a>(_options: &BoundedMaxOptions, input_dtype: &'a DType) ->
.map(|_| input_dtype)
}

fn truncate_max(value: Scalar, max_bytes: usize) -> VortexResult<Option<Scalar>> {
/// Returns the upper bound and whether it was truncated, or `None` if no bound fits `max_bytes`.
fn truncate_max(value: Scalar, max_bytes: usize) -> VortexResult<Option<(Scalar, bool)>> {
let nullability = value.dtype().nullability();
match value.dtype() {
DType::Utf8(_) => {
Ok(
upper_bound(BufferString::from_scalar(value)?, max_bytes, nullability)
.map(|(bound, _)| bound),
)
}
DType::Binary(_) => {
Ok(
upper_bound(ByteBuffer::from_scalar(value)?, max_bytes, nullability)
.map(|(bound, _)| bound),
)
}
_ => Ok(Some(value)),
DType::Utf8(_) => Ok(upper_bound(
BufferString::from_scalar(value)?,
max_bytes,
nullability,
)),
DType::Binary(_) => Ok(upper_bound(
ByteBuffer::from_scalar(value)?,
max_bytes,
nullability,
)),
_ => Ok(Some((value, false))),
}
}

Expand Down Expand Up @@ -584,4 +638,67 @@ mod tests {
assert_eq!(roundtrip, options);
Ok(())
}

#[test]
fn bounded_max_exactly_satisfies_max_only_when_untruncated() -> VortexResult<()> {
let mut ctx = array_session().create_execution_ctx();
let requested = Max.bind(NumericalAggregateOpts::skip_nans());
let new_acc = || {
Accumulator::try_new(
BoundedMax,
BoundedMaxOptions {
max_bytes: max_bytes(4),
},
VarBinViewArray::from_iter_str(["x"]).dtype().clone(),
)
};

// Every value fits the bound, so the bound is the true extremum.
let mut acc = new_acc()?;
acc.accumulate(
&VarBinViewArray::from_iter_str(["b", "abc"]).into_array(),
&mut ctx,
)?;
assert_eq!(acc.can_satisfy(&requested), AggregateFnSatisfaction::Exact);

// A truncated value that isn't the extremum doesn't affect exactness.
acc.accumulate(
&VarBinViewArray::from_iter_str(["aaaaa"]).into_array(),
&mut ctx,
)?;
assert_eq!(acc.can_satisfy(&requested), AggregateFnSatisfaction::Exact);

// A truncated extremum makes the bound approximate.
acc.accumulate(
&VarBinViewArray::from_iter_str(["bbbbb"]).into_array(),
&mut ctx,
)?;
assert_eq!(
acc.can_satisfy(&requested),
AggregateFnSatisfaction::Approximate
);

// Exactness isn't serialized, so a winning partial combined from its scalar form is
// approximate even though it was exact where it was computed.
let mut source = new_acc()?;
source.accumulate(
&VarBinViewArray::from_iter_str(["b"]).into_array(),
&mut ctx,
)?;
assert_eq!(
source.can_satisfy(&requested),
AggregateFnSatisfaction::Exact
);
let mut acc = new_acc()?;
acc.accumulate(
&VarBinViewArray::from_iter_str(["a"]).into_array(),
&mut ctx,
)?;
acc.combine_partials(&source.partial_scalar()?)?;
assert_eq!(
acc.can_satisfy(&requested),
AggregateFnSatisfaction::Approximate
);
Ok(())
}
}
Loading
Loading