Skip to main content

parquet/util/
bit_pack.rs

1// Licensed to the Apache Software Foundation (ASF) under one
2// or more contributor license agreements.  See the NOTICE file
3// distributed with this work for additional information
4// regarding copyright ownership.  The ASF licenses this file
5// to you under the Apache License, Version 2.0 (the
6// "License"); you may not use this file except in compliance
7// with the License.  You may obtain a copy of the License at
8//
9//   http://www.apache.org/licenses/LICENSE-2.0
10//
11// Unless required by applicable law or agreed to in writing,
12// software distributed under the License is distributed on an
13// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
14// KIND, either express or implied.  See the License for the
15// specific language governing permissions and limitations
16// under the License.
17
18//! Vectorised bit-packing utilities
19
20/// Macro that generates an unpack function taking the number of bits as a const generic
21macro_rules! unpack_impl {
22    ($t:ty, $bytes:literal, $bits:tt) => {
23        pub fn unpack<const NUM_BITS: usize>(input: &[u8], output: &mut [$t; $bits]) {
24            if NUM_BITS == 0 {
25                for out in output {
26                    *out = 0;
27                }
28                return;
29            }
30
31            assert!(NUM_BITS <= $bytes * 8);
32
33            let mask = match NUM_BITS {
34                $bits => <$t>::MAX,
35                _ => ((1 << NUM_BITS) - 1),
36            };
37
38            assert!(input.len() >= NUM_BITS * $bytes);
39
40            let r = |output_idx: usize| {
41                <$t>::from_le_bytes(
42                    input[output_idx * $bytes..output_idx * $bytes + $bytes]
43                        .try_into()
44                        .unwrap(),
45                )
46            };
47
48            seq_macro::seq!(i in 0..$bits {
49                let start_bit = i * NUM_BITS;
50                let end_bit = start_bit + NUM_BITS;
51
52                let start_bit_offset = start_bit % $bits;
53                let end_bit_offset = end_bit % $bits;
54                let start_byte = start_bit / $bits;
55                let end_byte = end_bit / $bits;
56                if start_byte != end_byte && end_bit_offset != 0 {
57                    let val = r(start_byte);
58                    let a = val >> start_bit_offset;
59                    let val = r(end_byte);
60                    let b = val << (NUM_BITS - end_bit_offset);
61
62                    output[i] = a | (b & mask);
63                } else {
64                    let val = r(start_byte);
65                    output[i] = (val >> start_bit_offset) & mask;
66                }
67            });
68        }
69    };
70}
71
72/// Macro that generates unpack functions that accept num_bits as a parameter
73macro_rules! unpack {
74    ($name:ident, $t:ty, $bytes:literal, $bits:tt) => {
75        mod $name {
76            unpack_impl!($t, $bytes, $bits);
77        }
78
79        /// Unpack packed `input` into `output` with a bit width of `num_bits`
80        pub fn $name(input: &[u8], output: &mut [$t; $bits], num_bits: usize) {
81            // This will get optimised into a jump table
82            seq_macro::seq!(i in 0..=$bits {
83                if i == num_bits {
84                    return $name::unpack::<i>(input, output);
85                }
86            });
87            unreachable!("invalid num_bits {}", num_bits);
88        }
89    };
90}
91
92unpack!(unpack8, u8, 1, 8);
93unpack!(unpack16, u16, 2, 16);
94unpack!(unpack32, u32, 4, 32);
95unpack!(unpack64, u64, 8, 64);
96
97/// Macro that generates a pack function taking the number of bits as a const generic
98macro_rules! pack_impl {
99    ($t:ty, $bytes:literal, $bits:tt) => {
100        #[inline(never)]
101        pub fn pack<const NUM_BITS: usize>(input: &[$t; $bits], output: &mut [u8]) {
102            if NUM_BITS == 0 {
103                return;
104            }
105
106            assert!(NUM_BITS <= $bytes * 8);
107            assert!(output.len() >= NUM_BITS * $bytes);
108
109            let mask = match NUM_BITS {
110                $bits => <$t>::MAX,
111                _ => ((1 << NUM_BITS) - 1),
112            };
113
114            // Accumulate into locals so the packed words stay in registers. Only the
115            // first NUM_BITS entries are used, `[$t; NUM_BITS]` needs generic_const_exprs
116            let mut words = [0; $bits];
117
118            seq_macro::seq!(i in 0..$bits {
119                let value = input[i] & mask;
120
121                let start_bit = i * NUM_BITS;
122                let end_bit = start_bit + NUM_BITS;
123
124                let start_bit_offset = start_bit % $bits;
125                let end_bit_offset = end_bit % $bits;
126                let start_word = start_bit / $bits;
127                let end_word = end_bit / $bits;
128
129                words[start_word] |= value << start_bit_offset;
130                if start_word != end_word && end_bit_offset != 0 {
131                    words[end_word] |= value >> (NUM_BITS - end_bit_offset);
132                }
133            });
134
135            seq_macro::seq!(w in 0..$bits {
136                if w < NUM_BITS {
137                    output[w * $bytes..(w + 1) * $bytes].copy_from_slice(&words[w].to_le_bytes());
138                }
139            });
140        }
141
142        pub fn pack_blocks<const NUM_BITS: usize>(input: &[$t], output: &mut [u8]) {
143            if NUM_BITS == 0 {
144                return;
145            }
146            let block_bytes = NUM_BITS * $bytes;
147            let blocks = input.len() / $bits;
148            assert!(output.len() >= blocks * block_bytes);
149            for (input, output) in input
150                .chunks_exact($bits)
151                .zip(output.chunks_exact_mut(block_bytes))
152            {
153                pack::<NUM_BITS>(input.try_into().unwrap(), output);
154            }
155        }
156    };
157}
158
159/// Macro that generates pack functions that accept num_bits as a parameter
160macro_rules! pack {
161    ($name:ident, $blocks:ident, $t:ty, $bytes:literal, $bits:tt) => {
162        mod $name {
163            pack_impl!($t, $bytes, $bits);
164        }
165
166        /// Pack `input` into `output` with a bit width of `num_bits`
167        ///
168        /// Only the `num_bits` least significant bits of each value are written,
169        /// and `output` must contain at least `num_bits * size_of::<T>()` bytes
170        pub fn $name(input: &[$t; $bits], output: &mut [u8], num_bits: usize) {
171            // This will get optimised into a jump table
172            seq_macro::seq!(i in 0..=$bits {
173                if i == num_bits {
174                    return $name::pack::<i>(input, output);
175                }
176            });
177            unreachable!("invalid num_bits {}", num_bits);
178        }
179
180        #[inline(never)]
181        pub(crate) fn $blocks(input: &[$t], output: &mut [u8], num_bits: usize) {
182            seq_macro::seq!(i in 0..=$bits {
183                if i == num_bits {
184                    return $name::pack_blocks::<i>(input, output);
185                }
186            });
187            unreachable!("invalid num_bits {}", num_bits);
188        }
189    };
190}
191
192pack!(pack8, pack8_blocks, u8, 1, 8);
193pack!(pack16, pack16_blocks, u16, 2, 16);
194pack!(pack32, pack32_blocks, u32, 4, 32);
195pack!(pack64, pack64_blocks, u64, 8, 64);
196
197#[cfg(test)]
198mod tests {
199    use super::*;
200
201    #[test]
202    fn test_basic() {
203        let input = [0xFF; 4096];
204
205        for i in 0..=8 {
206            let mut output = [0; 8];
207            unpack8(&input, &mut output, i);
208            for (idx, out) in output.iter().enumerate() {
209                assert_eq!(out.trailing_ones() as usize, i, "out[{idx}] = {out}");
210            }
211        }
212
213        for i in 0..=16 {
214            let mut output = [0; 16];
215            unpack16(&input, &mut output, i);
216            for (idx, out) in output.iter().enumerate() {
217                assert_eq!(out.trailing_ones() as usize, i, "out[{idx}] = {out}");
218            }
219        }
220
221        for i in 0..=32 {
222            let mut output = [0; 32];
223            unpack32(&input, &mut output, i);
224            for (idx, out) in output.iter().enumerate() {
225                assert_eq!(out.trailing_ones() as usize, i, "out[{idx}] = {out}");
226            }
227        }
228
229        for i in 0..=64 {
230            let mut output = [0; 64];
231            unpack64(&input, &mut output, i);
232            for (idx, out) in output.iter().enumerate() {
233                assert_eq!(out.trailing_ones() as usize, i, "out[{idx}] = {out}");
234            }
235        }
236    }
237
238    #[test]
239    fn test_pack_all_ones() {
240        // Packing all-ones values must set every bit of the packed block and
241        // touch nothing beyond it
242        let mut output = [0u8; 4096];
243
244        for i in 0..=8 {
245            output.fill(0);
246            pack8(&[u8::MAX; 8], &mut output, i);
247            assert!(output[..i].iter().all(|&b| b == u8::MAX), "num_bits = {i}");
248            assert!(output[i..].iter().all(|&b| b == 0), "num_bits = {i}");
249        }
250
251        for i in 0..=16 {
252            output.fill(0);
253            pack16(&[u16::MAX; 16], &mut output, i);
254            assert!(
255                output[..2 * i].iter().all(|&b| b == u8::MAX),
256                "num_bits = {i}"
257            );
258            assert!(output[2 * i..].iter().all(|&b| b == 0), "num_bits = {i}");
259        }
260
261        for i in 0..=32 {
262            output.fill(0);
263            pack32(&[u32::MAX; 32], &mut output, i);
264            assert!(
265                output[..4 * i].iter().all(|&b| b == u8::MAX),
266                "num_bits = {i}"
267            );
268            assert!(output[4 * i..].iter().all(|&b| b == 0), "num_bits = {i}");
269        }
270
271        for i in 0..=64 {
272            output.fill(0);
273            pack64(&[u64::MAX; 64], &mut output, i);
274            assert!(
275                output[..8 * i].iter().all(|&b| b == u8::MAX),
276                "num_bits = {i}"
277            );
278            assert!(output[8 * i..].iter().all(|&b| b == 0), "num_bits = {i}");
279        }
280    }
281
282    #[test]
283    fn test_pack_round_trip() {
284        use crate::util::test_common::rand_gen::random_numbers;
285
286        // Values are deliberately not masked, pack must ignore the high bits
287        for i in 0..=8 {
288            let input: [u8; 8] = random_numbers(8).try_into().unwrap();
289            let mut packed = vec![0u8; i];
290            pack8(&input, &mut packed, i);
291            let mut output = [0; 8];
292            unpack8(&packed, &mut output, i);
293            let mask = ((1u16 << i) - 1) as u8;
294            for (idx, (&v, &out)) in input.iter().zip(output.iter()).enumerate() {
295                assert_eq!(v & mask, out, "num_bits = {i}, index = {idx}");
296            }
297        }
298
299        for i in 0..=16 {
300            let input: [u16; 16] = random_numbers(16).try_into().unwrap();
301            let mut packed = vec![0u8; 2 * i];
302            pack16(&input, &mut packed, i);
303            let mut output = [0; 16];
304            unpack16(&packed, &mut output, i);
305            let mask = ((1u32 << i) - 1) as u16;
306            for (idx, (&v, &out)) in input.iter().zip(output.iter()).enumerate() {
307                assert_eq!(v & mask, out, "num_bits = {i}, index = {idx}");
308            }
309        }
310
311        for i in 0..=32 {
312            let input: [u32; 32] = random_numbers(32).try_into().unwrap();
313            let mut packed = vec![0u8; 4 * i];
314            pack32(&input, &mut packed, i);
315            let mut output = [0; 32];
316            unpack32(&packed, &mut output, i);
317            let mask = ((1u64 << i) - 1) as u32;
318            for (idx, (&v, &out)) in input.iter().zip(output.iter()).enumerate() {
319                assert_eq!(v & mask, out, "num_bits = {i}, index = {idx}");
320            }
321        }
322
323        for i in 0..=64 {
324            let input: [u64; 64] = random_numbers(64).try_into().unwrap();
325            let mut packed = vec![0u8; 8 * i];
326            pack64(&input, &mut packed, i);
327            let mut output = [0; 64];
328            unpack64(&packed, &mut output, i);
329            let mask = ((1u128 << i) - 1) as u64;
330            for (idx, (&v, &out)) in input.iter().zip(output.iter()).enumerate() {
331                assert_eq!(v & mask, out, "num_bits = {i}, index = {idx}");
332            }
333        }
334    }
335}