#[cfg(target_arch = "x86_64")]
use std::arch::x86_64::*;
#[cfg(target_arch = "x86_64")]
pub fn standard_distance_simd(source: &str, target: &str) -> usize {
let source_len = source.chars().count();
let target_len = target.chars().count();
if source_len < 16 && target_len < 16 {
return crate::distance::standard_distance_impl(source, target);
}
if is_x86_feature_detected!("avx2") {
unsafe { standard_distance_avx2(source, target) }
} else if is_x86_feature_detected!("sse4.1") {
unsafe { standard_distance_sse41(source, target) }
} else {
crate::distance::standard_distance_impl(source, target)
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
unsafe fn standard_distance_avx2(source: &str, target: &str) -> usize {
use smallvec::SmallVec;
let source_chars: SmallVec<[char; 32]> = source.chars().collect();
let target_chars: SmallVec<[char; 32]> = target.chars().collect();
let m = source_chars.len();
let n = target_chars.len();
if m == 0 {
return n;
}
if n == 0 {
return m;
}
if n < 16 {
return crate::distance::standard_distance_impl(source, target);
}
let mut prev_row = vec![0u32; n + 1];
let mut curr_row = vec![0u32; n + 1];
init_row_simd_avx2(&mut prev_row);
let one_vec = _mm256_set1_epi32(1);
for i in 1..=m {
curr_row[0] = i as u32;
let source_char = source_chars[i - 1] as u32;
let mut j = 1;
while j + 8 <= n {
let mut target_buf = [0u32; 8];
for k in 0..8 {
target_buf[k] = target_chars[j - 1 + k] as u32;
}
let target_vec = _mm256_loadu_si256(target_buf.as_ptr() as *const __m256i);
let source_vec = _mm256_set1_epi32(source_char as i32);
let eq_mask = _mm256_cmpeq_epi32(source_vec, target_vec);
let costs = _mm256_andnot_si256(eq_mask, one_vec);
let prev_same = _mm256_loadu_si256(prev_row.as_ptr().add(j) as *const __m256i);
let prev_diag = _mm256_loadu_si256(prev_row.as_ptr().add(j - 1) as *const __m256i);
let deletion = _mm256_add_epi32(prev_same, one_vec);
let substitution = _mm256_add_epi32(prev_diag, costs);
if j == 1 {
let cost = if source_char == target_chars[0] as u32 {
0
} else {
1
};
curr_row[1] = min3_scalar(prev_row[1] + 1, curr_row[0] + 1, prev_row[0] + cost);
j += 1;
continue;
}
let min_del_sub = _mm256_min_epu32(deletion, substitution);
let mut partial = [0u32; 8];
_mm256_storeu_si256(partial.as_mut_ptr() as *mut __m256i, min_del_sub);
for k in 0..8.min(n + 1 - j) {
let insertion = curr_row[j + k - 1] + 1;
curr_row[j + k] = partial[k].min(insertion);
}
j += 8;
}
for j in j..=n {
let cost = if source_char == target_chars[j - 1] as u32 {
0
} else {
1
};
curr_row[j] = min3_scalar(prev_row[j] + 1, curr_row[j - 1] + 1, prev_row[j - 1] + cost);
}
std::mem::swap(&mut prev_row, &mut curr_row);
}
prev_row[n] as usize
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
unsafe fn init_row_simd_avx2(row: &mut [u32]) {
let n = row.len();
let simd_count = n / 8;
let base = _mm256_setr_epi32(0, 1, 2, 3, 4, 5, 6, 7);
let increment = _mm256_set1_epi32(8);
let mut current = base;
for i in 0..simd_count {
_mm256_storeu_si256(row.as_mut_ptr().add(i * 8) as *mut __m256i, current);
current = _mm256_add_epi32(current, increment);
}
for (i, cell) in row
.iter_mut()
.enumerate()
.skip(simd_count * 8)
.take(n - simd_count * 8)
{
*cell = i as u32;
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "sse4.1")]
unsafe fn init_row_simd_sse41(row: &mut [u32]) {
let n = row.len();
let simd_count = n / 4;
let base = _mm_setr_epi32(0, 1, 2, 3);
let increment = _mm_set1_epi32(4);
let mut current = base;
for i in 0..simd_count {
_mm_storeu_si128(row.as_mut_ptr().add(i * 4) as *mut __m128i, current);
current = _mm_add_epi32(current, increment);
}
for (i, cell) in row
.iter_mut()
.enumerate()
.skip(simd_count * 4)
.take(n - simd_count * 4)
{
*cell = i as u32;
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "sse4.1")]
unsafe fn standard_distance_sse41(source: &str, target: &str) -> usize {
use smallvec::SmallVec;
let source_chars: SmallVec<[char; 32]> = source.chars().collect();
let target_chars: SmallVec<[char; 32]> = target.chars().collect();
let m = source_chars.len();
let n = target_chars.len();
if m == 0 {
return n;
}
if n == 0 {
return m;
}
if n < 8 {
return crate::distance::standard_distance_impl(source, target);
}
let mut prev_row = vec![0u32; n + 1];
let mut curr_row = vec![0u32; n + 1];
init_row_simd_sse41(&mut prev_row);
let one_vec = _mm_set1_epi32(1);
for i in 1..=m {
curr_row[0] = i as u32;
let source_char = source_chars[i - 1] as u32;
let mut j = 1;
while j + 4 <= n {
let mut target_buf = [0u32; 4];
for k in 0..4 {
target_buf[k] = target_chars[j - 1 + k] as u32;
}
let target_vec = _mm_loadu_si128(target_buf.as_ptr() as *const __m128i);
let source_vec = _mm_set1_epi32(source_char as i32);
let eq_mask = _mm_cmpeq_epi32(source_vec, target_vec);
let costs = _mm_andnot_si128(eq_mask, one_vec);
let prev_same = _mm_loadu_si128(prev_row.as_ptr().add(j) as *const __m128i);
let prev_diag = _mm_loadu_si128(prev_row.as_ptr().add(j - 1) as *const __m128i);
let deletion = _mm_add_epi32(prev_same, one_vec);
let substitution = _mm_add_epi32(prev_diag, costs);
if j == 1 {
let cost = if source_char == target_chars[0] as u32 {
0
} else {
1
};
curr_row[1] = min3_scalar(prev_row[1] + 1, curr_row[0] + 1, prev_row[0] + cost);
j += 1;
continue;
}
let min_del_sub = _mm_min_epu32(deletion, substitution);
let mut partial = [0u32; 4];
_mm_storeu_si128(partial.as_mut_ptr() as *mut __m128i, min_del_sub);
for k in 0..4.min(n + 1 - j) {
let insertion = curr_row[j + k - 1] + 1;
curr_row[j + k] = partial[k].min(insertion);
}
j += 4;
}
for j in j..=n {
let cost = if source_char == target_chars[j - 1] as u32 {
0
} else {
1
};
curr_row[j] = min3_scalar(prev_row[j] + 1, curr_row[j - 1] + 1, prev_row[j - 1] + cost);
}
std::mem::swap(&mut prev_row, &mut curr_row);
}
prev_row[n] as usize
}
#[inline(always)]
fn min3_scalar(a: u32, b: u32, c: u32) -> u32 {
a.min(b).min(c)
}
#[allow(dead_code)]
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
unsafe fn min3_avx2(a: __m256i, b: __m256i, c: __m256i) -> __m256i {
let ab_min = _mm256_min_epu32(a, b);
_mm256_min_epu32(ab_min, c)
}
pub fn strip_common_affixes_simd(a: &str, b: &str) -> (usize, usize, usize) {
use smallvec::SmallVec;
let a_chars: SmallVec<[char; 32]> = a.chars().collect();
let b_chars: SmallVec<[char; 32]> = b.chars().collect();
let len_a = a_chars.len();
let len_b = b_chars.len();
if len_a == 0 || len_b == 0 {
return (0, len_a, len_b);
}
let min_len = len_a.min(len_b);
let prefix_len = if is_x86_feature_detected!("avx2") && min_len >= 8 {
unsafe { find_common_prefix_avx2(&a_chars, &b_chars, min_len) }
} else if is_x86_feature_detected!("sse4.1") && min_len >= 4 {
unsafe { find_common_prefix_sse41(&a_chars, &b_chars, min_len) }
} else {
find_common_prefix_scalar(&a_chars, &b_chars, min_len)
};
if prefix_len == min_len {
return (prefix_len, len_a - prefix_len, len_b - prefix_len);
}
let suffix_len = if is_x86_feature_detected!("avx2") && (min_len - prefix_len) >= 8 {
unsafe { find_common_suffix_avx2(&a_chars, &b_chars, len_a, len_b, min_len, prefix_len) }
} else if is_x86_feature_detected!("sse4.1") && (min_len - prefix_len) >= 4 {
unsafe { find_common_suffix_sse41(&a_chars, &b_chars, len_a, len_b, min_len, prefix_len) }
} else {
find_common_suffix_scalar(&a_chars, &b_chars, len_a, len_b, min_len, prefix_len)
};
(
prefix_len,
len_a - prefix_len - suffix_len,
len_b - prefix_len - suffix_len,
)
}
#[inline(always)]
fn find_common_prefix_scalar(a: &[char], b: &[char], min_len: usize) -> usize {
let mut prefix_len = 0;
while prefix_len < min_len && a[prefix_len] == b[prefix_len] {
prefix_len += 1;
}
prefix_len
}
#[inline(always)]
fn find_common_suffix_scalar(
a: &[char],
b: &[char],
len_a: usize,
len_b: usize,
min_len: usize,
prefix_len: usize,
) -> usize {
let mut suffix_len = 0;
while suffix_len < (min_len - prefix_len)
&& a[len_a - 1 - suffix_len] == b[len_b - 1 - suffix_len]
{
suffix_len += 1;
}
suffix_len
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
unsafe fn find_common_prefix_avx2(a: &[char], b: &[char], min_len: usize) -> usize {
let mut prefix_len = 0;
while prefix_len + 8 <= min_len {
let mut a_buf = [0u32; 8];
let mut b_buf = [0u32; 8];
for i in 0..8 {
a_buf[i] = a[prefix_len + i] as u32;
b_buf[i] = b[prefix_len + i] as u32;
}
let a_vec = _mm256_loadu_si256(a_buf.as_ptr() as *const __m256i);
let b_vec = _mm256_loadu_si256(b_buf.as_ptr() as *const __m256i);
let eq_mask = _mm256_cmpeq_epi32(a_vec, b_vec);
let mask = _mm256_movemask_ps(_mm256_castsi256_ps(eq_mask));
if mask != 0xFF {
for i in 0..8 {
if a_buf[i] != b_buf[i] {
return prefix_len + i;
}
}
}
prefix_len += 8;
}
while prefix_len < min_len && a[prefix_len] == b[prefix_len] {
prefix_len += 1;
}
prefix_len
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "sse4.1")]
unsafe fn find_common_prefix_sse41(a: &[char], b: &[char], min_len: usize) -> usize {
let mut prefix_len = 0;
while prefix_len + 4 <= min_len {
let mut a_buf = [0u32; 4];
let mut b_buf = [0u32; 4];
for i in 0..4 {
a_buf[i] = a[prefix_len + i] as u32;
b_buf[i] = b[prefix_len + i] as u32;
}
let a_vec = _mm_loadu_si128(a_buf.as_ptr() as *const __m128i);
let b_vec = _mm_loadu_si128(b_buf.as_ptr() as *const __m128i);
let eq_mask = _mm_cmpeq_epi32(a_vec, b_vec);
let mask = _mm_movemask_ps(_mm_castsi128_ps(eq_mask));
if mask != 0xF {
for i in 0..4 {
if a_buf[i] != b_buf[i] {
return prefix_len + i;
}
}
}
prefix_len += 4;
}
while prefix_len < min_len && a[prefix_len] == b[prefix_len] {
prefix_len += 1;
}
prefix_len
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
unsafe fn find_common_suffix_avx2(
a: &[char],
b: &[char],
len_a: usize,
len_b: usize,
min_len: usize,
prefix_len: usize,
) -> usize {
let mut suffix_len = 0;
let max_suffix = min_len - prefix_len;
while suffix_len + 8 <= max_suffix {
let mut a_buf = [0u32; 8];
let mut b_buf = [0u32; 8];
for i in 0..8 {
a_buf[i] = a[len_a - 1 - suffix_len - (7 - i)] as u32;
b_buf[i] = b[len_b - 1 - suffix_len - (7 - i)] as u32;
}
let a_vec = _mm256_loadu_si256(a_buf.as_ptr() as *const __m256i);
let b_vec = _mm256_loadu_si256(b_buf.as_ptr() as *const __m256i);
let eq_mask = _mm256_cmpeq_epi32(a_vec, b_vec);
let mask = _mm256_movemask_ps(_mm256_castsi256_ps(eq_mask));
if mask != 0xFF {
for i in (0..8).rev() {
if a_buf[i] != b_buf[i] {
return suffix_len + (7 - i);
}
}
}
suffix_len += 8;
}
while suffix_len < max_suffix && a[len_a - 1 - suffix_len] == b[len_b - 1 - suffix_len] {
suffix_len += 1;
}
suffix_len
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "sse4.1")]
unsafe fn find_common_suffix_sse41(
a: &[char],
b: &[char],
len_a: usize,
len_b: usize,
min_len: usize,
prefix_len: usize,
) -> usize {
let mut suffix_len = 0;
let max_suffix = min_len - prefix_len;
while suffix_len + 4 <= max_suffix {
let mut a_buf = [0u32; 4];
let mut b_buf = [0u32; 4];
for i in 0..4 {
a_buf[i] = a[len_a - 1 - suffix_len - (3 - i)] as u32;
b_buf[i] = b[len_b - 1 - suffix_len - (3 - i)] as u32;
}
let a_vec = _mm_loadu_si128(a_buf.as_ptr() as *const __m128i);
let b_vec = _mm_loadu_si128(b_buf.as_ptr() as *const __m128i);
let eq_mask = _mm_cmpeq_epi32(a_vec, b_vec);
let mask = _mm_movemask_ps(_mm_castsi128_ps(eq_mask));
if mask != 0xF {
for i in (0..4).rev() {
if a_buf[i] != b_buf[i] {
return suffix_len + (3 - i);
}
}
}
suffix_len += 4;
}
while suffix_len < max_suffix && a[len_a - 1 - suffix_len] == b[len_b - 1 - suffix_len] {
suffix_len += 1;
}
suffix_len
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
#[cfg(target_arch = "x86_64")]
fn test_simd_basic() {
assert_eq!(standard_distance_simd("", ""), 0);
assert_eq!(standard_distance_simd("abc", ""), 3);
assert_eq!(standard_distance_simd("", "abc"), 3);
assert_eq!(standard_distance_simd("abc", "abc"), 0);
}
#[test]
#[cfg(target_arch = "x86_64")]
fn test_simd_vs_scalar() {
let test_cases = vec![
("kitten", "sitting"),
("saturday", "sunday"),
("book", "back"),
("hello world", "hallo welt"),
];
for (source, target) in test_cases {
let simd_result = standard_distance_simd(source, target);
let scalar_result = crate::distance::standard_distance_impl(source, target);
assert_eq!(
simd_result, scalar_result,
"SIMD and scalar results differ for ('{}', '{}')",
source, target
);
}
}
#[test]
fn test_strip_common_affixes_simd() {
let test_cases = vec![
("", "", (0, 0, 0)),
("abc", "", (0, 3, 0)),
("", "abc", (0, 0, 3)),
("abc", "abc", (3, 0, 0)), ("abcdef", "abc", (3, 3, 0)), ("abc", "abcdef", (3, 0, 3)), ("prefix_middle_suffix", "prefix_other_suffix", (7, 6, 5)), ("hello", "world", (0, 5, 5)), ("abcdefghij", "abcdefghij", (10, 0, 0)), ("test_prefix_abc", "test_prefix_xyz", (12, 3, 3)), ("abc_suffix", "xyz_suffix", (0, 3, 3)), ];
for (a, b, expected) in test_cases {
let simd_result = strip_common_affixes_simd(a, b);
let scalar_result = crate::distance::strip_common_affixes(a, b);
assert_eq!(
simd_result, expected,
"SIMD result incorrect for ('{}', '{}')",
a, b
);
assert_eq!(
simd_result, scalar_result,
"SIMD and scalar results differ for ('{}', '{}')",
a, b
);
}
}
}