1#[inline]
33pub fn collect_bool_word_scalar<F>(len: usize, mut f: F) -> u64
34where
35 F: FnMut(usize) -> bool,
36{
37 assert!(len <= 64, "cannot pack {len} bits into a u64 word");
38
39 let mut packed = 0;
40 for bit_idx in 0..len {
41 packed |= (f(bit_idx) as u64) << bit_idx;
42 }
43 packed
44}
45
46#[allow(clippy::inline_always)]
60#[inline(always)]
61pub(crate) fn collect_bool_words_inline<F>(words: &mut [u64], len: usize, f: F)
62where
63 F: FnMut(usize) -> bool,
64{
65 #[cfg(all(
66 target_arch = "x86_64",
67 target_feature = "avx512f",
68 target_feature = "avx512bw",
69 not(miri)
70 ))]
71 {
72 collect_bool_words_with(words, len, f, |bools| unsafe {
76 pack_bool_word_avx512(bools)
77 })
78 }
79 #[cfg(all(
80 target_arch = "x86_64",
81 target_feature = "avx2",
82 not(all(target_feature = "avx512f", target_feature = "avx512bw")),
83 not(miri)
84 ))]
85 {
86 collect_bool_words_with(words, len, f, |bools| unsafe { pack_bool_word_avx2(bools) })
88 }
89 #[cfg(all(target_arch = "x86_64", not(target_feature = "avx2"), not(miri)))]
90 {
91 collect_bool_words_with(words, len, f, |bools| unsafe { pack_bool_word_sse2(bools) })
93 }
94 #[cfg(all(target_arch = "aarch64", not(miri)))]
95 {
96 collect_bool_words_with(words, len, f, |bools| unsafe { pack_bool_word_neon(bools) })
98 }
99 #[cfg(any(not(any(target_arch = "x86_64", target_arch = "aarch64")), miri))]
100 collect_bool_words_with(words, len, f, pack_bool_word_swar)
101}
102
103#[inline]
118pub fn collect_bool_words_multiversioned<F>(words: &mut [u64], len: usize, f: F)
119where
120 F: FnMut(usize) -> bool,
121{
122 let num_words = len.div_ceil(64);
123 assert!(
124 words.len() >= num_words,
125 "words slice has {} entries, need at least {num_words}",
126 words.len(),
127 );
128
129 if len < 64 {
132 return collect_bool_words_inline(words, len, f);
133 }
134
135 #[cfg(all(target_arch = "x86_64", not(miri)))]
136 {
137 if is_x86_feature_detected!("avx512f") && is_x86_feature_detected!("avx512bw") {
138 return unsafe { collect_bool_words_avx512(words, len, f) };
140 }
141 if is_x86_feature_detected!("avx2") {
142 return unsafe { collect_bool_words_avx2(words, len, f) };
144 }
145 }
146 collect_bool_words_inline(words, len, f)
147}
148
149#[allow(clippy::inline_always)]
155#[inline(always)]
156fn collect_bool_words_with<F, P>(words: &mut [u64], len: usize, mut f: F, pack: P)
157where
158 F: FnMut(usize) -> bool,
159 P: Fn(&[bool; 64]) -> u64,
160{
161 let full = len / 64;
162 let remainder = len % 64;
163
164 for (word_idx, word) in words[..full].iter_mut().enumerate() {
165 let offset = word_idx * 64;
166 let mut bools = [false; 64];
167 for (bit_idx, b) in bools.iter_mut().enumerate() {
168 *b = f(offset + bit_idx);
169 }
170 *word = pack(&bools);
171 }
172
173 if remainder != 0 {
174 let offset = full * 64;
175 words[full] = collect_bool_word_scalar(remainder, |bit_idx| f(offset + bit_idx));
176 }
177}
178
179#[cfg(target_arch = "x86_64")]
187#[target_feature(enable = "sse2")]
188pub unsafe fn collect_bool_words_sse2<F: FnMut(usize) -> bool>(
189 words: &mut [u64],
190 len: usize,
191 f: F,
192) {
193 collect_bool_words_with(words, len, f, |bools| unsafe { pack_bool_word_sse2(bools) })
195}
196
197#[cfg(target_arch = "x86_64")]
205#[target_feature(enable = "avx2")]
206pub unsafe fn collect_bool_words_avx2<F: FnMut(usize) -> bool>(
207 words: &mut [u64],
208 len: usize,
209 f: F,
210) {
211 collect_bool_words_with(words, len, f, |bools| unsafe { pack_bool_word_avx2(bools) })
213}
214
215#[cfg(target_arch = "x86_64")]
223#[target_feature(enable = "avx512f,avx512bw")]
224pub unsafe fn collect_bool_words_avx512<F: FnMut(usize) -> bool>(
225 words: &mut [u64],
226 len: usize,
227 f: F,
228) {
229 collect_bool_words_with(words, len, f, |bools| unsafe {
231 pack_bool_word_avx512(bools)
232 })
233}
234
235#[cfg(target_arch = "aarch64")]
243#[target_feature(enable = "neon")]
244pub unsafe fn collect_bool_words_neon<F: FnMut(usize) -> bool>(
245 words: &mut [u64],
246 len: usize,
247 f: F,
248) {
249 collect_bool_words_with(words, len, f, |bools| unsafe { pack_bool_word_neon(bools) })
251}
252
253#[inline]
260pub fn pack_bool_word_swar(bools: &[bool; 64]) -> u64 {
261 const MAGIC: u64 = 0x0102_0408_1020_4080;
262
263 let (chunks, rest) = bools.as_chunks::<8>();
264 debug_assert!(rest.is_empty());
265
266 let mut packed = 0u64;
267 for (chunk_idx, chunk) in chunks.iter().enumerate() {
268 let word = u64::from_le_bytes(chunk.map(|b| b as u8));
269 packed |= (word.wrapping_mul(MAGIC) >> 56) << (8 * chunk_idx);
270 }
271 packed
272}
273
274#[cfg(target_arch = "x86_64")]
280#[inline]
281#[target_feature(enable = "sse2")]
282pub unsafe fn pack_bool_word_sse2(bools: &[bool; 64]) -> u64 {
283 use std::arch::x86_64::__m128i;
284 use std::arch::x86_64::_mm_cmpeq_epi8;
285 use std::arch::x86_64::_mm_loadu_si128;
286 use std::arch::x86_64::_mm_movemask_epi8;
287 use std::arch::x86_64::_mm_setzero_si128;
288
289 let ptr = bools.as_ptr().cast::<u8>();
290 let zero = _mm_setzero_si128();
291
292 let mut packed = 0u64;
293 for lane in 0..4 {
294 let chunk = unsafe { _mm_loadu_si128(ptr.add(lane * 16).cast::<__m128i>()) };
296 let zero_mask = _mm_movemask_epi8(_mm_cmpeq_epi8(chunk, zero)) as u32 as u64;
298 packed |= (!zero_mask & 0xFFFF) << (16 * lane);
299 }
300 packed
301}
302
303#[cfg(target_arch = "x86_64")]
309#[inline]
310#[target_feature(enable = "avx2")]
311pub unsafe fn pack_bool_word_avx2(bools: &[bool; 64]) -> u64 {
312 use std::arch::x86_64::__m256i;
313 use std::arch::x86_64::_mm256_cmpeq_epi8;
314 use std::arch::x86_64::_mm256_loadu_si256;
315 use std::arch::x86_64::_mm256_movemask_epi8;
316 use std::arch::x86_64::_mm256_setzero_si256;
317
318 let ptr = bools.as_ptr().cast::<u8>();
319 let zero = _mm256_setzero_si256();
320
321 let lo = unsafe { _mm256_loadu_si256(ptr.cast::<__m256i>()) };
323 let hi = unsafe { _mm256_loadu_si256(ptr.add(32).cast::<__m256i>()) };
325
326 let lo_mask = !(_mm256_movemask_epi8(_mm256_cmpeq_epi8(lo, zero)) as u32) as u64;
328 let hi_mask = !(_mm256_movemask_epi8(_mm256_cmpeq_epi8(hi, zero)) as u32) as u64;
329 lo_mask | (hi_mask << 32)
330}
331
332#[cfg(target_arch = "x86_64")]
338#[inline]
339#[target_feature(enable = "avx512f,avx512bw")]
340pub unsafe fn pack_bool_word_avx512(bools: &[bool; 64]) -> u64 {
341 use std::arch::x86_64::__m512i;
342 use std::arch::x86_64::_mm512_loadu_si512;
343 use std::arch::x86_64::_mm512_test_epi8_mask;
344
345 let chunk = unsafe { _mm512_loadu_si512(bools.as_ptr().cast::<__m512i>()) };
347 _mm512_test_epi8_mask(chunk, chunk)
349}
350
351#[cfg(target_arch = "aarch64")]
358#[inline]
359#[target_feature(enable = "neon")]
360pub unsafe fn pack_bool_word_neon(bools: &[bool; 64]) -> u64 {
361 use std::arch::aarch64::vgetq_lane_u64;
362 use std::arch::aarch64::vld1q_s8;
363 use std::arch::aarch64::vld1q_u8;
364 use std::arch::aarch64::vpaddq_u8;
365 use std::arch::aarch64::vreinterpretq_u64_u8;
366 use std::arch::aarch64::vshlq_u8;
367
368 const BIT_SHIFTS: [i8; 16] = [0, 1, 2, 3, 4, 5, 6, 7, 0, 1, 2, 3, 4, 5, 6, 7];
369
370 let ptr = bools.as_ptr().cast::<u8>();
371 unsafe {
374 let shifts = vld1q_s8(BIT_SHIFTS.as_ptr());
375
376 let m0 = vshlq_u8(vld1q_u8(ptr), shifts);
378 let m1 = vshlq_u8(vld1q_u8(ptr.add(16)), shifts);
379 let m2 = vshlq_u8(vld1q_u8(ptr.add(32)), shifts);
380 let m3 = vshlq_u8(vld1q_u8(ptr.add(48)), shifts);
381
382 let sum01 = vpaddq_u8(m0, m1);
385 let sum23 = vpaddq_u8(m2, m3);
386 let sum = vpaddq_u8(sum01, sum23);
387 let sum = vpaddq_u8(sum, sum);
388 vgetq_lane_u64::<0>(vreinterpretq_u64_u8(sum))
389 }
390}
391
392#[cfg(test)]
393mod tests {
394 use rstest::rstest;
395
396 use super::collect_bool_word_scalar;
397 use super::pack_bool_word_swar;
398
399 fn patterns() -> Vec<[bool; 64]> {
400 let mut patterns = vec![
401 [false; 64],
402 [true; 64],
403 std::array::from_fn(|i| i % 2 == 0),
404 std::array::from_fn(|i| i % 3 == 0),
405 std::array::from_fn(|i| i < 32),
406 std::array::from_fn(|i| i == 0 || i == 63),
407 ];
408 let mut state = 0x9E37_79B9_7F4A_7C15u64;
410 for _ in 0..8 {
411 patterns.push(std::array::from_fn(|_| {
412 state = state
413 .wrapping_mul(6364136223846793005)
414 .wrapping_add(1442695040888963407);
415 (state >> 33) & 1 == 1
416 }));
417 }
418 patterns
419 }
420
421 fn reference(bools: &[bool; 64]) -> u64 {
422 collect_bool_word_scalar(64, |i| bools[i])
423 }
424
425 #[test]
426 fn swar_matches_scalar() {
427 for bools in patterns() {
428 assert_eq!(pack_bool_word_swar(&bools), reference(&bools));
429 }
430 }
431
432 #[test]
433 fn dispatch_matches_scalar() {
434 for bools in patterns() {
435 assert_eq!(
436 crate::bit::collect_bool_word(64, |i| bools[i]),
437 reference(&bools)
438 );
439 }
440 }
441
442 #[cfg(all(target_arch = "x86_64", not(miri)))]
443 #[test]
444 fn sse2_matches_scalar() {
445 for bools in patterns() {
446 assert_eq!(
448 unsafe { super::pack_bool_word_sse2(&bools) },
449 reference(&bools)
450 );
451 }
452 }
453
454 #[cfg(all(target_arch = "x86_64", not(miri)))]
455 #[test]
456 fn avx2_matches_scalar() {
457 if !is_x86_feature_detected!("avx2") {
458 return;
459 }
460 for bools in patterns() {
461 assert_eq!(
463 unsafe { super::pack_bool_word_avx2(&bools) },
464 reference(&bools)
465 );
466 }
467 }
468
469 #[cfg(all(target_arch = "x86_64", not(miri)))]
470 #[test]
471 fn avx512_matches_scalar() {
472 if !(is_x86_feature_detected!("avx512f") && is_x86_feature_detected!("avx512bw")) {
473 return;
474 }
475 for bools in patterns() {
476 assert_eq!(
478 unsafe { super::pack_bool_word_avx512(&bools) },
479 reference(&bools)
480 );
481 }
482 }
483
484 #[cfg(all(target_arch = "aarch64", not(miri)))]
485 #[test]
486 fn neon_matches_scalar() {
487 for bools in patterns() {
488 assert_eq!(
490 unsafe { super::pack_bool_word_neon(&bools) },
491 reference(&bools)
492 );
493 }
494 }
495
496 #[rstest]
497 #[case(0)]
498 #[case(1)]
499 #[case(63)]
500 #[case(64)]
501 #[case(65)]
502 #[case(200)]
503 fn multiversioned_matches_inline(#[case] len: usize) {
504 let pattern = |i: usize| i.is_multiple_of(3) || i.is_multiple_of(7);
505 let num_words = len.div_ceil(64);
506 let mut multiversioned = vec![0u64; num_words];
507 super::collect_bool_words_multiversioned(&mut multiversioned, len, pattern);
508 let mut inline = vec![0u64; num_words];
509 super::collect_bool_words_inline(&mut inline, len, pattern);
510 assert_eq!(multiversioned, inline);
511 }
512
513 #[rstest]
514 #[case(0)]
515 #[case(1)]
516 #[case(5)]
517 #[case(63)]
518 #[case(64)]
519 fn collect_bool_word_partial_lens_match(#[case] len: usize) {
520 let expected = collect_bool_word_scalar(len, |i| i % 3 == 0);
521 assert_eq!(crate::bit::collect_bool_word(len, |i| i % 3 == 0), expected);
522 }
523}