#[inline]
pub fn dot_f64(a: &[f32], b: &[f32]) -> f64 {
#[cfg(target_arch = "x86_64")]
{
if std::arch::is_x86_feature_detected!("avx2") && std::arch::is_x86_feature_detected!("fma")
{
return unsafe { dot_f64_avx2_fma(a, b) };
}
}
dot_f64_unrolled(a, b)
}
#[inline]
fn dot_f64_unrolled(a: &[f32], b: &[f32]) -> f64 {
let n = a.len().min(b.len());
let chunks = n / 4;
let (mut s0, mut s1, mut s2, mut s3) = (0.0f64, 0.0f64, 0.0f64, 0.0f64);
for i in 0..chunks {
let base = i * 4;
s0 += a[base] as f64 * b[base] as f64;
s1 += a[base + 1] as f64 * b[base + 1] as f64;
s2 += a[base + 2] as f64 * b[base + 2] as f64;
s3 += a[base + 3] as f64 * b[base + 3] as f64;
}
let mut tail = 0.0f64;
for i in (chunks * 4)..n {
tail += a[i] as f64 * b[i] as f64;
}
(s0 + s1) + (s2 + s3) + tail
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2,fma")]
unsafe fn dot_f64_avx2_fma(a: &[f32], b: &[f32]) -> f64 {
use std::arch::x86_64::*;
let n = a.len().min(b.len());
let mut acc0 = _mm256_setzero_pd();
let mut acc1 = _mm256_setzero_pd();
let chunks = n / 8;
let ap = a.as_ptr();
let bp = b.as_ptr();
for i in 0..chunks {
let base = i * 8;
let a_lo = _mm256_cvtps_pd(_mm_loadu_ps(ap.add(base)));
let b_lo = _mm256_cvtps_pd(_mm_loadu_ps(bp.add(base)));
acc0 = _mm256_fmadd_pd(a_lo, b_lo, acc0);
let a_hi = _mm256_cvtps_pd(_mm_loadu_ps(ap.add(base + 4)));
let b_hi = _mm256_cvtps_pd(_mm_loadu_ps(bp.add(base + 4)));
acc1 = _mm256_fmadd_pd(a_hi, b_hi, acc1);
}
let acc = _mm256_add_pd(acc0, acc1);
let lo = _mm256_castpd256_pd128(acc);
let hi = _mm256_extractf128_pd(acc, 1);
let sum2 = _mm_add_pd(lo, hi);
let sum1 = _mm_add_sd(sum2, _mm_unpackhi_pd(sum2, sum2));
let mut total = _mm_cvtsd_f64(sum1);
for i in (chunks * 8)..n {
total += *a.get_unchecked(i) as f64 * *b.get_unchecked(i) as f64;
}
total
}
#[cfg(test)]
mod tests {
use super::*;
fn dot_reference(a: &[f32], b: &[f32]) -> f64 {
a.iter()
.zip(b.iter())
.map(|(&x, &y)| x as f64 * y as f64)
.sum()
}
fn test_vec(seed: f32, n: usize) -> Vec<f32> {
(0..n)
.map(|i| ((seed + i as f32) * 0.7311 + (i as f32) * 0.311).sin())
.collect()
}
#[test]
fn kernels_match_sequential_reference_across_dims_and_tails() {
for n in [0, 1, 3, 5, 7, 8, 9, 15, 16, 63, 64, 100, 256, 384, 385, 512] {
let a = test_vec(1.0, n);
let b = test_vec(9.0, n);
let reference = dot_reference(&a, &b);
let dispatched = dot_f64(&a, &b);
let unrolled = dot_f64_unrolled(&a, &b);
assert!(
(dispatched - reference).abs() <= 1e-9 * (1.0 + reference.abs()),
"dispatched kernel diverged at n={n}: {dispatched} vs {reference}"
);
assert!(
(unrolled - reference).abs() <= 1e-9 * (1.0 + reference.abs()),
"unrolled kernel diverged at n={n}: {unrolled} vs {reference}"
);
}
}
#[test]
#[ignore]
fn kernel_timing_before_after() {
use std::time::Instant;
const DIM: usize = 384;
const N: usize = 20_000;
const REPS: usize = 5;
let vecs: Vec<Vec<f32>> = (0..N)
.map(|i| {
(0..DIM)
.map(|j| ((i * 31 + j * 7) as f32 * 0.001).sin())
.collect()
})
.collect();
let query: Vec<f32> = (0..DIM).map(|j| ((j * 13) as f32 * 0.002).cos()).collect();
let hist = |a: &[f32], b: &[f32]| -> f64 {
let dot: f64 = a.iter().zip(b).map(|(&x, &y)| x as f64 * y as f64).sum();
let na: f64 = a
.iter()
.map(|&x| (x as f64) * (x as f64))
.sum::<f64>()
.sqrt();
let nb: f64 = b
.iter()
.map(|&x| (x as f64) * (x as f64))
.sum::<f64>()
.sqrt();
if !(na > 0.0) || !(nb > 0.0) {
return 1.0;
}
(1.0 - (dot / (na * nb))).clamp(0.0, 2.0)
};
let mut best_before = f64::MAX;
let mut sink = 0.0f64;
for _ in 0..REPS {
let t = Instant::now();
for v in &vecs {
sink += hist(&query, v);
}
best_before = best_before.min(t.elapsed().as_secs_f64());
}
let norms: Vec<f64> = vecs
.iter()
.map(|v| crate::vector::hnsw::norm_f64(v))
.collect();
let qnorm = crate::vector::hnsw::norm_f64(&query);
let mut best_after = f64::MAX;
for _ in 0..REPS {
let t = Instant::now();
for (v, &n) in vecs.iter().zip(&norms) {
sink += crate::vector::hnsw::dist_from(dot_f64(&query, v), qnorm, n);
}
best_after = best_after.min(t.elapsed().as_secs_f64());
}
let per_before = best_before / N as f64 * 1e9;
let per_after = best_after / N as f64 * 1e9;
println!("=== distance-path timing (dim={DIM}, N={N}, best of {REPS}) ===");
println!("BEFORE (3-pass scalar): {per_before:8.1} ns/comparison");
println!("AFTER (norms + SIMD dispatch): {per_after:8.1} ns/comparison");
println!("SPEEDUP: {:.2}x (sink={sink:.3})", per_before / per_after);
}
#[cfg(target_arch = "x86_64")]
#[test]
fn avx2_kernel_matches_reference_when_available() {
if !(std::arch::is_x86_feature_detected!("avx2")
&& std::arch::is_x86_feature_detected!("fma"))
{
eprintln!("skipping: avx2+fma not available on this CPU");
return;
}
for n in [7, 8, 64, 384, 385] {
let a = test_vec(2.0, n);
let b = test_vec(5.0, n);
let reference = dot_reference(&a, &b);
let simd = unsafe { dot_f64_avx2_fma(&a, &b) };
assert!(
(simd - reference).abs() <= 1e-9 * (1.0 + reference.abs()),
"avx2 kernel diverged at n={n}: {simd} vs {reference}"
);
}
}
}