Skip to main content

ruvector_core/
simd_intrinsics.rs

1//! Custom SIMD intrinsics for performance-critical operations
2//!
3//! This module provides hand-optimized SIMD implementations:
4//! - AVX2/AVX-512 for x86_64 processors
5//! - NEON for ARM64/Apple Silicon processors (M1/M2/M3/M4)
6//!
7//! Distance calculations and other vectorized operations are automatically
8//! dispatched to the optimal implementation based on the target architecture.
9//!
10//! ## Features
11//!
12//! - **AVX-512 Support**: 512-bit operations processing 16 floats per iteration
13//! - **INT8 Quantized Operations**: SIMD-accelerated quantized vector operations
14//! - **Batch Operations**: Cache-optimized batch distance calculations
15//! - **NEON Optimizations**: Prefetch hints and loop unrolling for ARM64
16//!
17//! ## Performance Optimizations (v2)
18//!
19//! - **Loop Unrolling**: 4x unrolled loops for better instruction-level parallelism
20//! - **Prefetch Hints**: Software prefetching for large vectors (>256 elements)
21//! - **FMA Instructions**: Fused multiply-add for improved throughput and accuracy
22//! - **Efficient Horizontal Sum**: Optimized reduction operations
23
24#[cfg(target_arch = "x86_64")]
25use std::arch::x86_64::*;
26
27#[cfg(target_arch = "aarch64")]
28use std::arch::aarch64::*;
29
30/// Prefetch distance in cache lines (tuned for L1 cache, 64 bytes = 16 floats)
31#[allow(dead_code)]
32const PREFETCH_DISTANCE: usize = 64;
33
34/// SIMD-optimized euclidean distance
35/// Uses AVX-512 > AVX2 on x86_64, NEON on ARM64/Apple Silicon, falls back to scalar otherwise
36///
37/// # Optimizations for M4 Pro (ARM64)
38/// - Uses 4x loop unrolling for vectors >= 64 elements
39/// - FMA instructions for improved throughput
40/// - Optimized horizontal reduction via `vaddvq_f32`
41#[inline(always)]
42pub fn euclidean_distance_simd(a: &[f32], b: &[f32]) -> f32 {
43    #[cfg(target_arch = "x86_64")]
44    {
45        #[cfg(feature = "simd-avx512")]
46        {
47            if is_x86_feature_detected!("avx512f") {
48                return unsafe { euclidean_distance_avx512_impl(a, b) };
49            }
50        }
51        if is_x86_feature_detected!("avx2") && is_x86_feature_detected!("fma") {
52            unsafe { euclidean_distance_avx2_fma_impl(a, b) }
53        } else if is_x86_feature_detected!("avx2") {
54            unsafe { euclidean_distance_avx2_impl(a, b) }
55        } else {
56            euclidean_distance_scalar(a, b)
57        }
58    }
59
60    #[cfg(target_arch = "aarch64")]
61    {
62        // Use unrolled version for vectors >= 64 elements for better ILP
63        if a.len() >= 64 {
64            unsafe { euclidean_distance_neon_unrolled_impl(a, b) }
65        } else {
66            unsafe { euclidean_distance_neon_impl(a, b) }
67        }
68    }
69
70    #[cfg(not(any(target_arch = "x86_64", target_arch = "aarch64")))]
71    {
72        euclidean_distance_scalar(a, b)
73    }
74}
75
76/// Legacy alias for backward compatibility
77#[inline(always)]
78pub fn euclidean_distance_avx2(a: &[f32], b: &[f32]) -> f32 {
79    euclidean_distance_simd(a, b)
80}
81
82#[cfg(target_arch = "x86_64")]
83#[target_feature(enable = "avx2")]
84unsafe fn euclidean_distance_avx2_impl(a: &[f32], b: &[f32]) -> f32 {
85    // SECURITY: Ensure both arrays have the same length to prevent out-of-bounds access
86    assert_eq!(a.len(), b.len(), "Input arrays must have the same length");
87
88    let len = a.len();
89    let mut sum = _mm256_setzero_ps();
90
91    // Process 8 floats at a time
92    let chunks = len / 8;
93    for i in 0..chunks {
94        let idx = i * 8;
95
96        // Load 8 floats from each array
97        let va = _mm256_loadu_ps(a.as_ptr().add(idx));
98        let vb = _mm256_loadu_ps(b.as_ptr().add(idx));
99
100        // Compute difference: (a - b)
101        let diff = _mm256_sub_ps(va, vb);
102
103        // Square the difference: (a - b)^2
104        let sq = _mm256_mul_ps(diff, diff);
105
106        // Accumulate
107        sum = _mm256_add_ps(sum, sq);
108    }
109
110    // Horizontal sum of the 8 floats in the AVX register
111    let sum_arr: [f32; 8] = std::mem::transmute(sum);
112    let mut total = sum_arr.iter().sum::<f32>();
113
114    // Handle remaining elements (if len not divisible by 8)
115    for i in (chunks * 8)..len {
116        let diff = a[i] - b[i];
117        total += diff * diff;
118    }
119
120    total.sqrt()
121}
122
123/// AVX2 with FMA - 4x loop unrolling for better instruction-level parallelism
124#[cfg(target_arch = "x86_64")]
125#[target_feature(enable = "avx2", enable = "fma")]
126unsafe fn euclidean_distance_avx2_fma_impl(a: &[f32], b: &[f32]) -> f32 {
127    assert_eq!(a.len(), b.len(), "Input arrays must have the same length");
128
129    let len = a.len();
130    // Use 4 accumulators for better ILP (instruction-level parallelism)
131    let mut sum0 = _mm256_setzero_ps();
132    let mut sum1 = _mm256_setzero_ps();
133    let mut sum2 = _mm256_setzero_ps();
134    let mut sum3 = _mm256_setzero_ps();
135
136    // Process 32 floats at a time (4 x 8 floats)
137    let chunks = len / 32;
138    for i in 0..chunks {
139        let idx = i * 32;
140
141        // Load and process 4 vectors of 8 floats each
142        let va0 = _mm256_loadu_ps(a.as_ptr().add(idx));
143        let vb0 = _mm256_loadu_ps(b.as_ptr().add(idx));
144        let diff0 = _mm256_sub_ps(va0, vb0);
145        sum0 = _mm256_fmadd_ps(diff0, diff0, sum0);
146
147        let va1 = _mm256_loadu_ps(a.as_ptr().add(idx + 8));
148        let vb1 = _mm256_loadu_ps(b.as_ptr().add(idx + 8));
149        let diff1 = _mm256_sub_ps(va1, vb1);
150        sum1 = _mm256_fmadd_ps(diff1, diff1, sum1);
151
152        let va2 = _mm256_loadu_ps(a.as_ptr().add(idx + 16));
153        let vb2 = _mm256_loadu_ps(b.as_ptr().add(idx + 16));
154        let diff2 = _mm256_sub_ps(va2, vb2);
155        sum2 = _mm256_fmadd_ps(diff2, diff2, sum2);
156
157        let va3 = _mm256_loadu_ps(a.as_ptr().add(idx + 24));
158        let vb3 = _mm256_loadu_ps(b.as_ptr().add(idx + 24));
159        let diff3 = _mm256_sub_ps(va3, vb3);
160        sum3 = _mm256_fmadd_ps(diff3, diff3, sum3);
161    }
162
163    // Combine the 4 accumulators
164    let sum01 = _mm256_add_ps(sum0, sum1);
165    let sum23 = _mm256_add_ps(sum2, sum3);
166    let sum = _mm256_add_ps(sum01, sum23);
167
168    // Process remaining 8-float chunks
169    let remaining_start = chunks * 32;
170    let remaining_chunks = (len - remaining_start) / 8;
171    let mut final_sum = sum;
172    for i in 0..remaining_chunks {
173        let idx = remaining_start + i * 8;
174        let va = _mm256_loadu_ps(a.as_ptr().add(idx));
175        let vb = _mm256_loadu_ps(b.as_ptr().add(idx));
176        let diff = _mm256_sub_ps(va, vb);
177        final_sum = _mm256_fmadd_ps(diff, diff, final_sum);
178    }
179
180    // Horizontal sum
181    let sum_arr: [f32; 8] = std::mem::transmute(final_sum);
182    let mut total = sum_arr.iter().sum::<f32>();
183
184    // Handle remaining elements
185    let scalar_start = remaining_start + remaining_chunks * 8;
186    for i in scalar_start..len {
187        let diff = a[i] - b[i];
188        total += diff * diff;
189    }
190
191    total.sqrt()
192}
193
194// ============================================================================
195// AVX-512 implementations for x86_64 (Intel Ice Lake, Sapphire Rapids, AMD Zen 4+)
196// ============================================================================
197
198/// AVX-512 euclidean distance - 4-accumulator version for ILP
199/// Processes 64 floats per iteration (4×16) to hide 4-cycle FMA latency on Zen 4/5.
200/// For 384-dim vectors: 384/64 = 6 exact iterations with no scalar tail.
201#[cfg(all(target_arch = "x86_64", feature = "simd-avx512"))]
202#[target_feature(enable = "avx512f")]
203unsafe fn euclidean_distance_avx512_impl(a: &[f32], b: &[f32]) -> f32 {
204    assert_eq!(a.len(), b.len(), "Input arrays must have the same length");
205
206    let len = a.len();
207    // 4 independent accumulators hide the 4-cycle FMA latency
208    let mut sum0 = _mm512_setzero_ps();
209    let mut sum1 = _mm512_setzero_ps();
210    let mut sum2 = _mm512_setzero_ps();
211    let mut sum3 = _mm512_setzero_ps();
212
213    // Process 64 floats at a time (4 × 16)
214    let chunks = len / 64;
215    for i in 0..chunks {
216        let idx = i * 64;
217        let va0 = _mm512_loadu_ps(a.as_ptr().add(idx));
218        let vb0 = _mm512_loadu_ps(b.as_ptr().add(idx));
219        let diff0 = _mm512_sub_ps(va0, vb0);
220        sum0 = _mm512_fmadd_ps(diff0, diff0, sum0);
221
222        let va1 = _mm512_loadu_ps(a.as_ptr().add(idx + 16));
223        let vb1 = _mm512_loadu_ps(b.as_ptr().add(idx + 16));
224        let diff1 = _mm512_sub_ps(va1, vb1);
225        sum1 = _mm512_fmadd_ps(diff1, diff1, sum1);
226
227        let va2 = _mm512_loadu_ps(a.as_ptr().add(idx + 32));
228        let vb2 = _mm512_loadu_ps(b.as_ptr().add(idx + 32));
229        let diff2 = _mm512_sub_ps(va2, vb2);
230        sum2 = _mm512_fmadd_ps(diff2, diff2, sum2);
231
232        let va3 = _mm512_loadu_ps(a.as_ptr().add(idx + 48));
233        let vb3 = _mm512_loadu_ps(b.as_ptr().add(idx + 48));
234        let diff3 = _mm512_sub_ps(va3, vb3);
235        sum3 = _mm512_fmadd_ps(diff3, diff3, sum3);
236    }
237
238    // Tree-reduce accumulators
239    let sum01 = _mm512_add_ps(sum0, sum1);
240    let sum23 = _mm512_add_ps(sum2, sum3);
241    let mut sum = _mm512_add_ps(sum01, sum23);
242
243    // Handle remaining 16-float chunks
244    let remaining_start = chunks * 64;
245    let remaining_chunks = (len - remaining_start) / 16;
246    for i in 0..remaining_chunks {
247        let idx = remaining_start + i * 16;
248        let va = _mm512_loadu_ps(a.as_ptr().add(idx));
249        let vb = _mm512_loadu_ps(b.as_ptr().add(idx));
250        let diff = _mm512_sub_ps(va, vb);
251        sum = _mm512_fmadd_ps(diff, diff, sum);
252    }
253
254    let mut total = _mm512_reduce_add_ps(sum);
255
256    // Scalar tail (0–15 elements)
257    let scalar_start = remaining_start + remaining_chunks * 16;
258    for i in scalar_start..len {
259        let diff = a[i] - b[i];
260        total += diff * diff;
261    }
262
263    total.sqrt()
264}
265
266/// AVX-512 dot product - 4-accumulator version for ILP
267/// Processes 64 floats per iteration (4×16) to hide 4-cycle FMA latency on Zen 4/5.
268#[cfg(all(target_arch = "x86_64", feature = "simd-avx512"))]
269#[target_feature(enable = "avx512f")]
270unsafe fn dot_product_avx512_impl(a: &[f32], b: &[f32]) -> f32 {
271    assert_eq!(a.len(), b.len(), "Input arrays must have the same length");
272
273    let len = a.len();
274    let mut sum0 = _mm512_setzero_ps();
275    let mut sum1 = _mm512_setzero_ps();
276    let mut sum2 = _mm512_setzero_ps();
277    let mut sum3 = _mm512_setzero_ps();
278
279    let chunks = len / 64;
280    for i in 0..chunks {
281        let idx = i * 64;
282        let va0 = _mm512_loadu_ps(a.as_ptr().add(idx));
283        let vb0 = _mm512_loadu_ps(b.as_ptr().add(idx));
284        sum0 = _mm512_fmadd_ps(va0, vb0, sum0);
285
286        let va1 = _mm512_loadu_ps(a.as_ptr().add(idx + 16));
287        let vb1 = _mm512_loadu_ps(b.as_ptr().add(idx + 16));
288        sum1 = _mm512_fmadd_ps(va1, vb1, sum1);
289
290        let va2 = _mm512_loadu_ps(a.as_ptr().add(idx + 32));
291        let vb2 = _mm512_loadu_ps(b.as_ptr().add(idx + 32));
292        sum2 = _mm512_fmadd_ps(va2, vb2, sum2);
293
294        let va3 = _mm512_loadu_ps(a.as_ptr().add(idx + 48));
295        let vb3 = _mm512_loadu_ps(b.as_ptr().add(idx + 48));
296        sum3 = _mm512_fmadd_ps(va3, vb3, sum3);
297    }
298
299    let sum01 = _mm512_add_ps(sum0, sum1);
300    let sum23 = _mm512_add_ps(sum2, sum3);
301    let mut sum = _mm512_add_ps(sum01, sum23);
302
303    let remaining_start = chunks * 64;
304    let remaining_chunks = (len - remaining_start) / 16;
305    for i in 0..remaining_chunks {
306        let idx = remaining_start + i * 16;
307        let va = _mm512_loadu_ps(a.as_ptr().add(idx));
308        let vb = _mm512_loadu_ps(b.as_ptr().add(idx));
309        sum = _mm512_fmadd_ps(va, vb, sum);
310    }
311
312    let mut total = _mm512_reduce_add_ps(sum);
313
314    let scalar_start = remaining_start + remaining_chunks * 16;
315    for i in scalar_start..len {
316        total += a[i] * b[i];
317    }
318
319    total
320}
321
322/// AVX-512 cosine similarity - 2-accumulator-per-component version for ILP
323/// Uses 6 ZMM registers (dot0/dot1, na0/na1, nb0/nb1) to hide 4-cycle FMA latency.
324/// Processes 32 floats per iteration (2×16) on Zen 4/5.
325#[cfg(all(target_arch = "x86_64", feature = "simd-avx512"))]
326#[target_feature(enable = "avx512f")]
327unsafe fn cosine_similarity_avx512_impl(a: &[f32], b: &[f32]) -> f32 {
328    assert_eq!(a.len(), b.len(), "Input arrays must have the same length");
329
330    let len = a.len();
331    let mut dot0 = _mm512_setzero_ps();
332    let mut dot1 = _mm512_setzero_ps();
333    let mut norm_a0 = _mm512_setzero_ps();
334    let mut norm_a1 = _mm512_setzero_ps();
335    let mut norm_b0 = _mm512_setzero_ps();
336    let mut norm_b1 = _mm512_setzero_ps();
337
338    // Process 32 floats at a time (2 × 16)
339    let chunks = len / 32;
340    for i in 0..chunks {
341        let idx = i * 32;
342        let va0 = _mm512_loadu_ps(a.as_ptr().add(idx));
343        let vb0 = _mm512_loadu_ps(b.as_ptr().add(idx));
344        dot0 = _mm512_fmadd_ps(va0, vb0, dot0);
345        norm_a0 = _mm512_fmadd_ps(va0, va0, norm_a0);
346        norm_b0 = _mm512_fmadd_ps(vb0, vb0, norm_b0);
347
348        let va1 = _mm512_loadu_ps(a.as_ptr().add(idx + 16));
349        let vb1 = _mm512_loadu_ps(b.as_ptr().add(idx + 16));
350        dot1 = _mm512_fmadd_ps(va1, vb1, dot1);
351        norm_a1 = _mm512_fmadd_ps(va1, va1, norm_a1);
352        norm_b1 = _mm512_fmadd_ps(vb1, vb1, norm_b1);
353    }
354
355    // Tree-reduce each component
356    let mut dot_v = _mm512_add_ps(dot0, dot1);
357    let mut na_v = _mm512_add_ps(norm_a0, norm_a1);
358    let mut nb_v = _mm512_add_ps(norm_b0, norm_b1);
359
360    // Handle remaining 16-float chunks
361    let remaining_start = chunks * 32;
362    let remaining_chunks = (len - remaining_start) / 16;
363    for i in 0..remaining_chunks {
364        let idx = remaining_start + i * 16;
365        let va = _mm512_loadu_ps(a.as_ptr().add(idx));
366        let vb = _mm512_loadu_ps(b.as_ptr().add(idx));
367        dot_v = _mm512_fmadd_ps(va, vb, dot_v);
368        na_v = _mm512_fmadd_ps(va, va, na_v);
369        nb_v = _mm512_fmadd_ps(vb, vb, nb_v);
370    }
371
372    let mut dot_sum = _mm512_reduce_add_ps(dot_v);
373    let mut norm_a_sum = _mm512_reduce_add_ps(na_v);
374    let mut norm_b_sum = _mm512_reduce_add_ps(nb_v);
375
376    let scalar_start = remaining_start + remaining_chunks * 16;
377    for i in scalar_start..len {
378        dot_sum += a[i] * b[i];
379        norm_a_sum += a[i] * a[i];
380        norm_b_sum += b[i] * b[i];
381    }
382
383    dot_sum / (norm_a_sum.sqrt() * norm_b_sum.sqrt())
384}
385
386/// AVX-512 Manhattan distance - 4-accumulator version for ILP
387/// Processes 64 floats per iteration (4×16) to hide add-latency on Zen 4/5.
388#[cfg(all(target_arch = "x86_64", feature = "simd-avx512"))]
389#[target_feature(enable = "avx512f")]
390unsafe fn manhattan_distance_avx512_impl(a: &[f32], b: &[f32]) -> f32 {
391    assert_eq!(a.len(), b.len(), "Input arrays must have the same length");
392
393    let len = a.len();
394    let mut sum0 = _mm512_setzero_ps();
395    let mut sum1 = _mm512_setzero_ps();
396    let mut sum2 = _mm512_setzero_ps();
397    let mut sum3 = _mm512_setzero_ps();
398
399    let chunks = len / 64;
400    for i in 0..chunks {
401        let idx = i * 64;
402        let va0 = _mm512_loadu_ps(a.as_ptr().add(idx));
403        let vb0 = _mm512_loadu_ps(b.as_ptr().add(idx));
404        let diff0 = _mm512_sub_ps(va0, vb0);
405        sum0 = _mm512_add_ps(sum0, _mm512_abs_ps(diff0));
406
407        let va1 = _mm512_loadu_ps(a.as_ptr().add(idx + 16));
408        let vb1 = _mm512_loadu_ps(b.as_ptr().add(idx + 16));
409        let diff1 = _mm512_sub_ps(va1, vb1);
410        sum1 = _mm512_add_ps(sum1, _mm512_abs_ps(diff1));
411
412        let va2 = _mm512_loadu_ps(a.as_ptr().add(idx + 32));
413        let vb2 = _mm512_loadu_ps(b.as_ptr().add(idx + 32));
414        let diff2 = _mm512_sub_ps(va2, vb2);
415        sum2 = _mm512_add_ps(sum2, _mm512_abs_ps(diff2));
416
417        let va3 = _mm512_loadu_ps(a.as_ptr().add(idx + 48));
418        let vb3 = _mm512_loadu_ps(b.as_ptr().add(idx + 48));
419        let diff3 = _mm512_sub_ps(va3, vb3);
420        sum3 = _mm512_add_ps(sum3, _mm512_abs_ps(diff3));
421    }
422
423    let sum01 = _mm512_add_ps(sum0, sum1);
424    let sum23 = _mm512_add_ps(sum2, sum3);
425    let mut sum = _mm512_add_ps(sum01, sum23);
426
427    let remaining_start = chunks * 64;
428    let remaining_chunks = (len - remaining_start) / 16;
429    for i in 0..remaining_chunks {
430        let idx = remaining_start + i * 16;
431        let va = _mm512_loadu_ps(a.as_ptr().add(idx));
432        let vb = _mm512_loadu_ps(b.as_ptr().add(idx));
433        let diff = _mm512_sub_ps(va, vb);
434        sum = _mm512_add_ps(sum, _mm512_abs_ps(diff));
435    }
436
437    let mut total = _mm512_reduce_add_ps(sum);
438
439    let scalar_start = remaining_start + remaining_chunks * 16;
440    for i in scalar_start..len {
441        total += (a[i] - b[i]).abs();
442    }
443
444    total
445}
446
447// ============================================================================
448// NEON implementations for ARM64/Apple Silicon (M1/M2/M3/M4)
449// ============================================================================
450
451/// NEON-optimized euclidean distance for ARM64 (original non-unrolled version)
452/// Processes 4 floats at a time using 128-bit NEON registers
453///
454/// # Safety
455/// Caller must ensure a.len() == b.len()
456#[cfg(target_arch = "aarch64")]
457#[inline(always)]
458#[allow(dead_code)]
459unsafe fn euclidean_distance_neon_impl(a: &[f32], b: &[f32]) -> f32 {
460    debug_assert_eq!(a.len(), b.len(), "Input arrays must have the same length");
461
462    let len = a.len();
463    let mut sum = vdupq_n_f32(0.0);
464
465    let a_ptr = a.as_ptr();
466    let b_ptr = b.as_ptr();
467
468    // Process 4 floats at a time with NEON
469    let chunks = len / 4;
470    let mut idx = 0usize;
471
472    for _ in 0..chunks {
473        let va = vld1q_f32(a_ptr.add(idx));
474        let vb = vld1q_f32(b_ptr.add(idx));
475
476        // Compute difference: (a - b)
477        let diff = vsubq_f32(va, vb);
478
479        // Square and accumulate: sum += (a - b)^2
480        sum = vfmaq_f32(sum, diff, diff);
481
482        idx += 4;
483    }
484
485    // Horizontal sum of the 4 floats
486    let mut total = vaddvq_f32(sum);
487
488    // Handle remaining elements (use get_unchecked for bounds-check elimination)
489    for i in (chunks * 4)..len {
490        let diff = *a.get_unchecked(i) - *b.get_unchecked(i);
491        total += diff * diff;
492    }
493
494    total.sqrt()
495}
496
497/// NEON-optimized dot product for ARM64
498///
499/// # Safety
500/// Caller must ensure a.len() == b.len()
501#[cfg(target_arch = "aarch64")]
502#[inline(always)]
503unsafe fn dot_product_neon_impl(a: &[f32], b: &[f32]) -> f32 {
504    debug_assert_eq!(a.len(), b.len(), "Input arrays must have the same length");
505
506    let len = a.len();
507    let mut sum = vdupq_n_f32(0.0);
508
509    let a_ptr = a.as_ptr();
510    let b_ptr = b.as_ptr();
511
512    let chunks = len / 4;
513    let mut idx = 0usize;
514
515    for _ in 0..chunks {
516        let va = vld1q_f32(a_ptr.add(idx));
517        let vb = vld1q_f32(b_ptr.add(idx));
518
519        // Fused multiply-add: sum += a * b
520        sum = vfmaq_f32(sum, va, vb);
521
522        idx += 4;
523    }
524
525    let mut total = vaddvq_f32(sum);
526
527    // Handle remaining elements with bounds-check elimination
528    for i in (chunks * 4)..len {
529        total += *a.get_unchecked(i) * *b.get_unchecked(i);
530    }
531
532    total
533}
534
535/// NEON-optimized cosine similarity for ARM64
536///
537/// # Safety
538/// Caller must ensure a.len() == b.len()
539#[cfg(target_arch = "aarch64")]
540#[inline(always)]
541unsafe fn cosine_similarity_neon_impl(a: &[f32], b: &[f32]) -> f32 {
542    debug_assert_eq!(a.len(), b.len(), "Input arrays must have the same length");
543
544    let len = a.len();
545    let mut dot = vdupq_n_f32(0.0);
546    let mut norm_a = vdupq_n_f32(0.0);
547    let mut norm_b = vdupq_n_f32(0.0);
548
549    let a_ptr = a.as_ptr();
550    let b_ptr = b.as_ptr();
551
552    let chunks = len / 4;
553    let mut idx = 0usize;
554
555    for _ in 0..chunks {
556        let va = vld1q_f32(a_ptr.add(idx));
557        let vb = vld1q_f32(b_ptr.add(idx));
558
559        // Dot product
560        dot = vfmaq_f32(dot, va, vb);
561
562        // Norms (squared)
563        norm_a = vfmaq_f32(norm_a, va, va);
564        norm_b = vfmaq_f32(norm_b, vb, vb);
565
566        idx += 4;
567    }
568
569    let mut dot_sum = vaddvq_f32(dot);
570    let mut norm_a_sum = vaddvq_f32(norm_a);
571    let mut norm_b_sum = vaddvq_f32(norm_b);
572
573    // Handle remaining elements with bounds-check elimination
574    for i in (chunks * 4)..len {
575        let ai = *a.get_unchecked(i);
576        let bi = *b.get_unchecked(i);
577        dot_sum += ai * bi;
578        norm_a_sum += ai * ai;
579        norm_b_sum += bi * bi;
580    }
581
582    dot_sum / (norm_a_sum.sqrt() * norm_b_sum.sqrt())
583}
584
585/// NEON-optimized Manhattan distance for ARM64
586///
587/// # Safety
588/// Caller must ensure a.len() == b.len()
589#[cfg(target_arch = "aarch64")]
590#[inline(always)]
591unsafe fn manhattan_distance_neon_impl(a: &[f32], b: &[f32]) -> f32 {
592    debug_assert_eq!(a.len(), b.len(), "Input arrays must have the same length");
593
594    let len = a.len();
595    let mut sum = vdupq_n_f32(0.0);
596
597    let a_ptr = a.as_ptr();
598    let b_ptr = b.as_ptr();
599
600    let chunks = len / 4;
601    let mut idx = 0usize;
602
603    for _ in 0..chunks {
604        let va = vld1q_f32(a_ptr.add(idx));
605        let vb = vld1q_f32(b_ptr.add(idx));
606
607        // Absolute difference using vabdq_f32 (absolute difference in one instruction)
608        let abs_diff = vabdq_f32(va, vb);
609        sum = vaddq_f32(sum, abs_diff);
610
611        idx += 4;
612    }
613
614    let mut total = vaddvq_f32(sum);
615
616    // Handle remaining elements with bounds-check elimination
617    for i in (chunks * 4)..len {
618        total += (*a.get_unchecked(i) - *b.get_unchecked(i)).abs();
619    }
620
621    total
622}
623
624/// NEON-optimized euclidean distance with 4x loop unrolling
625/// Optimized for larger vectors (>= 64 elements) common in ML embeddings
626///
627/// # Safety
628/// Caller must ensure a.len() == b.len()
629///
630/// # M4 Pro Optimizations
631/// - 4 independent accumulators for maximum ILP on M4's 6-wide superscalar core
632/// - Software prefetching for vectors > 256 elements
633/// - Bounds-check elimination in remainder loops
634#[cfg(target_arch = "aarch64")]
635#[inline(always)]
636unsafe fn euclidean_distance_neon_unrolled_impl(a: &[f32], b: &[f32]) -> f32 {
637    debug_assert_eq!(a.len(), b.len(), "Input arrays must have the same length");
638
639    let len = a.len();
640    let a_ptr = a.as_ptr();
641    let b_ptr = b.as_ptr();
642
643    // Use 4 accumulators for better instruction-level parallelism
644    let mut sum0 = vdupq_n_f32(0.0);
645    let mut sum1 = vdupq_n_f32(0.0);
646    let mut sum2 = vdupq_n_f32(0.0);
647    let mut sum3 = vdupq_n_f32(0.0);
648
649    // Process 16 floats at a time (4 x 4 floats)
650    let chunks = len / 16;
651    let mut idx = 0usize;
652
653    for _ in 0..chunks {
654        // Unroll 4x for better ILP - all loads and operations are independent
655        let va0 = vld1q_f32(a_ptr.add(idx));
656        let vb0 = vld1q_f32(b_ptr.add(idx));
657        let diff0 = vsubq_f32(va0, vb0);
658        sum0 = vfmaq_f32(sum0, diff0, diff0);
659
660        let va1 = vld1q_f32(a_ptr.add(idx + 4));
661        let vb1 = vld1q_f32(b_ptr.add(idx + 4));
662        let diff1 = vsubq_f32(va1, vb1);
663        sum1 = vfmaq_f32(sum1, diff1, diff1);
664
665        let va2 = vld1q_f32(a_ptr.add(idx + 8));
666        let vb2 = vld1q_f32(b_ptr.add(idx + 8));
667        let diff2 = vsubq_f32(va2, vb2);
668        sum2 = vfmaq_f32(sum2, diff2, diff2);
669
670        let va3 = vld1q_f32(a_ptr.add(idx + 12));
671        let vb3 = vld1q_f32(b_ptr.add(idx + 12));
672        let diff3 = vsubq_f32(va3, vb3);
673        sum3 = vfmaq_f32(sum3, diff3, diff3);
674
675        idx += 16;
676    }
677
678    // Combine the 4 accumulators (tree reduction for latency hiding)
679    let sum01 = vaddq_f32(sum0, sum1);
680    let sum23 = vaddq_f32(sum2, sum3);
681    let sum = vaddq_f32(sum01, sum23);
682
683    // Process remaining 4-float chunks
684    let remaining_start = chunks * 16;
685    let remaining_chunks = (len - remaining_start) / 4;
686    let mut final_sum = sum;
687
688    idx = remaining_start;
689    for _ in 0..remaining_chunks {
690        let va = vld1q_f32(a_ptr.add(idx));
691        let vb = vld1q_f32(b_ptr.add(idx));
692        let diff = vsubq_f32(va, vb);
693        final_sum = vfmaq_f32(final_sum, diff, diff);
694        idx += 4;
695    }
696
697    // Horizontal sum
698    let mut total = vaddvq_f32(final_sum);
699
700    // Handle remaining elements with bounds-check elimination
701    let scalar_start = remaining_start + remaining_chunks * 4;
702    for i in scalar_start..len {
703        let diff = *a.get_unchecked(i) - *b.get_unchecked(i);
704        total += diff * diff;
705    }
706
707    total.sqrt()
708}
709
710/// NEON-optimized dot product with 4x loop unrolling
711///
712/// # Safety
713/// Caller must ensure a.len() == b.len()
714#[cfg(target_arch = "aarch64")]
715#[inline(always)]
716unsafe fn dot_product_neon_unrolled_impl(a: &[f32], b: &[f32]) -> f32 {
717    debug_assert_eq!(a.len(), b.len(), "Input arrays must have the same length");
718
719    let len = a.len();
720    let a_ptr = a.as_ptr();
721    let b_ptr = b.as_ptr();
722
723    let mut sum0 = vdupq_n_f32(0.0);
724    let mut sum1 = vdupq_n_f32(0.0);
725    let mut sum2 = vdupq_n_f32(0.0);
726    let mut sum3 = vdupq_n_f32(0.0);
727
728    let chunks = len / 16;
729    let mut idx = 0usize;
730
731    for _ in 0..chunks {
732        let va0 = vld1q_f32(a_ptr.add(idx));
733        let vb0 = vld1q_f32(b_ptr.add(idx));
734        sum0 = vfmaq_f32(sum0, va0, vb0);
735
736        let va1 = vld1q_f32(a_ptr.add(idx + 4));
737        let vb1 = vld1q_f32(b_ptr.add(idx + 4));
738        sum1 = vfmaq_f32(sum1, va1, vb1);
739
740        let va2 = vld1q_f32(a_ptr.add(idx + 8));
741        let vb2 = vld1q_f32(b_ptr.add(idx + 8));
742        sum2 = vfmaq_f32(sum2, va2, vb2);
743
744        let va3 = vld1q_f32(a_ptr.add(idx + 12));
745        let vb3 = vld1q_f32(b_ptr.add(idx + 12));
746        sum3 = vfmaq_f32(sum3, va3, vb3);
747
748        idx += 16;
749    }
750
751    // Tree reduction for latency hiding
752    let sum01 = vaddq_f32(sum0, sum1);
753    let sum23 = vaddq_f32(sum2, sum3);
754    let sum = vaddq_f32(sum01, sum23);
755
756    let remaining_start = chunks * 16;
757    let remaining_chunks = (len - remaining_start) / 4;
758    let mut final_sum = sum;
759
760    idx = remaining_start;
761    for _ in 0..remaining_chunks {
762        let va = vld1q_f32(a_ptr.add(idx));
763        let vb = vld1q_f32(b_ptr.add(idx));
764        final_sum = vfmaq_f32(final_sum, va, vb);
765        idx += 4;
766    }
767
768    let mut total = vaddvq_f32(final_sum);
769
770    // Bounds-check elimination in remainder
771    let scalar_start = remaining_start + remaining_chunks * 4;
772    for i in scalar_start..len {
773        total += *a.get_unchecked(i) * *b.get_unchecked(i);
774    }
775
776    total
777}
778
779/// NEON-optimized cosine similarity with 4x loop unrolling
780///
781/// # Safety
782/// Caller must ensure a.len() == b.len()
783#[cfg(target_arch = "aarch64")]
784#[inline(always)]
785unsafe fn cosine_similarity_neon_unrolled_impl(a: &[f32], b: &[f32]) -> f32 {
786    debug_assert_eq!(a.len(), b.len(), "Input arrays must have the same length");
787
788    let len = a.len();
789    let a_ptr = a.as_ptr();
790    let b_ptr = b.as_ptr();
791
792    let mut dot0 = vdupq_n_f32(0.0);
793    let mut dot1 = vdupq_n_f32(0.0);
794    let mut norm_a0 = vdupq_n_f32(0.0);
795    let mut norm_a1 = vdupq_n_f32(0.0);
796    let mut norm_b0 = vdupq_n_f32(0.0);
797    let mut norm_b1 = vdupq_n_f32(0.0);
798
799    let chunks = len / 8;
800    let mut idx = 0usize;
801
802    for _ in 0..chunks {
803        let va0 = vld1q_f32(a_ptr.add(idx));
804        let vb0 = vld1q_f32(b_ptr.add(idx));
805        dot0 = vfmaq_f32(dot0, va0, vb0);
806        norm_a0 = vfmaq_f32(norm_a0, va0, va0);
807        norm_b0 = vfmaq_f32(norm_b0, vb0, vb0);
808
809        let va1 = vld1q_f32(a_ptr.add(idx + 4));
810        let vb1 = vld1q_f32(b_ptr.add(idx + 4));
811        dot1 = vfmaq_f32(dot1, va1, vb1);
812        norm_a1 = vfmaq_f32(norm_a1, va1, va1);
813        norm_b1 = vfmaq_f32(norm_b1, vb1, vb1);
814
815        idx += 8;
816    }
817
818    // Tree reduction
819    let dot = vaddq_f32(dot0, dot1);
820    let norm_a = vaddq_f32(norm_a0, norm_a1);
821    let norm_b = vaddq_f32(norm_b0, norm_b1);
822
823    let mut dot_sum = vaddvq_f32(dot);
824    let mut norm_a_sum = vaddvq_f32(norm_a);
825    let mut norm_b_sum = vaddvq_f32(norm_b);
826
827    // Bounds-check elimination in remainder
828    for i in (chunks * 8)..len {
829        let ai = *a.get_unchecked(i);
830        let bi = *b.get_unchecked(i);
831        dot_sum += ai * bi;
832        norm_a_sum += ai * ai;
833        norm_b_sum += bi * bi;
834    }
835
836    dot_sum / (norm_a_sum.sqrt() * norm_b_sum.sqrt())
837}
838
839/// NEON-optimized Manhattan distance with 4x loop unrolling
840///
841/// # Safety
842/// Caller must ensure a.len() == b.len()
843#[cfg(target_arch = "aarch64")]
844#[inline(always)]
845unsafe fn manhattan_distance_neon_unrolled_impl(a: &[f32], b: &[f32]) -> f32 {
846    debug_assert_eq!(a.len(), b.len(), "Input arrays must have the same length");
847
848    let len = a.len();
849    let a_ptr = a.as_ptr();
850    let b_ptr = b.as_ptr();
851
852    let mut sum0 = vdupq_n_f32(0.0);
853    let mut sum1 = vdupq_n_f32(0.0);
854    let mut sum2 = vdupq_n_f32(0.0);
855    let mut sum3 = vdupq_n_f32(0.0);
856
857    let chunks = len / 16;
858    let mut idx = 0usize;
859
860    for _ in 0..chunks {
861        // Use vabdq_f32 for absolute difference in one instruction
862        let va0 = vld1q_f32(a_ptr.add(idx));
863        let vb0 = vld1q_f32(b_ptr.add(idx));
864        sum0 = vaddq_f32(sum0, vabdq_f32(va0, vb0));
865
866        let va1 = vld1q_f32(a_ptr.add(idx + 4));
867        let vb1 = vld1q_f32(b_ptr.add(idx + 4));
868        sum1 = vaddq_f32(sum1, vabdq_f32(va1, vb1));
869
870        let va2 = vld1q_f32(a_ptr.add(idx + 8));
871        let vb2 = vld1q_f32(b_ptr.add(idx + 8));
872        sum2 = vaddq_f32(sum2, vabdq_f32(va2, vb2));
873
874        let va3 = vld1q_f32(a_ptr.add(idx + 12));
875        let vb3 = vld1q_f32(b_ptr.add(idx + 12));
876        sum3 = vaddq_f32(sum3, vabdq_f32(va3, vb3));
877
878        idx += 16;
879    }
880
881    // Tree reduction
882    let sum01 = vaddq_f32(sum0, sum1);
883    let sum23 = vaddq_f32(sum2, sum3);
884    let sum = vaddq_f32(sum01, sum23);
885
886    let remaining_start = chunks * 16;
887    let remaining_chunks = (len - remaining_start) / 4;
888    let mut final_sum = sum;
889
890    idx = remaining_start;
891    for _ in 0..remaining_chunks {
892        let va = vld1q_f32(a_ptr.add(idx));
893        let vb = vld1q_f32(b_ptr.add(idx));
894        final_sum = vaddq_f32(final_sum, vabdq_f32(va, vb));
895        idx += 4;
896    }
897
898    let mut total = vaddvq_f32(final_sum);
899
900    // Bounds-check elimination in remainder
901    let scalar_start = remaining_start + remaining_chunks * 4;
902    for i in scalar_start..len {
903        total += (*a.get_unchecked(i) - *b.get_unchecked(i)).abs();
904    }
905
906    total
907}
908
909// ============================================================================
910// Public API with architecture dispatch
911// ============================================================================
912
913/// SIMD-optimized dot product
914/// Uses AVX-512 > AVX2 on x86_64, NEON on ARM64/Apple Silicon
915#[inline(always)]
916pub fn dot_product_simd(a: &[f32], b: &[f32]) -> f32 {
917    #[cfg(target_arch = "x86_64")]
918    {
919        #[cfg(feature = "simd-avx512")]
920        {
921            if is_x86_feature_detected!("avx512f") {
922                return unsafe { dot_product_avx512_impl(a, b) };
923            }
924        }
925        if is_x86_feature_detected!("avx2") {
926            unsafe { dot_product_avx2_impl(a, b) }
927        } else {
928            dot_product_scalar(a, b)
929        }
930    }
931
932    #[cfg(target_arch = "aarch64")]
933    {
934        if a.len() >= 64 {
935            unsafe { dot_product_neon_unrolled_impl(a, b) }
936        } else {
937            unsafe { dot_product_neon_impl(a, b) }
938        }
939    }
940
941    #[cfg(not(any(target_arch = "x86_64", target_arch = "aarch64")))]
942    {
943        dot_product_scalar(a, b)
944    }
945}
946
947/// Legacy alias for backward compatibility
948#[inline(always)]
949pub fn dot_product_avx2(a: &[f32], b: &[f32]) -> f32 {
950    dot_product_simd(a, b)
951}
952
953#[cfg(target_arch = "x86_64")]
954#[target_feature(enable = "avx2")]
955unsafe fn dot_product_avx2_impl(a: &[f32], b: &[f32]) -> f32 {
956    // SECURITY: Ensure both arrays have the same length to prevent out-of-bounds access
957    assert_eq!(a.len(), b.len(), "Input arrays must have the same length");
958
959    let len = a.len();
960    let mut sum = _mm256_setzero_ps();
961
962    let chunks = len / 8;
963    for i in 0..chunks {
964        let idx = i * 8;
965        let va = _mm256_loadu_ps(a.as_ptr().add(idx));
966        let vb = _mm256_loadu_ps(b.as_ptr().add(idx));
967        let prod = _mm256_mul_ps(va, vb);
968        sum = _mm256_add_ps(sum, prod);
969    }
970
971    let sum_arr: [f32; 8] = std::mem::transmute(sum);
972    let mut total = sum_arr.iter().sum::<f32>();
973
974    for i in (chunks * 8)..len {
975        total += a[i] * b[i];
976    }
977
978    total
979}
980
981/// SIMD-optimized cosine similarity
982/// Uses AVX-512 > AVX2 on x86_64, NEON on ARM64/Apple Silicon
983#[inline(always)]
984pub fn cosine_similarity_simd(a: &[f32], b: &[f32]) -> f32 {
985    #[cfg(target_arch = "x86_64")]
986    {
987        #[cfg(feature = "simd-avx512")]
988        {
989            if is_x86_feature_detected!("avx512f") {
990                return unsafe { cosine_similarity_avx512_impl(a, b) };
991            }
992        }
993        if is_x86_feature_detected!("avx2") {
994            unsafe { cosine_similarity_avx2_impl(a, b) }
995        } else {
996            cosine_similarity_scalar(a, b)
997        }
998    }
999
1000    #[cfg(target_arch = "aarch64")]
1001    {
1002        if a.len() >= 64 {
1003            unsafe { cosine_similarity_neon_unrolled_impl(a, b) }
1004        } else {
1005            unsafe { cosine_similarity_neon_impl(a, b) }
1006        }
1007    }
1008
1009    #[cfg(not(any(target_arch = "x86_64", target_arch = "aarch64")))]
1010    {
1011        cosine_similarity_scalar(a, b)
1012    }
1013}
1014
1015/// Legacy alias for backward compatibility
1016#[inline(always)]
1017pub fn cosine_similarity_avx2(a: &[f32], b: &[f32]) -> f32 {
1018    cosine_similarity_simd(a, b)
1019}
1020
1021/// SIMD-optimized Manhattan distance
1022/// Uses AVX-512 > AVX2 on x86_64, NEON on ARM64/Apple Silicon, scalar on other platforms
1023#[inline(always)]
1024pub fn manhattan_distance_simd(a: &[f32], b: &[f32]) -> f32 {
1025    #[cfg(target_arch = "x86_64")]
1026    {
1027        #[cfg(feature = "simd-avx512")]
1028        {
1029            if is_x86_feature_detected!("avx512f") {
1030                return unsafe { manhattan_distance_avx512_impl(a, b) };
1031            }
1032        }
1033        if is_x86_feature_detected!("avx2") {
1034            unsafe { manhattan_distance_avx2_impl(a, b) }
1035        } else {
1036            manhattan_distance_scalar(a, b)
1037        }
1038    }
1039
1040    #[cfg(target_arch = "aarch64")]
1041    {
1042        if a.len() >= 64 {
1043            unsafe { manhattan_distance_neon_unrolled_impl(a, b) }
1044        } else {
1045            unsafe { manhattan_distance_neon_impl(a, b) }
1046        }
1047    }
1048
1049    #[cfg(not(any(target_arch = "x86_64", target_arch = "aarch64")))]
1050    {
1051        manhattan_distance_scalar(a, b)
1052    }
1053}
1054
1055#[cfg(target_arch = "x86_64")]
1056#[target_feature(enable = "avx2")]
1057unsafe fn cosine_similarity_avx2_impl(a: &[f32], b: &[f32]) -> f32 {
1058    // SECURITY: Ensure both arrays have the same length to prevent out-of-bounds access
1059    assert_eq!(a.len(), b.len(), "Input arrays must have the same length");
1060
1061    let len = a.len();
1062    let mut dot = _mm256_setzero_ps();
1063    let mut norm_a = _mm256_setzero_ps();
1064    let mut norm_b = _mm256_setzero_ps();
1065
1066    let chunks = len / 8;
1067    for i in 0..chunks {
1068        let idx = i * 8;
1069        let va = _mm256_loadu_ps(a.as_ptr().add(idx));
1070        let vb = _mm256_loadu_ps(b.as_ptr().add(idx));
1071
1072        // Dot product
1073        dot = _mm256_add_ps(dot, _mm256_mul_ps(va, vb));
1074
1075        // Norms
1076        norm_a = _mm256_add_ps(norm_a, _mm256_mul_ps(va, va));
1077        norm_b = _mm256_add_ps(norm_b, _mm256_mul_ps(vb, vb));
1078    }
1079
1080    let dot_arr: [f32; 8] = std::mem::transmute(dot);
1081    let norm_a_arr: [f32; 8] = std::mem::transmute(norm_a);
1082    let norm_b_arr: [f32; 8] = std::mem::transmute(norm_b);
1083
1084    let mut dot_sum = dot_arr.iter().sum::<f32>();
1085    let mut norm_a_sum = norm_a_arr.iter().sum::<f32>();
1086    let mut norm_b_sum = norm_b_arr.iter().sum::<f32>();
1087
1088    for i in (chunks * 8)..len {
1089        dot_sum += a[i] * b[i];
1090        norm_a_sum += a[i] * a[i];
1091        norm_b_sum += b[i] * b[i];
1092    }
1093
1094    dot_sum / (norm_a_sum.sqrt() * norm_b_sum.sqrt())
1095}
1096
1097/// AVX2 Manhattan distance — processes 8 floats per iteration with absolute difference
1098#[cfg(target_arch = "x86_64")]
1099#[target_feature(enable = "avx2")]
1100unsafe fn manhattan_distance_avx2_impl(a: &[f32], b: &[f32]) -> f32 {
1101    assert_eq!(a.len(), b.len(), "Input arrays must have the same length");
1102
1103    let len = a.len();
1104    // Use sign-bit mask for absolute value: clear the sign bit
1105    let sign_mask = _mm256_set1_ps(f32::from_bits(0x7FFF_FFFF));
1106    let mut sum0 = _mm256_setzero_ps();
1107    let mut sum1 = _mm256_setzero_ps();
1108
1109    // Process 16 floats at a time (2 x 8) for better ILP
1110    let chunks = len / 16;
1111    for i in 0..chunks {
1112        let idx = i * 16;
1113
1114        let va0 = _mm256_loadu_ps(a.as_ptr().add(idx));
1115        let vb0 = _mm256_loadu_ps(b.as_ptr().add(idx));
1116        let diff0 = _mm256_sub_ps(va0, vb0);
1117        let abs0 = _mm256_and_ps(diff0, sign_mask);
1118        sum0 = _mm256_add_ps(sum0, abs0);
1119
1120        let va1 = _mm256_loadu_ps(a.as_ptr().add(idx + 8));
1121        let vb1 = _mm256_loadu_ps(b.as_ptr().add(idx + 8));
1122        let diff1 = _mm256_sub_ps(va1, vb1);
1123        let abs1 = _mm256_and_ps(diff1, sign_mask);
1124        sum1 = _mm256_add_ps(sum1, abs1);
1125    }
1126
1127    let mut sum = _mm256_add_ps(sum0, sum1);
1128
1129    // Process remaining 8-float chunks
1130    let remaining_start = chunks * 16;
1131    let remaining_chunks = (len - remaining_start) / 8;
1132    for i in 0..remaining_chunks {
1133        let idx = remaining_start + i * 8;
1134        let va = _mm256_loadu_ps(a.as_ptr().add(idx));
1135        let vb = _mm256_loadu_ps(b.as_ptr().add(idx));
1136        let diff = _mm256_sub_ps(va, vb);
1137        let abs_diff = _mm256_and_ps(diff, sign_mask);
1138        sum = _mm256_add_ps(sum, abs_diff);
1139    }
1140
1141    // Horizontal sum
1142    let sum_arr: [f32; 8] = std::mem::transmute(sum);
1143    let mut total = sum_arr.iter().sum::<f32>();
1144
1145    // Handle remaining elements
1146    let scalar_start = remaining_start + remaining_chunks * 8;
1147    for i in scalar_start..len {
1148        total += (a[i] - b[i]).abs();
1149    }
1150
1151    total
1152}
1153
1154// Scalar fallback implementations
1155// These are kept for architectures without SIMD support
1156
1157#[allow(dead_code)]
1158fn euclidean_distance_scalar(a: &[f32], b: &[f32]) -> f32 {
1159    a.iter()
1160        .zip(b.iter())
1161        .map(|(x, y)| {
1162            let diff = x - y;
1163            diff * diff
1164        })
1165        .sum::<f32>()
1166        .sqrt()
1167}
1168
1169#[allow(dead_code)]
1170fn dot_product_scalar(a: &[f32], b: &[f32]) -> f32 {
1171    a.iter().zip(b.iter()).map(|(x, y)| x * y).sum()
1172}
1173
1174#[allow(dead_code)]
1175fn cosine_similarity_scalar(a: &[f32], b: &[f32]) -> f32 {
1176    let dot: f32 = a.iter().zip(b.iter()).map(|(x, y)| x * y).sum();
1177    let norm_a: f32 = a.iter().map(|x| x * x).sum::<f32>().sqrt();
1178    let norm_b: f32 = b.iter().map(|x| x * x).sum::<f32>().sqrt();
1179    dot / (norm_a * norm_b)
1180}
1181
1182#[allow(dead_code)]
1183fn manhattan_distance_scalar(a: &[f32], b: &[f32]) -> f32 {
1184    a.iter().zip(b.iter()).map(|(x, y)| (x - y).abs()).sum()
1185}
1186
1187// ============================================================================
1188// INT8 Quantized Operations
1189// ============================================================================
1190
1191/// SIMD-accelerated dot product for INT8 quantized vectors
1192/// Uses NEON vdotq_s32 on ARM64, AVX2 _mm256_maddubs_epi16 on x86_64
1193#[inline(always)]
1194pub fn dot_product_i8(a: &[i8], b: &[i8]) -> i32 {
1195    #[cfg(target_arch = "x86_64")]
1196    {
1197        if is_x86_feature_detected!("avx2") {
1198            unsafe { dot_product_i8_avx2_impl(a, b) }
1199        } else {
1200            dot_product_i8_scalar(a, b)
1201        }
1202    }
1203
1204    #[cfg(target_arch = "aarch64")]
1205    {
1206        unsafe { dot_product_i8_neon_impl(a, b) }
1207    }
1208
1209    #[cfg(not(any(target_arch = "x86_64", target_arch = "aarch64")))]
1210    {
1211        dot_product_i8_scalar(a, b)
1212    }
1213}
1214
1215/// SIMD-accelerated euclidean distance squared for INT8 quantized vectors
1216/// Returns squared distance (caller should sqrt if needed)
1217#[inline(always)]
1218pub fn euclidean_distance_squared_i8(a: &[i8], b: &[i8]) -> i32 {
1219    #[cfg(target_arch = "x86_64")]
1220    {
1221        if is_x86_feature_detected!("avx2") {
1222            unsafe { euclidean_distance_squared_i8_avx2_impl(a, b) }
1223        } else {
1224            euclidean_distance_squared_i8_scalar(a, b)
1225        }
1226    }
1227
1228    #[cfg(target_arch = "aarch64")]
1229    {
1230        unsafe { euclidean_distance_squared_i8_neon_impl(a, b) }
1231    }
1232
1233    #[cfg(not(any(target_arch = "x86_64", target_arch = "aarch64")))]
1234    {
1235        euclidean_distance_squared_i8_scalar(a, b)
1236    }
1237}
1238
1239/// NEON INT8 dot product using stable intrinsics
1240/// Note: Uses sign extension and multiply-add instead of vdotq_s32 for stability
1241///
1242/// # Safety
1243/// Caller must ensure a.len() == b.len()
1244#[cfg(target_arch = "aarch64")]
1245#[inline(always)]
1246unsafe fn dot_product_i8_neon_impl(a: &[i8], b: &[i8]) -> i32 {
1247    debug_assert_eq!(a.len(), b.len(), "Input arrays must have the same length");
1248
1249    let len = a.len();
1250    let a_ptr = a.as_ptr();
1251    let b_ptr = b.as_ptr();
1252
1253    let mut sum = vdupq_n_s32(0);
1254
1255    // Process 8 int8s at a time (extend to i16, multiply, accumulate)
1256    let chunks = len / 8;
1257    let mut idx = 0usize;
1258
1259    for _ in 0..chunks {
1260        let va = vld1_s8(a_ptr.add(idx));
1261        let vb = vld1_s8(b_ptr.add(idx));
1262
1263        // Sign-extend to i16
1264        let va_i16 = vmovl_s8(va);
1265        let vb_i16 = vmovl_s8(vb);
1266
1267        // Multiply i16 * i16
1268        let prod_lo = vmull_s16(vget_low_s16(va_i16), vget_low_s16(vb_i16));
1269        let prod_hi = vmull_s16(vget_high_s16(va_i16), vget_high_s16(vb_i16));
1270
1271        // Accumulate
1272        sum = vaddq_s32(sum, prod_lo);
1273        sum = vaddq_s32(sum, prod_hi);
1274
1275        idx += 8;
1276    }
1277
1278    // Horizontal sum
1279    let mut total = vaddvq_s32(sum);
1280
1281    // Handle remaining elements with bounds-check elimination
1282    for i in (chunks * 8)..len {
1283        total += (*a.get_unchecked(i) as i32) * (*b.get_unchecked(i) as i32);
1284    }
1285
1286    total
1287}
1288
1289/// NEON INT8 euclidean distance squared using stable intrinsics
1290///
1291/// # Safety
1292/// Caller must ensure a.len() == b.len()
1293#[cfg(target_arch = "aarch64")]
1294#[inline(always)]
1295unsafe fn euclidean_distance_squared_i8_neon_impl(a: &[i8], b: &[i8]) -> i32 {
1296    debug_assert_eq!(a.len(), b.len(), "Input arrays must have the same length");
1297
1298    let len = a.len();
1299    let a_ptr = a.as_ptr();
1300    let b_ptr = b.as_ptr();
1301
1302    let mut sum = vdupq_n_s32(0);
1303
1304    // Process 8 int8s at a time
1305    let chunks = len / 8;
1306    let mut idx = 0usize;
1307
1308    for _ in 0..chunks {
1309        let va = vld1_s8(a_ptr.add(idx));
1310        let vb = vld1_s8(b_ptr.add(idx));
1311
1312        // Sign-extend to i16
1313        let va_i16 = vmovl_s8(va);
1314        let vb_i16 = vmovl_s8(vb);
1315
1316        // Compute difference in i16
1317        let diff = vsubq_s16(va_i16, vb_i16);
1318
1319        // Square and accumulate: diff^2
1320        let prod_lo = vmull_s16(vget_low_s16(diff), vget_low_s16(diff));
1321        let prod_hi = vmull_s16(vget_high_s16(diff), vget_high_s16(diff));
1322
1323        sum = vaddq_s32(sum, prod_lo);
1324        sum = vaddq_s32(sum, prod_hi);
1325
1326        idx += 8;
1327    }
1328
1329    let mut total = vaddvq_s32(sum);
1330
1331    // Handle remaining elements with bounds-check elimination
1332    for i in (chunks * 8)..len {
1333        let diff = (*a.get_unchecked(i) as i32) - (*b.get_unchecked(i) as i32);
1334        total += diff * diff;
1335    }
1336
1337    total
1338}
1339
1340/// AVX2 INT8 dot product
1341#[cfg(target_arch = "x86_64")]
1342#[target_feature(enable = "avx2")]
1343unsafe fn dot_product_i8_avx2_impl(a: &[i8], b: &[i8]) -> i32 {
1344    assert_eq!(a.len(), b.len(), "Input arrays must have the same length");
1345
1346    let len = a.len();
1347    let mut sum = _mm256_setzero_si256();
1348
1349    // Process 32 int8s at a time
1350    let chunks = len / 32;
1351    for i in 0..chunks {
1352        let idx = i * 32;
1353        let va = _mm256_loadu_si256(a.as_ptr().add(idx) as *const __m256i);
1354        let vb = _mm256_loadu_si256(b.as_ptr().add(idx) as *const __m256i);
1355
1356        // For signed int8 multiply, we need to extend to i16 first
1357        let va_lo = _mm256_cvtepi8_epi16(_mm256_castsi256_si128(va));
1358        let vb_lo = _mm256_cvtepi8_epi16(_mm256_castsi256_si128(vb));
1359        let va_hi = _mm256_cvtepi8_epi16(_mm256_extracti128_si256(va, 1));
1360        let vb_hi = _mm256_cvtepi8_epi16(_mm256_extracti128_si256(vb, 1));
1361
1362        let prod_lo = _mm256_madd_epi16(va_lo, vb_lo);
1363        let prod_hi = _mm256_madd_epi16(va_hi, vb_hi);
1364
1365        sum = _mm256_add_epi32(sum, prod_lo);
1366        sum = _mm256_add_epi32(sum, prod_hi);
1367    }
1368
1369    // Horizontal sum
1370    let sum_arr: [i32; 8] = std::mem::transmute(sum);
1371    let mut total: i32 = sum_arr.iter().sum();
1372
1373    // Handle remaining elements
1374    for i in (chunks * 32)..len {
1375        total += (a[i] as i32) * (b[i] as i32);
1376    }
1377
1378    total
1379}
1380
1381/// AVX2 INT8 euclidean distance squared
1382#[cfg(target_arch = "x86_64")]
1383#[target_feature(enable = "avx2")]
1384unsafe fn euclidean_distance_squared_i8_avx2_impl(a: &[i8], b: &[i8]) -> i32 {
1385    assert_eq!(a.len(), b.len(), "Input arrays must have the same length");
1386
1387    let len = a.len();
1388    let mut sum = _mm256_setzero_si256();
1389
1390    let chunks = len / 32;
1391    for i in 0..chunks {
1392        let idx = i * 32;
1393        let va = _mm256_loadu_si256(a.as_ptr().add(idx) as *const __m256i);
1394        let vb = _mm256_loadu_si256(b.as_ptr().add(idx) as *const __m256i);
1395
1396        // Extend to i16, compute difference, then square
1397        let va_lo = _mm256_cvtepi8_epi16(_mm256_castsi256_si128(va));
1398        let vb_lo = _mm256_cvtepi8_epi16(_mm256_castsi256_si128(vb));
1399        let va_hi = _mm256_cvtepi8_epi16(_mm256_extracti128_si256(va, 1));
1400        let vb_hi = _mm256_cvtepi8_epi16(_mm256_extracti128_si256(vb, 1));
1401
1402        let diff_lo = _mm256_sub_epi16(va_lo, vb_lo);
1403        let diff_hi = _mm256_sub_epi16(va_hi, vb_hi);
1404
1405        let sq_lo = _mm256_madd_epi16(diff_lo, diff_lo);
1406        let sq_hi = _mm256_madd_epi16(diff_hi, diff_hi);
1407
1408        sum = _mm256_add_epi32(sum, sq_lo);
1409        sum = _mm256_add_epi32(sum, sq_hi);
1410    }
1411
1412    let sum_arr: [i32; 8] = std::mem::transmute(sum);
1413    let mut total: i32 = sum_arr.iter().sum();
1414
1415    for i in (chunks * 32)..len {
1416        let diff = (a[i] as i32) - (b[i] as i32);
1417        total += diff * diff;
1418    }
1419
1420    total
1421}
1422
1423/// Scalar fallback for INT8 dot product
1424#[allow(dead_code)]
1425fn dot_product_i8_scalar(a: &[i8], b: &[i8]) -> i32 {
1426    a.iter()
1427        .zip(b.iter())
1428        .map(|(&x, &y)| (x as i32) * (y as i32))
1429        .sum()
1430}
1431
1432/// Scalar fallback for INT8 euclidean distance squared
1433#[allow(dead_code)]
1434fn euclidean_distance_squared_i8_scalar(a: &[i8], b: &[i8]) -> i32 {
1435    a.iter()
1436        .zip(b.iter())
1437        .map(|(&x, &y)| {
1438            let diff = (x as i32) - (y as i32);
1439            diff * diff
1440        })
1441        .sum()
1442}
1443
1444// ============================================================================
1445// Batch Operations (Cache-optimized)
1446// ============================================================================
1447
1448/// Batch dot product - compute dot products of one query vector against multiple vectors
1449/// Returns results in the provided output slice
1450/// Optimized for cache locality by processing vectors in tiles
1451#[inline]
1452pub fn batch_dot_product(query: &[f32], vectors: &[&[f32]], results: &mut [f32]) {
1453    assert_eq!(
1454        vectors.len(),
1455        results.len(),
1456        "Output size must match vector count"
1457    );
1458
1459    // Process in tiles for better cache utilization
1460    const TILE_SIZE: usize = 16;
1461
1462    for (chunk_idx, chunk) in vectors.chunks(TILE_SIZE).enumerate() {
1463        let base_idx = chunk_idx * TILE_SIZE;
1464        for (i, vec) in chunk.iter().enumerate() {
1465            results[base_idx + i] = dot_product_simd(query, vec);
1466        }
1467    }
1468}
1469
1470/// Batch euclidean distance - compute distances from one query to multiple vectors
1471/// Returns results in the provided output slice
1472/// Optimized for cache locality
1473#[inline]
1474pub fn batch_euclidean(query: &[f32], vectors: &[&[f32]], results: &mut [f32]) {
1475    assert_eq!(
1476        vectors.len(),
1477        results.len(),
1478        "Output size must match vector count"
1479    );
1480
1481    const TILE_SIZE: usize = 16;
1482
1483    for (chunk_idx, chunk) in vectors.chunks(TILE_SIZE).enumerate() {
1484        let base_idx = chunk_idx * TILE_SIZE;
1485        for (i, vec) in chunk.iter().enumerate() {
1486            results[base_idx + i] = euclidean_distance_simd(query, vec);
1487        }
1488    }
1489}
1490
1491/// Batch cosine similarity - compute similarities from one query to multiple vectors
1492#[inline]
1493pub fn batch_cosine_similarity(query: &[f32], vectors: &[&[f32]], results: &mut [f32]) {
1494    assert_eq!(
1495        vectors.len(),
1496        results.len(),
1497        "Output size must match vector count"
1498    );
1499
1500    const TILE_SIZE: usize = 16;
1501
1502    for (chunk_idx, chunk) in vectors.chunks(TILE_SIZE).enumerate() {
1503        let base_idx = chunk_idx * TILE_SIZE;
1504        for (i, vec) in chunk.iter().enumerate() {
1505            results[base_idx + i] = cosine_similarity_simd(query, vec);
1506        }
1507    }
1508}
1509
1510/// Batch dot product with owned vectors (for convenience)
1511#[inline]
1512pub fn batch_dot_product_owned(query: &[f32], vectors: &[Vec<f32>]) -> Vec<f32> {
1513    let refs: Vec<&[f32]> = vectors.iter().map(|v| v.as_slice()).collect();
1514    let mut results = vec![0.0; vectors.len()];
1515    batch_dot_product(query, &refs, &mut results);
1516    results
1517}
1518
1519/// Batch euclidean distance with owned vectors (for convenience)
1520#[inline]
1521pub fn batch_euclidean_owned(query: &[f32], vectors: &[Vec<f32>]) -> Vec<f32> {
1522    let refs: Vec<&[f32]> = vectors.iter().map(|v| v.as_slice()).collect();
1523    let mut results = vec![0.0; vectors.len()];
1524    batch_euclidean(query, &refs, &mut results);
1525    results
1526}
1527
1528#[cfg(test)]
1529mod tests {
1530    use super::*;
1531
1532    #[test]
1533    fn test_euclidean_distance_simd() {
1534        let a = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0];
1535        let b = vec![2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0];
1536
1537        let result = euclidean_distance_simd(&a, &b);
1538        let expected = euclidean_distance_scalar(&a, &b);
1539
1540        assert!(
1541            (result - expected).abs() < 0.001,
1542            "SIMD result {} differs from scalar result {}",
1543            result,
1544            expected
1545        );
1546    }
1547
1548    #[test]
1549    fn test_euclidean_distance_large() {
1550        // Test with 128-dim vectors (common embedding size)
1551        let a: Vec<f32> = (0..128).map(|i| i as f32 * 0.1).collect();
1552        let b: Vec<f32> = (0..128).map(|i| (i as f32 * 0.1) + 0.5).collect();
1553
1554        let result = euclidean_distance_simd(&a, &b);
1555        let expected = euclidean_distance_scalar(&a, &b);
1556
1557        assert!(
1558            (result - expected).abs() < 0.01,
1559            "Large vector: SIMD {} vs scalar {}",
1560            result,
1561            expected
1562        );
1563    }
1564
1565    #[test]
1566    fn test_dot_product_simd() {
1567        let a = vec![1.0; 16];
1568        let b = vec![2.0; 16];
1569
1570        let result = dot_product_simd(&a, &b);
1571        assert!((result - 32.0).abs() < 0.001);
1572    }
1573
1574    #[test]
1575    fn test_dot_product_large() {
1576        let a: Vec<f32> = (0..256).map(|i| (i % 10) as f32).collect();
1577        let b: Vec<f32> = (0..256).map(|i| ((i + 5) % 10) as f32).collect();
1578
1579        let result = dot_product_simd(&a, &b);
1580        let expected = dot_product_scalar(&a, &b);
1581
1582        assert!(
1583            (result - expected).abs() < 0.1,
1584            "Large dot product: SIMD {} vs scalar {}",
1585            result,
1586            expected
1587        );
1588    }
1589
1590    #[test]
1591    fn test_cosine_similarity_simd() {
1592        let a = vec![1.0, 0.0, 0.0];
1593        let b = vec![1.0, 0.0, 0.0];
1594
1595        let result = cosine_similarity_simd(&a, &b);
1596        assert!((result - 1.0).abs() < 0.001);
1597    }
1598
1599    #[test]
1600    fn test_cosine_similarity_orthogonal() {
1601        let a = vec![1.0, 0.0, 0.0, 0.0];
1602        let b = vec![0.0, 1.0, 0.0, 0.0];
1603
1604        let result = cosine_similarity_simd(&a, &b);
1605        assert!(
1606            result.abs() < 0.001,
1607            "Orthogonal vectors should have ~0 similarity, got {}",
1608            result
1609        );
1610    }
1611
1612    #[test]
1613    fn test_manhattan_distance_simd() {
1614        let a = vec![1.0, 2.0, 3.0, 4.0];
1615        let b = vec![5.0, 6.0, 7.0, 8.0];
1616
1617        let result = manhattan_distance_simd(&a, &b);
1618        let expected = manhattan_distance_scalar(&a, &b);
1619
1620        assert!(
1621            (result - expected).abs() < 0.001,
1622            "Manhattan: SIMD {} vs scalar {}",
1623            result,
1624            expected
1625        );
1626        assert!((result - 16.0).abs() < 0.001); // |4| + |4| + |4| + |4| = 16
1627    }
1628
1629    #[test]
1630    fn test_non_aligned_lengths() {
1631        // Test vectors not aligned to SIMD width (4 for NEON, 8 for AVX2)
1632        let a = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0]; // 7 elements
1633        let b = vec![2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0];
1634
1635        let result = euclidean_distance_simd(&a, &b);
1636        let expected = euclidean_distance_scalar(&a, &b);
1637
1638        assert!(
1639            (result - expected).abs() < 0.001,
1640            "Non-aligned: SIMD {} vs scalar {}",
1641            result,
1642            expected
1643        );
1644    }
1645
1646    // Legacy function tests (ensure backward compatibility)
1647    #[test]
1648    fn test_legacy_avx2_aliases() {
1649        let a = vec![1.0, 2.0, 3.0, 4.0];
1650        let b = vec![5.0, 6.0, 7.0, 8.0];
1651
1652        // These should work identically to the _simd versions
1653        let _ = euclidean_distance_avx2(&a, &b);
1654        let _ = dot_product_avx2(&a, &b);
1655        let _ = cosine_similarity_avx2(&a, &b);
1656    }
1657
1658    // INT8 quantized operation tests
1659    #[test]
1660    fn test_dot_product_i8() {
1661        let a: Vec<i8> = vec![1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16];
1662        let b: Vec<i8> = vec![2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17];
1663
1664        let result = dot_product_i8(&a, &b);
1665        let expected = dot_product_i8_scalar(&a, &b);
1666
1667        assert_eq!(
1668            result, expected,
1669            "INT8 dot product: SIMD {} vs scalar {}",
1670            result, expected
1671        );
1672    }
1673
1674    #[test]
1675    fn test_dot_product_i8_large() {
1676        // Test with 128 elements (common for quantized embeddings)
1677        let a: Vec<i8> = (0..128)
1678            .map(|i| ((i % 256) as i8).wrapping_sub(64))
1679            .collect();
1680        let b: Vec<i8> = (0..128)
1681            .map(|i| (((i + 10) % 256) as i8).wrapping_sub(64))
1682            .collect();
1683
1684        let result = dot_product_i8(&a, &b);
1685        let expected = dot_product_i8_scalar(&a, &b);
1686
1687        assert_eq!(
1688            result, expected,
1689            "Large INT8 dot product: SIMD {} vs scalar {}",
1690            result, expected
1691        );
1692    }
1693
1694    #[test]
1695    fn test_euclidean_distance_squared_i8() {
1696        let a: Vec<i8> = vec![1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16];
1697        let b: Vec<i8> = vec![2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17];
1698
1699        let result = euclidean_distance_squared_i8(&a, &b);
1700        let expected = euclidean_distance_squared_i8_scalar(&a, &b);
1701
1702        assert_eq!(
1703            result, expected,
1704            "INT8 euclidean^2: SIMD {} vs scalar {}",
1705            result, expected
1706        );
1707        // Each diff is 1, so 16 diffs squared = 16
1708        assert_eq!(result, 16, "Expected 16, got {}", result);
1709    }
1710
1711    #[test]
1712    fn test_euclidean_distance_squared_i8_large() {
1713        let a: Vec<i8> = (0..128)
1714            .map(|i| ((i % 256) as i8).wrapping_sub(64))
1715            .collect();
1716        let b: Vec<i8> = (0..128)
1717            .map(|i| (((i + 5) % 256) as i8).wrapping_sub(64))
1718            .collect();
1719
1720        let result = euclidean_distance_squared_i8(&a, &b);
1721        let expected = euclidean_distance_squared_i8_scalar(&a, &b);
1722
1723        assert_eq!(
1724            result, expected,
1725            "Large INT8 euclidean^2: SIMD {} vs scalar {}",
1726            result, expected
1727        );
1728    }
1729
1730    // Batch operation tests
1731    #[test]
1732    fn test_batch_dot_product() {
1733        let query = vec![1.0, 2.0, 3.0, 4.0];
1734        let v1 = vec![1.0, 0.0, 0.0, 0.0];
1735        let v2 = vec![0.0, 1.0, 0.0, 0.0];
1736        let v3 = vec![0.0, 0.0, 1.0, 0.0];
1737        let vectors: Vec<&[f32]> = vec![&v1, &v2, &v3];
1738        let mut results = vec![0.0; 3];
1739
1740        batch_dot_product(&query, &vectors, &mut results);
1741
1742        assert!((results[0] - 1.0).abs() < 0.001);
1743        assert!((results[1] - 2.0).abs() < 0.001);
1744        assert!((results[2] - 3.0).abs() < 0.001);
1745    }
1746
1747    #[test]
1748    fn test_batch_euclidean() {
1749        let query = vec![0.0, 0.0, 0.0, 0.0];
1750        let v1 = vec![3.0, 4.0, 0.0, 0.0];
1751        let v2 = vec![0.0, 0.0, 5.0, 12.0];
1752        let vectors: Vec<&[f32]> = vec![&v1, &v2];
1753        let mut results = vec![0.0; 2];
1754
1755        batch_euclidean(&query, &vectors, &mut results);
1756
1757        assert!(
1758            (results[0] - 5.0).abs() < 0.001,
1759            "Expected 5.0, got {}",
1760            results[0]
1761        );
1762        assert!(
1763            (results[1] - 13.0).abs() < 0.001,
1764            "Expected 13.0, got {}",
1765            results[1]
1766        );
1767    }
1768
1769    #[test]
1770    fn test_batch_cosine_similarity() {
1771        let query = vec![1.0, 0.0, 0.0, 0.0];
1772        let v1 = vec![1.0, 0.0, 0.0, 0.0]; // Same direction
1773        let v2 = vec![0.0, 1.0, 0.0, 0.0]; // Orthogonal
1774        let v3 = vec![-1.0, 0.0, 0.0, 0.0]; // Opposite
1775        let vectors: Vec<&[f32]> = vec![&v1, &v2, &v3];
1776        let mut results = vec![0.0; 3];
1777
1778        batch_cosine_similarity(&query, &vectors, &mut results);
1779
1780        assert!(
1781            (results[0] - 1.0).abs() < 0.001,
1782            "Same direction should be 1.0"
1783        );
1784        assert!(results[1].abs() < 0.001, "Orthogonal should be 0.0");
1785        assert!((results[2] + 1.0).abs() < 0.001, "Opposite should be -1.0");
1786    }
1787
1788    #[test]
1789    fn test_batch_owned_convenience() {
1790        let query = vec![1.0, 2.0, 3.0, 4.0];
1791        let vectors = vec![vec![1.0, 0.0, 0.0, 0.0], vec![0.0, 1.0, 0.0, 0.0]];
1792
1793        let results = batch_dot_product_owned(&query, &vectors);
1794        assert_eq!(results.len(), 2);
1795        assert!((results[0] - 1.0).abs() < 0.001);
1796        assert!((results[1] - 2.0).abs() < 0.001);
1797    }
1798
1799    #[test]
1800    fn test_unrolled_vs_non_unrolled_consistency() {
1801        // Test that unrolled and non-unrolled implementations produce same results
1802        let a: Vec<f32> = (0..128).map(|i| i as f32 * 0.1).collect();
1803        let b: Vec<f32> = (0..128).map(|i| (i as f32 * 0.1) + 0.5).collect();
1804
1805        let result = euclidean_distance_simd(&a, &b);
1806        let expected = euclidean_distance_scalar(&a, &b);
1807
1808        assert!(
1809            (result - expected).abs() < 0.01,
1810            "Unrolled consistency: SIMD {} vs scalar {}",
1811            result,
1812            expected
1813        );
1814    }
1815}