use pulp::Simd;
#[inline]
pub fn dot_product_simd(a: &[f32], b: &[f32]) -> f32 {
assert_eq!(a.len(), b.len(), "Vector dimensions must match");
let simd = pulp::Arch::new();
simd.dispatch(|| dot_product_simd_impl(simd, a, b))
}
#[inline]
pub fn l2_distance_simd(a: &[f32], b: &[f32]) -> f32 {
assert_eq!(a.len(), b.len(), "Vector dimensions must match");
let simd = pulp::Arch::new();
simd.dispatch(|| l2_distance_simd_impl(simd, a, b))
}
#[inline]
pub fn l2_squared_distance_simd(a: &[f32], b: &[f32]) -> f32 {
assert_eq!(a.len(), b.len(), "Vector dimensions must match");
let simd = pulp::Arch::new();
simd.dispatch(|| l2_squared_distance_simd_impl(simd, a, b))
}
#[inline]
pub fn cosine_similarity_simd(a: &[f32], b: &[f32]) -> f32 {
assert_eq!(a.len(), b.len(), "Vector dimensions must match");
if a.is_empty() {
return 1.0; }
let simd = pulp::Arch::new();
simd.dispatch(|| {
let dot = dot_product_simd_impl(simd, a, b);
let norm_a = magnitude_simd_impl(simd, a);
let norm_b = magnitude_simd_impl(simd, b);
if norm_a == 0.0 || norm_b == 0.0 {
return 0.0;
}
let similarity = dot / (norm_a * norm_b);
similarity.clamp(-1.0, 1.0)
})
}
#[inline]
pub fn norm_simd(v: &[f32]) -> f32 {
let simd = pulp::Arch::new();
simd.dispatch(Magnitude { vector: v })
}
#[inline]
pub fn cosine_distance_prenorm(a: &[f32], b: &[f32], norm_b: f32) -> f32 {
debug_assert_eq!(a.len(), b.len());
if norm_b == 0.0 {
return 1.0;
}
let simd = pulp::Arch::new();
simd.dispatch(FusedCosineDistance { a, b, norm_b })
}
struct FusedCosineDistance<'a> {
a: &'a [f32],
b: &'a [f32],
norm_b: f32,
}
impl pulp::WithSimd for FusedCosineDistance<'_> {
type Output = f32;
#[inline(always)]
fn with_simd<S: Simd>(self, simd: S) -> Self::Output {
let a = self.a;
let b = self.b;
let norm_b = self.norm_b;
let (a_chunks, a_tail) = S::as_simd_f32s(a);
let (b_chunks, b_tail) = S::as_simd_f32s(b);
let mut dot_acc = simd.splat_f32s(0.0);
let mut norm_a_acc = simd.splat_f32s(0.0);
for (&a_vec, &b_vec) in a_chunks.iter().zip(b_chunks.iter()) {
dot_acc = simd.mul_add_e_f32s(a_vec, b_vec, dot_acc);
norm_a_acc = simd.mul_add_e_f32s(a_vec, a_vec, norm_a_acc);
}
let mut dot = simd.reduce_sum_f32s(dot_acc);
let mut norm_a_sq = simd.reduce_sum_f32s(norm_a_acc);
debug_assert_eq!(a_tail.len(), b_tail.len());
for (&a_scalar, &b_scalar) in a_tail.iter().zip(b_tail.iter()) {
dot += a_scalar * b_scalar;
norm_a_sq += a_scalar * a_scalar;
}
let norm_a = norm_a_sq.sqrt();
if norm_a == 0.0 {
return 1.0;
}
let similarity = dot / (norm_a * norm_b);
1.0 - similarity.clamp(-1.0, 1.0)
}
}
#[inline(always)]
fn dot_product_simd_impl(simd: pulp::Arch, a: &[f32], b: &[f32]) -> f32 {
struct DotProduct<'a> {
a: &'a [f32],
b: &'a [f32],
}
impl pulp::WithSimd for DotProduct<'_> {
type Output = f32;
#[inline(always)]
fn with_simd<S: Simd>(self, simd: S) -> Self::Output {
let a = self.a;
let b = self.b;
let (a_chunks, a_tail) = S::as_simd_f32s(a);
let (b_chunks, b_tail) = S::as_simd_f32s(b);
let mut sum = simd.splat_f32s(0.0);
for (&a_vec, &b_vec) in a_chunks.iter().zip(b_chunks.iter()) {
sum = simd.mul_add_e_f32s(a_vec, b_vec, sum);
}
let mut result = simd.reduce_sum_f32s(sum);
debug_assert_eq!(a_tail.len(), b_tail.len());
for (&a_scalar, &b_scalar) in a_tail.iter().zip(b_tail.iter()) {
result += a_scalar * b_scalar;
}
result
}
}
simd.dispatch(DotProduct { a, b })
}
#[inline(always)]
fn l2_distance_simd_impl(simd: pulp::Arch, a: &[f32], b: &[f32]) -> f32 {
struct L2Distance<'a> {
a: &'a [f32],
b: &'a [f32],
}
impl pulp::WithSimd for L2Distance<'_> {
type Output = f32;
#[inline(always)]
fn with_simd<S: Simd>(self, simd: S) -> Self::Output {
let a = self.a;
let b = self.b;
let (a_chunks, a_tail) = S::as_simd_f32s(a);
let (b_chunks, b_tail) = S::as_simd_f32s(b);
let mut sum_squares = simd.splat_f32s(0.0);
for (&a_vec, &b_vec) in a_chunks.iter().zip(b_chunks.iter()) {
let diff = simd.sub_f32s(a_vec, b_vec);
sum_squares = simd.mul_add_e_f32s(diff, diff, sum_squares);
}
let mut result = simd.reduce_sum_f32s(sum_squares);
debug_assert_eq!(a_tail.len(), b_tail.len());
for (&a_scalar, &b_scalar) in a_tail.iter().zip(b_tail.iter()) {
let diff = a_scalar - b_scalar;
result += diff * diff;
}
result.sqrt()
}
}
simd.dispatch(L2Distance { a, b })
}
#[inline(always)]
fn l2_squared_distance_simd_impl(simd: pulp::Arch, a: &[f32], b: &[f32]) -> f32 {
struct L2Squared<'a> {
a: &'a [f32],
b: &'a [f32],
}
impl pulp::WithSimd for L2Squared<'_> {
type Output = f32;
#[inline(always)]
fn with_simd<S: Simd>(self, simd: S) -> Self::Output {
let a = self.a;
let b = self.b;
let (a_chunks, a_tail) = S::as_simd_f32s(a);
let (b_chunks, b_tail) = S::as_simd_f32s(b);
let mut sum_squares = simd.splat_f32s(0.0);
for (&a_vec, &b_vec) in a_chunks.iter().zip(b_chunks.iter()) {
let diff = simd.sub_f32s(a_vec, b_vec);
sum_squares = simd.mul_add_e_f32s(diff, diff, sum_squares);
}
let mut result = simd.reduce_sum_f32s(sum_squares);
debug_assert_eq!(a_tail.len(), b_tail.len());
for (&a_scalar, &b_scalar) in a_tail.iter().zip(b_tail.iter()) {
let diff = a_scalar - b_scalar;
result += diff * diff;
}
result }
}
simd.dispatch(L2Squared { a, b })
}
struct Magnitude<'a> {
vector: &'a [f32],
}
impl pulp::WithSimd for Magnitude<'_> {
type Output = f32;
#[inline(always)]
fn with_simd<S: Simd>(self, simd: S) -> Self::Output {
let vector = self.vector;
let (chunks, tail) = S::as_simd_f32s(vector);
let mut sum_squares = simd.splat_f32s(0.0);
for &vector_chunk in chunks {
sum_squares = simd.mul_add_e_f32s(vector_chunk, vector_chunk, sum_squares);
}
let mut result = simd.reduce_sum_f32s(sum_squares);
for &value in tail {
result += value * value;
}
result.sqrt()
}
}
#[inline(always)]
fn magnitude_simd_impl(simd: pulp::Arch, vector: &[f32]) -> f32 {
simd.dispatch(Magnitude { vector })
}
#[cfg(test)]
mod tests {
use super::*;
use crate::vector::ops::{cosine_similarity, dot_product, l2_distance};
const EPSILON: f32 = 1e-5;
#[test]
fn test_dot_product_simd_basic() {
let a = vec![1.0, 2.0, 3.0];
let b = vec![4.0, 5.0, 6.0];
let result = dot_product_simd(&a, &b);
let expected = dot_product(&a, &b).unwrap();
assert!((result - expected).abs() < EPSILON);
assert!((result - 32.0).abs() < EPSILON);
}
#[test]
fn test_dot_product_simd_zero() {
let a = vec![1.0, 0.0, 0.0];
let b = vec![0.0, 1.0, 0.0];
let result = dot_product_simd(&a, &b);
assert!((result - 0.0).abs() < EPSILON);
}
#[test]
fn test_dot_product_simd_negative() {
let a = vec![1.0, 2.0, 3.0];
let b = vec![-1.0, -2.0, -3.0];
let result = dot_product_simd(&a, &b);
let expected = dot_product(&a, &b).unwrap();
assert!((result - expected).abs() < EPSILON);
}
#[test]
fn test_dot_product_simd_various_sizes() {
for size in [
1, 2, 3, 4, 5, 7, 8, 15, 16, 31, 32, 63, 64, 127, 128, 383, 384, 767, 768,
] {
let a: Vec<f32> = (0..size).map(|i| i as f32).collect();
let b: Vec<f32> = (0..size).map(|i| (i * 2) as f32).collect();
let simd_result = dot_product_simd(&a, &b);
let scalar_result = dot_product(&a, &b).unwrap();
let epsilon = if scalar_result.abs() > 1000.0 {
scalar_result.abs() * 1e-5 } else {
EPSILON
};
assert!(
(simd_result - scalar_result).abs() < epsilon,
"Size {}: SIMD={}, Scalar={}",
size,
simd_result,
scalar_result
);
}
}
#[test]
fn test_dot_product_simd_misaligned_subslice_regression() {
let size = 257;
let a_storage: Vec<f32> = (0..(size + 3))
.map(|i| ((i as f32) - 90.0) * 0.03125)
.collect();
let b_storage: Vec<f32> = (0..(size + 4))
.map(|i| ((size + 4 - i) as f32 - 120.0) * 0.0625)
.collect();
let a = &a_storage[1..(1 + size)];
let b = &b_storage[2..(2 + size)];
let simd_result = dot_product_simd(a, b);
let scalar_result = dot_product(a, b).unwrap();
assert!((simd_result - scalar_result).abs() < 1e-4);
}
#[test]
fn test_l2_distance_simd_basic() {
let a = vec![0.0, 0.0];
let b = vec![3.0, 4.0];
let result = l2_distance_simd(&a, &b);
let expected = l2_distance(&a, &b).unwrap();
assert!((result - expected).abs() < EPSILON);
assert!((result - 5.0).abs() < EPSILON);
}
#[test]
fn test_l2_distance_simd_zero() {
let a = vec![1.0, 2.0, 3.0];
let b = vec![1.0, 2.0, 3.0];
let result = l2_distance_simd(&a, &b);
assert!(result < EPSILON);
}
#[test]
fn test_l2_distance_simd_various_sizes() {
for size in [
1, 2, 3, 4, 5, 7, 8, 15, 16, 31, 32, 63, 64, 127, 128, 383, 384, 767, 768,
] {
let a: Vec<f32> = (0..size).map(|i| i as f32).collect();
let b: Vec<f32> = (0..size).map(|i| (i * 2) as f32).collect();
let simd_result = l2_distance_simd(&a, &b);
let scalar_result = l2_distance(&a, &b).unwrap();
let epsilon = if scalar_result.abs() > 1000.0 {
scalar_result.abs() * 1e-5 } else {
EPSILON
};
assert!(
(simd_result - scalar_result).abs() < epsilon,
"Size {}: SIMD={}, Scalar={}",
size,
simd_result,
scalar_result
);
}
}
#[test]
fn test_l2_distance_simd_misaligned_subslice_regression() {
let size = 257;
let a_storage: Vec<f32> = (0..(size + 3))
.map(|i| ((i as f32) - 30.0) * 0.125)
.collect();
let b_storage: Vec<f32> = (0..(size + 4))
.map(|i| ((i as f32) - 170.0) * -0.09375)
.collect();
let a = &a_storage[1..(1 + size)];
let b = &b_storage[2..(2 + size)];
let simd_result = l2_distance_simd(a, b);
let scalar_result = l2_distance(a, b).unwrap();
assert!((simd_result - scalar_result).abs() < 1e-4);
}
#[test]
fn test_cosine_similarity_simd_basic() {
let a = vec![1.0, 0.0, 0.0];
let b = vec![0.0, 1.0, 0.0];
let result = cosine_similarity_simd(&a, &b);
let expected = cosine_similarity(&a, &b).unwrap();
assert!((result - expected).abs() < EPSILON);
assert!((result - 0.0).abs() < EPSILON);
}
#[test]
fn test_cosine_similarity_simd_identical() {
let a = vec![1.0, 2.0, 3.0];
let b = vec![1.0, 2.0, 3.0];
let result = cosine_similarity_simd(&a, &b);
assert!((result - 1.0).abs() < EPSILON);
}
#[test]
fn test_cosine_similarity_simd_opposite() {
let a = vec![1.0, 2.0, 3.0];
let b = vec![-1.0, -2.0, -3.0];
let result = cosine_similarity_simd(&a, &b);
assert!((result - (-1.0)).abs() < EPSILON);
}
#[test]
fn test_cosine_similarity_simd_zero_vector() {
let a = vec![0.0, 0.0, 0.0];
let b = vec![1.0, 2.0, 3.0];
let result = cosine_similarity_simd(&a, &b);
assert!((result - 0.0).abs() < EPSILON);
}
#[test]
fn test_cosine_similarity_simd_various_sizes() {
for size in [
1, 2, 3, 4, 5, 7, 8, 15, 16, 31, 32, 63, 64, 127, 128, 383, 384, 767, 768, 1023, 1024,
] {
let a: Vec<f32> = (0..size).map(|i| (i as f32) / (size as f32)).collect();
let b: Vec<f32> = (0..size)
.map(|i| 1.0 - (i as f32) / (size as f32))
.collect();
let simd_result = cosine_similarity_simd(&a, &b);
let scalar_result = cosine_similarity(&a, &b).unwrap();
let epsilon = if size > 100 { 1e-4 } else { EPSILON };
assert!(
(simd_result - scalar_result).abs() < epsilon,
"Size {}: SIMD={}, Scalar={}",
size,
simd_result,
scalar_result
);
}
}
#[test]
fn test_cosine_similarity_simd_misaligned_subslice_regression() {
let size = 257;
let a_storage: Vec<f32> = (0..(size + 3))
.map(|i| (((i as f32) % 17.0) - 8.0) * 0.37)
.collect();
let b_storage: Vec<f32> = (0..(size + 4))
.map(|i| (((i as f32) % 19.0) - 9.0) * -0.29)
.collect();
let a = &a_storage[1..(1 + size)];
let b = &b_storage[2..(2 + size)];
let simd_result = cosine_similarity_simd(a, b);
let scalar_result = cosine_similarity(a, b).unwrap();
assert!((simd_result - scalar_result).abs() < 1e-4);
}
#[test]
fn test_simd_numerical_stability() {
let a = vec![1e6; 384];
let b = vec![2e6; 384];
let simd_result = cosine_similarity_simd(&a, &b);
let scalar_result = cosine_similarity(&a, &b).unwrap();
assert!((simd_result - scalar_result).abs() < EPSILON);
assert!((-1.0..=1.0).contains(&simd_result));
let a = vec![1e-6; 384];
let b = vec![2e-6; 384];
let simd_result = cosine_similarity_simd(&a, &b);
let scalar_result = cosine_similarity(&a, &b).unwrap();
assert!((simd_result - scalar_result).abs() < EPSILON);
assert!((-1.0..=1.0).contains(&simd_result));
}
#[test]
#[should_panic(expected = "Vector dimensions must match")]
fn test_dot_product_simd_dimension_mismatch() {
let a = vec![1.0, 2.0];
let b = vec![1.0, 2.0, 3.0];
let _ = dot_product_simd(&a, &b);
}
#[test]
#[should_panic(expected = "Vector dimensions must match")]
fn test_l2_distance_simd_dimension_mismatch() {
let a = vec![1.0, 2.0];
let b = vec![1.0, 2.0, 3.0];
let _ = l2_distance_simd(&a, &b);
}
#[test]
#[should_panic(expected = "Vector dimensions must match")]
fn test_cosine_similarity_simd_dimension_mismatch() {
let a = vec![1.0, 2.0];
let b = vec![1.0, 2.0, 3.0];
let _ = cosine_similarity_simd(&a, &b);
}
#[test]
fn test_norm_simd() {
let v = vec![3.0, 4.0];
let norm = norm_simd(&v);
assert!((norm - 5.0).abs() < EPSILON);
}
#[test]
fn test_cosine_distance_prenorm_matches_old() {
for size in [3, 4, 8, 16, 32, 64, 128, 384, 768] {
let a: Vec<f32> = (0..size).map(|i| (i as f32) / (size as f32)).collect();
let b: Vec<f32> = (0..size)
.map(|i| 1.0 - (i as f32) / (size as f32))
.collect();
let old_dist = 1.0 - cosine_similarity_simd(&a, &b);
let norm_b = norm_simd(&b);
let new_dist = cosine_distance_prenorm(&a, &b, norm_b);
let epsilon = if size > 100 { 1e-4 } else { EPSILON };
assert!(
(old_dist - new_dist).abs() < epsilon,
"Size {}: old={}, new={}",
size,
old_dist,
new_dist
);
}
}
#[test]
fn test_cosine_distance_prenorm_misaligned_subslice_regression() {
let size = 257;
let a_storage: Vec<f32> = (0..(size + 3))
.map(|i| (((i as f32) % 13.0) - 6.0) * 0.41)
.collect();
let b_storage: Vec<f32> = (0..(size + 4))
.map(|i| (((i as f32) % 11.0) - 5.0) * -0.23)
.collect();
let a = &a_storage[1..(1 + size)];
let b = &b_storage[2..(2 + size)];
let norm_b = norm_simd(b);
let simd_result = cosine_distance_prenorm(a, b, norm_b);
let scalar_result = 1.0 - cosine_similarity(a, b).unwrap();
assert!((simd_result - scalar_result).abs() < 1e-4);
}
#[test]
fn test_cosine_distance_prenorm_zero_vectors() {
let a = vec![0.0, 0.0, 0.0];
let b = vec![1.0, 2.0, 3.0];
let norm_b = norm_simd(&b);
assert!((cosine_distance_prenorm(&a, &b, norm_b) - 1.0).abs() < EPSILON);
assert!((cosine_distance_prenorm(&b, &a, 0.0) - 1.0).abs() < EPSILON);
}
#[test]
fn test_cosine_distance_prenorm_identical() {
let a = vec![1.0, 2.0, 3.0];
let norm_a = norm_simd(&a);
let dist = cosine_distance_prenorm(&a, &a, norm_a);
assert!(
dist.abs() < EPSILON,
"Identical vectors should have distance ~0, got {}",
dist
);
}
#[test]
fn test_cosine_distance_prenorm_opposite() {
let a = vec![1.0, 2.0, 3.0];
let b: Vec<f32> = a.iter().map(|x| -x).collect();
let norm_b = norm_simd(&b);
let dist = cosine_distance_prenorm(&a, &b, norm_b);
assert!(
(dist - 2.0).abs() < EPSILON,
"Opposite vectors should have distance ~2, got {}",
dist
);
}
}
#[inline]
pub fn sq8_asymmetric_l2_simd(query: &[f32], codes: &[u8], min: &[f32], scale: &[f32]) -> f32 {
debug_assert_eq!(query.len(), codes.len());
debug_assert_eq!(query.len(), min.len());
debug_assert_eq!(query.len(), scale.len());
#[cfg(target_arch = "x86_64")]
{
if is_x86_feature_detected!("avx2") && is_x86_feature_detected!("fma") {
return unsafe { sq8_asymmetric_l2_avx2(query, codes, min, scale) };
}
}
sq8_asymmetric_l2_scalar(query, codes, min, scale)
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2", enable = "fma")]
unsafe fn sq8_asymmetric_l2_avx2(query: &[f32], codes: &[u8], min: &[f32], scale: &[f32]) -> f32 {
use std::arch::x86_64::*;
let n = codes.len();
let mut acc = _mm256_setzero_ps();
let mut i = 0;
while i + 8 <= n {
let c8 = _mm_loadl_epi64(codes.as_ptr().add(i) as *const __m128i);
let c32 = _mm256_cvtepu8_epi32(c8);
let cf = _mm256_cvtepi32_ps(c32);
let s = _mm256_loadu_ps(scale.as_ptr().add(i));
let m = _mm256_loadu_ps(min.as_ptr().add(i));
let q = _mm256_loadu_ps(query.as_ptr().add(i));
let deq = _mm256_fmadd_ps(cf, s, m);
let d = _mm256_sub_ps(q, deq);
acc = _mm256_fmadd_ps(d, d, acc);
i += 8;
}
let hi = _mm256_extractf128_ps(acc, 1);
let lo = _mm256_castps256_ps128(acc);
let mut sum128 = _mm_add_ps(hi, lo);
sum128 = _mm_hadd_ps(sum128, sum128);
sum128 = _mm_hadd_ps(sum128, sum128);
let mut total = _mm_cvtss_f32(sum128);
for j in i..n {
let deq = min[j] + codes[j] as f32 * scale[j];
let d = query[j] - deq;
total += d * d;
}
total
}
fn sq8_asymmetric_l2_scalar(query: &[f32], codes: &[u8], min: &[f32], scale: &[f32]) -> f32 {
let mut acc = 0.0f32;
for i in 0..codes.len() {
let deq = min[i] + codes[i] as f32 * scale[i];
let d = query[i] - deq;
acc += d * d;
}
acc
}
#[inline]
pub fn sq8_asymmetric_dot_simd(query: &[f32], codes: &[u8], min: &[f32], scale: &[f32]) -> f32 {
debug_assert_eq!(query.len(), codes.len());
debug_assert_eq!(query.len(), min.len());
debug_assert_eq!(query.len(), scale.len());
#[cfg(target_arch = "x86_64")]
{
if is_x86_feature_detected!("avx2") && is_x86_feature_detected!("fma") {
return unsafe { sq8_asymmetric_dot_avx2(query, codes, min, scale) };
}
}
sq8_asymmetric_dot_scalar(query, codes, min, scale)
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2", enable = "fma")]
unsafe fn sq8_asymmetric_dot_avx2(query: &[f32], codes: &[u8], min: &[f32], scale: &[f32]) -> f32 {
use std::arch::x86_64::*;
let n = codes.len();
let mut acc = _mm256_setzero_ps();
let mut i = 0;
while i + 8 <= n {
let c8 = _mm_loadl_epi64(codes.as_ptr().add(i) as *const __m128i);
let c32 = _mm256_cvtepu8_epi32(c8);
let cf = _mm256_cvtepi32_ps(c32);
let s = _mm256_loadu_ps(scale.as_ptr().add(i));
let m = _mm256_loadu_ps(min.as_ptr().add(i));
let q = _mm256_loadu_ps(query.as_ptr().add(i));
let deq = _mm256_fmadd_ps(cf, s, m);
acc = _mm256_fmadd_ps(q, deq, acc);
i += 8;
}
let hi = _mm256_extractf128_ps(acc, 1);
let lo = _mm256_castps256_ps128(acc);
let mut sum128 = _mm_add_ps(hi, lo);
sum128 = _mm_hadd_ps(sum128, sum128);
sum128 = _mm_hadd_ps(sum128, sum128);
let mut total = _mm_cvtss_f32(sum128);
for j in i..n {
let deq = min[j] + codes[j] as f32 * scale[j];
total += query[j] * deq;
}
total
}
fn sq8_asymmetric_dot_scalar(query: &[f32], codes: &[u8], min: &[f32], scale: &[f32]) -> f32 {
let mut acc = 0.0f32;
for i in 0..codes.len() {
let deq = min[i] + codes[i] as f32 * scale[i];
acc += query[i] * deq;
}
acc
}
#[inline]
pub fn rabitq_asymmetric_l2_simd(
rq: &[f32],
bits: &[u8],
dtc_sq: f32,
est_factor: f32,
qn_sq: f32,
) -> f32 {
debug_assert_eq!(bits.len(), rq.len().div_ceil(8));
#[cfg(target_arch = "x86_64")]
{
if is_x86_feature_detected!("avx2") && is_x86_feature_detected!("fma") {
let s = unsafe { rabitq_signed_sum_avx2(rq, bits) };
return (dtc_sq + qn_sq - 2.0 * est_factor * s).max(0.0);
}
}
let s = rabitq_signed_sum_scalar(rq, bits);
(dtc_sq + qn_sq - 2.0 * est_factor * s).max(0.0)
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2", enable = "fma")]
unsafe fn rabitq_signed_sum_avx2(rq: &[f32], bits: &[u8]) -> f32 {
use std::arch::x86_64::*;
let n = rq.len();
let mut acc = _mm256_setzero_ps();
let mut i = 0;
let bit_masks = _mm256_set_epi32(128, 64, 32, 16, 8, 4, 2, 1);
let zero = _mm256_setzero_si256();
let ones = _mm256_set1_ps(1.0);
let neg_ones = _mm256_set1_ps(-1.0);
while i + 8 <= n {
let byte = bits[i / 8] as i32;
let byte_bcast = _mm256_set1_epi32(byte);
let anded = _mm256_and_si256(byte_bcast, bit_masks);
let is_set = _mm256_cmpgt_epi32(anded, zero);
let signs = _mm256_blendv_ps(neg_ones, ones, _mm256_castsi256_ps(is_set));
let rq_vec = _mm256_loadu_ps(rq.as_ptr().add(i));
acc = _mm256_fmadd_ps(rq_vec, signs, acc);
i += 8;
}
let hi = _mm256_extractf128_ps(acc, 1);
let lo = _mm256_castps256_ps128(acc);
let mut sum128 = _mm_add_ps(hi, lo);
sum128 = _mm_hadd_ps(sum128, sum128);
sum128 = _mm_hadd_ps(sum128, sum128);
let mut total = _mm_cvtss_f32(sum128);
for j in i..n {
let bit = (bits[j / 8] >> (j % 8)) & 1;
total += if bit == 1 { rq[j] } else { -rq[j] };
}
total
}
fn rabitq_signed_sum_scalar(rq: &[f32], bits: &[u8]) -> f32 {
let mut s = 0.0f32;
for (i, &rqi) in rq.iter().enumerate() {
let bit = (bits[i / 8] >> (i % 8)) & 1;
s += if bit == 1 { rqi } else { -rqi };
}
s
}
#[cfg(test)]
mod rabitq_asymmetric_tests {
use super::*;
#[test]
fn avx2_matches_scalar() {
let dim: usize = 131; let rq: Vec<f32> = (0..dim).map(|i| (i as f32 * 0.29).cos() * 7.0).collect();
let n_bytes = dim.div_ceil(8);
let bits: Vec<u8> = (0..n_bytes).map(|i| ((i * 37 + 11) % 256) as u8).collect();
let dtc_sq = 4.5f32;
let est_factor = 0.83f32;
let qn_sq = 6.1f32;
let dispatched = rabitq_asymmetric_l2_simd(&rq, &bits, dtc_sq, est_factor, qn_sq);
let s = rabitq_signed_sum_scalar(&rq, &bits);
let scalar = (dtc_sq + qn_sq - 2.0 * est_factor * s).max(0.0);
let rel = (dispatched - scalar).abs() / scalar.abs().max(1e-6);
assert!(
rel < 1e-5,
"AVX2 kernel disagrees with scalar: {dispatched} vs {scalar} (rel {rel:.2e})"
);
}
#[test]
fn matches_rabitq_module_reference_estimator() {
use crate::vector::rabitq::RaBitQuantizer;
use rand::{rngs::StdRng, RngExt, SeedableRng};
let mut rng = StdRng::seed_from_u64(21);
let dim = 97; let train: Vec<Vec<f32>> = (0..300)
.map(|_| (0..dim).map(|_| rng.random::<f32>() * 2.0 - 1.0).collect())
.collect();
let q = RaBitQuantizer::fit(&train);
let o: Vec<f32> = (0..dim).map(|_| rng.random::<f32>() * 2.0 - 1.0).collect();
let query: Vec<f32> = (0..dim).map(|_| rng.random::<f32>() * 2.0 - 1.0).collect();
let code = q.encode(&o);
let prepared = q.prepare_query(&query);
let reference = q.estimate_dist_sq(&prepared, &code);
let kernel = rabitq_asymmetric_l2_simd(
prepared.rq(),
&code.bits,
code.dtc_sq,
code.est_factor,
prepared.qn_sq(),
);
let rel = (kernel - reference).abs() / reference.abs().max(1e-6);
assert!(
rel < 1e-4,
"kernel disagrees with rabitq.rs reference estimator: {kernel} vs {reference} (rel {rel:.2e})"
);
}
}
#[cfg(test)]
mod sq8_asymmetric_tests {
use super::*;
#[test]
fn avx2_matches_scalar() {
let dim = 131; let query: Vec<f32> = (0..dim).map(|i| (i as f32 * 0.37).sin() * 12.0).collect();
let codes: Vec<u8> = (0..dim).map(|i| ((i * 7 + 3) % 256) as u8).collect();
let min: Vec<f32> = (0..dim).map(|i| -(i as f32) * 0.11).collect();
let scale: Vec<f32> = (0..dim).map(|i| 0.01 + (i % 5) as f32 * 0.03).collect();
let dispatched = sq8_asymmetric_l2_simd(&query, &codes, &min, &scale);
let scalar = sq8_asymmetric_l2_scalar(&query, &codes, &min, &scale);
let rel = (dispatched - scalar).abs() / scalar.abs().max(1e-6);
assert!(
rel < 1e-5,
"AVX2 kernel disagrees with scalar: {dispatched} vs {scalar} (rel {rel:.2e})"
);
}
#[test]
fn dot_avx2_matches_scalar() {
let dim = 131;
let query: Vec<f32> = (0..dim).map(|i| (i as f32 * 0.37).sin() * 12.0).collect();
let codes: Vec<u8> = (0..dim).map(|i| ((i * 7 + 3) % 256) as u8).collect();
let min: Vec<f32> = (0..dim).map(|i| -(i as f32) * 0.11).collect();
let scale: Vec<f32> = (0..dim).map(|i| 0.01 + (i % 5) as f32 * 0.03).collect();
let dispatched = sq8_asymmetric_dot_simd(&query, &codes, &min, &scale);
let scalar = sq8_asymmetric_dot_scalar(&query, &codes, &min, &scale);
let rel = (dispatched - scalar).abs() / scalar.abs().max(1e-6);
assert!(
rel < 1e-5,
"AVX2 dot kernel disagrees with scalar: {dispatched} vs {scalar} (rel {rel:.2e})"
);
}
}