Skip to content

Commit 71c3d0e

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

10 files changed

Lines changed: 928 additions & 103 deletions

File tree

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

Lines changed: 48 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -20,8 +20,10 @@ use crate::arrays::PiecewiseSequence;
2020
use crate::arrays::PrimitiveArray;
2121
use crate::arrays::bool::BoolArrayExt;
2222
use crate::arrays::dict::TakeExecute;
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::builtins::ArrayBuiltins;
2628
use crate::dtype::UnsignedPType;
2729
use crate::executor::ExecutionCtx;
@@ -76,16 +78,32 @@ fn take_piecewise_sequence(
7678
let Some((starts, lengths)) = execute_unit_multiplier_index_arrays(indices, ctx)? else {
7779
return Ok(None);
7880
};
79-
let buffer = match_each_unsigned_integer_ptype!(starts.ptype(), |S| {
80-
match_each_unsigned_integer_ptype!(lengths.ptype(), |L| {
81-
take_piecewise_bits(
82-
&array.to_bit_buffer(),
83-
starts.as_slice::<S>(),
84-
lengths.as_slice::<L>(),
85-
indices_ref.len(),
86-
)?
87-
})
88-
});
81+
let source = array.to_bit_buffer();
82+
let output_len = indices_ref.len();
83+
let buffer = match &lengths {
84+
UnitMultiplierLengths::Constant(length) => {
85+
match_each_unsigned_integer_ptype!(starts.ptype(), |S| {
86+
take_piecewise_bits_constant_length(
87+
&source,
88+
starts.as_slice::<S>(),
89+
*length,
90+
output_len,
91+
)?
92+
})
93+
}
94+
UnitMultiplierLengths::Array(lengths) => {
95+
match_each_unsigned_integer_ptype!(starts.ptype(), |S| {
96+
match_each_unsigned_integer_ptype!(lengths.ptype(), |L| {
97+
take_piecewise_bits(
98+
&source,
99+
starts.as_slice::<S>(),
100+
lengths.as_slice::<L>(),
101+
output_len,
102+
)?
103+
})
104+
})
105+
}
106+
};
89107

90108
Ok(Some(
91109
BoolArray::new(buffer, array.validity()?.take(indices_ref)?).into_array(),
@@ -119,6 +137,26 @@ fn take_bool_impl<I: AsPrimitive<usize>>(bools: BitBufferView<'_>, indices: &[I]
119137
})
120138
}
121139

140+
fn take_piecewise_bits_constant_length<S>(
141+
source: &BitBuffer,
142+
starts: &[S],
143+
length: usize,
144+
output_len: usize,
145+
) -> VortexResult<BitBuffer>
146+
where
147+
S: UnsignedPType,
148+
{
149+
validate_index_ranges_constant(source.len(), starts, length, output_len)?;
150+
151+
let mut values = BitBufferMut::with_capacity(output_len);
152+
for &start in starts {
153+
let start = start.as_();
154+
values.append_buffer(&source.slice(start..start + length));
155+
}
156+
157+
Ok(values.freeze())
158+
}
159+
122160
fn take_piecewise_bits<S, L>(
123161
source: &BitBuffer,
124162
starts: &[S],

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

Lines changed: 71 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -14,8 +14,10 @@ use crate::arrays::DecimalArray;
1414
use crate::arrays::PiecewiseSequence;
1515
use crate::arrays::PrimitiveArray;
1616
use crate::arrays::dict::TakeExecute;
17+
use crate::arrays::piecewise_sequence::UnitMultiplierLengths;
1718
use crate::arrays::piecewise_sequence::execute_unit_multiplier_index_arrays;
1819
use crate::arrays::piecewise_sequence::validate_index_ranges;
20+
use crate::arrays::piecewise_sequence::validate_index_ranges_constant;
1921
use crate::dtype::IntegerPType;
2022
use crate::dtype::NativeDecimalType;
2123
use crate::dtype::UnsignedPType;
@@ -67,10 +69,57 @@ fn take_piecewise_sequence(
6769
};
6870
let validity = array.validity()?.take(indices_ref)?;
6971
let output_len = indices_ref.len();
70-
let taken = take_piecewise_sequence_lengths(array, &starts, &lengths, validity, output_len)?;
72+
let taken = match lengths {
73+
UnitMultiplierLengths::Constant(length) => {
74+
take_piecewise_sequence_constant_length(array, &starts, length, validity, output_len)?
75+
}
76+
UnitMultiplierLengths::Array(lengths) => {
77+
take_piecewise_sequence_lengths(array, &starts, &lengths, validity, output_len)?
78+
}
79+
};
7180
Ok(Some(taken))
7281
}
7382

83+
fn take_piecewise_sequence_constant_length(
84+
array: ArrayView<'_, Decimal>,
85+
starts: &PrimitiveArray,
86+
length: usize,
87+
validity: Validity,
88+
output_len: usize,
89+
) -> VortexResult<ArrayRef> {
90+
match_each_decimal_value_type!(array.values_type(), |D| {
91+
take_piecewise_sequence_constant_length_typed::<D>(
92+
array, starts, length, validity, output_len,
93+
)
94+
})
95+
}
96+
97+
fn take_piecewise_sequence_constant_length_typed<D>(
98+
array: ArrayView<'_, Decimal>,
99+
starts: &PrimitiveArray,
100+
length: usize,
101+
validity: Validity,
102+
output_len: usize,
103+
) -> VortexResult<ArrayRef>
104+
where
105+
D: NativeDecimalType,
106+
{
107+
match_each_unsigned_integer_ptype!(starts.ptype(), |S| {
108+
let values = take_piecewise_constant_length_to_buffer::<S, D>(
109+
starts.as_slice::<S>(),
110+
length,
111+
array.buffer::<D>().as_slice(),
112+
output_len,
113+
)?;
114+
115+
// SAFETY: contiguous gather preserves the decimal dtype and value representation.
116+
Ok(
117+
unsafe { DecimalArray::new_unchecked(values, array.decimal_dtype(), validity) }
118+
.into_array(),
119+
)
120+
})
121+
}
122+
74123
fn take_piecewise_sequence_lengths(
75124
array: ArrayView<'_, Decimal>,
76125
starts: &PrimitiveArray,
@@ -131,6 +180,27 @@ fn take_to_buffer<I: IntegerPType, T: NativeDecimalType>(indices: &[I], values:
131180
indices.iter().map(|idx| values[idx.as_()]).collect()
132181
}
133182

183+
fn take_piecewise_constant_length_to_buffer<S, T>(
184+
starts: &[S],
185+
length: usize,
186+
values: &[T],
187+
output_len: usize,
188+
) -> VortexResult<Buffer<T>>
189+
where
190+
S: UnsignedPType,
191+
T: NativeDecimalType,
192+
{
193+
validate_index_ranges_constant(values.len(), starts, length, output_len)?;
194+
195+
let mut result = BufferMut::<T>::with_capacity(output_len);
196+
for &start in starts {
197+
let start = start.as_();
198+
result.extend_from_slice(&values[start..start + length]);
199+
}
200+
201+
Ok(result.freeze())
202+
}
203+
134204
fn take_piecewise_to_buffer<S, L, T>(
135205
starts: &[S],
136206
lengths: &[L],

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

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -84,7 +84,7 @@ fn take_empty_fsl(
8484
FixedSizeListArray::new_unchecked(new_elements, array.list_size(), new_validity, new_len)
8585
}
8686
.into_array()
87-
.optimize()
87+
.optimize_ctx(ctx.session())
8888
}
8989

9090
fn take_non_empty_fsl(
@@ -147,7 +147,7 @@ fn take_non_empty_degenerate_fsl(
147147
)
148148
}
149149
.into_array()
150-
.optimize()
150+
.optimize_ctx(ctx.session())
151151
}
152152

153153
fn take_non_empty_non_degenerate_fsl<I: IntegerPType>(
@@ -174,7 +174,7 @@ fn take_non_empty_non_degenerate_fsl<I: IntegerPType>(
174174
FixedSizeListArray::new_unchecked(new_elements, array.list_size(), new_validity, new_len)
175175
}
176176
.into_array()
177-
.optimize()
177+
.optimize_ctx(ctx.session())
178178
}
179179

180180
fn take_non_empty_non_degenerate_elements<I: IntegerPType>(

0 commit comments

Comments
 (0)