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..c43aabe914f 100644 --- a/vortex-array/src/arrays/interleave/execute/mod.rs +++ b/vortex-array/src/arrays/interleave/execute/mod.rs @@ -5,18 +5,19 @@ //! //! 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; -use super::InterleaveArrayExt; use crate::array::Array; use crate::executor::ExecutionCtx; use crate::executor::ExecutionResult; @@ -26,14 +27,48 @@ pub(super) fn execute( array: Array, ctx: &mut ExecutionCtx, ) -> VortexResult { - if array.value(0).dtype().is_boolean() { + if array.dtype().is_boolean() { bool::execute(array, ctx) + } else if array.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", - value_dtype + "interleave execution is not implemented for value dtype {}", + array.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..f483bc015be --- /dev/null +++ b/vortex-array/src/arrays/interleave/execute/primitive.rs @@ -0,0 +1,165 @@ +// 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::AnyColumnar; +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 { + array = require_child!(array, array.value(i), i + 2 => AnyColumnar); + } + + let validity = array.as_ref().validity()?; + let output = match_each_native_ptype!(array.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(T), +} + +impl PrimitiveValues { + /// Returns the physical value at `index` without bounds checking. + /// + /// # Safety + /// + /// For [`Self::Buffer`], `index` must be less than the buffer length. [`Self::Constant`] + /// accepts any index because it has a single physical value. + unsafe fn value_unchecked(&self, index: usize) -> T { + match self { + // SAFETY: the caller guarantees that `index` is in bounds. + Self::Buffer(values) => *unsafe { values.get_unchecked(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( + constant + .scalar() + .as_primitive() + .typed_value::() + // Validity carries nullness; a null constant's payload is never observed. + .unwrap_or_default(), + ) + } 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::(array, &values, branches.as_slice::(), rows) + }) +} + +fn gather_rows( + array: &Array, + values: &[PrimitiveValues], + branches: &[A], + rows: ArrayView<'_, Primitive>, +) -> VortexResult> +where + T: NativePType, + A: AsPrimitive, +{ + match_each_unsigned_integer_ptype!(rows.ptype(), |R| { + gather(array, values, branches, rows.as_slice::()) + }) +} + +fn gather( + array: &Array, + values: &[PrimitiveValues], + branches: &[A], + rows: &[R], +) -> VortexResult> +where + T: NativePType, + A: AsPrimitive, + R: AsPrimitive, +{ + let len = validate_selectors( + values.len(), + |branch| array.value(branch).len(), + branches, + rows, + )?; + + // SAFETY: `validate_selectors` proved both selector lengths and every logical source bound. + // Each `Buffer` has the same length as its source array, while `Constant` ignores the row. + Ok(unsafe { gather_unchecked(len, values, branches, rows) }) +} + +/// Gathers one primitive value per output from `values[branches[i]]` at position `rows[i]`. +/// +/// # Safety +/// +/// `branches` and `rows` must both contain at least `len` elements. For every `i < len`, +/// `branches[i] < values.len()` and, when the selected value is a [`PrimitiveValues::Buffer`], +/// `rows[i]` must be less than that buffer's length. +unsafe fn gather_unchecked( + len: usize, + values: &[PrimitiveValues], + branches: &[A], + rows: &[R], +) -> Buffer +where + T: NativePType, + A: AsPrimitive, + R: AsPrimitive, +{ + let mut output = BufferMut::with_capacity(len); + for ((branch, row), slot) in branches.iter().zip(rows).zip(output.spare_capacity_mut()) { + let branch = (*branch).as_(); + let row = (*row).as_(); + // SAFETY: the caller guarantees that the selected branch and row are in bounds for + // `values` and the selected physical value buffer. + slot.write(unsafe { values.get_unchecked(branch).value_unchecked(row) }); + } + + // SAFETY: the caller guarantees both selector slices have at least `len` elements, so the loop + // initialized exactly `len` output slots. + unsafe { output.set_len(len) }; + 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(()) } }