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 ;
46use vortex_error:: VortexExpect ;
57use vortex_error:: VortexResult ;
8+ use vortex_error:: vortex_ensure;
9+ use vortex_error:: vortex_err;
610
711use crate :: ArrayRef ;
812use crate :: IntoArray ;
913use crate :: array:: ArrayView ;
1014use crate :: arrays:: List ;
1115use crate :: arrays:: ListArray ;
16+ use crate :: arrays:: PiecewiseSequence ;
17+ use crate :: arrays:: PiecewiseSequenceArray ;
1218use crate :: arrays:: Primitive ;
1319use crate :: arrays:: PrimitiveArray ;
1420use crate :: arrays:: dict:: TakeExecute ;
1521use crate :: arrays:: list:: ListArrayExt ;
22+ use crate :: arrays:: piecewise_sequence:: execute_index_arrays;
23+ use crate :: arrays:: piecewise_sequence:: validate_index_ranges;
1624use crate :: arrays:: primitive:: PrimitiveArrayExt ;
1725use crate :: builders:: ArrayBuilder ;
1826use crate :: builders:: PrimitiveBuilder ;
1927use crate :: dtype:: IntegerPType ;
2028use crate :: dtype:: Nullability ;
29+ use crate :: dtype:: UnsignedPType ;
2130use crate :: executor:: ExecutionCtx ;
2231use crate :: match_each_unsigned_integer_ptype;
2332use 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