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}