Skip to main content

vortex_fastlanes/bitpacking/compute/
between.rs

1// SPDX-License-Identifier: Apache-2.0
2// SPDX-FileCopyrightText: Copyright the Vortex contributors
3
4//! Block-streaming between kernel for [`BitPackedArray`] against constant bounds.
5//!
6//! Reuses the same single-block scratch buffer as the compare kernel and folds a
7//! `lower op_l v op_u upper` predicate per element, so the full primitive never
8//! materialises.
9//!
10//! [`BitPackedArray`]: crate::BitPackedArray
11
12use vortex_array::ArrayRef;
13use vortex_array::ArrayView;
14use vortex_array::ExecutionCtx;
15use vortex_array::dtype::NativePType;
16use vortex_array::dtype::Nullability;
17use vortex_array::match_each_integer_ptype;
18use vortex_array::scalar_fn::fns::between::BetweenKernel;
19use vortex_array::scalar_fn::fns::between::BetweenOptions;
20use vortex_array::scalar_fn::fns::between::StrictComparison;
21use vortex_error::VortexExpect;
22use vortex_error::VortexResult;
23
24use crate::BitPacked;
25use crate::bitpacking::compute::stream_predicate::stream_predicate;
26
27impl BetweenKernel for BitPacked {
28    fn between(
29        array: ArrayView<'_, Self>,
30        lower: &ArrayRef,
31        upper: &ArrayRef,
32        options: &BetweenOptions,
33        ctx: &mut ExecutionCtx,
34    ) -> VortexResult<Option<ArrayRef>> {
35        // Only accelerate constant-bounds between; vary-by-row bounds fall through to the
36        // default `compare + and` pipeline.
37        let (Some(lower_const), Some(upper_const)) = (lower.as_constant(), upper.as_constant())
38        else {
39            return Ok(None);
40        };
41        let (Some(lower_prim), Some(upper_prim)) = (
42            lower_const.as_primitive_opt(),
43            upper_const.as_primitive_opt(),
44        ) else {
45            return Ok(None);
46        };
47
48        let nullability =
49            array.dtype().nullability() | lower.dtype().nullability() | upper.dtype().nullability();
50        let arr_ptype = array.dtype().as_ptype();
51        if lower_prim.ptype() != arr_ptype || upper_prim.ptype() != arr_ptype {
52            return Ok(None);
53        }
54
55        let result = match_each_integer_ptype!(arr_ptype, |T| {
56            let lo: T = lower_prim
57                .typed_value::<T>()
58                .vortex_expect("between precondition strips null lower");
59            let up: T = upper_prim
60                .typed_value::<T>()
61                .vortex_expect("between precondition strips null upper");
62            between_constant_typed::<T>(array, lo, up, options, nullability, ctx)?
63        });
64        Ok(Some(result))
65    }
66}
67
68fn between_constant_typed<T>(
69    array: ArrayView<'_, BitPacked>,
70    lower: T,
71    upper: T,
72    options: &BetweenOptions,
73    nullability: Nullability,
74    ctx: &mut ExecutionCtx,
75) -> VortexResult<ArrayRef>
76where
77    T: NativePType + Copy + crate::unpack_iter::BitPacked,
78{
79    // Branch on strictness once at the top so each call into `between_impl` monomorphises
80    // a single tight predicate — same shape as `Primitive::between` in `vortex-array`.
81    match (options.lower_strict, options.upper_strict) {
82        (StrictComparison::Strict, StrictComparison::Strict) => between_impl(
83            array,
84            lower,
85            NativePType::is_lt,
86            upper,
87            NativePType::is_lt,
88            nullability,
89            ctx,
90        ),
91        (StrictComparison::Strict, StrictComparison::NonStrict) => between_impl(
92            array,
93            lower,
94            NativePType::is_lt,
95            upper,
96            NativePType::is_le,
97            nullability,
98            ctx,
99        ),
100        (StrictComparison::NonStrict, StrictComparison::Strict) => between_impl(
101            array,
102            lower,
103            NativePType::is_le,
104            upper,
105            NativePType::is_lt,
106            nullability,
107            ctx,
108        ),
109        (StrictComparison::NonStrict, StrictComparison::NonStrict) => between_impl(
110            array,
111            lower,
112            NativePType::is_le,
113            upper,
114            NativePType::is_le,
115            nullability,
116            ctx,
117        ),
118    }
119}
120
121fn between_impl<T, Lo, Up>(
122    array: ArrayView<'_, BitPacked>,
123    lower: T,
124    lower_fn: Lo,
125    upper: T,
126    upper_fn: Up,
127    nullability: Nullability,
128    ctx: &mut ExecutionCtx,
129) -> VortexResult<ArrayRef>
130where
131    T: NativePType + Copy + crate::unpack_iter::BitPacked,
132    Lo: Fn(T, T) -> bool,
133    Up: Fn(T, T) -> bool,
134{
135    stream_predicate::<T, _>(
136        array,
137        nullability,
138        |v| lower_fn(lower, v) & upper_fn(v, upper),
139        ctx,
140    )
141}
142
143#[cfg(test)]
144mod tests {
145    use std::sync::LazyLock;
146
147    use rstest::rstest;
148    use vortex_array::IntoArray;
149    use vortex_array::VortexSessionExecute;
150    use vortex_array::arrays::BoolArray;
151    use vortex_array::arrays::ConstantArray;
152    use vortex_array::arrays::PrimitiveArray;
153    use vortex_array::assert_arrays_eq;
154    use vortex_array::builtins::ArrayBuiltins;
155    use vortex_array::scalar_fn::fns::between::BetweenOptions;
156    use vortex_array::scalar_fn::fns::between::StrictComparison;
157    use vortex_array::validity::Validity;
158    use vortex_buffer::BufferMut;
159    use vortex_error::VortexResult;
160    use vortex_session::VortexSession;
161
162    use crate::BitPackedArrayExt;
163    use crate::BitPackedData;
164
165    static SESSION: LazyLock<VortexSession> = LazyLock::new(|| {
166        let session = vortex_array::array_session();
167        crate::initialize(&session);
168        session
169    });
170
171    fn opts(lower: StrictComparison, upper: StrictComparison) -> BetweenOptions {
172        BetweenOptions {
173            lower_strict: lower,
174            upper_strict: upper,
175        }
176    }
177
178    #[rstest]
179    #[case(StrictComparison::NonStrict, StrictComparison::NonStrict)]
180    #[case(StrictComparison::Strict, StrictComparison::NonStrict)]
181    #[case(StrictComparison::NonStrict, StrictComparison::Strict)]
182    #[case(StrictComparison::Strict, StrictComparison::Strict)]
183    fn multi_chunk_against_primitive_baseline(
184        #[case] lower_strict: StrictComparison,
185        #[case] upper_strict: StrictComparison,
186    ) -> VortexResult<()> {
187        let mut ctx = SESSION.create_execution_ctx();
188        let values: BufferMut<u32> = (0..3000u32).map(|i| i % 257).collect();
189        let prim = PrimitiveArray::new(values.freeze(), Validity::NonNullable);
190        let packed = BitPackedData::encode(&prim.clone().into_array(), 9, &mut ctx)?;
191
192        let lower = ConstantArray::new(40u32, prim.len()).into_array();
193        let upper = ConstantArray::new(200u32, prim.len()).into_array();
194        let options = opts(lower_strict, upper_strict);
195
196        let expected = prim
197            .into_array()
198            .between(lower.clone(), upper.clone(), options.clone())?
199            .execute::<BoolArray>(&mut ctx)?;
200        let actual = packed
201            .into_array()
202            .between(lower, upper, options)?
203            .execute::<BoolArray>(&mut ctx)?;
204
205        assert_arrays_eq!(actual, expected, &mut ctx);
206        Ok(())
207    }
208
209    #[test]
210    fn signed_with_patches_against_primitive_baseline() -> VortexResult<()> {
211        let mut ctx = SESSION.create_execution_ctx();
212        let values: Vec<i32> = (0..1500)
213            .map(|i| if i % 73 == 0 { 100_000 + i } else { i % 100 })
214            .collect();
215        let prim = PrimitiveArray::from_iter(values);
216        let packed = BitPackedData::encode(&prim.clone().into_array(), 7, &mut ctx)?;
217        assert!(packed.patches().is_some(), "test setup expects patches");
218
219        let lower = ConstantArray::new(20i32, prim.len()).into_array();
220        let upper = ConstantArray::new(80i32, prim.len()).into_array();
221        let options = opts(StrictComparison::NonStrict, StrictComparison::NonStrict);
222
223        let expected = prim
224            .into_array()
225            .between(lower.clone(), upper.clone(), options.clone())?
226            .execute::<BoolArray>(&mut ctx)?;
227        let actual = packed
228            .into_array()
229            .between(lower, upper, options)?
230            .execute::<BoolArray>(&mut ctx)?;
231
232        assert_arrays_eq!(actual, expected, &mut ctx);
233        Ok(())
234    }
235
236    #[test]
237    fn nullable_propagates_validity() -> VortexResult<()> {
238        let mut ctx = SESSION.create_execution_ctx();
239        let prim =
240            PrimitiveArray::from_option_iter([Some(1u32), None, Some(3), Some(4), None, Some(6)]);
241        let packed = BitPackedData::encode(&prim.clone().into_array(), 3, &mut ctx)?;
242
243        let lower = ConstantArray::new(2u32, packed.len()).into_array();
244        let upper = ConstantArray::new(5u32, packed.len()).into_array();
245        let options = opts(StrictComparison::NonStrict, StrictComparison::NonStrict);
246
247        let actual = packed
248            .into_array()
249            .between(lower.clone(), upper.clone(), options.clone())?
250            .execute::<BoolArray>(&mut ctx)?;
251        let expected = prim
252            .into_array()
253            .between(lower, upper, options)?
254            .execute::<BoolArray>(&mut ctx)?;
255        assert_arrays_eq!(actual, expected, &mut ctx);
256        Ok(())
257    }
258}