Skip to content

Commit ffb78fb

Browse files
committed
Split PiecewiseSequence take dispatch
Signed-off-by: Daniel King <dan@spiraldb.com>
1 parent 1ded1f2 commit ffb78fb

2 files changed

Lines changed: 121 additions & 35 deletions

File tree

  • vortex-array/src/arrays

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

Lines changed: 62 additions & 19 deletions
Original file line numberDiff line numberDiff line change
@@ -18,10 +18,12 @@ use crate::arrays::piecewise_sequence::execute_unit_multiplier_index_arrays;
1818
use crate::arrays::piecewise_sequence::validate_index_ranges;
1919
use crate::dtype::IntegerPType;
2020
use crate::dtype::NativeDecimalType;
21+
use crate::dtype::UnsignedPType;
2122
use crate::executor::ExecutionCtx;
2223
use crate::match_each_decimal_value_type;
2324
use crate::match_each_integer_ptype;
2425
use crate::match_each_unsigned_integer_ptype;
26+
use crate::validity::Validity;
2527

2628
impl TakeExecute for Decimal {
2729
fn take(
@@ -63,24 +65,65 @@ fn take_piecewise_sequence(
6365
let Some((starts, lengths)) = execute_unit_multiplier_index_arrays(indices, ctx)? else {
6466
return Ok(None);
6567
};
68+
let validity = array.validity()?.take(indices_ref)?;
69+
let output_len = indices_ref.len();
70+
let taken = take_piecewise_sequence_lengths(array, &starts, &lengths, validity, output_len)?;
71+
Ok(Some(taken))
72+
}
73+
74+
fn take_piecewise_sequence_lengths(
75+
array: ArrayView<'_, Decimal>,
76+
starts: &PrimitiveArray,
77+
lengths: &PrimitiveArray,
78+
validity: Validity,
79+
output_len: usize,
80+
) -> VortexResult<ArrayRef> {
6681
match_each_decimal_value_type!(array.values_type(), |D| {
67-
match_each_unsigned_integer_ptype!(starts.ptype(), |S| {
68-
match_each_unsigned_integer_ptype!(lengths.ptype(), |L| {
69-
let values = take_piecewise_to_buffer::<S, L, D>(
70-
starts.as_slice::<S>(),
71-
lengths.as_slice::<L>(),
72-
array.buffer::<D>().as_slice(),
73-
indices_ref.len(),
74-
)?;
75-
let validity = array.validity()?.take(indices_ref)?;
76-
77-
// SAFETY: contiguous gather preserves the decimal dtype and value representation.
78-
Ok(Some(
79-
unsafe { DecimalArray::new_unchecked(values, array.decimal_dtype(), validity) }
80-
.into_array(),
81-
))
82-
})
83-
})
82+
take_piecewise_sequence_lengths_typed::<D>(array, starts, lengths, validity, output_len)
83+
})
84+
}
85+
86+
fn take_piecewise_sequence_lengths_typed<D>(
87+
array: ArrayView<'_, Decimal>,
88+
starts: &PrimitiveArray,
89+
lengths: &PrimitiveArray,
90+
validity: Validity,
91+
output_len: usize,
92+
) -> VortexResult<ArrayRef>
93+
where
94+
D: NativeDecimalType,
95+
{
96+
match_each_unsigned_integer_ptype!(starts.ptype(), |S| {
97+
take_piecewise_sequence_lengths_start_typed::<D, S>(
98+
array, starts, lengths, validity, output_len,
99+
)
100+
})
101+
}
102+
103+
fn take_piecewise_sequence_lengths_start_typed<D, S>(
104+
array: ArrayView<'_, Decimal>,
105+
starts: &PrimitiveArray,
106+
lengths: &PrimitiveArray,
107+
validity: Validity,
108+
output_len: usize,
109+
) -> VortexResult<ArrayRef>
110+
where
111+
D: NativeDecimalType,
112+
S: UnsignedPType,
113+
{
114+
match_each_unsigned_integer_ptype!(lengths.ptype(), |L| {
115+
let values = take_piecewise_to_buffer::<S, L, D>(
116+
starts.as_slice::<S>(),
117+
lengths.as_slice::<L>(),
118+
array.buffer::<D>().as_slice(),
119+
output_len,
120+
)?;
121+
122+
// SAFETY: contiguous gather preserves the decimal dtype and value representation.
123+
Ok(
124+
unsafe { DecimalArray::new_unchecked(values, array.decimal_dtype(), validity) }
125+
.into_array(),
126+
)
84127
})
85128
}
86129

@@ -95,8 +138,8 @@ fn take_piecewise_to_buffer<S, L, T>(
95138
output_len: usize,
96139
) -> VortexResult<Buffer<T>>
97140
where
98-
S: crate::dtype::UnsignedPType,
99-
L: crate::dtype::UnsignedPType,
141+
S: UnsignedPType,
142+
L: UnsignedPType,
100143
T: NativeDecimalType,
101144
{
102145
validate_index_ranges(values.len(), starts, lengths, output_len)?;

vortex-array/src/arrays/primitive/compute/take/mod.rs

Lines changed: 59 additions & 16 deletions
Original file line numberDiff line numberDiff line change
@@ -26,6 +26,8 @@ use crate::arrays::piecewise_sequence::validate_index_ranges;
2626
use crate::builtins::ArrayBuiltins;
2727
use crate::dtype::DType;
2828
use crate::dtype::IntegerPType;
29+
use crate::dtype::NativePType;
30+
use crate::dtype::UnsignedPType;
2931
use crate::executor::ExecutionCtx;
3032
use crate::match_each_integer_ptype;
3133
use crate::match_each_native_ptype;
@@ -143,20 +145,10 @@ fn take_piecewise_sequence(
143145
let Some((starts, lengths)) = execute_unit_multiplier_index_arrays(indices, ctx)? else {
144146
return Ok(None);
145147
};
146-
match_each_native_ptype!(array.ptype(), |T| {
147-
match_each_unsigned_integer_ptype!(starts.ptype(), |S| {
148-
match_each_unsigned_integer_ptype!(lengths.ptype(), |L| {
149-
let values = primitive_piecewise_values::<T, S, L>(
150-
array.as_slice::<T>(),
151-
starts.as_slice::<S>(),
152-
lengths.as_slice::<L>(),
153-
indices_ref.len(),
154-
)?;
155-
let validity = array.validity()?.take(indices_ref)?;
156-
Ok(Some(PrimitiveArray::new(values, validity).into_array()))
157-
})
158-
})
159-
})
148+
let validity = array.validity()?.take(indices_ref)?;
149+
let output_len = indices_ref.len();
150+
let taken = take_piecewise_sequence_lengths(array, &starts, &lengths, validity, output_len)?;
151+
Ok(Some(taken))
160152
}
161153

162154
// Compiler may see this as unused based on enabled features
@@ -180,6 +172,57 @@ fn take_primitive_scalar<T: Copy, I: IntegerPType>(buffer: &[T], indices: &[I])
180172
result.freeze()
181173
}
182174

175+
fn take_piecewise_sequence_lengths(
176+
array: ArrayView<'_, Primitive>,
177+
starts: &PrimitiveArray,
178+
lengths: &PrimitiveArray,
179+
validity: Validity,
180+
output_len: usize,
181+
) -> VortexResult<ArrayRef> {
182+
match_each_native_ptype!(array.ptype(), |T| {
183+
take_piecewise_sequence_lengths_typed::<T>(array, starts, lengths, validity, output_len)
184+
})
185+
}
186+
187+
fn take_piecewise_sequence_lengths_typed<T>(
188+
array: ArrayView<'_, Primitive>,
189+
starts: &PrimitiveArray,
190+
lengths: &PrimitiveArray,
191+
validity: Validity,
192+
output_len: usize,
193+
) -> VortexResult<ArrayRef>
194+
where
195+
T: NativePType,
196+
{
197+
match_each_unsigned_integer_ptype!(starts.ptype(), |S| {
198+
take_piecewise_sequence_lengths_start_typed::<T, S>(
199+
array, starts, lengths, validity, output_len,
200+
)
201+
})
202+
}
203+
204+
fn take_piecewise_sequence_lengths_start_typed<T, S>(
205+
array: ArrayView<'_, Primitive>,
206+
starts: &PrimitiveArray,
207+
lengths: &PrimitiveArray,
208+
validity: Validity,
209+
output_len: usize,
210+
) -> VortexResult<ArrayRef>
211+
where
212+
T: NativePType,
213+
S: UnsignedPType,
214+
{
215+
match_each_unsigned_integer_ptype!(lengths.ptype(), |L| {
216+
let values = primitive_piecewise_values::<T, S, L>(
217+
array.as_slice::<T>(),
218+
starts.as_slice::<S>(),
219+
lengths.as_slice::<L>(),
220+
output_len,
221+
)?;
222+
Ok(PrimitiveArray::new(values, validity).into_array())
223+
})
224+
}
225+
183226
fn primitive_piecewise_values<T, S, L>(
184227
source: &[T],
185228
starts: &[S],
@@ -188,8 +231,8 @@ fn primitive_piecewise_values<T, S, L>(
188231
) -> VortexResult<Buffer<T>>
189232
where
190233
T: Copy,
191-
S: crate::dtype::UnsignedPType,
192-
L: crate::dtype::UnsignedPType,
234+
S: UnsignedPType,
235+
L: UnsignedPType,
193236
{
194237
validate_index_ranges(source.len(), starts, lengths, output_len)?;
195238

0 commit comments

Comments
 (0)