Skip to content

Commit bf09ced

Browse files
committed
refactor(vector): split SIMD distance dispatch into submodules
Move simd.rs into a simd/ directory with dedicated files for each target (avx2, avx512, neon, scalar, hamming, runtime). Add a length-mismatch assertion in the top-level distance() entry point so mismatched vector dimensions panic immediately with a clear message rather than producing silent wrong results.
1 parent aa5f1ee commit bf09ced

9 files changed

Lines changed: 607 additions & 0 deletions

File tree

nodedb-vector/src/distance/mod.rs

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -13,6 +13,13 @@ pub use scalar::*;
1313
/// feature is enabled; otherwise uses scalar implementations.
1414
#[inline]
1515
pub fn distance(a: &[f32], b: &[f32], metric: DistanceMetric) -> f32 {
16+
assert_eq!(
17+
a.len(),
18+
b.len(),
19+
"distance: length mismatch (a.len()={}, b.len()={})",
20+
a.len(),
21+
b.len()
22+
);
1623
#[cfg(feature = "simd")]
1724
{
1825
let rt = simd::runtime();
Lines changed: 114 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,114 @@
1+
//! AVX2+FMA kernels for x86_64.
2+
3+
#![cfg(target_arch = "x86_64")]
4+
5+
pub fn l2_squared(a: &[f32], b: &[f32]) -> f32 {
6+
assert_eq!(a.len(), b.len(), "avx2 l2: length mismatch");
7+
// SAFETY: caller verified avx2+fma via is_x86_feature_detected.
8+
unsafe { l2_squared_impl(a, b) }
9+
}
10+
11+
#[target_feature(enable = "avx2,fma")]
12+
unsafe fn l2_squared_impl(a: &[f32], b: &[f32]) -> f32 {
13+
assert_eq!(a.len(), b.len(), "avx2 l2_impl: length mismatch");
14+
unsafe {
15+
use std::arch::x86_64::*;
16+
let n = a.len();
17+
let mut sum = _mm256_setzero_ps();
18+
let chunks = n / 8;
19+
for i in 0..chunks {
20+
let off = i * 8;
21+
let va = _mm256_loadu_ps(a.as_ptr().add(off));
22+
let vb = _mm256_loadu_ps(b.as_ptr().add(off));
23+
let diff = _mm256_sub_ps(va, vb);
24+
sum = _mm256_fmadd_ps(diff, diff, sum);
25+
}
26+
let mut result = hsum256(sum);
27+
for i in (chunks * 8)..n {
28+
let d = a[i] - b[i];
29+
result += d * d;
30+
}
31+
result
32+
}
33+
}
34+
35+
pub fn cosine_distance(a: &[f32], b: &[f32]) -> f32 {
36+
assert_eq!(a.len(), b.len(), "avx2 cosine: length mismatch");
37+
unsafe { cosine_impl(a, b) }
38+
}
39+
40+
#[target_feature(enable = "avx2,fma")]
41+
unsafe fn cosine_impl(a: &[f32], b: &[f32]) -> f32 {
42+
assert_eq!(a.len(), b.len(), "avx2 cosine_impl: length mismatch");
43+
unsafe {
44+
use std::arch::x86_64::*;
45+
let n = a.len();
46+
let mut vdot = _mm256_setzero_ps();
47+
let mut vna = _mm256_setzero_ps();
48+
let mut vnb = _mm256_setzero_ps();
49+
let chunks = n / 8;
50+
for i in 0..chunks {
51+
let off = i * 8;
52+
let va = _mm256_loadu_ps(a.as_ptr().add(off));
53+
let vb = _mm256_loadu_ps(b.as_ptr().add(off));
54+
vdot = _mm256_fmadd_ps(va, vb, vdot);
55+
vna = _mm256_fmadd_ps(va, va, vna);
56+
vnb = _mm256_fmadd_ps(vb, vb, vnb);
57+
}
58+
let mut dot = hsum256(vdot);
59+
let mut na = hsum256(vna);
60+
let mut nb = hsum256(vnb);
61+
for i in (chunks * 8)..n {
62+
dot += a[i] * b[i];
63+
na += a[i] * a[i];
64+
nb += b[i] * b[i];
65+
}
66+
let denom = (na * nb).sqrt();
67+
if denom < f32::EPSILON {
68+
1.0
69+
} else {
70+
(1.0 - dot / denom).max(0.0)
71+
}
72+
}
73+
}
74+
75+
pub fn neg_inner_product(a: &[f32], b: &[f32]) -> f32 {
76+
assert_eq!(a.len(), b.len(), "avx2 ip: length mismatch");
77+
unsafe { ip_impl(a, b) }
78+
}
79+
80+
#[target_feature(enable = "avx2,fma")]
81+
unsafe fn ip_impl(a: &[f32], b: &[f32]) -> f32 {
82+
assert_eq!(a.len(), b.len(), "avx2 ip_impl: length mismatch");
83+
unsafe {
84+
use std::arch::x86_64::*;
85+
let n = a.len();
86+
let mut vdot = _mm256_setzero_ps();
87+
let chunks = n / 8;
88+
for i in 0..chunks {
89+
let off = i * 8;
90+
let va = _mm256_loadu_ps(a.as_ptr().add(off));
91+
let vb = _mm256_loadu_ps(b.as_ptr().add(off));
92+
vdot = _mm256_fmadd_ps(va, vb, vdot);
93+
}
94+
let mut dot = hsum256(vdot);
95+
for i in (chunks * 8)..n {
96+
dot += a[i] * b[i];
97+
}
98+
-dot
99+
}
100+
}
101+
102+
/// Horizontal sum of 8 × f32 in a __m256.
103+
#[target_feature(enable = "avx2")]
104+
unsafe fn hsum256(v: std::arch::x86_64::__m256) -> f32 {
105+
use std::arch::x86_64::*;
106+
let hi = _mm256_extractf128_ps(v, 1);
107+
let lo = _mm256_castps256_ps128(v);
108+
let sum128 = _mm_add_ps(lo, hi);
109+
let shuf = _mm_movehdup_ps(sum128);
110+
let sums = _mm_add_ps(sum128, shuf);
111+
let shuf2 = _mm_movehl_ps(sums, sums);
112+
let sums2 = _mm_add_ss(sums, shuf2);
113+
_mm_cvtss_f32(sums2)
114+
}
Lines changed: 99 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,99 @@
1+
//! AVX-512 kernels for x86_64.
2+
3+
#![cfg(target_arch = "x86_64")]
4+
5+
pub fn l2_squared(a: &[f32], b: &[f32]) -> f32 {
6+
assert_eq!(a.len(), b.len(), "avx512 l2: length mismatch");
7+
unsafe { l2_impl(a, b) }
8+
}
9+
10+
#[target_feature(enable = "avx512f")]
11+
unsafe fn l2_impl(a: &[f32], b: &[f32]) -> f32 {
12+
assert_eq!(a.len(), b.len(), "avx512 l2_impl: length mismatch");
13+
unsafe {
14+
use std::arch::x86_64::*;
15+
let n = a.len();
16+
let mut sum = _mm512_setzero_ps();
17+
let chunks = n / 16;
18+
for i in 0..chunks {
19+
let off = i * 16;
20+
let va = _mm512_loadu_ps(a.as_ptr().add(off));
21+
let vb = _mm512_loadu_ps(b.as_ptr().add(off));
22+
let diff = _mm512_sub_ps(va, vb);
23+
sum = _mm512_fmadd_ps(diff, diff, sum);
24+
}
25+
let mut result = _mm512_reduce_add_ps(sum);
26+
for i in (chunks * 16)..n {
27+
let d = a[i] - b[i];
28+
result += d * d;
29+
}
30+
result
31+
}
32+
}
33+
34+
pub fn cosine_distance(a: &[f32], b: &[f32]) -> f32 {
35+
assert_eq!(a.len(), b.len(), "avx512 cosine: length mismatch");
36+
unsafe { cosine_impl(a, b) }
37+
}
38+
39+
#[target_feature(enable = "avx512f")]
40+
unsafe fn cosine_impl(a: &[f32], b: &[f32]) -> f32 {
41+
assert_eq!(a.len(), b.len(), "avx512 cosine_impl: length mismatch");
42+
unsafe {
43+
use std::arch::x86_64::*;
44+
let n = a.len();
45+
let mut vdot = _mm512_setzero_ps();
46+
let mut vna = _mm512_setzero_ps();
47+
let mut vnb = _mm512_setzero_ps();
48+
let chunks = n / 16;
49+
for i in 0..chunks {
50+
let off = i * 16;
51+
let va = _mm512_loadu_ps(a.as_ptr().add(off));
52+
let vb = _mm512_loadu_ps(b.as_ptr().add(off));
53+
vdot = _mm512_fmadd_ps(va, vb, vdot);
54+
vna = _mm512_fmadd_ps(va, va, vna);
55+
vnb = _mm512_fmadd_ps(vb, vb, vnb);
56+
}
57+
let mut dot = _mm512_reduce_add_ps(vdot);
58+
let mut na = _mm512_reduce_add_ps(vna);
59+
let mut nb = _mm512_reduce_add_ps(vnb);
60+
for i in (chunks * 16)..n {
61+
dot += a[i] * b[i];
62+
na += a[i] * a[i];
63+
nb += b[i] * b[i];
64+
}
65+
let denom = (na * nb).sqrt();
66+
if denom < f32::EPSILON {
67+
1.0
68+
} else {
69+
(1.0 - dot / denom).max(0.0)
70+
}
71+
}
72+
}
73+
74+
pub fn neg_inner_product(a: &[f32], b: &[f32]) -> f32 {
75+
assert_eq!(a.len(), b.len(), "avx512 ip: length mismatch");
76+
unsafe { ip_impl(a, b) }
77+
}
78+
79+
#[target_feature(enable = "avx512f")]
80+
unsafe fn ip_impl(a: &[f32], b: &[f32]) -> f32 {
81+
assert_eq!(a.len(), b.len(), "avx512 ip_impl: length mismatch");
82+
unsafe {
83+
use std::arch::x86_64::*;
84+
let n = a.len();
85+
let mut vdot = _mm512_setzero_ps();
86+
let chunks = n / 16;
87+
for i in 0..chunks {
88+
let off = i * 16;
89+
let va = _mm512_loadu_ps(a.as_ptr().add(off));
90+
let vb = _mm512_loadu_ps(b.as_ptr().add(off));
91+
vdot = _mm512_fmadd_ps(va, vb, vdot);
92+
}
93+
let mut dot = _mm512_reduce_add_ps(vdot);
94+
for i in (chunks * 16)..n {
95+
dot += a[i] * b[i];
96+
}
97+
-dot
98+
}
99+
}
Lines changed: 35 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,35 @@
1+
//! Fast Hamming distance using u64 POPCNT.
2+
3+
pub fn fast_hamming(a: &[u8], b: &[u8]) -> u32 {
4+
assert_eq!(a.len(), b.len(), "fast_hamming: length mismatch");
5+
let mut dist = 0u32;
6+
let chunks = a.len() / 8;
7+
for i in 0..chunks {
8+
let off = i * 8;
9+
let xa = u64::from_le_bytes([
10+
a[off],
11+
a[off + 1],
12+
a[off + 2],
13+
a[off + 3],
14+
a[off + 4],
15+
a[off + 5],
16+
a[off + 6],
17+
a[off + 7],
18+
]);
19+
let xb = u64::from_le_bytes([
20+
b[off],
21+
b[off + 1],
22+
b[off + 2],
23+
b[off + 3],
24+
b[off + 4],
25+
b[off + 5],
26+
b[off + 6],
27+
b[off + 7],
28+
]);
29+
dist += (xa ^ xb).count_ones();
30+
}
31+
for i in (chunks * 8)..a.len() {
32+
dist += (a[i] ^ b[i]).count_ones();
33+
}
34+
dist
35+
}
Lines changed: 14 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,14 @@
1+
//! Runtime SIMD dispatch for vector distance and bitmap operations.
2+
3+
pub mod hamming;
4+
pub mod runtime;
5+
pub mod scalar;
6+
7+
#[cfg(target_arch = "x86_64")]
8+
pub mod avx2;
9+
#[cfg(target_arch = "x86_64")]
10+
pub mod avx512;
11+
#[cfg(target_arch = "aarch64")]
12+
pub mod neon;
13+
14+
pub use runtime::{SimdRuntime, runtime};
Lines changed: 96 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,96 @@
1+
//! NEON kernels for ARM64.
2+
3+
#![cfg(target_arch = "aarch64")]
4+
5+
pub fn l2_squared(a: &[f32], b: &[f32]) -> f32 {
6+
assert_eq!(a.len(), b.len(), "neon l2: length mismatch");
7+
unsafe { l2_impl(a, b) }
8+
}
9+
10+
unsafe fn l2_impl(a: &[f32], b: &[f32]) -> f32 {
11+
assert_eq!(a.len(), b.len(), "neon l2_impl: length mismatch");
12+
unsafe {
13+
use std::arch::aarch64::*;
14+
let n = a.len();
15+
let mut sum = vdupq_n_f32(0.0);
16+
let chunks = n / 4;
17+
for i in 0..chunks {
18+
let off = i * 4;
19+
let va = vld1q_f32(a.as_ptr().add(off));
20+
let vb = vld1q_f32(b.as_ptr().add(off));
21+
let diff = vsubq_f32(va, vb);
22+
sum = vfmaq_f32(sum, diff, diff);
23+
}
24+
let mut result = vaddvq_f32(sum);
25+
for i in (chunks * 4)..n {
26+
let d = a[i] - b[i];
27+
result += d * d;
28+
}
29+
result
30+
}
31+
}
32+
33+
pub fn cosine_distance(a: &[f32], b: &[f32]) -> f32 {
34+
assert_eq!(a.len(), b.len(), "neon cosine: length mismatch");
35+
unsafe { cosine_impl(a, b) }
36+
}
37+
38+
unsafe fn cosine_impl(a: &[f32], b: &[f32]) -> f32 {
39+
assert_eq!(a.len(), b.len(), "neon cosine_impl: length mismatch");
40+
unsafe {
41+
use std::arch::aarch64::*;
42+
let n = a.len();
43+
let mut vdot = vdupq_n_f32(0.0);
44+
let mut vna = vdupq_n_f32(0.0);
45+
let mut vnb = vdupq_n_f32(0.0);
46+
let chunks = n / 4;
47+
for i in 0..chunks {
48+
let off = i * 4;
49+
let va = vld1q_f32(a.as_ptr().add(off));
50+
let vb = vld1q_f32(b.as_ptr().add(off));
51+
vdot = vfmaq_f32(vdot, va, vb);
52+
vna = vfmaq_f32(vna, va, va);
53+
vnb = vfmaq_f32(vnb, vb, vb);
54+
}
55+
let mut dot = vaddvq_f32(vdot);
56+
let mut na = vaddvq_f32(vna);
57+
let mut nb = vaddvq_f32(vnb);
58+
for i in (chunks * 4)..n {
59+
dot += a[i] * b[i];
60+
na += a[i] * a[i];
61+
nb += b[i] * b[i];
62+
}
63+
let denom = (na * nb).sqrt();
64+
if denom < f32::EPSILON {
65+
1.0
66+
} else {
67+
(1.0 - dot / denom).max(0.0)
68+
}
69+
}
70+
}
71+
72+
pub fn neg_inner_product(a: &[f32], b: &[f32]) -> f32 {
73+
assert_eq!(a.len(), b.len(), "neon ip: length mismatch");
74+
unsafe { ip_impl(a, b) }
75+
}
76+
77+
unsafe fn ip_impl(a: &[f32], b: &[f32]) -> f32 {
78+
assert_eq!(a.len(), b.len(), "neon ip_impl: length mismatch");
79+
unsafe {
80+
use std::arch::aarch64::*;
81+
let n = a.len();
82+
let mut vdot = vdupq_n_f32(0.0);
83+
let chunks = n / 4;
84+
for i in 0..chunks {
85+
let off = i * 4;
86+
let va = vld1q_f32(a.as_ptr().add(off));
87+
let vb = vld1q_f32(b.as_ptr().add(off));
88+
vdot = vfmaq_f32(vdot, va, vb);
89+
}
90+
let mut dot = vaddvq_f32(vdot);
91+
for i in (chunks * 4)..n {
92+
dot += a[i] * b[i];
93+
}
94+
-dot
95+
}
96+
}

0 commit comments

Comments
 (0)