use core::cmp::Ordering;
use core::{fmt::Debug, hint};
use cfg_if::cfg_if;
use crate::int::KEYS_BYTES;
cfg_if! {
if #[cfg(all(
any(target_arch = "x86", target_arch = "x86_64"),
target_feature = "avx512bw",
target_feature = "popcnt",
))] {
mod avx512;
} else if #[cfg(all(
any(target_arch = "x86", target_arch = "x86_64"),
target_feature = "avx2",
target_feature = "popcnt",
))] {
mod avx2;
} else if #[cfg(all(
any(target_arch = "x86", target_arch = "x86_64"),
target_feature = "sse2",
target_feature = "popcnt",
))] {
mod sse2_popcnt;
} else if #[cfg(all(
any(target_arch = "x86", target_arch = "x86_64"),
target_feature = "sse2",
))] {
mod sse2;
} else if #[cfg(all(
target_arch = "aarch64",
target_feature = "sve",
))] {
mod sve;
} else if #[cfg(all(
target_arch = "aarch64",
target_feature = "neon",
))] {
mod neon;
} else if #[cfg(all(
any(target_arch = "riscv32", target_arch = "riscv64"),
target_feature = "v",
))] {
mod rvv;
} else {
impl SimdSearch for u8 {}
impl SimdSearch for u16 {}
impl SimdSearch for u32 {}
impl SimdSearch for u64 {}
impl SimdSearch for u128 {}
impl SimdSearch for i8 {}
impl SimdSearch for i16 {}
impl SimdSearch for i32 {}
impl SimdSearch for i64 {}
impl SimdSearch for i128 {}
}
}
pub(crate) trait Int: Ord + Copy + Debug {
const ZERO: Self;
fn wrapping_add(self, other: Self) -> Self;
}
macro_rules! impl_zero {
($($int:ident,)*) => {
$(
impl Int for $int {
const ZERO: Self = 0;
#[inline]
fn wrapping_add(self, other: Self) -> Self {
self.wrapping_add(other)
}
}
)*
};
}
impl_zero! {
u8,
u16,
u32,
u64,
u128,
i8,
i16,
i32,
i64,
i128,
}
pub(crate) trait SimdSearch: Int {
const SIMD_WIDTH: usize = 1;
const BIAS: Self = Self::ZERO;
#[inline]
fn bias_cmp(a: Self, b: Self) -> Ordering {
Ord::cmp(&a.wrapping_add(Self::BIAS), &b.wrapping_add(Self::BIAS))
}
#[inline]
unsafe fn search(keys: &[Self], search: Self) -> usize {
debug_assert!(keys.len() >= 2);
debug_assert!(keys.len() >= Self::SIMD_WIDTH);
debug_assert!(keys.len().is_power_of_two());
debug_assert_eq!(keys.as_ptr().addr() % KEYS_BYTES, 0);
let mut len = keys.len();
let mut base = 0;
while len > Self::SIMD_WIDTH {
let mid = base + len / 2;
let key = unsafe { *keys.get_unchecked(mid - 1) };
base = hint::select_unpredictable(Self::bias_cmp(search, key).is_gt(), mid, base);
len /= 2;
}
debug_assert_eq!(len, Self::SIMD_WIDTH);
debug_assert_eq!(base % Self::SIMD_WIDTH, 0);
base + unsafe { Self::simd_search(keys.as_ptr().add(base), search) }
}
#[inline]
unsafe fn simd_search(keys: *const Self, search: Self) -> usize {
assert_eq!(Self::SIMD_WIDTH, 1);
debug_assert!(Self::bias_cmp(search, unsafe { keys.read() }).is_le());
0
}
}
#[inline]
#[allow(dead_code)]
unsafe fn exact_div_unchecked(a: usize, b: usize) -> usize {
unsafe {
hint::assert_unchecked(a.is_multiple_of(b));
a / core::num::NonZero::new_unchecked(b)
}
}
#[cfg(test)]
mod tests {
use super::SimdSearch;
use crate::int::{AlignedKeys, KEYS_BYTES};
fn generic_search<T: SimdSearch>(keys: &[T], search: T) -> usize {
keys[..keys.len() - 1].partition_point(|&key| T::bias_cmp(key, search).is_lt())
}
fn test_search<T: SimdSearch>(encode: impl Fn(usize) -> T, max: T) {
let len = KEYS_BYTES / std::mem::size_of::<T>();
let mut keys: AlignedKeys<[T; KEYS_BYTES]> = unsafe { std::mem::zeroed() };
for i in 0..len {
keys.0[i] = encode(i & !1);
}
keys.0[len - 1] = max;
for i in 0..len {
assert_eq!(generic_search(&keys.0[..len], encode(i)), unsafe {
T::search(&keys.0[..len], encode(i))
});
}
}
#[test]
fn test_search_u8() {
test_search(
|i| (i as u8).wrapping_add(SimdSearch::BIAS),
u8::MAX.wrapping_add(SimdSearch::BIAS),
);
}
#[test]
fn test_search_u16() {
test_search(
|i| (i as u16).wrapping_add(SimdSearch::BIAS),
u16::MAX.wrapping_add(SimdSearch::BIAS),
);
}
#[test]
fn test_search_u32() {
test_search(
|i| (i as u32).wrapping_add(SimdSearch::BIAS),
u32::MAX.wrapping_add(SimdSearch::BIAS),
);
}
#[test]
fn test_search_u64() {
test_search(
|i| (i as u64).wrapping_add(SimdSearch::BIAS),
u64::MAX.wrapping_add(SimdSearch::BIAS),
);
}
#[test]
fn test_search_u128() {
test_search(
|i| (i as u128).wrapping_add(SimdSearch::BIAS),
u128::MAX.wrapping_add(SimdSearch::BIAS),
);
}
#[test]
fn test_search_i8() {
test_search(
|i| (i as i8).wrapping_add(SimdSearch::BIAS),
i8::MAX.wrapping_add(SimdSearch::BIAS),
);
}
#[test]
fn test_search_i16() {
test_search(
|i| (i as i16).wrapping_add(SimdSearch::BIAS),
i16::MAX.wrapping_add(SimdSearch::BIAS),
);
}
#[test]
fn test_search_i32() {
test_search(
|i| (i as i32).wrapping_add(SimdSearch::BIAS),
i32::MAX.wrapping_add(SimdSearch::BIAS),
);
}
#[test]
fn test_search_i64() {
test_search(
|i| (i as i64).wrapping_add(SimdSearch::BIAS),
i64::MAX.wrapping_add(SimdSearch::BIAS),
);
}
#[test]
fn test_search_i128() {
test_search(
|i| (i as i128).wrapping_add(SimdSearch::BIAS),
i128::MAX.wrapping_add(SimdSearch::BIAS),
);
}
}
#[cfg(feature = "internal_benches")]
mod bench {
use super::SimdSearch;
use crate::int::{AlignedKeys, KEYS_BYTES};
#[divan::bench(types = [
u8,
u16,
u32,
u64,
u128,
i8,
i16,
i32,
i64,
i128,
])]
fn search<T: SimdSearch>(bencher: divan::Bencher) {
let keys: AlignedKeys<[T; KEYS_BYTES]> = unsafe { std::mem::zeroed() };
bencher.bench_local(|| {
let zero: T = unsafe { std::mem::zeroed() };
let len = KEYS_BYTES / std::mem::size_of::<T>();
unsafe { T::search(&keys.0[..len], divan::black_box(zero)) }
});
}
#[divan::bench(types = [
u8,
u16,
u32,
u64,
u128,
i8,
i16,
i32,
i64,
i128,
])]
fn generic_search<T: SimdSearch>(bencher: divan::Bencher) {
fn generic_search<T: SimdSearch>(keys: &[T], search: T) -> usize {
keys[..keys.len() - 1].partition_point(|&key| T::bias_cmp(key, search).is_lt())
}
let keys: AlignedKeys<[T; KEYS_BYTES]> = unsafe { std::mem::zeroed() };
bencher.bench_local(|| {
let zero: T = unsafe { std::mem::zeroed() };
let len = KEYS_BYTES / std::mem::size_of::<T>();
generic_search(&keys.0[..len], divan::black_box(zero))
});
}
}