#[cfg(target_arch = "x86_64")]
use std::arch::x86_64::*;
#[cfg(target_arch = "x86_64")]
pub fn characteristic_vector_simd<'a>(
dict_char: u8,
query: &[u8],
window_size: usize,
offset: usize,
buffer: &'a mut [bool; 8],
) -> &'a [bool] {
let len = window_size.min(8);
if len >= 8 && is_x86_feature_detected!("avx2") {
unsafe { characteristic_vector_avx2(dict_char, query, len, offset, buffer) }
} else if len >= 4 && is_x86_feature_detected!("sse4.1") {
unsafe { characteristic_vector_sse41(dict_char, query, len, offset, buffer) }
} else {
characteristic_vector_scalar(dict_char, query, len, offset, buffer)
}
}
#[inline(always)]
fn characteristic_vector_scalar<'a>(
dict_char: u8,
query: &[u8],
len: usize,
offset: usize,
buffer: &'a mut [bool; 8],
) -> &'a [bool] {
for (i, item) in buffer.iter_mut().enumerate().take(len) {
let query_idx = offset + i;
*item = query_idx < query.len() && query[query_idx] == dict_char;
}
&buffer[..len]
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
unsafe fn characteristic_vector_avx2<'a>(
dict_char: u8,
query: &[u8],
len: usize,
offset: usize,
buffer: &'a mut [bool; 8],
) -> &'a [bool] {
debug_assert!(len == 8, "AVX2 path expects exactly 8 elements");
let dict_vec = _mm256_set1_epi8(dict_char as i8);
let mut query_buf = [0u8; 32]; for (i, cell) in query_buf.iter_mut().enumerate().take(8) {
let query_idx = offset + i;
*cell = if query_idx < query.len() {
query[query_idx]
} else {
0xFF };
}
let query_vec = _mm256_loadu_si256(query_buf.as_ptr() as *const __m256i);
let cmp_result = _mm256_cmpeq_epi8(dict_vec, query_vec);
let mask = _mm256_movemask_epi8(cmp_result);
for (i, cell) in buffer.iter_mut().enumerate().take(8) {
*cell = (mask & (1 << i)) != 0;
}
&buffer[..len]
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "sse4.1")]
unsafe fn characteristic_vector_sse41<'a>(
dict_char: u8,
query: &[u8],
len: usize,
offset: usize,
buffer: &'a mut [bool; 8],
) -> &'a [bool] {
debug_assert!(len >= 4, "SSE4.1 path expects at least 4 elements");
let dict_vec = _mm_set1_epi8(dict_char as i8);
let mut query_buf = [0u8; 4];
for (i, cell) in query_buf.iter_mut().enumerate().take(4) {
let query_idx = offset + i;
*cell = if query_idx < query.len() {
query[query_idx]
} else {
0 };
}
let query_val = u32::from_le_bytes(query_buf);
let query_vec = _mm_set1_epi32(query_val as i32);
let cmp_result = _mm_cmpeq_epi8(dict_vec, query_vec);
let mask = _mm_movemask_epi8(cmp_result);
for (i, cell) in buffer.iter_mut().enumerate().take(4) {
*cell = (mask & (1 << i)) != 0;
}
for (i, cell) in buffer.iter_mut().enumerate().take(len).skip(4) {
let query_idx = offset + i;
*cell = query_idx < query.len() && query[query_idx] == dict_char;
}
&buffer[..len]
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
#[cfg(target_arch = "x86_64")]
fn test_characteristic_vector_simd() {
let test_cases = vec![
(b'a', b"aaaaaaa".as_slice(), 7, 0, vec![true; 7]),
(
b'a',
b"abcdefg".as_slice(),
7,
0,
vec![true, false, false, false, false, false, false],
),
(b'x', b"abcdefg".as_slice(), 7, 0, vec![false; 7]),
(
b'a',
b"abc".as_slice(),
5,
0,
vec![true, false, false, false, false],
), (
b'c',
b"abcdefg".as_slice(),
5,
2,
vec![true, false, false, false, false],
), (
b'a',
b"aaa".as_slice(),
8,
0,
vec![true, true, true, false, false, false, false, false],
), ];
for (dict_char, query, window_size, offset, expected) in test_cases {
let mut buffer = [false; 8];
let result =
characteristic_vector_simd(dict_char, query, window_size, offset, &mut buffer);
assert_eq!(
result,
&expected[..],
"SIMD result mismatch for dict_char={}, query={:?}, window_size={}, offset={}",
dict_char as char,
query,
window_size,
offset
);
let mut scalar_buffer = [false; 8];
let scalar_result = characteristic_vector_scalar(
dict_char,
query,
window_size.min(8),
offset,
&mut scalar_buffer,
);
assert_eq!(
result, scalar_result,
"SIMD vs scalar mismatch for dict_char={}, query={:?}, window_size={}, offset={}",
dict_char as char, query, window_size, offset
);
}
}
}
#[cfg(target_arch = "x86_64")]
pub fn check_subsumption_simd<'a>(
lhs_term_indices: &[usize],
lhs_errors: &[usize],
rhs_term_indices: &[usize],
rhs_errors: &[usize],
count: usize,
results: &'a mut [bool; 8],
) -> &'a [bool] {
debug_assert!(count <= 8, "count must be <= 8");
debug_assert!(lhs_term_indices.len() >= count);
debug_assert!(lhs_errors.len() >= count);
debug_assert!(rhs_term_indices.len() >= count);
debug_assert!(rhs_errors.len() >= count);
if count == 8 && is_x86_feature_detected!("avx2") {
unsafe {
check_subsumption_avx2(
lhs_term_indices,
lhs_errors,
rhs_term_indices,
rhs_errors,
results,
)
}
} else if count >= 4 && is_x86_feature_detected!("sse4.1") {
unsafe {
check_subsumption_sse41(
lhs_term_indices,
lhs_errors,
rhs_term_indices,
rhs_errors,
count,
results,
)
}
} else {
check_subsumption_scalar(
lhs_term_indices,
lhs_errors,
rhs_term_indices,
rhs_errors,
count,
results,
)
}
}
#[inline(always)]
fn check_subsumption_scalar<'a>(
lhs_term_indices: &[usize],
lhs_errors: &[usize],
rhs_term_indices: &[usize],
rhs_errors: &[usize],
count: usize,
results: &'a mut [bool; 8],
) -> &'a [bool] {
for idx in 0..count {
let i = lhs_term_indices[idx];
let e = lhs_errors[idx];
let j = rhs_term_indices[idx];
let f = rhs_errors[idx];
if e > f {
results[idx] = false;
continue;
}
let index_diff = i.abs_diff(j);
let error_diff = f - e;
results[idx] = index_diff <= error_diff;
}
&results[..count]
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
unsafe fn check_subsumption_avx2<'a>(
lhs_term_indices: &[usize],
lhs_errors: &[usize],
rhs_term_indices: &[usize],
rhs_errors: &[usize],
results: &'a mut [bool; 8],
) -> &'a [bool] {
let mut lhs_i_buf = [0u32; 8];
let mut lhs_e_buf = [0u32; 8];
let mut rhs_j_buf = [0u32; 8];
let mut rhs_f_buf = [0u32; 8];
for idx in 0..8 {
lhs_i_buf[idx] = lhs_term_indices[idx] as u32;
lhs_e_buf[idx] = lhs_errors[idx] as u32;
rhs_j_buf[idx] = rhs_term_indices[idx] as u32;
rhs_f_buf[idx] = rhs_errors[idx] as u32;
}
let i_vec = _mm256_loadu_si256(lhs_i_buf.as_ptr() as *const __m256i);
let e_vec = _mm256_loadu_si256(lhs_e_buf.as_ptr() as *const __m256i);
let j_vec = _mm256_loadu_si256(rhs_j_buf.as_ptr() as *const __m256i);
let f_vec = _mm256_loadu_si256(rhs_f_buf.as_ptr() as *const __m256i);
let e_gt_f = _mm256_cmpgt_epi32(e_vec, f_vec);
let i_sub_j = _mm256_sub_epi32(i_vec, j_vec);
let j_sub_i = _mm256_sub_epi32(j_vec, i_vec);
let abs_diff = _mm256_max_epi32(i_sub_j, j_sub_i);
let error_diff = _mm256_sub_epi32(f_vec, e_vec);
let abs_gt_error = _mm256_cmpgt_epi32(abs_diff, error_diff);
let subsumes_mask = _mm256_andnot_si256(abs_gt_error, _mm256_set1_epi32(-1));
let final_mask = _mm256_andnot_si256(e_gt_f, subsumes_mask);
let mask = _mm256_movemask_ps(_mm256_castsi256_ps(final_mask));
for (idx, result) in results.iter_mut().enumerate().take(8) {
*result = (mask & (1 << idx)) != 0;
}
&results[..8]
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "sse4.1")]
unsafe fn check_subsumption_sse41<'a>(
lhs_term_indices: &[usize],
lhs_errors: &[usize],
rhs_term_indices: &[usize],
rhs_errors: &[usize],
count: usize,
results: &'a mut [bool; 8],
) -> &'a [bool] {
debug_assert!((4..=8).contains(&count));
let mut lhs_i_buf = [0u32; 4];
let mut lhs_e_buf = [0u32; 4];
let mut rhs_j_buf = [0u32; 4];
let mut rhs_f_buf = [0u32; 4];
for idx in 0..4 {
lhs_i_buf[idx] = lhs_term_indices[idx] as u32;
lhs_e_buf[idx] = lhs_errors[idx] as u32;
rhs_j_buf[idx] = rhs_term_indices[idx] as u32;
rhs_f_buf[idx] = rhs_errors[idx] as u32;
}
let i_vec = _mm_loadu_si128(lhs_i_buf.as_ptr() as *const __m128i);
let e_vec = _mm_loadu_si128(lhs_e_buf.as_ptr() as *const __m128i);
let j_vec = _mm_loadu_si128(rhs_j_buf.as_ptr() as *const __m128i);
let f_vec = _mm_loadu_si128(rhs_f_buf.as_ptr() as *const __m128i);
let e_gt_f = _mm_cmpgt_epi32(e_vec, f_vec);
let i_sub_j = _mm_sub_epi32(i_vec, j_vec);
let j_sub_i = _mm_sub_epi32(j_vec, i_vec);
let abs_diff = _mm_max_epi32(i_sub_j, j_sub_i);
let error_diff = _mm_sub_epi32(f_vec, e_vec);
let abs_gt_error = _mm_cmpgt_epi32(abs_diff, error_diff);
let subsumes_mask = _mm_andnot_si128(abs_gt_error, _mm_set1_epi32(-1));
let final_mask = _mm_andnot_si128(e_gt_f, subsumes_mask);
let mask = _mm_movemask_ps(_mm_castsi128_ps(final_mask));
for (idx, result) in results.iter_mut().enumerate().take(4) {
*result = (mask & (1 << idx)) != 0;
}
for idx in 4..count {
let i = lhs_term_indices[idx];
let e = lhs_errors[idx];
let j = rhs_term_indices[idx];
let f = rhs_errors[idx];
results[idx] = e <= f && i.abs_diff(j) <= (f - e);
}
&results[..count]
}
#[cfg(test)]
mod subsumption_tests {
use super::*;
#[test]
#[cfg(target_arch = "x86_64")]
fn test_subsumption_simd_basic() {
let test_cases = vec![
(5, 2, 5, 3, true),
(5, 2, 4, 3, true),
(3, 2, 3, 2, true),
(3, 3, 5, 2, false),
(10, 1, 8, 3, true),
(10, 1, 5, 3, false),
(0, 0, 0, 0, true),
(0, 0, 1, 1, true),
];
for (lhs_i, lhs_e, rhs_j, rhs_f, expected) in test_cases {
let lhs_indices = [lhs_i];
let lhs_errs = [lhs_e];
let rhs_indices = [rhs_j];
let rhs_errs = [rhs_f];
let mut results = [false; 8];
let result = check_subsumption_simd(
&lhs_indices,
&lhs_errs,
&rhs_indices,
&rhs_errs,
1,
&mut results,
);
assert_eq!(
result[0], expected,
"SIMD result mismatch for ({}, {}) vs ({}, {}): expected {}, got {}",
lhs_i, lhs_e, rhs_j, rhs_f, expected, result[0]
);
}
}
#[test]
#[cfg(target_arch = "x86_64")]
fn test_subsumption_simd_batch() {
let lhs_indices = [5, 5, 3, 3, 10, 10, 0, 0];
let lhs_errs = [2, 2, 2, 3, 1, 1, 0, 0];
let rhs_indices = [5, 4, 3, 5, 8, 5, 0, 1];
let rhs_errs = [3, 3, 2, 2, 3, 3, 0, 1];
let expected = [true, true, true, false, true, false, true, true];
let mut results = [false; 8];
let result = check_subsumption_simd(
&lhs_indices,
&lhs_errs,
&rhs_indices,
&rhs_errs,
8,
&mut results,
);
for i in 0..8 {
assert_eq!(
result[i], expected[i],
"Batch SIMD result mismatch at index {}: ({}, {}) vs ({}, {}) - expected {}, got {}",
i, lhs_indices[i], lhs_errs[i], rhs_indices[i], rhs_errs[i], expected[i], result[i]
);
}
}
#[test]
#[cfg(target_arch = "x86_64")]
fn test_subsumption_simd_vs_scalar() {
let test_cases = vec![
([0, 1, 5, 10, 3, 7, 2, 15], [0, 0, 2, 1, 3, 2, 1, 5]),
([0, 1, 2, 3, 4, 5, 6, 7], [1, 1, 1, 1, 1, 1, 1, 1]),
([10, 10, 10, 10, 10, 10, 10, 10], [0, 1, 2, 3, 4, 5, 6, 7]),
];
let rhs_cases = vec![
([0, 2, 6, 12, 5, 9, 4, 20], [0, 1, 3, 3, 4, 4, 2, 8]),
([0, 1, 2, 3, 4, 5, 6, 7], [1, 2, 2, 2, 2, 2, 2, 2]),
([10, 11, 12, 13, 14, 15, 16, 17], [1, 2, 3, 4, 5, 6, 7, 8]),
];
for (lhs_i, lhs_e) in &test_cases {
for (rhs_j, rhs_f) in &rhs_cases {
let mut simd_results = [false; 8];
let mut scalar_results = [false; 8];
let simd_result =
check_subsumption_simd(lhs_i, lhs_e, rhs_j, rhs_f, 8, &mut simd_results);
let scalar_result =
check_subsumption_scalar(lhs_i, lhs_e, rhs_j, rhs_f, 8, &mut scalar_results);
assert_eq!(
simd_result, scalar_result,
"SIMD vs scalar mismatch:\nlhs_i: {:?}\nlhs_e: {:?}\nrhs_j: {:?}\nrhs_f: {:?}\nSIMD: {:?}\nScalar: {:?}",
lhs_i, lhs_e, rhs_j, rhs_f, simd_result, scalar_result
);
}
}
}
#[test]
#[cfg(target_arch = "x86_64")]
fn test_subsumption_simd_edge_cases() {
let lhs_indices = [1000, 2000, 5000, 10000, 100, 200, 300, 400];
let lhs_errs = [10, 20, 50, 100, 5, 10, 15, 20];
let rhs_indices = [1010, 2025, 5060, 10150, 110, 215, 325, 440];
let rhs_errs = [20, 50, 120, 250, 20, 25, 40, 60];
let mut results = [false; 8];
let result = check_subsumption_simd(
&lhs_indices,
&lhs_errs,
&rhs_indices,
&rhs_errs,
8,
&mut results,
);
let mut scalar_results = [false; 8];
check_subsumption_scalar(
&lhs_indices,
&lhs_errs,
&rhs_indices,
&rhs_errs,
8,
&mut scalar_results,
);
assert_eq!(
result,
&scalar_results[..8],
"SIMD vs scalar mismatch on large indices"
);
let lhs_indices = [0, 1, 2, 3, 4, 5, 6, 7];
let lhs_errs = [0, 0, 0, 0, 0, 0, 0, 0];
let rhs_indices = [0, 1, 2, 3, 4, 5, 6, 7];
let rhs_errs = [0, 0, 0, 0, 0, 0, 0, 0];
let mut results = [false; 8];
let result = check_subsumption_simd(
&lhs_indices,
&lhs_errs,
&rhs_indices,
&rhs_errs,
8,
&mut results,
);
assert_eq!(
result,
&[true, true, true, true, true, true, true, true],
"Zero errors edge case failed"
);
}
#[test]
#[cfg(target_arch = "x86_64")]
fn test_subsumption_simd_partial_batches() {
let lhs_indices = [5, 5, 3, 3, 10];
let lhs_errs = [2, 2, 2, 3, 1];
let rhs_indices = [5, 4, 3, 5, 8];
let rhs_errs = [3, 3, 2, 2, 3];
let expected = [true, true, true, false, true];
let mut results = [false; 8];
let result = check_subsumption_simd(
&lhs_indices,
&lhs_errs,
&rhs_indices,
&rhs_errs,
5,
&mut results,
);
for i in 0..5 {
assert_eq!(
result[i], expected[i],
"Partial batch (count=5) mismatch at index {}: expected {}, got {}",
i, expected[i], result[i]
);
}
let mut results = [false; 8];
let result = check_subsumption_simd(
&lhs_indices,
&lhs_errs,
&rhs_indices,
&rhs_errs,
4,
&mut results,
);
for i in 0..4 {
assert_eq!(
result[i], expected[i],
"Partial batch (count=4) mismatch at index {}: expected {}, got {}",
i, expected[i], result[i]
);
}
let mut results = [false; 8];
let result = check_subsumption_simd(
&lhs_indices,
&lhs_errs,
&rhs_indices,
&rhs_errs,
2,
&mut results,
);
for i in 0..2 {
assert_eq!(
result[i], expected[i],
"Partial batch (count=2) mismatch at index {}: expected {}, got {}",
i, expected[i], result[i]
);
}
}
}
#[cfg(target_arch = "x86_64")]
pub fn find_minimum_simd(values: &[usize], count: usize) -> usize {
debug_assert!(count > 0 && count <= 8, "count must be in range 1..=8");
debug_assert!(values.len() >= count);
if count == 1 {
return values[0];
}
if count == 8 && is_x86_feature_detected!("avx2") {
unsafe { find_minimum_avx2(values) }
} else if count >= 4 && is_x86_feature_detected!("sse4.1") {
unsafe { find_minimum_sse41(values, count) }
} else {
find_minimum_scalar(values, count)
}
}
#[inline(always)]
fn find_minimum_scalar(values: &[usize], count: usize) -> usize {
values[0..count]
.iter()
.copied()
.min()
.expect("find_minimum_scalar: count > 0 (caller invariant)")
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
unsafe fn find_minimum_avx2(values: &[usize]) -> usize {
debug_assert!(values.len() >= 8);
let mut buf = [0u32; 8];
for i in 0..8 {
buf[i] = values[i] as u32;
}
let vec = _mm256_loadu_si256(buf.as_ptr() as *const __m256i);
let high = _mm256_extracti128_si256(vec, 1);
let low = _mm256_castsi256_si128(vec);
let min_half = _mm_min_epu32(low, high);
let shuffled = _mm_shuffle_epi32(min_half, 0b01_00_11_10);
let min_pairs = _mm_min_epu32(min_half, shuffled);
let final_shuffle = _mm_shuffle_epi32(min_pairs, 0b00_00_00_01);
let final_min = _mm_min_epu32(min_pairs, final_shuffle);
_mm_extract_epi32(final_min, 0) as usize
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "sse4.1")]
unsafe fn find_minimum_sse41(values: &[usize], count: usize) -> usize {
debug_assert!((4..=8).contains(&count));
let mut buf = [0u32; 4];
for i in 0..4 {
buf[i] = values[i] as u32;
}
let vec = _mm_loadu_si128(buf.as_ptr() as *const __m128i);
let shuffled = _mm_shuffle_epi32(vec, 0b01_00_11_10);
let min_pairs = _mm_min_epu32(vec, shuffled);
let final_shuffle = _mm_shuffle_epi32(min_pairs, 0b00_00_00_01);
let final_min = _mm_min_epu32(min_pairs, final_shuffle);
let mut min_val = _mm_extract_epi32(final_min, 0) as usize;
for value in values.iter().take(count).skip(4) {
min_val = min_val.min(*value);
}
min_val
}
#[cfg(test)]
mod minimum_tests {
use super::*;
#[test]
#[cfg(target_arch = "x86_64")]
fn test_find_minimum_simd_basic() {
let test_cases = vec![
(vec![5], 1, 5),
(vec![10, 3], 2, 3),
(vec![1, 20], 2, 1),
(vec![5, 2, 8, 1], 4, 1),
(vec![100, 50, 75, 25], 4, 25),
(vec![10, 3, 7, 2, 15, 1, 9, 5], 8, 1),
(vec![100, 200, 50, 300, 25, 400, 150, 75], 8, 25),
(vec![1, 2, 3, 4, 5, 6, 7, 8], 8, 1), (vec![8, 7, 6, 5, 4, 3, 2, 1], 8, 1), (vec![5, 4, 3, 2, 1, 2, 3, 4], 8, 1), ];
for (values, count, expected) in test_cases {
let result = find_minimum_simd(&values, count);
assert_eq!(
result, expected,
"SIMD minimum mismatch for values={:?}, count={}: expected {}, got {}",
values, count, expected, result
);
}
}
#[test]
#[cfg(target_arch = "x86_64")]
fn test_find_minimum_simd_vs_scalar() {
let test_cases = vec![
vec![5, 2, 8, 1, 15, 3, 9, 6],
vec![100, 200, 50, 300, 25, 400, 150, 75],
vec![10, 10, 10, 10, 10, 10, 10, 10], vec![0, 1, 2, 3, 4, 5, 6, 7], vec![1000, 500, 250, 125, 62, 31, 15, 7], ];
for values in test_cases {
for count in 1..=8 {
let simd_result = find_minimum_simd(&values, count);
let scalar_result = find_minimum_scalar(&values, count);
assert_eq!(
simd_result, scalar_result,
"SIMD vs scalar mismatch for values={:?}, count={}: SIMD={}, scalar={}",
values, count, simd_result, scalar_result
);
}
}
}
#[test]
#[cfg(target_arch = "x86_64")]
fn test_find_minimum_edge_cases() {
let values = vec![42, 42, 42, 42, 42, 42, 42, 42];
assert_eq!(find_minimum_simd(&values, 8), 42);
let values = vec![
1000000, 2000000, 500000, 3000000, 250000, 4000000, 1500000, 750000,
];
assert_eq!(find_minimum_simd(&values, 8), 250000);
let values = vec![10, 5, 0, 3, 7, 2, 8, 4];
assert_eq!(find_minimum_simd(&values, 8), 0);
let values = vec![10, 5, 3, 7, 2];
assert_eq!(find_minimum_simd(&values, 5), 2);
assert_eq!(find_minimum_simd(&values, 3), 3);
}
#[test]
#[cfg(target_arch = "x86_64")]
fn test_find_minimum_real_world() {
let errors = vec![2, 1];
assert_eq!(find_minimum_simd(&errors, 2), 1);
let errors = vec![3, 1, 2, 4];
assert_eq!(find_minimum_simd(&errors, 4), 1);
let errors = vec![5, 2, 7, 1, 3, 6, 4, 8];
assert_eq!(find_minimum_simd(&errors, 8), 1);
let errors = vec![0, 1, 2, 3];
assert_eq!(find_minimum_simd(&errors, 4), 0);
}
}
#[cfg(target_arch = "x86_64")]
pub fn find_edge_label_simd<T>(edges: &[(u8, T)], target_label: u8) -> Option<usize> {
let count = edges.len();
if count < 12 {
return find_edge_label_scalar(edges, target_label);
}
if count <= 16 && is_x86_feature_detected!("sse4.1") {
return unsafe { find_edge_label_sse41(edges, target_label, count) };
}
find_edge_label_scalar(edges, target_label)
}
#[inline]
fn find_edge_label_scalar<T>(edges: &[(u8, T)], target_label: u8) -> Option<usize> {
edges.iter().position(|(label, _)| *label == target_label)
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "sse4.1")]
unsafe fn find_edge_label_sse41<T>(
edges: &[(u8, T)],
target_label: u8,
count: usize,
) -> Option<usize> {
use std::arch::x86_64::*;
let mut labels = [0xFFu8; 16]; for (i, (label, _)) in edges.iter().enumerate().take(16.min(count)) {
labels[i] = *label;
}
let labels_vec = _mm_loadu_si128(labels.as_ptr() as *const __m128i);
let target_vec = _mm_set1_epi8(target_label as i8);
let cmp_result = _mm_cmpeq_epi8(labels_vec, target_vec);
let mask = _mm_movemask_epi8(cmp_result);
if mask != 0 {
let index = mask.trailing_zeros() as usize;
if index < count {
return Some(index);
}
}
None
}
#[allow(dead_code)]
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
unsafe fn find_edge_label_avx2<T>(
edges: &[(u8, T)],
target_label: u8,
count: usize,
) -> Option<usize> {
use std::arch::x86_64::*;
let mut labels = [0xFFu8; 32]; for (i, (label, _)) in edges.iter().enumerate().take(32.min(count)) {
labels[i] = *label;
}
let labels_vec = _mm256_loadu_si256(labels.as_ptr() as *const __m256i);
let target_vec = _mm256_set1_epi8(target_label as i8);
let cmp_result = _mm256_cmpeq_epi8(labels_vec, target_vec);
let mask = _mm256_movemask_epi8(cmp_result);
if mask != 0 {
let index = mask.trailing_zeros() as usize;
if index < count {
return Some(index);
}
}
None
}
#[cfg(test)]
mod edge_lookup_tests {
use super::*;
#[test]
fn test_edge_not_found() {
let edges = vec![(b'a', 1), (b'c', 2), (b'e', 3), (b'g', 4)];
assert_eq!(find_edge_label_simd(&edges, b'z'), None);
assert_eq!(find_edge_label_simd(&edges, b'b'), None);
}
#[test]
fn test_edge_at_beginning() {
let edges = vec![(b'a', 10), (b'b', 20), (b'c', 30), (b'd', 40)];
assert_eq!(find_edge_label_simd(&edges, b'a'), Some(0));
}
#[test]
fn test_edge_at_middle() {
let edges = vec![(b'a', 10), (b'b', 20), (b'c', 30), (b'd', 40)];
assert_eq!(find_edge_label_simd(&edges, b'b'), Some(1));
assert_eq!(find_edge_label_simd(&edges, b'c'), Some(2));
}
#[test]
fn test_edge_at_end() {
let edges = vec![(b'a', 10), (b'b', 20), (b'c', 30), (b'd', 40)];
assert_eq!(find_edge_label_simd(&edges, b'd'), Some(3));
}
#[test]
fn test_empty_edges() {
let edges: Vec<(u8, usize)> = vec![];
assert_eq!(find_edge_label_simd(&edges, b'a'), None);
}
#[test]
fn test_single_edge() {
let edges = vec![(b'x', 100)];
assert_eq!(find_edge_label_simd(&edges, b'x'), Some(0));
assert_eq!(find_edge_label_simd(&edges, b'y'), None);
}
#[test]
fn test_exactly_4_edges() {
let edges = vec![(b'a', 1), (b'b', 2), (b'c', 3), (b'd', 4)];
assert_eq!(find_edge_label_simd(&edges, b'a'), Some(0));
assert_eq!(find_edge_label_simd(&edges, b'c'), Some(2));
assert_eq!(find_edge_label_simd(&edges, b'd'), Some(3));
assert_eq!(find_edge_label_simd(&edges, b'z'), None);
}
#[test]
fn test_8_edges() {
let edges = vec![
(b'a', 1),
(b'b', 2),
(b'c', 3),
(b'd', 4),
(b'e', 5),
(b'f', 6),
(b'g', 7),
(b'h', 8),
];
assert_eq!(find_edge_label_simd(&edges, b'a'), Some(0));
assert_eq!(find_edge_label_simd(&edges, b'e'), Some(4));
assert_eq!(find_edge_label_simd(&edges, b'h'), Some(7));
assert_eq!(find_edge_label_simd(&edges, b'z'), None);
}
#[test]
fn test_exactly_16_edges() {
let edges: Vec<(u8, usize)> = (0..16).map(|i| (b'a' + i, i as usize)).collect();
assert_eq!(find_edge_label_simd(&edges, b'a'), Some(0)); assert_eq!(find_edge_label_simd(&edges, b'h'), Some(7)); assert_eq!(find_edge_label_simd(&edges, b'p'), Some(15)); assert_eq!(find_edge_label_simd(&edges, b'z'), None); }
#[test]
fn test_32_edges() {
let edges: Vec<(u8, usize)> = (0..32).map(|i| (i as u8, i as usize)).collect();
assert_eq!(find_edge_label_simd(&edges, 0), Some(0)); assert_eq!(find_edge_label_simd(&edges, 16), Some(16)); assert_eq!(find_edge_label_simd(&edges, 31), Some(31)); assert_eq!(find_edge_label_simd(&edges, 255), None); }
#[test]
fn test_boundary_15_edges() {
let edges: Vec<(u8, usize)> = (0..15).map(|i| (b'a' + i, i as usize)).collect();
assert_eq!(find_edge_label_simd(&edges, b'a'), Some(0));
assert_eq!(find_edge_label_simd(&edges, b'o'), Some(14)); assert_eq!(find_edge_label_simd(&edges, b'z'), None);
}
#[test]
fn test_boundary_17_edges() {
let edges: Vec<(u8, usize)> = (0..17).map(|i| (b'a' + i, i as usize)).collect();
assert_eq!(find_edge_label_simd(&edges, b'a'), Some(0));
assert_eq!(find_edge_label_simd(&edges, b'q'), Some(16)); assert_eq!(find_edge_label_simd(&edges, b'z'), None);
}
#[test]
fn test_all_positions_in_16_edges() {
let edges: Vec<(u8, usize)> = (0..16).map(|i| (b'a' + i, i as usize * 10)).collect();
for i in 0..16 {
let label = b'a' + i;
assert_eq!(
find_edge_label_simd(&edges, label),
Some(i as usize),
"Failed to find label '{}' at position {}",
label as char,
i
);
}
}
#[test]
fn test_realistic_english_letters() {
let edges = vec![
(b'a', 1),
(b'e', 2),
(b'i', 3),
(b'n', 4),
(b'o', 5),
(b'r', 6),
(b's', 7),
(b't', 8),
];
assert_eq!(find_edge_label_simd(&edges, b'a'), Some(0));
assert_eq!(find_edge_label_simd(&edges, b'e'), Some(1));
assert_eq!(find_edge_label_simd(&edges, b't'), Some(7));
assert_eq!(find_edge_label_simd(&edges, b'x'), None);
}
}