diff --git a/vortex-array/src/arrays/filter/mod.rs b/vortex-array/src/arrays/filter/mod.rs index 145aeb889bd..fb7dc6a4141 100644 --- a/vortex-array/src/arrays/filter/mod.rs +++ b/vortex-array/src/arrays/filter/mod.rs @@ -11,6 +11,7 @@ pub use vtable::FilterArray; mod execute; pub(crate) use execute::buffer::filter_buffer; +pub(crate) use execute::buffer::prepare_mask_for_reuse; pub(crate) use execute::filter_validity; mod kernel; diff --git a/vortex-array/src/arrays/scalar_fn/rules.rs b/vortex-array/src/arrays/scalar_fn/rules.rs index 52373d80a33..71f1a628eb3 100644 --- a/vortex-array/src/arrays/scalar_fn/rules.rs +++ b/vortex-array/src/arrays/scalar_fn/rules.rs @@ -14,6 +14,7 @@ use crate::arrays::ScalarFn; use crate::arrays::ScalarFnArray; use crate::arrays::Slice; use crate::arrays::StructArray; +use crate::arrays::filter::prepare_mask_for_reuse; use crate::arrays::scalar_fn::ScalarFnArrayExt; use crate::optimizer::rules::ArrayParentReduceRule; use crate::optimizer::rules::ArrayReduceRule; @@ -27,7 +28,7 @@ pub(super) const RULES: ReduceRuleSet = ReduceRuleSet::new(&[&ScalarFnPackToStructRule, &ScalarFnAbstractReduceRule]); pub(super) const PARENT_RULES: ParentRuleSet = ParentRuleSet::new(&[ - ParentRuleSet::lift(&ScalarFnUnaryFilterPushDownRule), + ParentRuleSet::lift(&ScalarFilterPushdownRule), ParentRuleSet::lift(&ScalarFnSliceReduceRule), ]); @@ -95,9 +96,9 @@ impl ArrayReduceRule for ScalarFnAbstractReduceRule { } #[derive(Debug)] -struct ScalarFnUnaryFilterPushDownRule; +struct ScalarFilterPushdownRule; -impl ArrayParentReduceRule for ScalarFnUnaryFilterPushDownRule { +impl ArrayParentReduceRule for ScalarFilterPushdownRule { type Parent = Filter; fn reduce_parent( @@ -106,48 +107,145 @@ impl ArrayParentReduceRule for ScalarFnUnaryFilterPushDownRule { parent: ArrayView<'_, Filter>, _child_idx: usize, ) -> VortexResult> { - // If we only have one non-constant child, then it is _always_ cheaper to push down the - // filter over the children of the scalar function array. - if child + let nchildren = child .iter_children() .filter(|c| !c.is::()) - .count() - == 1 + .count(); + if nchildren > 1 + && let Some(values) = parent.filter_mask().values() { - let new_children: Vec<_> = child - .iter_children() - .map(|c| match c.as_opt::() { - Some(array) => { - Ok(ConstantArray::new(array.scalar().clone(), parent.len()).into_array()) - } - None => c.filter(parent.filter_mask().clone()), - }) - .try_collect()?; - - let new_array = - ScalarFnArray::try_new(child.scalar_fn().clone(), new_children)?.into_array(); - - return Ok(Some(new_array)); + prepare_mask_for_reuse(values, nchildren); } - Ok(None) + let new_children: Vec<_> = child + .iter_children() + .map(|c| match c.as_opt::() { + Some(array) => { + Ok(ConstantArray::new(array.scalar().clone(), parent.len()).into_array()) + } + None => c.filter(parent.filter_mask().clone()), + }) + .try_collect()?; + + Ok(Some( + ScalarFnArray::try_new_with_len(child.scalar_fn().clone(), new_children, parent.len())? + .into_array(), + )) } } #[cfg(test)] mod tests { + use rstest::rstest; use vortex_error::VortexExpect; + use vortex_error::VortexResult; + use vortex_error::vortex_err; + use vortex_mask::Mask; + use super::ScalarFilterPushdownRule; + use crate::VortexSessionExecute; use crate::array::IntoArray; + use crate::array_session; use crate::arrays::ChunkedArray; + use crate::arrays::Constant; + use crate::arrays::Filter; + use crate::arrays::FilterArray; use crate::arrays::PrimitiveArray; + use crate::arrays::ScalarFn; + use crate::arrays::ScalarFnArray; + use crate::arrays::scalar_fn::ScalarFnArrayExt; use crate::arrays::scalar_fn::rules::ConstantArray; + use crate::assert_arrays_eq; use crate::dtype::DType; use crate::dtype::Nullability; use crate::dtype::PType; use crate::expr::cast; use crate::expr::is_null; use crate::expr::root; + use crate::optimizer::rules::ArrayParentReduceRule; + use crate::scalar::Scalar; + use crate::scalar_fn::TypedScalarFnInstance; + use crate::scalar_fn::fns::binary::Binary; + use crate::scalar_fn::fns::literal::Literal; + use crate::scalar_fn::fns::operators::Operator; + + #[rstest] + #[case(0, [14, 14, 14, 14, 14])] + #[case(1, [7, 27, 47, 67, 87])] + #[case(2, [0, 40, 80, 120, 160])] + fn test_filter_pushdown( + #[case] nchildren: usize, + #[case] expected: [i32; 5], + ) -> VortexResult<()> { + let children = (0..2) + .map(|index| { + if index < nchildren { + PrimitiveArray::from_iter(0i32..100).into_array() + } else { + ConstantArray::new(7i32, 100).into_array() + } + }) + .collect(); + let array = ScalarFnArray::try_new( + TypedScalarFnInstance::new(Binary, Operator::Add).erased(), + children, + )?; + let mask = Mask::from_iter((0..100).map(|index| index % 20 == 0)); + let values = mask + .values() + .ok_or_else(|| vortex_err!("expected mask values"))?; + assert!(values.cached_indices().is_none()); + let parent = FilterArray::try_new(array.clone().into_array(), mask.clone())?; + + let result = ScalarFilterPushdownRule + .reduce_parent(array.as_view(), parent.as_view(), 0)? + .ok_or_else(|| vortex_err!("expected filter pushdown"))?; + + assert_eq!(values.cached_indices().is_some(), nchildren > 1); + let scalar_fn = result.as_::(); + assert_eq!(scalar_fn.len(), 5); + for (index, child) in scalar_fn.iter_children().enumerate() { + assert_eq!(child.len(), 5); + if index < nchildren { + assert!(child.is::()); + } else { + assert!(child.is::()); + } + } + let expected = PrimitiveArray::from_iter(expected); + assert_arrays_eq!( + result, + expected, + &mut array_session().create_execution_ctx() + ); + Ok(()) + } + + #[test] + fn test_filter_pushdown_without_children() -> VortexResult<()> { + let array = ScalarFnArray::try_new_with_len( + TypedScalarFnInstance::new(Literal, Scalar::from(7i32)).erased(), + vec![], + 3, + )?; + let parent = FilterArray::try_new( + array.clone().into_array(), + Mask::from_iter([true, false, true]), + )?; + + let result = ScalarFilterPushdownRule + .reduce_parent(array.as_view(), parent.as_view(), 0)? + .ok_or_else(|| vortex_err!("expected filter pushdown"))?; + + assert_eq!(result.as_::().nchildren(), 0); + assert_eq!(result.len(), 2); + assert_arrays_eq!( + result, + ConstantArray::new(7i32, 2), + &mut array_session().create_execution_ctx() + ); + Ok(()) + } #[test] fn test_empty_constants() {