use std::sync::OnceLock;
#[cfg(target_arch = "aarch64")]
use std::arch::is_aarch64_feature_detected;
const DEFAULT_SIMD_MIN_LEN: usize = 32;
fn simd_min_len() -> usize {
static CACHE: OnceLock<usize> = OnceLock::new();
*CACHE.get_or_init(|| {
std::env::var("SEMANTEX_SIMD_MIN_LEN")
.ok()
.and_then(|v| v.parse::<usize>().ok())
.filter(|&v| v > 0)
.unwrap_or(DEFAULT_SIMD_MIN_LEN)
})
}
#[inline]
fn assert_same_len(a_len: usize, b_len: usize) {
assert!(
a_len == b_len,
"simd kernel requires equal-length slices (got {a_len} and {b_len})"
);
}
#[inline]
pub fn dot_f32(a: &[f32], b: &[f32]) -> f32 {
assert_same_len(a.len(), b.len());
if a.len() < simd_min_len() {
return scalar::dot_f32(a, b);
}
#[cfg(target_arch = "x86_64")]
{
if is_x86_feature_detected!("avx2") {
return unsafe { avx2::dot_f32(a, b) };
}
}
#[cfg(target_arch = "aarch64")]
{
if is_aarch64_feature_detected!("neon") {
return unsafe { neon::dot_f32(a, b) };
}
}
scalar::dot_f32(a, b)
}
#[inline]
pub fn cosine_f32(a: &[f32], b: &[f32]) -> f32 {
assert_same_len(a.len(), b.len());
if a.len() < simd_min_len() {
return scalar::cosine_f32(a, b);
}
#[cfg(target_arch = "x86_64")]
{
if is_x86_feature_detected!("avx2") {
return unsafe { avx2::cosine_f32(a, b) };
}
}
#[cfg(target_arch = "aarch64")]
{
if is_aarch64_feature_detected!("neon") {
return unsafe { neon::cosine_f32(a, b) };
}
}
scalar::cosine_f32(a, b)
}
#[inline]
pub fn l2_f32(a: &[f32], b: &[f32]) -> f32 {
assert_same_len(a.len(), b.len());
if a.len() < simd_min_len() {
return scalar::l2_f32(a, b);
}
#[cfg(target_arch = "x86_64")]
{
if is_x86_feature_detected!("avx2") {
return unsafe { avx2::l2_f32(a, b) };
}
}
#[cfg(target_arch = "aarch64")]
{
if is_aarch64_feature_detected!("neon") {
return unsafe { neon::l2_f32(a, b) };
}
}
scalar::l2_f32(a, b)
}
#[inline]
pub fn dot_i8(a: &[i8], b: &[i8]) -> f32 {
assert_same_len(a.len(), b.len());
if a.len() < simd_min_len() {
return scalar::dot_i8(a, b);
}
#[cfg(target_arch = "x86_64")]
{
if is_x86_feature_detected!("avx2") {
return unsafe { avx2::dot_i8(a, b) };
}
}
#[cfg(target_arch = "aarch64")]
{
if is_aarch64_feature_detected!("neon") {
return unsafe { neon::dot_i8(a, b) };
}
}
scalar::dot_i8(a, b)
}
#[inline]
pub fn cosine_i8(a: &[i8], b: &[i8]) -> f32 {
assert_same_len(a.len(), b.len());
if a.len() < simd_min_len() {
return scalar::cosine_i8(a, b);
}
#[cfg(target_arch = "x86_64")]
{
if is_x86_feature_detected!("avx2") {
return unsafe { avx2::cosine_i8(a, b) };
}
}
#[cfg(target_arch = "aarch64")]
{
if is_aarch64_feature_detected!("neon") {
return unsafe { neon::cosine_i8(a, b) };
}
}
scalar::cosine_i8(a, b)
}
mod scalar {
#[inline]
pub fn dot_f32(a: &[f32], b: &[f32]) -> f32 {
a.iter().zip(b).map(|(x, y)| x * y).sum()
}
#[inline]
pub fn l2_f32(a: &[f32], b: &[f32]) -> f32 {
a.iter()
.zip(b)
.map(|(x, y)| {
let d = x - y;
d * d
})
.sum::<f32>()
.sqrt()
}
#[inline]
pub fn cosine_f32(a: &[f32], b: &[f32]) -> f32 {
let mut dot = 0.0f32;
let mut na = 0.0f32;
let mut nb = 0.0f32;
for (&x, &y) in a.iter().zip(b) {
dot += x * y;
na += x * x;
nb += y * y;
}
let denom = na.sqrt() * nb.sqrt();
if denom == 0.0 { 0.0 } else { dot / denom }
}
#[inline]
pub fn dot_i8(a: &[i8], b: &[i8]) -> f32 {
let mut acc: i32 = 0;
for (&x, &y) in a.iter().zip(b) {
acc += i32::from(x) * i32::from(y);
}
acc as f32
}
#[inline]
pub fn cosine_i8(a: &[i8], b: &[i8]) -> f32 {
let mut dot: i32 = 0;
let mut na: i32 = 0;
let mut nb: i32 = 0;
for (&x, &y) in a.iter().zip(b) {
let (xi, yi) = (i32::from(x), i32::from(y));
dot += xi * yi;
na += xi * xi;
nb += yi * yi;
}
let denom = (na as f32).sqrt() * (nb as f32).sqrt();
if denom == 0.0 {
0.0
} else {
dot as f32 / denom
}
}
}
#[cfg(target_arch = "x86_64")]
mod avx2 {
#![allow(clippy::similar_names, clippy::many_single_char_names)]
#[allow(clippy::wildcard_imports)]
use std::arch::x86_64::*;
#[target_feature(enable = "avx2")]
unsafe fn hsum256_ps(v: __m256) -> f32 {
let low = _mm256_castps256_ps128(v);
let high = _mm256_extractf128_ps(v, 1);
let sum128 = _mm_add_ps(low, high);
let shuf = _mm_movehdup_ps(sum128); let sums = _mm_add_ps(sum128, shuf); let hi64 = _mm_movehl_ps(shuf, sums); let final_sum = _mm_add_ss(sums, hi64);
_mm_cvtss_f32(final_sum)
}
#[target_feature(enable = "avx2")]
pub unsafe fn dot_f32(a: &[f32], b: &[f32]) -> f32 {
unsafe {
let len = a.len();
let chunks = len & !7; let mut acc = _mm256_setzero_ps();
let mut i = 0;
while i < chunks {
let va = _mm256_loadu_ps(a.as_ptr().add(i));
let vb = _mm256_loadu_ps(b.as_ptr().add(i));
acc = _mm256_fmadd_ps(va, vb, acc);
i += 8;
}
let mut result = hsum256_ps(acc);
while i < len {
result += a[i] * b[i];
i += 1;
}
result
}
}
#[target_feature(enable = "avx2")]
pub unsafe fn l2_f32(a: &[f32], b: &[f32]) -> f32 {
unsafe {
let len = a.len();
let chunks = len & !7;
let mut acc = _mm256_setzero_ps();
let mut i = 0;
while i < chunks {
let va = _mm256_loadu_ps(a.as_ptr().add(i));
let vb = _mm256_loadu_ps(b.as_ptr().add(i));
let diff = _mm256_sub_ps(va, vb);
acc = _mm256_fmadd_ps(diff, diff, acc); i += 8;
}
let mut sumsq = hsum256_ps(acc);
while i < len {
let d = a[i] - b[i];
sumsq += d * d;
i += 1;
}
sumsq.sqrt()
}
}
#[target_feature(enable = "avx2")]
pub unsafe fn cosine_f32(a: &[f32], b: &[f32]) -> f32 {
unsafe {
let len = a.len();
let chunks = len & !7;
let mut dot = _mm256_setzero_ps();
let mut na = _mm256_setzero_ps();
let mut nb = _mm256_setzero_ps();
let mut i = 0;
while i < chunks {
let va = _mm256_loadu_ps(a.as_ptr().add(i));
let vb = _mm256_loadu_ps(b.as_ptr().add(i));
dot = _mm256_fmadd_ps(va, vb, dot);
na = _mm256_fmadd_ps(va, va, na);
nb = _mm256_fmadd_ps(vb, vb, nb);
i += 8;
}
let mut dot_s = hsum256_ps(dot);
let mut na_s = hsum256_ps(na);
let mut nb_s = hsum256_ps(nb);
while i < len {
let (x, y) = (a[i], b[i]);
dot_s += x * y;
na_s += x * x;
nb_s += y * y;
i += 1;
}
let denom = na_s.sqrt() * nb_s.sqrt();
if denom == 0.0 { 0.0 } else { dot_s / denom }
}
}
#[target_feature(enable = "avx2")]
unsafe fn hsum256_epi32(v: __m256i) -> i32 {
let low = _mm256_castsi256_si128(v);
let high = _mm256_extracti128_si256(v, 1);
let sum128 = _mm_add_epi32(low, high);
let hi64 = _mm_unpackhi_epi64(sum128, sum128);
let sum64 = _mm_add_epi32(sum128, hi64);
let hi32 = _mm_shuffle_epi32(sum64, 0b01);
let sum32 = _mm_add_epi32(sum64, hi32);
_mm_cvtsi128_si32(sum32)
}
#[target_feature(enable = "avx2")]
unsafe fn widen_i8_to_i16(v: __m128i) -> __m256i {
_mm256_cvtepi8_epi16(v) }
#[target_feature(enable = "avx2")]
#[allow(clippy::cast_precision_loss)] #[allow(clippy::cast_ptr_alignment)] pub unsafe fn dot_i8(a: &[i8], b: &[i8]) -> f32 {
unsafe {
let len = a.len();
let chunks = len & !15; let mut acc = _mm256_setzero_si256();
let mut i = 0;
while i < chunks {
let va = widen_i8_to_i16(_mm_loadu_si128(a.as_ptr().add(i).cast::<__m128i>()));
let vb = widen_i8_to_i16(_mm_loadu_si128(b.as_ptr().add(i).cast::<__m128i>()));
let prod = _mm256_madd_epi16(va, vb);
acc = _mm256_add_epi32(acc, prod);
i += 16;
}
let mut total = hsum256_epi32(acc);
while i < len {
total += i32::from(a[i]) * i32::from(b[i]);
i += 1;
}
total as f32
}
}
#[target_feature(enable = "avx2")]
#[allow(clippy::cast_precision_loss)] #[allow(clippy::cast_ptr_alignment)] pub unsafe fn cosine_i8(a: &[i8], b: &[i8]) -> f32 {
unsafe {
let len = a.len();
let chunks = len & !15;
let mut dot = _mm256_setzero_si256();
let mut na = _mm256_setzero_si256();
let mut nb = _mm256_setzero_si256();
let mut i = 0;
while i < chunks {
let va = widen_i8_to_i16(_mm_loadu_si128(a.as_ptr().add(i).cast::<__m128i>()));
let vb = widen_i8_to_i16(_mm_loadu_si128(b.as_ptr().add(i).cast::<__m128i>()));
dot = _mm256_add_epi32(dot, _mm256_madd_epi16(va, vb));
na = _mm256_add_epi32(na, _mm256_madd_epi16(va, va));
nb = _mm256_add_epi32(nb, _mm256_madd_epi16(vb, vb));
i += 16;
}
let mut dot_s = hsum256_epi32(dot);
let mut na_s = hsum256_epi32(na);
let mut nb_s = hsum256_epi32(nb);
while i < len {
let (x, y) = (i32::from(a[i]), i32::from(b[i]));
dot_s += x * y;
na_s += x * x;
nb_s += y * y;
i += 1;
}
let denom = (na_s as f32).sqrt() * (nb_s as f32).sqrt();
if denom == 0.0 {
0.0
} else {
dot_s as f32 / denom
}
}
}
}
#[cfg(target_arch = "aarch64")]
mod neon {
#![allow(clippy::similar_names, clippy::many_single_char_names)]
#[allow(clippy::wildcard_imports)]
use std::arch::aarch64::*;
#[target_feature(enable = "neon")]
pub unsafe fn dot_f32(a: &[f32], b: &[f32]) -> f32 {
unsafe {
let len = a.len();
let chunks = len & !3; let mut acc = vdupq_n_f32(0.0);
let mut i = 0;
while i < chunks {
let va = vld1q_f32(a.as_ptr().add(i));
let vb = vld1q_f32(b.as_ptr().add(i));
acc = vfmaq_f32(acc, va, vb); i += 4;
}
let mut result = vaddvq_f32(acc); while i < len {
result += a[i] * b[i];
i += 1;
}
result
}
}
#[target_feature(enable = "neon")]
pub unsafe fn cosine_f32(a: &[f32], b: &[f32]) -> f32 {
unsafe {
let len = a.len();
let chunks = len & !3;
let mut dot = vdupq_n_f32(0.0);
let mut na = vdupq_n_f32(0.0);
let mut nb = vdupq_n_f32(0.0);
let mut i = 0;
while i < chunks {
let va = vld1q_f32(a.as_ptr().add(i));
let vb = vld1q_f32(b.as_ptr().add(i));
dot = vfmaq_f32(dot, va, vb);
na = vfmaq_f32(na, va, va);
nb = vfmaq_f32(nb, vb, vb);
i += 4;
}
let mut dot_s = vaddvq_f32(dot);
let mut na_s = vaddvq_f32(na);
let mut nb_s = vaddvq_f32(nb);
while i < len {
let (x, y) = (a[i], b[i]);
dot_s += x * y;
na_s += x * x;
nb_s += y * y;
i += 1;
}
let denom = na_s.sqrt() * nb_s.sqrt();
if denom == 0.0 { 0.0 } else { dot_s / denom }
}
}
#[target_feature(enable = "neon")]
pub unsafe fn l2_f32(a: &[f32], b: &[f32]) -> f32 {
unsafe {
let len = a.len();
let chunks = len & !3;
let mut acc = vdupq_n_f32(0.0);
let mut i = 0;
while i < chunks {
let va = vld1q_f32(a.as_ptr().add(i));
let vb = vld1q_f32(b.as_ptr().add(i));
let diff = vsubq_f32(va, vb);
acc = vfmaq_f32(acc, diff, diff); i += 4;
}
let mut sumsq = vaddvq_f32(acc);
while i < len {
let d = a[i] - b[i];
sumsq += d * d;
i += 1;
}
sumsq.sqrt()
}
}
#[target_feature(enable = "neon")]
#[allow(clippy::cast_precision_loss)] pub unsafe fn dot_i8(a: &[i8], b: &[i8]) -> f32 {
unsafe {
let len = a.len();
let chunks = len & !7; let mut acc = vdupq_n_s32(0);
let mut i = 0;
while i < chunks {
let va = vld1_s8(a.as_ptr().add(i)); let vb = vld1_s8(b.as_ptr().add(i));
let prod = vmull_s8(va, vb); acc = vaddq_s32(acc, vpaddlq_s16(prod)); i += 8;
}
let mut total = vaddvq_s32(acc); while i < len {
total += i32::from(a[i]) * i32::from(b[i]);
i += 1;
}
total as f32
}
}
#[target_feature(enable = "neon")]
#[allow(clippy::cast_precision_loss)] pub unsafe fn cosine_i8(a: &[i8], b: &[i8]) -> f32 {
unsafe {
let len = a.len();
let chunks = len & !7;
let mut dot = vdupq_n_s32(0);
let mut na = vdupq_n_s32(0);
let mut nb = vdupq_n_s32(0);
let mut i = 0;
while i < chunks {
let va = vld1_s8(a.as_ptr().add(i));
let vb = vld1_s8(b.as_ptr().add(i));
dot = vaddq_s32(dot, vpaddlq_s16(vmull_s8(va, vb)));
na = vaddq_s32(na, vpaddlq_s16(vmull_s8(va, va)));
nb = vaddq_s32(nb, vpaddlq_s16(vmull_s8(vb, vb)));
i += 8;
}
let mut dot_s = vaddvq_s32(dot);
let mut na_s = vaddvq_s32(na);
let mut nb_s = vaddvq_s32(nb);
while i < len {
let (x, y) = (i32::from(a[i]), i32::from(b[i]));
dot_s += x * y;
na_s += x * x;
nb_s += y * y;
i += 1;
}
let denom = (na_s as f32).sqrt() * (nb_s as f32).sqrt();
if denom == 0.0 {
0.0
} else {
dot_s as f32 / denom
}
}
}
}
#[cfg(test)]
mod tests {
#![allow(clippy::similar_names, clippy::many_single_char_names)]
use super::*;
fn assert_close(a: f32, b: f32) {
let diff = (a - b).abs();
let tol = 1e-6_f32 * a.abs().max(b.abs()).max(1.0);
assert!(
diff <= tol,
"values differ beyond 1e-6 tolerance: a={a}, b={b}, diff={diff}, tol={tol}"
);
}
fn make_vec(len: usize, seed: u64) -> Vec<f32> {
let mut s = seed
.wrapping_mul(2_862_933_555_777_941_757)
.wrapping_add(3_037_000_493);
(0..len)
.map(|_| {
s = s
.wrapping_mul(6_364_136_223_846_793_005)
.wrapping_add(1_442_695_040_888_963_407);
((s >> 33) as f32 / (1u64 << 31) as f32) - 1.0
})
.collect()
}
#[test]
fn scalar_dot_matches_hand_computation() {
let a = [1.0f32, 2.0, 3.0];
let b = [4.0f32, 5.0, 6.0];
assert_close(scalar::dot_f32(&a, &b), 32.0);
}
#[test]
fn scalar_l2_matches_hand_computation() {
let a = [0.0f32, 0.0];
let b = [3.0f32, 4.0];
assert_close(scalar::l2_f32(&a, &b), 5.0);
}
#[test]
fn scalar_cosine_orthogonal_is_zero_and_parallel_is_one() {
assert_close(scalar::cosine_f32(&[1.0, 0.0], &[0.0, 1.0]), 0.0);
assert_close(scalar::cosine_f32(&[1.0, 2.0, 3.0], &[2.0, 4.0, 6.0]), 1.0);
}
#[test]
fn cosine_zero_norm_returns_zero() {
assert_close(cosine_f32(&[0.0; 768], &make_vec(768, 1)), 0.0);
assert_close(scalar::cosine_i8(&[0i8; 768], &[1i8; 768]), 0.0);
}
#[test]
fn scalar_i8_dot_matches_hand_computation() {
let a = [1i8, -2, 3];
let b = [4i8, 5, -6];
assert_close(scalar::dot_i8(&a, &b), -24.0);
}
#[test]
fn dispatch_matches_scalar_across_sizes() {
for &len in &[1usize, 3, 7, 8, 16, 31, 32, 33, 100, 768, 769] {
let a = make_vec(len, 0xABCD ^ len as u64);
let b = make_vec(len, 0x1234 ^ len as u64);
assert_close(dot_f32(&a, &b), scalar::dot_f32(&a, &b));
assert_close(cosine_f32(&a, &b), scalar::cosine_f32(&a, &b));
assert_close(l2_f32(&a, &b), scalar::l2_f32(&a, &b));
let ai: Vec<i8> = a.iter().map(|x| (x * 100.0) as i8).collect();
let bi: Vec<i8> = b.iter().map(|x| (x * 100.0) as i8).collect();
assert_close(dot_i8(&ai, &bi), scalar::dot_i8(&ai, &bi));
assert_close(cosine_i8(&ai, &bi), scalar::cosine_i8(&ai, &bi));
}
}
#[cfg(target_arch = "x86_64")]
#[test]
fn avx2_dot_l2_cosine_match_scalar() {
if !is_x86_feature_detected!("avx2") {
eprintln!("skipping avx2 parity: AVX2 not present on this host");
return;
}
for &len in &[8usize, 15, 16, 17, 64, 768, 769] {
let a = make_vec(len, 0x55 ^ len as u64);
let b = make_vec(len, 0xAA ^ len as u64);
let (d, l, c) = unsafe {
(
avx2::dot_f32(&a, &b),
avx2::l2_f32(&a, &b),
avx2::cosine_f32(&a, &b),
)
};
assert_close(d, scalar::dot_f32(&a, &b));
assert_close(l, scalar::l2_f32(&a, &b));
assert_close(c, scalar::cosine_f32(&a, &b));
}
}
#[cfg(target_arch = "aarch64")]
#[test]
fn neon_dot_l2_cosine_match_scalar() {
for &len in &[4usize, 7, 8, 9, 64, 768, 769] {
let a = make_vec(len, 0x55 ^ len as u64);
let b = make_vec(len, 0xAA ^ len as u64);
let (d, l, c) = unsafe {
(
neon::dot_f32(&a, &b),
neon::l2_f32(&a, &b),
neon::cosine_f32(&a, &b),
)
};
assert_close(d, scalar::dot_f32(&a, &b));
assert_close(l, scalar::l2_f32(&a, &b));
assert_close(c, scalar::cosine_f32(&a, &b));
}
}
fn i8_vecs(len: usize, seed: u64) -> (Vec<i8>, Vec<i8>) {
let af = make_vec(len, seed);
let bf = make_vec(len, seed ^ 0xFFFF);
(
af.iter().map(|x| (x * 127.0) as i8).collect(),
bf.iter().map(|x| (x * 127.0) as i8).collect(),
)
}
#[cfg(target_arch = "x86_64")]
#[test]
fn avx2_i8_match_scalar() {
if !is_x86_feature_detected!("avx2") {
eprintln!("skipping avx2 i8 parity: AVX2 not present");
return;
}
for &len in &[16usize, 17, 31, 32, 768, 769] {
let (a, b) = i8_vecs(len, 0x9 ^ len as u64);
let (d, c) = unsafe { (avx2::dot_i8(&a, &b), avx2::cosine_i8(&a, &b)) };
assert_eq!(d, scalar::dot_i8(&a, &b), "i8 dot must be integer-exact");
assert_close(c, scalar::cosine_i8(&a, &b));
}
}
#[cfg(target_arch = "aarch64")]
#[test]
fn neon_i8_match_scalar() {
for &len in &[8usize, 9, 15, 16, 768, 769] {
let (a, b) = i8_vecs(len, 0x9 ^ len as u64);
let (d, c) = unsafe { (neon::dot_i8(&a, &b), neon::cosine_i8(&a, &b)) };
assert_eq!(d, scalar::dot_i8(&a, &b), "i8 dot must be integer-exact");
assert_close(c, scalar::cosine_i8(&a, &b));
}
}
#[test]
#[should_panic(expected = "equal-length")]
fn mismatched_lengths_panic() {
let _ = dot_f32(&[1.0, 2.0], &[1.0]);
}
}