Skip to main content

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