Skip to main content

stream_vbyte/decode/
ssse3.rs

1use std::cmp;
2
3use std::arch::x86_64::{__m128i, _mm_loadu_si128, _mm_shuffle_epi8, _mm_storeu_si128};
4
5use super::{DecodeQuadSink, Decoder, WriteQuadToSlice};
6use crate::tables;
7
8/// Decoder using SSSE3 instructions.
9pub struct Ssse3;
10
11impl Decoder for Ssse3 {
12    type DecodedQuad = __m128i;
13
14    fn decode_quads<S: DecodeQuadSink<Self>>(
15        control_bytes: &[u8],
16        encoded_nums: &[u8],
17        control_bytes_to_decode: usize,
18        nums_already_decoded: usize,
19        sink: &mut S,
20    ) -> (usize, usize) {
21        let mut bytes_read: usize = 0;
22        let mut nums_decoded: usize = nums_already_decoded;
23
24        // Decoding reads 16 bytes at a time from input, so we won't be able to read the last few
25        // control byte's worth because they may be encoded at 1 byte per number, so we need 3
26        // additional control bytes' worth of numbers to provide the extra 12 bytes.
27        // However, if control_bytes_to_decode is short enough, we can decode all the requested
28        // numbers because we'll have un-processed input to ensure we can read 16 bytes.
29        let control_byte_limit = cmp::min(
30            control_bytes_to_decode,
31            control_bytes.len().saturating_sub(3),
32        );
33
34        // need to ensure that we can copy 16 encoded bytes, so last few quads will be handled
35        // by a slower loop
36        for &control_byte in control_bytes[0..control_byte_limit].iter() {
37            let length = tables::DECODE_LENGTH_PER_QUAD_TABLE[control_byte as usize];
38            let mask_bytes = tables::X86_SSSE3_DECODE_SHUFFLE_TABLE[control_byte as usize];
39            // we'll read 16 bytes from this always, so using explicit slice size to make sure it's
40            // ok to read unsafe
41            let next_4 = &encoded_nums[bytes_read..(bytes_read + 16)];
42
43            let mask;
44            let data;
45            unsafe {
46                // TODO load mask unaligned once https://github.com/rust-lang/rust/issues/33626
47                // hits stable
48                mask = _mm_loadu_si128(mask_bytes.as_ptr() as *const __m128i);
49                data = _mm_loadu_si128(next_4.as_ptr() as *const __m128i);
50            }
51
52            let decompressed = unsafe { _mm_shuffle_epi8(data, mask) };
53
54            sink.on_quad(decompressed, nums_decoded);
55
56            bytes_read += length as usize;
57            nums_decoded += 4;
58        }
59
60        (nums_decoded - nums_already_decoded, bytes_read)
61    }
62}
63
64impl WriteQuadToSlice for Ssse3 {
65    #[inline]
66    fn write_quad_to_slice(quad: Self::DecodedQuad, slice: &mut [u32]) {
67        unsafe { _mm_storeu_si128(slice.as_ptr() as *mut __m128i, quad) }
68    }
69}
70
71#[cfg(test)]
72mod tests {
73    use super::*;
74    use crate::{cumulative_encoded_len, decode::SliceDecodeSink, encode::encode, scalar::Scalar};
75
76    #[test]
77    fn reads_all_requested_control_bytes_when_12_extra_input_bytes() {
78        let nums: Vec<u32> = (0..64).map(|i| i * 100).collect();
79        let mut encoded = Vec::new();
80        let mut decoded: Vec<u32> = Vec::new();
81        encoded.resize(nums.len() * 5, 0xFF);
82
83        encode::<Scalar>(&nums, &mut encoded);
84
85        // 16 control bytes
86        let control_bytes = &encoded[0..16];
87        let encoded_nums = &encoded[16..];
88
89        for control_bytes_to_decode in 0..14 {
90            decoded.clear();
91            decoded.resize(nums.len(), 54321);
92
93            // requesting 13 or fewer control bytes decodes all requested bytes
94            let (nums_decoded, bytes_read) = Ssse3::decode_quads(
95                &control_bytes,
96                &encoded_nums,
97                control_bytes_to_decode,
98                0,
99                &mut SliceDecodeSink::new(&mut decoded),
100            );
101            assert_eq!(control_bytes_to_decode * 4, nums_decoded);
102            assert_eq!(
103                cumulative_encoded_len(&control_bytes[0..control_bytes_to_decode]),
104                bytes_read
105            );
106            assert_eq!(&nums[0..nums_decoded], &decoded[0..nums_decoded]);
107            assert!(&decoded[nums_decoded..].iter().all(|&i| i == 54321_u32));
108        }
109
110        for control_bytes_to_decode in 14..17 {
111            decoded.clear();
112            decoded.resize(nums.len(), 54321);
113
114            // requesting more than 13 gets capped to 13 because there may not be enough encoded
115            // nums to read 16 bytes at a time
116            let (nums_decoded, bytes_read) = Ssse3::decode_quads(
117                &control_bytes,
118                &encoded_nums,
119                control_bytes_to_decode,
120                0,
121                &mut SliceDecodeSink::new(&mut decoded),
122            );
123            assert_eq!(13 * 4, nums_decoded);
124            assert_eq!(
125                cumulative_encoded_len(&control_bytes[0..(nums_decoded / 4)]),
126                bytes_read
127            );
128            assert_eq!(&nums[0..nums_decoded], &decoded[0..nums_decoded]);
129            assert!(&decoded[nums_decoded..].iter().all(|&i| i == 54321_u32));
130        }
131    }
132}