1#[cfg(target_arch = "x86_64")]
25use std::arch::x86_64::*;
26
27#[cfg(target_arch = "aarch64")]
28use std::arch::aarch64::*;
29
30#[allow(dead_code)]
32const PREFETCH_DISTANCE: usize = 64;
33
34#[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 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#[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 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 let chunks = len / 8;
93 for i in 0..chunks {
94 let idx = i * 8;
95
96 let va = _mm256_loadu_ps(a.as_ptr().add(idx));
98 let vb = _mm256_loadu_ps(b.as_ptr().add(idx));
99
100 let diff = _mm256_sub_ps(va, vb);
102
103 let sq = _mm256_mul_ps(diff, diff);
105
106 sum = _mm256_add_ps(sum, sq);
108 }
109
110 let sum_arr: [f32; 8] = std::mem::transmute(sum);
112 let mut total = sum_arr.iter().sum::<f32>();
113
114 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#[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 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 let chunks = len / 32;
138 for i in 0..chunks {
139 let idx = i * 32;
140
141 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 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 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 let sum_arr: [f32; 8] = std::mem::transmute(final_sum);
182 let mut total = sum_arr.iter().sum::<f32>();
183
184 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#[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 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 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 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 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 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#[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#[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 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 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 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#[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#[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 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 let diff = vsubq_f32(va, vb);
478
479 sum = vfmaq_f32(sum, diff, diff);
481
482 idx += 4;
483 }
484
485 let mut total = vaddvq_f32(sum);
487
488 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#[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 sum = vfmaq_f32(sum, va, vb);
521
522 idx += 4;
523 }
524
525 let mut total = vaddvq_f32(sum);
526
527 for i in (chunks * 4)..len {
529 total += *a.get_unchecked(i) * *b.get_unchecked(i);
530 }
531
532 total
533}
534
535#[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 = vfmaq_f32(dot, va, vb);
561
562 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 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#[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 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 for i in (chunks * 4)..len {
618 total += (*a.get_unchecked(i) - *b.get_unchecked(i)).abs();
619 }
620
621 total
622}
623
624#[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 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 let chunks = len / 16;
651 let mut idx = 0usize;
652
653 for _ in 0..chunks {
654 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 let sum01 = vaddq_f32(sum0, sum1);
680 let sum23 = vaddq_f32(sum2, sum3);
681 let sum = vaddq_f32(sum01, sum23);
682
683 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 let mut total = vaddvq_f32(final_sum);
699
700 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#[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 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 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#[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 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 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#[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 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 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 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#[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#[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 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#[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#[inline(always)]
1017pub fn cosine_similarity_avx2(a: &[f32], b: &[f32]) -> f32 {
1018 cosine_similarity_simd(a, b)
1019}
1020
1021#[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 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 = _mm256_add_ps(dot, _mm256_mul_ps(va, vb));
1074
1075 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#[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 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 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 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 let sum_arr: [f32; 8] = std::mem::transmute(sum);
1143 let mut total = sum_arr.iter().sum::<f32>();
1144
1145 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#[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#[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#[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#[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 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 let va_i16 = vmovl_s8(va);
1265 let vb_i16 = vmovl_s8(vb);
1266
1267 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 sum = vaddq_s32(sum, prod_lo);
1273 sum = vaddq_s32(sum, prod_hi);
1274
1275 idx += 8;
1276 }
1277
1278 let mut total = vaddvq_s32(sum);
1280
1281 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#[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 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 let va_i16 = vmovl_s8(va);
1314 let vb_i16 = vmovl_s8(vb);
1315
1316 let diff = vsubq_s16(va_i16, vb_i16);
1318
1319 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 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#[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 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 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 let sum_arr: [i32; 8] = std::mem::transmute(sum);
1371 let mut total: i32 = sum_arr.iter().sum();
1372
1373 for i in (chunks * 32)..len {
1375 total += (a[i] as i32) * (b[i] as i32);
1376 }
1377
1378 total
1379}
1380
1381#[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 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#[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#[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#[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 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#[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#[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#[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#[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 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); }
1628
1629 #[test]
1630 fn test_non_aligned_lengths() {
1631 let a = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0]; 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 #[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 let _ = euclidean_distance_avx2(&a, &b);
1654 let _ = dot_product_avx2(&a, &b);
1655 let _ = cosine_similarity_avx2(&a, &b);
1656 }
1657
1658 #[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 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 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 #[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]; let v2 = vec![0.0, 1.0, 0.0, 0.0]; let v3 = vec![-1.0, 0.0, 0.0, 0.0]; 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 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}