Skip to main content

lance_bitpacking/bitpacker_internal/
bitpacker8x.rs

1// SPDX-License-Identifier: Apache-2.0
2// SPDX-FileCopyrightText: Copyright The Lance Authors
3
4use super::BitPacker;
5
6use crate::bitpacker_internal::{Available, UnsafeBitPacker};
7
8const BLOCK_LEN: usize = 32 * 8;
9
10#[cfg(target_arch = "x86_64")]
11mod avx2 {
12    use super::BLOCK_LEN;
13    use crate::bitpacker_internal::Available;
14
15    use std::arch::x86_64::__m256i as DataType;
16    use std::arch::x86_64::_mm256_and_si256 as op_and;
17    use std::arch::x86_64::_mm256_lddqu_si256 as load_unaligned;
18    use std::arch::x86_64::_mm256_or_si256 as op_or;
19    use std::arch::x86_64::_mm256_set1_epi32 as set1;
20    use std::arch::x86_64::_mm256_slli_epi32 as left_shift_32;
21    use std::arch::x86_64::_mm256_srli_epi32 as right_shift_32;
22    use std::arch::x86_64::_mm256_storeu_si256 as store_unaligned;
23
24    use std::arch::x86_64::{
25        _mm256_add_epi32, _mm256_extract_epi32, _mm256_permute2f128_si256, _mm256_shuffle_epi32,
26        _mm256_slli_si256, _mm256_srli_si256, _mm256_sub_epi32,
27    };
28
29    #[allow(non_snake_case)]
30    unsafe fn or_collapse_to_u32(accumulator: DataType) -> u32 {
31        let a__b__c__d__e__f__g__h_ = accumulator;
32        let ______a__b________e__f = _mm256_srli_si256(a__b__c__d__e__f__g__h_, 8);
33        let a__b__ca_db_e__f__ge_hf = op_or(a__b__c__d__e__f__g__h_, ______a__b________e__f);
34        let ___a__b__ca____e__f__ge = _mm256_srli_si256(a__b__ca_db_e__f__ge_hf, 4);
35        let _________cadb______gehf = op_or(a__b__ca_db_e__f__ge_hf, ___a__b__ca____e__f__ge);
36        let cadb = _mm256_extract_epi32(_________cadb______gehf, 0);
37        let gehf = _mm256_extract_epi32(_________cadb______gehf, 4);
38        (cadb | gehf) as u32
39    }
40
41    unsafe fn compute_delta(curr: DataType, prev: DataType) -> DataType {
42        let left_shift = _mm256_slli_si256(curr, 4);
43        let curr_shift = _mm256_srli_si256(curr, 12);
44        let curr_right_only = _mm256_permute2f128_si256(curr_shift, curr_shift, 8);
45        let prev_shift = _mm256_srli_si256(prev, 12);
46        let sub_left = _mm256_permute2f128_si256(prev_shift, prev_shift, 3 | (8 << 4));
47        let diff = op_or(left_shift, op_or(curr_right_only, sub_left));
48        _mm256_sub_epi32(curr, diff)
49    }
50
51    #[allow(non_snake_case)]
52    unsafe fn integrate_delta(prev: DataType, delta: DataType) -> DataType {
53        let offset_repeat = _mm256_shuffle_epi32(prev, 0xff);
54        let offset = _mm256_permute2f128_si256(offset_repeat, offset_repeat, 3 | (8 << 4));
55        let a__b__c__d__e__f__g__h__ = delta;
56        let ______a__b________e__f__ = _mm256_slli_si256(delta, 8);
57        let a__b__ca_db_e__f__ge_fh_ =
58            _mm256_add_epi32(a__b__c__d__e__f__g__h__, ______a__b________e__f__);
59        let ___a__b__ca____e__f__ge_ = _mm256_slli_si256(a__b__ca_db_e__f__ge_fh_, 4);
60        let halved_prefix_sum =
61            _mm256_add_epi32(___a__b__ca____e__f__ge_, a__b__ca_db_e__f__ge_fh_);
62        let offsetted_halved_prefix_sum = _mm256_add_epi32(halved_prefix_sum, offset);
63        let select_last_low = _mm256_shuffle_epi32(offsetted_halved_prefix_sum, 0xff);
64        let high_offset = _mm256_permute2f128_si256(select_last_low, select_last_low, 8);
65        _mm256_add_epi32(high_offset, offsetted_halved_prefix_sum)
66    }
67
68    unsafe fn add(left: DataType, right: DataType) -> DataType {
69        _mm256_add_epi32(left, right)
70    }
71
72    unsafe fn sub(left: DataType, right: DataType) -> DataType {
73        _mm256_sub_epi32(left, right)
74    }
75
76    declare_bitpacker!(target_feature(enable = "avx2"));
77
78    impl Available for UnsafeBitPackerImpl {
79        fn available() -> bool {
80            is_x86_feature_detected!("avx2")
81        }
82    }
83}
84
85#[cfg(all(target_arch = "aarch64", target_endian = "little"))]
86mod neon {
87    use super::BLOCK_LEN;
88    use crate::bitpacker_internal::Available;
89    use std::arch::aarch64::{
90        uint32x4_t, vaddq_u32, vandq_u32, vdupq_n_u32, vextq_u32, vgetq_lane_u32, vld1q_u32,
91        vorrq_u32, vshlq_n_u32, vshrq_n_u32, vst1q_u32, vsubq_u32,
92    };
93
94    pub(crate) type DataType = [uint32x4_t; 2];
95
96    #[inline]
97    unsafe fn set1(el: i32) -> DataType {
98        let lanes = vdupq_n_u32(el as u32);
99        [lanes, lanes]
100    }
101
102    #[inline]
103    unsafe fn right_shift_32<const N: i32>(el: DataType) -> DataType {
104        const {
105            assert!(N >= 0);
106            assert!(N <= 32);
107        }
108
109        match N {
110            0 => el,
111            1 => [vshrq_n_u32::<1>(el[0]), vshrq_n_u32::<1>(el[1])],
112            2 => [vshrq_n_u32::<2>(el[0]), vshrq_n_u32::<2>(el[1])],
113            3 => [vshrq_n_u32::<3>(el[0]), vshrq_n_u32::<3>(el[1])],
114            4 => [vshrq_n_u32::<4>(el[0]), vshrq_n_u32::<4>(el[1])],
115            5 => [vshrq_n_u32::<5>(el[0]), vshrq_n_u32::<5>(el[1])],
116            6 => [vshrq_n_u32::<6>(el[0]), vshrq_n_u32::<6>(el[1])],
117            7 => [vshrq_n_u32::<7>(el[0]), vshrq_n_u32::<7>(el[1])],
118            8 => [vshrq_n_u32::<8>(el[0]), vshrq_n_u32::<8>(el[1])],
119            9 => [vshrq_n_u32::<9>(el[0]), vshrq_n_u32::<9>(el[1])],
120            10 => [vshrq_n_u32::<10>(el[0]), vshrq_n_u32::<10>(el[1])],
121            11 => [vshrq_n_u32::<11>(el[0]), vshrq_n_u32::<11>(el[1])],
122            12 => [vshrq_n_u32::<12>(el[0]), vshrq_n_u32::<12>(el[1])],
123            13 => [vshrq_n_u32::<13>(el[0]), vshrq_n_u32::<13>(el[1])],
124            14 => [vshrq_n_u32::<14>(el[0]), vshrq_n_u32::<14>(el[1])],
125            15 => [vshrq_n_u32::<15>(el[0]), vshrq_n_u32::<15>(el[1])],
126            16 => [vshrq_n_u32::<16>(el[0]), vshrq_n_u32::<16>(el[1])],
127            17 => [vshrq_n_u32::<17>(el[0]), vshrq_n_u32::<17>(el[1])],
128            18 => [vshrq_n_u32::<18>(el[0]), vshrq_n_u32::<18>(el[1])],
129            19 => [vshrq_n_u32::<19>(el[0]), vshrq_n_u32::<19>(el[1])],
130            20 => [vshrq_n_u32::<20>(el[0]), vshrq_n_u32::<20>(el[1])],
131            21 => [vshrq_n_u32::<21>(el[0]), vshrq_n_u32::<21>(el[1])],
132            22 => [vshrq_n_u32::<22>(el[0]), vshrq_n_u32::<22>(el[1])],
133            23 => [vshrq_n_u32::<23>(el[0]), vshrq_n_u32::<23>(el[1])],
134            24 => [vshrq_n_u32::<24>(el[0]), vshrq_n_u32::<24>(el[1])],
135            25 => [vshrq_n_u32::<25>(el[0]), vshrq_n_u32::<25>(el[1])],
136            26 => [vshrq_n_u32::<26>(el[0]), vshrq_n_u32::<26>(el[1])],
137            27 => [vshrq_n_u32::<27>(el[0]), vshrq_n_u32::<27>(el[1])],
138            28 => [vshrq_n_u32::<28>(el[0]), vshrq_n_u32::<28>(el[1])],
139            29 => [vshrq_n_u32::<29>(el[0]), vshrq_n_u32::<29>(el[1])],
140            30 => [vshrq_n_u32::<30>(el[0]), vshrq_n_u32::<30>(el[1])],
141            31 => [vshrq_n_u32::<31>(el[0]), vshrq_n_u32::<31>(el[1])],
142            32 => set1(0),
143            _ => core::hint::unreachable_unchecked(),
144        }
145    }
146
147    #[inline]
148    unsafe fn left_shift_32<const N: i32>(el: DataType) -> DataType {
149        const {
150            assert!(N >= 0);
151            assert!(N <= 32);
152        }
153
154        match N {
155            0 => el,
156            1 => [vshlq_n_u32::<1>(el[0]), vshlq_n_u32::<1>(el[1])],
157            2 => [vshlq_n_u32::<2>(el[0]), vshlq_n_u32::<2>(el[1])],
158            3 => [vshlq_n_u32::<3>(el[0]), vshlq_n_u32::<3>(el[1])],
159            4 => [vshlq_n_u32::<4>(el[0]), vshlq_n_u32::<4>(el[1])],
160            5 => [vshlq_n_u32::<5>(el[0]), vshlq_n_u32::<5>(el[1])],
161            6 => [vshlq_n_u32::<6>(el[0]), vshlq_n_u32::<6>(el[1])],
162            7 => [vshlq_n_u32::<7>(el[0]), vshlq_n_u32::<7>(el[1])],
163            8 => [vshlq_n_u32::<8>(el[0]), vshlq_n_u32::<8>(el[1])],
164            9 => [vshlq_n_u32::<9>(el[0]), vshlq_n_u32::<9>(el[1])],
165            10 => [vshlq_n_u32::<10>(el[0]), vshlq_n_u32::<10>(el[1])],
166            11 => [vshlq_n_u32::<11>(el[0]), vshlq_n_u32::<11>(el[1])],
167            12 => [vshlq_n_u32::<12>(el[0]), vshlq_n_u32::<12>(el[1])],
168            13 => [vshlq_n_u32::<13>(el[0]), vshlq_n_u32::<13>(el[1])],
169            14 => [vshlq_n_u32::<14>(el[0]), vshlq_n_u32::<14>(el[1])],
170            15 => [vshlq_n_u32::<15>(el[0]), vshlq_n_u32::<15>(el[1])],
171            16 => [vshlq_n_u32::<16>(el[0]), vshlq_n_u32::<16>(el[1])],
172            17 => [vshlq_n_u32::<17>(el[0]), vshlq_n_u32::<17>(el[1])],
173            18 => [vshlq_n_u32::<18>(el[0]), vshlq_n_u32::<18>(el[1])],
174            19 => [vshlq_n_u32::<19>(el[0]), vshlq_n_u32::<19>(el[1])],
175            20 => [vshlq_n_u32::<20>(el[0]), vshlq_n_u32::<20>(el[1])],
176            21 => [vshlq_n_u32::<21>(el[0]), vshlq_n_u32::<21>(el[1])],
177            22 => [vshlq_n_u32::<22>(el[0]), vshlq_n_u32::<22>(el[1])],
178            23 => [vshlq_n_u32::<23>(el[0]), vshlq_n_u32::<23>(el[1])],
179            24 => [vshlq_n_u32::<24>(el[0]), vshlq_n_u32::<24>(el[1])],
180            25 => [vshlq_n_u32::<25>(el[0]), vshlq_n_u32::<25>(el[1])],
181            26 => [vshlq_n_u32::<26>(el[0]), vshlq_n_u32::<26>(el[1])],
182            27 => [vshlq_n_u32::<27>(el[0]), vshlq_n_u32::<27>(el[1])],
183            28 => [vshlq_n_u32::<28>(el[0]), vshlq_n_u32::<28>(el[1])],
184            29 => [vshlq_n_u32::<29>(el[0]), vshlq_n_u32::<29>(el[1])],
185            30 => [vshlq_n_u32::<30>(el[0]), vshlq_n_u32::<30>(el[1])],
186            31 => [vshlq_n_u32::<31>(el[0]), vshlq_n_u32::<31>(el[1])],
187            32 => set1(0),
188            _ => core::hint::unreachable_unchecked(),
189        }
190    }
191
192    #[inline]
193    unsafe fn op_or(left: DataType, right: DataType) -> DataType {
194        [vorrq_u32(left[0], right[0]), vorrq_u32(left[1], right[1])]
195    }
196
197    #[inline]
198    unsafe fn op_and(left: DataType, right: DataType) -> DataType {
199        [vandq_u32(left[0], right[0]), vandq_u32(left[1], right[1])]
200    }
201
202    #[inline]
203    unsafe fn load_unaligned(addr: *const DataType) -> DataType {
204        let ptr = addr.cast::<u32>();
205        [vld1q_u32(ptr), vld1q_u32(ptr.add(4))]
206    }
207
208    #[inline]
209    unsafe fn store_unaligned(addr: *mut DataType, data: DataType) {
210        let ptr = addr.cast::<u32>();
211        vst1q_u32(ptr, data[0]);
212        vst1q_u32(ptr.add(4), data[1]);
213    }
214
215    #[inline]
216    unsafe fn or_collapse_to_u32(accumulator: DataType) -> u32 {
217        vgetq_lane_u32(accumulator[0], 0)
218            | vgetq_lane_u32(accumulator[0], 1)
219            | vgetq_lane_u32(accumulator[0], 2)
220            | vgetq_lane_u32(accumulator[0], 3)
221            | vgetq_lane_u32(accumulator[1], 0)
222            | vgetq_lane_u32(accumulator[1], 1)
223            | vgetq_lane_u32(accumulator[1], 2)
224            | vgetq_lane_u32(accumulator[1], 3)
225    }
226
227    #[inline]
228    unsafe fn compute_delta(curr: DataType, prev: DataType) -> DataType {
229        [
230            vsubq_u32(curr[0], vextq_u32(prev[1], curr[0], 3)),
231            vsubq_u32(curr[1], vextq_u32(curr[0], curr[1], 3)),
232        ]
233    }
234
235    #[allow(non_snake_case)]
236    #[inline]
237    unsafe fn integrate_half(base: u32, delta: uint32x4_t) -> uint32x4_t {
238        let base = vdupq_n_u32(base);
239        let zero = vdupq_n_u32(0);
240        let a__b__c__d_ = delta;
241        let ______a__b_ = vextq_u32(zero, a__b__c__d_, 2);
242        let a__b__ca_db = vaddq_u32(______a__b_, a__b__c__d_);
243        let ___a__b__ca = vextq_u32(zero, a__b__ca_db, 3);
244        let a_ab_abc_abcd = vaddq_u32(___a__b__ca, a__b__ca_db);
245        vaddq_u32(base, a_ab_abc_abcd)
246    }
247
248    #[inline]
249    unsafe fn integrate_delta(prev: DataType, delta: DataType) -> DataType {
250        let low = integrate_half(vgetq_lane_u32(prev[1], 3), delta[0]);
251        let high = integrate_half(vgetq_lane_u32(low, 3), delta[1]);
252        [low, high]
253    }
254
255    #[inline]
256    unsafe fn add(left: DataType, right: DataType) -> DataType {
257        [vaddq_u32(left[0], right[0]), vaddq_u32(left[1], right[1])]
258    }
259
260    #[inline]
261    unsafe fn sub(left: DataType, right: DataType) -> DataType {
262        [vsubq_u32(left[0], right[0]), vsubq_u32(left[1], right[1])]
263    }
264
265    declare_bitpacker!(target_feature(enable = "neon"));
266
267    impl Available for UnsafeBitPackerImpl {
268        fn available() -> bool {
269            std::arch::is_aarch64_feature_detected!("neon")
270        }
271    }
272}
273
274mod scalar {
275    use super::BLOCK_LEN;
276    use crate::bitpacker_internal::Available;
277    use std::ptr;
278
279    pub(crate) type DataType = [u32; 8];
280
281    pub(crate) fn set1(el: i32) -> DataType {
282        [el as u32; 8]
283    }
284
285    pub(crate) fn right_shift_32<const N: i32>(el: DataType) -> DataType {
286        [
287            el[0] >> N,
288            el[1] >> N,
289            el[2] >> N,
290            el[3] >> N,
291            el[4] >> N,
292            el[5] >> N,
293            el[6] >> N,
294            el[7] >> N,
295        ]
296    }
297
298    pub(crate) fn left_shift_32<const N: i32>(el: DataType) -> DataType {
299        [
300            el[0] << N,
301            el[1] << N,
302            el[2] << N,
303            el[3] << N,
304            el[4] << N,
305            el[5] << N,
306            el[6] << N,
307            el[7] << N,
308        ]
309    }
310
311    pub(crate) fn op_or(left: DataType, right: DataType) -> DataType {
312        [
313            left[0] | right[0],
314            left[1] | right[1],
315            left[2] | right[2],
316            left[3] | right[3],
317            left[4] | right[4],
318            left[5] | right[5],
319            left[6] | right[6],
320            left[7] | right[7],
321        ]
322    }
323
324    pub(crate) fn op_and(left: DataType, right: DataType) -> DataType {
325        [
326            left[0] & right[0],
327            left[1] & right[1],
328            left[2] & right[2],
329            left[3] & right[3],
330            left[4] & right[4],
331            left[5] & right[5],
332            left[6] & right[6],
333            left[7] & right[7],
334        ]
335    }
336
337    pub(crate) unsafe fn load_unaligned(addr: *const DataType) -> DataType {
338        ptr::read_unaligned(addr)
339    }
340
341    pub(crate) unsafe fn store_unaligned(addr: *mut DataType, data: DataType) {
342        ptr::write_unaligned(addr, data);
343    }
344
345    pub(crate) fn or_collapse_to_u32(accumulator: DataType) -> u32 {
346        ((accumulator[0] | accumulator[1]) | (accumulator[2] | accumulator[3]))
347            | ((accumulator[4] | accumulator[5]) | (accumulator[6] | accumulator[7]))
348    }
349
350    fn compute_delta(curr: DataType, prev: DataType) -> DataType {
351        [
352            curr[0].wrapping_sub(prev[7]),
353            curr[1].wrapping_sub(curr[0]),
354            curr[2].wrapping_sub(curr[1]),
355            curr[3].wrapping_sub(curr[2]),
356            curr[4].wrapping_sub(curr[3]),
357            curr[5].wrapping_sub(curr[4]),
358            curr[6].wrapping_sub(curr[5]),
359            curr[7].wrapping_sub(curr[6]),
360        ]
361    }
362
363    fn integrate_delta(offset: DataType, delta: DataType) -> DataType {
364        let el0 = offset[7].wrapping_add(delta[0]);
365        let el1 = el0.wrapping_add(delta[1]);
366        let el2 = el1.wrapping_add(delta[2]);
367        let el3 = el2.wrapping_add(delta[3]);
368        let el4 = el3.wrapping_add(delta[4]);
369        let el5 = el4.wrapping_add(delta[5]);
370        let el6 = el5.wrapping_add(delta[6]);
371        let el7 = el6.wrapping_add(delta[7]);
372        [el0, el1, el2, el3, el4, el5, el6, el7]
373    }
374
375    pub(crate) fn add(left: DataType, right: DataType) -> DataType {
376        [
377            left[0].wrapping_add(right[0]),
378            left[1].wrapping_add(right[1]),
379            left[2].wrapping_add(right[2]),
380            left[3].wrapping_add(right[3]),
381            left[4].wrapping_add(right[4]),
382            left[5].wrapping_add(right[5]),
383            left[6].wrapping_add(right[6]),
384            left[7].wrapping_add(right[7]),
385        ]
386    }
387
388    pub(crate) fn sub(left: DataType, right: DataType) -> DataType {
389        [
390            left[0].wrapping_sub(right[0]),
391            left[1].wrapping_sub(right[1]),
392            left[2].wrapping_sub(right[2]),
393            left[3].wrapping_sub(right[3]),
394            left[4].wrapping_sub(right[4]),
395            left[5].wrapping_sub(right[5]),
396            left[6].wrapping_sub(right[6]),
397            left[7].wrapping_sub(right[7]),
398        ]
399    }
400
401    // The `allow(unused)` is here to put an attribute that has no effect.
402    //
403    // For other bitpackers, we enable a specific CPU instruction set, but for
404    // the scalar bitpacker none is required.
405    declare_bitpacker!(allow(unused));
406
407    impl Available for UnsafeBitPackerImpl {
408        fn available() -> bool {
409            true
410        }
411    }
412}
413
414#[derive(Clone, Copy)]
415enum InstructionSet {
416    #[cfg(target_arch = "x86_64")]
417    AVX2,
418    #[cfg(all(target_arch = "aarch64", target_endian = "little"))]
419    NEON,
420    Scalar,
421}
422
423/// 8-wide bitpacker implementation.
424///
425/// One block contains 256 integers.
426#[derive(Clone, Copy)]
427pub struct BitPacker8x(InstructionSet);
428
429impl BitPacker8x {
430    #[cfg(target_arch = "x86_64")]
431    pub(crate) fn new_avx2() -> Option<Self> {
432        avx2::UnsafeBitPackerImpl::available().then_some(BitPacker8x(InstructionSet::AVX2))
433    }
434
435    #[cfg(not(target_arch = "x86_64"))]
436    pub(crate) fn new_avx2() -> Option<Self> {
437        None
438    }
439
440    #[cfg(all(target_arch = "aarch64", target_endian = "little"))]
441    pub(crate) fn new_neon() -> Option<Self> {
442        neon::UnsafeBitPackerImpl::available().then_some(BitPacker8x(InstructionSet::NEON))
443    }
444
445    #[cfg(not(all(target_arch = "aarch64", target_endian = "little")))]
446    pub(crate) fn new_neon() -> Option<Self> {
447        None
448    }
449
450    pub(crate) fn new_scalar() -> Self {
451        BitPacker8x(InstructionSet::Scalar)
452    }
453}
454
455impl BitPacker for BitPacker8x {
456    const BLOCK_LEN: usize = BLOCK_LEN;
457
458    fn new() -> Self {
459        Self::new_avx2()
460            .or_else(Self::new_neon)
461            .unwrap_or_else(Self::new_scalar)
462    }
463
464    fn compress(&self, decompressed: &[u32], compressed: &mut [u8], num_bits: u8) -> usize {
465        unsafe {
466            match self.0 {
467                #[cfg(target_arch = "x86_64")]
468                InstructionSet::AVX2 => {
469                    avx2::UnsafeBitPackerImpl::compress(decompressed, compressed, num_bits)
470                }
471                #[cfg(all(target_arch = "aarch64", target_endian = "little"))]
472                InstructionSet::NEON => {
473                    neon::UnsafeBitPackerImpl::compress(decompressed, compressed, num_bits)
474                }
475                InstructionSet::Scalar => {
476                    scalar::UnsafeBitPackerImpl::compress(decompressed, compressed, num_bits)
477                }
478            }
479        }
480    }
481
482    fn compress_sorted(
483        &self,
484        initial: u32,
485        decompressed: &[u32],
486        compressed: &mut [u8],
487        num_bits: u8,
488    ) -> usize {
489        unsafe {
490            match self.0 {
491                #[cfg(target_arch = "x86_64")]
492                InstructionSet::AVX2 => avx2::UnsafeBitPackerImpl::compress_sorted(
493                    initial,
494                    decompressed,
495                    compressed,
496                    num_bits,
497                ),
498                #[cfg(all(target_arch = "aarch64", target_endian = "little"))]
499                InstructionSet::NEON => neon::UnsafeBitPackerImpl::compress_sorted(
500                    initial,
501                    decompressed,
502                    compressed,
503                    num_bits,
504                ),
505                InstructionSet::Scalar => scalar::UnsafeBitPackerImpl::compress_sorted(
506                    initial,
507                    decompressed,
508                    compressed,
509                    num_bits,
510                ),
511            }
512        }
513    }
514
515    fn compress_strictly_sorted(
516        &self,
517        initial: Option<u32>,
518        decompressed: &[u32],
519        compressed: &mut [u8],
520        num_bits: u8,
521    ) -> usize {
522        unsafe {
523            match self.0 {
524                #[cfg(target_arch = "x86_64")]
525                InstructionSet::AVX2 => avx2::UnsafeBitPackerImpl::compress_strictly_sorted(
526                    initial,
527                    decompressed,
528                    compressed,
529                    num_bits,
530                ),
531                #[cfg(all(target_arch = "aarch64", target_endian = "little"))]
532                InstructionSet::NEON => neon::UnsafeBitPackerImpl::compress_strictly_sorted(
533                    initial,
534                    decompressed,
535                    compressed,
536                    num_bits,
537                ),
538                InstructionSet::Scalar => scalar::UnsafeBitPackerImpl::compress_strictly_sorted(
539                    initial,
540                    decompressed,
541                    compressed,
542                    num_bits,
543                ),
544            }
545        }
546    }
547
548    fn decompress(&self, compressed: &[u8], decompressed: &mut [u32], num_bits: u8) -> usize {
549        unsafe {
550            match self.0 {
551                #[cfg(target_arch = "x86_64")]
552                InstructionSet::AVX2 => {
553                    avx2::UnsafeBitPackerImpl::decompress(compressed, decompressed, num_bits)
554                }
555                #[cfg(all(target_arch = "aarch64", target_endian = "little"))]
556                InstructionSet::NEON => {
557                    neon::UnsafeBitPackerImpl::decompress(compressed, decompressed, num_bits)
558                }
559                InstructionSet::Scalar => {
560                    scalar::UnsafeBitPackerImpl::decompress(compressed, decompressed, num_bits)
561                }
562            }
563        }
564    }
565
566    fn decompress_sorted(
567        &self,
568        initial: u32,
569        compressed: &[u8],
570        decompressed: &mut [u32],
571        num_bits: u8,
572    ) -> usize {
573        unsafe {
574            match self.0 {
575                #[cfg(target_arch = "x86_64")]
576                InstructionSet::AVX2 => avx2::UnsafeBitPackerImpl::decompress_sorted(
577                    initial,
578                    compressed,
579                    decompressed,
580                    num_bits,
581                ),
582                #[cfg(all(target_arch = "aarch64", target_endian = "little"))]
583                InstructionSet::NEON => neon::UnsafeBitPackerImpl::decompress_sorted(
584                    initial,
585                    compressed,
586                    decompressed,
587                    num_bits,
588                ),
589                InstructionSet::Scalar => scalar::UnsafeBitPackerImpl::decompress_sorted(
590                    initial,
591                    compressed,
592                    decompressed,
593                    num_bits,
594                ),
595            }
596        }
597    }
598
599    fn decompress_strictly_sorted(
600        &self,
601        initial: Option<u32>,
602        compressed: &[u8],
603        decompressed: &mut [u32],
604        num_bits: u8,
605    ) -> usize {
606        unsafe {
607            match self.0 {
608                #[cfg(target_arch = "x86_64")]
609                InstructionSet::AVX2 => avx2::UnsafeBitPackerImpl::decompress_strictly_sorted(
610                    initial,
611                    compressed,
612                    decompressed,
613                    num_bits,
614                ),
615                #[cfg(all(target_arch = "aarch64", target_endian = "little"))]
616                InstructionSet::NEON => neon::UnsafeBitPackerImpl::decompress_strictly_sorted(
617                    initial,
618                    compressed,
619                    decompressed,
620                    num_bits,
621                ),
622                InstructionSet::Scalar => scalar::UnsafeBitPackerImpl::decompress_strictly_sorted(
623                    initial,
624                    compressed,
625                    decompressed,
626                    num_bits,
627                ),
628            }
629        }
630    }
631
632    fn num_bits(&self, decompressed: &[u32]) -> u8 {
633        unsafe {
634            match self.0 {
635                #[cfg(target_arch = "x86_64")]
636                InstructionSet::AVX2 => avx2::UnsafeBitPackerImpl::num_bits(decompressed),
637                #[cfg(all(target_arch = "aarch64", target_endian = "little"))]
638                InstructionSet::NEON => neon::UnsafeBitPackerImpl::num_bits(decompressed),
639                InstructionSet::Scalar => scalar::UnsafeBitPackerImpl::num_bits(decompressed),
640            }
641        }
642    }
643
644    fn num_bits_sorted(&self, initial: u32, decompressed: &[u32]) -> u8 {
645        unsafe {
646            match self.0 {
647                #[cfg(target_arch = "x86_64")]
648                InstructionSet::AVX2 => {
649                    avx2::UnsafeBitPackerImpl::num_bits_sorted(initial, decompressed)
650                }
651                #[cfg(all(target_arch = "aarch64", target_endian = "little"))]
652                InstructionSet::NEON => {
653                    neon::UnsafeBitPackerImpl::num_bits_sorted(initial, decompressed)
654                }
655                InstructionSet::Scalar => {
656                    scalar::UnsafeBitPackerImpl::num_bits_sorted(initial, decompressed)
657                }
658            }
659        }
660    }
661
662    fn num_bits_strictly_sorted(&self, initial: Option<u32>, decompressed: &[u32]) -> u8 {
663        unsafe {
664            match self.0 {
665                #[cfg(target_arch = "x86_64")]
666                InstructionSet::AVX2 => {
667                    avx2::UnsafeBitPackerImpl::num_bits_strictly_sorted(initial, decompressed)
668                }
669                #[cfg(all(target_arch = "aarch64", target_endian = "little"))]
670                InstructionSet::NEON => {
671                    neon::UnsafeBitPackerImpl::num_bits_strictly_sorted(initial, decompressed)
672                }
673                InstructionSet::Scalar => {
674                    scalar::UnsafeBitPackerImpl::num_bits_strictly_sorted(initial, decompressed)
675                }
676            }
677        }
678    }
679}
680
681#[cfg(test)]
682mod tests {
683    use super::BitPacker8x;
684    use crate::bitpacker_internal::BitPacker;
685    use bitpacking::{BitPacker as ExternalBitPacker, BitPacker8x as ExternalBitPacker8x};
686
687    fn mask_for_width(width: u8) -> u32 {
688        match width {
689            0 => 0,
690            32 => u32::MAX,
691            _ => (1u32 << width) - 1,
692        }
693    }
694
695    fn raw_values(width: u8, seed: u64) -> Vec<u32> {
696        let mask = mask_for_width(width);
697        let mut state = seed;
698        (0..BitPacker8x::BLOCK_LEN)
699            .map(|idx| {
700                state ^= state << 13;
701                state ^= state >> 7;
702                state ^= state << 17;
703                match seed % 4 {
704                    0 => 0,
705                    1 => mask,
706                    2 => idx as u32 & mask,
707                    _ => state as u32 & mask,
708                }
709            })
710            .collect()
711    }
712
713    fn sorted_values(width: u8, seed: u64) -> (u32, Vec<u32>) {
714        if width == 0 {
715            return (17, vec![17; BitPacker8x::BLOCK_LEN]);
716        }
717        if width == 32 {
718            return (0, vec![u32::MAX; BitPacker8x::BLOCK_LEN]);
719        }
720
721        let mask = mask_for_width(width).min(127);
722        let mut state = seed;
723        let mut current = 17u32;
724        let values = (0..BitPacker8x::BLOCK_LEN)
725            .map(|_| {
726                state ^= state << 13;
727                state ^= state >> 7;
728                state ^= state << 17;
729                current += state as u32 & mask;
730                current
731            })
732            .collect();
733        (17, values)
734    }
735
736    fn strictly_sorted_values(width: u8, seed: u64) -> (Option<u32>, Vec<u32>) {
737        let mask = mask_for_width(width).min(127);
738        let mut state = seed;
739        let mut current = 0u32;
740        let values = (0..BitPacker8x::BLOCK_LEN)
741            .map(|idx| {
742                if idx == 0 {
743                    current = 0;
744                } else {
745                    state ^= state << 13;
746                    state ^= state >> 7;
747                    state ^= state << 17;
748                    current += 1 + (state as u32 & mask);
749                }
750                current
751            })
752            .collect();
753        (None, values)
754    }
755
756    fn assert_raw_compatible(ours: BitPacker8x, external: ExternalBitPacker8x) {
757        for width in 0..=32 {
758            for seed in [0, 1, 2, 123456789] {
759                let values = raw_values(width, seed);
760                assert_eq!(ours.num_bits(&values), external.num_bits(&values));
761
762                let mut actual = vec![0u8; BitPacker8x::compressed_block_size(width)];
763                let actual_len = ours.compress(&values, &mut actual, width);
764
765                let mut expected = vec![0u8; ExternalBitPacker8x::compressed_block_size(width)];
766                let expected_len = external.compress(&values, &mut expected, width);
767
768                assert_eq!(actual_len, expected_len);
769                assert_eq!(actual, expected, "raw width {width} seed {seed}");
770
771                let mut decoded = vec![0u32; BitPacker8x::BLOCK_LEN];
772                assert_eq!(ours.decompress(&actual, &mut decoded, width), actual_len);
773                assert_eq!(decoded, values);
774            }
775        }
776    }
777
778    fn assert_sorted_compatible(ours: BitPacker8x, external: ExternalBitPacker8x) {
779        for width in 0..=32 {
780            for seed in [0, 1, 2, 123456789] {
781                let (initial, values) = sorted_values(width, seed);
782                assert_eq!(
783                    ours.num_bits_sorted(initial, &values),
784                    external.num_bits_sorted(initial, &values)
785                );
786
787                let mut actual = vec![0u8; BitPacker8x::compressed_block_size(width)];
788                let actual_len = ours.compress_sorted(initial, &values, &mut actual, width);
789
790                let mut expected = vec![0u8; ExternalBitPacker8x::compressed_block_size(width)];
791                let expected_len = external.compress_sorted(initial, &values, &mut expected, width);
792
793                assert_eq!(actual_len, expected_len);
794                assert_eq!(actual, expected, "sorted width {width} seed {seed}");
795
796                let mut decoded = vec![0u32; BitPacker8x::BLOCK_LEN];
797                assert_eq!(
798                    ours.decompress_sorted(initial, &actual, &mut decoded, width),
799                    actual_len
800                );
801                assert_eq!(decoded, values);
802            }
803        }
804    }
805
806    fn assert_strictly_sorted_compatible(ours: BitPacker8x, external: ExternalBitPacker8x) {
807        for width in 0..=16 {
808            for seed in [0, 1, 2, 123456789] {
809                let (initial, values) = strictly_sorted_values(width, seed);
810                let num_bits = external.num_bits_strictly_sorted(initial, &values);
811                assert_eq!(ours.num_bits_strictly_sorted(initial, &values), num_bits);
812
813                let mut actual = vec![0u8; BitPacker8x::compressed_block_size(num_bits)];
814                let actual_len =
815                    ours.compress_strictly_sorted(initial, &values, &mut actual, num_bits);
816
817                let mut expected = vec![0u8; ExternalBitPacker8x::compressed_block_size(num_bits)];
818                let expected_len =
819                    external.compress_strictly_sorted(initial, &values, &mut expected, num_bits);
820
821                assert_eq!(actual_len, expected_len);
822                assert_eq!(actual, expected, "strict width {width} seed {seed}");
823
824                let mut decoded = vec![0u32; BitPacker8x::BLOCK_LEN];
825                assert_eq!(
826                    ours.decompress_strictly_sorted(initial, &actual, &mut decoded, num_bits),
827                    actual_len
828                );
829                assert_eq!(decoded, values);
830            }
831        }
832    }
833
834    #[test]
835    fn bitpacker8x_raw_compatible_with_external_bitpacking() {
836        assert_raw_compatible(BitPacker8x::new(), ExternalBitPacker8x::new());
837    }
838
839    #[test]
840    fn bitpacker8x_sorted_compatible_with_external_bitpacking() {
841        assert_sorted_compatible(BitPacker8x::new(), ExternalBitPacker8x::new());
842    }
843
844    #[test]
845    fn scalar_backend_matches_external_bitpacker8x() {
846        let scalar = BitPacker8x::new_scalar();
847        let external = ExternalBitPacker8x::new();
848
849        assert_raw_compatible(scalar, external);
850        assert_sorted_compatible(scalar, external);
851    }
852
853    #[test]
854    fn scalar_backend_matches_external_strictly_sorted_bitpacker8x() {
855        let scalar = BitPacker8x::new_scalar();
856        let external = ExternalBitPacker8x::new();
857
858        assert_strictly_sorted_compatible(scalar, external);
859    }
860
861    #[cfg(all(target_arch = "aarch64", target_endian = "little"))]
862    #[test]
863    fn neon_backend_matches_external_bitpacker8x() {
864        if let Some(neon) = BitPacker8x::new_neon() {
865            let external = ExternalBitPacker8x::new();
866
867            assert_raw_compatible(neon, external);
868            assert_sorted_compatible(neon, external);
869            assert_strictly_sorted_compatible(neon, external);
870        }
871    }
872}