diff --git a/vortex-array/src/arrays/interleave/execute/bool.rs b/vortex-array/src/arrays/interleave/execute/bool.rs index fde5b161dfd..a051f55ec5d 100644 --- a/vortex-array/src/arrays/interleave/execute/bool.rs +++ b/vortex-array/src/arrays/interleave/execute/bool.rs @@ -7,10 +7,10 @@ use num_traits::AsPrimitive; use vortex_buffer::BitBuffer; use vortex_buffer::BitBufferMut; use vortex_error::VortexResult; -use vortex_error::vortex_ensure; use super::super::Interleave; use super::super::InterleaveArrayExt; +use super::validate_selectors; use crate::array::Array; use crate::arrays::Bool; use crate::arrays::BoolArray; @@ -71,46 +71,18 @@ fn gather, R: AsPrimitive>( branches: &[A], rows: &[R], ) -> VortexResult { - let len = validate_selectors(value_bits, branches, rows)?; + let len = validate_selectors( + value_bits.len(), + |branch| value_bits[branch].len(), + branches, + rows, + )?; // SAFETY: `validate_selectors` proved `branches.len() == rows.len() == len`, and for every // `i < len` that `branches[i] < value_bits.len()` and `rows[i] < value_bits[branches[i]].len()`. Ok(unsafe { gather_bits(len, value_bits, branches, rows) }) } -/// Validates the per-row selector bounds, returning the output length (`branches.len()`). -/// -/// On success, `rows.len() == branches.len() == len` and, for every `i < len`, -/// `branches[i] < value_bits.len()` and `rows[i] < value_bits[branches[i]].len()` — exactly the -/// preconditions of [`gather_bits`]. Errors (rather than panics) on any out-of-bounds selector. -fn validate_selectors, R: AsPrimitive>( - value_bits: &[BitBuffer], - branches: &[A], - rows: &[R], -) -> VortexResult { - // The two selectors are validated to equal length at construction, which is the output length. - let len = branches.len(); - vortex_ensure!( - rows.len() == len, - "interleave selectors differ in length: array_indices {len}, row_indices {}", - rows.len() - ); - - for i in 0..len { - let branch = branches[i].as_(); - vortex_ensure!( - branch < value_bits.len(), - "interleave array index out of bounds" - ); - vortex_ensure!( - rows[i].as_() < value_bits[branch].len(), - "interleave row index out of bounds" - ); - } - - Ok(len) -} - /// Gathers one bit per output from `bits[branches[i]]` at position `rows[i]`, packing 64 results per /// word with [`BitBufferMut::collect_bool`]. /// diff --git a/vortex-array/src/arrays/interleave/execute/mod.rs b/vortex-array/src/arrays/interleave/execute/mod.rs index 05dcd161f62..e4ac1f5ef43 100644 --- a/vortex-array/src/arrays/interleave/execute/mod.rs +++ b/vortex-array/src/arrays/interleave/execute/mod.rs @@ -5,14 +5,16 @@ //! //! All values share a type (validated in [`Interleave::check`]), so the //! physical gather kernel is chosen from the first value. The selector types are an orthogonal -//! concern handled within each kernel. Only boolean values are implemented today (see the [`bool`] module). +//! concern handled within each kernel. //! //! [`Interleave::check`]: super::Interleave::check -//! [`bool`]: module@crate::arrays::interleave::execute::bool mod bool; +mod primitive; +use num_traits::AsPrimitive; use vortex_error::VortexResult; +use vortex_error::vortex_ensure; use vortex_error::vortex_panic; use super::Interleave; @@ -28,12 +30,47 @@ pub(super) fn execute( ) -> VortexResult { if array.value(0).dtype().is_boolean() { bool::execute(array, ctx) + } else if array.value(0).dtype().is_primitive() { + primitive::execute(array, ctx) } else { let value_dtype = array.value(0).dtype().clone(); vortex_panic!( - "interleave execution is only implemented for boolean values; value dtype {} is not \ - yet supported", + "interleave execution is not implemented for value dtype {}", value_dtype ) } } + +/// Validate selector lengths and bounds, returning the common output length. +/// +/// On success, `branches.len() == rows.len() == len`; for every `i < len`, +/// `branches[i] < num_values` and `rows[i] < value_len(branches[i])`. +fn validate_selectors( + num_values: usize, + value_len: F, + branches: &[A], + rows: &[R], +) -> VortexResult +where + A: AsPrimitive, + R: AsPrimitive, + F: Fn(usize) -> usize, +{ + let len = branches.len(); + vortex_ensure!( + rows.len() == len, + "interleave selectors differ in length: array_indices {len}, row_indices {}", + rows.len() + ); + + for i in 0..len { + let branch = branches[i].as_(); + vortex_ensure!(branch < num_values, "interleave array index out of bounds"); + vortex_ensure!( + rows[i].as_() < value_len(branch), + "interleave row index out of bounds" + ); + } + + Ok(len) +} diff --git a/vortex-array/src/arrays/interleave/execute/primitive.rs b/vortex-array/src/arrays/interleave/execute/primitive.rs new file mode 100644 index 00000000000..a42039e1f74 --- /dev/null +++ b/vortex-array/src/arrays/interleave/execute/primitive.rs @@ -0,0 +1,128 @@ +// SPDX-License-Identifier: Apache-2.0 +// SPDX-FileCopyrightText: Copyright the Vortex contributors + +//! Execution for primitive [`Interleave`] values. + +use num_traits::AsPrimitive; +use vortex_buffer::Buffer; +use vortex_buffer::BufferMut; +use vortex_error::VortexResult; + +use super::super::Interleave; +use super::super::InterleaveArrayExt; +use super::validate_selectors; +use crate::array::Array; +use crate::array::ArrayView; +use crate::arrays::Constant; +use crate::arrays::Primitive; +use crate::arrays::PrimitiveArray; +use crate::arrays::primitive::PrimitiveArrayExt; +use crate::dtype::NativePType; +use crate::executor::ExecutionCtx; +use crate::executor::ExecutionResult; +use crate::match_each_native_ptype; +use crate::match_each_unsigned_integer_ptype; +use crate::require_child; + +pub(super) fn execute( + mut array: Array, + _ctx: &mut ExecutionCtx, +) -> VortexResult { + let num_values = array.num_values(); + array = require_child!(array, array.array_indices(), 0 => Primitive); + array = require_child!(array, array.row_indices(), 1 => Primitive); + for i in 0..num_values { + if array.value(i).as_opt::().is_none() { + array = require_child!(array, array.value(i), i + 2 => Primitive); + } + } + + let validity = array.as_ref().validity()?; + let output = match_each_native_ptype!(array.value(0).dtype().as_ptype(), |T| { + let values = gather_values::(&array)?; + VortexResult::Ok(PrimitiveArray::new(values, validity)) + })?; + + Ok(ExecutionResult::done(output)) +} + +/// Physical primitive values; nullness remains in the source array's validity. +enum PrimitiveValues { + Buffer(Buffer), + Constant { value: T, len: usize }, +} + +impl PrimitiveValues { + fn len(&self) -> usize { + match self { + Self::Buffer(values) => values.len(), + Self::Constant { len, .. } => *len, + } + } + + fn value(&self, index: usize) -> T { + match self { + Self::Buffer(values) => values[index], + Self::Constant { value, .. } => *value, + } + } +} + +fn gather_values(array: &Array) -> VortexResult> { + let values = (0..array.num_values()) + .map(|i| { + let value = array.value(i); + if let Some(constant) = value.as_opt::() { + PrimitiveValues::Constant { + value: constant + .scalar() + .as_primitive() + .typed_value::() + // Validity carries nullness; a null constant's payload is never observed. + .unwrap_or_default(), + len: value.len(), + } + } else { + PrimitiveValues::Buffer(value.as_::().to_buffer::()) + } + }) + .collect::>(); + let branches = array.array_indices().as_::(); + let rows = array.row_indices().as_::(); + + match_each_unsigned_integer_ptype!(branches.ptype(), |A| { + gather_rows::(&values, branches.as_slice::(), rows) + }) +} + +fn gather_rows( + values: &[PrimitiveValues], + branches: &[A], + rows: ArrayView<'_, Primitive>, +) -> VortexResult> +where + T: NativePType, + A: AsPrimitive, +{ + match_each_unsigned_integer_ptype!(rows.ptype(), |R| { + gather(values, branches, rows.as_slice::()) + }) +} + +fn gather( + values: &[PrimitiveValues], + branches: &[A], + rows: &[R], +) -> VortexResult> +where + T: NativePType, + A: AsPrimitive, + R: AsPrimitive, +{ + let len = validate_selectors(values.len(), |branch| values[branch].len(), branches, rows)?; + let mut output = BufferMut::with_capacity(len); + for i in 0..len { + output.push(values[branches[i].as_()].value(rows[i].as_())); + } + Ok(output.freeze()) +} diff --git a/vortex-array/src/arrays/interleave/mod.rs b/vortex-array/src/arrays/interleave/mod.rs index bff03ab055f..d29981249d4 100644 --- a/vortex-array/src/arrays/interleave/mod.rs +++ b/vortex-array/src/arrays/interleave/mod.rs @@ -463,6 +463,7 @@ mod tests { use crate::arrays::BoolArray; use crate::arrays::PrimitiveArray; use crate::assert_arrays_eq; + use crate::dtype::PType; /// Reference (oracle) implementation of the interleave spec, used only to validate the optimized /// [execute](super::execute) path. It is intentionally simple and slow: it pulls each output @@ -719,17 +720,55 @@ mod tests { } #[test] - #[should_panic(expected = "only implemented for boolean values")] - fn non_boolean_value_execution_panics() { - // Execution dispatches on the value type: primitive values have no kernel yet. - let v0 = PrimitiveArray::from_iter([1u32]).into_array(); - let v1 = PrimitiveArray::from_iter([2u32]).into_array(); - let array_indices = PrimitiveArray::from_iter([0u32, 1]).into_array(); - let row_indices = PrimitiveArray::from_iter([0u32, 0]).into_array(); - let interleaved = InterleaveArray::try_new(vec![v0, v1], array_indices, row_indices) - .vortex_expect("primitive values should construct") - .into_array(); + fn executes_primitive_values() -> VortexResult<()> { + let v0 = PrimitiveArray::from_iter([1.0f64, 2.0]).into_array(); + let v1 = PrimitiveArray::from_option_iter([Some(10.0f64), None]).into_array(); + let array_indices = PrimitiveArray::from_iter([0u8, 1, 0, 1]).into_array(); + let row_indices = PrimitiveArray::from_iter([0u32, 0, 1, 1]).into_array(); + let interleaved = + InterleaveArray::try_new(vec![v0, v1], array_indices, row_indices)?.into_array(); + let expected = + PrimitiveArray::from_option_iter([Some(1.0f64), Some(10.0), Some(2.0), None]) + .into_array(); + let mut ctx = array_session().create_execution_ctx(); + assert_arrays_eq!(interleaved, expected, &mut ctx); + Ok(()) + } + + #[test] + fn executes_primitive_constant_values() -> VortexResult<()> { + let constant = ConstantArray::new(1.0f64, 2).into_array(); + let column = PrimitiveArray::from_iter([10.0f64, 20.0]).into_array(); + let array_indices = PrimitiveArray::from_iter([0u8, 1, 0, 1]).into_array(); + let row_indices = PrimitiveArray::from_iter([0u32, 0, 1, 1]).into_array(); + let interleaved = + InterleaveArray::try_new(vec![constant, column], array_indices, row_indices)? + .into_array(); + let expected = PrimitiveArray::from_iter([1.0f64, 10.0, 1.0, 20.0]).into_array(); let mut ctx = array_session().create_execution_ctx(); - interleaved.execute::(&mut ctx).ok(); + + assert_arrays_eq!(interleaved, expected, &mut ctx); + Ok(()) + } + + #[test] + fn executes_null_primitive_constant_values() -> VortexResult<()> { + let constant = ConstantArray::new( + Scalar::null(DType::Primitive(PType::F64, Nullability::Nullable)), + 2, + ) + .into_array(); + let column = PrimitiveArray::from_iter([10.0f64, 20.0]).into_array(); + let array_indices = PrimitiveArray::from_iter([0u8, 1, 0, 1]).into_array(); + let row_indices = PrimitiveArray::from_iter([0u32, 0, 1, 1]).into_array(); + let interleaved = + InterleaveArray::try_new(vec![constant, column], array_indices, row_indices)? + .into_array(); + let expected = + PrimitiveArray::from_option_iter([None, Some(10.0f64), None, Some(20.0)]).into_array(); + let mut ctx = array_session().create_execution_ctx(); + + assert_arrays_eq!(interleaved, expected, &mut ctx); + Ok(()) } }