use super::distance::{estimate_distance, QueryContext};
use crate::{Metric, QuantizedVector};
#[inline]
pub fn estimate_distance_fast(
ctx: &QueryContext,
centroid: &[f32],
quantized: &QuantizedVector,
metric: Metric,
) -> f32 {
#[cfg(target_arch = "x86_64")]
{
#[cfg(target_feature = "avx512f")]
{
if is_x86_feature_detected!("avx512f") {
return unsafe { estimate_distance_avx512(ctx, centroid, quantized, metric) };
}
}
if is_x86_feature_detected!("avx2") && is_x86_feature_detected!("fma") {
return unsafe { estimate_distance_avx2(ctx, centroid, quantized, metric) };
}
}
estimate_distance(ctx, centroid, quantized, metric)
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2", enable = "fma")]
unsafe fn estimate_distance_avx2(
ctx: &QueryContext,
centroid: &[f32],
quantized: &QuantizedVector,
metric: Metric,
) -> f32 {
let g_add = match metric {
Metric::L2 => l2_distance_sqr_avx2(ctx.query, centroid),
Metric::InnerProduct => -dot_avx2(ctx.query, centroid),
};
let binary_code = quantized.unpack_binary_code();
let binary_dot = binary_u8_dot_f32_avx2(ctx.query, &binary_code);
let binary_term = binary_dot + ctx.c1 * ctx.sum_query;
let distance_1bit = quantized.f_add + g_add + quantized.f_rescale * binary_term;
if ctx.ex_bits > 0 {
let ex_code = quantized.unpack_ex_code();
let ex_code_u8: Vec<u8> = ex_code.iter().map(|&x| x.min(255) as u8).collect();
let ex_dot = ex_u8_dot_f32_avx2(ctx.query, &ex_code_u8);
let total_term = ctx.binary_scale * binary_dot + ex_dot + ctx.cb * ctx.sum_query;
quantized.f_add_ex + g_add + quantized.f_rescale_ex * total_term
} else {
distance_1bit
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
unsafe fn binary_u8_dot_f32_avx2(query: &[f32], binary_code: &[u8]) -> f32 {
use std::arch::x86_64::*;
let len = query.len().min(binary_code.len());
let mut sum = _mm256_setzero_ps();
let chunks = len / 8;
for i in 0..chunks {
let offset = i * 8;
let q = _mm256_loadu_ps(query.as_ptr().add(offset));
let b_u8 = _mm_loadl_epi64(binary_code.as_ptr().add(offset) as *const __m128i);
let b_i32 = _mm256_cvtepu8_epi32(b_u8);
let b = _mm256_cvtepi32_ps(b_i32);
sum = _mm256_fmadd_ps(q, b, sum);
}
let sum = _mm256_hadd_ps(sum, sum);
let sum = _mm256_hadd_ps(sum, sum);
let lo = _mm256_extractf128_ps(sum, 0);
let hi = _mm256_extractf128_ps(sum, 1);
let sum128 = _mm_add_ss(lo, hi);
let mut result = _mm_cvtss_f32(sum128);
for i in (chunks * 8)..len {
result += query[i] * (binary_code[i] as f32);
}
result
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
unsafe fn ex_u8_dot_f32_avx2(query: &[f32], ex_code: &[u8]) -> f32 {
use std::arch::x86_64::*;
let len = query.len().min(ex_code.len());
let mut sum = _mm256_setzero_ps();
let chunks = len / 8;
for i in 0..chunks {
let offset = i * 8;
let q = _mm256_loadu_ps(query.as_ptr().add(offset));
let ex_u8 = _mm_loadl_epi64(ex_code.as_ptr().add(offset) as *const __m128i);
let ex_i32 = _mm256_cvtepu8_epi32(ex_u8);
let ex = _mm256_cvtepi32_ps(ex_i32);
sum = _mm256_fmadd_ps(q, ex, sum);
}
let sum = _mm256_hadd_ps(sum, sum);
let sum = _mm256_hadd_ps(sum, sum);
let lo = _mm256_extractf128_ps(sum, 0);
let hi = _mm256_extractf128_ps(sum, 1);
let sum128 = _mm_add_ss(lo, hi);
let mut result = _mm_cvtss_f32(sum128);
for i in (chunks * 8)..len {
result += query[i] * (ex_code[i] as f32);
}
result
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2", enable = "fma")]
unsafe fn l2_distance_sqr_avx2(a: &[f32], b: &[f32]) -> f32 {
use std::arch::x86_64::*;
let len = a.len().min(b.len());
let mut sum = _mm256_setzero_ps();
let chunks = len / 8;
for i in 0..chunks {
let offset = i * 8;
let a_vec = _mm256_loadu_ps(a.as_ptr().add(offset));
let b_vec = _mm256_loadu_ps(b.as_ptr().add(offset));
let diff = _mm256_sub_ps(a_vec, b_vec);
sum = _mm256_fmadd_ps(diff, diff, sum);
}
let sum = _mm256_hadd_ps(sum, sum);
let sum = _mm256_hadd_ps(sum, sum);
let lo = _mm256_extractf128_ps(sum, 0);
let hi = _mm256_extractf128_ps(sum, 1);
let sum128 = _mm_add_ss(lo, hi);
let mut result = _mm_cvtss_f32(sum128);
for i in (chunks * 8)..len {
let diff = a[i] - b[i];
result += diff * diff;
}
result
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2", enable = "fma")]
unsafe fn dot_avx2(a: &[f32], b: &[f32]) -> f32 {
use std::arch::x86_64::*;
let len = a.len().min(b.len());
let mut sum = _mm256_setzero_ps();
let chunks = len / 8;
for i in 0..chunks {
let offset = i * 8;
let a_vec = _mm256_loadu_ps(a.as_ptr().add(offset));
let b_vec = _mm256_loadu_ps(b.as_ptr().add(offset));
sum = _mm256_fmadd_ps(a_vec, b_vec, sum);
}
let sum = _mm256_hadd_ps(sum, sum);
let sum = _mm256_hadd_ps(sum, sum);
let lo = _mm256_extractf128_ps(sum, 0);
let hi = _mm256_extractf128_ps(sum, 1);
let sum128 = _mm_add_ss(lo, hi);
let mut result = _mm_cvtss_f32(sum128);
for i in (chunks * 8)..len {
result += a[i] * b[i];
}
result
}
#[cfg(all(target_arch = "x86_64", target_feature = "avx512f"))]
#[target_feature(enable = "avx512f")]
unsafe fn binary_u8_dot_f32_avx512(query: &[f32], binary_code: &[u8]) -> f32 {
use std::arch::x86_64::*;
let len = query.len().min(binary_code.len());
let mut sum = _mm512_setzero_ps();
let chunks = len / 16;
for i in 0..chunks {
let offset = i * 16;
let q = _mm512_loadu_ps(query.as_ptr().add(offset));
let b_u8 = _mm_loadu_si128(binary_code.as_ptr().add(offset) as *const __m128i);
let b_i32 = _mm512_cvtepu8_epi32(b_u8);
let b = _mm512_cvtepi32_ps(b_i32);
sum = _mm512_fmadd_ps(q, b, sum);
}
let mut result = _mm512_reduce_add_ps(sum);
for i in (chunks * 16)..len {
result += query[i] * (binary_code[i] as f32);
}
result
}
#[cfg(all(target_arch = "x86_64", target_feature = "avx512f"))]
#[target_feature(enable = "avx512f")]
unsafe fn ex_u8_dot_f32_avx512(query: &[f32], ex_code: &[u8]) -> f32 {
use std::arch::x86_64::*;
let len = query.len().min(ex_code.len());
let mut sum = _mm512_setzero_ps();
let chunks = len / 16;
for i in 0..chunks {
let offset = i * 16;
let q = _mm512_loadu_ps(query.as_ptr().add(offset));
let ex_u8 = _mm_loadu_si128(ex_code.as_ptr().add(offset) as *const __m128i);
let ex_i32 = _mm512_cvtepu8_epi32(ex_u8);
let ex = _mm512_cvtepi32_ps(ex_i32);
sum = _mm512_fmadd_ps(q, ex, sum);
}
let mut result = _mm512_reduce_add_ps(sum);
for i in (chunks * 16)..len {
result += query[i] * (ex_code[i] as f32);
}
result
}
#[cfg(all(target_arch = "x86_64", target_feature = "avx512f"))]
#[target_feature(enable = "avx512f")]
unsafe fn l2_distance_sqr_avx512(a: &[f32], b: &[f32]) -> f32 {
use std::arch::x86_64::*;
let len = a.len().min(b.len());
let mut sum = _mm512_setzero_ps();
let chunks = len / 16;
for i in 0..chunks {
let offset = i * 16;
let a_vec = _mm512_loadu_ps(a.as_ptr().add(offset));
let b_vec = _mm512_loadu_ps(b.as_ptr().add(offset));
let diff = _mm512_sub_ps(a_vec, b_vec);
sum = _mm512_fmadd_ps(diff, diff, sum);
}
let mut result = _mm512_reduce_add_ps(sum);
for i in (chunks * 16)..len {
let diff = a[i] - b[i];
result += diff * diff;
}
result
}
#[cfg(all(target_arch = "x86_64", target_feature = "avx512f"))]
#[target_feature(enable = "avx512f")]
unsafe fn dot_avx512(a: &[f32], b: &[f32]) -> f32 {
use std::arch::x86_64::*;
let len = a.len().min(b.len());
let mut sum = _mm512_setzero_ps();
let chunks = len / 16;
for i in 0..chunks {
let offset = i * 16;
let a_vec = _mm512_loadu_ps(a.as_ptr().add(offset));
let b_vec = _mm512_loadu_ps(b.as_ptr().add(offset));
sum = _mm512_fmadd_ps(a_vec, b_vec, sum);
}
let mut result = _mm512_reduce_add_ps(sum);
for i in (chunks * 16)..len {
result += a[i] * b[i];
}
result
}
#[cfg(all(target_arch = "x86_64", target_feature = "avx512f"))]
#[target_feature(enable = "avx512f")]
unsafe fn estimate_distance_avx512(
ctx: &QueryContext,
centroid: &[f32],
quantized: &QuantizedVector,
metric: Metric,
) -> f32 {
let g_add = match metric {
Metric::L2 => l2_distance_sqr_avx512(ctx.query, centroid),
Metric::InnerProduct => -dot_avx512(ctx.query, centroid),
};
let binary_code = quantized.unpack_binary_code();
let binary_dot = binary_u8_dot_f32_avx512(ctx.query, &binary_code);
let binary_term = binary_dot + ctx.c1 * ctx.sum_query;
let distance_1bit = quantized.f_add + g_add + quantized.f_rescale * binary_term;
if ctx.ex_bits > 0 {
let ex_code = quantized.unpack_ex_code();
let ex_code_u8: Vec<u8> = ex_code.iter().map(|&x| x.min(255) as u8).collect();
let ex_dot = ex_u8_dot_f32_avx512(ctx.query, &ex_code_u8);
let total_term = ctx.binary_scale * binary_dot + ex_dot + ctx.cb * ctx.sum_query;
quantized.f_add_ex + g_add + quantized.f_rescale_ex * total_term
} else {
distance_1bit
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::quantizer::{quantize_with_centroid, RabitqConfig};
use crate::Metric;
use rand::prelude::*;
#[test]
fn test_simd_vs_scalar() {
let mut rng = StdRng::seed_from_u64(42);
let dim = 960;
let query: Vec<f32> = (0..dim).map(|_| rng.gen()).collect();
let centroid: Vec<f32> = (0..dim).map(|_| rng.gen()).collect();
let vector: Vec<f32> = (0..dim).map(|_| rng.gen()).collect();
let config = RabitqConfig::faster(dim, 7, 42);
let quantized = quantize_with_centroid(&vector, ¢roid, &config, Metric::L2);
let ex_bits = config.total_bits.saturating_sub(1) as u8;
let ctx = QueryContext::new(&query, ex_bits);
let result_scalar = estimate_distance(&ctx, ¢roid, &quantized, Metric::L2);
let result_simd = estimate_distance_fast(&ctx, ¢roid, &quantized, Metric::L2);
let diff = (result_scalar - result_simd).abs();
assert!(
diff < 0.01,
"SIMD and scalar results differ: {} vs {} (diff: {})",
result_scalar,
result_simd,
diff
);
}
#[test]
#[cfg(target_arch = "x86_64")]
fn test_binary_dot_simd() {
if !is_x86_feature_detected!("avx2") {
println!("Skipping AVX2 test on non-AVX2 CPU");
return;
}
let query = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0];
let binary = vec![1, 0, 1, 0, 1, 0, 1, 0, 1, 0];
let result_simd = unsafe { binary_u8_dot_f32_avx2(&query, &binary) };
let expected = 25.0;
assert!(
(result_simd - expected).abs() < 0.001,
"Binary dot SIMD: got {}, expected {}",
result_simd,
expected
);
}
}