diff --git a/encodings/fastlanes/Cargo.toml b/encodings/fastlanes/Cargo.toml index 9085390b67b..1127bf37bed 100644 --- a/encodings/fastlanes/Cargo.toml +++ b/encodings/fastlanes/Cargo.toml @@ -48,6 +48,10 @@ _test-harness = ["dep:rand"] name = "bitpacking_take" harness = false +[[bench]] +name = "bitpacking_list_contains" +harness = false + [[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..7479555a321 --- /dev/null +++ b/encodings/fastlanes/benches/bitpacking_list_contains.rs @@ -0,0 +1,173 @@ +// SPDX-License-Identifier: Apache-2.0 +// SPDX-FileCopyrightText: Copyright the Vortex contributors + +//! Compares compressed list membership with the canonical fallback. +//! +//! 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`. + +#![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::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::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_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() +} + +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 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(), + 10, + &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, + ); + 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)) + .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); +} + +#[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 new file mode 100644 index 00000000000..8c7138a2313 --- /dev/null +++ b/encodings/fastlanes/src/bitpacking/compute/list_contains/mod.rs @@ -0,0 +1,142 @@ +// 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::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; + +impl ListContainsElementKernel for BitPacked { + fn list_contains( + list: &ArrayRef, + element: ArrayView<'_, Self>, + ctx: &mut ExecutionCtx, + ) -> VortexResult> { + list_contains_compressed(list, element, ctx) + } +} + +fn list_contains_compressed( + list: &ArrayRef, + 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 { + 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 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::>(); + + 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; + 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() { + 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(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..4cfd5ebe2c7 --- /dev/null +++ b/encodings/fastlanes/src/bitpacking/compute/list_contains/tests.rs @@ -0,0 +1,340 @@ +// 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; +#[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; +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::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(); + 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_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(); + 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] +#[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::>(); + 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/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), + ); } 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();