Skip to content

Commit 7ac5b5d

Browse files
committed
Port PiecewiseSequence run take consumers
Signed-off-by: Daniel King <dan@spiraldb.com>
1 parent b615895 commit 7ac5b5d

4 files changed

Lines changed: 410 additions & 85 deletions

File tree

vortex-array/src/arrays/extension/compute/rules.rs

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -11,6 +11,7 @@ use crate::arrays::ConstantArray;
1111
use crate::arrays::Extension;
1212
use crate::arrays::ExtensionArray;
1313
use crate::arrays::Filter;
14+
use crate::arrays::dict::TakeReduceAdaptor;
1415
use crate::arrays::extension::ExtensionArrayExt;
1516
use crate::arrays::filter::FilterReduceAdaptor;
1617
use crate::arrays::slice::SliceReduceAdaptor;
@@ -50,6 +51,7 @@ pub(crate) const PARENT_RULES: ParentRuleSet<Extension> = ParentRuleSet::new(&[
5051
ParentRuleSet::lift(&FilterReduceAdaptor(Extension)),
5152
ParentRuleSet::lift(&MaskReduceAdaptor(Extension)),
5253
ParentRuleSet::lift(&SliceReduceAdaptor(Extension)),
54+
ParentRuleSet::lift(&TakeReduceAdaptor(Extension)),
5355
]);
5456

5557
/// Push filter operations into the storage array of an extension array.

vortex-array/src/arrays/extension/compute/take.rs

Lines changed: 16 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -10,8 +10,24 @@ use crate::array::ArrayView;
1010
use crate::arrays::Extension;
1111
use crate::arrays::ExtensionArray;
1212
use crate::arrays::dict::TakeExecute;
13+
use crate::arrays::dict::TakeReduce;
1314
use crate::arrays::extension::ExtensionArrayExt;
1415

16+
impl TakeReduce for Extension {
17+
fn take(array: ArrayView<'_, Extension>, indices: &ArrayRef) -> VortexResult<Option<ArrayRef>> {
18+
let taken_storage = array.storage_array().take(indices.clone())?;
19+
Ok(Some(
20+
ExtensionArray::new(
21+
array
22+
.ext_dtype()
23+
.with_nullability(taken_storage.dtype().nullability()),
24+
taken_storage,
25+
)
26+
.into_array(),
27+
))
28+
}
29+
}
30+
1531
impl TakeExecute for Extension {
1632
fn take(
1733
array: ArrayView<'_, Extension>,

vortex-array/src/arrays/list/compute/take.rs

Lines changed: 249 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,26 +1,36 @@
11
// SPDX-License-Identifier: Apache-2.0
22
// SPDX-FileCopyrightText: Copyright the Vortex contributors
33

4+
use itertools::Itertools as _;
5+
use vortex_buffer::BufferMut;
46
use vortex_error::VortexExpect;
57
use vortex_error::VortexResult;
8+
use vortex_error::vortex_ensure;
9+
use vortex_error::vortex_err;
610

711
use crate::ArrayRef;
812
use crate::IntoArray;
913
use crate::array::ArrayView;
1014
use crate::arrays::List;
1115
use crate::arrays::ListArray;
16+
use crate::arrays::PiecewiseSequence;
17+
use crate::arrays::PiecewiseSequenceArray;
1218
use crate::arrays::Primitive;
1319
use crate::arrays::PrimitiveArray;
1420
use crate::arrays::dict::TakeExecute;
1521
use crate::arrays::list::ListArrayExt;
22+
use crate::arrays::piecewise_sequence::execute_index_arrays;
23+
use crate::arrays::piecewise_sequence::validate_index_ranges;
1624
use crate::arrays::primitive::PrimitiveArrayExt;
1725
use crate::builders::ArrayBuilder;
1826
use crate::builders::PrimitiveBuilder;
1927
use crate::dtype::IntegerPType;
2028
use crate::dtype::Nullability;
29+
use crate::dtype::UnsignedPType;
2130
use crate::executor::ExecutionCtx;
2231
use crate::match_each_unsigned_integer_ptype;
2332
use crate::match_smallest_offset_type;
33+
use crate::validity::Validity;
2434

2535
// TODO(connor)[ListView]: Re-revert to the version where we simply convert to a `ListView` and call
2636
// the `ListView::take` compute function once `ListView` is more stable.
@@ -37,6 +47,12 @@ impl TakeExecute for List {
3747
indices: &ArrayRef,
3848
ctx: &mut ExecutionCtx,
3949
) -> VortexResult<Option<ArrayRef>> {
50+
if let Some(piecewise_indices) = indices.as_opt::<PiecewiseSequence>()
51+
&& let Some(taken) = take_piecewise_sequence(array, piecewise_indices, indices, ctx)?
52+
{
53+
return Ok(Some(taken));
54+
}
55+
4056
let indices = indices.clone().execute::<PrimitiveArray>(ctx)?;
4157
let indices = indices.reinterpret_cast(indices.ptype().to_unsigned());
4258
let offsets = array.offsets().clone().execute::<PrimitiveArray>(ctx)?;
@@ -127,6 +143,191 @@ fn _take<I: IntegerPType, O: IntegerPType, OutputOffsetType: IntegerPType>(
127143
.into_array())
128144
}
129145

146+
fn take_piecewise_sequence(
147+
array: ArrayView<'_, List>,
148+
indices: ArrayView<'_, PiecewiseSequence>,
149+
indices_ref: &ArrayRef,
150+
ctx: &mut ExecutionCtx,
151+
) -> VortexResult<Option<ArrayRef>> {
152+
let data_validity = array
153+
.list_validity()
154+
.execute_mask(array.as_ref().len(), ctx)?;
155+
if !data_validity.all_true() {
156+
return Ok(None);
157+
}
158+
159+
let (starts, lengths) = execute_index_arrays(indices, ctx)?;
160+
let offsets = array.offsets().clone().execute::<PrimitiveArray>(ctx)?;
161+
let offsets = offsets.reinterpret_cast(offsets.ptype().to_unsigned());
162+
163+
match_each_unsigned_integer_ptype!(starts.ptype(), |S| {
164+
match_each_unsigned_integer_ptype!(lengths.ptype(), |L| {
165+
match_each_unsigned_integer_ptype!(offsets.ptype(), |O| {
166+
take_piecewise_sequence_typed::<S, L, O>(
167+
array,
168+
starts.as_slice::<S>(),
169+
lengths.as_slice::<L>(),
170+
offsets.as_slice::<O>(),
171+
indices_ref,
172+
)
173+
})
174+
})
175+
})
176+
.map(Some)
177+
}
178+
179+
fn take_piecewise_sequence_typed<S, L, Offset>(
180+
array: ArrayView<'_, List>,
181+
starts: &[S],
182+
lengths: &[L],
183+
offsets: &[Offset],
184+
indices_ref: &ArrayRef,
185+
) -> VortexResult<ArrayRef>
186+
where
187+
S: UnsignedPType,
188+
L: UnsignedPType,
189+
Offset: UnsignedPType,
190+
{
191+
validate_index_ranges(array.len(), starts, lengths, indices_ref.len())?;
192+
let total_elements =
193+
piecewise_list_elements_len(array.elements().len(), offsets, starts, lengths)?;
194+
195+
match_smallest_offset_type!(total_elements, |OutputOffset| {
196+
let gathered = gather_piecewise_list::<S, L, Offset, OutputOffset>(
197+
array.elements(),
198+
offsets,
199+
starts,
200+
lengths,
201+
indices_ref.len(),
202+
total_elements,
203+
)?;
204+
let validity = array.validity()?.take(indices_ref)?;
205+
206+
// SAFETY: output offsets are rebuilt from valid monotonic source offsets; output elements
207+
// are exactly the gathered child ranges referenced by those offsets; validity has one bit
208+
// per output row.
209+
Ok(
210+
unsafe { ListArray::new_unchecked(gathered.elements, gathered.offsets, validity) }
211+
.into_array(),
212+
)
213+
})
214+
}
215+
216+
struct GatheredList {
217+
elements: ArrayRef,
218+
offsets: ArrayRef,
219+
}
220+
221+
fn piecewise_list_elements_len<S, L, Offset>(
222+
elements_len: usize,
223+
offsets: &[Offset],
224+
starts: &[S],
225+
lengths: &[L],
226+
) -> VortexResult<usize>
227+
where
228+
S: UnsignedPType,
229+
L: UnsignedPType,
230+
Offset: UnsignedPType,
231+
{
232+
let mut total = 0usize;
233+
for (&start, &length) in starts.iter().zip_eq(lengths) {
234+
let start: usize = start.as_();
235+
let length: usize = length.as_();
236+
let end = start + length;
237+
if length == 0 {
238+
continue;
239+
}
240+
241+
let element_start: usize = offsets[start].as_();
242+
let element_end: usize = offsets[end].as_();
243+
vortex_ensure!(
244+
element_start <= element_end && element_end <= elements_len,
245+
"List offsets range {element_start}..{element_end} exceeds elements length {elements_len}",
246+
);
247+
total = total
248+
.checked_add(element_end - element_start)
249+
.ok_or_else(|| vortex_err!("List take output elements length overflow"))?;
250+
}
251+
Ok(total)
252+
}
253+
254+
fn gather_piecewise_list<S, L, Offset, OutputOffset>(
255+
elements: &ArrayRef,
256+
offsets: &[Offset],
257+
starts: &[S],
258+
lengths: &[L],
259+
output_len: usize,
260+
total_elements: usize,
261+
) -> VortexResult<GatheredList>
262+
where
263+
S: UnsignedPType,
264+
L: UnsignedPType,
265+
Offset: UnsignedPType,
266+
OutputOffset: IntegerPType,
267+
{
268+
let offsets_capacity = output_len
269+
.checked_add(1)
270+
.ok_or_else(|| vortex_err!("List take offsets length overflow"))?;
271+
let mut new_offsets = BufferMut::<OutputOffset>::with_capacity(offsets_capacity);
272+
let mut element_starts = BufferMut::<u64>::with_capacity(starts.len());
273+
let mut element_lengths = BufferMut::<u64>::with_capacity(lengths.len());
274+
let mut output_elements = 0usize;
275+
276+
new_offsets.push(OutputOffset::zero());
277+
for (&start, &length) in starts.iter().zip_eq(lengths) {
278+
let start: usize = start.as_();
279+
let length: usize = length.as_();
280+
let end = start + length;
281+
if length == 0 {
282+
continue;
283+
}
284+
285+
let element_start: usize = offsets[start].as_();
286+
let element_end: usize = offsets[end].as_();
287+
for &offset in &offsets[start + 1..=end] {
288+
let offset: usize = offset.as_();
289+
let relative = offset
290+
.checked_sub(element_start)
291+
.ok_or_else(|| vortex_err!("List offsets are not monotonic at offset {offset}"))?;
292+
let output_offset = output_elements
293+
.checked_add(relative)
294+
.ok_or_else(|| vortex_err!("List take output elements length overflow"))?;
295+
new_offsets.push(new_offset_value::<OutputOffset>(output_offset)?);
296+
}
297+
298+
let element_length = element_end - element_start;
299+
element_starts.push(element_start as u64);
300+
element_lengths.push(element_length as u64);
301+
output_elements = output_elements
302+
.checked_add(element_length)
303+
.ok_or_else(|| vortex_err!("List take output elements length overflow"))?;
304+
}
305+
debug_assert_eq!(output_elements, total_elements);
306+
307+
let offsets = PrimitiveArray::new(new_offsets.freeze(), Validity::NonNullable).into_array();
308+
// SAFETY: element ranges are derived from validated source list offsets, and total_elements is
309+
// the sum of the gathered element range lengths.
310+
let element_indices = unsafe {
311+
PiecewiseSequenceArray::new_unchecked(
312+
element_starts.into_array(),
313+
element_lengths.into_array(),
314+
total_elements,
315+
)
316+
};
317+
let elements = elements.take(element_indices.into_array())?;
318+
319+
Ok(GatheredList { elements, offsets })
320+
}
321+
322+
fn new_offset_value<T: IntegerPType>(value: usize) -> VortexResult<T> {
323+
T::from_usize(value).ok_or_else(|| {
324+
vortex_err!(
325+
"List take offset value {value} does not fit in {}",
326+
T::PTYPE
327+
)
328+
})
329+
}
330+
130331
// Kept out-of-line: as a single-callsite generic helper it would otherwise be inlined into every
131332
// monomorphization of `_take`, duplicating the entire nullable path across all specializations.
132333
#[inline(never)]
@@ -217,6 +418,7 @@ mod test {
217418
use crate::arrays::BoolArray;
218419
use crate::arrays::ListArray;
219420
use crate::arrays::ListViewArray;
421+
use crate::arrays::PiecewiseSequenceArray;
220422
use crate::arrays::PrimitiveArray;
221423
use crate::compute::conformance::take::test_take_conformance;
222424
use crate::dtype::DType;
@@ -403,6 +605,53 @@ mod test {
403605
);
404606
}
405607

608+
#[test]
609+
fn piecewise_sequence_take() {
610+
let mut ctx = array_session().create_execution_ctx();
611+
let list = ListArray::try_new(
612+
buffer![0i32, 1, 2, 3, 4, 5, 6].into_array(),
613+
buffer![0u32, 2, 5, 5, 7].into_array(),
614+
Validity::NonNullable,
615+
)
616+
.unwrap()
617+
.into_array();
618+
let idx = PiecewiseSequenceArray::try_new(
619+
buffer![1u64, 0].into_array(),
620+
buffer![2u64, 1].into_array(),
621+
3,
622+
)
623+
.unwrap()
624+
.into_array();
625+
626+
let result = list
627+
.take(idx)
628+
.unwrap()
629+
.execute::<ListViewArray>(&mut ctx)
630+
.unwrap();
631+
632+
let element_dtype: Arc<DType> = Arc::new(I32.into());
633+
assert_eq!(
634+
result.execute_scalar(0, &mut ctx).unwrap(),
635+
Scalar::list(
636+
Arc::clone(&element_dtype),
637+
vec![2i32.into(), 3.into(), 4.into()],
638+
Nullability::NonNullable
639+
)
640+
);
641+
assert_eq!(
642+
result.execute_scalar(1, &mut ctx).unwrap(),
643+
Scalar::list(Arc::clone(&element_dtype), vec![], Nullability::NonNullable)
644+
);
645+
assert_eq!(
646+
result.execute_scalar(2, &mut ctx).unwrap(),
647+
Scalar::list(
648+
element_dtype,
649+
vec![0i32.into(), 1.into()],
650+
Nullability::NonNullable
651+
)
652+
);
653+
}
654+
406655
#[test]
407656
fn test_take_empty_array() {
408657
let list = ListArray::try_new(

0 commit comments

Comments
 (0)