From 1e3454437ebc04046bba26eac76a97e9e064b53e Mon Sep 17 00:00:00 2001 From: Will Manning Date: Thu, 27 Aug 2026 17:17:26 -0400 Subject: [PATCH 1/8] perf(fastlanes): Evaluate constant list membership Signed-off-by: Will Manning --- encodings/fastlanes/Cargo.toml | 5 + .../benches/bitpacking_list_contains.rs | 214 ++++++++++++ .../bitpacking/compute/list_contains/mod.rs | 148 +++++++++ .../bitpacking/compute/list_contains/tests.rs | 309 ++++++++++++++++++ .../fastlanes/src/bitpacking/compute/mod.rs | 1 + encodings/fastlanes/src/bitpacking/mod.rs | 4 + .../src/bitpacking/vtable/kernels.rs | 7 + 7 files changed, 688 insertions(+) create mode 100644 encodings/fastlanes/benches/bitpacking_list_contains.rs create mode 100644 encodings/fastlanes/src/bitpacking/compute/list_contains/mod.rs create mode 100644 encodings/fastlanes/src/bitpacking/compute/list_contains/tests.rs diff --git a/encodings/fastlanes/Cargo.toml b/encodings/fastlanes/Cargo.toml index 9085390b67b..c8efeebc43d 100644 --- a/encodings/fastlanes/Cargo.toml +++ b/encodings/fastlanes/Cargo.toml @@ -48,6 +48,11 @@ _test-harness = ["dep:rand"] name = "bitpacking_take" harness = false +[[bench]] +name = "bitpacking_list_contains" +harness = false +required-features = ["_test-harness"] + [[bench]] name = "canonicalize_bench" harness = false diff --git a/encodings/fastlanes/benches/bitpacking_list_contains.rs b/encodings/fastlanes/benches/bitpacking_list_contains.rs new file mode 100644 index 00000000000..127d3c28d6e --- /dev/null +++ b/encodings/fastlanes/benches/bitpacking_list_contains.rs @@ -0,0 +1,214 @@ +// SPDX-License-Identifier: Apache-2.0 +// SPDX-FileCopyrightText: Copyright the Vortex contributors + +//! Measures the linear-search and binary-search crossover for constant-list membership. +//! +//! The strategy benchmarks isolate lookup cost across mostly-missing and mixed probes. The kernel +//! benchmark includes scalar extraction, sorting, FastLanes decoding, and result construction. +//! The isolated lookup benchmark excludes sorting, so it favors binary search. Use the forced +//! full-kernel results to select the production strategy. +//! +//! Run with `cargo bench -p vortex-fastlanes --bench bitpacking_list_contains`. + +#![expect(clippy::cast_possible_truncation)] +#![expect(clippy::unwrap_used)] + +use std::hint::black_box; +use std::sync::Arc; + +use divan::Bencher; +use divan::counter::ItemsCount; +use vortex_array::ArrayRef; +use vortex_array::IntoArray; +use vortex_array::VortexSessionExecute; +use vortex_array::array_session; +use vortex_array::arrays::ConstantArray; +use vortex_array::arrays::PrimitiveArray; +use vortex_array::dtype::DType; +use vortex_array::dtype::Nullability; +use vortex_array::dtype::PType; +use vortex_array::scalar::Scalar; +use vortex_array::scalar_fn::fns::list_contains::ListContainsElementKernel; +use vortex_array::validity::Validity; +use vortex_buffer::Alignment; +use vortex_buffer::BufferMut; +use vortex_fastlanes::BitPacked; +use vortex_fastlanes::BitPackedArray; +use vortex_fastlanes::BitPackedData; +use vortex_fastlanes::list_contains_test_harness; +use vortex_fastlanes::list_contains_test_harness::MembershipSearch; + +const LEN: usize = 64 * 1_024; +const MEMBER_COUNTS: &[usize] = &[1, 2, 3, 4, 6, 8, 9, 12, 16, 24, 32, 64]; +const STRATEGY_MEMBER_COUNTS: &[usize] = &[3, 4, 6, 8, 9, 12, 16, 24, 32, 64]; + +fn main() { + divan::main(); +} + +fn members(count: usize) -> Vec { + (0..count).map(|index| index as u32 * 2).collect() +} + +fn mostly_missing_values() -> Vec { + (0..LEN) + .map(|index| ((index as u32 * 17) % 4_096) | 1) + .collect() +} + +fn mixed_values(members: &[u32]) -> Vec { + (0..LEN) + .map(|index| { + if index.is_multiple_of(2) { + members[(index / 2) % members.len()] + } else { + ((index as u32 * 17) % 4_096) | 1 + } + }) + .collect() +} + +fn count_linear(values: &[u32], members: &[u32]) -> usize { + values + .iter() + .filter(|value| members.contains(black_box(value))) + .count() +} + +fn count_binary(values: &[u32], members: &[u32]) -> usize { + values + .iter() + .filter(|value| members.binary_search(black_box(value)).is_ok()) + .count() +} + +#[divan::bench(args = MEMBER_COUNTS)] +fn linear_mostly_missing(bencher: Bencher, member_count: usize) { + let values = mostly_missing_values(); + let members = members(member_count); + bencher + .counter(ItemsCount::new(LEN)) + .bench_local(|| black_box(count_linear(&values, &members))); +} + +#[divan::bench(args = MEMBER_COUNTS)] +fn binary_mostly_missing(bencher: Bencher, member_count: usize) { + let values = mostly_missing_values(); + let members = members(member_count); + bencher + .counter(ItemsCount::new(LEN)) + .bench_local(|| black_box(count_binary(&values, &members))); +} + +#[divan::bench(args = MEMBER_COUNTS)] +fn linear_mixed(bencher: Bencher, member_count: usize) { + let members = members(member_count); + let values = mixed_values(&members); + bencher + .counter(ItemsCount::new(LEN)) + .bench_local(|| black_box(count_linear(&values, &members))); +} + +#[divan::bench(args = MEMBER_COUNTS)] +fn binary_mixed(bencher: Bencher, member_count: usize) { + let members = members(member_count); + let values = mixed_values(&members); + bencher + .counter(ItemsCount::new(LEN)) + .bench_local(|| black_box(count_binary(&values, &members))); +} + +fn page_aligned(array: BitPackedArray) -> BitPackedArray { + let ptype = array.dtype().as_ptype(); + let parts = BitPacked::into_parts(array); + BitPacked::try_new( + parts.packed.ensure_aligned(Alignment::new(4_096)).unwrap(), + ptype, + parts.validity, + parts.patches, + parts.bit_width, + parts.len, + parts.offset, + ) + .unwrap() +} + +fn kernel_inputs(member_count: usize) -> (BitPackedArray, ArrayRef) { + let mut ctx = array_session().create_execution_ctx(); + let values: BufferMut = (0..LEN).map(|index| (index as u32 * 17) % 1_024).collect(); + let packed = page_aligned( + BitPackedData::encode( + &PrimitiveArray::new(values.freeze(), Validity::NonNullable).into_array(), + 10, + &mut ctx, + ) + .unwrap(), + ); + let member_scalars = members(member_count) + .into_iter() + .map(|value| Scalar::primitive(value, Nullability::NonNullable)) + .collect(); + let list = ConstantArray::new( + Scalar::list( + Arc::new(DType::Primitive(PType::U32, Nullability::NonNullable)), + member_scalars, + Nullability::NonNullable, + ), + LEN, + ) + .into_array(); + (packed, list) +} + +#[divan::bench(args = MEMBER_COUNTS)] +fn bitpacked_kernel(bencher: Bencher, member_count: usize) { + let (packed, list) = kernel_inputs(member_count); + let mut ctx = array_session().create_execution_ctx(); + bencher.counter(ItemsCount::new(LEN)).bench_local(|| { + black_box( + ::list_contains( + &list, + packed.as_view(), + &mut ctx, + ) + .unwrap() + .unwrap(), + ) + }); +} + +#[divan::bench(args = STRATEGY_MEMBER_COUNTS)] +fn bitpacked_linear(bencher: Bencher, member_count: usize) { + let (packed, list) = kernel_inputs(member_count); + let mut ctx = array_session().create_execution_ctx(); + bencher.counter(ItemsCount::new(LEN)).bench_local(|| { + black_box( + list_contains_test_harness::list_contains( + &list, + packed.as_view(), + MembershipSearch::Linear, + &mut ctx, + ) + .unwrap() + .unwrap(), + ) + }); +} + +#[divan::bench(args = STRATEGY_MEMBER_COUNTS)] +fn bitpacked_binary(bencher: Bencher, member_count: usize) { + let (packed, list) = kernel_inputs(member_count); + let mut ctx = array_session().create_execution_ctx(); + bencher.counter(ItemsCount::new(LEN)).bench_local(|| { + black_box( + list_contains_test_harness::list_contains( + &list, + packed.as_view(), + MembershipSearch::Binary, + &mut ctx, + ) + .unwrap() + .unwrap(), + ) + }); +} diff --git a/encodings/fastlanes/src/bitpacking/compute/list_contains/mod.rs b/encodings/fastlanes/src/bitpacking/compute/list_contains/mod.rs new file mode 100644 index 00000000000..f93c09af30f --- /dev/null +++ b/encodings/fastlanes/src/bitpacking/compute/list_contains/mod.rs @@ -0,0 +1,148 @@ +// SPDX-License-Identifier: Apache-2.0 +// SPDX-FileCopyrightText: Copyright the Vortex contributors + +use vortex_array::ArrayRef; +use vortex_array::ArrayView; +use vortex_array::ExecutionCtx; +use vortex_array::IntoArray; +use vortex_array::arrays::ConstantArray; +use vortex_array::dtype::DType; +use vortex_array::dtype::NativePType; +use vortex_array::match_each_integer_ptype; +use vortex_array::scalar::Scalar; +use vortex_array::scalar_fn::fns::list_contains::ListContainsElementKernel; +use vortex_error::VortexResult; +use vortex_error::vortex_err; + +use super::compare_fused::stream_compare_fused; +use crate::BitPacked; + +#[derive(Clone, Copy)] +enum SearchStrategy { + Linear, + Binary, +} + +impl ListContainsElementKernel for BitPacked { + fn list_contains( + list: &ArrayRef, + element: ArrayView<'_, Self>, + ctx: &mut ExecutionCtx, + ) -> VortexResult> { + list_contains_with_strategy(list, element, SearchStrategy::Binary, ctx) + } +} + +fn list_contains_with_strategy( + list: &ArrayRef, + element: ArrayView<'_, BitPacked>, + strategy: SearchStrategy, + ctx: &mut ExecutionCtx, +) -> VortexResult> { + let Some(list_scalar) = list.as_constant() else { + return Ok(None); + }; + let DType::List(member_dtype, _) = list.dtype() else { + return Ok(None); + }; + if !member_dtype.eq_ignore_nullability(element.dtype()) { + return Ok(None); + } + + let nullability = list.dtype().nullability() | element.dtype().nullability(); + let Some(elements) = list_scalar.as_list().elements() else { + return Ok(Some( + ConstantArray::new(Scalar::null(DType::Bool(nullability)), element.len()).into_array(), + )); + }; + + let result = match_each_integer_ptype!(element.dtype().as_ptype(), |T| { + let mut members = elements + .iter() + .map(|value| { + value + .as_primitive_opt() + .ok_or_else(|| vortex_err!("List member is not a primitive scalar"))? + .try_typed_value::() + }) + .collect::>>>()? + .into_iter() + .flatten() + .collect::>(); + + match members.as_slice() { + [] => ConstantArray::new(Scalar::bool(false, nullability), element.len()).into_array(), + [member] => { + let member = *member; + stream_compare_fused::(element, member, nullability, NativePType::is_eq, ctx)? + } + [first, second] => { + let (first, second) = (*first, *second); + stream_compare_fused::( + element, + first, + nullability, + move |value, _| value.is_eq(first) | value.is_eq(second), + ctx, + )? + } + _ if matches!(strategy, SearchStrategy::Linear) => stream_compare_fused::( + element, + members[0], + nullability, + |value, _| members.contains(&value), + ctx, + )?, + _ => { + members.sort_unstable(); + members.dedup(); + stream_compare_fused::( + element, + members[0], + nullability, + |value, _| members.binary_search(&value).is_ok(), + ctx, + )? + } + } + }); + Ok(Some(result)) +} + +#[cfg(feature = "_test-harness")] +pub mod test_harness { + use vortex_array::ArrayRef; + use vortex_array::ArrayView; + use vortex_array::ExecutionCtx; + use vortex_error::VortexResult; + + use super::SearchStrategy; + use super::list_contains_with_strategy; + use crate::BitPacked; + + /// Selects the membership lookup strategy for a benchmark invocation. + #[derive(Clone, Copy)] + pub enum MembershipSearch { + /// Scans list members in order. + Linear, + /// Sorts list members and uses binary search. + Binary, + } + + /// Executes the BitPacked membership kernel with a fixed lookup strategy. + pub fn list_contains( + list: &ArrayRef, + element: ArrayView<'_, BitPacked>, + strategy: MembershipSearch, + ctx: &mut ExecutionCtx, + ) -> VortexResult> { + let strategy = match strategy { + MembershipSearch::Linear => SearchStrategy::Linear, + MembershipSearch::Binary => SearchStrategy::Binary, + }; + list_contains_with_strategy(list, element, strategy, ctx) + } +} + +#[cfg(test)] +mod tests; diff --git a/encodings/fastlanes/src/bitpacking/compute/list_contains/tests.rs b/encodings/fastlanes/src/bitpacking/compute/list_contains/tests.rs new file mode 100644 index 00000000000..d237d583e4e --- /dev/null +++ b/encodings/fastlanes/src/bitpacking/compute/list_contains/tests.rs @@ -0,0 +1,309 @@ +// SPDX-License-Identifier: Apache-2.0 +// SPDX-FileCopyrightText: Copyright the Vortex contributors + +use std::sync::Arc; +use std::sync::LazyLock; + +use rstest::rstest; +use vortex_array::ArrayRef; +use vortex_array::IntoArray; +use vortex_array::VortexSessionExecute; +use vortex_array::arrays::BoolArray; +use vortex_array::arrays::ConstantArray; +use vortex_array::arrays::PrimitiveArray; +use vortex_array::arrays::slice::SliceKernel; +use vortex_array::assert_arrays_eq; +use vortex_array::dtype::DType; +use vortex_array::dtype::NativePType; +use vortex_array::dtype::Nullability; +use vortex_array::expr::list_contains; +use vortex_array::expr::lit; +use vortex_array::expr::root; +use vortex_array::scalar::PValue; +use vortex_array::scalar::Scalar; +use vortex_array::scalar_fn::fns::list_contains::ListContainsElementKernel; +use vortex_array::test_harness::trace::TraceOptions; +use vortex_array::test_harness::trace::TraceResolution; +use vortex_array::test_harness::trace::trace_op_with; +use vortex_array::validity::Validity; +use vortex_error::VortexResult; +use vortex_error::vortex_err; +use vortex_session::VortexSession; + +use crate::BitPacked; +use crate::BitPackedArray; +use crate::BitPackedArrayExt; +use crate::BitPackedData; + +static SESSION: LazyLock = LazyLock::new(|| { + let session = vortex_array::array_session(); + crate::initialize(&session); + session +}); + +fn member_list( + values: impl IntoIterator>, + member_nullability: Nullability, +) -> Scalar +where + T: NativePType + Into, +{ + let member_dtype = DType::Primitive(T::PTYPE, member_nullability); + let members = values + .into_iter() + .map(|value| { + value + .map(|value| Scalar::primitive(value, member_nullability)) + .unwrap_or_else(|| Scalar::null(member_dtype.clone())) + }) + .collect(); + Scalar::list(Arc::new(member_dtype), members, Nullability::NonNullable) +} + +fn list_array(list: Scalar, len: usize) -> ArrayRef { + ConstantArray::new(list, len).into_array() +} + +fn execute_direct( + list: &ArrayRef, + element: &BitPackedArray, + ctx: &mut vortex_array::ExecutionCtx, +) -> VortexResult { + ::list_contains(list, element.as_view(), ctx)? + .ok_or_else(|| vortex_err!("BitPacked list_contains kernel declined a supported input"))? + .execute::(ctx) +} + +macro_rules! integer_type_test { + ($name:ident, $T:ty, $bit_width:expr) => { + #[test] + fn $name() -> VortexResult<()> { + let mut ctx = SESSION.create_execution_ctx(); + let values = (0..2_048) + .map(|value| (value % 64) as $T) + .collect::>(); + let members = [1 as $T, 3 as $T, 63 as $T]; + let primitive = PrimitiveArray::from_iter(values.iter().copied()); + let packed = BitPackedData::encode(&primitive.into_array(), $bit_width, &mut ctx)?; + let list = list_array( + member_list(members.into_iter().map(Some), Nullability::NonNullable), + packed.len(), + ); + + let actual = execute_direct(&list, &packed, &mut ctx)?; + let expected = + BoolArray::from_iter(values.into_iter().map(|value| members.contains(&value))); + assert_arrays_eq!(actual, expected, &mut ctx); + Ok(()) + } + }; +} + +integer_type_test!(test_integer_type_u8, u8, 6); +integer_type_test!(test_integer_type_u16, u16, 6); +integer_type_test!(test_integer_type_u32, u32, 6); +integer_type_test!(test_integer_type_u64, u64, 6); +integer_type_test!(test_integer_type_i8, i8, 6); +integer_type_test!(test_integer_type_i16, i16, 6); +integer_type_test!(test_integer_type_i32, i32, 6); +integer_type_test!(test_integer_type_i64, i64, 6); + +#[rstest] +#[case::empty(vec![])] +#[case::one(vec![3])] +#[case::two(vec![3, 7])] +#[case::three(vec![3, 7, 11])] +#[case::seven((0..7).map(|value| value * 2 + 1).collect())] +#[case::eight((0..8).map(|value| value * 2 + 1).collect())] +#[case::nine((0..9).map(|value| value * 2 + 1).collect())] +#[case::larger((0..32).map(|value| value * 3).collect())] +#[case::duplicates(vec![3, 3, 7, 7, 11, 11, 15, 15, 15])] +fn test_member_cardinalities(#[case] members: Vec) -> VortexResult<()> { + let mut ctx = SESSION.create_execution_ctx(); + let values = (0..2_048).map(|value| value % 128).collect::>(); + let primitive = PrimitiveArray::from_iter(values.iter().copied()); + let packed = BitPackedData::encode(&primitive.into_array(), 7, &mut ctx)?; + let list = list_array( + member_list(members.iter().copied().map(Some), Nullability::NonNullable), + packed.len(), + ); + + let actual = execute_direct(&list, &packed, &mut ctx)?; + let expected = BoolArray::from_iter(values.into_iter().map(|value| members.contains(&value))); + assert_arrays_eq!(actual, expected, &mut ctx); + Ok(()) +} + +#[test] +fn test_patches() -> VortexResult<()> { + let mut ctx = SESSION.create_execution_ctx(); + let values = (0..2_048) + .map(|index| { + if index % 97 == 0 { + 100_000 + index + } else { + index % 100 + } + }) + .collect::>(); + let primitive = PrimitiveArray::from_iter(values.iter().copied()); + let packed = BitPackedData::encode(&primitive.into_array(), 7, &mut ctx)?; + assert!(packed.patches().is_some(), "test setup requires patches"); + let members = [3, 100_097]; + let list = list_array( + member_list(members.into_iter().map(Some), Nullability::NonNullable), + packed.len(), + ); + + let actual = execute_direct(&list, &packed, &mut ctx)?; + let expected = BoolArray::from_iter(values.into_iter().map(|value| members.contains(&value))); + assert_arrays_eq!(actual, expected, &mut ctx); + Ok(()) +} + +#[test] +fn test_sliced_array() -> VortexResult<()> { + let mut ctx = SESSION.create_execution_ctx(); + let values = (0..5_000).map(|value| value % 128).collect::>(); + let primitive = PrimitiveArray::from_iter(values.iter().copied()); + let packed = BitPackedData::encode(&primitive.into_array(), 7, &mut ctx)?; + let range = 333..4_333; + let sliced = ::slice(packed.as_view(), range.clone(), &mut ctx)? + .ok_or_else(|| vortex_err!("BitPacked slice kernel declined a supported input"))?; + let members = [1, 63, 127]; + let list = list_array( + member_list(members.into_iter().map(Some), Nullability::NonNullable), + sliced.len(), + ); + + let actual = ::list_contains( + &list, + sliced.as_::(), + &mut ctx, + )? + .ok_or_else(|| vortex_err!("BitPacked list_contains kernel declined a sliced input"))? + .execute::(&mut ctx)?; + let expected = BoolArray::from_iter(values[range].iter().map(|value| members.contains(value))); + assert_arrays_eq!(actual, expected, &mut ctx); + Ok(()) +} + +#[test] +fn test_null_needles() -> VortexResult<()> { + let mut ctx = SESSION.create_execution_ctx(); + let values = [Some(1i32), None, Some(2), Some(3), None]; + let primitive = PrimitiveArray::from_option_iter(values); + let packed = BitPackedData::encode(&primitive.into_array(), 2, &mut ctx)?; + let list = list_array( + member_list([Some(1), Some(3)], Nullability::NonNullable), + packed.len(), + ); + + let actual = execute_direct(&list, &packed, &mut ctx)?; + let expected = BoolArray::from_iter([Some(true), None, Some(false), Some(true), None]); + assert_arrays_eq!(actual, expected, &mut ctx); + Ok(()) +} + +#[test] +fn test_null_list() -> VortexResult<()> { + let mut ctx = SESSION.create_execution_ctx(); + let primitive = PrimitiveArray::from_iter([1i32, 2, 3]); + let packed = BitPackedData::encode(&primitive.into_array(), 2, &mut ctx)?; + let list_dtype = DType::List( + Arc::new(DType::Primitive(i32::PTYPE, Nullability::NonNullable)), + Nullability::Nullable, + ); + let list = list_array(Scalar::null(list_dtype), packed.len()); + + let actual = execute_direct(&list, &packed, &mut ctx)?; + let expected = BoolArray::new( + [false, false, false].into_iter().collect(), + Validity::AllInvalid, + ); + assert_arrays_eq!(actual, expected, &mut ctx); + Ok(()) +} + +#[test] +fn test_nullable_members_are_ignored() -> VortexResult<()> { + let mut ctx = SESSION.create_execution_ctx(); + let primitive = PrimitiveArray::from_iter([1i32, 2, 3, 4]); + let packed = BitPackedData::encode(&primitive.into_array(), 3, &mut ctx)?; + let list = list_array( + member_list([Some(1), None, Some(3)], Nullability::Nullable), + packed.len(), + ); + + let actual = execute_direct(&list, &packed, &mut ctx)?; + let expected = BoolArray::from_iter([true, false, true, false]); + assert_arrays_eq!(actual, expected, &mut ctx); + Ok(()) +} + +#[test] +fn test_wrong_integer_type_declines_without_panic() -> VortexResult<()> { + let mut ctx = SESSION.create_execution_ctx(); + let primitive = PrimitiveArray::from_iter([1i32, 2, 3]); + let packed = BitPackedData::encode(&primitive.into_array(), 2, &mut ctx)?; + let list = list_array( + member_list([Some(1i64), Some(3)], Nullability::NonNullable), + packed.len(), + ); + + let result = + ::list_contains(&list, packed.as_view(), &mut ctx)?; + assert!(result.is_none()); + Ok(()) +} + +#[test] +fn test_noninteger_list_declines_without_panic() -> VortexResult<()> { + let mut ctx = SESSION.create_execution_ctx(); + let primitive = PrimitiveArray::from_iter([1i32, 2, 3]); + let packed = BitPackedData::encode(&primitive.into_array(), 2, &mut ctx)?; + let list = list_array( + Scalar::list( + Arc::new(DType::Utf8(Nullability::NonNullable)), + vec![Scalar::utf8("one", Nullability::NonNullable)], + Nullability::NonNullable, + ), + packed.len(), + ); + + let result = + ::list_contains(&list, packed.as_view(), &mut ctx)?; + assert!(result.is_none()); + Ok(()) +} + +#[test] +fn test_registered_kernel_executes_through_expression() -> VortexResult<()> { + let mut ctx = SESSION.create_execution_ctx(); + let values = (0..2_048).map(|value| value % 128).collect::>(); + let primitive = PrimitiveArray::from_iter(values.iter().copied()); + let packed = BitPackedData::encode(&primitive.into_array(), 7, &mut ctx)?; + let members = [0, 99]; + let expression = list_contains( + lit(member_list( + members.into_iter().map(Some), + Nullability::NonNullable, + )), + root(), + ); + let contains = packed.into_array().apply(&expression)?; + + let traced = trace_op_with( + TraceOptions { + resolution: TraceResolution::Attempts, + }, + || contains.execute::(&mut ctx), + )?; + let trace = traced.trace.to_string(); + assert!(trace.contains("parent=vortex.list.contains"), "{trace}"); + assert!(trace.contains("source=session"), "{trace}"); + + let expected = BoolArray::from_iter(values.into_iter().map(|value| members.contains(&value))); + assert_arrays_eq!(traced.output, expected, &mut ctx); + Ok(()) +} diff --git a/encodings/fastlanes/src/bitpacking/compute/mod.rs b/encodings/fastlanes/src/bitpacking/compute/mod.rs index 38f86f781bb..f5986711d73 100644 --- a/encodings/fastlanes/src/bitpacking/compute/mod.rs +++ b/encodings/fastlanes/src/bitpacking/compute/mod.rs @@ -7,6 +7,7 @@ mod compare; mod compare_fused; mod filter; pub(crate) mod is_constant; +pub(crate) mod list_contains; mod slice; mod stream_predicate; mod take; diff --git a/encodings/fastlanes/src/bitpacking/mod.rs b/encodings/fastlanes/src/bitpacking/mod.rs index efa0677a91e..d27efc9dc92 100644 --- a/encodings/fastlanes/src/bitpacking/mod.rs +++ b/encodings/fastlanes/src/bitpacking/mod.rs @@ -13,6 +13,10 @@ pub use array::unpack_iter; pub(crate) mod compute; +#[cfg(feature = "_test-harness")] +#[doc(hidden)] +pub use compute::list_contains::test_harness as list_contains_test_harness; + mod plugin; mod vtable; diff --git a/encodings/fastlanes/src/bitpacking/vtable/kernels.rs b/encodings/fastlanes/src/bitpacking/vtable/kernels.rs index eb0dd9b7a23..9a0add2130b 100644 --- a/encodings/fastlanes/src/bitpacking/vtable/kernels.rs +++ b/encodings/fastlanes/src/bitpacking/vtable/kernels.rs @@ -16,6 +16,8 @@ use vortex_array::scalar_fn::fns::binary::Binary; use vortex_array::scalar_fn::fns::binary::CompareExecuteAdaptor; use vortex_array::scalar_fn::fns::cast::Cast; use vortex_array::scalar_fn::fns::cast::CastExecuteAdaptor; +use vortex_array::scalar_fn::fns::list_contains::ListContains; +use vortex_array::scalar_fn::fns::list_contains::ListContainsElementExecuteAdaptor; use vortex_session::VortexSession; use crate::BitPacked; @@ -36,4 +38,9 @@ pub(crate) fn initialize(session: &VortexSession) { kernels.register_execute_parent_kernel(Filter.id(), BitPacked, FilterExecuteAdaptor(BitPacked)); kernels.register_execute_parent_kernel(Slice.id(), BitPacked, SliceExecuteAdaptor(BitPacked)); kernels.register_execute_parent_kernel(Dict.id(), BitPacked, TakeExecuteAdaptor(BitPacked)); + kernels.register_execute_parent_kernel( + ListContains.id(), + BitPacked, + ListContainsElementExecuteAdaptor(BitPacked), + ); } From 4c49d32f358c547fb776850d70a3543dafc21285 Mon Sep 17 00:00:00 2001 From: Will Manning Date: Fri, 28 Aug 2026 18:16:24 -0400 Subject: [PATCH 2/8] perf(fastlanes): Refine constant list membership Signed-off-by: Will Manning --- encodings/fastlanes/Cargo.toml | 1 - .../benches/bitpacking_list_contains.rs | 247 ++++++++---------- .../bitpacking/compute/list_contains/mod.rs | 110 ++++---- .../bitpacking/compute/list_contains/tests.rs | 37 ++- encodings/fastlanes/src/bitpacking/mod.rs | 4 - .../arrays/primitive/compute/list_contains.rs | 155 +++++++++++ .../src/arrays/primitive/compute/mod.rs | 1 + .../src/arrays/primitive/vtable/kernel.rs | 7 + .../fns/list_contains/integer_membership.rs | 184 +++++++++++++ .../src/scalar_fn/fns/list_contains/mod.rs | 97 ++++++- 10 files changed, 628 insertions(+), 215 deletions(-) create mode 100644 vortex-array/src/arrays/primitive/compute/list_contains.rs create mode 100644 vortex-array/src/scalar_fn/fns/list_contains/integer_membership.rs diff --git a/encodings/fastlanes/Cargo.toml b/encodings/fastlanes/Cargo.toml index c8efeebc43d..1127bf37bed 100644 --- a/encodings/fastlanes/Cargo.toml +++ b/encodings/fastlanes/Cargo.toml @@ -51,7 +51,6 @@ harness = false [[bench]] name = "bitpacking_list_contains" harness = false -required-features = ["_test-harness"] [[bench]] name = "canonicalize_bench" diff --git a/encodings/fastlanes/benches/bitpacking_list_contains.rs b/encodings/fastlanes/benches/bitpacking_list_contains.rs index 127d3c28d6e..7479555a321 100644 --- a/encodings/fastlanes/benches/bitpacking_list_contains.rs +++ b/encodings/fastlanes/benches/bitpacking_list_contains.rs @@ -1,12 +1,11 @@ // SPDX-License-Identifier: Apache-2.0 // SPDX-FileCopyrightText: Copyright the Vortex contributors -//! Measures the linear-search and binary-search crossover for constant-list membership. +//! Compares compressed list membership with the canonical fallback. //! -//! The strategy benchmarks isolate lookup cost across mostly-missing and mixed probes. The kernel -//! benchmark includes scalar extraction, sorting, FastLanes decoding, and result construction. -//! The isolated lookup benchmark excludes sorting, so it favors binary search. Use the forced -//! full-kernel results to select the production strategy. +//! The specialized session evaluates membership while it decodes FastLanes lanes. The fallback +//! session decodes the complete array before the generic membership operation. +//! Density cases stress the 4 KiB lookup-table boundary. Sparse cases exceed that boundary. //! //! Run with `cargo bench -p vortex-fastlanes --bench bitpacking_list_contains`. @@ -22,100 +21,57 @@ use vortex_array::ArrayRef; use vortex_array::IntoArray; use vortex_array::VortexSessionExecute; use vortex_array::array_session; -use vortex_array::arrays::ConstantArray; +use vortex_array::arrays::BoolArray; use vortex_array::arrays::PrimitiveArray; use vortex_array::dtype::DType; use vortex_array::dtype::Nullability; use vortex_array::dtype::PType; +use vortex_array::expr::list_contains; +use vortex_array::expr::lit; +use vortex_array::expr::root; use vortex_array::scalar::Scalar; -use vortex_array::scalar_fn::fns::list_contains::ListContainsElementKernel; +use vortex_array::session::ArraySessionExt; use vortex_array::validity::Validity; use vortex_buffer::Alignment; use vortex_buffer::BufferMut; use vortex_fastlanes::BitPacked; use vortex_fastlanes::BitPackedArray; use vortex_fastlanes::BitPackedData; -use vortex_fastlanes::list_contains_test_harness; -use vortex_fastlanes::list_contains_test_harness::MembershipSearch; - -const LEN: usize = 64 * 1_024; -const MEMBER_COUNTS: &[usize] = &[1, 2, 3, 4, 6, 8, 9, 12, 16, 24, 32, 64]; -const STRATEGY_MEMBER_COUNTS: &[usize] = &[3, 4, 6, 8, 9, 12, 16, 24, 32, 64]; +use vortex_session::VortexSession; + +const DENSE_CASES: &[(usize, usize)] = &[ + (64, 1), + (64, 4), + (64, 8), + (64, 32), + (64, 64), + (1_024, 1), + (1_024, 4), + (1_024, 8), + (1_024, 32), + (1_024, 64), + (65_536, 1), + (65_536, 4), + (65_536, 8), + (65_536, 32), + (65_536, 64), +]; +const SPARSE_CASES: &[(usize, usize)] = &[(1_024, 8), (1_024, 64), (65_536, 8), (65_536, 64)]; +const DENSITY_CASES: &[(usize, usize, u32)] = &[ + (64, 5, 1_000), + (64, 8, 512), + (64, 64, 64), + (65_536, 5, 1_000), + (65_536, 8, 512), + (65_536, 64, 64), +]; fn main() { divan::main(); } -fn members(count: usize) -> Vec { - (0..count).map(|index| index as u32 * 2).collect() -} - -fn mostly_missing_values() -> Vec { - (0..LEN) - .map(|index| ((index as u32 * 17) % 4_096) | 1) - .collect() -} - -fn mixed_values(members: &[u32]) -> Vec { - (0..LEN) - .map(|index| { - if index.is_multiple_of(2) { - members[(index / 2) % members.len()] - } else { - ((index as u32 * 17) % 4_096) | 1 - } - }) - .collect() -} - -fn count_linear(values: &[u32], members: &[u32]) -> usize { - values - .iter() - .filter(|value| members.contains(black_box(value))) - .count() -} - -fn count_binary(values: &[u32], members: &[u32]) -> usize { - values - .iter() - .filter(|value| members.binary_search(black_box(value)).is_ok()) - .count() -} - -#[divan::bench(args = MEMBER_COUNTS)] -fn linear_mostly_missing(bencher: Bencher, member_count: usize) { - let values = mostly_missing_values(); - let members = members(member_count); - bencher - .counter(ItemsCount::new(LEN)) - .bench_local(|| black_box(count_linear(&values, &members))); -} - -#[divan::bench(args = MEMBER_COUNTS)] -fn binary_mostly_missing(bencher: Bencher, member_count: usize) { - let values = mostly_missing_values(); - let members = members(member_count); - bencher - .counter(ItemsCount::new(LEN)) - .bench_local(|| black_box(count_binary(&values, &members))); -} - -#[divan::bench(args = MEMBER_COUNTS)] -fn linear_mixed(bencher: Bencher, member_count: usize) { - let members = members(member_count); - let values = mixed_values(&members); - bencher - .counter(ItemsCount::new(LEN)) - .bench_local(|| black_box(count_linear(&values, &members))); -} - -#[divan::bench(args = MEMBER_COUNTS)] -fn binary_mixed(bencher: Bencher, member_count: usize) { - let members = members(member_count); - let values = mixed_values(&members); - bencher - .counter(ItemsCount::new(LEN)) - .bench_local(|| black_box(count_binary(&values, &members))); +fn members(count: usize, stride: u32) -> Vec { + (0..count).map(|index| index as u32 * stride).collect() } fn page_aligned(array: BitPackedArray) -> BitPackedArray { @@ -133,9 +89,21 @@ fn page_aligned(array: BitPackedArray) -> BitPackedArray { .unwrap() } -fn kernel_inputs(member_count: usize) -> (BitPackedArray, ArrayRef) { - let mut ctx = array_session().create_execution_ctx(); - let values: BufferMut = (0..LEN).map(|index| (index as u32 * 17) % 1_024).collect(); +fn benchmark_input( + len: usize, + member_count: usize, + member_stride: u32, + specialized: bool, +) -> (ArrayRef, VortexSession) { + let session = array_session(); + if specialized { + vortex_fastlanes::initialize(&session); + } else { + session.arrays().register(BitPacked); + } + + let mut ctx = session.create_execution_ctx(); + let values: BufferMut = (0..len).map(|index| (index as u32 * 17) % 1_024).collect(); let packed = page_aligned( BitPackedData::encode( &PrimitiveArray::new(values.freeze(), Validity::NonNullable).into_array(), @@ -144,71 +112,62 @@ fn kernel_inputs(member_count: usize) -> (BitPackedArray, ArrayRef) { ) .unwrap(), ); - let member_scalars = members(member_count) + let member_scalars = members(member_count, member_stride) .into_iter() .map(|value| Scalar::primitive(value, Nullability::NonNullable)) .collect(); - let list = ConstantArray::new( - Scalar::list( - Arc::new(DType::Primitive(PType::U32, Nullability::NonNullable)), - member_scalars, - Nullability::NonNullable, - ), - LEN, - ) - .into_array(); - (packed, list) + let list = Scalar::list( + Arc::new(DType::Primitive(PType::U32, Nullability::NonNullable)), + member_scalars, + Nullability::NonNullable, + ); + let contains = packed + .into_array() + .apply(&list_contains(lit(list), root())) + .unwrap(); + (contains, session) } -#[divan::bench(args = MEMBER_COUNTS)] -fn bitpacked_kernel(bencher: Bencher, member_count: usize) { - let (packed, list) = kernel_inputs(member_count); - let mut ctx = array_session().create_execution_ctx(); - bencher.counter(ItemsCount::new(LEN)).bench_local(|| { - black_box( - ::list_contains( - &list, - packed.as_view(), - &mut ctx, - ) - .unwrap() - .unwrap(), - ) - }); +fn bench_contains( + bencher: Bencher, + len: usize, + member_count: usize, + member_stride: u32, + specialized: bool, +) { + let (contains, session) = benchmark_input(len, member_count, member_stride, specialized); + let mut ctx = session.create_execution_ctx(); + bencher + .counter(ItemsCount::new(len)) + .bench_local(|| black_box(contains.clone().execute::(&mut ctx).unwrap())); } -#[divan::bench(args = STRATEGY_MEMBER_COUNTS)] -fn bitpacked_linear(bencher: Bencher, member_count: usize) { - let (packed, list) = kernel_inputs(member_count); - let mut ctx = array_session().create_execution_ctx(); - bencher.counter(ItemsCount::new(LEN)).bench_local(|| { - black_box( - list_contains_test_harness::list_contains( - &list, - packed.as_view(), - MembershipSearch::Linear, - &mut ctx, - ) - .unwrap() - .unwrap(), - ) - }); +#[divan::bench(args = DENSE_CASES)] +fn compressed_dense(bencher: Bencher, (len, member_count): (usize, usize)) { + bench_contains(bencher, len, member_count, 2, true); } -#[divan::bench(args = STRATEGY_MEMBER_COUNTS)] -fn bitpacked_binary(bencher: Bencher, member_count: usize) { - let (packed, list) = kernel_inputs(member_count); - let mut ctx = array_session().create_execution_ctx(); - bencher.counter(ItemsCount::new(LEN)).bench_local(|| { - black_box( - list_contains_test_harness::list_contains( - &list, - packed.as_view(), - MembershipSearch::Binary, - &mut ctx, - ) - .unwrap() - .unwrap(), - ) - }); +#[divan::bench(args = DENSE_CASES)] +fn canonical_dense(bencher: Bencher, (len, member_count): (usize, usize)) { + bench_contains(bencher, len, member_count, 2, false); +} + +#[divan::bench(args = SPARSE_CASES)] +fn compressed_sparse(bencher: Bencher, (len, member_count): (usize, usize)) { + bench_contains(bencher, len, member_count, 10_000, true); +} + +#[divan::bench(args = SPARSE_CASES)] +fn canonical_sparse(bencher: Bencher, (len, member_count): (usize, usize)) { + bench_contains(bencher, len, member_count, 10_000, false); +} + +#[divan::bench(args = DENSITY_CASES)] +fn compressed_density(bencher: Bencher, (len, member_count, member_stride): (usize, usize, u32)) { + bench_contains(bencher, len, member_count, member_stride, true); +} + +#[divan::bench(args = DENSITY_CASES)] +fn canonical_density(bencher: Bencher, (len, member_count, member_stride): (usize, usize, u32)) { + bench_contains(bencher, len, member_count, member_stride, false); } diff --git a/encodings/fastlanes/src/bitpacking/compute/list_contains/mod.rs b/encodings/fastlanes/src/bitpacking/compute/list_contains/mod.rs index f93c09af30f..8c7138a2313 100644 --- a/encodings/fastlanes/src/bitpacking/compute/list_contains/mod.rs +++ b/encodings/fastlanes/src/bitpacking/compute/list_contains/mod.rs @@ -5,38 +5,35 @@ use vortex_array::ArrayRef; use vortex_array::ArrayView; use vortex_array::ExecutionCtx; use vortex_array::IntoArray; +use vortex_array::arrays::BoolArray; use vortex_array::arrays::ConstantArray; +use vortex_array::arrays::PrimitiveArray; use vortex_array::dtype::DType; use vortex_array::dtype::NativePType; use vortex_array::match_each_integer_ptype; use vortex_array::scalar::Scalar; +use vortex_array::scalar_fn::fns::list_contains::IntegerMembership; use vortex_array::scalar_fn::fns::list_contains::ListContainsElementKernel; +use vortex_buffer::BitBuffer; use vortex_error::VortexResult; use vortex_error::vortex_err; use super::compare_fused::stream_compare_fused; use crate::BitPacked; -#[derive(Clone, Copy)] -enum SearchStrategy { - Linear, - Binary, -} - impl ListContainsElementKernel for BitPacked { fn list_contains( list: &ArrayRef, element: ArrayView<'_, Self>, ctx: &mut ExecutionCtx, ) -> VortexResult> { - list_contains_with_strategy(list, element, SearchStrategy::Binary, ctx) + list_contains_compressed(list, element, ctx) } } -fn list_contains_with_strategy( +fn list_contains_compressed( list: &ArrayRef, element: ArrayView<'_, BitPacked>, - strategy: SearchStrategy, ctx: &mut ExecutionCtx, ) -> VortexResult> { let Some(list_scalar) = list.as_constant() else { @@ -57,7 +54,7 @@ fn list_contains_with_strategy( }; let result = match_each_integer_ptype!(element.dtype().as_ptype(), |T| { - let mut members = elements + let members = elements .iter() .map(|value| { value @@ -70,7 +67,15 @@ fn list_contains_with_strategy( .flatten() .collect::>(); - match members.as_slice() { + if members.is_empty() && !elements.is_empty() { + let validity = element.validity()?.union_nullability(nullability); + return Ok(Some( + BoolArray::new(BitBuffer::new_unset(element.len()), validity).into_array(), + )); + } + let membership = IntegerMembership::new(members); + + match membership.members() { [] => ConstantArray::new(Scalar::bool(false, nullability), element.len()).into_array(), [member] => { let member = *member; @@ -86,63 +91,52 @@ fn list_contains_with_strategy( ctx, )? } - _ if matches!(strategy, SearchStrategy::Linear) => stream_compare_fused::( - element, - members[0], - nullability, - |value, _| members.contains(&value), - ctx, - )?, - _ => { - members.sort_unstable(); - members.dedup(); + [first, second, third] => { + let (first, second, third) = (*first, *second, *third); stream_compare_fused::( element, - members[0], + first, + nullability, + move |value, _| value.is_eq(first) | value.is_eq(second) | value.is_eq(third), + ctx, + )? + } + [first, second, third, fourth] => { + let (first, second, third, fourth) = (*first, *second, *third, *fourth); + stream_compare_fused::( + element, + first, nullability, - |value, _| members.binary_search(&value).is_ok(), + move |value, _| { + value.is_eq(first) + | value.is_eq(second) + | value.is_eq(third) + | value.is_eq(fourth) + }, ctx, )? } + _ => { + if membership.uses_dense_table() { + stream_compare_fused::( + element, + membership.members()[0], + nullability, + |value, _| membership.contains(value), + ctx, + )? + } else { + let primitive = element + .into_owned() + .into_array() + .execute::(ctx)?; + membership.evaluate_primitive(primitive.as_view(), nullability)? + } + } } }); Ok(Some(result)) } -#[cfg(feature = "_test-harness")] -pub mod test_harness { - use vortex_array::ArrayRef; - use vortex_array::ArrayView; - use vortex_array::ExecutionCtx; - use vortex_error::VortexResult; - - use super::SearchStrategy; - use super::list_contains_with_strategy; - use crate::BitPacked; - - /// Selects the membership lookup strategy for a benchmark invocation. - #[derive(Clone, Copy)] - pub enum MembershipSearch { - /// Scans list members in order. - Linear, - /// Sorts list members and uses binary search. - Binary, - } - - /// Executes the BitPacked membership kernel with a fixed lookup strategy. - pub fn list_contains( - list: &ArrayRef, - element: ArrayView<'_, BitPacked>, - strategy: MembershipSearch, - ctx: &mut ExecutionCtx, - ) -> VortexResult> { - let strategy = match strategy { - MembershipSearch::Linear => SearchStrategy::Linear, - MembershipSearch::Binary => SearchStrategy::Binary, - }; - list_contains_with_strategy(list, element, strategy, ctx) - } -} - #[cfg(test)] mod tests; diff --git a/encodings/fastlanes/src/bitpacking/compute/list_contains/tests.rs b/encodings/fastlanes/src/bitpacking/compute/list_contains/tests.rs index d237d583e4e..4cfd5ebe2c7 100644 --- a/encodings/fastlanes/src/bitpacking/compute/list_contains/tests.rs +++ b/encodings/fastlanes/src/bitpacking/compute/list_contains/tests.rs @@ -16,14 +16,20 @@ use vortex_array::assert_arrays_eq; use vortex_array::dtype::DType; use vortex_array::dtype::NativePType; use vortex_array::dtype::Nullability; +#[cfg(not(codspeed))] use vortex_array::expr::list_contains; +#[cfg(not(codspeed))] use vortex_array::expr::lit; +#[cfg(not(codspeed))] use vortex_array::expr::root; use vortex_array::scalar::PValue; use vortex_array::scalar::Scalar; use vortex_array::scalar_fn::fns::list_contains::ListContainsElementKernel; +#[cfg(not(codspeed))] use vortex_array::test_harness::trace::TraceOptions; +#[cfg(not(codspeed))] use vortex_array::test_harness::trace::TraceResolution; +#[cfg(not(codspeed))] use vortex_array::test_harness::trace::trace_op_with; use vortex_array::validity::Validity; use vortex_error::VortexResult; @@ -113,10 +119,9 @@ integer_type_test!(test_integer_type_i64, i64, 6); #[case::one(vec![3])] #[case::two(vec![3, 7])] #[case::three(vec![3, 7, 11])] -#[case::seven((0..7).map(|value| value * 2 + 1).collect())] -#[case::eight((0..8).map(|value| value * 2 + 1).collect())] -#[case::nine((0..9).map(|value| value * 2 + 1).collect())] +#[case::four(vec![3, 7, 11, 15])] #[case::larger((0..32).map(|value| value * 3).collect())] +#[case::sparse((0..32).map(|value| value * 10_000).collect())] #[case::duplicates(vec![3, 3, 7, 7, 11, 11, 15, 15, 15])] fn test_member_cardinalities(#[case] members: Vec) -> VortexResult<()> { let mut ctx = SESSION.create_execution_ctx(); @@ -241,6 +246,31 @@ fn test_nullable_members_are_ignored() -> VortexResult<()> { Ok(()) } +#[test] +fn test_empty_and_all_null_members_with_null_needles() -> VortexResult<()> { + let mut ctx = SESSION.create_execution_ctx(); + let values = [Some(1i32), None, Some(2)]; + let primitive = PrimitiveArray::from_option_iter(values); + let packed = BitPackedData::encode(&primitive.into_array(), 2, &mut ctx)?; + + let empty_list = list_array( + member_list(std::iter::empty::>(), Nullability::Nullable), + packed.len(), + ); + let actual = execute_direct(&empty_list, &packed, &mut ctx)?; + let expected = BoolArray::from_iter([Some(false), Some(false), Some(false)]); + assert_arrays_eq!(actual, expected, &mut ctx); + + let all_null_list = list_array( + member_list([None::], Nullability::Nullable), + packed.len(), + ); + let actual = execute_direct(&all_null_list, &packed, &mut ctx)?; + let expected = BoolArray::from_iter([Some(false), None, Some(false)]); + assert_arrays_eq!(actual, expected, &mut ctx); + Ok(()) +} + #[test] fn test_wrong_integer_type_declines_without_panic() -> VortexResult<()> { let mut ctx = SESSION.create_execution_ctx(); @@ -278,6 +308,7 @@ fn test_noninteger_list_declines_without_panic() -> VortexResult<()> { } #[test] +#[cfg(not(codspeed))] fn test_registered_kernel_executes_through_expression() -> VortexResult<()> { let mut ctx = SESSION.create_execution_ctx(); let values = (0..2_048).map(|value| value % 128).collect::>(); diff --git a/encodings/fastlanes/src/bitpacking/mod.rs b/encodings/fastlanes/src/bitpacking/mod.rs index d27efc9dc92..efa0677a91e 100644 --- a/encodings/fastlanes/src/bitpacking/mod.rs +++ b/encodings/fastlanes/src/bitpacking/mod.rs @@ -13,10 +13,6 @@ pub use array::unpack_iter; pub(crate) mod compute; -#[cfg(feature = "_test-harness")] -#[doc(hidden)] -pub use compute::list_contains::test_harness as list_contains_test_harness; - mod plugin; mod vtable; diff --git a/vortex-array/src/arrays/primitive/compute/list_contains.rs b/vortex-array/src/arrays/primitive/compute/list_contains.rs new file mode 100644 index 00000000000..b32e9900d0e --- /dev/null +++ b/vortex-array/src/arrays/primitive/compute/list_contains.rs @@ -0,0 +1,155 @@ +// SPDX-License-Identifier: Apache-2.0 +// SPDX-FileCopyrightText: Copyright the Vortex contributors + +use vortex_error::VortexExpect; +use vortex_error::VortexResult; + +use crate::ArrayRef; +use crate::ArrayView; +use crate::ExecutionCtx; +use crate::IntoArray; +use crate::arrays::ConstantArray; +use crate::arrays::Primitive; +use crate::dtype::DType; +use crate::match_each_integer_ptype; +use crate::scalar::Scalar; +use crate::scalar_fn::fns::list_contains::IntegerMembership; +use crate::scalar_fn::fns::list_contains::ListContainsElementKernel; + +impl ListContainsElementKernel for Primitive { + fn list_contains( + list: &ArrayRef, + element: ArrayView<'_, Self>, + _ctx: &mut ExecutionCtx, + ) -> VortexResult> { + let Some(list_scalar) = list.as_constant() else { + return Ok(None); + }; + let DType::List(member_dtype, _) = list.dtype() else { + return Ok(None); + }; + if !member_dtype.eq_ignore_nullability(element.dtype()) || !element.ptype().is_int() { + return Ok(None); + } + + let nullability = list.dtype().nullability() | element.dtype().nullability(); + let Some(elements) = list_scalar.as_list().elements() else { + return Ok(Some( + ConstantArray::new(Scalar::null(DType::Bool(nullability)), element.len()) + .into_array(), + )); + }; + if elements.is_empty() { + return Ok(Some( + ConstantArray::new(Scalar::bool(false, nullability), element.len()).into_array(), + )); + } + + let result = match_each_integer_ptype!(element.ptype(), |T| { + let members = elements + .iter() + .map(|value| { + value + .as_primitive_opt() + .vortex_expect("list dtype was checked before member extraction") + .try_typed_value::() + }) + .collect::>>>()? + .into_iter() + .flatten() + .collect::>(); + + IntegerMembership::new(members).evaluate_primitive(element, nullability)? + }); + + Ok(Some(result)) + } +} + +#[cfg(test)] +mod tests { + use std::sync::Arc; + + use rstest::rstest; + use vortex_buffer::BitBuffer; + + use super::*; + use crate::VortexSessionExecute; + use crate::arrays::BoolArray; + use crate::arrays::PrimitiveArray; + use crate::assert_arrays_eq; + use crate::dtype::Nullability; + use crate::dtype::PType::I32; + + fn list(values: impl IntoIterator, len: usize) -> ArrayRef { + ConstantArray::new( + Scalar::list( + Arc::new(DType::Primitive(I32, Nullability::NonNullable)), + values + .into_iter() + .map(|value| Scalar::primitive(value, Nullability::NonNullable)) + .collect(), + Nullability::NonNullable, + ), + len, + ) + .into_array() + } + + #[rstest] + #[case::empty(vec![])] + #[case::one(vec![3])] + #[case::four(vec![3, 7, 11, 15])] + #[case::dense((0..32).map(|value| value * 3).collect())] + #[case::sparse((0..32).map(|value| value * 10_000).collect())] + fn test_membership_plans(#[case] members: Vec) -> VortexResult<()> { + let mut ctx = crate::array_session().create_execution_ctx(); + let values = [0, 3, 7, 15, 31, 90_000, 310_000]; + let element = PrimitiveArray::from_iter(values); + let expected = BoolArray::from_iter(values.map(|value| members.contains(&value))); + + let actual = ::list_contains( + &list(members, element.len()), + element.as_view(), + &mut ctx, + )? + .vortex_expect("integer constant-list membership is supported"); + + assert_arrays_eq!(actual, expected, &mut ctx); + Ok(()) + } + + #[test] + fn test_null_needles() -> VortexResult<()> { + let mut ctx = crate::array_session().create_execution_ctx(); + let element = PrimitiveArray::from_option_iter([Some(1), None, Some(2)]); + let expected = BoolArray::from_iter([Some(true), None, Some(false)]); + + let actual = ::list_contains( + &list([1, 3], element.len()), + element.as_view(), + &mut ctx, + )? + .vortex_expect("integer constant-list membership is supported"); + + assert_arrays_eq!(actual, expected, &mut ctx); + Ok(()) + } + + #[test] + fn test_empty_list_ignores_needle_validity() -> VortexResult<()> { + let mut ctx = crate::array_session().create_execution_ctx(); + let element = PrimitiveArray::from_option_iter([Some(1i32), None, Some(2)]); + let expected = BoolArray::new(BitBuffer::new_unset(3), crate::validity::Validity::AllValid); + + let actual = ::list_contains( + &list([], element.len()), + element.as_view(), + &mut ctx, + )? + .vortex_expect("integer constant-list membership is supported"); + + assert_arrays_eq!(actual, expected, &mut ctx); + Ok(()) + } +} diff --git a/vortex-array/src/arrays/primitive/compute/mod.rs b/vortex-array/src/arrays/primitive/compute/mod.rs index 382b42ee6e2..7f1dcdcb4cf 100644 --- a/vortex-array/src/arrays/primitive/compute/mod.rs +++ b/vortex-array/src/arrays/primitive/compute/mod.rs @@ -5,6 +5,7 @@ mod between; mod cast; mod fill_null; mod fixed_width; +mod list_contains; mod mask; pub(crate) mod rules; mod slice; diff --git a/vortex-array/src/arrays/primitive/vtable/kernel.rs b/vortex-array/src/arrays/primitive/vtable/kernel.rs index 6382ea73794..3f13282c334 100644 --- a/vortex-array/src/arrays/primitive/vtable/kernel.rs +++ b/vortex-array/src/arrays/primitive/vtable/kernel.rs @@ -15,6 +15,8 @@ use crate::scalar_fn::fns::cast::Cast; use crate::scalar_fn::fns::cast::CastExecuteAdaptor; use crate::scalar_fn::fns::fill_null::FillNull; use crate::scalar_fn::fns::fill_null::FillNullExecuteAdaptor; +use crate::scalar_fn::fns::list_contains::ListContains; +use crate::scalar_fn::fns::list_contains::ListContainsElementExecuteAdaptor; use crate::scalar_fn::fns::zip::Zip; use crate::scalar_fn::fns::zip::ZipExecuteAdaptor; @@ -31,6 +33,11 @@ pub(crate) fn initialize(session: &VortexSession) { Primitive, FillNullExecuteAdaptor(Primitive), ); + kernels.register_execute_parent_kernel( + ListContains.id(), + Primitive, + ListContainsElementExecuteAdaptor(Primitive), + ); kernels.register_execute_parent_kernel(Dict.id(), Primitive, TakeExecuteAdaptor(Primitive)); kernels.register_execute_parent_kernel(Zip.id(), Primitive, ZipExecuteAdaptor(Primitive)); } diff --git a/vortex-array/src/scalar_fn/fns/list_contains/integer_membership.rs b/vortex-array/src/scalar_fn/fns/list_contains/integer_membership.rs new file mode 100644 index 00000000000..5a6d17eb2a2 --- /dev/null +++ b/vortex-array/src/scalar_fn/fns/list_contains/integer_membership.rs @@ -0,0 +1,184 @@ +// SPDX-License-Identifier: Apache-2.0 +// SPDX-FileCopyrightText: Copyright the Vortex contributors + +use vortex_buffer::BitBuffer; +use vortex_error::VortexExpect; +use vortex_error::VortexResult; +use vortex_error::vortex_ensure; + +use crate::ArrayRef; +use crate::ArrayView; +use crate::IntoArray; +use crate::arrays::BoolArray; +use crate::arrays::Primitive; +use crate::dtype::IntegerPType; +use crate::dtype::NativePType; +use crate::dtype::Nullability; + +const MAX_DENSE_SPAN: usize = 4_096; + +/// A prepared integer set for constant-list membership kernels. +/// +/// The set sorts and deduplicates lists with more than four members. It builds a byte table when +/// the member span fits the bounded table. +pub struct IntegerMembership { + members: Box<[T]>, + dense: Option, +} + +impl IntegerMembership { + /// Prepares a membership set from integer values. + pub fn new(mut members: Vec) -> Self { + if members.len() > 4 { + members.sort_unstable(); + members.dedup(); + } + let dense = DenseIntegerMembership::try_new(&members); + + Self { + members: members.into_boxed_slice(), + dense, + } + } + + /// Returns the normalized members. + pub fn members(&self) -> &[T] { + &self.members + } + + /// Returns true when this set uses a dense lookup table. + pub fn uses_dense_table(&self) -> bool { + self.dense.is_some() + } + + /// Tests membership through the selected lookup representation. + pub fn contains(&self, value: T) -> bool { + self.dense.as_ref().map_or_else( + || { + if self.members.len() <= 4 { + self.members.contains(&value) + } else { + self.members.binary_search(&value).is_ok() + } + }, + |dense| dense.contains(value), + ) + } + + /// Evaluates this set against a primitive array of the same integer type. + pub fn evaluate_primitive( + &self, + element: ArrayView<'_, Primitive>, + nullability: Nullability, + ) -> VortexResult { + vortex_ensure!( + element.ptype() == T::PTYPE, + "Membership type {} does not match array type {}", + T::PTYPE, + element.ptype(), + ); + let values = element.as_slice::(); + let bits = match self.members() { + [] => BitBuffer::new_unset(values.len()), + [member] => collect_direct(values, move |value| value.is_eq(*member)), + [first, second] => collect_direct(values, move |value| { + value.is_eq(*first) | value.is_eq(*second) + }), + [first, second, third] => collect_direct(values, move |value| { + value.is_eq(*first) | value.is_eq(*second) | value.is_eq(*third) + }), + [first, second, third, fourth] => collect_direct(values, move |value| { + value.is_eq(*first) + | value.is_eq(*second) + | value.is_eq(*third) + | value.is_eq(*fourth) + }), + _ => collect_many(values, self), + }; + + Ok(BoolArray::new(bits, element.validity()?.union_nullability(nullability)).into_array()) + } +} + +fn collect_direct(values: &[T], mut predicate: impl FnMut(T) -> bool) -> BitBuffer { + BitBuffer::collect_bool_multiversioned(values.len(), |index| { + // SAFETY: collect_bool_multiversioned visits each valid index once. + predicate(unsafe { *values.get_unchecked(index) }) + }) +} + +fn collect_many(values: &[T], membership: &IntegerMembership) -> BitBuffer { + if let Some(dense) = membership.dense.as_ref() { + return BitBuffer::collect_bool(values.len(), |index| { + // SAFETY: collect_bool visits each valid index once. + let value = unsafe { *values.get_unchecked(index) }; + dense.contains(value) + }); + } + + BitBuffer::collect_bool(values.len(), |index| { + // SAFETY: collect_bool visits each valid index once. + let value = unsafe { *values.get_unchecked(index) }; + membership.contains(value) + }) +} + +/// A bounded byte table for dense integer membership. +struct DenseIntegerMembership { + minimum: i128, + table: Box<[u8]>, +} + +impl DenseIntegerMembership { + fn try_new(members: &[T]) -> Option { + if members.len() <= 4 { + return None; + } + + let minimum = members[0].to_i128()?; + let maximum = members[members.len() - 1].to_i128()?; + let span = usize::try_from(maximum - minimum + 1).ok()?; + if span > MAX_DENSE_SPAN { + return None; + } + + let mut table = vec![0u8; span]; + for member in members { + let index = usize::try_from( + member.to_i128().vortex_expect("integer converts to i128") - minimum, + ) + .vortex_expect("member lies inside the dense span"); + table[index] = 1; + } + + Some(Self { + minimum, + table: table.into_boxed_slice(), + }) + } + + /// Tests whether the table contains an integer value. + fn contains(&self, value: T) -> bool { + let offset = value.to_i128().vortex_expect("integer converts to i128") - self.minimum; + usize::try_from(offset) + .ok() + .and_then(|offset| self.table.get(offset)) + .copied() + .unwrap_or(0) + != 0 + } +} + +#[cfg(test)] +mod tests { + use super::IntegerMembership; + + #[test] + fn small_unsorted_set_contains_members() { + let membership = IntegerMembership::new(vec![7i32, 3]); + + assert!(membership.contains(3)); + assert!(membership.contains(7)); + assert!(!membership.contains(5)); + } +} diff --git a/vortex-array/src/scalar_fn/fns/list_contains/mod.rs b/vortex-array/src/scalar_fn/fns/list_contains/mod.rs index d2508014089..dcce13bce38 100644 --- a/vortex-array/src/scalar_fn/fns/list_contains/mod.rs +++ b/vortex-array/src/scalar_fn/fns/list_contains/mod.rs @@ -1,11 +1,13 @@ // SPDX-License-Identifier: Apache-2.0 // SPDX-FileCopyrightText: Copyright the Vortex contributors +mod integer_membership; mod kernel; use std::ops::BitOr; use arrow_buffer::bit_iterator::BitIndexIterator; +pub use integer_membership::IntegerMembership; pub use kernel::*; use num_traits::Zero; use vortex_buffer::BitBuffer; @@ -13,6 +15,7 @@ use vortex_error::VortexExpect; use vortex_error::VortexResult; use vortex_error::vortex_bail; use vortex_error::vortex_err; +use vortex_mask::Mask; use vortex_session::VortexSession; use vortex_session::registry::CachedId; use vortex_utils::iter::ReduceBalancedIterExt; @@ -24,6 +27,7 @@ use crate::arrays::BoolArray; use crate::arrays::Constant; use crate::arrays::ConstantArray; use crate::arrays::ListViewArray; +use crate::arrays::Primitive; use crate::arrays::PrimitiveArray; use crate::arrays::ScalarFnArray; use crate::arrays::bool::BoolArrayExt; @@ -146,8 +150,7 @@ impl ScalarFnVTable for ListContains { fn compute_contains_scalar(list: &Scalar, needle: &Scalar) -> VortexResult { let nullability = list.dtype().nullability() | needle.dtype().nullability(); - // Handle null list or null needle - if list.is_null() || needle.is_null() { + if list.is_null() { return Ok(Scalar::null(DType::Bool(nullability))); } @@ -155,6 +158,12 @@ fn compute_contains_scalar(list: &Scalar, needle: &Scalar) -> VortexResult(ctx)?; + if let Some(result) = + ::list_contains(array, value.as_view(), ctx)? + { + return Ok(result); + } + } + + if array.all_invalid(ctx)? { return Ok(ConstantArray::new( Scalar::null(DType::Bool(Nullability::Nullable)), array.len(), @@ -206,6 +226,10 @@ fn constant_list_scalar_contains( let len = values.len(); let false_scalar = Scalar::bool(false, nullability); + if elements.is_empty() { + return Ok(ConstantArray::new(false_scalar, len).into_array()); + } + let result = elements .iter() .map(|element| { @@ -221,7 +245,9 @@ fn constant_list_scalar_contains( .into_iter() .try_reduce_balanced(|acc, res| acc.binary(res, Operator::Or))?; - Ok(result.unwrap_or_else(|| ConstantArray::new(false_scalar, len).into_array())) + result + .unwrap_or_else(|| ConstantArray::new(false_scalar, len).into_array()) + .mask(values.validity()?.to_array(len)) } /// Returns a [`BoolArray`] where each bit represents if a list contains the scalar. @@ -244,6 +270,9 @@ fn list_contains_scalar( // Must return false when a list is empty (but valid), or null when the list itself is null. return list_false_or_null(&list_array, nullability); } + if value.is_null() { + return list_false_if_empty_else_null(&list_array, nullability, ctx); + } let rhs = ConstantArray::new(value.clone(), elems.len()); let matching_elements = @@ -305,6 +334,25 @@ fn list_contains_scalar( .into_array()) } +/// Returns false for valid empty lists and null for all other lists. +fn list_false_if_empty_else_null( + list_array: &ListViewArray, + nullability: Nullability, + ctx: &mut ExecutionCtx, +) -> VortexResult { + let sizes = list_array.sizes().clone().execute::(ctx)?; + let empty = match_each_integer_ptype!(sizes.ptype(), |S| { + Mask::from_iter(sizes.as_slice::().iter().map(|size| size.is_zero())) + }); + let valid = list_array.validity()?.execute_mask(list_array.len(), ctx)? & ∅ + + Ok(BoolArray::new( + BitBuffer::new_unset(list_array.len()), + Validity::from_mask(valid, nullability), + ) + .into_array()) +} + /// Returns a [`BitBuffer`] where each bit represents if a list contains the scalar, derived from a /// [`BoolArray`] of matches on the child elements array. fn process_matches( @@ -749,7 +797,7 @@ mod tests { #[case( null_strings(vec![vec![], vec![None, None], vec![None, None, None]]), None, - bool_array(vec![false, true, true], Validity::AllInvalid) + BoolArray::from_iter([Some(false), None, None]) )] #[case( null_strings(vec![vec![], vec![None, None], vec![None, None, None]]), @@ -796,6 +844,45 @@ mod tests { assert_arrays_eq!(contains, expected, &mut ctx); } + #[rstest] + #[case::empty( + Vec::>::new(), + [Some(false), Some(false), Some(false)] + )] + #[case::nonempty( + vec![Some(1), Some(3)], + [Some(true), None, Some(false)] + )] + #[case::all_null( + vec![None, None], + [Some(false), None, Some(false)] + )] + fn test_constant_list_nullable_needles( + #[case] members: Vec>, + #[case] expected: [Option; 3], + ) { + let mut ctx = array_session().create_execution_ctx(); + let member_dtype = DType::Primitive(I32, Nullability::Nullable); + let list = Scalar::list( + Arc::new(member_dtype.clone()), + members + .into_iter() + .map(|member| { + member + .map(|value| Scalar::primitive(value, Nullability::Nullable)) + .unwrap_or_else(|| Scalar::null(member_dtype.clone())) + }) + .collect(), + Nullability::NonNullable, + ); + let needles = PrimitiveArray::from_option_iter([Some(1), None, Some(2)]).into_array(); + + let result = needles.apply(&list_contains(lit(list), root())).unwrap(); + let expected = BoolArray::from_iter(expected); + + assert_arrays_eq!(result, expected, &mut ctx); + } + #[test] fn test_all_nulls() { let mut ctx = array_session().create_execution_ctx(); From 2a78653575de7a1f77919d09047b6a73cfa8361e Mon Sep 17 00:00:00 2001 From: Will Manning Date: Sat, 29 Aug 2026 13:27:35 -0400 Subject: [PATCH 3/8] fix(list_contains): Harden specialized membership Signed-off-by: Will Manning --- .../benches/bitpacking_list_contains.rs | 480 ++++++++++++++---- .../bitpacking/compute/list_contains/mod.rs | 26 +- .../bitpacking/compute/list_contains/tests.rs | 197 ++++--- .../sequence/src/compute/list_contains.rs | 49 +- .../arrays/primitive/compute/list_contains.rs | 95 +++- vortex-array/src/expr/exprs.rs | 2 + .../fns/list_contains/integer_membership.rs | 28 +- .../src/scalar_fn/fns/list_contains/kernel.rs | 36 ++ .../src/scalar_fn/fns/list_contains/mod.rs | 158 ++++-- vortex-datafusion/src/convert/exprs.rs | 90 +++- vortex-datafusion/src/persistent/tests.rs | 40 ++ vortex-duckdb/src/convert/expr.rs | 42 +- vortex-duckdb/src/duckdb/value.rs | 6 +- .../src/e2e_test/vortex_scan_test.rs | 14 + 14 files changed, 958 insertions(+), 305 deletions(-) diff --git a/encodings/fastlanes/benches/bitpacking_list_contains.rs b/encodings/fastlanes/benches/bitpacking_list_contains.rs index 7479555a321..2770bd24106 100644 --- a/encodings/fastlanes/benches/bitpacking_list_contains.rs +++ b/encodings/fastlanes/benches/bitpacking_list_contains.rs @@ -1,17 +1,19 @@ // SPDX-License-Identifier: Apache-2.0 // SPDX-FileCopyrightText: Copyright the Vortex contributors -//! Compares compressed list membership with the canonical fallback. +//! Compares compressed list membership with two explicit fallback paths. //! -//! The specialized session evaluates membership while it decodes FastLanes lanes. The fallback -//! session decodes the complete array before the generic membership operation. -//! Density cases stress the 4 KiB lookup-table boundary. Sparse cases exceed that boundary. +//! The decode-once path measures the lower bound for materializing a Primitive array before +//! membership evaluation. The old-generic path freezes the former balanced equality-expression +//! fallback. Primitive cases isolate the benefit of the prepared integer membership set. //! //! Run with `cargo bench -p vortex-fastlanes --bench bitpacking_list_contains`. #![expect(clippy::cast_possible_truncation)] #![expect(clippy::unwrap_used)] +use std::fmt::Display; +use std::fmt::Formatter; use std::hint::black_box; use std::sync::Arc; @@ -22,58 +24,199 @@ use vortex_array::IntoArray; use vortex_array::VortexSessionExecute; use vortex_array::array_session; use vortex_array::arrays::BoolArray; +use vortex_array::arrays::ConstantArray; +use vortex_array::arrays::Primitive; use vortex_array::arrays::PrimitiveArray; +use vortex_array::builtins::ArrayBuiltins; use vortex_array::dtype::DType; +use vortex_array::dtype::IntegerPType; use vortex_array::dtype::Nullability; use vortex_array::dtype::PType; use vortex_array::expr::list_contains; use vortex_array::expr::lit; use vortex_array::expr::root; use vortex_array::scalar::Scalar; -use vortex_array::session::ArraySessionExt; +use vortex_array::scalar_fn::fns::binary::Binary; +use vortex_array::scalar_fn::fns::list_contains::ListContainsElementKernel; +use vortex_array::scalar_fn::fns::operators::Operator; use vortex_array::validity::Validity; use vortex_buffer::Alignment; use vortex_buffer::BufferMut; +use vortex_error::VortexExpect; use vortex_fastlanes::BitPacked; use vortex_fastlanes::BitPackedArray; use vortex_fastlanes::BitPackedData; use vortex_session::VortexSession; -const DENSE_CASES: &[(usize, usize)] = &[ - (64, 1), - (64, 4), - (64, 8), - (64, 32), - (64, 64), - (1_024, 1), - (1_024, 4), - (1_024, 8), - (1_024, 32), - (1_024, 64), - (65_536, 1), - (65_536, 4), - (65_536, 8), - (65_536, 32), - (65_536, 64), -]; -const SPARSE_CASES: &[(usize, usize)] = &[(1_024, 8), (1_024, 64), (65_536, 8), (65_536, 64)]; -const DENSITY_CASES: &[(usize, usize, u32)] = &[ - (64, 5, 1_000), - (64, 8, 512), - (64, 64, 64), - (65_536, 5, 1_000), - (65_536, 8, 512), - (65_536, 64, 64), -]; - fn main() { divan::main(); } -fn members(count: usize, stride: u32) -> Vec { - (0..count).map(|index| index as u32 * stride).collect() +trait BenchInt: IntegerPType + Copy + Into { + fn from_counter(value: u64) -> Self; +} + +macro_rules! impl_bench_int { + ($($T:ty),+) => { + $(impl BenchInt for $T { + fn from_counter(value: u64) -> Self { + value as $T + } + })+ + }; +} + +impl_bench_int!(u8, u16, u32, u64); + +#[derive(Clone, Copy)] +enum MemberSpec { + Explicit(&'static [u64]), + Stride { count: usize, stride: u64 }, +} + +impl MemberSpec { + fn values(self) -> Vec { + match self { + Self::Explicit(values) => values.to_vec(), + Self::Stride { count, stride } => { + (0..count).map(|index| index as u64 * stride).collect() + } + } + } +} + +#[derive(Clone, Copy)] +struct PackedCase { + name: &'static str, + ptype: PType, + bit_width: u8, + len: usize, + members: MemberSpec, + hit_percent: u8, +} + +impl Display for PackedCase { + fn fmt(&self, formatter: &mut Formatter<'_>) -> std::fmt::Result { + write!( + formatter, + "{}_{}_w{}_n{}_hit{}", + self.name, self.ptype, self.bit_width, self.len, self.hit_percent + ) + } } +const FOUR_BOUNDARY_MEMBERS: &[u64] = &[0, 1, 2, 4_095]; +const FIVE_DENSE_BOUNDARY_MEMBERS: &[u64] = &[0, 1, 2, 3, 4_095]; +const FIVE_SPARSE_BOUNDARY_MEMBERS: &[u64] = &[0, 1, 2, 3, 4_096]; + +const PACKED_CASES: &[PackedCase] = &[ + PackedCase { + name: "short_direct_m4", + ptype: PType::U32, + bit_width: 10, + len: 1_024, + members: MemberSpec::Stride { + count: 4, + stride: 2, + }, + hit_percent: 50, + }, + PackedCase { + name: "four_member_span4096", + ptype: PType::U32, + bit_width: 13, + len: 2_048, + members: MemberSpec::Explicit(FOUR_BOUNDARY_MEMBERS), + hit_percent: 50, + }, + PackedCase { + name: "before_length_gate_m5_span4096", + ptype: PType::U32, + bit_width: 13, + len: 2_047, + members: MemberSpec::Explicit(FIVE_DENSE_BOUNDARY_MEMBERS), + hit_percent: 50, + }, + PackedCase { + name: "at_length_gate_m5_span4096", + ptype: PType::U32, + bit_width: 13, + len: 2_048, + members: MemberSpec::Explicit(FIVE_DENSE_BOUNDARY_MEMBERS), + hit_percent: 50, + }, + PackedCase { + name: "above_span_gate_m5_span4097", + ptype: PType::U32, + bit_width: 13, + len: 2_048, + members: MemberSpec::Explicit(FIVE_SPARSE_BOUNDARY_MEMBERS), + hit_percent: 50, + }, + PackedCase { + name: "long_dense_m32", + ptype: PType::U32, + bit_width: 10, + len: 65_536, + members: MemberSpec::Stride { + count: 32, + stride: 2, + }, + hit_percent: 50, + }, + PackedCase { + name: "long_sparse_m32", + ptype: PType::U64, + bit_width: 40, + len: 65_536, + members: MemberSpec::Stride { + count: 32, + stride: 10_000, + }, + hit_percent: 50, + }, + PackedCase { + name: "zero_hit_m8", + ptype: PType::U8, + bit_width: 6, + len: 65_536, + members: MemberSpec::Stride { + count: 8, + stride: 2, + }, + hit_percent: 0, + }, + PackedCase { + name: "full_hit_m8", + ptype: PType::U16, + bit_width: 12, + len: 65_536, + members: MemberSpec::Stride { + count: 8, + stride: 2, + }, + hit_percent: 100, + }, + PackedCase { + name: "wide_packed_m8", + ptype: PType::U32, + bit_width: 31, + len: 65_536, + members: MemberSpec::Stride { + count: 8, + stride: 2, + }, + hit_percent: 50, + }, +]; + +const OLD_GENERIC_CASES: &[PackedCase] = &[ + PACKED_CASES[0], + PACKED_CASES[3], + PACKED_CASES[5], + PACKED_CASES[6], +]; + fn page_aligned(array: BitPackedArray) -> BitPackedArray { let ptype = array.dtype().as_ptype(); let parts = BitPacked::into_parts(array); @@ -89,85 +232,248 @@ fn page_aligned(array: BitPackedArray) -> BitPackedArray { .unwrap() } -fn benchmark_input( - len: usize, - member_count: usize, - member_stride: u32, - specialized: bool, -) -> (ArrayRef, VortexSession) { - let session = array_session(); - if specialized { - vortex_fastlanes::initialize(&session); - } else { - session.arrays().register(BitPacked); - } +fn generated_values(case: PackedCase, members: &[u64]) -> Vec { + let domain_size = 1u64 << case.bit_width; + (0..case.len) + .map(|index| { + let is_hit = match case.hit_percent { + 0 => false, + 100 => true, + percent => index % 100 < usize::from(percent), + }; + if is_hit { + members[index % members.len()] + } else { + let mut candidate = (index as u64 * 17 + 11) % domain_size; + while members.contains(&candidate) { + candidate = (candidate + 1) % domain_size; + } + candidate + } + }) + .collect() +} + +fn list_scalar(members: &[u64]) -> Scalar { + Scalar::list( + Arc::new(DType::Primitive(T::PTYPE, Nullability::NonNullable)), + members + .iter() + .map(|value| T::from_counter(*value).into()) + .collect(), + Nullability::NonNullable, + ) +} +fn packed_input(case: PackedCase) -> (BitPackedArray, Scalar, VortexSession) { + let session = array_session(); + vortex_fastlanes::initialize(&session); let mut ctx = session.create_execution_ctx(); - let values: BufferMut = (0..len).map(|index| (index as u32 * 17) % 1_024).collect(); + let members = case.members.values(); + let values: BufferMut = generated_values(case, &members) + .into_iter() + .map(T::from_counter) + .collect(); let packed = page_aligned( BitPackedData::encode( &PrimitiveArray::new(values.freeze(), Validity::NonNullable).into_array(), - 10, + case.bit_width, &mut ctx, ) .unwrap(), ); - let member_scalars = members(member_count, member_stride) - .into_iter() - .map(|value| Scalar::primitive(value, Nullability::NonNullable)) - .collect(); - let list = Scalar::list( - Arc::new(DType::Primitive(PType::U32, Nullability::NonNullable)), - member_scalars, - Nullability::NonNullable, - ); + (packed, list_scalar::(&members), session) +} + +fn old_generic_contains(values: ArrayRef, list: &Scalar) -> ArrayRef { + let false_scalar = Scalar::bool(false, values.dtype().nullability()); + let mut level = list + .as_list() + .elements() + .vortex_expect("benchmark list is non-null") + .iter() + .map(|member| { + Binary::try_new( + ConstantArray::new(member.clone(), values.len()).into_array(), + values.clone(), + Operator::Eq, + ) + .unwrap() + .into_array() + .fill_null(false_scalar.clone()) + .unwrap() + }) + .collect::>(); + + while level.len() > 1 { + let mut next = Vec::with_capacity(level.len().div_ceil(2)); + let mut arrays = level.into_iter(); + while let Some(left) = arrays.next() { + next.push(if let Some(right) = arrays.next() { + left.binary(right, Operator::Or).unwrap() + } else { + left + }); + } + level = next; + } + + level.pop().vortex_expect("benchmark list is nonempty") +} + +fn bench_packed_specialized(bencher: Bencher, case: PackedCase) { + let (packed, list, session) = packed_input::(case); let contains = packed .into_array() .apply(&list_contains(lit(list), root())) .unwrap(); - (contains, session) -} - -fn bench_contains( - bencher: Bencher, - len: usize, - member_count: usize, - member_stride: u32, - specialized: bool, -) { - let (contains, session) = benchmark_input(len, member_count, member_stride, specialized); let mut ctx = session.create_execution_ctx(); bencher - .counter(ItemsCount::new(len)) + .counter(ItemsCount::new(case.len)) .bench_local(|| black_box(contains.clone().execute::(&mut ctx).unwrap())); } -#[divan::bench(args = DENSE_CASES)] -fn compressed_dense(bencher: Bencher, (len, member_count): (usize, usize)) { - bench_contains(bencher, len, member_count, 2, true); +fn bench_packed_decode_once(bencher: Bencher, case: PackedCase) { + let (packed, list, session) = packed_input::(case); + let list = ConstantArray::new(list, case.len).into_array(); + let mut ctx = session.create_execution_ctx(); + bencher.counter(ItemsCount::new(case.len)).bench_local(|| { + let primitive = packed + .clone() + .into_array() + .execute::(&mut ctx) + .unwrap(); + let result = ::list_contains( + &list, + primitive.as_view(), + &mut ctx, + ) + .unwrap() + .unwrap(); + black_box(result.execute::(&mut ctx).unwrap()) + }); +} + +fn bench_packed_old_generic(bencher: Bencher, case: PackedCase) { + let (packed, list, session) = packed_input::(case); + let mut ctx = session.create_execution_ctx(); + bencher.counter(ItemsCount::new(case.len)).bench_local(|| { + let result = old_generic_contains(packed.clone().into_array(), &list); + black_box(result.execute::(&mut ctx).unwrap()) + }); +} + +macro_rules! dispatch_packed { + ($bencher:expr, $case:expr, $function:ident) => { + match $case.ptype { + PType::U8 => $function::($bencher, $case), + PType::U16 => $function::($bencher, $case), + PType::U32 => $function::($bencher, $case), + PType::U64 => $function::($bencher, $case), + _ => unreachable!("benchmark case uses an unsigned integer type"), + } + }; +} + +#[divan::bench(args = PACKED_CASES)] +fn packed_specialized(bencher: Bencher, case: PackedCase) { + dispatch_packed!(bencher, case, bench_packed_specialized); } -#[divan::bench(args = DENSE_CASES)] -fn canonical_dense(bencher: Bencher, (len, member_count): (usize, usize)) { - bench_contains(bencher, len, member_count, 2, false); +#[divan::bench(args = PACKED_CASES)] +fn packed_decode_once(bencher: Bencher, case: PackedCase) { + dispatch_packed!(bencher, case, bench_packed_decode_once); } -#[divan::bench(args = SPARSE_CASES)] -fn compressed_sparse(bencher: Bencher, (len, member_count): (usize, usize)) { - bench_contains(bencher, len, member_count, 10_000, true); +#[divan::bench(args = OLD_GENERIC_CASES)] +fn packed_old_generic(bencher: Bencher, case: PackedCase) { + dispatch_packed!(bencher, case, bench_packed_old_generic); } -#[divan::bench(args = SPARSE_CASES)] -fn canonical_sparse(bencher: Bencher, (len, member_count): (usize, usize)) { - bench_contains(bencher, len, member_count, 10_000, false); +#[cfg(not(codspeed))] +fn length_sweep_cases() -> Vec { + [ + 2_048, 2_049, 2_304, 2_560, 3_072, 4_095, 4_096, 4_097, 6_144, 8_192, + ] + .map(|len| PackedCase { + name: "length_sweep_m5_span4096", + ptype: PType::U32, + bit_width: 13, + len, + members: MemberSpec::Explicit(FIVE_DENSE_BOUNDARY_MEMBERS), + hit_percent: 50, + }) + .to_vec() } -#[divan::bench(args = DENSITY_CASES)] -fn compressed_density(bencher: Bencher, (len, member_count, member_stride): (usize, usize, u32)) { - bench_contains(bencher, len, member_count, member_stride, true); +#[cfg(not(codspeed))] +#[divan::bench(args = length_sweep_cases())] +fn length_sweep_specialized(bencher: Bencher, case: PackedCase) { + bench_packed_specialized::(bencher, case); } -#[divan::bench(args = DENSITY_CASES)] -fn canonical_density(bencher: Bencher, (len, member_count, member_stride): (usize, usize, u32)) { - bench_contains(bencher, len, member_count, member_stride, false); +#[cfg(not(codspeed))] +#[divan::bench(args = length_sweep_cases())] +fn length_sweep_decode_once(bencher: Bencher, case: PackedCase) { + bench_packed_decode_once::(bencher, case); } + +fn primitive_input() -> (PrimitiveArray, Scalar, VortexSession) { + const LEN: usize = 65_536; + let case = PackedCase { + name: "primitive", + ptype: T::PTYPE, + bit_width: 12, + len: LEN, + members: MemberSpec::Stride { + count: 8, + stride: 2, + }, + hit_percent: 50, + }; + let members = case.members.values(); + let values = generated_values(case, &members) + .into_iter() + .map(T::from_counter) + .collect::(); + (values, list_scalar::(&members), array_session()) +} + +macro_rules! primitive_benchmarks { + ($module:ident, $T:ty) => { + mod $module { + use super::*; + + #[divan::bench] + fn specialized(bencher: Bencher) { + let (values, list, session) = primitive_input::<$T>(); + let len = values.len(); + let contains = values + .into_array() + .apply(&list_contains(lit(list), root())) + .unwrap(); + let mut ctx = session.create_execution_ctx(); + bencher.counter(ItemsCount::new(len)).bench_local(|| { + black_box(contains.clone().execute::(&mut ctx).unwrap()) + }); + } + + #[divan::bench] + fn old_generic(bencher: Bencher) { + let (values, list, session) = primitive_input::<$T>(); + let len = values.len(); + let values = values.into_array(); + let mut ctx = session.create_execution_ctx(); + bencher.counter(ItemsCount::new(len)).bench_local(|| { + let result = old_generic_contains(values.clone(), &list); + black_box(result.execute::(&mut ctx).unwrap()) + }); + } + } + }; +} + +primitive_benchmarks!(primitive_u8, u8); +primitive_benchmarks!(primitive_u16, u16); +primitive_benchmarks!(primitive_u32, u32); +primitive_benchmarks!(primitive_u64, u64); diff --git a/encodings/fastlanes/src/bitpacking/compute/list_contains/mod.rs b/encodings/fastlanes/src/bitpacking/compute/list_contains/mod.rs index 8c7138a2313..405978de24c 100644 --- a/encodings/fastlanes/src/bitpacking/compute/list_contains/mod.rs +++ b/encodings/fastlanes/src/bitpacking/compute/list_contains/mod.rs @@ -6,12 +6,10 @@ use vortex_array::ArrayView; use vortex_array::ExecutionCtx; use vortex_array::IntoArray; use vortex_array::arrays::BoolArray; -use vortex_array::arrays::ConstantArray; use vortex_array::arrays::PrimitiveArray; use vortex_array::dtype::DType; use vortex_array::dtype::NativePType; use vortex_array::match_each_integer_ptype; -use vortex_array::scalar::Scalar; use vortex_array::scalar_fn::fns::list_contains::IntegerMembership; use vortex_array::scalar_fn::fns::list_contains::ListContainsElementKernel; use vortex_buffer::BitBuffer; @@ -21,6 +19,9 @@ use vortex_error::vortex_err; use super::compare_fused::stream_compare_fused; use crate::BitPacked; +// Decode short batches once because their fixed fusion overhead exceeds the saved materialization. +const MIN_DENSE_FUSION_LEN: usize = 2_048; + impl ListContainsElementKernel for BitPacked { fn list_contains( list: &ArrayRef, @@ -48,10 +49,11 @@ fn list_contains_compressed( let nullability = list.dtype().nullability() | element.dtype().nullability(); let Some(elements) = list_scalar.as_list().elements() else { - return Ok(Some( - ConstantArray::new(Scalar::null(DType::Bool(nullability)), element.len()).into_array(), - )); + return Ok(None); }; + if elements.is_empty() { + return Ok(None); + } let result = match_each_integer_ptype!(element.dtype().as_ptype(), |T| { let members = elements @@ -67,16 +69,14 @@ fn list_contains_compressed( .flatten() .collect::>(); - if members.is_empty() && !elements.is_empty() { - let validity = element.validity()?.union_nullability(nullability); - return Ok(Some( - BoolArray::new(BitBuffer::new_unset(element.len()), validity).into_array(), - )); - } let membership = IntegerMembership::new(members); match membership.members() { - [] => ConstantArray::new(Scalar::bool(false, nullability), element.len()).into_array(), + [] => BoolArray::new( + BitBuffer::new_unset(element.len()), + element.validity()?.union_nullability(nullability), + ) + .into_array(), [member] => { let member = *member; stream_compare_fused::(element, member, nullability, NativePType::is_eq, ctx)? @@ -117,7 +117,7 @@ fn list_contains_compressed( )? } _ => { - if membership.uses_dense_table() { + if membership.uses_dense_table() && element.len() >= MIN_DENSE_FUSION_LEN { stream_compare_fused::( element, membership.members()[0], diff --git a/encodings/fastlanes/src/bitpacking/compute/list_contains/tests.rs b/encodings/fastlanes/src/bitpacking/compute/list_contains/tests.rs index 4cfd5ebe2c7..4eb377b4df8 100644 --- a/encodings/fastlanes/src/bitpacking/compute/list_contains/tests.rs +++ b/encodings/fastlanes/src/bitpacking/compute/list_contains/tests.rs @@ -10,6 +10,7 @@ use vortex_array::IntoArray; use vortex_array::VortexSessionExecute; use vortex_array::arrays::BoolArray; use vortex_array::arrays::ConstantArray; +use vortex_array::arrays::ListArray; use vortex_array::arrays::PrimitiveArray; use vortex_array::arrays::slice::SliceKernel; use vortex_array::assert_arrays_eq; @@ -26,12 +27,7 @@ use vortex_array::scalar::PValue; use vortex_array::scalar::Scalar; use vortex_array::scalar_fn::fns::list_contains::ListContainsElementKernel; #[cfg(not(codspeed))] -use vortex_array::test_harness::trace::TraceOptions; -#[cfg(not(codspeed))] -use vortex_array::test_harness::trace::TraceResolution; -#[cfg(not(codspeed))] -use vortex_array::test_harness::trace::trace_op_with; -use vortex_array::validity::Validity; +use vortex_array::test_harness::trace::trace_op; use vortex_error::VortexResult; use vortex_error::vortex_err; use vortex_session::VortexSession; @@ -115,11 +111,11 @@ integer_type_test!(test_integer_type_i32, i32, 6); integer_type_test!(test_integer_type_i64, i64, 6); #[rstest] -#[case::empty(vec![])] #[case::one(vec![3])] #[case::two(vec![3, 7])] #[case::three(vec![3, 7, 11])] #[case::four(vec![3, 7, 11, 15])] +#[case::five(vec![3, 7, 11, 15, 19])] #[case::larger((0..32).map(|value| value * 3).collect())] #[case::sparse((0..32).map(|value| value * 10_000).collect())] #[case::duplicates(vec![3, 3, 7, 7, 11, 11, 15, 15, 15])] @@ -139,43 +135,62 @@ fn test_member_cardinalities(#[case] members: Vec) -> VortexResult<()> { Ok(()) } -#[test] -fn test_patches() -> VortexResult<()> { +#[rstest] +#[case::present([true; 128], vec![0])] +#[case::absent([false; 128], vec![1])] +fn test_zero_bit_width( + #[case] expected: [bool; 128], + #[case] members: Vec, +) -> VortexResult<()> { let mut ctx = SESSION.create_execution_ctx(); - let values = (0..2_048) - .map(|index| { - if index % 97 == 0 { - 100_000 + index - } else { - index % 100 - } - }) - .collect::>(); - let primitive = PrimitiveArray::from_iter(values.iter().copied()); - let packed = BitPackedData::encode(&primitive.into_array(), 7, &mut ctx)?; - assert!(packed.patches().is_some(), "test setup requires patches"); - let members = [3, 100_097]; + let primitive = PrimitiveArray::from_iter([0i32; 128]); + let packed = BitPackedData::encode(&primitive.into_array(), 0, &mut ctx)?; let list = list_array( member_list(members.into_iter().map(Some), Nullability::NonNullable), packed.len(), ); let actual = execute_direct(&list, &packed, &mut ctx)?; - let expected = BoolArray::from_iter(values.into_iter().map(|value| members.contains(&value))); + let expected = BoolArray::from_iter(expected); + assert_arrays_eq!(actual, expected, &mut ctx); + Ok(()) +} + +#[test] +fn test_empty_array() -> VortexResult<()> { + let mut ctx = SESSION.create_execution_ctx(); + let primitive = PrimitiveArray::from_iter(std::iter::empty::()); + let packed = BitPackedData::encode(&primitive.into_array(), 1, &mut ctx)?; + let list = list_array( + member_list([Some(0)], Nullability::NonNullable), + packed.len(), + ); + + let actual = execute_direct(&list, &packed, &mut ctx)?; + let expected = BoolArray::from_iter(std::iter::empty::()); assert_arrays_eq!(actual, expected, &mut ctx); Ok(()) } #[test] -fn test_sliced_array() -> VortexResult<()> { +fn test_sliced_patched_array() -> VortexResult<()> { let mut ctx = SESSION.create_execution_ctx(); - let values = (0..5_000).map(|value| value % 128).collect::>(); + let values = (0..5_000) + .map(|index| { + if index % 97 == 0 { + 100_000 + index + } else { + index % 100 + } + }) + .collect::>(); let primitive = PrimitiveArray::from_iter(values.iter().copied()); let packed = BitPackedData::encode(&primitive.into_array(), 7, &mut ctx)?; + assert!(packed.patches().is_some(), "test setup requires patches"); let range = 333..4_333; let sliced = ::slice(packed.as_view(), range.clone(), &mut ctx)? .ok_or_else(|| vortex_err!("BitPacked slice kernel declined a supported input"))?; - let members = [1, 63, 127]; + let members = [3, 100_388]; let list = list_array( member_list(members.into_iter().map(Some), Nullability::NonNullable), sliced.len(), @@ -193,80 +208,38 @@ fn test_sliced_array() -> VortexResult<()> { Ok(()) } -#[test] -fn test_null_needles() -> VortexResult<()> { +#[rstest] +#[case::nullable_needles( + vec![Some(1), Some(3)], + Nullability::NonNullable, + vec![Some(1), None, Some(2)], + vec![Some(true), None, Some(false)], +)] +#[case::nullable_members( + vec![Some(1), None, Some(3)], + Nullability::Nullable, + vec![Some(1), Some(2), Some(3)], + vec![Some(true), Some(false), Some(true)], +)] +#[case::all_null_members( + vec![None, None], + Nullability::Nullable, + vec![Some(1), None, Some(2)], + vec![Some(false), None, Some(false)], +)] +fn test_null_semantics( + #[case] members: Vec>, + #[case] member_nullability: Nullability, + #[case] values: Vec>, + #[case] expected: Vec>, +) -> VortexResult<()> { let mut ctx = SESSION.create_execution_ctx(); - let values = [Some(1i32), None, Some(2), Some(3), None]; let primitive = PrimitiveArray::from_option_iter(values); - let packed = BitPackedData::encode(&primitive.into_array(), 2, &mut ctx)?; - let list = list_array( - member_list([Some(1), Some(3)], Nullability::NonNullable), - packed.len(), - ); - - let actual = execute_direct(&list, &packed, &mut ctx)?; - let expected = BoolArray::from_iter([Some(true), None, Some(false), Some(true), None]); - assert_arrays_eq!(actual, expected, &mut ctx); - Ok(()) -} - -#[test] -fn test_null_list() -> VortexResult<()> { - let mut ctx = SESSION.create_execution_ctx(); - let primitive = PrimitiveArray::from_iter([1i32, 2, 3]); - let packed = BitPackedData::encode(&primitive.into_array(), 2, &mut ctx)?; - let list_dtype = DType::List( - Arc::new(DType::Primitive(i32::PTYPE, Nullability::NonNullable)), - Nullability::Nullable, - ); - let list = list_array(Scalar::null(list_dtype), packed.len()); - - let actual = execute_direct(&list, &packed, &mut ctx)?; - let expected = BoolArray::new( - [false, false, false].into_iter().collect(), - Validity::AllInvalid, - ); - assert_arrays_eq!(actual, expected, &mut ctx); - Ok(()) -} - -#[test] -fn test_nullable_members_are_ignored() -> VortexResult<()> { - let mut ctx = SESSION.create_execution_ctx(); - let primitive = PrimitiveArray::from_iter([1i32, 2, 3, 4]); let packed = BitPackedData::encode(&primitive.into_array(), 3, &mut ctx)?; - let list = list_array( - member_list([Some(1), None, Some(3)], Nullability::Nullable), - packed.len(), - ); + let list = list_array(member_list(members, member_nullability), packed.len()); let actual = execute_direct(&list, &packed, &mut ctx)?; - let expected = BoolArray::from_iter([true, false, true, false]); - assert_arrays_eq!(actual, expected, &mut ctx); - Ok(()) -} - -#[test] -fn test_empty_and_all_null_members_with_null_needles() -> VortexResult<()> { - let mut ctx = SESSION.create_execution_ctx(); - let values = [Some(1i32), None, Some(2)]; - let primitive = PrimitiveArray::from_option_iter(values); - let packed = BitPackedData::encode(&primitive.into_array(), 2, &mut ctx)?; - - let empty_list = list_array( - member_list(std::iter::empty::>(), Nullability::Nullable), - packed.len(), - ); - let actual = execute_direct(&empty_list, &packed, &mut ctx)?; - let expected = BoolArray::from_iter([Some(false), Some(false), Some(false)]); - assert_arrays_eq!(actual, expected, &mut ctx); - - let all_null_list = list_array( - member_list([None::], Nullability::Nullable), - packed.len(), - ); - let actual = execute_direct(&all_null_list, &packed, &mut ctx)?; - let expected = BoolArray::from_iter([Some(false), None, Some(false)]); + let expected = BoolArray::from_iter(expected); assert_arrays_eq!(actual, expected, &mut ctx); Ok(()) } @@ -288,18 +261,15 @@ fn test_wrong_integer_type_declines_without_panic() -> VortexResult<()> { } #[test] -fn test_noninteger_list_declines_without_panic() -> VortexResult<()> { +fn test_nonconstant_list_declines() -> VortexResult<()> { let mut ctx = SESSION.create_execution_ctx(); let primitive = PrimitiveArray::from_iter([1i32, 2, 3]); let packed = BitPackedData::encode(&primitive.into_array(), 2, &mut ctx)?; - let list = list_array( - Scalar::list( - Arc::new(DType::Utf8(Nullability::NonNullable)), - vec![Scalar::utf8("one", Nullability::NonNullable)], - Nullability::NonNullable, - ), - packed.len(), - ); + let list = ListArray::from_iter_slow::( + vec![vec![1i32], vec![2], vec![3]], + Arc::new(DType::Primitive(i32::PTYPE, Nullability::NonNullable)), + )? + .into_array(); let result = ::list_contains(&list, packed.as_view(), &mut ctx)?; @@ -324,15 +294,18 @@ fn test_registered_kernel_executes_through_expression() -> VortexResult<()> { ); let contains = packed.into_array().apply(&expression)?; - let traced = trace_op_with( - TraceOptions { - resolution: TraceResolution::Attempts, - }, - || contains.execute::(&mut ctx), - )?; + let traced = trace_op(|| contains.execute::(&mut ctx))?; let trace = traced.trace.to_string(); - assert!(trace.contains("parent=vortex.list.contains"), "{trace}"); - assert!(trace.contains("source=session"), "{trace}"); + let applied = trace + .lines() + .filter(|line| { + line.contains("child_execute_parent session[") + && line.contains("slot=1") + && line.contains("parent=vortex.list.contains") + && line.contains("child=fastlanes.bitpacked") + }) + .collect::>(); + assert_eq!(applied.len(), 1, "{trace}"); let expected = BoolArray::from_iter(values.into_iter().map(|value| members.contains(&value))); assert_arrays_eq!(traced.output, expected, &mut ctx); diff --git a/encodings/sequence/src/compute/list_contains.rs b/encodings/sequence/src/compute/list_contains.rs index 80ffcad24cd..d2350a1c35a 100644 --- a/encodings/sequence/src/compute/list_contains.rs +++ b/encodings/sequence/src/compute/list_contains.rs @@ -6,9 +6,9 @@ use vortex_array::ArrayView; use vortex_array::IntoArray; use vortex_array::arrays::BoolArray; use vortex_array::arrays::ConstantArray; +use vortex_array::dtype::DType; use vortex_array::scalar::Scalar; use vortex_array::scalar_fn::fns::list_contains::ListContainsElementReduce; -use vortex_error::VortexExpect; use vortex_error::VortexResult; use crate::array::Sequence; @@ -23,11 +23,19 @@ impl ListContainsElementReduce for Sequence { let Some(list_scalar) = list.as_constant() else { return Ok(None); }; + let DType::List(member_dtype, _) = list.dtype() else { + return Ok(None); + }; + if !member_dtype.eq_ignore_nullability(element.dtype()) { + return Ok(None); + } - let list_elements = list_scalar - .as_list() - .elements() - .vortex_expect("non-null element (checked in entry)"); + let Some(list_elements) = list_scalar.as_list().elements() else { + return Ok(None); + }; + if list_elements.is_empty() { + return Ok(None); + } let nullability = list.dtype().nullability() | element.dtype().nullability(); @@ -65,10 +73,13 @@ mod tests { use std::sync::Arc; use std::sync::LazyLock; + use rstest::rstest; use vortex_array::IntoArray; use vortex_array::VortexSessionExecute; use vortex_array::arrays::BoolArray; + use vortex_array::arrays::Constant; use vortex_array::assert_arrays_eq; + use vortex_array::dtype::DType; use vortex_array::dtype::Nullability; use vortex_array::dtype::PType::I32; use vortex_array::expr::list_contains; @@ -139,4 +150,32 @@ mod tests { let expected = BoolArray::from_iter([Some(true), Some(true), Some(true)]); assert_arrays_eq!(result, expected, &mut SESSION.create_execution_ctx()); } + + #[rstest] + #[case::null_list( + Scalar::null(DType::List(Arc::new(I32.into()), Nullability::Nullable)), + [None, None, None] + )] + #[case::empty_list( + Scalar::list(Arc::new(I32.into()), vec![], Nullability::Nullable), + [Some(false), Some(false), Some(false)] + )] + fn test_constant_list_semantics( + #[case] list_scalar: Scalar, + #[case] expected: [Option; 3], + ) { + let array = Sequence::try_new_typed(1i32, 1, Nullability::NonNullable, 3) + .unwrap() + .into_array(); + let expr = list_contains(lit(list_scalar), root()); + + let result = array.apply(&expr).unwrap(); + + assert!(result.is::()); + assert_arrays_eq!( + result, + BoolArray::from_iter(expected), + &mut SESSION.create_execution_ctx() + ); + } } diff --git a/vortex-array/src/arrays/primitive/compute/list_contains.rs b/vortex-array/src/arrays/primitive/compute/list_contains.rs index b32e9900d0e..740441358c4 100644 --- a/vortex-array/src/arrays/primitive/compute/list_contains.rs +++ b/vortex-array/src/arrays/primitive/compute/list_contains.rs @@ -7,12 +7,9 @@ use vortex_error::VortexResult; use crate::ArrayRef; use crate::ArrayView; use crate::ExecutionCtx; -use crate::IntoArray; -use crate::arrays::ConstantArray; use crate::arrays::Primitive; use crate::dtype::DType; use crate::match_each_integer_ptype; -use crate::scalar::Scalar; use crate::scalar_fn::fns::list_contains::IntegerMembership; use crate::scalar_fn::fns::list_contains::ListContainsElementKernel; @@ -34,15 +31,10 @@ impl ListContainsElementKernel for Primitive { let nullability = list.dtype().nullability() | element.dtype().nullability(); let Some(elements) = list_scalar.as_list().elements() else { - return Ok(Some( - ConstantArray::new(Scalar::null(DType::Bool(nullability)), element.len()) - .into_array(), - )); + return Ok(None); }; if elements.is_empty() { - return Ok(Some( - ConstantArray::new(Scalar::bool(false, nullability), element.len()).into_array(), - )); + return Ok(None); } let result = match_each_integer_ptype!(element.ptype(), |T| { @@ -71,15 +63,25 @@ mod tests { use std::sync::Arc; use rstest::rstest; - use vortex_buffer::BitBuffer; use super::*; + use crate::IntoArray; use crate::VortexSessionExecute; use crate::arrays::BoolArray; + use crate::arrays::ConstantArray; use crate::arrays::PrimitiveArray; use crate::assert_arrays_eq; use crate::dtype::Nullability; use crate::dtype::PType::I32; + #[cfg(not(codspeed))] + use crate::expr::list_contains; + #[cfg(not(codspeed))] + use crate::expr::lit; + #[cfg(not(codspeed))] + use crate::expr::root; + use crate::scalar::Scalar; + #[cfg(not(codspeed))] + use crate::test_harness::trace::trace_op; fn list(values: impl IntoIterator, len: usize) -> ArrayRef { ConstantArray::new( @@ -97,8 +99,9 @@ mod tests { } #[rstest] - #[case::empty(vec![])] #[case::one(vec![3])] + #[case::two(vec![3, 7])] + #[case::three(vec![3, 7, 11])] #[case::four(vec![3, 7, 11, 15])] #[case::dense((0..32).map(|value| value * 3).collect())] #[case::sparse((0..32).map(|value| value * 10_000).collect())] @@ -119,6 +122,38 @@ mod tests { Ok(()) } + #[test] + #[cfg(not(codspeed))] + fn test_registered_kernel_executes_through_expression() -> VortexResult<()> { + let mut ctx = crate::array_session().create_execution_ctx(); + let values = [0i32, 1, 2, 3]; + let element = PrimitiveArray::from_iter(values); + let members = [1, 3]; + let contains = element.into_array().apply(&list_contains( + lit(list(members, values.len()) + .as_constant() + .vortex_expect("constant list")), + root(), + ))?; + + let traced = trace_op(|| contains.execute::(&mut ctx))?; + let trace = traced.trace.to_string(); + let applied = trace + .lines() + .filter(|line| { + line.contains("child_execute_parent session[") + && line.contains("slot=1") + && line.contains("parent=vortex.list.contains") + && line.contains("child=vortex.primitive") + }) + .collect::>(); + assert_eq!(applied.len(), 1, "{trace}"); + + let expected = BoolArray::from_iter(values.map(|value| members.contains(&value))); + assert_arrays_eq!(traced.output, expected, &mut ctx); + Ok(()) + } + #[test] fn test_null_needles() -> VortexResult<()> { let mut ctx = crate::array_session().create_execution_ctx(); @@ -136,20 +171,44 @@ mod tests { Ok(()) } - #[test] - fn test_empty_list_ignores_needle_validity() -> VortexResult<()> { + #[rstest] + #[case::mixed( + vec![Some(1), None, Some(3)], + [Some(true), None, Some(true)] + )] + #[case::all_null(vec![None, None], [Some(false), None, Some(false)])] + fn test_nullable_members( + #[case] members: Vec>, + #[case] expected: [Option; 3], + ) -> VortexResult<()> { let mut ctx = crate::array_session().create_execution_ctx(); - let element = PrimitiveArray::from_option_iter([Some(1i32), None, Some(2)]); - let expected = BoolArray::new(BitBuffer::new_unset(3), crate::validity::Validity::AllValid); + let member_dtype = DType::Primitive(I32, Nullability::Nullable); + let list = ConstantArray::new( + Scalar::list( + Arc::new(member_dtype.clone()), + members + .into_iter() + .map(|member| { + member + .map(|value| Scalar::primitive(value, Nullability::Nullable)) + .unwrap_or_else(|| Scalar::null(member_dtype.clone())) + }) + .collect(), + Nullability::NonNullable, + ), + 3, + ) + .into_array(); + let element = PrimitiveArray::from_option_iter([Some(1), None, Some(3)]); let actual = ::list_contains( - &list([], element.len()), + &list, element.as_view(), &mut ctx, )? .vortex_expect("integer constant-list membership is supported"); - assert_arrays_eq!(actual, expected, &mut ctx); + assert_arrays_eq!(actual, BoolArray::from_iter(expected), &mut ctx); Ok(()) } } diff --git a/vortex-array/src/expr/exprs.rs b/vortex-array/src/expr/exprs.rs index fb8bfe227aa..16dcd3256a7 100644 --- a/vortex-array/src/expr/exprs.rs +++ b/vortex-array/src/expr/exprs.rs @@ -1100,6 +1100,8 @@ pub fn bound_dynamic( /// Creates an expression that checks if a value is contained in a list. /// /// Returns a boolean array indicating whether the value appears in each list. +/// A null list produces null. An empty list produces false, including for a null value. +/// A null value produces null for a nonempty list. Null list members do not match any value. /// /// ```rust /// # use vortex_array::expr::{list_contains, lit, root}; diff --git a/vortex-array/src/scalar_fn/fns/list_contains/integer_membership.rs b/vortex-array/src/scalar_fn/fns/list_contains/integer_membership.rs index 5a6d17eb2a2..b5005c72aa7 100644 --- a/vortex-array/src/scalar_fn/fns/list_contains/integer_membership.rs +++ b/vortex-array/src/scalar_fn/fns/list_contains/integer_membership.rs @@ -174,11 +174,35 @@ mod tests { use super::IntegerMembership; #[test] - fn small_unsorted_set_contains_members() { - let membership = IntegerMembership::new(vec![7i32, 3]); + fn normalizes_large_unsorted_duplicates() { + let membership = IntegerMembership::new(vec![7i32, 3, 7, 1, 9, 3, 1]); + assert_eq!(membership.members(), &[1, 3, 7, 9]); + assert!(membership.contains(1)); assert!(membership.contains(3)); assert!(membership.contains(7)); assert!(!membership.contains(5)); } + + #[test] + fn dense_table_span_boundary() { + let at_limit = IntegerMembership::new(vec![0i32, 1, 2, 3, 4_095]); + let above_limit = IntegerMembership::new(vec![0i32, 1, 2, 3, 4_096]); + + assert!(at_limit.uses_dense_table()); + assert!(!above_limit.uses_dense_table()); + } + + #[test] + fn integer_extremes_do_not_overflow() { + let signed = IntegerMembership::new(vec![i64::MAX, 0, i64::MIN, -1, 1]); + assert!(signed.contains(i64::MIN)); + assert!(signed.contains(i64::MAX)); + assert!(!signed.uses_dense_table()); + + let unsigned = IntegerMembership::new(vec![u64::MAX, 0, 1, 2, 3]); + assert!(unsigned.contains(0)); + assert!(unsigned.contains(u64::MAX)); + assert!(!unsigned.uses_dense_table()); + } } diff --git a/vortex-array/src/scalar_fn/fns/list_contains/kernel.rs b/vortex-array/src/scalar_fn/fns/list_contains/kernel.rs index 563600bfeee..fc50c8e70ac 100644 --- a/vortex-array/src/scalar_fn/fns/list_contains/kernel.rs +++ b/vortex-array/src/scalar_fn/fns/list_contains/kernel.rs @@ -6,16 +6,42 @@ use vortex_error::VortexResult; use crate::ArrayRef; use crate::ExecutionCtx; +use crate::IntoArray; use crate::array::ArrayView; use crate::array::VTable; +use crate::arrays::ConstantArray; use crate::arrays::ScalarFn; use crate::arrays::scalar_fn::ExactScalarFn; use crate::arrays::scalar_fn::ScalarFnArrayExt; use crate::arrays::scalar_fn::ScalarFnArrayView; +use crate::dtype::DType; use crate::kernel::ExecuteParentKernel; use crate::optimizer::rules::ArrayParentReduceRule; +use crate::scalar::Scalar; use crate::scalar_fn::fns::list_contains::ListContains as ListContainsExpr; +fn constant_list_result( + list: &ArrayRef, + element_len: usize, + element_nullability: crate::dtype::Nullability, +) -> Option { + let list_scalar = list.as_constant()?; + let DType::List(_, list_nullability) = list.dtype() else { + return None; + }; + let nullability = *list_nullability | element_nullability; + + match list_scalar.as_list().elements() { + None => Some( + ConstantArray::new(Scalar::null(DType::Bool(nullability)), element_len).into_array(), + ), + Some(elements) if elements.is_empty() => { + Some(ConstantArray::new(Scalar::bool(false, nullability), element_len).into_array()) + } + Some(_) => None, + } +} + /// Check list-contains without reading buffers (metadata-only). /// /// This trait dispatches on the **element** (needle) child at index 1 of the `ListContains` @@ -25,6 +51,8 @@ use crate::scalar_fn::fns::list_contains::ListContains as ListContainsExpr; /// A future `ListContainsListReduce` could dispatch on the list side (child 0) for encodings /// with specialized list representations. /// +/// The parent adaptor resolves null and empty constant lists before delegation. +/// /// Return `None` if the operation cannot be resolved from metadata alone. pub trait ListContainsElementReduce: VTable { fn list_contains( @@ -38,6 +66,8 @@ pub trait ListContainsElementReduce: VTable { /// Like [`ListContainsElementReduce`], this dispatches on the **element** (needle) child at /// index 1. Unlike the reduce variant, implementations may read and execute on buffers via /// the provided [`ExecutionCtx`]. +/// +/// The parent adaptor resolves null and empty constant lists before delegation. pub trait ListContainsElementKernel: VTable { fn list_contains( list: &ArrayRef, @@ -70,6 +100,9 @@ where .as_opt::() .vortex_expect("ExactScalarFn matcher confirmed ScalarFnArray"); let list = scalar_fn_array.get_child(0); + if let Some(result) = constant_list_result(list, array.len(), array.dtype().nullability()) { + return Ok(Some(result)); + } ::list_contains(list, array) } } @@ -99,6 +132,9 @@ where .as_opt::() .vortex_expect("ExactScalarFn matcher confirmed ScalarFnArray"); let list = scalar_fn_array.get_child(0); + if let Some(result) = constant_list_result(list, array.len(), array.dtype().nullability()) { + return Ok(Some(result)); + } ::list_contains(list, array, ctx) } } diff --git a/vortex-array/src/scalar_fn/fns/list_contains/mod.rs b/vortex-array/src/scalar_fn/fns/list_contains/mod.rs index dcce13bce38..470b81964a4 100644 --- a/vortex-array/src/scalar_fn/fns/list_contains/mod.rs +++ b/vortex-array/src/scalar_fn/fns/list_contains/mod.rs @@ -27,7 +27,6 @@ use crate::arrays::BoolArray; use crate::arrays::Constant; use crate::arrays::ConstantArray; use crate::arrays::ListViewArray; -use crate::arrays::Primitive; use crate::arrays::PrimitiveArray; use crate::arrays::ScalarFnArray; use crate::arrays::bool::BoolArrayExt; @@ -60,7 +59,8 @@ impl ListContains { /// /// # Errors /// - /// Returns an error if the children have different lengths or `list` is not a list array. + /// Returns an error if the children have different lengths, `list` is not a list array, or + /// the list member type differs from the needle type. pub fn try_new(list: ArrayRef, needle: ArrayRef) -> VortexResult { ScalarFnArray::try_new(ListContains.bind(EmptyOptions), vec![list, needle]) } @@ -104,16 +104,17 @@ impl ScalarFnVTable for ListContains { let list_dtype = &arg_dtypes[0]; let needle_dtype = &arg_dtypes[1]; - let nullability = match list_dtype { - DType::List(_, list_nullability) => list_nullability, - _ => { - vortex_bail!( - "First argument to ListContains must be a List, got {:?}", - list_dtype - ); - } + let DType::List(member_dtype, list_nullability) = list_dtype else { + vortex_bail!("First argument to ListContains must be a List, got {list_dtype}"); + }; + if !member_dtype.eq_ignore_nullability(needle_dtype) { + vortex_bail!( + "Element type {} of list does not match search value {}", + member_dtype, + needle_dtype + ); } - .bitor(needle_dtype.nullability()); + let nullability = list_nullability.bitor(needle_dtype.nullability()); Ok(DType::Bool(nullability)) } @@ -185,17 +186,6 @@ fn compute_list_contains( ); } - if matches!(value.dtype(), DType::Primitive(ptype, _) if ptype.is_int()) - && array.as_constant().is_some() - { - let value = value.clone().execute::(ctx)?; - if let Some(result) = - ::list_contains(array, value.as_view(), ctx)? - { - return Ok(result); - } - } - if array.all_invalid(ctx)? { return Ok(ConstantArray::new( Scalar::null(DType::Bool(Nullability::Nullable)), @@ -462,6 +452,10 @@ mod tests { use crate::IntoArray; use crate::VortexSessionExecute; use crate::array_session; + #[cfg(not(codspeed))] + use crate::arrays::Dict; + #[cfg(not(codspeed))] + use crate::arrays::DictArray; use crate::arrays::ListArray; use crate::arrays::VarBinArray; use crate::assert_arrays_eq; @@ -483,6 +477,7 @@ mod tests { use crate::scalar::Scalar; use crate::scalar_fn::fns::list_contains::BoolArray; use crate::scalar_fn::fns::list_contains::ConstantArray; + use crate::scalar_fn::fns::list_contains::ListContains; use crate::scalar_fn::fns::list_contains::ListViewArray; use crate::scalar_fn::fns::list_contains::PrimitiveArray; use crate::stats::StatsSession; @@ -635,6 +630,54 @@ mod tests { ); } + #[test] + fn test_return_type_rejects_mismatched_member_type() { + let list = ConstantArray::new( + Scalar::list( + Arc::new(DType::Primitive(I32, Nullability::NonNullable)), + vec![], + Nullability::NonNullable, + ), + 1, + ) + .into_array(); + let needle = + ConstantArray::new(Scalar::utf8("needle", Nullability::NonNullable), 1).into_array(); + + let error = ListContains::try_new(list, needle).unwrap_err(); + + assert!( + error + .to_string() + .contains("Element type i32 of list does not match search value utf8") + ); + } + + #[test] + #[cfg(not(codspeed))] + fn test_dictionary_needles_preserve_dictionary_pushdown() -> VortexResult<()> { + let mut ctx = array_session().create_execution_ctx(); + let values = PrimitiveArray::from_iter([1i32, 2, 3]).into_array(); + let codes = PrimitiveArray::from_iter([0u8, 1, 2, 0]).into_array(); + let needles = DictArray::try_new(codes, values)?.into_array(); + let list = Scalar::list( + Arc::new(DType::Primitive(I32, Nullability::NonNullable)), + vec![1.into(), 3.into()], + Nullability::NonNullable, + ); + let contains = needles.apply(&list_contains(lit(list), root()))?; + + assert!(contains.is::()); + let actual = contains.execute::(&mut ctx)?; + + assert_arrays_eq!( + actual, + BoolArray::from_iter([true, false, true, true]), + &mut ctx + ); + Ok(()) + } + #[test] pub fn list_falsification() -> VortexResult<()> { let expr = list_contains( @@ -719,6 +762,40 @@ mod tests { ); } + #[rstest] + #[case::null_list(true, vec![], Some(1), None)] + #[case::empty_list_null_needle(false, vec![], None, Some(false))] + #[case::nonempty_list_null_needle(false, vec![1], None, None)] + fn test_constant_scalar_null_semantics( + #[case] null_list: bool, + #[case] members: Vec, + #[case] needle: Option, + #[case] expected: Option, + ) { + let member_dtype = DType::Primitive(I32, Nullability::NonNullable); + let list_dtype = DType::List(Arc::new(member_dtype.clone()), Nullability::Nullable); + let list = if null_list { + Scalar::null(list_dtype) + } else { + Scalar::list( + Arc::new(member_dtype), + members.into_iter().map(Scalar::from).collect(), + Nullability::Nullable, + ) + }; + let needle = needle + .map(|value| Scalar::primitive(value, Nullability::Nullable)) + .unwrap_or_else(|| Scalar::null(DType::Primitive(I32, Nullability::Nullable))); + let expected = expected + .map(|value| Scalar::bool(value, Nullability::Nullable)) + .unwrap_or_else(|| Scalar::null(DType::Bool(Nullability::Nullable))); + + assert_eq!( + super::compute_contains_scalar(&list, &needle).unwrap(), + expected + ); + } + // -- Tests migrated from compute/list_contains.rs -- fn nonnull_strings(values: Vec>) -> ArrayRef { @@ -825,57 +902,42 @@ mod tests { assert_arrays_eq!(result, expected, &mut ctx); } - #[test] - fn test_constant_list() { - let mut ctx = array_session().create_execution_ctx(); - let list_array = ConstantArray::new( - Scalar::list( - Arc::new(DType::Primitive(I32, Nullability::NonNullable)), - vec![1i32.into(), 2i32.into(), 3i32.into()], - Nullability::NonNullable, - ), - 2, - ) - .into_array(); - - let expr = list_contains(root(), lit(2i32)); - let contains = list_array.apply(&expr).unwrap(); - let expected = BoolArray::from_iter([true, true]); - assert_arrays_eq!(contains, expected, &mut ctx); - } - #[rstest] #[case::empty( - Vec::>::new(), + Vec::>::new(), [Some(false), Some(false), Some(false)] )] #[case::nonempty( - vec![Some(1), Some(3)], + vec![Some("a"), Some("c")], [Some(true), None, Some(false)] )] #[case::all_null( vec![None, None], [Some(false), None, Some(false)] )] - fn test_constant_list_nullable_needles( - #[case] members: Vec>, + fn test_constant_string_list_nullable_needles( + #[case] members: Vec>, #[case] expected: [Option; 3], ) { let mut ctx = array_session().create_execution_ctx(); - let member_dtype = DType::Primitive(I32, Nullability::Nullable); + let member_dtype = DType::Utf8(Nullability::Nullable); let list = Scalar::list( Arc::new(member_dtype.clone()), members .into_iter() .map(|member| { member - .map(|value| Scalar::primitive(value, Nullability::Nullable)) + .map(|value| Scalar::utf8(value, Nullability::Nullable)) .unwrap_or_else(|| Scalar::null(member_dtype.clone())) }) .collect(), Nullability::NonNullable, ); - let needles = PrimitiveArray::from_option_iter([Some(1), None, Some(2)]).into_array(); + let needles = VarBinArray::from_iter( + [Some("a"), None, Some("b")], + DType::Utf8(Nullability::Nullable), + ) + .into_array(); let result = needles.apply(&list_contains(lit(list), root())).unwrap(); let expected = BoolArray::from_iter(expected); diff --git a/vortex-datafusion/src/convert/exprs.rs b/vortex-datafusion/src/convert/exprs.rs index 2ab975ecfd7..f62567c25c9 100644 --- a/vortex-datafusion/src/convert/exprs.rs +++ b/vortex-datafusion/src/convert/exprs.rs @@ -373,14 +373,23 @@ impl ExpressionConvertor for DefaultExpressionConvertor { if let Some(in_list) = df.downcast_ref::() { let value = self.convert(in_list.expr().as_ref())?; + if in_list.is_empty() { + return Err(exec_datafusion_err!("Cannot convert an empty IN list")); + } let list_elements: Vec<_> = in_list .list() .iter() .map(|e| { if let Some(lit) = e.downcast_ref::() { - Ok(scalar_from_df(lit.value(), &self.session)) + if lit.value().is_null() { + Err(exec_datafusion_err!( + "Cannot push down an IN list that contains NULL" + )) + } else { + Ok(scalar_from_df(lit.value(), &self.session)) + } } else { - Err(exec_datafusion_err!("Failed to cast sub-expression")) + Err(exec_datafusion_err!("IN list member is not a literal")) } }) .try_collect()?; @@ -433,6 +442,19 @@ impl ExpressionConvertor for DefaultExpressionConvertor { return Ok(TreeNodeRecursion::Stop); } + if let Some(in_list) = node.downcast_ref::() + && !can_in_list_be_pushed_down(in_list, input_schema) + { + scan_projection.extend( + collect_columns(node) + .into_iter() + .map(|c| (c.name().to_string(), get_item(c.name(), root()))), + ); + + leftover_projection.push(projection_expr.clone()); + return Ok(TreeNodeRecursion::Stop); + } + // DataFusion assumes different decimal types can be coerced. // Vortex expects a perfect match so we don't push it down. if let Some(binary_expr) = node.downcast_ref::() @@ -555,11 +577,7 @@ fn can_be_pushed_down_impl(expr: &Arc, schema: &Schema) -> boo } else if let Some(is_not_null) = expr.downcast_ref::() { can_be_pushed_down_impl(is_not_null.arg(), schema) } else if let Some(in_list) = expr.downcast_ref::() { - can_be_pushed_down_impl(in_list.expr(), schema) - && in_list - .list() - .iter() - .all(|e| can_be_pushed_down_impl(e, schema)) + can_in_list_be_pushed_down(in_list, schema) } else if let Some(scalar_fn) = expr.downcast_ref::() { can_scalar_fn_be_pushed_down(scalar_fn, schema) } else if let Some(case_expr) = expr.downcast_ref::() { @@ -570,6 +588,17 @@ fn can_be_pushed_down_impl(expr: &Arc, schema: &Schema) -> boo } } +fn can_in_list_be_pushed_down(in_list: &df_expr::InListExpr, schema: &Schema) -> bool { + can_be_pushed_down_impl(in_list.expr(), schema) + && !in_list.is_empty() + && in_list.list().iter().all(|expr| { + expr.downcast_ref::() + .is_some_and(|literal| { + !literal.value().is_null() && supported_data_types(&literal.value().data_type()) + }) + }) +} + /// Checks if an expression type is one that convert() can handle. /// This is less restrictive than can_be_pushed_down since it only checks /// expression types, not data type support. @@ -872,6 +901,53 @@ mod tests { assert_snapshot!(result.display_tree().to_string(), @"vortex.literal(42i32)"); } + #[rstest] + #[case::in_list(false)] + #[case::not_in_list(true)] + fn test_null_in_list_is_not_pushed_down(test_schema: Schema, #[case] negated: bool) { + let value = Arc::new(df_expr::Column::new("id", 0)) as Arc; + let list = vec![ + Arc::new(df_expr::Literal::new(ScalarValue::Int32(Some(1)))) as Arc, + Arc::new(df_expr::Literal::new(ScalarValue::Int32(None))) as Arc, + ]; + let expr = + Arc::new(df_expr::InListExpr::try_new(value, list, negated, &test_schema).unwrap()) + as Arc; + let convertor = DefaultExpressionConvertor::default(); + + assert!(!convertor.can_be_pushed_down(&expr, &test_schema)); + assert!( + convertor + .convert(expr.as_ref()) + .unwrap_err() + .to_string() + .contains("IN list that contains NULL") + ); + } + + #[test] + fn test_expr_from_df_in_list() { + let schema = test_schema(); + let value = Arc::new(df_expr::Column::new("id", 0)) as Arc; + let list = [1, 3] + .map(|value| { + Arc::new(df_expr::Literal::new(ScalarValue::Int32(Some(value)))) + as Arc + }) + .to_vec(); + let expr = Arc::new(df_expr::InListExpr::try_new(value, list, false, &schema).unwrap()) + as Arc; + let convertor = DefaultExpressionConvertor::default(); + + assert!(convertor.can_be_pushed_down(&expr, &schema)); + assert_snapshot!(convertor.convert(expr.as_ref()).unwrap().display_tree().to_string(), @r" + vortex.list.contains() + ├── list: vortex.literal([1i32, 3i32]) + └── needle: vortex.get_item(id) + └── input: vortex.root() + "); + } + #[test] fn test_expr_from_df_binary() { let left = Arc::new(df_expr::Column::new("left", 0)) as Arc; diff --git a/vortex-datafusion/src/persistent/tests.rs b/vortex-datafusion/src/persistent/tests.rs index 35dd745461d..cb6dc818b8e 100644 --- a/vortex-datafusion/src/persistent/tests.rs +++ b/vortex-datafusion/src/persistent/tests.rs @@ -235,6 +235,46 @@ async fn test_octet_length_pushdown() -> anyhow::Result<()> { Ok(()) } +#[tokio::test] +async fn test_nullable_in_projection_falls_back() -> anyhow::Result<()> { + let ctx = TestSessionContext::new(true); + + ctx.session + .sql( + "CREATE EXTERNAL TABLE nullable_in (id INT) \ + STORED AS vortex LOCATION '/nullable_in/'", + ) + .await?; + ctx.session + .sql("INSERT INTO nullable_in VALUES (1), (2), (NULL)") + .await? + .collect() + .await?; + + let result = ctx + .session + .sql( + "SELECT id, id IN (1, NULL) AS in_result, \ + id NOT IN (1, NULL) AS not_in_result \ + FROM nullable_in ORDER BY id NULLS LAST", + ) + .await? + .collect() + .await?; + + assert_snapshot!(pretty_format_batches(&result)?, @r" + +----+-----------+---------------+ + | id | in_result | not_in_result | + +----+-----------+---------------+ + | 1 | true | false | + | 2 | | | + | | | | + +----+-----------+---------------+ + "); + + Ok(()) +} + #[tokio::test] async fn create_table_ordered_by() -> anyhow::Result<()> { let ctx = TestSessionContext::default(); diff --git a/vortex-duckdb/src/convert/expr.rs b/vortex-duckdb/src/convert/expr.rs index 51aeebebb6d..814aa029f90 100644 --- a/vortex-duckdb/src/convert/expr.rs +++ b/vortex-duckdb/src/convert/expr.rs @@ -27,7 +27,6 @@ use vortex::error::VortexExpect; use vortex::error::VortexResult; use vortex::error::vortex_bail; use vortex::error::vortex_ensure; -use vortex::error::vortex_err; use vortex::expr::Expression; use vortex::expr::and_collect; use vortex::expr::byte_length; @@ -52,7 +51,6 @@ use vortex::scalar_fn::fns::between::StrictComparison; use vortex::scalar_fn::fns::binary::Binary; use vortex::scalar_fn::fns::like::Like; use vortex::scalar_fn::fns::like::LikeOptions; -use vortex::scalar_fn::fns::literal::Literal; use vortex::scalar_fn::fns::operators::Operator; use vortex_spatial::extension::LineString; use vortex_spatial::extension::MultiLineString; @@ -443,7 +441,27 @@ pub fn can_push_expression(value: &duckdb::ExpressionRef) -> bool { ) { return false; } - op.children().all(can_push_expression) + if matches!( + op.op, + DUCKDB_VX_EXPR_TYPE::DUCKDB_VX_EXPR_TYPE_COMPARE_IN + | DUCKDB_VX_EXPR_TYPE::DUCKDB_VX_EXPR_TYPE_COMPARE_NOT_IN + ) { + let mut children = op.children(); + let Some(element) = children.next() else { + return false; + }; + can_push_expression(element) + && children.all(|child| { + matches!( + child.as_class(), + Some(BoundConstant(constant)) + if Scalar::try_from(constant.value) + .is_ok_and(|scalar| !scalar.is_null()) + ) + }) + } else { + op.children().all(can_push_expression) + } } ExpressionClass::BoundAggregate(_) => false, } @@ -717,7 +735,9 @@ fn try_from_compare_in( ) -> VortexResult> { // First child is element, rest form the list. let children: Vec<_> = operator.children().collect(); - assert!(children.len() >= 2); + if children.len() < 2 { + return Ok(None); + } let Some(element) = try_from_expression_inner(children[0], ctx)? else { return Ok(None); }; @@ -725,16 +745,14 @@ fn try_from_compare_in( let Some(list_elements) = children .iter() .skip(1) - .map(|c| { - let Some(value) = try_from_expression_inner(c, ctx)? else { + .map(|child| { + let Some(BoundConstant(constant)) = child.as_class() else { return Ok(None); }; - Ok(Some( - value - .as_opt::() - .ok_or_else(|| vortex_err!("cannot have a non literal in a in_list"))? - .clone(), - )) + if constant.value.is_null() { + return Ok(None); + } + Ok(Some(Scalar::try_from(constant.value)?)) }) .collect::>>>()? else { diff --git a/vortex-duckdb/src/duckdb/value.rs b/vortex-duckdb/src/duckdb/value.rs index 8ea5b253a01..b21d38f7c31 100644 --- a/vortex-duckdb/src/duckdb/value.rs +++ b/vortex-duckdb/src/duckdb/value.rs @@ -28,6 +28,10 @@ use crate::lifetime_wrapper; lifetime_wrapper!(Value, cpp::duckdb_value, cpp::duckdb_destroy_value); impl ValueRef { + pub fn is_null(&self) -> bool { + unsafe { cpp::duckdb_is_null_value(self.as_ptr()) } + } + pub fn logical_type(&self) -> &LogicalTypeRef { unsafe { LogicalType::borrow(cpp::duckdb_get_value_type(self.as_ptr())) } } @@ -41,7 +45,7 @@ impl ValueRef { /// Extracts the value from the DuckDB `Value` into a `ExtractedValue`. pub fn extract(&self) -> ExtractedValue { - if unsafe { cpp::duckdb_is_null_value(self.as_ptr()) } { + if self.is_null() { return ExtractedValue::Null; } match self.logical_type().as_type_id() { diff --git a/vortex-duckdb/src/e2e_test/vortex_scan_test.rs b/vortex-duckdb/src/e2e_test/vortex_scan_test.rs index 0876be1ca4c..49b97fa2f80 100644 --- a/vortex-duckdb/src/e2e_test/vortex_scan_test.rs +++ b/vortex-duckdb/src/e2e_test/vortex_scan_test.rs @@ -281,6 +281,20 @@ fn test_issue_5927_not_in_does_not_panic() { assert_eq!(sum, -4); } +#[test] +fn test_not_in_with_null_is_not_pushed_down() { + let file = RUNTIME.block_on(async { + let numbers = buffer![1i32, 42, 100, -5, 0]; + write_single_column_vortex_file("number", numbers).await + }); + let count: i64 = scan_vortex_file_single_row::( + file, + "SELECT COUNT(*) FROM ? WHERE number NOT IN (42, NULL)", + 0, + ); + assert_eq!(count, 0); +} + #[test] fn test_vortex_scan_floats() { let file = RUNTIME.block_on(async { From 289f55e8fbfd353c217b07d05b295d3cfa5e11c6 Mon Sep 17 00:00:00 2001 From: Will Manning Date: Sat, 29 Aug 2026 17:57:26 -0400 Subject: [PATCH 4/8] perf(list_contains): Specialize integer membership Signed-off-by: Will Manning --- Cargo.lock | 1 + encodings/fastlanes/Cargo.toml | 1 + .../benches/bitpacking_list_contains.rs | 424 ++++-------------- .../src/bitpacking/compute/compare_fused.rs | 56 ++- .../bitpacking/compute/list_contains/mod.rs | 209 +++++---- .../bitpacking/compute/list_contains/tests.rs | 84 +++- vortex-array/Cargo.toml | 4 + vortex-array/benches/list_contains.rs | 195 ++++++++ .../arrays/primitive/compute/list_contains.rs | 202 +++++++-- .../src/arrays/primitive/compute/mod.rs | 2 + vortex-array/src/arrays/primitive/mod.rs | 4 + .../fns/list_contains/integer_membership.rs | 185 +++----- .../src/scalar_fn/fns/list_contains/mod.rs | 143 +++--- vortex-datafusion/src/convert/exprs.rs | 90 +--- vortex-datafusion/src/persistent/tests.rs | 40 -- vortex-duckdb/src/convert/expr.rs | 42 +- vortex-duckdb/src/duckdb/value.rs | 6 +- .../src/e2e_test/vortex_scan_test.rs | 14 - 18 files changed, 835 insertions(+), 867 deletions(-) create mode 100644 vortex-array/benches/list_contains.rs diff --git a/Cargo.lock b/Cargo.lock index 2448cb9aff8..70d7eb4f883 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -11009,6 +11009,7 @@ dependencies = [ "rstest", "vortex-alp", "vortex-array", + "vortex-bench-support", "vortex-buffer", "vortex-error", "vortex-fastlanes", diff --git a/encodings/fastlanes/Cargo.toml b/encodings/fastlanes/Cargo.toml index 1127bf37bed..9fb317ec8c8 100644 --- a/encodings/fastlanes/Cargo.toml +++ b/encodings/fastlanes/Cargo.toml @@ -39,6 +39,7 @@ rand = { workspace = true } rstest = { workspace = true } vortex-alp = { path = "../alp" } vortex-array = { workspace = true, features = ["_test-harness"] } +vortex-bench-support = { workspace = true } vortex-fastlanes = { path = ".", features = ["_test-harness"] } [features] diff --git a/encodings/fastlanes/benches/bitpacking_list_contains.rs b/encodings/fastlanes/benches/bitpacking_list_contains.rs index 2770bd24106..72b4c1beada 100644 --- a/encodings/fastlanes/benches/bitpacking_list_contains.rs +++ b/encodings/fastlanes/benches/bitpacking_list_contains.rs @@ -1,15 +1,17 @@ // SPDX-License-Identifier: Apache-2.0 // SPDX-FileCopyrightText: Copyright the Vortex contributors -//! Compares compressed list membership with two explicit fallback paths. +//! Measures compressed constant-list membership. //! -//! The decode-once path measures the lower bound for materializing a Primitive array before -//! membership evaluation. The old-generic path freezes the former balanced equality-expression -//! fallback. Primitive cases isolate the benefit of the prepared integer membership set. +//! FastLanes evaluates constant lists with at most four distinct non-null members during unpacking. +//! Mid-size lists use repeated packed comparisons. Larger lists decode once at a threshold that +//! depends on the physical integer width and array length. Every path runs on each real CPU feature +//! leg in CodSpeed. +//! To recalculate the thresholds, temporarily replace `min_decode_source_members` with a constant. +//! Return `usize::MAX` to force repeated comparisons. Return `5` to force decode-once. //! //! Run with `cargo bench -p vortex-fastlanes --bench bitpacking_list_contains`. -#![expect(clippy::cast_possible_truncation)] #![expect(clippy::unwrap_used)] use std::fmt::Display; @@ -19,15 +21,12 @@ use std::sync::Arc; use divan::Bencher; use divan::counter::ItemsCount; -use vortex_array::ArrayRef; use vortex_array::IntoArray; use vortex_array::VortexSessionExecute; use vortex_array::array_session; use vortex_array::arrays::BoolArray; -use vortex_array::arrays::ConstantArray; -use vortex_array::arrays::Primitive; use vortex_array::arrays::PrimitiveArray; -use vortex_array::builtins::ArrayBuiltins; +use vortex_array::assert_arrays_eq; use vortex_array::dtype::DType; use vortex_array::dtype::IntegerPType; use vortex_array::dtype::Nullability; @@ -36,13 +35,9 @@ use vortex_array::expr::list_contains; use vortex_array::expr::lit; use vortex_array::expr::root; use vortex_array::scalar::Scalar; -use vortex_array::scalar_fn::fns::binary::Binary; -use vortex_array::scalar_fn::fns::list_contains::ListContainsElementKernel; -use vortex_array::scalar_fn::fns::operators::Operator; use vortex_array::validity::Validity; use vortex_buffer::Alignment; use vortex_buffer::BufferMut; -use vortex_error::VortexExpect; use vortex_fastlanes::BitPacked; use vortex_fastlanes::BitPackedArray; use vortex_fastlanes::BitPackedData; @@ -56,32 +51,27 @@ trait BenchInt: IntegerPType + Copy + Into { fn from_counter(value: u64) -> Self; } -macro_rules! impl_bench_int { - ($($T:ty),+) => { - $(impl BenchInt for $T { - fn from_counter(value: u64) -> Self { - value as $T - } - })+ - }; +impl BenchInt for u8 { + fn from_counter(value: u64) -> Self { + Self::try_from(value).unwrap() + } } -impl_bench_int!(u8, u16, u32, u64); +impl BenchInt for u16 { + fn from_counter(value: u64) -> Self { + Self::try_from(value).unwrap() + } +} -#[derive(Clone, Copy)] -enum MemberSpec { - Explicit(&'static [u64]), - Stride { count: usize, stride: u64 }, +impl BenchInt for u32 { + fn from_counter(value: u64) -> Self { + Self::try_from(value).unwrap() + } } -impl MemberSpec { - fn values(self) -> Vec { - match self { - Self::Explicit(values) => values.to_vec(), - Self::Stride { count, stride } => { - (0..count).map(|index| index as u64 * stride).collect() - } - } +impl BenchInt for u64 { + fn from_counter(value: u64) -> Self { + value } } @@ -91,130 +81,62 @@ struct PackedCase { ptype: PType, bit_width: u8, len: usize, - members: MemberSpec, - hit_percent: u8, + member_count: usize, + member_stride: u64, } impl Display for PackedCase { fn fmt(&self, formatter: &mut Formatter<'_>) -> std::fmt::Result { write!( formatter, - "{}_{}_w{}_n{}_hit{}", - self.name, self.ptype, self.bit_width, self.len, self.hit_percent + "{}_{}_w{}_n{}", + self.name, self.ptype, self.bit_width, self.len ) } } -const FOUR_BOUNDARY_MEMBERS: &[u64] = &[0, 1, 2, 4_095]; -const FIVE_DENSE_BOUNDARY_MEMBERS: &[u64] = &[0, 1, 2, 3, 4_095]; -const FIVE_SPARSE_BOUNDARY_MEMBERS: &[u64] = &[0, 1, 2, 3, 4_096]; - -const PACKED_CASES: &[PackedCase] = &[ - PackedCase { - name: "short_direct_m4", - ptype: PType::U32, - bit_width: 10, - len: 1_024, - members: MemberSpec::Stride { - count: 4, - stride: 2, - }, - hit_percent: 50, - }, - PackedCase { - name: "four_member_span4096", - ptype: PType::U32, - bit_width: 13, - len: 2_048, - members: MemberSpec::Explicit(FOUR_BOUNDARY_MEMBERS), - hit_percent: 50, - }, - PackedCase { - name: "before_length_gate_m5_span4096", - ptype: PType::U32, - bit_width: 13, - len: 2_047, - members: MemberSpec::Explicit(FIVE_DENSE_BOUNDARY_MEMBERS), - hit_percent: 50, - }, - PackedCase { - name: "at_length_gate_m5_span4096", - ptype: PType::U32, - bit_width: 13, - len: 2_048, - members: MemberSpec::Explicit(FIVE_DENSE_BOUNDARY_MEMBERS), - hit_percent: 50, - }, - PackedCase { - name: "above_span_gate_m5_span4097", - ptype: PType::U32, - bit_width: 13, - len: 2_048, - members: MemberSpec::Explicit(FIVE_SPARSE_BOUNDARY_MEMBERS), - hit_percent: 50, - }, - PackedCase { - name: "long_dense_m32", - ptype: PType::U32, - bit_width: 10, - len: 65_536, - members: MemberSpec::Stride { - count: 32, - stride: 2, - }, - hit_percent: 50, - }, - PackedCase { - name: "long_sparse_m32", - ptype: PType::U64, - bit_width: 40, - len: 65_536, - members: MemberSpec::Stride { - count: 32, - stride: 10_000, - }, - hit_percent: 50, - }, - PackedCase { - name: "zero_hit_m8", - ptype: PType::U8, - bit_width: 6, - len: 65_536, - members: MemberSpec::Stride { - count: 8, - stride: 2, - }, - hit_percent: 0, - }, - PackedCase { - name: "full_hit_m8", - ptype: PType::U16, - bit_width: 12, - len: 65_536, - members: MemberSpec::Stride { - count: 8, - stride: 2, - }, - hit_percent: 100, - }, +const fn strided_case( + name: &'static str, + ptype: PType, + bit_width: u8, + len: usize, + count: usize, + stride: u64, +) -> PackedCase { PackedCase { - name: "wide_packed_m8", - ptype: PType::U32, - bit_width: 31, - len: 65_536, - members: MemberSpec::Stride { - count: 8, - stride: 2, - }, - hit_percent: 50, - }, -]; + name, + ptype, + bit_width, + len, + member_count: count, + member_stride: stride, + } +} -const OLD_GENERIC_CASES: &[PackedCase] = &[ - PACKED_CASES[0], - PACKED_CASES[3], - PACKED_CASES[5], - PACKED_CASES[6], +const PACKED_CASES: &[PackedCase] = &[ + strided_case("direct_u8_m4", PType::U8, 6, 65_536, 4, 2), + strided_case("direct_u16_m4", PType::U16, 12, 65_536, 4, 2), + strided_case("direct_u32_m4", PType::U32, 20, 65_536, 4, 2), + strided_case("direct_u64_m4", PType::U64, 40, 65_536, 4, 2), + strided_case("generic_u8_m29", PType::U8, 6, 65_536, 29, 2), + strided_case("decode_u8_m30", PType::U8, 6, 65_536, 30, 2), + strided_case("generic_u16_m24", PType::U16, 8, 65_536, 24, 2), + strided_case("decode_u16_m25", PType::U16, 8, 65_536, 25, 2), + strided_case("generic_u32_m12", PType::U32, 8, 65_536, 12, 2), + strided_case("decode_u32_m13", PType::U32, 8, 65_536, 13, 2), + strided_case("decode_u64_m5", PType::U64, 40, 65_536, 5, 2), + strided_case("short_direct_u32_m4", PType::U32, 10, 1_024, 4, 2), + strided_case("short_generic_u8_m9", PType::U8, 6, 8_192, 9, 2), + strided_case("short_decode_u8_m10", PType::U8, 6, 8_192, 10, 2), + strided_case("short_generic_u16_m9", PType::U16, 8, 8_192, 9, 2), + strided_case("short_decode_u16_m10", PType::U16, 8, 8_192, 10, 2), + strided_case("short_generic_u32_m10", PType::U32, 8, 16_384, 10, 2), + strided_case("short_decode_u32_m11", PType::U32, 8, 16_384, 11, 2), + strided_case("longer_generic_u8_m10", PType::U8, 6, 16_384, 10, 2), + strided_case("longer_generic_u16_m10", PType::U16, 8, 16_384, 10, 2), + strided_case("longer_generic_u32_m11", PType::U32, 8, 32_768, 11, 2), + strided_case("short_direct_u64_m4", PType::U64, 8, 8_192, 4, 2), + strided_case("short_decode_u64_m5", PType::U64, 8, 8_192, 5, 2), ]; fn page_aligned(array: BitPackedArray) -> BitPackedArray { @@ -234,17 +156,19 @@ fn page_aligned(array: BitPackedArray) -> BitPackedArray { fn generated_values(case: PackedCase, members: &[u64]) -> Vec { let domain_size = 1u64 << case.bit_width; + let mut state = 0x9E37_79B9_7F4A_7C15u64; (0..case.len) - .map(|index| { - let is_hit = match case.hit_percent { - 0 => false, - 100 => true, - percent => index % 100 < usize::from(percent), - }; + .map(|_| { + state = state + .wrapping_mul(6_364_136_223_846_793_005) + .wrapping_add(1_442_695_040_888_963_407); + let is_hit = (state >> 32).is_multiple_of(2); if is_hit { - members[index % members.len()] + let member_index = + usize::try_from(state % u64::try_from(members.len()).unwrap()).unwrap(); + members[member_index] } else { - let mut candidate = (index as u64 * 17 + 11) % domain_size; + let mut candidate = state.rotate_left(17) % domain_size; while members.contains(&candidate) { candidate = (candidate + 1) % domain_size; } @@ -265,15 +189,18 @@ fn list_scalar(members: &[u64]) -> Scalar { ) } -fn packed_input(case: PackedCase) -> (BitPackedArray, Scalar, VortexSession) { +fn packed_input( + case: PackedCase, +) -> (BitPackedArray, Scalar, BoolArray, VortexSession) { let session = array_session(); vortex_fastlanes::initialize(&session); let mut ctx = session.create_execution_ctx(); - let members = case.members.values(); - let values: BufferMut = generated_values(case, &members) - .into_iter() - .map(T::from_counter) - .collect(); + let members = (0..case.member_count) + .map(|index| u64::try_from(index).unwrap() * case.member_stride) + .collect::>(); + let generated = generated_values(case, &members); + let expected = BoolArray::from_iter(generated.iter().map(|value| members.contains(value))); + let values: BufferMut = generated.into_iter().map(T::from_counter).collect(); let packed = page_aligned( BitPackedData::encode( &PrimitiveArray::new(values.freeze(), Validity::NonNullable).into_array(), @@ -282,87 +209,23 @@ fn packed_input(case: PackedCase) -> (BitPackedArray, Scalar, Vorte ) .unwrap(), ); - (packed, list_scalar::(&members), session) + (packed, list_scalar::(&members), expected, session) } -fn old_generic_contains(values: ArrayRef, list: &Scalar) -> ArrayRef { - let false_scalar = Scalar::bool(false, values.dtype().nullability()); - let mut level = list - .as_list() - .elements() - .vortex_expect("benchmark list is non-null") - .iter() - .map(|member| { - Binary::try_new( - ConstantArray::new(member.clone(), values.len()).into_array(), - values.clone(), - Operator::Eq, - ) - .unwrap() - .into_array() - .fill_null(false_scalar.clone()) - .unwrap() - }) - .collect::>(); - - while level.len() > 1 { - let mut next = Vec::with_capacity(level.len().div_ceil(2)); - let mut arrays = level.into_iter(); - while let Some(left) = arrays.next() { - next.push(if let Some(right) = arrays.next() { - left.binary(right, Operator::Or).unwrap() - } else { - left - }); - } - level = next; - } - - level.pop().vortex_expect("benchmark list is nonempty") -} - -fn bench_packed_specialized(bencher: Bencher, case: PackedCase) { - let (packed, list, session) = packed_input::(case); +fn bench_packed_current(bencher: Bencher, case: PackedCase) { + let (packed, list, expected, session) = packed_input::(case); let contains = packed .into_array() .apply(&list_contains(lit(list), root())) .unwrap(); let mut ctx = session.create_execution_ctx(); + let actual = contains.clone().execute::(&mut ctx).unwrap(); + assert_arrays_eq!(actual, expected, &mut ctx); bencher .counter(ItemsCount::new(case.len)) .bench_local(|| black_box(contains.clone().execute::(&mut ctx).unwrap())); } -fn bench_packed_decode_once(bencher: Bencher, case: PackedCase) { - let (packed, list, session) = packed_input::(case); - let list = ConstantArray::new(list, case.len).into_array(); - let mut ctx = session.create_execution_ctx(); - bencher.counter(ItemsCount::new(case.len)).bench_local(|| { - let primitive = packed - .clone() - .into_array() - .execute::(&mut ctx) - .unwrap(); - let result = ::list_contains( - &list, - primitive.as_view(), - &mut ctx, - ) - .unwrap() - .unwrap(); - black_box(result.execute::(&mut ctx).unwrap()) - }); -} - -fn bench_packed_old_generic(bencher: Bencher, case: PackedCase) { - let (packed, list, session) = packed_input::(case); - let mut ctx = session.create_execution_ctx(); - bencher.counter(ItemsCount::new(case.len)).bench_local(|| { - let result = old_generic_contains(packed.clone().into_array(), &list); - black_box(result.execute::(&mut ctx).unwrap()) - }); -} - macro_rules! dispatch_packed { ($bencher:expr, $case:expr, $function:ident) => { match $case.ptype { @@ -375,105 +238,8 @@ macro_rules! dispatch_packed { }; } +#[vortex_bench_support::cpu_features] #[divan::bench(args = PACKED_CASES)] -fn packed_specialized(bencher: Bencher, case: PackedCase) { - dispatch_packed!(bencher, case, bench_packed_specialized); -} - -#[divan::bench(args = PACKED_CASES)] -fn packed_decode_once(bencher: Bencher, case: PackedCase) { - dispatch_packed!(bencher, case, bench_packed_decode_once); -} - -#[divan::bench(args = OLD_GENERIC_CASES)] -fn packed_old_generic(bencher: Bencher, case: PackedCase) { - dispatch_packed!(bencher, case, bench_packed_old_generic); -} - -#[cfg(not(codspeed))] -fn length_sweep_cases() -> Vec { - [ - 2_048, 2_049, 2_304, 2_560, 3_072, 4_095, 4_096, 4_097, 6_144, 8_192, - ] - .map(|len| PackedCase { - name: "length_sweep_m5_span4096", - ptype: PType::U32, - bit_width: 13, - len, - members: MemberSpec::Explicit(FIVE_DENSE_BOUNDARY_MEMBERS), - hit_percent: 50, - }) - .to_vec() -} - -#[cfg(not(codspeed))] -#[divan::bench(args = length_sweep_cases())] -fn length_sweep_specialized(bencher: Bencher, case: PackedCase) { - bench_packed_specialized::(bencher, case); -} - -#[cfg(not(codspeed))] -#[divan::bench(args = length_sweep_cases())] -fn length_sweep_decode_once(bencher: Bencher, case: PackedCase) { - bench_packed_decode_once::(bencher, case); -} - -fn primitive_input() -> (PrimitiveArray, Scalar, VortexSession) { - const LEN: usize = 65_536; - let case = PackedCase { - name: "primitive", - ptype: T::PTYPE, - bit_width: 12, - len: LEN, - members: MemberSpec::Stride { - count: 8, - stride: 2, - }, - hit_percent: 50, - }; - let members = case.members.values(); - let values = generated_values(case, &members) - .into_iter() - .map(T::from_counter) - .collect::(); - (values, list_scalar::(&members), array_session()) +fn packed_current(bencher: Bencher, case: PackedCase) { + dispatch_packed!(bencher, case, bench_packed_current); } - -macro_rules! primitive_benchmarks { - ($module:ident, $T:ty) => { - mod $module { - use super::*; - - #[divan::bench] - fn specialized(bencher: Bencher) { - let (values, list, session) = primitive_input::<$T>(); - let len = values.len(); - let contains = values - .into_array() - .apply(&list_contains(lit(list), root())) - .unwrap(); - let mut ctx = session.create_execution_ctx(); - bencher.counter(ItemsCount::new(len)).bench_local(|| { - black_box(contains.clone().execute::(&mut ctx).unwrap()) - }); - } - - #[divan::bench] - fn old_generic(bencher: Bencher) { - let (values, list, session) = primitive_input::<$T>(); - let len = values.len(); - let values = values.into_array(); - let mut ctx = session.create_execution_ctx(); - bencher.counter(ItemsCount::new(len)).bench_local(|| { - let result = old_generic_contains(values.clone(), &list); - black_box(result.execute::(&mut ctx).unwrap()) - }); - } - } - }; -} - -primitive_benchmarks!(primitive_u8, u8); -primitive_benchmarks!(primitive_u16, u16); -primitive_benchmarks!(primitive_u32, u32); -primitive_benchmarks!(primitive_u64, u64); diff --git a/encodings/fastlanes/src/bitpacking/compute/compare_fused.rs b/encodings/fastlanes/src/bitpacking/compute/compare_fused.rs index f27384b898c..82a43647fc0 100644 --- a/encodings/fastlanes/src/bitpacking/compute/compare_fused.rs +++ b/encodings/fastlanes/src/bitpacking/compute/compare_fused.rs @@ -1,12 +1,12 @@ // SPDX-License-Identifier: Apache-2.0 // SPDX-FileCopyrightText: Copyright the Vortex contributors -//! Fused compare kernel for [`BitPackedArray`] against a constant. +//! Fused predicate kernel for [`BitPackedArray`]. //! //! Where [`super::stream_predicate`] unpacks a full 1024-element FastLanes block into a scratch //! buffer and *then* folds a predicate over it, this path hands the comparison down into the -//! FastLanes [`BitPackingCompare::unchecked_unpack_cmp`] kernel, which compares each value against -//! the constant *as it is unpacked*, accumulating the boolean results straight into a 1024-bit +//! FastLanes [`BitPackingCompare::unchecked_unpack_cmp`] kernel, which evaluates each value +//! *as it is unpacked*, accumulating the boolean results straight into a 1024-bit //! mask (`[u64; 16]`) in transposed FastLanes lane order - one register-resident word per lane, no //! `[bool; 1024]` or `[T; 1024]` scratch. A single SIMD [`transpose_bits`] per block then rotates //! that mask into logical row order. @@ -21,7 +21,7 @@ //! slot with no per-block temporary and only one shared scratch `[u64; 16]`. The leading `offset` //! garbage rows are represented as the final [`BitBuffer`] bit offset, which naturally handles //! sub-byte slices without copy-aligning. Inline patches are spliced in afterwards by overwriting -//! the bits at the patched indices with `cmp(patch_value, rhs)`. +//! the bits at the patched indices with the predicate result. //! //! [`BitPackedArray`]: crate::BitPackedArray //! [`BitBuffer`]: vortex_buffer::BitBuffer @@ -70,6 +70,46 @@ pub(super) fn stream_compare_fused( cmp: F, ctx: &mut ExecutionCtx, ) -> VortexResult +where + T: NativePType + + BitPackedIter + + FastLanesComparable::Physical>, + ::Physical: BitPacking + NativePType + BitPackingCompare, + F: Fn(T, T) -> bool + Copy, +{ + stream_compare_fused_inner(array, rhs, nullability, cmp, ctx) +} + +/// Evaluates `predicate` while FastLanes unpacks each value. +pub(super) fn stream_predicate_fused( + array: ArrayView<'_, BitPacked>, + nullability: Nullability, + predicate: F, + ctx: &mut ExecutionCtx, +) -> VortexResult +where + T: NativePType + + BitPackedIter + + FastLanesComparable::Physical>, + ::Physical: BitPacking + NativePType + BitPackingCompare, + F: Fn(T) -> bool + Copy, +{ + stream_compare_fused_inner( + array, + T::default(), + nullability, + move |value, _| predicate(value), + ctx, + ) +} + +fn stream_compare_fused_inner( + array: ArrayView<'_, BitPacked>, + rhs: T, + nullability: Nullability, + cmp: F, + ctx: &mut ExecutionCtx, +) -> VortexResult where T: NativePType + BitPackedIter @@ -84,7 +124,7 @@ where // A degenerate width has no packed payload for the fused kernel to consume; defer to the scalar // streaming predicate, which handles every layout (including the empty array). if len == 0 || bit_width == 0 { - return stream_predicate::(array, nullability, move |v| cmp(v, rhs), ctx); + return stream_predicate::(array, nullability, move |value| cmp(value, rhs), ctx); } // Over-allocate to whole 1024-bit blocks in padded coordinates so every block - including the @@ -119,12 +159,12 @@ where let mut bits = BitBufferMut::from_buffer(words.into_byte_buffer(), offset, len); - // Patched indices hold placeholder packed values, so their fused result is meaningless; - // overwrite each with the comparison against the real patch value. + // Patched indices hold placeholder packed values, so their fused result is meaningless. + // Overwrite each result with the predicate for the real patch value. // TODO(joe): apply patches per `packed_chunked`. if let Some(p) = array.patches() { let p_idx = p.indices().clone().execute::(ctx)?; - // TODO(joe): push down cmp?? + // TODO(joe): push down the predicate. let p_val = p.values().clone().execute::(ctx)?; let p_off = p.offset(); match_each_unsigned_integer_ptype!(p_idx.ptype(), |I| { diff --git a/encodings/fastlanes/src/bitpacking/compute/list_contains/mod.rs b/encodings/fastlanes/src/bitpacking/compute/list_contains/mod.rs index 405978de24c..d708c60a89e 100644 --- a/encodings/fastlanes/src/bitpacking/compute/list_contains/mod.rs +++ b/encodings/fastlanes/src/bitpacking/compute/list_contains/mod.rs @@ -1,26 +1,53 @@ // SPDX-License-Identifier: Apache-2.0 // SPDX-FileCopyrightText: Copyright the Vortex contributors +use fastlanes::BitPacking; +use fastlanes::BitPackingCompare; +use fastlanes::FastLanesComparable; use vortex_array::ArrayRef; use vortex_array::ArrayView; use vortex_array::ExecutionCtx; use vortex_array::IntoArray; use vortex_array::arrays::BoolArray; use vortex_array::arrays::PrimitiveArray; -use vortex_array::dtype::DType; +use vortex_array::arrays::primitive::evaluate_prepared_integer_membership; +use vortex_array::arrays::primitive::integer_membership_binary_search_min; +use vortex_array::dtype::IntegerPType; use vortex_array::dtype::NativePType; +use vortex_array::dtype::PType; +use vortex_array::dtype::PhysicalPType; use vortex_array::match_each_integer_ptype; use vortex_array::scalar_fn::fns::list_contains::IntegerMembership; use vortex_array::scalar_fn::fns::list_contains::ListContainsElementKernel; use vortex_buffer::BitBuffer; use vortex_error::VortexResult; -use vortex_error::vortex_err; -use super::compare_fused::stream_compare_fused; +use super::compare_fused::stream_predicate_fused; use crate::BitPacked; +use crate::unpack_iter::BitPacked as BitPackedIter; -// Decode short batches once because their fixed fusion overhead exceeds the saved materialization. -const MIN_DENSE_FUSION_LEN: usize = 2_048; +const MAX_FUSED_DISTINCT_MEMBERS: usize = 4; +const SHORT_ARRAY_MAX_ROWS_8_16: usize = 8_192; +const SHORT_ARRAY_MAX_ROWS_32: usize = 16_384; +fn min_decode_source_members(ptype: PType, len: usize) -> usize { + // The generic fallback scans the packed child once per source member. Decode before repeated + // packed scans become more expensive than one decode plus Primitive membership evaluation. + let short_array_max_rows = if ptype.bit_width() == 32 { + SHORT_ARRAY_MAX_ROWS_32 + } else { + SHORT_ARRAY_MAX_ROWS_8_16 + }; + if len <= short_array_max_rows && ptype.bit_width() < 64 { + return integer_membership_binary_search_min(ptype); + } + match ptype.bit_width() { + 8 => 30, + 16 => 25, + 32 => 13, + 64 => 5, + _ => 5, + } +} impl ListContainsElementKernel for BitPacked { fn list_contains( @@ -37,104 +64,92 @@ fn list_contains_compressed( element: ArrayView<'_, BitPacked>, ctx: &mut ExecutionCtx, ) -> VortexResult> { - let Some(list_scalar) = list.as_constant() else { - return Ok(None); - }; - let DType::List(member_dtype, _) = list.dtype() else { - return Ok(None); - }; - if !member_dtype.eq_ignore_nullability(element.dtype()) { - return Ok(None); - } - let nullability = list.dtype().nullability() | element.dtype().nullability(); - let Some(elements) = list_scalar.as_list().elements() else { + + match_each_integer_ptype!(element.dtype().as_ptype(), |T| { + list_contains_typed::(list, element, nullability, ctx) + }) +} + +fn list_contains_typed( + list: &ArrayRef, + element: ArrayView<'_, BitPacked>, + nullability: vortex_array::dtype::Nullability, + ctx: &mut ExecutionCtx, +) -> VortexResult> +where + T: IntegerPType + + BitPackedIter + + FastLanesComparable::Physical>, + ::Physical: BitPacking + NativePType + BitPackingCompare, +{ + let Some(membership) = IntegerMembership::::try_from_constant_list(list, element.dtype())? + else { return Ok(None); }; - if elements.is_empty() { - return Ok(None); + if membership.members().len() > MAX_FUSED_DISTINCT_MEMBERS { + if membership.non_null_source_len() + < min_decode_source_members(element.dtype().as_ptype(), element.len()) + { + return Ok(None); + } + // The generic list implementation expands membership into one comparison per source + // member. Each comparison scans the packed child. Decode once before applying the + // Primitive membership policy when repeated packed scans become more expensive. + let primitive = element.array().clone().execute::(ctx)?; + return evaluate_prepared_integer_membership(membership, primitive.as_view(), nullability) + .map(Some); } - let result = match_each_integer_ptype!(element.dtype().as_ptype(), |T| { - let members = elements - .iter() - .map(|value| { - value - .as_primitive_opt() - .ok_or_else(|| vortex_err!("List member is not a primitive scalar"))? - .try_typed_value::() - }) - .collect::>>>()? - .into_iter() - .flatten() - .collect::>(); - - let membership = IntegerMembership::new(members); - - match membership.members() { - [] => BoolArray::new( - BitBuffer::new_unset(element.len()), - element.validity()?.union_nullability(nullability), - ) - .into_array(), - [member] => { - let member = *member; - stream_compare_fused::(element, member, nullability, NativePType::is_eq, ctx)? - } - [first, second] => { - let (first, second) = (*first, *second); - stream_compare_fused::( - element, - first, - nullability, - move |value, _| value.is_eq(first) | value.is_eq(second), - ctx, - )? - } - [first, second, third] => { - let (first, second, third) = (*first, *second, *third); - stream_compare_fused::( - element, - first, - nullability, - move |value, _| value.is_eq(first) | value.is_eq(second) | value.is_eq(third), - ctx, - )? - } - [first, second, third, fourth] => { - let (first, second, third, fourth) = (*first, *second, *third, *fourth); - stream_compare_fused::( - element, - first, - nullability, - move |value, _| { - value.is_eq(first) - | value.is_eq(second) - | value.is_eq(third) - | value.is_eq(fourth) - }, - ctx, - )? - } - _ => { - if membership.uses_dense_table() && element.len() >= MIN_DENSE_FUSION_LEN { - stream_compare_fused::( - element, - membership.members()[0], - nullability, - |value, _| membership.contains(value), - ctx, - )? - } else { - let primitive = element - .into_owned() - .into_array() - .execute::(ctx)?; - membership.evaluate_primitive(primitive.as_view(), nullability)? - } - } + let result = match membership.members() { + [] => BoolArray::new( + BitBuffer::new_unset(element.len()), + element.validity()?.union_nullability(nullability), + ) + .into_array(), + [member] => { + let member = *member; + stream_predicate_fused::( + element, + nullability, + move |value| value.is_eq(member), + ctx, + )? + } + [first, second] => { + let (first, second) = (*first, *second); + stream_predicate_fused::( + element, + nullability, + move |value| value.is_eq(first) | value.is_eq(second), + ctx, + )? } - }); + [first, second, third] => { + let (first, second, third) = (*first, *second, *third); + stream_predicate_fused::( + element, + nullability, + move |value| value.is_eq(first) | value.is_eq(second) | value.is_eq(third), + ctx, + )? + } + [first, second, third, fourth] => { + let (first, second, third, fourth) = (*first, *second, *third, *fourth); + stream_predicate_fused::( + element, + nullability, + move |value| { + value.is_eq(first) + | value.is_eq(second) + | value.is_eq(third) + | value.is_eq(fourth) + }, + ctx, + )? + } + _ => return Ok(None), + }; Ok(Some(result)) } diff --git a/encodings/fastlanes/src/bitpacking/compute/list_contains/tests.rs b/encodings/fastlanes/src/bitpacking/compute/list_contains/tests.rs index 4eb377b4df8..435de238896 100644 --- a/encodings/fastlanes/src/bitpacking/compute/list_contains/tests.rs +++ b/encodings/fastlanes/src/bitpacking/compute/list_contains/tests.rs @@ -17,17 +17,15 @@ use vortex_array::assert_arrays_eq; use vortex_array::dtype::DType; use vortex_array::dtype::NativePType; use vortex_array::dtype::Nullability; -#[cfg(not(codspeed))] use vortex_array::expr::list_contains; -#[cfg(not(codspeed))] use vortex_array::expr::lit; -#[cfg(not(codspeed))] use vortex_array::expr::root; use vortex_array::scalar::PValue; use vortex_array::scalar::Scalar; use vortex_array::scalar_fn::fns::list_contains::ListContainsElementKernel; #[cfg(not(codspeed))] use vortex_array::test_harness::trace::trace_op; +use vortex_error::VortexExpect; use vortex_error::VortexResult; use vortex_error::vortex_err; use vortex_session::VortexSession; @@ -115,10 +113,7 @@ integer_type_test!(test_integer_type_i64, i64, 6); #[case::two(vec![3, 7])] #[case::three(vec![3, 7, 11])] #[case::four(vec![3, 7, 11, 15])] -#[case::five(vec![3, 7, 11, 15, 19])] -#[case::larger((0..32).map(|value| value * 3).collect())] -#[case::sparse((0..32).map(|value| value * 10_000).collect())] -#[case::duplicates(vec![3, 3, 7, 7, 11, 11, 15, 15, 15])] +#[case::duplicate_source(vec![3, 3, 7, 7, 11, 11, 15, 15, 15])] fn test_member_cardinalities(#[case] members: Vec) -> VortexResult<()> { let mut ctx = SESSION.create_execution_ctx(); let values = (0..2_048).map(|value| value % 128).collect::>(); @@ -135,6 +130,64 @@ fn test_member_cardinalities(#[case] members: Vec) -> VortexResult<()> { Ok(()) } +#[rstest] +#[case::generic_five((0..5).map(|value| value * 2).collect())] +#[case::decoded_many((0..32).map(|value| value * 2).collect())] +fn test_many_member_public_expression_paths(#[case] members: Vec) -> VortexResult<()> { + let mut ctx = SESSION.create_execution_ctx(); + let values = (0..4_096).map(|value| value % 128).collect::>(); + let primitive = PrimitiveArray::from_iter(values.iter().copied()); + let packed = BitPackedData::encode(&primitive.into_array(), 7, &mut ctx)?; + let expression = list_contains( + lit(member_list( + members.iter().copied().map(Some), + Nullability::NonNullable, + )), + root(), + ); + + let actual = packed + .into_array() + .apply(&expression)? + .execute::(&mut ctx)?; + let expected = BoolArray::from_iter(values.into_iter().map(|value| members.contains(&value))); + assert_arrays_eq!(actual, expected, &mut ctx); + Ok(()) +} + +#[test] +fn test_many_member_kernel_policy() -> VortexResult<()> { + let mut ctx = SESSION.create_execution_ctx(); + let values = [0i32, 7, 99]; + let primitive = PrimitiveArray::from_iter(values); + let packed = BitPackedData::encode(&primitive.into_array(), 7, &mut ctx)?; + let decode_threshold = + super::min_decode_source_members(vortex_array::dtype::PType::I32, packed.len()); + + for (member_count, expected_supported) in + [(decode_threshold - 1, false), (decode_threshold, true)] + { + let member_count = i32::try_from(member_count).vortex_expect("member count fits in an i32"); + let list = list_array( + member_list((0..member_count).map(Some), Nullability::NonNullable), + packed.len(), + ); + let actual = ::list_contains( + &list, + packed.as_view(), + &mut ctx, + )?; + + assert_eq!(actual.is_some(), expected_supported); + if let Some(actual) = actual { + let expected = + BoolArray::from_iter(values.map(|value| (0..member_count).contains(&value))); + assert_arrays_eq!(actual, expected, &mut ctx); + } + } + Ok(()) +} + #[rstest] #[case::present([true; 128], vec![0])] #[case::absent([false; 128], vec![1])] @@ -156,22 +209,6 @@ fn test_zero_bit_width( Ok(()) } -#[test] -fn test_empty_array() -> VortexResult<()> { - let mut ctx = SESSION.create_execution_ctx(); - let primitive = PrimitiveArray::from_iter(std::iter::empty::()); - let packed = BitPackedData::encode(&primitive.into_array(), 1, &mut ctx)?; - let list = list_array( - member_list([Some(0)], Nullability::NonNullable), - packed.len(), - ); - - let actual = execute_direct(&list, &packed, &mut ctx)?; - let expected = BoolArray::from_iter(std::iter::empty::()); - assert_arrays_eq!(actual, expected, &mut ctx); - Ok(()) -} - #[test] fn test_sliced_patched_array() -> VortexResult<()> { let mut ctx = SESSION.create_execution_ctx(); @@ -305,6 +342,7 @@ fn test_registered_kernel_executes_through_expression() -> VortexResult<()> { && line.contains("child=fastlanes.bitpacked") }) .collect::>(); + // A silent fallback preserves values but loses compressed-domain execution. assert_eq!(applied.len(), 1, "{trace}"); let expected = BoolArray::from_iter(values.into_iter().map(|value| members.contains(&value))); diff --git a/vortex-array/Cargo.toml b/vortex-array/Cargo.toml index fb9ea651ad0..cd3dd8c11fe 100644 --- a/vortex-array/Cargo.toml +++ b/vortex-array/Cargo.toml @@ -290,6 +290,10 @@ harness = false name = "list_length" harness = false +[[bench]] +name = "list_contains" +harness = false + [[bench]] name = "list_sum" harness = false diff --git a/vortex-array/benches/list_contains.rs b/vortex-array/benches/list_contains.rs new file mode 100644 index 00000000000..9e67e13678b --- /dev/null +++ b/vortex-array/benches/list_contains.rs @@ -0,0 +1,195 @@ +// SPDX-License-Identifier: Apache-2.0 +// SPDX-FileCopyrightText: Copyright the Vortex contributors + +//! Compares the Primitive constant-list membership dispatch paths. +//! +//! Primitive arrays use direct comparisons for up to four distinct members. They use binary +//! search from 10 members for 8- and 16-bit integers. The 32- and 64-bit thresholds are 11 and 13 +//! members. Every path runs on each real CPU feature leg in CodSpeed. +//! To recalculate the thresholds, run this benchmark twice with temporary policy constants. Use a +//! high cutoff to force generic evaluation. Use `5` to force binary search above four members. +//! +//! Run with `cargo bench -p vortex-array --bench list_contains`. + +#![expect(clippy::unwrap_used)] + +use std::fmt::Display; +use std::fmt::Formatter; +use std::hint::black_box; +use std::sync::Arc; + +use divan::Bencher; +use divan::counter::ItemsCount; +use vortex_array::IntoArray; +use vortex_array::VortexSessionExecute; +use vortex_array::array_session; +use vortex_array::arrays::BoolArray; +use vortex_array::arrays::PrimitiveArray; +use vortex_array::assert_arrays_eq; +use vortex_array::dtype::DType; +use vortex_array::dtype::IntegerPType; +use vortex_array::dtype::Nullability; +use vortex_array::expr::list_contains; +use vortex_array::expr::lit; +use vortex_array::expr::root; +use vortex_array::scalar::Scalar; +use vortex_array::validity::Validity; +use vortex_session::VortexSession; + +fn main() { + divan::main(); +} + +trait BenchInt: IntegerPType + Copy + Into { + fn from_counter(value: u64) -> Self; +} + +impl BenchInt for u8 { + fn from_counter(value: u64) -> Self { + Self::try_from(value).unwrap() + } +} + +impl BenchInt for u16 { + fn from_counter(value: u64) -> Self { + Self::try_from(value).unwrap() + } +} + +impl BenchInt for u32 { + fn from_counter(value: u64) -> Self { + Self::try_from(value).unwrap() + } +} + +impl BenchInt for u64 { + fn from_counter(value: u64) -> Self { + value + } +} + +#[derive(Clone, Copy)] +struct PrimitiveCase { + name: &'static str, + len: usize, + member_count: usize, +} + +impl Display for PrimitiveCase { + fn fmt(&self, formatter: &mut Formatter<'_>) -> std::fmt::Result { + write!( + formatter, + "{}_m{}_n{}", + self.name, self.member_count, self.len + ) + } +} + +const fn primitive_case(name: &'static str, len: usize, member_count: usize) -> PrimitiveCase { + PrimitiveCase { + name, + len, + member_count, + } +} + +const LONG_M1: PrimitiveCase = primitive_case("long", 65_536, 1); +const LONG_M4: PrimitiveCase = primitive_case("long", 65_536, 4); +const LONG_M9: PrimitiveCase = primitive_case("long", 65_536, 9); +const LONG_M10: PrimitiveCase = primitive_case("long", 65_536, 10); +const LONG_M11: PrimitiveCase = primitive_case("long", 65_536, 11); +const LONG_M12: PrimitiveCase = primitive_case("long", 65_536, 12); +const LONG_M13: PrimitiveCase = primitive_case("long", 65_536, 13); +const LONG_M32: PrimitiveCase = primitive_case("long", 65_536, 32); +const SHORT_M11: PrimitiveCase = primitive_case("short", 1_024, 11); +const SHORT_M13: PrimitiveCase = primitive_case("short", 1_024, 13); + +const CURRENT_10: &[PrimitiveCase] = &[LONG_M1, LONG_M4, LONG_M9, LONG_M10, LONG_M32, SHORT_M11]; +const CURRENT_11: &[PrimitiveCase] = &[LONG_M1, LONG_M4, LONG_M10, LONG_M11, LONG_M32, SHORT_M11]; +const CURRENT_13: &[PrimitiveCase] = &[LONG_M1, LONG_M4, LONG_M12, LONG_M13, LONG_M32, SHORT_M13]; + +fn primitive_input( + case: PrimitiveCase, +) -> (PrimitiveArray, Scalar, BoolArray, VortexSession) { + let members = (0..case.member_count) + .map(|index| T::from_counter(u64::try_from(index).unwrap() * 2)) + .collect::>(); + let domain_bits = T::PTYPE.bit_width().min(12); + let domain_size = 1u64 << domain_bits; + let mut state = 0x9E37_79B9_7F4A_7C15u64; + let generated = (0..case.len) + .map(|_| { + state = state + .wrapping_mul(6_364_136_223_846_793_005) + .wrapping_add(1_442_695_040_888_963_407); + if (state >> 32).is_multiple_of(2) { + let member_index = + usize::try_from(state % u64::try_from(members.len()).unwrap()).unwrap(); + members[member_index] + } else { + let mut candidate = state.rotate_left(17) % domain_size; + while members.contains(&T::from_counter(candidate)) { + candidate = (candidate + 1) % domain_size; + } + T::from_counter(candidate) + } + }) + .collect::>(); + let expected = BoolArray::from_iter(generated.iter().map(|value| members.contains(value))); + let list = Scalar::list( + Arc::new(DType::Primitive(T::PTYPE, Nullability::NonNullable)), + members.iter().copied().map(Into::into).collect(), + Nullability::NonNullable, + ); + ( + PrimitiveArray::new::(generated, Validity::NonNullable), + list, + expected, + array_session(), + ) +} + +fn bench_current(bencher: Bencher, case: PrimitiveCase) { + let (array, list, expected, session) = primitive_input::(case); + let expression = list_contains(lit(list), root()); + let mut ctx = session.create_execution_ctx(); + let actual = array + .clone() + .into_array() + .apply(&expression) + .unwrap() + .execute::(&mut ctx) + .unwrap(); + assert_arrays_eq!(actual, expected, &mut ctx); + + bencher.counter(ItemsCount::new(case.len)).bench_local(|| { + black_box( + array + .clone() + .into_array() + .apply(&expression) + .unwrap() + .execute::(&mut ctx) + .unwrap(), + ) + }); +} + +macro_rules! primitive_benchmarks { + ($type_name:ident, $ty:ty, $current:ident) => { + mod $type_name { + use super::*; + + #[vortex_bench_support::cpu_features] + #[divan::bench(args = $current)] + fn current(bencher: Bencher, case: PrimitiveCase) { + bench_current::<$ty>(bencher, case); + } + } + }; +} + +primitive_benchmarks!(u8_cases, u8, CURRENT_10); +primitive_benchmarks!(u16_cases, u16, CURRENT_10); +primitive_benchmarks!(u32_cases, u32, CURRENT_11); +primitive_benchmarks!(u64_cases, u64, CURRENT_13); diff --git a/vortex-array/src/arrays/primitive/compute/list_contains.rs b/vortex-array/src/arrays/primitive/compute/list_contains.rs index 740441358c4..f488734e281 100644 --- a/vortex-array/src/arrays/primitive/compute/list_contains.rs +++ b/vortex-array/src/arrays/primitive/compute/list_contains.rs @@ -1,61 +1,87 @@ // SPDX-License-Identifier: Apache-2.0 // SPDX-FileCopyrightText: Copyright the Vortex contributors -use vortex_error::VortexExpect; use vortex_error::VortexResult; use crate::ArrayRef; use crate::ArrayView; use crate::ExecutionCtx; use crate::arrays::Primitive; -use crate::dtype::DType; +use crate::dtype::IntegerPType; +use crate::dtype::PType; use crate::match_each_integer_ptype; use crate::scalar_fn::fns::list_contains::IntegerMembership; use crate::scalar_fn::fns::list_contains::ListContainsElementKernel; +use crate::scalar_fn::fns::list_contains::constant_list_scalar_contains; + +/// Returns the source-member count where Primitive integer membership uses binary search. +#[doc(hidden)] +pub fn integer_membership_binary_search_min(ptype: PType) -> usize { + // The generic implementation evaluates one equality expression per source member. Use the + // prepared set once binary search becomes faster than the expression tree. + match ptype.bit_width() { + 8 | 16 => 10, + 32 => 11, + 64 => 13, + _ => 13, + } +} impl ListContainsElementKernel for Primitive { fn list_contains( list: &ArrayRef, element: ArrayView<'_, Self>, - _ctx: &mut ExecutionCtx, + ctx: &mut ExecutionCtx, ) -> VortexResult> { - let Some(list_scalar) = list.as_constant() else { - return Ok(None); - }; - let DType::List(member_dtype, _) = list.dtype() else { - return Ok(None); - }; - if !member_dtype.eq_ignore_nullability(element.dtype()) || !element.ptype().is_int() { - return Ok(None); - } + evaluate_constant_list_membership(list, element, ctx) + } +} - let nullability = list.dtype().nullability() | element.dtype().nullability(); - let Some(elements) = list_scalar.as_list().elements() else { - return Ok(None); - }; - if elements.is_empty() { - return Ok(None); - } - - let result = match_each_integer_ptype!(element.ptype(), |T| { - let members = elements - .iter() - .map(|value| { - value - .as_primitive_opt() - .vortex_expect("list dtype was checked before member extraction") - .try_typed_value::() - }) - .collect::>>>()? - .into_iter() - .flatten() - .collect::>(); - - IntegerMembership::new(members).evaluate_primitive(element, nullability)? - }); - - Ok(Some(result)) +fn evaluate_constant_list_membership( + list: &ArrayRef, + element: ArrayView<'_, Primitive>, + _ctx: &mut ExecutionCtx, +) -> VortexResult> { + if !element.ptype().is_int() { + return Ok(None); } + + let nullability = list.dtype().nullability() | element.dtype().nullability(); + + match_each_integer_ptype!(element.ptype(), |T| { + evaluate_integer_membership::(list, element, nullability) + }) +} + +fn evaluate_integer_membership( + list: &ArrayRef, + element: ArrayView<'_, Primitive>, + nullability: crate::dtype::Nullability, +) -> VortexResult> { + let Some(membership) = IntegerMembership::::try_from_constant_list(list, element.dtype())? + else { + return Ok(None); + }; + evaluate_prepared_integer_membership(membership, element, nullability).map(Some) +} + +/// Evaluates a prepared integer set against a Primitive array. +#[doc(hidden)] +pub fn evaluate_prepared_integer_membership( + membership: IntegerMembership, + element: ArrayView<'_, Primitive>, + nullability: crate::dtype::Nullability, +) -> VortexResult { + if membership.members().len() > 4 + && membership.non_null_source_len() < integer_membership_binary_search_min(element.ptype()) + { + return constant_list_scalar_contains( + &membership.source_list().as_list(), + element.array(), + nullability, + ); + } + membership.evaluate_primitive(element, nullability) } #[cfg(test)] @@ -63,21 +89,23 @@ mod tests { use std::sync::Arc; use rstest::rstest; + use vortex_error::VortexExpect; use super::*; use crate::IntoArray; use crate::VortexSessionExecute; use crate::arrays::BoolArray; + use crate::arrays::Constant; use crate::arrays::ConstantArray; use crate::arrays::PrimitiveArray; use crate::assert_arrays_eq; + use crate::dtype::DType; use crate::dtype::Nullability; + use crate::dtype::PType::F32; use crate::dtype::PType::I32; - #[cfg(not(codspeed))] + use crate::dtype::PType::I64; use crate::expr::list_contains; - #[cfg(not(codspeed))] use crate::expr::lit; - #[cfg(not(codspeed))] use crate::expr::root; use crate::scalar::Scalar; #[cfg(not(codspeed))] @@ -103,8 +131,10 @@ mod tests { #[case::two(vec![3, 7])] #[case::three(vec![3, 7, 11])] #[case::four(vec![3, 7, 11, 15])] - #[case::dense((0..32).map(|value| value * 3).collect())] - #[case::sparse((0..32).map(|value| value * 10_000).collect())] + #[case::five((0..5).map(|value| value * 3).collect())] + #[case::eleven((0..11).map(|value| value * 3).collect())] + #[case::many((0..32).map(|value| value * 3).collect())] + #[case::duplicate_heavy((0..32).map(|value| value % 5).collect())] fn test_membership_plans(#[case] members: Vec) -> VortexResult<()> { let mut ctx = crate::array_session().create_execution_ctx(); let values = [0, 3, 7, 15, 31, 90_000, 310_000]; @@ -122,6 +152,37 @@ mod tests { Ok(()) } + #[rstest] + #[case::small(5)] + #[case::many(13)] + fn test_i64_membership(#[case] member_count: usize) -> VortexResult<()> { + let mut ctx = crate::array_session().create_execution_ctx(); + let members = (0..member_count) + .map(|value| i64::try_from(value).vortex_expect("member count fits i64")) + .collect::>(); + let values = [0i64, 11, 99]; + let element = PrimitiveArray::from_iter(values); + let list = ConstantArray::new( + Scalar::list( + Arc::new(DType::Primitive(I64, Nullability::NonNullable)), + members.iter().copied().map(Scalar::from).collect(), + Nullability::NonNullable, + ), + element.len(), + ) + .into_array(); + + let actual = ::list_contains( + &list, + element.as_view(), + &mut ctx, + )? + .vortex_expect("integer constant-list membership is supported"); + let expected = BoolArray::from_iter(values.map(|value| members.contains(&value))); + assert_arrays_eq!(actual, expected, &mut ctx); + Ok(()) + } + #[test] #[cfg(not(codspeed))] fn test_registered_kernel_executes_through_expression() -> VortexResult<()> { @@ -147,6 +208,7 @@ mod tests { && line.contains("child=vortex.primitive") }) .collect::>(); + // A silent fallback preserves values but loses the membership optimization. assert_eq!(applied.len(), 1, "{trace}"); let expected = BoolArray::from_iter(values.map(|value| members.contains(&value))); @@ -154,6 +216,33 @@ mod tests { Ok(()) } + #[test] + fn test_float_falls_back_through_expression() -> VortexResult<()> { + let mut ctx = crate::array_session().create_execution_ctx(); + let values = [1.5f32, 2.5, 3.5]; + let element = PrimitiveArray::from_iter(values); + let members = [1.5f32, 3.5]; + let list = ConstantArray::new( + Scalar::list( + Arc::new(DType::Primitive(F32, Nullability::NonNullable)), + members.into_iter().map(Scalar::from).collect(), + Nullability::NonNullable, + ), + element.len(), + ) + .into_array(); + let list_scalar = list.as_constant().vortex_expect("list is constant"); + + let actual = element + .into_array() + .apply(&list_contains(lit(list_scalar), root()))? + .execute::(&mut ctx)?; + let expected = BoolArray::from_iter(values.map(|value| members.contains(&value))); + + assert_arrays_eq!(actual, expected, &mut ctx); + Ok(()) + } + #[test] fn test_null_needles() -> VortexResult<()> { let mut ctx = crate::array_session().create_execution_ctx(); @@ -171,6 +260,33 @@ mod tests { Ok(()) } + #[rstest] + #[case::null_list(true)] + #[case::empty_list(false)] + fn test_constant_list_adaptor(#[case] null_list: bool) -> VortexResult<()> { + let member_dtype = DType::Primitive(I32, Nullability::NonNullable); + let list = if null_list { + Scalar::null(DType::List(Arc::new(member_dtype), Nullability::Nullable)) + } else { + Scalar::list(Arc::new(member_dtype), vec![], Nullability::NonNullable) + }; + let needles = PrimitiveArray::from_option_iter([Some(1i32), None, Some(3)]).into_array(); + + let mut ctx = crate::array_session().create_execution_ctx(); + let contains = needles + .apply(&list_contains(lit(list), root()))? + .execute::(&mut ctx)?; + let expected = if null_list { + BoolArray::from_iter([None, None, None]) + } else { + BoolArray::from_iter([Some(false), Some(false), Some(false)]) + }; + + assert!(contains.is::()); + assert_arrays_eq!(contains, expected, &mut ctx); + Ok(()) + } + #[rstest] #[case::mixed( vec![Some(1), None, Some(3)], diff --git a/vortex-array/src/arrays/primitive/compute/mod.rs b/vortex-array/src/arrays/primitive/compute/mod.rs index 7f1dcdcb4cf..7769def7df5 100644 --- a/vortex-array/src/arrays/primitive/compute/mod.rs +++ b/vortex-array/src/arrays/primitive/compute/mod.rs @@ -6,6 +6,8 @@ mod cast; mod fill_null; mod fixed_width; mod list_contains; +pub use list_contains::evaluate_prepared_integer_membership; +pub use list_contains::integer_membership_binary_search_min; mod mask; pub(crate) mod rules; mod slice; diff --git a/vortex-array/src/arrays/primitive/mod.rs b/vortex-array/src/arrays/primitive/mod.rs index 748beda8339..b99ec51c0f8 100644 --- a/vortex-array/src/arrays/primitive/mod.rs +++ b/vortex-array/src/arrays/primitive/mod.rs @@ -14,6 +14,10 @@ pub use vtable::PrimitiveArray; pub(crate) mod compute; mod vtable; +#[doc(hidden)] +pub use compute::evaluate_prepared_integer_membership; +#[doc(hidden)] +pub use compute::integer_membership_binary_search_min; pub use compute::rules::PrimitiveMaskedValidityRule; pub use vtable::Primitive; diff --git a/vortex-array/src/scalar_fn/fns/list_contains/integer_membership.rs b/vortex-array/src/scalar_fn/fns/list_contains/integer_membership.rs index b5005c72aa7..1e206870a64 100644 --- a/vortex-array/src/scalar_fn/fns/list_contains/integer_membership.rs +++ b/vortex-array/src/scalar_fn/fns/list_contains/integer_membership.rs @@ -11,63 +11,90 @@ use crate::ArrayView; use crate::IntoArray; use crate::arrays::BoolArray; use crate::arrays::Primitive; +use crate::dtype::DType; use crate::dtype::IntegerPType; use crate::dtype::NativePType; use crate::dtype::Nullability; - -const MAX_DENSE_SPAN: usize = 4_096; +use crate::scalar::Scalar; /// A prepared integer set for constant-list membership kernels. /// -/// The set sorts and deduplicates lists with more than four members. It builds a byte table when -/// the member span fits the bounded table. +/// The set sorts and deduplicates its members. +#[doc(hidden)] pub struct IntegerMembership { members: Box<[T]>, - dense: Option, + non_null_source_len: usize, + source_list: Scalar, } impl IntegerMembership { - /// Prepares a membership set from integer values. - pub fn new(mut members: Vec) -> Self { - if members.len() > 4 { - members.sort_unstable(); - members.dedup(); - } - let dense = DenseIntegerMembership::try_new(&members); - + fn new(mut members: Vec, source_list: Scalar) -> Self { + let non_null_source_len = members.len(); + members.sort_unstable(); + members.dedup(); Self { members: members.into_boxed_slice(), - dense, + non_null_source_len, + source_list, + } + } + + /// Extracts an integer set from a compatible constant list. + pub fn try_from_constant_list( + list: &ArrayRef, + element_dtype: &DType, + ) -> VortexResult> { + let Some(list_scalar) = list.as_constant() else { + return Ok(None); + }; + let DType::List(member_dtype, _) = list.dtype() else { + return Ok(None); + }; + if !member_dtype.eq_ignore_nullability(element_dtype) { + return Ok(None); } + let Some(elements) = list_scalar.as_list().elements() else { + return Ok(None); + }; + + let members = elements + .iter() + .map(|value| { + value + .as_primitive_opt() + .vortex_expect("list member type was checked") + .try_typed_value::() + }) + .collect::>>>()? + .into_iter() + .flatten() + .collect(); + Ok(Some(Self::new(members, list_scalar))) } - /// Returns the normalized members. + /// Returns the prepared members. pub fn members(&self) -> &[T] { &self.members } - /// Returns true when this set uses a dense lookup table. - pub fn uses_dense_table(&self) -> bool { - self.dense.is_some() + /// Returns the number of non-null source members before deduplication. + #[doc(hidden)] + pub fn non_null_source_len(&self) -> usize { + self.non_null_source_len } - /// Tests membership through the selected lookup representation. - pub fn contains(&self, value: T) -> bool { - self.dense.as_ref().map_or_else( - || { - if self.members.len() <= 4 { - self.members.contains(&value) - } else { - self.members.binary_search(&value).is_ok() - } - }, - |dense| dense.contains(value), - ) + pub(crate) fn source_list(&self) -> &Scalar { + &self.source_list + } + + /// Tests whether the prepared set contains `value`. + pub(crate) fn contains(&self, value: T) -> bool { + self.members.binary_search(&value).is_ok() } /// Evaluates this set against a primitive array of the same integer type. - pub fn evaluate_primitive( - &self, + pub(crate) fn evaluate_primitive( + self, element: ArrayView<'_, Primitive>, nullability: Nullability, ) -> VortexResult { @@ -93,7 +120,7 @@ impl IntegerMembership { | value.is_eq(*third) | value.is_eq(*fourth) }), - _ => collect_many(values, self), + _ => collect_many(values, &self), }; Ok(BoolArray::new(bits, element.validity()?.union_nullability(nullability)).into_array()) @@ -108,101 +135,9 @@ fn collect_direct(values: &[T], mut predicate: impl FnMut(T) -> } fn collect_many(values: &[T], membership: &IntegerMembership) -> BitBuffer { - if let Some(dense) = membership.dense.as_ref() { - return BitBuffer::collect_bool(values.len(), |index| { - // SAFETY: collect_bool visits each valid index once. - let value = unsafe { *values.get_unchecked(index) }; - dense.contains(value) - }); - } - BitBuffer::collect_bool(values.len(), |index| { // SAFETY: collect_bool visits each valid index once. let value = unsafe { *values.get_unchecked(index) }; membership.contains(value) }) } - -/// A bounded byte table for dense integer membership. -struct DenseIntegerMembership { - minimum: i128, - table: Box<[u8]>, -} - -impl DenseIntegerMembership { - fn try_new(members: &[T]) -> Option { - if members.len() <= 4 { - return None; - } - - let minimum = members[0].to_i128()?; - let maximum = members[members.len() - 1].to_i128()?; - let span = usize::try_from(maximum - minimum + 1).ok()?; - if span > MAX_DENSE_SPAN { - return None; - } - - let mut table = vec![0u8; span]; - for member in members { - let index = usize::try_from( - member.to_i128().vortex_expect("integer converts to i128") - minimum, - ) - .vortex_expect("member lies inside the dense span"); - table[index] = 1; - } - - Some(Self { - minimum, - table: table.into_boxed_slice(), - }) - } - - /// Tests whether the table contains an integer value. - fn contains(&self, value: T) -> bool { - let offset = value.to_i128().vortex_expect("integer converts to i128") - self.minimum; - usize::try_from(offset) - .ok() - .and_then(|offset| self.table.get(offset)) - .copied() - .unwrap_or(0) - != 0 - } -} - -#[cfg(test)] -mod tests { - use super::IntegerMembership; - - #[test] - fn normalizes_large_unsorted_duplicates() { - let membership = IntegerMembership::new(vec![7i32, 3, 7, 1, 9, 3, 1]); - - assert_eq!(membership.members(), &[1, 3, 7, 9]); - assert!(membership.contains(1)); - assert!(membership.contains(3)); - assert!(membership.contains(7)); - assert!(!membership.contains(5)); - } - - #[test] - fn dense_table_span_boundary() { - let at_limit = IntegerMembership::new(vec![0i32, 1, 2, 3, 4_095]); - let above_limit = IntegerMembership::new(vec![0i32, 1, 2, 3, 4_096]); - - assert!(at_limit.uses_dense_table()); - assert!(!above_limit.uses_dense_table()); - } - - #[test] - fn integer_extremes_do_not_overflow() { - let signed = IntegerMembership::new(vec![i64::MAX, 0, i64::MIN, -1, 1]); - assert!(signed.contains(i64::MIN)); - assert!(signed.contains(i64::MAX)); - assert!(!signed.uses_dense_table()); - - let unsigned = IntegerMembership::new(vec![u64::MAX, 0, 1, 2, 3]); - assert!(unsigned.contains(0)); - assert!(unsigned.contains(u64::MAX)); - assert!(!unsigned.uses_dense_table()); - } -} diff --git a/vortex-array/src/scalar_fn/fns/list_contains/mod.rs b/vortex-array/src/scalar_fn/fns/list_contains/mod.rs index 470b81964a4..df1e37ef044 100644 --- a/vortex-array/src/scalar_fn/fns/list_contains/mod.rs +++ b/vortex-array/src/scalar_fn/fns/list_contains/mod.rs @@ -196,6 +196,11 @@ fn compute_list_contains( let nullability = array.dtype().nullability() | value.dtype().nullability(); + if value.all_invalid(ctx)? { + let list_array = array.clone().execute::(ctx)?; + return list_false_if_empty_else_null(&list_array, nullability, ctx); + } + if let Some(value_scalar) = value.as_constant() { list_contains_scalar(array, &value_scalar, nullability, ctx) } else if let Some(list_scalar) = array.as_constant() { @@ -206,7 +211,7 @@ fn compute_list_contains( } /// There is a constant list scalar (haystack) being compared to an array of needles. -fn constant_list_scalar_contains( +pub(crate) fn constant_list_scalar_contains( list_scalar: &ListScalar, values: &ArrayRef, nullability: Nullability, @@ -222,6 +227,7 @@ fn constant_list_scalar_contains( let result = elements .iter() + .filter(|element| !element.is_null()) .map(|element| { Binary::try_new( ConstantArray::new(element.clone(), len).into_array(), @@ -235,9 +241,12 @@ fn constant_list_scalar_contains( .into_iter() .try_reduce_balanced(|acc, res| acc.binary(res, Operator::Or))?; - result - .unwrap_or_else(|| ConstantArray::new(false_scalar, len).into_array()) - .mask(values.validity()?.to_array(len)) + let result = result.unwrap_or_else(|| ConstantArray::new(false_scalar, len).into_array()); + if values.dtype().is_nullable() { + result.mask(values.validity()?.to_array(len)) + } else { + Ok(result) + } } /// Returns a [`BoolArray`] where each bit represents if a list contains the scalar. @@ -285,13 +294,7 @@ fn list_contains_scalar( list_false_or_null(&list_array, nullability) } // No elements match, and all comparisons are valid (result in `false`). - Some(false) => { - // False, but match the nullability to the input list array. - Ok( - ConstantArray::new(Scalar::bool(false, nullability), list_array.len()) - .into_array(), - ) - } + Some(false) => list_false_or_null(&list_array, nullability), // All elements match, and all comparisons are valid (result in `true`). Some(true) => { // True, unless the list itself is empty or NULL. @@ -313,9 +316,9 @@ fn list_contains_scalar( // Process based on the offset and size types. let list_matches = match_each_unsigned_integer_ptype!(offsets.ptype(), |O| { match_each_unsigned_integer_ptype!(sizes.ptype(), |S| { - process_matches::(matches, list_array.len(), offsets, sizes) + process_matches::(&matches, list_array.len(), offsets, sizes, ctx) }) - }); + })?; Ok(BoolArray::new( list_matches, @@ -346,30 +349,36 @@ fn list_false_if_empty_else_null( /// Returns a [`BitBuffer`] where each bit represents if a list contains the scalar, derived from a /// [`BoolArray`] of matches on the child elements array. fn process_matches( - matches: BoolArray, + matches: &BoolArray, list_array_len: usize, offsets: PrimitiveArray, sizes: PrimitiveArray, -) -> BitBuffer + ctx: &mut ExecutionCtx, +) -> VortexResult where O: IntegerPType, S: IntegerPType, { let offsets_slice = offsets.as_slice::(); let sizes_slice = sizes.as_slice::(); - let bits = matches.bit_buffer_view(); + let value_bits = matches.to_bit_buffer(); + let valid_matches = match matches.validity()? { + Validity::NonNullable | Validity::AllValid => value_bits, + Validity::AllInvalid => BitBuffer::new_unset(matches.len()), + validity => value_bits & validity.execute_mask(matches.len(), ctx)?.into_bit_buffer(), + }; - (0..list_array_len) + Ok((0..list_array_len) .map(|i| { let offset = offsets_slice[i].as_(); let size = sizes_slice[i].as_(); // BitIndexIterator yields indices of true bits only. If `.next()` returns // `Some(_)`, at least one element in this list's range matches. - let mut set_bits = BitIndexIterator::new(bits.inner(), offset, size); + let mut set_bits = BitIndexIterator::new(valid_matches.inner(), offset, size); set_bits.next().is_some() }) - .collect::() + .collect::()) } /// Returns a `Bool` array with `false` for lists that are valid, @@ -452,9 +461,7 @@ mod tests { use crate::IntoArray; use crate::VortexSessionExecute; use crate::array_session; - #[cfg(not(codspeed))] use crate::arrays::Dict; - #[cfg(not(codspeed))] use crate::arrays::DictArray; use crate::arrays::ListArray; use crate::arrays::VarBinArray; @@ -583,8 +590,10 @@ mod tests { ); } - #[test] - pub fn test_nullable() { + #[rstest] + #[case::match_present(2, Some(true))] + #[case::match_absent(4, Some(false))] + pub fn test_nullable(#[case] needle: i32, #[case] expected_first: Option) { let arr = ListArray::try_new( PrimitiveArray::from_iter(vec![1, 1, 2, 2, 2]).into_array(), PrimitiveArray::from_iter(vec![0, 5, 5]).into_array(), @@ -593,18 +602,13 @@ mod tests { .unwrap() .into_array(); - let expr = list_contains(root(), lit(2)); + let expr = list_contains(root(), lit(needle)); let item = arr.apply(&expr).unwrap(); - assert_eq!( - item.execute_scalar(0, &mut array_session().create_execution_ctx()) - .unwrap(), - Scalar::bool(true, Nullability::Nullable) - ); - assert!( - !item - .is_valid(1, &mut array_session().create_execution_ctx()) - .unwrap() + assert_arrays_eq!( + item, + BoolArray::from_iter([expected_first, None]), + &mut array_session().create_execution_ctx() ); } @@ -654,7 +658,6 @@ mod tests { } #[test] - #[cfg(not(codspeed))] fn test_dictionary_needles_preserve_dictionary_pushdown() -> VortexResult<()> { let mut ctx = array_session().create_execution_ctx(); let values = PrimitiveArray::from_iter([1i32, 2, 3]).into_array(); @@ -667,6 +670,7 @@ mod tests { ); let contains = needles.apply(&list_contains(lit(list), root()))?; + // Dictionary preservation avoids materializing repeated needle values. assert!(contains.is::()); let actual = contains.execute::(&mut ctx)?; @@ -730,39 +734,9 @@ mod tests { assert_eq!(expr2.to_string(), "vortex.list.contains($, 42i32)"); } - #[test] - pub fn test_constant_scalars() { - let arr = test_array(); - - // Both list and needle are constants - should use scalar optimization - let list_scalar = Scalar::list( - Arc::new(DType::Primitive(I32, Nullability::NonNullable)), - vec![1.into(), 2.into(), 3.into()], - Nullability::NonNullable, - ); - - // Test contains true - let expr = list_contains(lit(list_scalar.clone()), lit(2i32)); - let result = arr.clone().apply(&expr).unwrap(); - assert_eq!( - result - .execute_scalar(0, &mut array_session().create_execution_ctx()) - .unwrap(), - Scalar::bool(true, Nullability::NonNullable) - ); - - // Test contains false - let expr = list_contains(lit(list_scalar), lit(42i32)); - let result = arr.apply(&expr).unwrap(); - assert_eq!( - result - .execute_scalar(0, &mut array_session().create_execution_ctx()) - .unwrap(), - Scalar::bool(false, Nullability::NonNullable) - ); - } - #[rstest] + #[case::present(false, vec![1, 2, 3], Some(2), Some(true))] + #[case::absent(false, vec![1, 2, 3], Some(42), Some(false))] #[case::null_list(true, vec![], Some(1), None)] #[case::empty_list_null_needle(false, vec![], None, Some(false))] #[case::nonempty_list_null_needle(false, vec![1], None, None)] @@ -771,7 +745,7 @@ mod tests { #[case] members: Vec, #[case] needle: Option, #[case] expected: Option, - ) { + ) -> VortexResult<()> { let member_dtype = DType::Primitive(I32, Nullability::NonNullable); let list_dtype = DType::List(Arc::new(member_dtype.clone()), Nullability::Nullable); let list = if null_list { @@ -790,10 +764,17 @@ mod tests { .map(|value| Scalar::bool(value, Nullability::Nullable)) .unwrap_or_else(|| Scalar::null(DType::Bool(Nullability::Nullable))); + let contains = ListContains::try_new( + ConstantArray::new(list, 1).into_array(), + ConstantArray::new(needle, 1).into_array(), + )? + .into_array(); + assert_eq!( - super::compute_contains_scalar(&list, &needle).unwrap(), + contains.execute_scalar(0, &mut array_session().create_execution_ctx())?, expected ); + Ok(()) } // -- Tests migrated from compute/list_contains.rs -- @@ -967,6 +948,27 @@ mod tests { assert_arrays_eq!(contains, expected, &mut ctx); } + #[test] + fn test_nonconstant_all_null_needles() -> VortexResult<()> { + let mut ctx = array_session().create_execution_ctx(); + let lists = ListArray::try_new( + PrimitiveArray::from_iter([1i32]).into_array(), + PrimitiveArray::from_iter([0u32, 0, 1, 1]).into_array(), + Validity::Array(BoolArray::from(BitBuffer::from(vec![true, true, false])).into_array()), + )? + .into_array(); + let needles = PrimitiveArray::from_option_iter::([None, None, None]).into_array(); + + let contains = ListContains::try_new(lists, needles)?.into_array(); + + assert_arrays_eq!( + contains, + BoolArray::from_iter([Some(false), None, None]), + &mut ctx + ); + Ok(()) + } + #[test] fn test_list_array_element() { let mut ctx = array_session().create_execution_ctx(); @@ -1036,8 +1038,9 @@ mod tests { ); assert_arrays_eq!(result, expected, &mut ctx); - // Searching for non-null - let expr2 = list_contains(root(), lit(42i32)); + // Null primitive payloads default to zero. Searching for zero verifies that invalid + // comparison values do not become matches. + let expr2 = list_contains(root(), lit(0i32)); let result2 = list_array.into_array().apply(&expr2).unwrap(); let expected2 = BoolArray::from_iter([false, false, false]); diff --git a/vortex-datafusion/src/convert/exprs.rs b/vortex-datafusion/src/convert/exprs.rs index f62567c25c9..2ab975ecfd7 100644 --- a/vortex-datafusion/src/convert/exprs.rs +++ b/vortex-datafusion/src/convert/exprs.rs @@ -373,23 +373,14 @@ impl ExpressionConvertor for DefaultExpressionConvertor { if let Some(in_list) = df.downcast_ref::() { let value = self.convert(in_list.expr().as_ref())?; - if in_list.is_empty() { - return Err(exec_datafusion_err!("Cannot convert an empty IN list")); - } let list_elements: Vec<_> = in_list .list() .iter() .map(|e| { if let Some(lit) = e.downcast_ref::() { - if lit.value().is_null() { - Err(exec_datafusion_err!( - "Cannot push down an IN list that contains NULL" - )) - } else { - Ok(scalar_from_df(lit.value(), &self.session)) - } + Ok(scalar_from_df(lit.value(), &self.session)) } else { - Err(exec_datafusion_err!("IN list member is not a literal")) + Err(exec_datafusion_err!("Failed to cast sub-expression")) } }) .try_collect()?; @@ -442,19 +433,6 @@ impl ExpressionConvertor for DefaultExpressionConvertor { return Ok(TreeNodeRecursion::Stop); } - if let Some(in_list) = node.downcast_ref::() - && !can_in_list_be_pushed_down(in_list, input_schema) - { - scan_projection.extend( - collect_columns(node) - .into_iter() - .map(|c| (c.name().to_string(), get_item(c.name(), root()))), - ); - - leftover_projection.push(projection_expr.clone()); - return Ok(TreeNodeRecursion::Stop); - } - // DataFusion assumes different decimal types can be coerced. // Vortex expects a perfect match so we don't push it down. if let Some(binary_expr) = node.downcast_ref::() @@ -577,7 +555,11 @@ fn can_be_pushed_down_impl(expr: &Arc, schema: &Schema) -> boo } else if let Some(is_not_null) = expr.downcast_ref::() { can_be_pushed_down_impl(is_not_null.arg(), schema) } else if let Some(in_list) = expr.downcast_ref::() { - can_in_list_be_pushed_down(in_list, schema) + can_be_pushed_down_impl(in_list.expr(), schema) + && in_list + .list() + .iter() + .all(|e| can_be_pushed_down_impl(e, schema)) } else if let Some(scalar_fn) = expr.downcast_ref::() { can_scalar_fn_be_pushed_down(scalar_fn, schema) } else if let Some(case_expr) = expr.downcast_ref::() { @@ -588,17 +570,6 @@ fn can_be_pushed_down_impl(expr: &Arc, schema: &Schema) -> boo } } -fn can_in_list_be_pushed_down(in_list: &df_expr::InListExpr, schema: &Schema) -> bool { - can_be_pushed_down_impl(in_list.expr(), schema) - && !in_list.is_empty() - && in_list.list().iter().all(|expr| { - expr.downcast_ref::() - .is_some_and(|literal| { - !literal.value().is_null() && supported_data_types(&literal.value().data_type()) - }) - }) -} - /// Checks if an expression type is one that convert() can handle. /// This is less restrictive than can_be_pushed_down since it only checks /// expression types, not data type support. @@ -901,53 +872,6 @@ mod tests { assert_snapshot!(result.display_tree().to_string(), @"vortex.literal(42i32)"); } - #[rstest] - #[case::in_list(false)] - #[case::not_in_list(true)] - fn test_null_in_list_is_not_pushed_down(test_schema: Schema, #[case] negated: bool) { - let value = Arc::new(df_expr::Column::new("id", 0)) as Arc; - let list = vec![ - Arc::new(df_expr::Literal::new(ScalarValue::Int32(Some(1)))) as Arc, - Arc::new(df_expr::Literal::new(ScalarValue::Int32(None))) as Arc, - ]; - let expr = - Arc::new(df_expr::InListExpr::try_new(value, list, negated, &test_schema).unwrap()) - as Arc; - let convertor = DefaultExpressionConvertor::default(); - - assert!(!convertor.can_be_pushed_down(&expr, &test_schema)); - assert!( - convertor - .convert(expr.as_ref()) - .unwrap_err() - .to_string() - .contains("IN list that contains NULL") - ); - } - - #[test] - fn test_expr_from_df_in_list() { - let schema = test_schema(); - let value = Arc::new(df_expr::Column::new("id", 0)) as Arc; - let list = [1, 3] - .map(|value| { - Arc::new(df_expr::Literal::new(ScalarValue::Int32(Some(value)))) - as Arc - }) - .to_vec(); - let expr = Arc::new(df_expr::InListExpr::try_new(value, list, false, &schema).unwrap()) - as Arc; - let convertor = DefaultExpressionConvertor::default(); - - assert!(convertor.can_be_pushed_down(&expr, &schema)); - assert_snapshot!(convertor.convert(expr.as_ref()).unwrap().display_tree().to_string(), @r" - vortex.list.contains() - ├── list: vortex.literal([1i32, 3i32]) - └── needle: vortex.get_item(id) - └── input: vortex.root() - "); - } - #[test] fn test_expr_from_df_binary() { let left = Arc::new(df_expr::Column::new("left", 0)) as Arc; diff --git a/vortex-datafusion/src/persistent/tests.rs b/vortex-datafusion/src/persistent/tests.rs index cb6dc818b8e..35dd745461d 100644 --- a/vortex-datafusion/src/persistent/tests.rs +++ b/vortex-datafusion/src/persistent/tests.rs @@ -235,46 +235,6 @@ async fn test_octet_length_pushdown() -> anyhow::Result<()> { Ok(()) } -#[tokio::test] -async fn test_nullable_in_projection_falls_back() -> anyhow::Result<()> { - let ctx = TestSessionContext::new(true); - - ctx.session - .sql( - "CREATE EXTERNAL TABLE nullable_in (id INT) \ - STORED AS vortex LOCATION '/nullable_in/'", - ) - .await?; - ctx.session - .sql("INSERT INTO nullable_in VALUES (1), (2), (NULL)") - .await? - .collect() - .await?; - - let result = ctx - .session - .sql( - "SELECT id, id IN (1, NULL) AS in_result, \ - id NOT IN (1, NULL) AS not_in_result \ - FROM nullable_in ORDER BY id NULLS LAST", - ) - .await? - .collect() - .await?; - - assert_snapshot!(pretty_format_batches(&result)?, @r" - +----+-----------+---------------+ - | id | in_result | not_in_result | - +----+-----------+---------------+ - | 1 | true | false | - | 2 | | | - | | | | - +----+-----------+---------------+ - "); - - Ok(()) -} - #[tokio::test] async fn create_table_ordered_by() -> anyhow::Result<()> { let ctx = TestSessionContext::default(); diff --git a/vortex-duckdb/src/convert/expr.rs b/vortex-duckdb/src/convert/expr.rs index 814aa029f90..51aeebebb6d 100644 --- a/vortex-duckdb/src/convert/expr.rs +++ b/vortex-duckdb/src/convert/expr.rs @@ -27,6 +27,7 @@ use vortex::error::VortexExpect; use vortex::error::VortexResult; use vortex::error::vortex_bail; use vortex::error::vortex_ensure; +use vortex::error::vortex_err; use vortex::expr::Expression; use vortex::expr::and_collect; use vortex::expr::byte_length; @@ -51,6 +52,7 @@ use vortex::scalar_fn::fns::between::StrictComparison; use vortex::scalar_fn::fns::binary::Binary; use vortex::scalar_fn::fns::like::Like; use vortex::scalar_fn::fns::like::LikeOptions; +use vortex::scalar_fn::fns::literal::Literal; use vortex::scalar_fn::fns::operators::Operator; use vortex_spatial::extension::LineString; use vortex_spatial::extension::MultiLineString; @@ -441,27 +443,7 @@ pub fn can_push_expression(value: &duckdb::ExpressionRef) -> bool { ) { return false; } - if matches!( - op.op, - DUCKDB_VX_EXPR_TYPE::DUCKDB_VX_EXPR_TYPE_COMPARE_IN - | DUCKDB_VX_EXPR_TYPE::DUCKDB_VX_EXPR_TYPE_COMPARE_NOT_IN - ) { - let mut children = op.children(); - let Some(element) = children.next() else { - return false; - }; - can_push_expression(element) - && children.all(|child| { - matches!( - child.as_class(), - Some(BoundConstant(constant)) - if Scalar::try_from(constant.value) - .is_ok_and(|scalar| !scalar.is_null()) - ) - }) - } else { - op.children().all(can_push_expression) - } + op.children().all(can_push_expression) } ExpressionClass::BoundAggregate(_) => false, } @@ -735,9 +717,7 @@ fn try_from_compare_in( ) -> VortexResult> { // First child is element, rest form the list. let children: Vec<_> = operator.children().collect(); - if children.len() < 2 { - return Ok(None); - } + assert!(children.len() >= 2); let Some(element) = try_from_expression_inner(children[0], ctx)? else { return Ok(None); }; @@ -745,14 +725,16 @@ fn try_from_compare_in( let Some(list_elements) = children .iter() .skip(1) - .map(|child| { - let Some(BoundConstant(constant)) = child.as_class() else { + .map(|c| { + let Some(value) = try_from_expression_inner(c, ctx)? else { return Ok(None); }; - if constant.value.is_null() { - return Ok(None); - } - Ok(Some(Scalar::try_from(constant.value)?)) + Ok(Some( + value + .as_opt::() + .ok_or_else(|| vortex_err!("cannot have a non literal in a in_list"))? + .clone(), + )) }) .collect::>>>()? else { diff --git a/vortex-duckdb/src/duckdb/value.rs b/vortex-duckdb/src/duckdb/value.rs index b21d38f7c31..8ea5b253a01 100644 --- a/vortex-duckdb/src/duckdb/value.rs +++ b/vortex-duckdb/src/duckdb/value.rs @@ -28,10 +28,6 @@ use crate::lifetime_wrapper; lifetime_wrapper!(Value, cpp::duckdb_value, cpp::duckdb_destroy_value); impl ValueRef { - pub fn is_null(&self) -> bool { - unsafe { cpp::duckdb_is_null_value(self.as_ptr()) } - } - pub fn logical_type(&self) -> &LogicalTypeRef { unsafe { LogicalType::borrow(cpp::duckdb_get_value_type(self.as_ptr())) } } @@ -45,7 +41,7 @@ impl ValueRef { /// Extracts the value from the DuckDB `Value` into a `ExtractedValue`. pub fn extract(&self) -> ExtractedValue { - if self.is_null() { + if unsafe { cpp::duckdb_is_null_value(self.as_ptr()) } { return ExtractedValue::Null; } match self.logical_type().as_type_id() { diff --git a/vortex-duckdb/src/e2e_test/vortex_scan_test.rs b/vortex-duckdb/src/e2e_test/vortex_scan_test.rs index 49b97fa2f80..0876be1ca4c 100644 --- a/vortex-duckdb/src/e2e_test/vortex_scan_test.rs +++ b/vortex-duckdb/src/e2e_test/vortex_scan_test.rs @@ -281,20 +281,6 @@ fn test_issue_5927_not_in_does_not_panic() { assert_eq!(sum, -4); } -#[test] -fn test_not_in_with_null_is_not_pushed_down() { - let file = RUNTIME.block_on(async { - let numbers = buffer![1i32, 42, 100, -5, 0]; - write_single_column_vortex_file("number", numbers).await - }); - let count: i64 = scan_vortex_file_single_row::( - file, - "SELECT COUNT(*) FROM ? WHERE number NOT IN (42, NULL)", - 0, - ); - assert_eq!(count, 0); -} - #[test] fn test_vortex_scan_floats() { let file = RUNTIME.block_on(async { From cc5bbc0548ff70ef6f9a5bcfbc2a51fbc5f3830e Mon Sep 17 00:00:00 2001 From: Will Manning Date: Sat, 29 Aug 2026 18:22:11 -0400 Subject: [PATCH 5/8] refactor(list_contains): Remove unused context plumbing Signed-off-by: Will Manning --- vortex-array/src/arrays/primitive/compute/list_contains.rs | 5 ++--- 1 file changed, 2 insertions(+), 3 deletions(-) diff --git a/vortex-array/src/arrays/primitive/compute/list_contains.rs b/vortex-array/src/arrays/primitive/compute/list_contains.rs index f488734e281..eb3965b5167 100644 --- a/vortex-array/src/arrays/primitive/compute/list_contains.rs +++ b/vortex-array/src/arrays/primitive/compute/list_contains.rs @@ -31,16 +31,15 @@ impl ListContainsElementKernel for Primitive { fn list_contains( list: &ArrayRef, element: ArrayView<'_, Self>, - ctx: &mut ExecutionCtx, + _ctx: &mut ExecutionCtx, ) -> VortexResult> { - evaluate_constant_list_membership(list, element, ctx) + evaluate_constant_list_membership(list, element) } } fn evaluate_constant_list_membership( list: &ArrayRef, element: ArrayView<'_, Primitive>, - _ctx: &mut ExecutionCtx, ) -> VortexResult> { if !element.ptype().is_int() { return Ok(None); From 8982831fb0ae12d56f8bb37ee5803501808435c3 Mon Sep 17 00:00:00 2001 From: Will Manning Date: Sun, 30 Aug 2026 10:09:00 -0400 Subject: [PATCH 6/8] perf(list_contains): Refine prepared membership Signed-off-by: Will Manning --- .../bitpacking/compute/list_contains/mod.rs | 9 ++- .../bitpacking/compute/list_contains/tests.rs | 17 ++++-- .../sequence/src/compute/list_contains.rs | 60 ++++++++++++++++++- .../arrays/primitive/compute/list_contains.rs | 12 ++-- .../src/arrays/primitive/compute/mod.rs | 1 - vortex-array/src/arrays/primitive/mod.rs | 2 - vortex-array/src/scalar/typed_view/list.rs | 5 ++ .../fns/list_contains/integer_membership.rs | 27 ++++----- .../src/scalar_fn/fns/list_contains/kernel.rs | 23 +++---- 9 files changed, 112 insertions(+), 44 deletions(-) diff --git a/encodings/fastlanes/src/bitpacking/compute/list_contains/mod.rs b/encodings/fastlanes/src/bitpacking/compute/list_contains/mod.rs index d708c60a89e..1160912c34c 100644 --- a/encodings/fastlanes/src/bitpacking/compute/list_contains/mod.rs +++ b/encodings/fastlanes/src/bitpacking/compute/list_contains/mod.rs @@ -11,7 +11,6 @@ use vortex_array::IntoArray; use vortex_array::arrays::BoolArray; use vortex_array::arrays::PrimitiveArray; use vortex_array::arrays::primitive::evaluate_prepared_integer_membership; -use vortex_array::arrays::primitive::integer_membership_binary_search_min; use vortex_array::dtype::IntegerPType; use vortex_array::dtype::NativePType; use vortex_array::dtype::PType; @@ -29,6 +28,8 @@ use crate::unpack_iter::BitPacked as BitPackedIter; const MAX_FUSED_DISTINCT_MEMBERS: usize = 4; const SHORT_ARRAY_MAX_ROWS_8_16: usize = 8_192; const SHORT_ARRAY_MAX_ROWS_32: usize = 16_384; +const SHORT_ARRAY_MIN_DECODE_MEMBERS_8_16: usize = 10; +const SHORT_ARRAY_MIN_DECODE_MEMBERS_32: usize = 11; fn min_decode_source_members(ptype: PType, len: usize) -> usize { // The generic fallback scans the packed child once per source member. Decode before repeated // packed scans become more expensive than one decode plus Primitive membership evaluation. @@ -38,7 +39,11 @@ fn min_decode_source_members(ptype: PType, len: usize) -> usize { SHORT_ARRAY_MAX_ROWS_8_16 }; if len <= short_array_max_rows && ptype.bit_width() < 64 { - return integer_membership_binary_search_min(ptype); + return match ptype.bit_width() { + 8 | 16 => SHORT_ARRAY_MIN_DECODE_MEMBERS_8_16, + 32 => SHORT_ARRAY_MIN_DECODE_MEMBERS_32, + _ => unreachable!("short-array policy only applies to 8-, 16-, and 32-bit integers"), + }; } match ptype.bit_width() { 8 => 30, diff --git a/encodings/fastlanes/src/bitpacking/compute/list_contains/tests.rs b/encodings/fastlanes/src/bitpacking/compute/list_contains/tests.rs index 435de238896..a694be3d743 100644 --- a/encodings/fastlanes/src/bitpacking/compute/list_contains/tests.rs +++ b/encodings/fastlanes/src/bitpacking/compute/list_contains/tests.rs @@ -17,6 +17,7 @@ use vortex_array::assert_arrays_eq; use vortex_array::dtype::DType; use vortex_array::dtype::NativePType; use vortex_array::dtype::Nullability; +use vortex_array::dtype::PType; use vortex_array::expr::list_contains; use vortex_array::expr::lit; use vortex_array::expr::root; @@ -161,8 +162,7 @@ fn test_many_member_kernel_policy() -> VortexResult<()> { let values = [0i32, 7, 99]; let primitive = PrimitiveArray::from_iter(values); let packed = BitPackedData::encode(&primitive.into_array(), 7, &mut ctx)?; - let decode_threshold = - super::min_decode_source_members(vortex_array::dtype::PType::I32, packed.len()); + let decode_threshold = super::min_decode_source_members(PType::I32, packed.len()); for (member_count, expected_supported) in [(decode_threshold - 1, false), (decode_threshold, true)] @@ -209,8 +209,14 @@ fn test_zero_bit_width( Ok(()) } -#[test] -fn test_sliced_patched_array() -> VortexResult<()> { +#[rstest] +#[case::fused(vec![3, 100_388])] +#[case::decoded({ + let mut members = (0..31).collect::>(); + members.push(100_388); + members +})] +fn test_sliced_patched_array(#[case] members: Vec) -> VortexResult<()> { let mut ctx = SESSION.create_execution_ctx(); let values = (0..5_000) .map(|index| { @@ -227,9 +233,8 @@ fn test_sliced_patched_array() -> VortexResult<()> { let range = 333..4_333; let sliced = ::slice(packed.as_view(), range.clone(), &mut ctx)? .ok_or_else(|| vortex_err!("BitPacked slice kernel declined a supported input"))?; - let members = [3, 100_388]; let list = list_array( - member_list(members.into_iter().map(Some), Nullability::NonNullable), + member_list(members.iter().copied().map(Some), Nullability::NonNullable), sliced.len(), ); diff --git a/encodings/sequence/src/compute/list_contains.rs b/encodings/sequence/src/compute/list_contains.rs index d2350a1c35a..f31a77ae311 100644 --- a/encodings/sequence/src/compute/list_contains.rs +++ b/encodings/sequence/src/compute/list_contains.rs @@ -5,6 +5,7 @@ use vortex_array::ArrayRef; use vortex_array::ArrayView; use vortex_array::IntoArray; use vortex_array::arrays::BoolArray; +use vortex_array::arrays::Constant; use vortex_array::arrays::ConstantArray; use vortex_array::dtype::DType; use vortex_array::scalar::Scalar; @@ -20,7 +21,7 @@ impl ListContainsElementReduce for Sequence { list: &ArrayRef, element: ArrayView<'_, Self>, ) -> VortexResult> { - let Some(list_scalar) = list.as_constant() else { + let Some(list_array) = list.as_opt::() else { return Ok(None); }; let DType::List(member_dtype, _) = list.dtype() else { @@ -30,7 +31,7 @@ impl ListContainsElementReduce for Sequence { return Ok(None); } - let Some(list_elements) = list_scalar.as_list().elements() else { + let Some(list_elements) = list_array.scalar().as_list().elements() else { return Ok(None); }; if list_elements.is_empty() { @@ -78,14 +79,19 @@ mod tests { use vortex_array::VortexSessionExecute; use vortex_array::arrays::BoolArray; use vortex_array::arrays::Constant; + use vortex_array::arrays::ConstantArray; use vortex_array::assert_arrays_eq; use vortex_array::dtype::DType; use vortex_array::dtype::Nullability; use vortex_array::dtype::PType::I32; + use vortex_array::dtype::PType::I64; use vortex_array::expr::list_contains; use vortex_array::expr::lit; use vortex_array::expr::root; use vortex_array::scalar::Scalar; + use vortex_array::scalar_fn::fns::list_contains::ListContainsElementReduce; + use vortex_error::VortexExpect; + use vortex_error::VortexResult; use vortex_session::VortexSession; use crate::Sequence; @@ -178,4 +184,54 @@ mod tests { &mut SESSION.create_execution_ctx() ); } + + #[test] + fn test_nullable_members() -> VortexResult<()> { + let member_dtype = DType::Primitive(I32, Nullability::Nullable); + let list = ConstantArray::new( + Scalar::list( + Arc::new(member_dtype.clone()), + vec![ + Scalar::primitive(1i32, Nullability::Nullable), + Scalar::null(member_dtype), + Scalar::primitive(3i32, Nullability::Nullable), + ], + Nullability::NonNullable, + ), + 3, + ) + .into_array(); + let sequence = Sequence::try_new_typed(1i32, 1, Nullability::NonNullable, 3)?; + + let result = + ::list_contains(&list, sequence.as_view())? + .vortex_expect("matching integer types are supported"); + + assert_arrays_eq!( + result, + BoolArray::from_iter([true, false, true]), + &mut SESSION.create_execution_ctx() + ); + Ok(()) + } + + #[test] + fn test_wrong_integer_type_declines() -> VortexResult<()> { + let list = ConstantArray::new( + Scalar::list( + Arc::new(DType::Primitive(I64, Nullability::NonNullable)), + vec![1i64.into(), 3i64.into()], + Nullability::NonNullable, + ), + 3, + ) + .into_array(); + let sequence = Sequence::try_new_typed(1i32, 1, Nullability::NonNullable, 3)?; + + let result = + ::list_contains(&list, sequence.as_view())?; + + assert!(result.is_none()); + Ok(()) + } } diff --git a/vortex-array/src/arrays/primitive/compute/list_contains.rs b/vortex-array/src/arrays/primitive/compute/list_contains.rs index eb3965b5167..63a1683a7d9 100644 --- a/vortex-array/src/arrays/primitive/compute/list_contains.rs +++ b/vortex-array/src/arrays/primitive/compute/list_contains.rs @@ -1,6 +1,7 @@ // SPDX-License-Identifier: Apache-2.0 // SPDX-FileCopyrightText: Copyright the Vortex contributors +use vortex_error::VortexExpect; use vortex_error::VortexResult; use crate::ArrayRef; @@ -14,9 +15,7 @@ use crate::scalar_fn::fns::list_contains::IntegerMembership; use crate::scalar_fn::fns::list_contains::ListContainsElementKernel; use crate::scalar_fn::fns::list_contains::constant_list_scalar_contains; -/// Returns the source-member count where Primitive integer membership uses binary search. -#[doc(hidden)] -pub fn integer_membership_binary_search_min(ptype: PType) -> usize { +fn integer_membership_binary_search_min(ptype: PType) -> usize { // The generic implementation evaluates one equality expression per source member. Use the // prepared set once binary search becomes faster than the expression tree. match ptype.bit_width() { @@ -75,7 +74,12 @@ pub fn evaluate_prepared_integer_membership( && membership.non_null_source_len() < integer_membership_binary_search_min(element.ptype()) { return constant_list_scalar_contains( - &membership.source_list().as_list(), + &membership + .source_list() + .as_opt::() + .vortex_expect("membership was prepared from a constant list") + .scalar() + .as_list(), element.array(), nullability, ); diff --git a/vortex-array/src/arrays/primitive/compute/mod.rs b/vortex-array/src/arrays/primitive/compute/mod.rs index 7769def7df5..fec7e246d37 100644 --- a/vortex-array/src/arrays/primitive/compute/mod.rs +++ b/vortex-array/src/arrays/primitive/compute/mod.rs @@ -7,7 +7,6 @@ mod fill_null; mod fixed_width; mod list_contains; pub use list_contains::evaluate_prepared_integer_membership; -pub use list_contains::integer_membership_binary_search_min; mod mask; pub(crate) mod rules; mod slice; diff --git a/vortex-array/src/arrays/primitive/mod.rs b/vortex-array/src/arrays/primitive/mod.rs index b99ec51c0f8..feb55149b08 100644 --- a/vortex-array/src/arrays/primitive/mod.rs +++ b/vortex-array/src/arrays/primitive/mod.rs @@ -16,8 +16,6 @@ pub(crate) mod compute; mod vtable; #[doc(hidden)] pub use compute::evaluate_prepared_integer_membership; -#[doc(hidden)] -pub use compute::integer_membership_binary_search_min; pub use compute::rules::PrimitiveMaskedValidityRule; pub use vtable::Primitive; diff --git a/vortex-array/src/scalar/typed_view/list.rs b/vortex-array/src/scalar/typed_view/list.rs index f97857c92ed..ef9948b4218 100644 --- a/vortex-array/src/scalar/typed_view/list.rs +++ b/vortex-array/src/scalar/typed_view/list.rs @@ -142,6 +142,11 @@ impl<'a> ListScalar<'a> { self.elements.is_none() } + #[inline] + pub(crate) fn values(&self) -> Option<&'a [Option]> { + self.elements + } + /// Returns the data type of the list's elements. pub fn element_dtype(&self) -> &DType { self.dtype diff --git a/vortex-array/src/scalar_fn/fns/list_contains/integer_membership.rs b/vortex-array/src/scalar_fn/fns/list_contains/integer_membership.rs index 1e206870a64..5ac0aac1b7c 100644 --- a/vortex-array/src/scalar_fn/fns/list_contains/integer_membership.rs +++ b/vortex-array/src/scalar_fn/fns/list_contains/integer_membership.rs @@ -2,7 +2,6 @@ // SPDX-FileCopyrightText: Copyright the Vortex contributors use vortex_buffer::BitBuffer; -use vortex_error::VortexExpect; use vortex_error::VortexResult; use vortex_error::vortex_ensure; @@ -10,12 +9,12 @@ use crate::ArrayRef; use crate::ArrayView; use crate::IntoArray; use crate::arrays::BoolArray; +use crate::arrays::Constant; use crate::arrays::Primitive; use crate::dtype::DType; use crate::dtype::IntegerPType; use crate::dtype::NativePType; use crate::dtype::Nullability; -use crate::scalar::Scalar; /// A prepared integer set for constant-list membership kernels. /// @@ -24,11 +23,11 @@ use crate::scalar::Scalar; pub struct IntegerMembership { members: Box<[T]>, non_null_source_len: usize, - source_list: Scalar, + source_list: ArrayRef, } impl IntegerMembership { - fn new(mut members: Vec, source_list: Scalar) -> Self { + fn new(mut members: Vec, source_list: ArrayRef) -> Self { let non_null_source_len = members.len(); members.sort_unstable(); members.dedup(); @@ -44,7 +43,7 @@ impl IntegerMembership { list: &ArrayRef, element_dtype: &DType, ) -> VortexResult> { - let Some(list_scalar) = list.as_constant() else { + let Some(list_array) = list.as_opt::() else { return Ok(None); }; let DType::List(member_dtype, _) = list.dtype() else { @@ -53,23 +52,19 @@ impl IntegerMembership { if !member_dtype.eq_ignore_nullability(element_dtype) { return Ok(None); } - let Some(elements) = list_scalar.as_list().elements() else { + let Some(elements) = list_array.scalar().as_list().values() else { return Ok(None); }; let members = elements .iter() + .filter_map(|value| value.as_ref()) .map(|value| { - value - .as_primitive_opt() - .vortex_expect("list member type was checked") - .try_typed_value::() + // The validated list scalar stores primitive values of `member_dtype`. + value.as_primitive().cast::() }) - .collect::>>>()? - .into_iter() - .flatten() - .collect(); - Ok(Some(Self::new(members, list_scalar))) + .collect::>>()?; + Ok(Some(Self::new(members, list.clone()))) } /// Returns the prepared members. @@ -83,7 +78,7 @@ impl IntegerMembership { self.non_null_source_len } - pub(crate) fn source_list(&self) -> &Scalar { + pub(crate) fn source_list(&self) -> &ArrayRef { &self.source_list } diff --git a/vortex-array/src/scalar_fn/fns/list_contains/kernel.rs b/vortex-array/src/scalar_fn/fns/list_contains/kernel.rs index fc50c8e70ac..ab04ac5ab52 100644 --- a/vortex-array/src/scalar_fn/fns/list_contains/kernel.rs +++ b/vortex-array/src/scalar_fn/fns/list_contains/kernel.rs @@ -9,6 +9,7 @@ use crate::ExecutionCtx; use crate::IntoArray; use crate::array::ArrayView; use crate::array::VTable; +use crate::arrays::Constant; use crate::arrays::ConstantArray; use crate::arrays::ScalarFn; use crate::arrays::scalar_fn::ExactScalarFn; @@ -25,21 +26,24 @@ fn constant_list_result( element_len: usize, element_nullability: crate::dtype::Nullability, ) -> Option { - let list_scalar = list.as_constant()?; + let list_array = list.as_opt::()?; + let list_scalar = list_array.scalar().as_list(); let DType::List(_, list_nullability) = list.dtype() else { return None; }; let nullability = *list_nullability | element_nullability; - match list_scalar.as_list().elements() { - None => Some( + if list_scalar.is_null() { + return Some( ConstantArray::new(Scalar::null(DType::Bool(nullability)), element_len).into_array(), - ), - Some(elements) if elements.is_empty() => { - Some(ConstantArray::new(Scalar::bool(false, nullability), element_len).into_array()) - } - Some(_) => None, + ); + } + if list_scalar.is_empty() { + return Some( + ConstantArray::new(Scalar::bool(false, nullability), element_len).into_array(), + ); } + None } /// Check list-contains without reading buffers (metadata-only). @@ -48,9 +52,6 @@ fn constant_list_result( /// expression. `Self::Array` is the concrete element encoding, while the list (haystack) is /// passed as an opaque `&ArrayRef`. /// -/// A future `ListContainsListReduce` could dispatch on the list side (child 0) for encodings -/// with specialized list representations. -/// /// The parent adaptor resolves null and empty constant lists before delegation. /// /// Return `None` if the operation cannot be resolved from metadata alone. From a2ae8d2b9e1692d2ee363672eb7e30a03d333f35 Mon Sep 17 00:00:00 2001 From: Will Manning Date: Tue, 1 Sep 2026 14:14:09 -0400 Subject: [PATCH 7/8] perf(list_contains): Focus integer membership specialization Keep direct Primitive and FastLanes membership at one to four source members. Preserve the generic expression path beyond that measured boundary. Add fair generic baselines and retain unsafe IN expressions in DataFusion and DuckDB. Signed-off-by: Will Manning --- .../benches/bitpacking_list_contains.rs | 153 +++++++++++++----- .../src/bitpacking/compute/compare_fused.rs | 2 +- .../bitpacking/compute/list_contains/mod.rs | 49 +----- .../bitpacking/compute/list_contains/tests.rs | 67 ++------ vortex-array/benches/list_contains.rs | 128 +++++++++------ .../arrays/primitive/compute/list_contains.rs | 137 ++++++++-------- .../src/arrays/primitive/compute/mod.rs | 1 - vortex-array/src/arrays/primitive/mod.rs | 2 - .../fns/list_contains/integer_membership.rs | 86 +--------- .../src/scalar_fn/fns/list_contains/mod.rs | 26 ++- vortex-datafusion/src/convert/exprs.rs | 92 +++++++++-- vortex-duckdb/src/convert/expr.rs | 37 ++++- .../src/e2e_test/vortex_scan_test.rs | 14 ++ 13 files changed, 448 insertions(+), 346 deletions(-) diff --git a/encodings/fastlanes/benches/bitpacking_list_contains.rs b/encodings/fastlanes/benches/bitpacking_list_contains.rs index 72b4c1beada..32eff584fec 100644 --- a/encodings/fastlanes/benches/bitpacking_list_contains.rs +++ b/encodings/fastlanes/benches/bitpacking_list_contains.rs @@ -3,12 +3,8 @@ //! Measures compressed constant-list membership. //! -//! FastLanes evaluates constant lists with at most four distinct non-null members during unpacking. -//! Mid-size lists use repeated packed comparisons. Larger lists decode once at a threshold that -//! depends on the physical integer width and array length. Every path runs on each real CPU feature -//! leg in CodSpeed. -//! To recalculate the thresholds, temporarily replace `min_decode_source_members` with a constant. -//! Return `usize::MAX` to force repeated comparisons. Return `5` to force decode-once. +//! FastLanes evaluates at most four distinct integer members during unpacking. Larger sets use the +//! frozen generic path. Every path runs on each real CPU feature leg in CodSpeed. //! //! Run with `cargo bench -p vortex-fastlanes --bench bitpacking_list_contains`. @@ -31,8 +27,11 @@ use vortex_array::dtype::DType; use vortex_array::dtype::IntegerPType; use vortex_array::dtype::Nullability; use vortex_array::dtype::PType; +use vortex_array::expr::Expression; +use vortex_array::expr::eq; use vortex_array::expr::list_contains; use vortex_array::expr::lit; +use vortex_array::expr::or; use vortex_array::expr::root; use vortex_array::scalar::Scalar; use vortex_array::validity::Validity; @@ -40,6 +39,7 @@ use vortex_buffer::Alignment; use vortex_buffer::BufferMut; use vortex_fastlanes::BitPacked; use vortex_fastlanes::BitPackedArray; +use vortex_fastlanes::BitPackedArrayExt; use vortex_fastlanes::BitPackedData; use vortex_session::VortexSession; @@ -83,6 +83,7 @@ struct PackedCase { len: usize, member_count: usize, member_stride: u64, + patch_every: Option, } impl Display for PackedCase { @@ -110,33 +111,43 @@ const fn strided_case( len, member_count: count, member_stride: stride, + patch_every: None, + } +} + +const fn patched_case( + name: &'static str, + ptype: PType, + bit_width: u8, + len: usize, + count: usize, + stride: u64, + patch_every: usize, +) -> PackedCase { + PackedCase { + name, + ptype, + bit_width, + len, + member_count: count, + member_stride: stride, + patch_every: Some(patch_every), } } const PACKED_CASES: &[PackedCase] = &[ strided_case("direct_u8_m4", PType::U8, 6, 65_536, 4, 2), - strided_case("direct_u16_m4", PType::U16, 12, 65_536, 4, 2), - strided_case("direct_u32_m4", PType::U32, 20, 65_536, 4, 2), + strided_case("fallback_u8_m5", PType::U8, 6, 65_536, 5, 2), + strided_case("direct_u16_m4", PType::U16, 8, 65_536, 4, 2), + strided_case("fallback_u16_m5", PType::U16, 8, 65_536, 5, 2), + strided_case("direct_u32_m4", PType::U32, 8, 65_536, 4, 2), + strided_case("fallback_u32_m5", PType::U32, 8, 65_536, 5, 2), strided_case("direct_u64_m4", PType::U64, 40, 65_536, 4, 2), - strided_case("generic_u8_m29", PType::U8, 6, 65_536, 29, 2), - strided_case("decode_u8_m30", PType::U8, 6, 65_536, 30, 2), - strided_case("generic_u16_m24", PType::U16, 8, 65_536, 24, 2), - strided_case("decode_u16_m25", PType::U16, 8, 65_536, 25, 2), - strided_case("generic_u32_m12", PType::U32, 8, 65_536, 12, 2), - strided_case("decode_u32_m13", PType::U32, 8, 65_536, 13, 2), - strided_case("decode_u64_m5", PType::U64, 40, 65_536, 5, 2), + strided_case("fallback_u64_m5", PType::U64, 40, 65_536, 5, 2), strided_case("short_direct_u32_m4", PType::U32, 10, 1_024, 4, 2), - strided_case("short_generic_u8_m9", PType::U8, 6, 8_192, 9, 2), - strided_case("short_decode_u8_m10", PType::U8, 6, 8_192, 10, 2), - strided_case("short_generic_u16_m9", PType::U16, 8, 8_192, 9, 2), - strided_case("short_decode_u16_m10", PType::U16, 8, 8_192, 10, 2), - strided_case("short_generic_u32_m10", PType::U32, 8, 16_384, 10, 2), - strided_case("short_decode_u32_m11", PType::U32, 8, 16_384, 11, 2), - strided_case("longer_generic_u8_m10", PType::U8, 6, 16_384, 10, 2), - strided_case("longer_generic_u16_m10", PType::U16, 8, 16_384, 10, 2), - strided_case("longer_generic_u32_m11", PType::U32, 8, 32_768, 11, 2), - strided_case("short_direct_u64_m4", PType::U64, 8, 8_192, 4, 2), - strided_case("short_decode_u64_m5", PType::U64, 8, 8_192, 5, 2), + strided_case("short_fallback_u32_m5", PType::U32, 10, 1_024, 5, 2), + strided_case("wide_direct_u32_m4", PType::U32, 31, 65_536, 4, 2), + patched_case("patch_sparse_u32_m4", PType::U32, 8, 65_536, 4, 2, 64), ]; fn page_aligned(array: BitPackedArray) -> BitPackedArray { @@ -154,22 +165,34 @@ fn page_aligned(array: BitPackedArray) -> BitPackedArray { .unwrap() } -fn generated_values(case: PackedCase, members: &[u64]) -> Vec { +fn generated_values(case: PackedCase, ordinary_members: &[u64]) -> Vec { let domain_size = 1u64 << case.bit_width; + let patch_hit = domain_size; + let patch_miss = patch_hit + 1; let mut state = 0x9E37_79B9_7F4A_7C15u64; (0..case.len) - .map(|_| { + .map(|index| { + if let Some(patch_every) = case.patch_every + && index.is_multiple_of(patch_every) + { + return if (index / patch_every).is_multiple_of(2) { + patch_hit + } else { + patch_miss + }; + } state = state .wrapping_mul(6_364_136_223_846_793_005) .wrapping_add(1_442_695_040_888_963_407); let is_hit = (state >> 32).is_multiple_of(2); if is_hit { let member_index = - usize::try_from(state % u64::try_from(members.len()).unwrap()).unwrap(); - members[member_index] + usize::try_from(state % u64::try_from(ordinary_members.len()).unwrap()) + .unwrap(); + ordinary_members[member_index] } else { let mut candidate = state.rotate_left(17) % domain_size; - while members.contains(&candidate) { + while ordinary_members.contains(&candidate) { candidate = (candidate + 1) % domain_size; } candidate @@ -189,16 +212,52 @@ fn list_scalar(members: &[u64]) -> Scalar { ) } +fn generic_membership_expression(members: &[u64]) -> Expression { + fn balanced_or(expressions: &[Expression]) -> Expression { + assert!(!expressions.is_empty()); + if let [expression] = expressions { + return expression.clone(); + } + let (left, right) = expressions.split_at(expressions.len() / 2); + or(balanced_or(left), balanced_or(right)) + } + + let comparisons = members + .iter() + .map(|member| eq(root(), lit(T::from_counter(*member)))) + .collect::>(); + balanced_or(&comparisons) +} + +fn execute_generic_baseline( + values: &BitPackedArray, + members: &[u64], + ctx: &mut vortex_array::ExecutionCtx, +) -> BoolArray { + values + .clone() + .into_array() + .apply(&generic_membership_expression::(members)) + .unwrap() + .execute::(ctx) + .unwrap() +} + fn packed_input( case: PackedCase, -) -> (BitPackedArray, Scalar, BoolArray, VortexSession) { +) -> (BitPackedArray, Vec, BoolArray, VortexSession) { let session = array_session(); vortex_fastlanes::initialize(&session); let mut ctx = session.create_execution_ctx(); - let members = (0..case.member_count) + let in_domain_member_count = case.member_count - usize::from(case.patch_every.is_some()); + let ordinary_members = (0..in_domain_member_count) .map(|index| u64::try_from(index).unwrap() * case.member_stride) .collect::>(); - let generated = generated_values(case, &members); + let mut members = ordinary_members.clone(); + if case.patch_every.is_some() { + members.push(1u64 << case.bit_width); + } + let generated = generated_values(case, &ordinary_members); let expected = BoolArray::from_iter(generated.iter().map(|value| members.contains(value))); let values: BufferMut = generated.into_iter().map(T::from_counter).collect(); let packed = page_aligned( @@ -209,14 +268,17 @@ fn packed_input( ) .unwrap(), ); - (packed, list_scalar::(&members), expected, session) + if case.patch_every.is_some() { + assert!(packed.patches().is_some()); + } + (packed, members, expected, session) } fn bench_packed_current(bencher: Bencher, case: PackedCase) { - let (packed, list, expected, session) = packed_input::(case); + let (packed, members, expected, session) = packed_input::(case); let contains = packed .into_array() - .apply(&list_contains(lit(list), root())) + .apply(&list_contains(lit(list_scalar::(&members)), root())) .unwrap(); let mut ctx = session.create_execution_ctx(); let actual = contains.clone().execute::(&mut ctx).unwrap(); @@ -226,6 +288,17 @@ fn bench_packed_current(bencher: Bencher, case: PackedCase) { .bench_local(|| black_box(contains.clone().execute::(&mut ctx).unwrap())); } +fn bench_packed_generic_baseline(bencher: Bencher, case: PackedCase) { + let (packed, members, expected, session) = packed_input::(case); + let mut ctx = session.create_execution_ctx(); + let actual = execute_generic_baseline::(&packed, &members, &mut ctx); + assert_arrays_eq!(actual, expected, &mut ctx); + // The pre-change implementation built this comparison tree during execution. + bencher + .counter(ItemsCount::new(case.len)) + .bench_local(|| black_box(execute_generic_baseline::(&packed, &members, &mut ctx))); +} + macro_rules! dispatch_packed { ($bencher:expr, $case:expr, $function:ident) => { match $case.ptype { @@ -243,3 +316,9 @@ macro_rules! dispatch_packed { fn packed_current(bencher: Bencher, case: PackedCase) { dispatch_packed!(bencher, case, bench_packed_current); } + +#[vortex_bench_support::cpu_features] +#[divan::bench(args = PACKED_CASES)] +fn packed_generic_baseline(bencher: Bencher, case: PackedCase) { + dispatch_packed!(bencher, case, bench_packed_generic_baseline); +} diff --git a/encodings/fastlanes/src/bitpacking/compute/compare_fused.rs b/encodings/fastlanes/src/bitpacking/compute/compare_fused.rs index 82a43647fc0..86dd77bff18 100644 --- a/encodings/fastlanes/src/bitpacking/compute/compare_fused.rs +++ b/encodings/fastlanes/src/bitpacking/compute/compare_fused.rs @@ -143,7 +143,7 @@ where .vortex_expect("over-allocated buffer holds a full block per chunk"); // SAFETY: `packed_chunk` holds exactly `128 * bit_width / size_of::()` packed // elements and `bit_width <= U::T`, satisfying `unchecked_unpack_cmp`'s contract. The - // kernel assigns every word in `transposed`, so its previous contents are irrelevant. + // kernel assigns every word in `lane_major`, so its previous contents are irrelevant. unsafe { <::Physical as BitPackingCompare>::unchecked_unpack_cmp::( bit_width, diff --git a/encodings/fastlanes/src/bitpacking/compute/list_contains/mod.rs b/encodings/fastlanes/src/bitpacking/compute/list_contains/mod.rs index 1160912c34c..a631a1b5891 100644 --- a/encodings/fastlanes/src/bitpacking/compute/list_contains/mod.rs +++ b/encodings/fastlanes/src/bitpacking/compute/list_contains/mod.rs @@ -9,15 +9,13 @@ use vortex_array::ArrayView; use vortex_array::ExecutionCtx; use vortex_array::IntoArray; use vortex_array::arrays::BoolArray; -use vortex_array::arrays::PrimitiveArray; -use vortex_array::arrays::primitive::evaluate_prepared_integer_membership; use vortex_array::dtype::IntegerPType; use vortex_array::dtype::NativePType; -use vortex_array::dtype::PType; use vortex_array::dtype::PhysicalPType; use vortex_array::match_each_integer_ptype; use vortex_array::scalar_fn::fns::list_contains::IntegerMembership; use vortex_array::scalar_fn::fns::list_contains::ListContainsElementKernel; +use vortex_array::scalar_fn::fns::list_contains::evaluate_constant_list_generic; use vortex_buffer::BitBuffer; use vortex_error::VortexResult; @@ -25,35 +23,6 @@ use super::compare_fused::stream_predicate_fused; use crate::BitPacked; use crate::unpack_iter::BitPacked as BitPackedIter; -const MAX_FUSED_DISTINCT_MEMBERS: usize = 4; -const SHORT_ARRAY_MAX_ROWS_8_16: usize = 8_192; -const SHORT_ARRAY_MAX_ROWS_32: usize = 16_384; -const SHORT_ARRAY_MIN_DECODE_MEMBERS_8_16: usize = 10; -const SHORT_ARRAY_MIN_DECODE_MEMBERS_32: usize = 11; -fn min_decode_source_members(ptype: PType, len: usize) -> usize { - // The generic fallback scans the packed child once per source member. Decode before repeated - // packed scans become more expensive than one decode plus Primitive membership evaluation. - let short_array_max_rows = if ptype.bit_width() == 32 { - SHORT_ARRAY_MAX_ROWS_32 - } else { - SHORT_ARRAY_MAX_ROWS_8_16 - }; - if len <= short_array_max_rows && ptype.bit_width() < 64 { - return match ptype.bit_width() { - 8 | 16 => SHORT_ARRAY_MIN_DECODE_MEMBERS_8_16, - 32 => SHORT_ARRAY_MIN_DECODE_MEMBERS_32, - _ => unreachable!("short-array policy only applies to 8-, 16-, and 32-bit integers"), - }; - } - match ptype.bit_width() { - 8 => 30, - 16 => 25, - 32 => 13, - 64 => 5, - _ => 5, - } -} - impl ListContainsElementKernel for BitPacked { fn list_contains( list: &ArrayRef, @@ -90,22 +59,8 @@ where { let Some(membership) = IntegerMembership::::try_from_constant_list(list, element.dtype())? else { - return Ok(None); + return evaluate_constant_list_generic(list, element.array(), nullability); }; - if membership.members().len() > MAX_FUSED_DISTINCT_MEMBERS { - if membership.non_null_source_len() - < min_decode_source_members(element.dtype().as_ptype(), element.len()) - { - return Ok(None); - } - // The generic list implementation expands membership into one comparison per source - // member. Each comparison scans the packed child. Decode once before applying the - // Primitive membership policy when repeated packed scans become more expensive. - let primitive = element.array().clone().execute::(ctx)?; - return evaluate_prepared_integer_membership(membership, primitive.as_view(), nullability) - .map(Some); - } - let result = match membership.members() { [] => BoolArray::new( BitBuffer::new_unset(element.len()), diff --git a/encodings/fastlanes/src/bitpacking/compute/list_contains/tests.rs b/encodings/fastlanes/src/bitpacking/compute/list_contains/tests.rs index a694be3d743..3ef71a2cee8 100644 --- a/encodings/fastlanes/src/bitpacking/compute/list_contains/tests.rs +++ b/encodings/fastlanes/src/bitpacking/compute/list_contains/tests.rs @@ -17,7 +17,6 @@ use vortex_array::assert_arrays_eq; use vortex_array::dtype::DType; use vortex_array::dtype::NativePType; use vortex_array::dtype::Nullability; -use vortex_array::dtype::PType; use vortex_array::expr::list_contains; use vortex_array::expr::lit; use vortex_array::expr::root; @@ -26,7 +25,6 @@ use vortex_array::scalar::Scalar; use vortex_array::scalar_fn::fns::list_contains::ListContainsElementKernel; #[cfg(not(codspeed))] use vortex_array::test_harness::trace::trace_op; -use vortex_error::VortexExpect; use vortex_error::VortexResult; use vortex_error::vortex_err; use vortex_session::VortexSession; @@ -104,6 +102,8 @@ integer_type_test!(test_integer_type_u8, u8, 6); integer_type_test!(test_integer_type_u16, u16, 6); integer_type_test!(test_integer_type_u32, u32, 6); integer_type_test!(test_integer_type_u64, u64, 6); +// BitPacked encoding rejects negative integers. These cases verify signed PType dispatch with the +// representable nonnegative domain. integer_type_test!(test_integer_type_i8, i8, 6); integer_type_test!(test_integer_type_i16, i16, 6); integer_type_test!(test_integer_type_i32, i32, 6); @@ -114,7 +114,7 @@ integer_type_test!(test_integer_type_i64, i64, 6); #[case::two(vec![3, 7])] #[case::three(vec![3, 7, 11])] #[case::four(vec![3, 7, 11, 15])] -#[case::duplicate_source(vec![3, 3, 7, 7, 11, 11, 15, 15, 15])] +#[case::duplicate_source(vec![3, 3, 7, 7])] fn test_member_cardinalities(#[case] members: Vec) -> VortexResult<()> { let mut ctx = SESSION.create_execution_ctx(); let values = (0..2_048).map(|value| value % 128).collect::>(); @@ -131,11 +131,10 @@ fn test_member_cardinalities(#[case] members: Vec) -> VortexResult<()> { Ok(()) } -#[rstest] -#[case::generic_five((0..5).map(|value| value * 2).collect())] -#[case::decoded_many((0..32).map(|value| value * 2).collect())] -fn test_many_member_public_expression_paths(#[case] members: Vec) -> VortexResult<()> { +#[test] +fn test_many_member_public_expression_path() -> VortexResult<()> { let mut ctx = SESSION.create_execution_ctx(); + let members = (0..5).map(|value| value * 2).collect::>(); let values = (0..4_096).map(|value| value % 128).collect::>(); let primitive = PrimitiveArray::from_iter(values.iter().copied()); let packed = BitPackedData::encode(&primitive.into_array(), 7, &mut ctx)?; @@ -156,38 +155,6 @@ fn test_many_member_public_expression_paths(#[case] members: Vec) -> Vortex Ok(()) } -#[test] -fn test_many_member_kernel_policy() -> VortexResult<()> { - let mut ctx = SESSION.create_execution_ctx(); - let values = [0i32, 7, 99]; - let primitive = PrimitiveArray::from_iter(values); - let packed = BitPackedData::encode(&primitive.into_array(), 7, &mut ctx)?; - let decode_threshold = super::min_decode_source_members(PType::I32, packed.len()); - - for (member_count, expected_supported) in - [(decode_threshold - 1, false), (decode_threshold, true)] - { - let member_count = i32::try_from(member_count).vortex_expect("member count fits in an i32"); - let list = list_array( - member_list((0..member_count).map(Some), Nullability::NonNullable), - packed.len(), - ); - let actual = ::list_contains( - &list, - packed.as_view(), - &mut ctx, - )?; - - assert_eq!(actual.is_some(), expected_supported); - if let Some(actual) = actual { - let expected = - BoolArray::from_iter(values.map(|value| (0..member_count).contains(&value))); - assert_arrays_eq!(actual, expected, &mut ctx); - } - } - Ok(()) -} - #[rstest] #[case::present([true; 128], vec![0])] #[case::absent([false; 128], vec![1])] @@ -211,8 +178,8 @@ fn test_zero_bit_width( #[rstest] #[case::fused(vec![3, 100_388])] -#[case::decoded({ - let mut members = (0..31).collect::>(); +#[case::fallback({ + let mut members = (0..4).collect::>(); members.push(100_388); members })] @@ -233,18 +200,12 @@ fn test_sliced_patched_array(#[case] members: Vec) -> VortexResult<()> { let range = 333..4_333; let sliced = ::slice(packed.as_view(), range.clone(), &mut ctx)? .ok_or_else(|| vortex_err!("BitPacked slice kernel declined a supported input"))?; - let list = list_array( - member_list(members.iter().copied().map(Some), Nullability::NonNullable), - sliced.len(), - ); + let list = member_list(members.iter().copied().map(Some), Nullability::NonNullable); - let actual = ::list_contains( - &list, - sliced.as_::(), - &mut ctx, - )? - .ok_or_else(|| vortex_err!("BitPacked list_contains kernel declined a sliced input"))? - .execute::(&mut ctx)?; + let actual = sliced + .into_array() + .apply(&list_contains(lit(list), root()))? + .execute::(&mut ctx)?; let expected = BoolArray::from_iter(values[range].iter().map(|value| members.contains(value))); assert_arrays_eq!(actual, expected, &mut ctx); Ok(()) @@ -348,7 +309,7 @@ fn test_registered_kernel_executes_through_expression() -> VortexResult<()> { }) .collect::>(); // A silent fallback preserves values but loses compressed-domain execution. - assert_eq!(applied.len(), 1, "{trace}"); + assert!(!applied.is_empty(), "{trace}"); let expected = BoolArray::from_iter(values.into_iter().map(|value| members.contains(&value))); assert_arrays_eq!(traced.output, expected, &mut ctx); diff --git a/vortex-array/benches/list_contains.rs b/vortex-array/benches/list_contains.rs index 9e67e13678b..511f524440e 100644 --- a/vortex-array/benches/list_contains.rs +++ b/vortex-array/benches/list_contains.rs @@ -3,11 +3,8 @@ //! Compares the Primitive constant-list membership dispatch paths. //! -//! Primitive arrays use direct comparisons for up to four distinct members. They use binary -//! search from 10 members for 8- and 16-bit integers. The 32- and 64-bit thresholds are 11 and 13 -//! members. Every path runs on each real CPU feature leg in CodSpeed. -//! To recalculate the thresholds, run this benchmark twice with temporary policy constants. Use a -//! high cutoff to force generic evaluation. Use `5` to force binary search above four members. +//! Primitive arrays use direct comparisons for at most four distinct integer members. Larger sets +//! use the frozen generic path. Every path runs on each real CPU feature leg in CodSpeed. //! //! Run with `cargo bench -p vortex-array --bench list_contains`. @@ -29,8 +26,11 @@ use vortex_array::assert_arrays_eq; use vortex_array::dtype::DType; use vortex_array::dtype::IntegerPType; use vortex_array::dtype::Nullability; +use vortex_array::expr::Expression; +use vortex_array::expr::eq; use vortex_array::expr::list_contains; use vortex_array::expr::lit; +use vortex_array::expr::or; use vortex_array::expr::root; use vortex_array::scalar::Scalar; use vortex_array::validity::Validity; @@ -95,22 +95,14 @@ const fn primitive_case(name: &'static str, len: usize, member_count: usize) -> const LONG_M1: PrimitiveCase = primitive_case("long", 65_536, 1); const LONG_M4: PrimitiveCase = primitive_case("long", 65_536, 4); -const LONG_M9: PrimitiveCase = primitive_case("long", 65_536, 9); -const LONG_M10: PrimitiveCase = primitive_case("long", 65_536, 10); -const LONG_M11: PrimitiveCase = primitive_case("long", 65_536, 11); -const LONG_M12: PrimitiveCase = primitive_case("long", 65_536, 12); -const LONG_M13: PrimitiveCase = primitive_case("long", 65_536, 13); -const LONG_M32: PrimitiveCase = primitive_case("long", 65_536, 32); -const SHORT_M11: PrimitiveCase = primitive_case("short", 1_024, 11); -const SHORT_M13: PrimitiveCase = primitive_case("short", 1_024, 13); - -const CURRENT_10: &[PrimitiveCase] = &[LONG_M1, LONG_M4, LONG_M9, LONG_M10, LONG_M32, SHORT_M11]; -const CURRENT_11: &[PrimitiveCase] = &[LONG_M1, LONG_M4, LONG_M10, LONG_M11, LONG_M32, SHORT_M11]; -const CURRENT_13: &[PrimitiveCase] = &[LONG_M1, LONG_M4, LONG_M12, LONG_M13, LONG_M32, SHORT_M13]; +const LONG_M5: PrimitiveCase = primitive_case("long", 65_536, 5); +const SHORT_M4: PrimitiveCase = primitive_case("short", 1_024, 4); + +const CURRENT: &[PrimitiveCase] = &[LONG_M1, LONG_M4, LONG_M5, SHORT_M4]; fn primitive_input( case: PrimitiveCase, -) -> (PrimitiveArray, Scalar, BoolArray, VortexSession) { +) -> (PrimitiveArray, Vec, BoolArray, VortexSession) { let members = (0..case.member_count) .map(|index| T::from_counter(u64::try_from(index).unwrap() * 2)) .collect::>(); @@ -136,43 +128,81 @@ fn primitive_input( }) .collect::>(); let expected = BoolArray::from_iter(generated.iter().map(|value| members.contains(value))); - let list = Scalar::list( - Arc::new(DType::Primitive(T::PTYPE, Nullability::NonNullable)), - members.iter().copied().map(Into::into).collect(), - Nullability::NonNullable, - ); ( PrimitiveArray::new::(generated, Validity::NonNullable), - list, + members, expected, array_session(), ) } -fn bench_current(bencher: Bencher, case: PrimitiveCase) { - let (array, list, expected, session) = primitive_input::(case); - let expression = list_contains(lit(list), root()); - let mut ctx = session.create_execution_ctx(); - let actual = array +fn list_scalar(members: &[T]) -> Scalar { + Scalar::list( + Arc::new(DType::Primitive(T::PTYPE, Nullability::NonNullable)), + members.iter().copied().map(Into::into).collect(), + Nullability::NonNullable, + ) +} + +fn generic_membership_expression(members: &[T]) -> Expression { + fn balanced_or(expressions: &[Expression]) -> Expression { + assert!(!expressions.is_empty()); + if let [expression] = expressions { + return expression.clone(); + } + let (left, right) = expressions.split_at(expressions.len() / 2); + or(balanced_or(left), balanced_or(right)) + } + + let comparisons = members + .iter() + .map(|member| { + let member: Scalar = (*member).into(); + eq(root(), lit(member)) + }) + .collect::>(); + balanced_or(&comparisons) +} + +fn execute_generic_baseline( + values: &PrimitiveArray, + members: &[T], + ctx: &mut vortex_array::ExecutionCtx, +) -> BoolArray { + values .clone() .into_array() - .apply(&expression) + .apply(&generic_membership_expression(members)) + .unwrap() + .execute::(ctx) .unwrap() - .execute::(&mut ctx) +} + +fn bench_current(bencher: Bencher, case: PrimitiveCase) { + let (array, members, expected, session) = primitive_input::(case); + let contains = array + .into_array() + .apply(&list_contains(lit(list_scalar(&members)), root())) .unwrap(); + let mut ctx = session.create_execution_ctx(); + let actual = contains.clone().execute::(&mut ctx).unwrap(); assert_arrays_eq!(actual, expected, &mut ctx); - bencher.counter(ItemsCount::new(case.len)).bench_local(|| { - black_box( - array - .clone() - .into_array() - .apply(&expression) - .unwrap() - .execute::(&mut ctx) - .unwrap(), - ) - }); + bencher + .counter(ItemsCount::new(case.len)) + .bench_local(|| black_box(contains.clone().execute::(&mut ctx).unwrap())); +} + +fn bench_generic_baseline(bencher: Bencher, case: PrimitiveCase) { + let (array, members, expected, session) = primitive_input::(case); + let mut ctx = session.create_execution_ctx(); + let actual = execute_generic_baseline(&array, &members, &mut ctx); + assert_arrays_eq!(actual, expected, &mut ctx); + + // The pre-change implementation built this comparison tree during execution. + bencher + .counter(ItemsCount::new(case.len)) + .bench_local(|| black_box(execute_generic_baseline(&array, &members, &mut ctx))); } macro_rules! primitive_benchmarks { @@ -185,11 +215,17 @@ macro_rules! primitive_benchmarks { fn current(bencher: Bencher, case: PrimitiveCase) { bench_current::<$ty>(bencher, case); } + + #[vortex_bench_support::cpu_features] + #[divan::bench(args = $current)] + fn generic_baseline(bencher: Bencher, case: PrimitiveCase) { + bench_generic_baseline::<$ty>(bencher, case); + } } }; } -primitive_benchmarks!(u8_cases, u8, CURRENT_10); -primitive_benchmarks!(u16_cases, u16, CURRENT_10); -primitive_benchmarks!(u32_cases, u32, CURRENT_11); -primitive_benchmarks!(u64_cases, u64, CURRENT_13); +primitive_benchmarks!(u8_cases, u8, CURRENT); +primitive_benchmarks!(u16_cases, u16, CURRENT); +primitive_benchmarks!(u32_cases, u32, CURRENT); +primitive_benchmarks!(u64_cases, u64, CURRENT); diff --git a/vortex-array/src/arrays/primitive/compute/list_contains.rs b/vortex-array/src/arrays/primitive/compute/list_contains.rs index 63a1683a7d9..243c87be23b 100644 --- a/vortex-array/src/arrays/primitive/compute/list_contains.rs +++ b/vortex-array/src/arrays/primitive/compute/list_contains.rs @@ -1,30 +1,21 @@ // SPDX-License-Identifier: Apache-2.0 // SPDX-FileCopyrightText: Copyright the Vortex contributors -use vortex_error::VortexExpect; +use vortex_buffer::BitBuffer; use vortex_error::VortexResult; use crate::ArrayRef; use crate::ArrayView; use crate::ExecutionCtx; +use crate::IntoArray; +use crate::arrays::BoolArray; use crate::arrays::Primitive; use crate::dtype::IntegerPType; -use crate::dtype::PType; +use crate::dtype::NativePType; use crate::match_each_integer_ptype; use crate::scalar_fn::fns::list_contains::IntegerMembership; use crate::scalar_fn::fns::list_contains::ListContainsElementKernel; -use crate::scalar_fn::fns::list_contains::constant_list_scalar_contains; - -fn integer_membership_binary_search_min(ptype: PType) -> usize { - // The generic implementation evaluates one equality expression per source member. Use the - // prepared set once binary search becomes faster than the expression tree. - match ptype.bit_width() { - 8 | 16 => 10, - 32 => 11, - 64 => 13, - _ => 13, - } -} +use crate::scalar_fn::fns::list_contains::evaluate_constant_list_generic; impl ListContainsElementKernel for Primitive { fn list_contains( @@ -44,47 +35,45 @@ fn evaluate_constant_list_membership( return Ok(None); } - let nullability = list.dtype().nullability() | element.dtype().nullability(); - match_each_integer_ptype!(element.ptype(), |T| { - evaluate_integer_membership::(list, element, nullability) + evaluate_integer_membership::(list, element) }) } fn evaluate_integer_membership( list: &ArrayRef, element: ArrayView<'_, Primitive>, - nullability: crate::dtype::Nullability, ) -> VortexResult> { + let nullability = list.dtype().nullability() | element.dtype().nullability(); let Some(membership) = IntegerMembership::::try_from_constant_list(list, element.dtype())? else { - return Ok(None); + return evaluate_constant_list_generic(list, element.array(), nullability); }; - evaluate_prepared_integer_membership(membership, element, nullability).map(Some) + let values = element.as_slice::(); + let bits = match membership.members() { + [] => BitBuffer::new_unset(values.len()), + [member] => collect_direct(values, move |value| value.is_eq(*member)), + [first, second] => collect_direct(values, move |value| { + value.is_eq(*first) | value.is_eq(*second) + }), + [first, second, third] => collect_direct(values, move |value| { + value.is_eq(*first) | value.is_eq(*second) | value.is_eq(*third) + }), + [first, second, third, fourth] => collect_direct(values, move |value| { + value.is_eq(*first) | value.is_eq(*second) | value.is_eq(*third) | value.is_eq(*fourth) + }), + _ => return Ok(None), + }; + Ok(Some( + BoolArray::new(bits, element.validity()?.union_nullability(nullability)).into_array(), + )) } -/// Evaluates a prepared integer set against a Primitive array. -#[doc(hidden)] -pub fn evaluate_prepared_integer_membership( - membership: IntegerMembership, - element: ArrayView<'_, Primitive>, - nullability: crate::dtype::Nullability, -) -> VortexResult { - if membership.members().len() > 4 - && membership.non_null_source_len() < integer_membership_binary_search_min(element.ptype()) - { - return constant_list_scalar_contains( - &membership - .source_list() - .as_opt::() - .vortex_expect("membership was prepared from a constant list") - .scalar() - .as_list(), - element.array(), - nullability, - ); - } - membership.evaluate_primitive(element, nullability) +fn collect_direct(values: &[T], mut predicate: impl FnMut(T) -> bool) -> BitBuffer { + BitBuffer::collect_bool(values.len(), |index| { + // SAFETY: collect_bool visits each valid index once. + predicate(unsafe { *values.get_unchecked(index) }) + }) } #[cfg(test)] @@ -106,7 +95,6 @@ mod tests { use crate::dtype::Nullability; use crate::dtype::PType::F32; use crate::dtype::PType::I32; - use crate::dtype::PType::I64; use crate::expr::list_contains; use crate::expr::lit; use crate::expr::root; @@ -134,10 +122,7 @@ mod tests { #[case::two(vec![3, 7])] #[case::three(vec![3, 7, 11])] #[case::four(vec![3, 7, 11, 15])] - #[case::five((0..5).map(|value| value * 3).collect())] - #[case::eleven((0..11).map(|value| value * 3).collect())] - #[case::many((0..32).map(|value| value * 3).collect())] - #[case::duplicate_heavy((0..32).map(|value| value % 5).collect())] + #[case::duplicate_source(vec![3, 3, 7, 7])] fn test_membership_plans(#[case] members: Vec) -> VortexResult<()> { let mut ctx = crate::array_session().create_execution_ctx(); let values = [0, 3, 7, 15, 31, 90_000, 310_000]; @@ -155,32 +140,20 @@ mod tests { Ok(()) } - #[rstest] - #[case::small(5)] - #[case::many(13)] - fn test_i64_membership(#[case] member_count: usize) -> VortexResult<()> { + #[test] + fn test_five_members_use_generic_fallback() -> VortexResult<()> { let mut ctx = crate::array_session().create_execution_ctx(); - let members = (0..member_count) - .map(|value| i64::try_from(value).vortex_expect("member count fits i64")) - .collect::>(); - let values = [0i64, 11, 99]; + let values = [0i32, 3, 99]; let element = PrimitiveArray::from_iter(values); - let list = ConstantArray::new( - Scalar::list( - Arc::new(DType::Primitive(I64, Nullability::NonNullable)), - members.iter().copied().map(Scalar::from).collect(), - Nullability::NonNullable, - ), - element.len(), - ) - .into_array(); + let members = [0, 3, 6, 9, 12]; let actual = ::list_contains( - &list, + &list(members, element.len()), element.as_view(), &mut ctx, )? - .vortex_expect("integer constant-list membership is supported"); + .vortex_expect("larger constant lists use the generic fallback"); + let expected = BoolArray::from_iter(values.map(|value| members.contains(&value))); assert_arrays_eq!(actual, expected, &mut ctx); Ok(()) @@ -212,7 +185,7 @@ mod tests { }) .collect::>(); // A silent fallback preserves values but loses the membership optimization. - assert_eq!(applied.len(), 1, "{trace}"); + assert!(!applied.is_empty(), "{trace}"); let expected = BoolArray::from_iter(values.map(|value| members.contains(&value))); assert_arrays_eq!(traced.output, expected, &mut ctx); @@ -263,6 +236,36 @@ mod tests { Ok(()) } + #[test] + fn test_nullable_list_preserves_output_nullability() -> VortexResult<()> { + let mut ctx = crate::array_session().create_execution_ctx(); + let list = ConstantArray::new( + Scalar::list( + Arc::new(DType::Primitive(I32, Nullability::NonNullable)), + [1, 3].into_iter().map(Scalar::from).collect(), + Nullability::Nullable, + ), + 3, + ) + .into_array(); + let element = PrimitiveArray::from_iter([1, 2, 3]); + + let actual = ::list_contains( + &list, + element.as_view(), + &mut ctx, + )? + .vortex_expect("integer constant-list membership is supported"); + + assert_eq!(actual.dtype(), &DType::Bool(Nullability::Nullable)); + assert_arrays_eq!( + actual, + BoolArray::from_iter([Some(true), Some(false), Some(true)]), + &mut ctx + ); + Ok(()) + } + #[rstest] #[case::null_list(true)] #[case::empty_list(false)] diff --git a/vortex-array/src/arrays/primitive/compute/mod.rs b/vortex-array/src/arrays/primitive/compute/mod.rs index fec7e246d37..7f1dcdcb4cf 100644 --- a/vortex-array/src/arrays/primitive/compute/mod.rs +++ b/vortex-array/src/arrays/primitive/compute/mod.rs @@ -6,7 +6,6 @@ mod cast; mod fill_null; mod fixed_width; mod list_contains; -pub use list_contains::evaluate_prepared_integer_membership; mod mask; pub(crate) mod rules; mod slice; diff --git a/vortex-array/src/arrays/primitive/mod.rs b/vortex-array/src/arrays/primitive/mod.rs index feb55149b08..748beda8339 100644 --- a/vortex-array/src/arrays/primitive/mod.rs +++ b/vortex-array/src/arrays/primitive/mod.rs @@ -14,8 +14,6 @@ pub use vtable::PrimitiveArray; pub(crate) mod compute; mod vtable; -#[doc(hidden)] -pub use compute::evaluate_prepared_integer_membership; pub use compute::rules::PrimitiveMaskedValidityRule; pub use vtable::Primitive; diff --git a/vortex-array/src/scalar_fn/fns/list_contains/integer_membership.rs b/vortex-array/src/scalar_fn/fns/list_contains/integer_membership.rs index 5ac0aac1b7c..83ea016b6c5 100644 --- a/vortex-array/src/scalar_fn/fns/list_contains/integer_membership.rs +++ b/vortex-array/src/scalar_fn/fns/list_contains/integer_membership.rs @@ -1,20 +1,14 @@ // SPDX-License-Identifier: Apache-2.0 // SPDX-FileCopyrightText: Copyright the Vortex contributors -use vortex_buffer::BitBuffer; use vortex_error::VortexResult; -use vortex_error::vortex_ensure; use crate::ArrayRef; -use crate::ArrayView; -use crate::IntoArray; -use crate::arrays::BoolArray; use crate::arrays::Constant; -use crate::arrays::Primitive; use crate::dtype::DType; use crate::dtype::IntegerPType; -use crate::dtype::NativePType; -use crate::dtype::Nullability; + +const MAX_SOURCE_MEMBERS: usize = 4; /// A prepared integer set for constant-list membership kernels. /// @@ -22,19 +16,14 @@ use crate::dtype::Nullability; #[doc(hidden)] pub struct IntegerMembership { members: Box<[T]>, - non_null_source_len: usize, - source_list: ArrayRef, } impl IntegerMembership { - fn new(mut members: Vec, source_list: ArrayRef) -> Self { - let non_null_source_len = members.len(); + fn new(mut members: Vec) -> Self { members.sort_unstable(); members.dedup(); Self { members: members.into_boxed_slice(), - non_null_source_len, - source_list, } } @@ -55,6 +44,9 @@ impl IntegerMembership { let Some(elements) = list_array.scalar().as_list().values() else { return Ok(None); }; + if elements.len() > MAX_SOURCE_MEMBERS { + return Ok(None); + } let members = elements .iter() @@ -64,75 +56,11 @@ impl IntegerMembership { value.as_primitive().cast::() }) .collect::>>()?; - Ok(Some(Self::new(members, list.clone()))) + Ok(Some(Self::new(members))) } /// Returns the prepared members. pub fn members(&self) -> &[T] { &self.members } - - /// Returns the number of non-null source members before deduplication. - #[doc(hidden)] - pub fn non_null_source_len(&self) -> usize { - self.non_null_source_len - } - - pub(crate) fn source_list(&self) -> &ArrayRef { - &self.source_list - } - - /// Tests whether the prepared set contains `value`. - pub(crate) fn contains(&self, value: T) -> bool { - self.members.binary_search(&value).is_ok() - } - - /// Evaluates this set against a primitive array of the same integer type. - pub(crate) fn evaluate_primitive( - self, - element: ArrayView<'_, Primitive>, - nullability: Nullability, - ) -> VortexResult { - vortex_ensure!( - element.ptype() == T::PTYPE, - "Membership type {} does not match array type {}", - T::PTYPE, - element.ptype(), - ); - let values = element.as_slice::(); - let bits = match self.members() { - [] => BitBuffer::new_unset(values.len()), - [member] => collect_direct(values, move |value| value.is_eq(*member)), - [first, second] => collect_direct(values, move |value| { - value.is_eq(*first) | value.is_eq(*second) - }), - [first, second, third] => collect_direct(values, move |value| { - value.is_eq(*first) | value.is_eq(*second) | value.is_eq(*third) - }), - [first, second, third, fourth] => collect_direct(values, move |value| { - value.is_eq(*first) - | value.is_eq(*second) - | value.is_eq(*third) - | value.is_eq(*fourth) - }), - _ => collect_many(values, &self), - }; - - Ok(BoolArray::new(bits, element.validity()?.union_nullability(nullability)).into_array()) - } -} - -fn collect_direct(values: &[T], mut predicate: impl FnMut(T) -> bool) -> BitBuffer { - BitBuffer::collect_bool_multiversioned(values.len(), |index| { - // SAFETY: collect_bool_multiversioned visits each valid index once. - predicate(unsafe { *values.get_unchecked(index) }) - }) -} - -fn collect_many(values: &[T], membership: &IntegerMembership) -> BitBuffer { - BitBuffer::collect_bool(values.len(), |index| { - // SAFETY: collect_bool visits each valid index once. - let value = unsafe { *values.get_unchecked(index) }; - membership.contains(value) - }) } diff --git a/vortex-array/src/scalar_fn/fns/list_contains/mod.rs b/vortex-array/src/scalar_fn/fns/list_contains/mod.rs index df1e37ef044..1ab67f156c9 100644 --- a/vortex-array/src/scalar_fn/fns/list_contains/mod.rs +++ b/vortex-array/src/scalar_fn/fns/list_contains/mod.rs @@ -210,8 +210,32 @@ fn compute_list_contains( } } +/// Evaluates the generic constant-list path for an encoding-specific kernel. +#[doc(hidden)] +pub fn evaluate_constant_list_generic( + list: &ArrayRef, + values: &ArrayRef, + nullability: Nullability, +) -> VortexResult> { + let Some(list_array) = list.as_opt::() else { + return Ok(None); + }; + let DType::List(member_dtype, _) = list.dtype() else { + return Ok(None); + }; + if !member_dtype.eq_ignore_nullability(values.dtype()) { + return Ok(None); + } + let list_scalar = list_array.scalar().as_list(); + if list_scalar.is_null() { + return Ok(None); + } + + constant_list_scalar_contains(&list_scalar, values, nullability).map(Some) +} + /// There is a constant list scalar (haystack) being compared to an array of needles. -pub(crate) fn constant_list_scalar_contains( +fn constant_list_scalar_contains( list_scalar: &ListScalar, values: &ArrayRef, nullability: Nullability, diff --git a/vortex-datafusion/src/convert/exprs.rs b/vortex-datafusion/src/convert/exprs.rs index 2ab975ecfd7..6356e82fb71 100644 --- a/vortex-datafusion/src/convert/exprs.rs +++ b/vortex-datafusion/src/convert/exprs.rs @@ -384,12 +384,16 @@ impl ExpressionConvertor for DefaultExpressionConvertor { } }) .try_collect()?; + let Some(first) = list_elements.first() else { + return Err(exec_datafusion_err!("Cannot push down an empty IN list")); + }; + if list_elements.iter().any(Scalar::is_null) { + return Err(exec_datafusion_err!( + "Cannot push down an IN list that contains null" + )); + } - let list = Scalar::list( - list_elements[0].dtype().clone(), - list_elements, - Nullability::Nullable, - ); + let list = Scalar::list(first.dtype().clone(), list_elements, Nullability::Nullable); let expr = list_contains(lit(list), value); return Ok(if in_list.negated() { not(expr) } else { expr }); @@ -433,6 +437,17 @@ impl ExpressionConvertor for DefaultExpressionConvertor { return Ok(TreeNodeRecursion::Stop); } + if let Some(in_list) = node.downcast_ref::() + && !can_in_list_be_pushed_down(in_list, input_schema) + { + scan_projection.extend(collect_columns(node).into_iter().map(|column| { + (column.name().to_string(), get_item(column.name(), root())) + })); + + leftover_projection.push(projection_expr.clone()); + return Ok(TreeNodeRecursion::Stop); + } + // DataFusion assumes different decimal types can be coerced. // Vortex expects a perfect match so we don't push it down. if let Some(binary_expr) = node.downcast_ref::() @@ -555,11 +570,7 @@ fn can_be_pushed_down_impl(expr: &Arc, schema: &Schema) -> boo } else if let Some(is_not_null) = expr.downcast_ref::() { can_be_pushed_down_impl(is_not_null.arg(), schema) } else if let Some(in_list) = expr.downcast_ref::() { - can_be_pushed_down_impl(in_list.expr(), schema) - && in_list - .list() - .iter() - .all(|e| can_be_pushed_down_impl(e, schema)) + can_in_list_be_pushed_down(in_list, schema) } else if let Some(scalar_fn) = expr.downcast_ref::() { can_scalar_fn_be_pushed_down(scalar_fn, schema) } else if let Some(case_expr) = expr.downcast_ref::() { @@ -570,6 +581,16 @@ fn can_be_pushed_down_impl(expr: &Arc, schema: &Schema) -> boo } } +fn can_in_list_be_pushed_down(in_list: &df_expr::InListExpr, schema: &Schema) -> bool { + can_be_pushed_down_impl(in_list.expr(), schema) + && !in_list.list().is_empty() + && in_list.list().iter().all(|expr| { + expr.downcast_ref::() + .is_some_and(|literal| !literal.value().is_null()) + && can_be_pushed_down_impl(expr, schema) + }) +} + /// Checks if an expression type is one that convert() can handle. /// This is less restrictive than can_be_pushed_down since it only checks /// expression types, not data type support. @@ -787,6 +808,19 @@ mod tests { ) } + fn in_list_expr( + values: impl IntoIterator, + negated: bool, + schema: &Schema, + ) -> Arc { + let value = Arc::new(df_expr::Column::new("id", 0)) as Arc; + let list = values + .into_iter() + .map(|value| Arc::new(df_expr::Literal::new(value)) as Arc) + .collect(); + Arc::new(df_expr::InListExpr::try_new(value, list, negated, schema).unwrap()) + } + #[test] fn test_make_vortex_predicate_empty() { let expr_convertor = DefaultExpressionConvertor::default(); @@ -1079,6 +1113,44 @@ mod tests { assert!(!can_be_pushed_down_impl(&binary_expr, &test_schema)); } + #[rstest] + #[case::in_nonempty(vec![ScalarValue::Int32(Some(1))], false, true)] + #[case::not_in_nonempty(vec![ScalarValue::Int32(Some(1))], true, true)] + #[case::in_empty(vec![], false, false)] + #[case::not_in_empty(vec![], true, false)] + #[case::in_null(vec![ScalarValue::Int32(None)], false, false)] + #[case::not_in_null(vec![ScalarValue::Int32(None)], true, false)] + fn test_can_be_pushed_down_in_list( + #[case] values: Vec, + #[case] negated: bool, + #[case] expected: bool, + test_schema: Schema, + ) { + let expression = in_list_expr(values, negated, &test_schema); + + assert_eq!(can_be_pushed_down_impl(&expression, &test_schema), expected); + } + + #[rstest] + #[case::empty(vec![], false)] + #[case::null(vec![ScalarValue::Int32(None)], true)] + fn test_split_projection_keeps_unsafe_in_list_in_datafusion( + #[case] values: Vec, + #[case] negated: bool, + test_schema: Schema, + ) { + let expression = in_list_expr(values, negated, &test_schema); + let source_projection = + ProjectionExprs::new([ProjectionExpr::new(expression, "matches".to_string())]); + let output_schema = Schema::new(vec![Field::new("matches", DataType::Boolean, true)]); + + let processed = DefaultExpressionConvertor::default() + .split_projection(source_projection.clone(), &test_schema, &output_schema) + .unwrap(); + + assert_eq!(processed.leftover_projection, source_projection); + } + #[rstest] fn test_can_be_pushed_down_like_supported(test_schema: Schema) { let expr = Arc::new(df_expr::Column::new("name", 1)) as Arc; diff --git a/vortex-duckdb/src/convert/expr.rs b/vortex-duckdb/src/convert/expr.rs index 51aeebebb6d..cc5d490350f 100644 --- a/vortex-duckdb/src/convert/expr.rs +++ b/vortex-duckdb/src/convert/expr.rs @@ -386,6 +386,31 @@ fn can_push_cast(cast: &duckdb::BoundCast<'_>, target: &duckdb::LogicalTypeRef) !cast.is_try && target.is_primitive_integer() && cast.child.return_type().is_primitive_integer() } +fn can_push_in_list(operator: &BoundOperator<'_>) -> bool { + let mut children = operator.children(); + let Some(element) = children.next() else { + return false; + }; + if !can_push_expression(element) { + return false; + } + + let mut member_count = 0; + for child in children { + let Some(BoundConstant(constant)) = child.as_class() else { + return false; + }; + let Ok(member) = Scalar::try_from(constant.value) else { + return false; + }; + if member.is_null() { + return false; + } + member_count += 1; + } + member_count > 0 +} + // Called before pushdown_complex_filter or a table filter expression call. // As we support complex filter pushdown, Duckdb pushes expressions to Vortex. // However, it doesn't know what type of expressions we can handle. Here we list @@ -433,13 +458,18 @@ pub fn can_push_expression(value: &duckdb::ExpressionRef) -> bool { // columns are native. } ExpressionClass::BoundOperator(op) => { + if matches!( + op.op, + DUCKDB_VX_EXPR_TYPE::DUCKDB_VX_EXPR_TYPE_COMPARE_IN + | DUCKDB_VX_EXPR_TYPE::DUCKDB_VX_EXPR_TYPE_COMPARE_NOT_IN + ) { + return can_push_in_list(&op); + } if !matches!( op.op, DUCKDB_VX_EXPR_TYPE::DUCKDB_VX_EXPR_TYPE_OPERATOR_NOT | DUCKDB_VX_EXPR_TYPE::DUCKDB_VX_EXPR_TYPE_OPERATOR_IS_NULL | DUCKDB_VX_EXPR_TYPE::DUCKDB_VX_EXPR_TYPE_OPERATOR_IS_NOT_NULL - | DUCKDB_VX_EXPR_TYPE::DUCKDB_VX_EXPR_TYPE_COMPARE_IN - | DUCKDB_VX_EXPR_TYPE::DUCKDB_VX_EXPR_TYPE_COMPARE_NOT_IN ) { return false; } @@ -740,6 +770,9 @@ fn try_from_compare_in( else { return Ok(None); }; + if list_elements.is_empty() || list_elements.iter().any(Scalar::is_null) { + return Ok(None); + } let list = Scalar::list( Arc::new(list_elements[0].dtype().clone()), list_elements, diff --git a/vortex-duckdb/src/e2e_test/vortex_scan_test.rs b/vortex-duckdb/src/e2e_test/vortex_scan_test.rs index 0876be1ca4c..49b97fa2f80 100644 --- a/vortex-duckdb/src/e2e_test/vortex_scan_test.rs +++ b/vortex-duckdb/src/e2e_test/vortex_scan_test.rs @@ -281,6 +281,20 @@ fn test_issue_5927_not_in_does_not_panic() { assert_eq!(sum, -4); } +#[test] +fn test_not_in_with_null_is_not_pushed_down() { + let file = RUNTIME.block_on(async { + let numbers = buffer![1i32, 42, 100, -5, 0]; + write_single_column_vortex_file("number", numbers).await + }); + let count: i64 = scan_vortex_file_single_row::( + file, + "SELECT COUNT(*) FROM ? WHERE number NOT IN (42, NULL)", + 0, + ); + assert_eq!(count, 0); +} + #[test] fn test_vortex_scan_floats() { let file = RUNTIME.block_on(async { From 129789661890ae0801aa3e9cbbdf211a4781b59d Mon Sep 17 00:00:00 2001 From: Will Manning Date: Tue, 1 Sep 2026 17:31:41 -0400 Subject: [PATCH 8/8] perf(list_contains): Refine specialization boundaries Signed-off-by: Will Manning --- .../benches/bitpacking_list_contains.rs | 54 ++++++++++++------- .../sequence/src/compute/list_contains.rs | 29 ++++++++++ vortex-array/benches/list_contains.rs | 52 +++++++++++------- .../src/scalar_fn/fns/list_contains/mod.rs | 42 +++++++++++++++ vortex-datafusion/src/convert/exprs.rs | 31 ++++++++--- 5 files changed, 166 insertions(+), 42 deletions(-) diff --git a/encodings/fastlanes/benches/bitpacking_list_contains.rs b/encodings/fastlanes/benches/bitpacking_list_contains.rs index 32eff584fec..27439139eda 100644 --- a/encodings/fastlanes/benches/bitpacking_list_contains.rs +++ b/encodings/fastlanes/benches/bitpacking_list_contains.rs @@ -17,26 +17,29 @@ use std::sync::Arc; use divan::Bencher; use divan::counter::ItemsCount; +use vortex_array::ArrayRef; use vortex_array::IntoArray; use vortex_array::VortexSessionExecute; use vortex_array::array_session; use vortex_array::arrays::BoolArray; +use vortex_array::arrays::ConstantArray; use vortex_array::arrays::PrimitiveArray; use vortex_array::assert_arrays_eq; +use vortex_array::builtins::ArrayBuiltins; use vortex_array::dtype::DType; use vortex_array::dtype::IntegerPType; use vortex_array::dtype::Nullability; use vortex_array::dtype::PType; -use vortex_array::expr::Expression; -use vortex_array::expr::eq; use vortex_array::expr::list_contains; use vortex_array::expr::lit; -use vortex_array::expr::or; use vortex_array::expr::root; use vortex_array::scalar::Scalar; +use vortex_array::scalar_fn::fns::binary::Binary; +use vortex_array::scalar_fn::fns::operators::Operator; use vortex_array::validity::Validity; use vortex_buffer::Alignment; use vortex_buffer::BufferMut; +use vortex_error::VortexResult; use vortex_fastlanes::BitPacked; use vortex_fastlanes::BitPackedArray; use vortex_fastlanes::BitPackedArrayExt; @@ -212,21 +215,39 @@ fn list_scalar(members: &[u64]) -> Scalar { ) } -fn generic_membership_expression(members: &[u64]) -> Expression { - fn balanced_or(expressions: &[Expression]) -> Expression { - assert!(!expressions.is_empty()); - if let [expression] = expressions { - return expression.clone(); +fn frozen_generic_membership( + values: ArrayRef, + members: &[u64], +) -> VortexResult { + fn balanced_or(arrays: &[ArrayRef]) -> VortexResult { + if let [array] = arrays { + return Ok(array.clone()); } - let (left, right) = expressions.split_at(expressions.len() / 2); - or(balanced_or(left), balanced_or(right)) + let (left, right) = arrays.split_at(arrays.len() / 2); + balanced_or(left)?.binary(balanced_or(right)?, Operator::Or) } + let len = values.len(); + let nullability = values.dtype().nullability(); + let false_scalar = Scalar::bool(false, nullability); let comparisons = members .iter() - .map(|member| eq(root(), lit(T::from_counter(*member)))) - .collect::>(); - balanced_or(&comparisons) + .map(|member| { + Binary::try_new( + ConstantArray::new(T::from_counter(*member).into(), len).into_array(), + values.clone(), + Operator::Eq, + )? + .into_array() + .fill_null(false_scalar.clone()) + }) + .collect::>>()?; + + if comparisons.is_empty() { + Ok(ConstantArray::new(false_scalar, len).into_array()) + } else { + balanced_or(&comparisons) + } } fn execute_generic_baseline( @@ -234,10 +255,7 @@ fn execute_generic_baseline( members: &[u64], ctx: &mut vortex_array::ExecutionCtx, ) -> BoolArray { - values - .clone() - .into_array() - .apply(&generic_membership_expression::(members)) + frozen_generic_membership::(values.clone().into_array(), members) .unwrap() .execute::(ctx) .unwrap() @@ -293,7 +311,7 @@ fn bench_packed_generic_baseline(bencher: Bencher, case: PackedCase let mut ctx = session.create_execution_ctx(); let actual = execute_generic_baseline::(&packed, &members, &mut ctx); assert_arrays_eq!(actual, expected, &mut ctx); - // The pre-change implementation built this comparison tree during execution. + // The frozen pre-change implementation built this array tree during execution. bencher .counter(ItemsCount::new(case.len)) .bench_local(|| black_box(execute_generic_baseline::(&packed, &members, &mut ctx))); diff --git a/encodings/sequence/src/compute/list_contains.rs b/encodings/sequence/src/compute/list_contains.rs index f31a77ae311..34cc5b276e4 100644 --- a/encodings/sequence/src/compute/list_contains.rs +++ b/encodings/sequence/src/compute/list_contains.rs @@ -63,6 +63,12 @@ impl ListContainsElementReduce for Sequence { } } + if set_indices.is_empty() { + return Ok(Some( + ConstantArray::new(Scalar::bool(false, nullability), element.len()).into_array(), + )); + } + Ok(Some( BoolArray::from_indices(element.len(), set_indices, nullability.into()).into_array(), )) @@ -157,6 +163,29 @@ mod tests { assert_arrays_eq!(result, expected, &mut SESSION.create_execution_ctx()); } + #[test] + fn test_no_intersection_reduces_to_constant() { + let list_scalar = Scalar::list( + Arc::new(I32.into()), + vec![7.into(), 42.into()], + Nullability::NonNullable, + ); + let array = Sequence::try_new_typed(1i32, 1, Nullability::NonNullable, 3) + .unwrap() + .into_array(); + + let result = array + .apply(&list_contains(lit(list_scalar), root())) + .unwrap(); + + assert!(result.is::()); + assert_arrays_eq!( + result, + BoolArray::from_iter([false, false, false]), + &mut SESSION.create_execution_ctx() + ); + } + #[rstest] #[case::null_list( Scalar::null(DType::List(Arc::new(I32.into()), Nullability::Nullable)), diff --git a/vortex-array/benches/list_contains.rs b/vortex-array/benches/list_contains.rs index 511f524440e..725c2125682 100644 --- a/vortex-array/benches/list_contains.rs +++ b/vortex-array/benches/list_contains.rs @@ -17,23 +17,26 @@ use std::sync::Arc; use divan::Bencher; use divan::counter::ItemsCount; +use vortex_array::ArrayRef; use vortex_array::IntoArray; use vortex_array::VortexSessionExecute; use vortex_array::array_session; use vortex_array::arrays::BoolArray; +use vortex_array::arrays::ConstantArray; use vortex_array::arrays::PrimitiveArray; use vortex_array::assert_arrays_eq; +use vortex_array::builtins::ArrayBuiltins; use vortex_array::dtype::DType; use vortex_array::dtype::IntegerPType; use vortex_array::dtype::Nullability; -use vortex_array::expr::Expression; -use vortex_array::expr::eq; use vortex_array::expr::list_contains; use vortex_array::expr::lit; -use vortex_array::expr::or; use vortex_array::expr::root; use vortex_array::scalar::Scalar; +use vortex_array::scalar_fn::fns::binary::Binary; +use vortex_array::scalar_fn::fns::operators::Operator; use vortex_array::validity::Validity; +use vortex_error::VortexResult; use vortex_session::VortexSession; fn main() { @@ -144,24 +147,40 @@ fn list_scalar(members: &[T]) -> Scalar { ) } -fn generic_membership_expression(members: &[T]) -> Expression { - fn balanced_or(expressions: &[Expression]) -> Expression { - assert!(!expressions.is_empty()); - if let [expression] = expressions { - return expression.clone(); +fn frozen_generic_membership( + values: ArrayRef, + members: &[T], +) -> VortexResult { + fn balanced_or(arrays: &[ArrayRef]) -> VortexResult { + if let [array] = arrays { + return Ok(array.clone()); } - let (left, right) = expressions.split_at(expressions.len() / 2); - or(balanced_or(left), balanced_or(right)) + let (left, right) = arrays.split_at(arrays.len() / 2); + balanced_or(left)?.binary(balanced_or(right)?, Operator::Or) } + let len = values.len(); + let nullability = values.dtype().nullability(); + let false_scalar = Scalar::bool(false, nullability); let comparisons = members .iter() .map(|member| { let member: Scalar = (*member).into(); - eq(root(), lit(member)) + Binary::try_new( + ConstantArray::new(member, len).into_array(), + values.clone(), + Operator::Eq, + )? + .into_array() + .fill_null(false_scalar.clone()) }) - .collect::>(); - balanced_or(&comparisons) + .collect::>>()?; + + if comparisons.is_empty() { + Ok(ConstantArray::new(false_scalar, len).into_array()) + } else { + balanced_or(&comparisons) + } } fn execute_generic_baseline( @@ -169,10 +188,7 @@ fn execute_generic_baseline( members: &[T], ctx: &mut vortex_array::ExecutionCtx, ) -> BoolArray { - values - .clone() - .into_array() - .apply(&generic_membership_expression(members)) + frozen_generic_membership(values.clone().into_array(), members) .unwrap() .execute::(ctx) .unwrap() @@ -199,7 +215,7 @@ fn bench_generic_baseline(bencher: Bencher, case: PrimitiveCase) { let actual = execute_generic_baseline(&array, &members, &mut ctx); assert_arrays_eq!(actual, expected, &mut ctx); - // The pre-change implementation built this comparison tree during execution. + // The frozen pre-change implementation built this array tree during execution. bencher .counter(ItemsCount::new(case.len)) .bench_local(|| black_box(execute_generic_baseline(&array, &members, &mut ctx))); diff --git a/vortex-array/src/scalar_fn/fns/list_contains/mod.rs b/vortex-array/src/scalar_fn/fns/list_contains/mod.rs index 1ab67f156c9..6e984c17b79 100644 --- a/vortex-array/src/scalar_fn/fns/list_contains/mod.rs +++ b/vortex-array/src/scalar_fn/fns/list_contains/mod.rs @@ -197,6 +197,14 @@ fn compute_list_contains( let nullability = array.dtype().nullability() | value.dtype().nullability(); if value.all_invalid(ctx)? { + if let Some(list_scalar) = array.as_constant() { + let result = if list_scalar.as_list().is_empty() { + Scalar::bool(false, nullability) + } else { + Scalar::null(DType::Bool(nullability)) + }; + return Ok(ConstantArray::new(result, array.len()).into_array()); + } let list_array = array.clone().execute::(ctx)?; return list_false_if_empty_else_null(&list_array, nullability, ctx); } @@ -485,6 +493,7 @@ mod tests { use crate::IntoArray; use crate::VortexSessionExecute; use crate::array_session; + use crate::arrays::Constant; use crate::arrays::Dict; use crate::arrays::DictArray; use crate::arrays::ListArray; @@ -950,6 +959,39 @@ mod tests { assert_arrays_eq!(result, expected, &mut ctx); } + #[rstest] + #[case::empty(Vec::<&'static str>::new(), Some(false))] + #[case::nonempty(vec!["a"], None)] + fn test_constant_string_list_all_null_needles_reduces_to_constant( + #[case] members: Vec<&'static str>, + #[case] expected: Option, + ) { + let mut ctx = array_session().create_execution_ctx(); + let list = Scalar::list( + Arc::new(DType::Utf8(Nullability::NonNullable)), + members.into_iter().map(Scalar::from).collect(), + Nullability::NonNullable, + ); + let needles = VarBinArray::from_iter( + [None::<&str>, None, None], + DType::Utf8(Nullability::Nullable), + ) + .into_array(); + + let result = needles + .apply(&list_contains(lit(list), root())) + .unwrap() + .execute::(&mut ctx) + .unwrap(); + + assert!(result.is::()); + assert_arrays_eq!( + result, + BoolArray::from_iter([expected, expected, expected]), + &mut ctx + ); + } + #[test] fn test_all_nulls() { let mut ctx = array_session().create_execution_ctx(); diff --git a/vortex-datafusion/src/convert/exprs.rs b/vortex-datafusion/src/convert/exprs.rs index 6356e82fb71..c227663f10d 100644 --- a/vortex-datafusion/src/convert/exprs.rs +++ b/vortex-datafusion/src/convert/exprs.rs @@ -583,17 +583,25 @@ fn can_be_pushed_down_impl(expr: &Arc, schema: &Schema) -> boo fn can_in_list_be_pushed_down(in_list: &df_expr::InListExpr, schema: &Schema) -> bool { can_be_pushed_down_impl(in_list.expr(), schema) - && !in_list.list().is_empty() + && is_convertible_in_list(in_list) + && in_list + .list() + .iter() + .all(|expr| can_be_pushed_down_impl(expr, schema)) +} + +fn is_convertible_in_list(in_list: &df_expr::InListExpr) -> bool { + !in_list.list().is_empty() && in_list.list().iter().all(|expr| { expr.downcast_ref::() .is_some_and(|literal| !literal.value().is_null()) - && can_be_pushed_down_impl(expr, schema) }) } -/// Checks if an expression type is one that convert() can handle. -/// This is less restrictive than can_be_pushed_down since it only checks -/// expression types, not data type support. +/// Checks if an expression can be converted without schema information. +/// +/// This is less restrictive than `can_be_pushed_down_impl` because it does not check data type +/// support. fn is_convertible_expr(expr: &Arc) -> bool { // Expression types that convert() handles expr.downcast_ref::().is_some() @@ -605,7 +613,9 @@ fn is_convertible_expr(expr: &Arc) -> bool { .is_some_and(|e| is_convertible_expr(e.expr())) || expr.downcast_ref::().is_some() || expr.downcast_ref::().is_some() - || expr.downcast_ref::().is_some() + || expr + .downcast_ref::() + .is_some_and(is_convertible_in_list) || expr.downcast_ref::().is_some_and(|sf| { ScalarFunctionExpr::try_downcast_func::(sf).is_some() || ScalarFunctionExpr::try_downcast_func::(sf).is_some() @@ -1127,8 +1137,17 @@ mod tests { test_schema: Schema, ) { let expression = in_list_expr(values, negated, &test_schema); + let cast_expression = Arc::new(df_expr::CastExpr::new_with_target_field( + Arc::clone(&expression), + Arc::new(Field::new("matches", DataType::Boolean, true)), + None, + )) as Arc; assert_eq!(can_be_pushed_down_impl(&expression, &test_schema), expected); + assert_eq!( + can_be_pushed_down_impl(&cast_expression, &test_schema), + expected + ); } #[rstest]