1macro_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
72macro_rules! unpack {
74 ($name:ident, $t:ty, $bytes:literal, $bits:tt) => {
75 mod $name {
76 unpack_impl!($t, $bytes, $bits);
77 }
78
79 pub fn $name(input: &[u8], output: &mut [$t; $bits], num_bits: usize) {
81 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
97macro_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 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
159macro_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 pub fn $name(input: &[$t; $bits], output: &mut [u8], num_bits: usize) {
171 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 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 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}