#![allow(unsafe_code)]
#[must_use]
#[inline]
pub fn next_match_from(ends: &[i32], from: usize, limit: usize) -> Option<usize> {
let limit = limit.min(ends.len());
let probe = limit.min(from.saturating_add(PROBE));
if let Some(m) = next_match_from_scalar(ends, from, probe) {
return Some(m);
}
next_match_from_unprobed(ends, probe, limit)
}
#[must_use]
#[inline]
pub fn next_match_from_unprobed(ends: &[i32], from: usize, limit: usize) -> Option<usize> {
let limit = limit.min(ends.len());
#[cfg(target_arch = "x86_64")]
{
match crate::isa::tier() {
crate::isa::Tier::Avx512 => {
return unsafe { next_match_from_avx512(ends, from, limit) };
}
crate::isa::Tier::Avx2 => return unsafe { next_match_from_avx2(ends, from, limit) },
crate::isa::Tier::Sse2 | crate::isa::Tier::Scalar => {}
}
}
next_match_from_scalar(ends, from, limit)
}
const PROBE: usize = 4;
const WIDE: usize = 8;
#[must_use]
pub fn takes_vector(ends: &[i32]) -> bool {
const SAMPLE: usize = 4096;
let n = ends.len().min(SAMPLE);
let mut a = 0usize;
let mut matches = 0usize;
while let Some(m) = next_match_from_scalar(ends, a, n) {
a = ends[m] as usize;
matches += 1;
}
matches == 0 || n / matches >= WIDE
}
#[must_use]
#[inline]
pub fn next_match_from_scalar(ends: &[i32], from: usize, limit: usize) -> Option<usize> {
let limit = limit.min(ends.len());
(from..limit).find(|&a| ends[a] > a as i32)
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx512f")]
unsafe fn next_match_from_avx512(ends: &[i32], from: usize, limit: usize) -> Option<usize> {
use core::arch::x86_64::{
_mm512_add_epi32, _mm512_cmpgt_epi32_mask, _mm512_loadu_si512, _mm512_set1_epi32,
_mm512_setr_epi32,
};
let lanes = _mm512_setr_epi32(0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15);
let mut a = from;
while a + 16 <= limit {
let v = unsafe { _mm512_loadu_si512(ends.as_ptr().add(a).cast()) };
let idx = _mm512_add_epi32(lanes, _mm512_set1_epi32(a as i32));
let mask: u16 = _mm512_cmpgt_epi32_mask(v, idx);
if mask != 0 {
return Some(a + mask.trailing_zeros() as usize);
}
a += 16;
}
next_match_from_scalar(ends, a, limit)
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
unsafe fn next_match_from_avx2(ends: &[i32], from: usize, limit: usize) -> Option<usize> {
use core::arch::x86_64::{
_mm256_add_epi32, _mm256_castsi256_ps, _mm256_cmpgt_epi32, _mm256_loadu_si256,
_mm256_movemask_ps, _mm256_set1_epi32, _mm256_setr_epi32,
};
let lanes = _mm256_setr_epi32(0, 1, 2, 3, 4, 5, 6, 7);
let mut a = from;
while a + 8 <= limit {
let v = unsafe { _mm256_loadu_si256(ends.as_ptr().add(a).cast()) };
let idx = _mm256_add_epi32(lanes, _mm256_set1_epi32(a as i32));
let gt = _mm256_cmpgt_epi32(v, idx);
let mask = _mm256_movemask_ps(_mm256_castsi256_ps(gt)) as u32;
if mask != 0 {
return Some(a + mask.trailing_zeros() as usize);
}
a += 8;
}
next_match_from_scalar(ends, a, limit)
}
#[cfg(test)]
mod tests {
use super::*;
fn sample_ends(n: usize, stride: u32) -> Vec<i32> {
(0..n)
.map(|a| {
let r = (a as u32).wrapping_mul(2_654_435_761) >> 13;
if r.is_multiple_of(7) {
a as i32
} else if r.is_multiple_of(stride) {
a as i32 + 1 + (r % 11) as i32
} else {
-1
}
})
.collect()
}
#[test]
fn dispatched_scan_matches_scalar() {
for stride in [2u32, 3, 8, 80, 4000] {
let ends = sample_ends(5003, stride);
for from in [0usize, 1, 7, 8, 15, 16, 17, 63, 64, 1000, 4999, 5003] {
for limit in [0usize, 1, 8, 16, 31, 32, 33, 1000, 5002, 5003] {
assert_eq!(
next_match_from(&ends, from, limit),
next_match_from_scalar(&ends, from, limit),
"stride {stride} from {from} limit {limit}"
);
}
}
}
}
#[test]
fn an_empty_or_inverted_range_finds_nothing() {
let ends = sample_ends(200, 3);
assert_eq!(next_match_from(&ends, 0, 0), None);
assert_eq!(next_match_from(&ends, 100, 100), None);
assert_eq!(next_match_from(&ends, 150, 20), None);
assert_eq!(next_match_from(&[], 0, 0), None);
assert_eq!(next_match_from(&ends, 0, 100_000), next_match_from_scalar(&ends, 0, 200));
assert_eq!(next_match_from(&[], 0, 64), None);
}
#[test]
fn the_leftmost_match_in_a_window_is_the_one_returned() {
let mut ends = vec![-1i32; 64];
ends[19] = 25;
ends[21] = 30;
assert_eq!(next_match_from(&ends, 0, 64), Some(19));
assert_eq!(next_match_from(&ends, 20, 64), Some(21));
assert_eq!(next_match_from(&ends, 22, 64), None);
}
fn walk_with(ends: &[i32], mut pick: impl FnMut(&[i32], usize, usize) -> Option<usize>) -> Vec<usize> {
let n = ends.len();
let mut a = 0usize;
let mut taken = Vec::new();
while let Some(m) = pick(ends, a, n) {
taken.push(m);
a = ends[m] as usize;
}
taken
}
#[test]
fn the_adaptive_walk_takes_what_the_scalar_walk_takes() {
for stride in [2u32, 3, 8, 80, 4000] {
let ends = sample_ends(5003, stride);
let chosen: fn(&[i32], usize, usize) -> Option<usize> =
if takes_vector(&ends) { next_match_from } else { next_match_from_scalar };
assert_eq!(
walk_with(&ends, chosen),
walk_with(&ends, next_match_from_scalar),
"uniform stride {stride}"
);
}
let mut mixed = sample_ends(4000, 300);
let dense: Vec<i32> = (4000..8000)
.map(|a| {
let r = (a as u32).wrapping_mul(2_654_435_761) >> 13;
if r.is_multiple_of(2) { (a + 1).min(7999) } else { -1 }
})
.collect();
mixed.extend_from_slice(&dense);
let chosen: fn(&[i32], usize, usize) -> Option<usize> =
if takes_vector(&mixed) { next_match_from } else { next_match_from_scalar };
assert_eq!(
walk_with(&mixed, chosen),
walk_with(&mixed, next_match_from_scalar),
"sparse then dense in one array"
);
}
#[cfg(target_arch = "x86_64")]
#[test]
fn each_simd_path_matches_scalar_on_large_input() {
let ends = sample_ends(9001, 60);
let probes: &[(usize, usize)] =
&[(0, 9001), (0, 16), (1, 9000), (17, 9001), (8000, 9001), (9001, 9001)];
if std::is_x86_feature_detected!("avx2") {
for &(from, limit) in probes {
assert_eq!(
unsafe { next_match_from_avx2(&ends, from, limit) },
next_match_from_scalar(&ends, from, limit),
"avx2 from {from} limit {limit}"
);
}
}
if std::is_x86_feature_detected!("avx512f") {
for &(from, limit) in probes {
assert_eq!(
unsafe { next_match_from_avx512(&ends, from, limit) },
next_match_from_scalar(&ends, from, limit),
"avx512 from {from} limit {limit}"
);
}
}
}
}