Skip to content

Commit 3ae4120

Browse files
committed
Specialize constant PiecewiseSequence runs
Signed-off-by: Daniel King <dan@spiraldb.com>
1 parent fdf6fb7 commit 3ae4120

3 files changed

Lines changed: 301 additions & 22 deletions

File tree

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

Lines changed: 269 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -20,8 +20,10 @@ use crate::arrays::Primitive;
2020
use crate::arrays::PrimitiveArray;
2121
use crate::arrays::dict::TakeExecute;
2222
use crate::arrays::list::ListArrayExt;
23+
use crate::arrays::piecewise_sequence::UnitMultiplierLengths;
2324
use crate::arrays::piecewise_sequence::execute_unit_multiplier_index_arrays;
2425
use crate::arrays::piecewise_sequence::validate_index_ranges;
26+
use crate::arrays::piecewise_sequence::validate_index_ranges_constant;
2527
use crate::arrays::primitive::PrimitiveArrayExt;
2628
use crate::builders::ArrayBuilder;
2729
use crate::builders::PrimitiveBuilder;
@@ -162,21 +164,174 @@ fn take_piecewise_sequence(
162164
};
163165
let offsets = array.offsets().clone().execute::<PrimitiveArray>(ctx)?;
164166
let offsets = offsets.reinterpret_cast(offsets.ptype().to_unsigned());
167+
let output_len = indices_ref.len();
168+
169+
let taken = match &lengths {
170+
UnitMultiplierLengths::Constant(length) => take_piecewise_sequence_constant_dispatch(
171+
array,
172+
&starts,
173+
*length,
174+
&offsets,
175+
indices_ref,
176+
output_len,
177+
)?,
178+
UnitMultiplierLengths::Array(lengths) => take_piecewise_sequence_lengths_dispatch(
179+
array,
180+
&starts,
181+
lengths,
182+
&offsets,
183+
indices_ref,
184+
output_len,
185+
)?,
186+
};
187+
Ok(Some(taken))
188+
}
165189

190+
fn take_piecewise_sequence_constant_dispatch(
191+
array: ArrayView<'_, List>,
192+
starts: &PrimitiveArray,
193+
length: usize,
194+
offsets: &PrimitiveArray,
195+
indices_ref: &ArrayRef,
196+
output_len: usize,
197+
) -> VortexResult<ArrayRef> {
166198
match_each_unsigned_integer_ptype!(starts.ptype(), |S| {
167-
match_each_unsigned_integer_ptype!(lengths.ptype(), |L| {
168-
match_each_unsigned_integer_ptype!(offsets.ptype(), |O| {
169-
take_piecewise_sequence_typed::<S, L, O>(
170-
array,
171-
starts.as_slice::<S>(),
172-
lengths.as_slice::<L>(),
173-
offsets.as_slice::<O>(),
174-
indices_ref,
175-
)
176-
})
177-
})
199+
take_piecewise_sequence_constant_start_dispatch::<S>(
200+
array,
201+
starts,
202+
length,
203+
offsets,
204+
indices_ref,
205+
output_len,
206+
)
207+
})
208+
}
209+
210+
fn take_piecewise_sequence_constant_start_dispatch<S>(
211+
array: ArrayView<'_, List>,
212+
starts: &PrimitiveArray,
213+
length: usize,
214+
offsets: &PrimitiveArray,
215+
indices_ref: &ArrayRef,
216+
output_len: usize,
217+
) -> VortexResult<ArrayRef>
218+
where
219+
S: UnsignedPType,
220+
{
221+
match_each_unsigned_integer_ptype!(offsets.ptype(), |O| {
222+
take_piecewise_sequence_constant_length::<S, O>(
223+
array,
224+
starts.as_slice::<S>(),
225+
length,
226+
offsets.as_slice::<O>(),
227+
indices_ref,
228+
output_len,
229+
)
230+
})
231+
}
232+
233+
fn take_piecewise_sequence_lengths_dispatch(
234+
array: ArrayView<'_, List>,
235+
starts: &PrimitiveArray,
236+
lengths: &PrimitiveArray,
237+
offsets: &PrimitiveArray,
238+
indices_ref: &ArrayRef,
239+
output_len: usize,
240+
) -> VortexResult<ArrayRef> {
241+
match_each_unsigned_integer_ptype!(starts.ptype(), |S| {
242+
take_piecewise_sequence_lengths_start_dispatch::<S>(
243+
array,
244+
starts,
245+
lengths,
246+
offsets,
247+
indices_ref,
248+
output_len,
249+
)
250+
})
251+
}
252+
253+
fn take_piecewise_sequence_lengths_start_dispatch<S>(
254+
array: ArrayView<'_, List>,
255+
starts: &PrimitiveArray,
256+
lengths: &PrimitiveArray,
257+
offsets: &PrimitiveArray,
258+
indices_ref: &ArrayRef,
259+
output_len: usize,
260+
) -> VortexResult<ArrayRef>
261+
where
262+
S: UnsignedPType,
263+
{
264+
match_each_unsigned_integer_ptype!(lengths.ptype(), |L| {
265+
take_piecewise_sequence_lengths_start_length_dispatch::<S, L>(
266+
array,
267+
starts,
268+
lengths,
269+
offsets,
270+
indices_ref,
271+
output_len,
272+
)
273+
})
274+
}
275+
276+
fn take_piecewise_sequence_lengths_start_length_dispatch<S, L>(
277+
array: ArrayView<'_, List>,
278+
starts: &PrimitiveArray,
279+
lengths: &PrimitiveArray,
280+
offsets: &PrimitiveArray,
281+
indices_ref: &ArrayRef,
282+
output_len: usize,
283+
) -> VortexResult<ArrayRef>
284+
where
285+
S: UnsignedPType,
286+
L: UnsignedPType,
287+
{
288+
match_each_unsigned_integer_ptype!(offsets.ptype(), |O| {
289+
take_piecewise_sequence_typed::<S, L, O>(
290+
array,
291+
starts.as_slice::<S>(),
292+
lengths.as_slice::<L>(),
293+
offsets.as_slice::<O>(),
294+
indices_ref,
295+
output_len,
296+
)
297+
})
298+
}
299+
300+
fn take_piecewise_sequence_constant_length<S, Offset>(
301+
array: ArrayView<'_, List>,
302+
starts: &[S],
303+
length: usize,
304+
offsets: &[Offset],
305+
indices_ref: &ArrayRef,
306+
output_len: usize,
307+
) -> VortexResult<ArrayRef>
308+
where
309+
S: UnsignedPType,
310+
Offset: UnsignedPType,
311+
{
312+
validate_index_ranges_constant(array.len(), starts, length, output_len)?;
313+
let total_elements =
314+
piecewise_list_elements_len_constant(array.elements().len(), offsets, starts, length)?;
315+
let validity = array.validity()?.take(indices_ref)?;
316+
317+
match_smallest_offset_type!(total_elements, |OutputOffset| {
318+
let gathered = gather_piecewise_list_constant_length::<S, Offset, OutputOffset>(
319+
array.elements(),
320+
offsets,
321+
starts,
322+
length,
323+
output_len,
324+
total_elements,
325+
)?;
326+
327+
// SAFETY: output offsets are rebuilt from valid monotonic source offsets; output elements
328+
// are exactly the gathered child ranges referenced by those offsets; validity has one bit
329+
// per output row.
330+
Ok(
331+
unsafe { ListArray::new_unchecked(gathered.elements, gathered.offsets, validity) }
332+
.into_array(),
333+
)
178334
})
179-
.map(Some)
180335
}
181336

182337
fn take_piecewise_sequence_typed<S, L, Offset>(
@@ -185,13 +340,14 @@ fn take_piecewise_sequence_typed<S, L, Offset>(
185340
lengths: &[L],
186341
offsets: &[Offset],
187342
indices_ref: &ArrayRef,
343+
output_len: usize,
188344
) -> VortexResult<ArrayRef>
189345
where
190346
S: UnsignedPType,
191347
L: UnsignedPType,
192348
Offset: UnsignedPType,
193349
{
194-
validate_index_ranges(array.len(), starts, lengths, indices_ref.len())?;
350+
validate_index_ranges(array.len(), starts, lengths, output_len)?;
195351
let total_elements =
196352
piecewise_list_elements_len(array.elements().len(), offsets, starts, lengths)?;
197353

@@ -201,7 +357,7 @@ where
201357
offsets,
202358
starts,
203359
lengths,
204-
indices_ref.len(),
360+
output_len,
205361
total_elements,
206362
)?;
207363
let validity = array.validity()?.take(indices_ref)?;
@@ -221,6 +377,37 @@ struct GatheredList {
221377
offsets: ArrayRef,
222378
}
223379

380+
fn piecewise_list_elements_len_constant<S, Offset>(
381+
elements_len: usize,
382+
offsets: &[Offset],
383+
starts: &[S],
384+
length: usize,
385+
) -> VortexResult<usize>
386+
where
387+
S: UnsignedPType,
388+
Offset: UnsignedPType,
389+
{
390+
let mut total = 0usize;
391+
for &start in starts {
392+
let start: usize = start.as_();
393+
let end = start + length;
394+
if length == 0 {
395+
continue;
396+
}
397+
398+
let element_start: usize = offsets[start].as_();
399+
let element_end: usize = offsets[end].as_();
400+
vortex_ensure!(
401+
element_start <= element_end && element_end <= elements_len,
402+
"List offsets range {element_start}..{element_end} exceeds elements length {elements_len}",
403+
);
404+
total = total
405+
.checked_add(element_end - element_start)
406+
.ok_or_else(|| vortex_err!("List take output elements length overflow"))?;
407+
}
408+
Ok(total)
409+
}
410+
224411
fn piecewise_list_elements_len<S, L, Offset>(
225412
elements_len: usize,
226413
offsets: &[Offset],
@@ -254,6 +441,74 @@ where
254441
Ok(total)
255442
}
256443

444+
fn gather_piecewise_list_constant_length<S, Offset, OutputOffset>(
445+
elements: &ArrayRef,
446+
offsets: &[Offset],
447+
starts: &[S],
448+
length: usize,
449+
output_len: usize,
450+
total_elements: usize,
451+
) -> VortexResult<GatheredList>
452+
where
453+
S: UnsignedPType,
454+
Offset: UnsignedPType,
455+
OutputOffset: IntegerPType,
456+
{
457+
let offsets_capacity = output_len
458+
.checked_add(1)
459+
.ok_or_else(|| vortex_err!("List take offsets length overflow"))?;
460+
let mut new_offsets = BufferMut::<OutputOffset>::with_capacity(offsets_capacity);
461+
let mut element_starts = BufferMut::<u64>::with_capacity(starts.len());
462+
let mut element_lengths = BufferMut::<u64>::with_capacity(starts.len());
463+
let mut output_elements = 0usize;
464+
465+
new_offsets.push(OutputOffset::zero());
466+
for &start in starts {
467+
let start: usize = start.as_();
468+
let end = start + length;
469+
if length == 0 {
470+
continue;
471+
}
472+
473+
let element_start: usize = offsets[start].as_();
474+
let element_end: usize = offsets[end].as_();
475+
for &offset in &offsets[start + 1..=end] {
476+
let offset: usize = offset.as_();
477+
let relative = offset
478+
.checked_sub(element_start)
479+
.ok_or_else(|| vortex_err!("List offsets are not monotonic at offset {offset}"))?;
480+
let output_offset = output_elements
481+
.checked_add(relative)
482+
.ok_or_else(|| vortex_err!("List take output elements length overflow"))?;
483+
new_offsets.push(new_offset_value::<OutputOffset>(output_offset)?);
484+
}
485+
486+
let element_length = element_end - element_start;
487+
element_starts.push(element_start as u64);
488+
element_lengths.push(element_length as u64);
489+
output_elements = output_elements
490+
.checked_add(element_length)
491+
.ok_or_else(|| vortex_err!("List take output elements length overflow"))?;
492+
}
493+
debug_assert_eq!(output_elements, total_elements);
494+
495+
let offsets = PrimitiveArray::new(new_offsets.freeze(), Validity::NonNullable).into_array();
496+
let multipliers = ConstantArray::new(1u64, element_starts.len()).into_array();
497+
// SAFETY: element ranges are derived from validated source list offsets, and total_elements is
498+
// the sum of the gathered element range lengths. Multiplier 1 preserves contiguous ranges.
499+
let element_indices = unsafe {
500+
PiecewiseSequenceArray::new_unchecked(
501+
element_starts.into_array(),
502+
element_lengths.into_array(),
503+
multipliers,
504+
total_elements,
505+
)
506+
};
507+
let elements = elements.take(element_indices.into_array())?;
508+
509+
Ok(GatheredList { elements, offsets })
510+
}
511+
257512
fn gather_piecewise_list<S, L, Offset, OutputOffset>(
258513
elements: &ArrayRef,
259514
offsets: &[Offset],

vortex-array/src/arrays/listview/rebuild.rs

Lines changed: 10 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -10,7 +10,6 @@ use vortex_mask::Mask;
1010

1111
use crate::ExecutionCtx;
1212
use crate::IntoArray;
13-
use crate::RecursiveCanonical;
1413
use crate::arrays::ConstantArray;
1514
use crate::arrays::ListViewArray;
1615
use crate::arrays::PiecewiseSequenceArray;
@@ -286,24 +285,27 @@ impl ListViewArray {
286285
lengths,
287286
elements_len,
288287
} = ranges;
288+
let constant_length = lengths
289+
.first()
290+
.copied()
291+
.filter(|first| lengths.iter().all(|length| *length == *first));
292+
let lengths = match constant_length {
293+
Some(length) => ConstantArray::new(length, starts.len()).into_array(),
294+
None => lengths.into_array(),
295+
};
289296

290297
// SAFETY: range starts and lengths are derived from valid ListView metadata; elements_len
291298
// is the sum of all generated range lengths. Multiplier 1 preserves contiguous ranges.
292299
let multipliers = ConstantArray::new(1u64, starts.len()).into_array();
293300
let element_indices = unsafe {
294301
PiecewiseSequenceArray::new_unchecked(
295302
starts.into_array(),
296-
lengths.into_array(),
303+
lengths,
297304
multipliers,
298305
elements_len,
299306
)
300307
};
301-
let elements = self
302-
.elements()
303-
.take(element_indices.into_array())?
304-
.execute::<RecursiveCanonical>(ctx)?
305-
.0
306-
.into_array();
308+
let elements = self.elements().take(element_indices.into_array())?;
307309

308310
// Built unsigned; reinterpret back to the signed-preserving result types.
309311
let offsets = PrimitiveArray::new(new_offsets.freeze(), Validity::NonNullable)

0 commit comments

Comments
 (0)