Skip to main content

rudb_vector/
bytes.rs

1//! Where some bytes are in a block of sixty four, one bit a byte.
2//!
3//! This is the step a structural scan starts from, the one simdjson and simdcsv take: compare a
4//! block against a few bytes of interest and get a `u64` per byte, first byte in the lowest bit.
5//! The compares vectorize by themselves. Turning sixteen compare results into sixteen bits does
6//! not, because the portable way to gather them is a multiply per eight bytes and the compiler does
7//! not know that `pmovmskb` does it in one instruction. On a `lineitem` load from CSV the gathering
8//! was more than half of the splitter's time.
9//!
10//! So on `x86_64` the gather is SSE2's `movemask`, which every `x86_64` processor has, and there is
11//! nothing to detect at run time. Everywhere else it is the portable multiply, which is also what
12//! the tests hold the SSE2 version to.
13
14/// One mask per byte in `needles`, with bit `i` of mask `n` set when `block[i] == needles[n]`.
15#[must_use]
16#[inline]
17pub fn masks<const N: usize>(block: &[u8; 64], needles: [u8; N]) -> [u64; N] {
18    #[cfg(target_arch = "x86_64")]
19    {
20        sse2(block, needles)
21    }
22    #[cfg(not(target_arch = "x86_64"))]
23    {
24        portable(block, needles)
25    }
26}
27
28/// [`masks`] with SSE2: four loads of sixteen bytes, then a compare and a `movemask` per needle and
29/// load.
30#[cfg(target_arch = "x86_64")]
31#[inline]
32#[allow(unsafe_code)]
33fn sse2<const N: usize>(block: &[u8; 64], needles: [u8; N]) -> [u64; N] {
34    use std::arch::x86_64::{
35        __m128i, _mm_cmpeq_epi8, _mm_loadu_si128, _mm_movemask_epi8, _mm_set1_epi8,
36    };
37    // SAFETY: SSE2 is part of the `x86_64` baseline, so every processor this code can run on has
38    // these instructions. Each load reads sixteen bytes starting at offset 0, 16, 32 or 48 of a
39    // sixty four byte array, all inside it, and `loadu` has no alignment requirement.
40    unsafe {
41        let base = block.as_ptr().cast::<__m128i>();
42        let lanes = [
43            _mm_loadu_si128(base),
44            _mm_loadu_si128(base.add(1)),
45            _mm_loadu_si128(base.add(2)),
46            _mm_loadu_si128(base.add(3)),
47        ];
48        needles.map(|needle| {
49            // The needle is a byte and `set1` takes the same bits as an `i8`.
50            #[allow(clippy::cast_possible_wrap)]
51            let splat = _mm_set1_epi8(needle as i8);
52            let mut mask = 0u64;
53            for (at, lane) in lanes.iter().enumerate() {
54                // `movemask` of bytes sets only the low sixteen bits, so the cast keeps them all.
55                #[allow(clippy::cast_sign_loss)]
56                let bits = _mm_movemask_epi8(_mm_cmpeq_epi8(*lane, splat)) as u32 as u64;
57                mask |= bits << (16 * at);
58            }
59            mask
60        })
61    }
62}
63
64/// [`masks`] without a platform's instructions. A compare per byte, which becomes vector compares,
65/// and then eight of those gathered into eight bits with one multiply.
66#[cfg_attr(target_arch = "x86_64", allow(dead_code))]
67#[inline]
68fn portable<const N: usize>(block: &[u8; 64], needles: [u8; N]) -> [u64; N] {
69    needles.map(|needle| {
70        let mut hits = [0u8; 64];
71        for (hit, &byte) in hits.iter_mut().zip(block) {
72            *hit = u8::from(byte == needle);
73        }
74        let mut mask = 0u64;
75        for (at, eight) in hits.chunks_exact(8).enumerate() {
76            let word = u64::from_le_bytes(eight.try_into().expect("eight bytes"));
77            // The multiply puts a copy of byte `i` at bit `56 + i` for every `i` at once, and every
78            // other copy it makes lands on a bit of its own below bit 56, so nothing carries into
79            // the top byte.
80            mask |= (word.wrapping_mul(0x0102_0408_1020_4080) >> 56) << (8 * at);
81        }
82        mask
83    })
84}
85
86#[cfg(test)]
87mod tests {
88    use super::*;
89
90    fn by_hand(block: &[u8; 64], needle: u8) -> u64 {
91        let mut mask = 0;
92        for (at, &byte) in block.iter().enumerate() {
93            if byte == needle {
94                mask |= 1 << at;
95            }
96        }
97        mask
98    }
99
100    #[test]
101    fn every_mask_has_a_bit_for_exactly_the_bytes_that_match() {
102        let mut seed = 0x9e37_79b9_7f4a_7c15u64;
103        for round in 0..2000 {
104            let mut block = [0u8; 64];
105            for byte in &mut block {
106                seed ^= seed << 13;
107                seed ^= seed >> 7;
108                seed ^= seed << 17;
109                // A small alphabet so that every needle turns up, with the high bit set half the
110                // time so that a byte that is negative as an `i8` is covered.
111                *byte = [b',', b'"', b'\n', b'\r', b'a', 0x80, 0xff, 0][(seed % 8) as usize];
112            }
113            let needles = [b',', b'"', b'\n', b'\r', 0xff, 0, b'z'];
114            let want = needles.map(|needle| by_hand(&block, needle));
115            assert_eq!(masks(&block, needles), want, "round {round}");
116            assert_eq!(portable(&block, needles), want, "round {round}");
117        }
118    }
119}