From 1e3454437ebc04046bba26eac76a97e9e064b53e Mon Sep 17 00:00:00 2001 From: Will Manning Date: Thu, 27 Aug 2026 17:17:26 -0400 Subject: [PATCH 1/5] 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/5] 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/5] 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/5] 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/5] 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);