1#[cfg(target_arch = "aarch64")]
19#[allow(unsafe_op_in_unsafe_fn)]
20mod neon {
21 use std::arch::aarch64::*;
22
23 #[target_feature(enable = "neon")]
25 pub unsafe fn unpack_8bit(input: &[u8], output: &mut [u32], count: usize) {
26 let chunks = count / 16;
27 let remainder = count % 16;
28
29 for chunk in 0..chunks {
30 let base = chunk * 16;
31 let in_ptr = input.as_ptr().add(base);
32
33 let bytes = vld1q_u8(in_ptr);
35
36 let low8 = vget_low_u8(bytes);
38 let high8 = vget_high_u8(bytes);
39
40 let low16 = vmovl_u8(low8);
41 let high16 = vmovl_u8(high8);
42
43 let v0 = vmovl_u16(vget_low_u16(low16));
44 let v1 = vmovl_u16(vget_high_u16(low16));
45 let v2 = vmovl_u16(vget_low_u16(high16));
46 let v3 = vmovl_u16(vget_high_u16(high16));
47
48 let out_ptr = output.as_mut_ptr().add(base);
49 vst1q_u32(out_ptr, v0);
50 vst1q_u32(out_ptr.add(4), v1);
51 vst1q_u32(out_ptr.add(8), v2);
52 vst1q_u32(out_ptr.add(12), v3);
53 }
54
55 let base = chunks * 16;
57 for i in 0..remainder {
58 output[base + i] = input[base + i] as u32;
59 }
60 }
61
62 #[target_feature(enable = "neon")]
64 pub unsafe fn unpack_16bit(input: &[u8], output: &mut [u32], count: usize) {
65 let chunks = count / 8;
66 let remainder = count % 8;
67
68 for chunk in 0..chunks {
69 let base = chunk * 8;
70 let in_ptr = input.as_ptr().add(base * 2) as *const u16;
71
72 let vals = vld1q_u16(in_ptr);
73 let low = vmovl_u16(vget_low_u16(vals));
74 let high = vmovl_u16(vget_high_u16(vals));
75
76 let out_ptr = output.as_mut_ptr().add(base);
77 vst1q_u32(out_ptr, low);
78 vst1q_u32(out_ptr.add(4), high);
79 }
80
81 let base = chunks * 8;
83 for i in 0..remainder {
84 let idx = (base + i) * 2;
85 output[base + i] = u16::from_le_bytes([input[idx], input[idx + 1]]) as u32;
86 }
87 }
88
89 #[target_feature(enable = "neon")]
91 pub unsafe fn unpack_32bit(input: &[u8], output: &mut [u32], count: usize) {
92 let chunks = count / 4;
93 let remainder = count % 4;
94
95 let in_ptr = input.as_ptr() as *const u32;
96 let out_ptr = output.as_mut_ptr();
97
98 for chunk in 0..chunks {
99 let vals = vld1q_u32(in_ptr.add(chunk * 4));
100 vst1q_u32(out_ptr.add(chunk * 4), vals);
101 }
102
103 let base = chunks * 4;
105 for i in 0..remainder {
106 let idx = (base + i) * 4;
107 output[base + i] =
108 u32::from_le_bytes([input[idx], input[idx + 1], input[idx + 2], input[idx + 3]]);
109 }
110 }
111
112 #[inline]
116 #[target_feature(enable = "neon")]
117 unsafe fn prefix_sum_4(v: uint32x4_t) -> uint32x4_t {
118 let shifted1 = vextq_u32(vdupq_n_u32(0), v, 3);
121 let sum1 = vaddq_u32(v, shifted1);
122
123 let shifted2 = vextq_u32(vdupq_n_u32(0), sum1, 2);
126 vaddq_u32(sum1, shifted2)
127 }
128
129 #[target_feature(enable = "neon")]
133 pub unsafe fn delta_decode(
134 output: &mut [u32],
135 deltas: &[u32],
136 first_doc_id: u32,
137 count: usize,
138 ) {
139 if count == 0 {
140 return;
141 }
142
143 output[0] = first_doc_id;
144 if count == 1 {
145 return;
146 }
147
148 let ones = vdupq_n_u32(1);
149 let mut carry = vdupq_n_u32(first_doc_id);
150
151 let full_groups = (count - 1) / 4;
152 let remainder = (count - 1) % 4;
153
154 for group in 0..full_groups {
155 let base = group * 4;
156
157 let d = vld1q_u32(deltas[base..].as_ptr());
159 let gaps = vaddq_u32(d, ones);
160
161 let prefix = prefix_sum_4(gaps);
163
164 let result = vaddq_u32(prefix, carry);
166
167 vst1q_u32(output[base + 1..].as_mut_ptr(), result);
169
170 carry = vdupq_n_u32(vgetq_lane_u32(result, 3));
172 }
173
174 let base = full_groups * 4;
176 let mut scalar_carry = vgetq_lane_u32(carry, 0);
177 for j in 0..remainder {
178 scalar_carry = scalar_carry.wrapping_add(deltas[base + j]).wrapping_add(1);
179 output[base + j + 1] = scalar_carry;
180 }
181 }
182
183 #[target_feature(enable = "neon")]
185 pub unsafe fn add_one(values: &mut [u32], count: usize) {
186 let ones = vdupq_n_u32(1);
187 let chunks = count / 4;
188 let remainder = count % 4;
189
190 for chunk in 0..chunks {
191 let base = chunk * 4;
192 let ptr = values.as_mut_ptr().add(base);
193 let v = vld1q_u32(ptr);
194 let result = vaddq_u32(v, ones);
195 vst1q_u32(ptr, result);
196 }
197
198 let base = chunks * 4;
199 for i in 0..remainder {
200 values[base + i] += 1;
201 }
202 }
203
204 #[target_feature(enable = "neon")]
207 pub unsafe fn unpack_8bit_delta_decode(
208 input: &[u8],
209 output: &mut [u32],
210 first_value: u32,
211 count: usize,
212 ) {
213 output[0] = first_value;
214 if count <= 1 {
215 return;
216 }
217
218 let ones = vdupq_n_u32(1);
219 let mut carry = vdupq_n_u32(first_value);
220
221 let full_groups = (count - 1) / 4;
222 let remainder = (count - 1) % 4;
223
224 for group in 0..full_groups {
225 let base = group * 4;
226
227 let raw = std::ptr::read_unaligned(input.as_ptr().add(base) as *const u32);
229 let bytes = vreinterpret_u8_u32(vdup_n_u32(raw));
230 let u16s = vmovl_u8(bytes); let d = vmovl_u16(vget_low_u16(u16s)); let gaps = vaddq_u32(d, ones);
235
236 let prefix = prefix_sum_4(gaps);
238
239 let result = vaddq_u32(prefix, carry);
241
242 vst1q_u32(output[base + 1..].as_mut_ptr(), result);
244
245 carry = vdupq_n_u32(vgetq_lane_u32(result, 3));
247 }
248
249 let base = full_groups * 4;
251 let mut scalar_carry = vgetq_lane_u32(carry, 0);
252 for j in 0..remainder {
253 scalar_carry = scalar_carry
254 .wrapping_add(input[base + j] as u32)
255 .wrapping_add(1);
256 output[base + j + 1] = scalar_carry;
257 }
258 }
259
260 #[target_feature(enable = "neon")]
262 pub unsafe fn unpack_16bit_delta_decode(
263 input: &[u8],
264 output: &mut [u32],
265 first_value: u32,
266 count: usize,
267 ) {
268 output[0] = first_value;
269 if count <= 1 {
270 return;
271 }
272
273 let ones = vdupq_n_u32(1);
274 let mut carry = vdupq_n_u32(first_value);
275
276 let full_groups = (count - 1) / 4;
277 let remainder = (count - 1) % 4;
278
279 for group in 0..full_groups {
280 let base = group * 4;
281 let in_ptr = input.as_ptr().add(base * 2) as *const u16;
282
283 let vals = vld1_u16(in_ptr);
285 let d = vmovl_u16(vals);
286
287 let gaps = vaddq_u32(d, ones);
289
290 let prefix = prefix_sum_4(gaps);
292
293 let result = vaddq_u32(prefix, carry);
295
296 vst1q_u32(output[base + 1..].as_mut_ptr(), result);
298
299 carry = vdupq_n_u32(vgetq_lane_u32(result, 3));
301 }
302
303 let base = full_groups * 4;
305 let mut scalar_carry = vgetq_lane_u32(carry, 0);
306 for j in 0..remainder {
307 let idx = (base + j) * 2;
308 let delta = u16::from_le_bytes([input[idx], input[idx + 1]]) as u32;
309 scalar_carry = scalar_carry.wrapping_add(delta).wrapping_add(1);
310 output[base + j + 1] = scalar_carry;
311 }
312 }
313
314 #[target_feature(enable = "neon")]
317 pub unsafe fn hamming_distance(a: &[u8], b: &[u8]) -> u32 {
318 let len = a.len();
319 let chunks16 = len / 16;
320 let mut total = 0u32;
321
322 let mut i = 0;
325 while i < chunks16 {
326 let batch_end = (i + 31).min(chunks16);
327 let mut acc = vdupq_n_u8(0);
328 for j in i..batch_end {
329 let off = j * 16;
330 let va = vld1q_u8(a.as_ptr().add(off));
331 let vb = vld1q_u8(b.as_ptr().add(off));
332 let popcnt = vcntq_u8(veorq_u8(va, vb));
333 acc = vaddq_u8(acc, popcnt);
334 }
335 let sum64 = vpaddlq_u32(vpaddlq_u16(vpaddlq_u8(acc)));
337 total += vgetq_lane_u64(sum64, 0) as u32 + vgetq_lane_u64(sum64, 1) as u32;
338 i = batch_end;
339 }
340
341 let base = chunks16 * 16;
343 for k in base..len {
344 total += (a[k] ^ b[k]).count_ones();
345 }
346
347 total
348 }
349
350 #[target_feature(enable = "neon")]
356 pub unsafe fn hamming_distance_x4(query: &[u8], rows: [&[u8]; 4]) -> [u32; 4] {
357 let len = query.len();
358 let chunks16 = len / 16;
359 let mut total = [0u32; 4];
360
361 let mut i = 0;
362 while i < chunks16 {
363 let batch_end = (i + 31).min(chunks16);
364 let mut acc = [vdupq_n_u8(0); 4];
365 for j in i..batch_end {
366 let off = j * 16;
367 let vq = vld1q_u8(query.as_ptr().add(off));
368 for r in 0..4 {
369 let vr = vld1q_u8(rows[r].as_ptr().add(off));
370 acc[r] = vaddq_u8(acc[r], vcntq_u8(veorq_u8(vq, vr)));
371 }
372 }
373 for r in 0..4 {
374 let sum64 = vpaddlq_u32(vpaddlq_u16(vpaddlq_u8(acc[r])));
375 total[r] += vgetq_lane_u64(sum64, 0) as u32 + vgetq_lane_u64(sum64, 1) as u32;
376 }
377 i = batch_end;
378 }
379
380 let base = chunks16 * 16;
383 if base < len {
384 let tail = &query[base..];
385 for r in 0..4 {
386 total[r] += super::hamming_distance_scalar(tail, &rows[r][base..]);
387 }
388 }
389
390 total
391 }
392
393 #[inline]
395 pub fn is_available() -> bool {
396 true
397 }
398}
399
400#[cfg(target_arch = "x86_64")]
405#[allow(unsafe_op_in_unsafe_fn)]
406mod sse {
407 use std::arch::x86_64::*;
408
409 #[target_feature(enable = "sse2", enable = "sse4.1")]
411 pub unsafe fn unpack_8bit(input: &[u8], output: &mut [u32], count: usize) {
412 let chunks = count / 16;
413 let remainder = count % 16;
414
415 for chunk in 0..chunks {
416 let base = chunk * 16;
417 let in_ptr = input.as_ptr().add(base);
418
419 let bytes = _mm_loadu_si128(in_ptr as *const __m128i);
420
421 let v0 = _mm_cvtepu8_epi32(bytes);
423 let v1 = _mm_cvtepu8_epi32(_mm_srli_si128(bytes, 4));
424 let v2 = _mm_cvtepu8_epi32(_mm_srli_si128(bytes, 8));
425 let v3 = _mm_cvtepu8_epi32(_mm_srli_si128(bytes, 12));
426
427 let out_ptr = output.as_mut_ptr().add(base);
428 _mm_storeu_si128(out_ptr as *mut __m128i, v0);
429 _mm_storeu_si128(out_ptr.add(4) as *mut __m128i, v1);
430 _mm_storeu_si128(out_ptr.add(8) as *mut __m128i, v2);
431 _mm_storeu_si128(out_ptr.add(12) as *mut __m128i, v3);
432 }
433
434 let base = chunks * 16;
435 for i in 0..remainder {
436 output[base + i] = input[base + i] as u32;
437 }
438 }
439
440 #[target_feature(enable = "sse2", enable = "sse4.1")]
442 pub unsafe fn unpack_16bit(input: &[u8], output: &mut [u32], count: usize) {
443 let chunks = count / 8;
444 let remainder = count % 8;
445
446 for chunk in 0..chunks {
447 let base = chunk * 8;
448 let in_ptr = input.as_ptr().add(base * 2);
449
450 let vals = _mm_loadu_si128(in_ptr as *const __m128i);
451 let low = _mm_cvtepu16_epi32(vals);
452 let high = _mm_cvtepu16_epi32(_mm_srli_si128(vals, 8));
453
454 let out_ptr = output.as_mut_ptr().add(base);
455 _mm_storeu_si128(out_ptr as *mut __m128i, low);
456 _mm_storeu_si128(out_ptr.add(4) as *mut __m128i, high);
457 }
458
459 let base = chunks * 8;
460 for i in 0..remainder {
461 let idx = (base + i) * 2;
462 output[base + i] = u16::from_le_bytes([input[idx], input[idx + 1]]) as u32;
463 }
464 }
465
466 #[target_feature(enable = "sse2")]
468 pub unsafe fn unpack_32bit(input: &[u8], output: &mut [u32], count: usize) {
469 let chunks = count / 4;
470 let remainder = count % 4;
471
472 let in_ptr = input.as_ptr() as *const __m128i;
473 let out_ptr = output.as_mut_ptr() as *mut __m128i;
474
475 for chunk in 0..chunks {
476 let vals = _mm_loadu_si128(in_ptr.add(chunk));
477 _mm_storeu_si128(out_ptr.add(chunk), vals);
478 }
479
480 let base = chunks * 4;
482 for i in 0..remainder {
483 let idx = (base + i) * 4;
484 output[base + i] =
485 u32::from_le_bytes([input[idx], input[idx + 1], input[idx + 2], input[idx + 3]]);
486 }
487 }
488
489 #[inline]
493 #[target_feature(enable = "sse2")]
494 unsafe fn prefix_sum_4(v: __m128i) -> __m128i {
495 let shifted1 = _mm_slli_si128(v, 4);
498 let sum1 = _mm_add_epi32(v, shifted1);
499
500 let shifted2 = _mm_slli_si128(sum1, 8);
503 _mm_add_epi32(sum1, shifted2)
504 }
505
506 #[target_feature(enable = "sse2", enable = "sse4.1")]
508 pub unsafe fn delta_decode(
509 output: &mut [u32],
510 deltas: &[u32],
511 first_doc_id: u32,
512 count: usize,
513 ) {
514 if count == 0 {
515 return;
516 }
517
518 output[0] = first_doc_id;
519 if count == 1 {
520 return;
521 }
522
523 let ones = _mm_set1_epi32(1);
524 let mut carry = _mm_set1_epi32(first_doc_id as i32);
525
526 let full_groups = (count - 1) / 4;
527 let remainder = (count - 1) % 4;
528
529 for group in 0..full_groups {
530 let base = group * 4;
531
532 let d = _mm_loadu_si128(deltas[base..].as_ptr() as *const __m128i);
534 let gaps = _mm_add_epi32(d, ones);
535
536 let prefix = prefix_sum_4(gaps);
538
539 let result = _mm_add_epi32(prefix, carry);
541
542 _mm_storeu_si128(output[base + 1..].as_mut_ptr() as *mut __m128i, result);
544
545 carry = _mm_shuffle_epi32(result, 0xFF); }
548
549 let base = full_groups * 4;
551 let mut scalar_carry = _mm_extract_epi32(carry, 0) as u32;
552 for j in 0..remainder {
553 scalar_carry = scalar_carry.wrapping_add(deltas[base + j]).wrapping_add(1);
554 output[base + j + 1] = scalar_carry;
555 }
556 }
557
558 #[target_feature(enable = "sse2")]
560 pub unsafe fn add_one(values: &mut [u32], count: usize) {
561 let ones = _mm_set1_epi32(1);
562 let chunks = count / 4;
563 let remainder = count % 4;
564
565 for chunk in 0..chunks {
566 let base = chunk * 4;
567 let ptr = values.as_mut_ptr().add(base) as *mut __m128i;
568 let v = _mm_loadu_si128(ptr);
569 let result = _mm_add_epi32(v, ones);
570 _mm_storeu_si128(ptr, result);
571 }
572
573 let base = chunks * 4;
574 for i in 0..remainder {
575 values[base + i] += 1;
576 }
577 }
578
579 #[target_feature(enable = "sse2", enable = "sse4.1")]
581 pub unsafe fn unpack_8bit_delta_decode(
582 input: &[u8],
583 output: &mut [u32],
584 first_value: u32,
585 count: usize,
586 ) {
587 output[0] = first_value;
588 if count <= 1 {
589 return;
590 }
591
592 let ones = _mm_set1_epi32(1);
593 let mut carry = _mm_set1_epi32(first_value as i32);
594
595 let full_groups = (count - 1) / 4;
596 let remainder = (count - 1) % 4;
597
598 for group in 0..full_groups {
599 let base = group * 4;
600
601 let bytes = _mm_cvtsi32_si128(std::ptr::read_unaligned(
603 input.as_ptr().add(base) as *const i32
604 ));
605 let d = _mm_cvtepu8_epi32(bytes);
606
607 let gaps = _mm_add_epi32(d, ones);
609
610 let prefix = prefix_sum_4(gaps);
612
613 let result = _mm_add_epi32(prefix, carry);
615
616 _mm_storeu_si128(output[base + 1..].as_mut_ptr() as *mut __m128i, result);
618
619 carry = _mm_shuffle_epi32(result, 0xFF);
621 }
622
623 let base = full_groups * 4;
625 let mut scalar_carry = _mm_extract_epi32(carry, 0) as u32;
626 for j in 0..remainder {
627 scalar_carry = scalar_carry
628 .wrapping_add(input[base + j] as u32)
629 .wrapping_add(1);
630 output[base + j + 1] = scalar_carry;
631 }
632 }
633
634 #[target_feature(enable = "sse2", enable = "sse4.1")]
636 pub unsafe fn unpack_16bit_delta_decode(
637 input: &[u8],
638 output: &mut [u32],
639 first_value: u32,
640 count: usize,
641 ) {
642 output[0] = first_value;
643 if count <= 1 {
644 return;
645 }
646
647 let ones = _mm_set1_epi32(1);
648 let mut carry = _mm_set1_epi32(first_value as i32);
649
650 let full_groups = (count - 1) / 4;
651 let remainder = (count - 1) % 4;
652
653 for group in 0..full_groups {
654 let base = group * 4;
655 let in_ptr = input.as_ptr().add(base * 2);
656
657 let vals = _mm_loadl_epi64(in_ptr as *const __m128i); let d = _mm_cvtepu16_epi32(vals);
660
661 let gaps = _mm_add_epi32(d, ones);
663
664 let prefix = prefix_sum_4(gaps);
666
667 let result = _mm_add_epi32(prefix, carry);
669
670 _mm_storeu_si128(output[base + 1..].as_mut_ptr() as *mut __m128i, result);
672
673 carry = _mm_shuffle_epi32(result, 0xFF);
675 }
676
677 let base = full_groups * 4;
679 let mut scalar_carry = _mm_extract_epi32(carry, 0) as u32;
680 for j in 0..remainder {
681 let idx = (base + j) * 2;
682 let delta = u16::from_le_bytes([input[idx], input[idx + 1]]) as u32;
683 scalar_carry = scalar_carry.wrapping_add(delta).wrapping_add(1);
684 output[base + j + 1] = scalar_carry;
685 }
686 }
687
688 #[inline]
690 pub fn is_available() -> bool {
691 is_x86_feature_detected!("sse4.1")
692 }
693}
694
695#[cfg(target_arch = "x86_64")]
700#[allow(unsafe_op_in_unsafe_fn)]
701mod avx2 {
702 use std::arch::x86_64::*;
703
704 #[target_feature(enable = "avx2")]
706 pub unsafe fn unpack_8bit(input: &[u8], output: &mut [u32], count: usize) {
707 let chunks = count / 32;
708 let remainder = count % 32;
709
710 for chunk in 0..chunks {
711 let base = chunk * 32;
712 let in_ptr = input.as_ptr().add(base);
713
714 let bytes_lo = _mm_loadu_si128(in_ptr as *const __m128i);
716 let bytes_hi = _mm_loadu_si128(in_ptr.add(16) as *const __m128i);
717
718 let v0 = _mm256_cvtepu8_epi32(bytes_lo);
720 let v1 = _mm256_cvtepu8_epi32(_mm_srli_si128(bytes_lo, 8));
721 let v2 = _mm256_cvtepu8_epi32(bytes_hi);
722 let v3 = _mm256_cvtepu8_epi32(_mm_srli_si128(bytes_hi, 8));
723
724 let out_ptr = output.as_mut_ptr().add(base);
725 _mm256_storeu_si256(out_ptr as *mut __m256i, v0);
726 _mm256_storeu_si256(out_ptr.add(8) as *mut __m256i, v1);
727 _mm256_storeu_si256(out_ptr.add(16) as *mut __m256i, v2);
728 _mm256_storeu_si256(out_ptr.add(24) as *mut __m256i, v3);
729 }
730
731 let base = chunks * 32;
733 for i in 0..remainder {
734 output[base + i] = input[base + i] as u32;
735 }
736 }
737
738 #[target_feature(enable = "avx2")]
740 pub unsafe fn unpack_16bit(input: &[u8], output: &mut [u32], count: usize) {
741 let chunks = count / 16;
742 let remainder = count % 16;
743
744 for chunk in 0..chunks {
745 let base = chunk * 16;
746 let in_ptr = input.as_ptr().add(base * 2);
747
748 let vals_lo = _mm_loadu_si128(in_ptr as *const __m128i);
750 let vals_hi = _mm_loadu_si128(in_ptr.add(16) as *const __m128i);
751
752 let v0 = _mm256_cvtepu16_epi32(vals_lo);
754 let v1 = _mm256_cvtepu16_epi32(vals_hi);
755
756 let out_ptr = output.as_mut_ptr().add(base);
757 _mm256_storeu_si256(out_ptr as *mut __m256i, v0);
758 _mm256_storeu_si256(out_ptr.add(8) as *mut __m256i, v1);
759 }
760
761 let base = chunks * 16;
763 for i in 0..remainder {
764 let idx = (base + i) * 2;
765 output[base + i] = u16::from_le_bytes([input[idx], input[idx + 1]]) as u32;
766 }
767 }
768
769 #[target_feature(enable = "avx2")]
771 pub unsafe fn unpack_32bit(input: &[u8], output: &mut [u32], count: usize) {
772 let chunks = count / 8;
773 let remainder = count % 8;
774
775 let in_ptr = input.as_ptr() as *const __m256i;
776 let out_ptr = output.as_mut_ptr() as *mut __m256i;
777
778 for chunk in 0..chunks {
779 let vals = _mm256_loadu_si256(in_ptr.add(chunk));
780 _mm256_storeu_si256(out_ptr.add(chunk), vals);
781 }
782
783 let base = chunks * 8;
785 for i in 0..remainder {
786 let idx = (base + i) * 4;
787 output[base + i] =
788 u32::from_le_bytes([input[idx], input[idx + 1], input[idx + 2], input[idx + 3]]);
789 }
790 }
791
792 #[target_feature(enable = "avx2")]
794 pub unsafe fn add_one(values: &mut [u32], count: usize) {
795 let ones = _mm256_set1_epi32(1);
796 let chunks = count / 8;
797 let remainder = count % 8;
798
799 for chunk in 0..chunks {
800 let base = chunk * 8;
801 let ptr = values.as_mut_ptr().add(base) as *mut __m256i;
802 let v = _mm256_loadu_si256(ptr);
803 let result = _mm256_add_epi32(v, ones);
804 _mm256_storeu_si256(ptr, result);
805 }
806
807 let base = chunks * 8;
808 for i in 0..remainder {
809 values[base + i] += 1;
810 }
811 }
812
813 #[inline]
817 #[target_feature(enable = "avx2")]
818 unsafe fn prefix_sum_8(v: __m256i) -> __m256i {
819 let s1 = _mm256_slli_si256(v, 4);
821 let r1 = _mm256_add_epi32(v, s1);
822
823 let s2 = _mm256_slli_si256(r1, 8);
825 let r2 = _mm256_add_epi32(r1, s2);
826
827 let lo_sum = _mm256_shuffle_epi32(r2, 0xFF);
830 let carry = _mm256_permute2x128_si256(lo_sum, lo_sum, 0x00);
832 let carry_hi = _mm256_blend_epi32::<0xF0>(_mm256_setzero_si256(), carry);
834 _mm256_add_epi32(r2, carry_hi)
835 }
836
837 #[target_feature(enable = "avx2")]
839 pub unsafe fn unpack_8bit_delta_decode(
840 input: &[u8],
841 output: &mut [u32],
842 first_value: u32,
843 count: usize,
844 ) {
845 output[0] = first_value;
846 if count <= 1 {
847 return;
848 }
849
850 let ones = _mm256_set1_epi32(1);
851 let mut carry = _mm256_set1_epi32(first_value as i32);
852 let broadcast_idx = _mm256_set1_epi32(7);
853
854 let full_groups = (count - 1) / 8;
855 let remainder = (count - 1) % 8;
856
857 for group in 0..full_groups {
858 let base = group * 8;
859
860 let bytes = _mm_loadl_epi64(input.as_ptr().add(base) as *const __m128i);
862 let d = _mm256_cvtepu8_epi32(bytes);
863
864 let gaps = _mm256_add_epi32(d, ones);
866
867 let prefix = prefix_sum_8(gaps);
869
870 let result = _mm256_add_epi32(prefix, carry);
872
873 _mm256_storeu_si256(output[base + 1..].as_mut_ptr() as *mut __m256i, result);
875
876 carry = _mm256_permutevar8x32_epi32(result, broadcast_idx);
878 }
879
880 let base = full_groups * 8;
882 let mut scalar_carry = _mm256_extract_epi32::<0>(carry) as u32;
883 for j in 0..remainder {
884 scalar_carry = scalar_carry
885 .wrapping_add(input[base + j] as u32)
886 .wrapping_add(1);
887 output[base + j + 1] = scalar_carry;
888 }
889 }
890
891 #[target_feature(enable = "avx2")]
893 pub unsafe fn unpack_16bit_delta_decode(
894 input: &[u8],
895 output: &mut [u32],
896 first_value: u32,
897 count: usize,
898 ) {
899 output[0] = first_value;
900 if count <= 1 {
901 return;
902 }
903
904 let ones = _mm256_set1_epi32(1);
905 let mut carry = _mm256_set1_epi32(first_value as i32);
906 let broadcast_idx = _mm256_set1_epi32(7);
907
908 let full_groups = (count - 1) / 8;
909 let remainder = (count - 1) % 8;
910
911 for group in 0..full_groups {
912 let base = group * 8;
913 let in_ptr = input.as_ptr().add(base * 2);
914
915 let vals = _mm_loadu_si128(in_ptr as *const __m128i);
917 let d = _mm256_cvtepu16_epi32(vals);
918
919 let gaps = _mm256_add_epi32(d, ones);
921
922 let prefix = prefix_sum_8(gaps);
924
925 let result = _mm256_add_epi32(prefix, carry);
927
928 _mm256_storeu_si256(output[base + 1..].as_mut_ptr() as *mut __m256i, result);
930
931 carry = _mm256_permutevar8x32_epi32(result, broadcast_idx);
933 }
934
935 let base = full_groups * 8;
937 let mut scalar_carry = _mm256_extract_epi32::<0>(carry) as u32;
938 for j in 0..remainder {
939 let idx = (base + j) * 2;
940 let delta = u16::from_le_bytes([input[idx], input[idx + 1]]) as u32;
941 scalar_carry = scalar_carry.wrapping_add(delta).wrapping_add(1);
942 output[base + j + 1] = scalar_carry;
943 }
944 }
945
946 #[target_feature(enable = "avx2")]
949 pub unsafe fn hamming_distance(a: &[u8], b: &[u8]) -> u32 {
950 let len = a.len();
951 let chunks32 = len / 32;
952 let low_mask = _mm256_set1_epi8(0x0f);
953 let lookup = _mm256_setr_epi8(
955 0, 1, 1, 2, 1, 2, 2, 3, 1, 2, 2, 3, 2, 3, 3, 4, 0, 1, 1, 2, 1, 2, 2, 3, 1, 2, 2, 3, 2,
956 3, 3, 4,
957 );
958 let mut total = 0u64;
959
960 let mut i = 0;
961 while i < chunks32 {
962 let batch_end = (i + 31).min(chunks32);
965 let mut acc = _mm256_setzero_si256();
966 for j in i..batch_end {
967 let off = j * 32;
968 let va = _mm256_loadu_si256(a.as_ptr().add(off) as *const __m256i);
969 let vb = _mm256_loadu_si256(b.as_ptr().add(off) as *const __m256i);
970 let xored = _mm256_xor_si256(va, vb);
971 let lo = _mm256_and_si256(xored, low_mask);
973 let hi = _mm256_and_si256(_mm256_srli_epi16(xored, 4), low_mask);
974 let popcnt = _mm256_add_epi8(
975 _mm256_shuffle_epi8(lookup, lo),
976 _mm256_shuffle_epi8(lookup, hi),
977 );
978 acc = _mm256_add_epi8(acc, popcnt);
979 }
980 let sad = _mm256_sad_epu8(acc, _mm256_setzero_si256());
982 total += _mm256_extract_epi64(sad, 0) as u64
983 + _mm256_extract_epi64(sad, 1) as u64
984 + _mm256_extract_epi64(sad, 2) as u64
985 + _mm256_extract_epi64(sad, 3) as u64;
986 i = batch_end;
987 }
988
989 let base = chunks32 * 32;
991 for k in base..len {
992 total += (a[k] ^ b[k]).count_ones() as u64;
993 }
994
995 total as u32
996 }
997
998 #[target_feature(enable = "avx2")]
1004 pub unsafe fn hamming_distance_x4(query: &[u8], rows: [&[u8]; 4]) -> [u32; 4] {
1005 let len = query.len();
1006 let chunks32 = len / 32;
1007 let low_mask = _mm256_set1_epi8(0x0f);
1008 let lookup = _mm256_setr_epi8(
1009 0, 1, 1, 2, 1, 2, 2, 3, 1, 2, 2, 3, 2, 3, 3, 4, 0, 1, 1, 2, 1, 2, 2, 3, 1, 2, 2, 3, 2,
1010 3, 3, 4,
1011 );
1012 let mut total = [0u64; 4];
1013
1014 let mut i = 0;
1015 while i < chunks32 {
1016 let batch_end = (i + 31).min(chunks32);
1017 let mut acc = [_mm256_setzero_si256(); 4];
1018 for j in i..batch_end {
1019 let off = j * 32;
1020 let vq = _mm256_loadu_si256(query.as_ptr().add(off) as *const __m256i);
1021 for r in 0..4 {
1022 let vr = _mm256_loadu_si256(rows[r].as_ptr().add(off) as *const __m256i);
1023 let xored = _mm256_xor_si256(vq, vr);
1024 let lo = _mm256_and_si256(xored, low_mask);
1025 let hi = _mm256_and_si256(_mm256_srli_epi16(xored, 4), low_mask);
1026 acc[r] = _mm256_add_epi8(
1027 acc[r],
1028 _mm256_add_epi8(
1029 _mm256_shuffle_epi8(lookup, lo),
1030 _mm256_shuffle_epi8(lookup, hi),
1031 ),
1032 );
1033 }
1034 }
1035 for r in 0..4 {
1036 let sad = _mm256_sad_epu8(acc[r], _mm256_setzero_si256());
1037 total[r] += _mm256_extract_epi64(sad, 0) as u64
1038 + _mm256_extract_epi64(sad, 1) as u64
1039 + _mm256_extract_epi64(sad, 2) as u64
1040 + _mm256_extract_epi64(sad, 3) as u64;
1041 }
1042 i = batch_end;
1043 }
1044
1045 let base = chunks32 * 32;
1048 if base < len {
1049 let tail = &query[base..];
1050 for r in 0..4 {
1051 total[r] += u64::from(super::hamming_distance_scalar(tail, &rows[r][base..]));
1052 }
1053 }
1054
1055 [
1056 total[0] as u32,
1057 total[1] as u32,
1058 total[2] as u32,
1059 total[3] as u32,
1060 ]
1061 }
1062
1063 #[inline]
1065 pub fn is_available() -> bool {
1066 is_x86_feature_detected!("avx2")
1067 }
1068}
1069
1070#[allow(dead_code)]
1075mod scalar {
1076 #[inline]
1078 pub fn unpack_8bit(input: &[u8], output: &mut [u32], count: usize) {
1079 for i in 0..count {
1080 output[i] = input[i] as u32;
1081 }
1082 }
1083
1084 #[inline]
1086 pub fn unpack_16bit(input: &[u8], output: &mut [u32], count: usize) {
1087 for (i, out) in output.iter_mut().enumerate().take(count) {
1088 let idx = i * 2;
1089 *out = u16::from_le_bytes([input[idx], input[idx + 1]]) as u32;
1090 }
1091 }
1092
1093 #[inline]
1095 pub fn unpack_32bit(input: &[u8], output: &mut [u32], count: usize) {
1096 for (i, out) in output.iter_mut().enumerate().take(count) {
1097 let idx = i * 4;
1098 *out = u32::from_le_bytes([input[idx], input[idx + 1], input[idx + 2], input[idx + 3]]);
1099 }
1100 }
1101
1102 #[inline]
1104 pub fn delta_decode(output: &mut [u32], deltas: &[u32], first_doc_id: u32, count: usize) {
1105 if count == 0 {
1106 return;
1107 }
1108
1109 output[0] = first_doc_id;
1110 let mut carry = first_doc_id;
1111
1112 for i in 0..count - 1 {
1113 carry = carry.wrapping_add(deltas[i]).wrapping_add(1);
1114 output[i + 1] = carry;
1115 }
1116 }
1117
1118 #[inline]
1120 pub fn add_one(values: &mut [u32], count: usize) {
1121 for val in values.iter_mut().take(count) {
1122 *val += 1;
1123 }
1124 }
1125}
1126
1127#[inline]
1133pub fn unpack_8bit(input: &[u8], output: &mut [u32], count: usize) {
1134 #[cfg(target_arch = "aarch64")]
1135 {
1136 if neon::is_available() {
1137 unsafe {
1138 neon::unpack_8bit(input, output, count);
1139 }
1140 return;
1141 }
1142 }
1143
1144 #[cfg(target_arch = "x86_64")]
1145 {
1146 if avx2::is_available() {
1148 unsafe {
1149 avx2::unpack_8bit(input, output, count);
1150 }
1151 return;
1152 }
1153 if sse::is_available() {
1154 unsafe {
1155 sse::unpack_8bit(input, output, count);
1156 }
1157 return;
1158 }
1159 }
1160
1161 scalar::unpack_8bit(input, output, count);
1162}
1163
1164#[inline]
1166pub fn unpack_16bit(input: &[u8], output: &mut [u32], count: usize) {
1167 #[cfg(target_arch = "aarch64")]
1168 {
1169 if neon::is_available() {
1170 unsafe {
1171 neon::unpack_16bit(input, output, count);
1172 }
1173 return;
1174 }
1175 }
1176
1177 #[cfg(target_arch = "x86_64")]
1178 {
1179 if avx2::is_available() {
1181 unsafe {
1182 avx2::unpack_16bit(input, output, count);
1183 }
1184 return;
1185 }
1186 if sse::is_available() {
1187 unsafe {
1188 sse::unpack_16bit(input, output, count);
1189 }
1190 return;
1191 }
1192 }
1193
1194 scalar::unpack_16bit(input, output, count);
1195}
1196
1197#[inline]
1199pub fn unpack_32bit(input: &[u8], output: &mut [u32], count: usize) {
1200 #[cfg(target_arch = "aarch64")]
1201 {
1202 if neon::is_available() {
1203 unsafe {
1204 neon::unpack_32bit(input, output, count);
1205 }
1206 return;
1207 }
1208 }
1209
1210 #[cfg(target_arch = "x86_64")]
1211 {
1212 if avx2::is_available() {
1214 unsafe {
1215 avx2::unpack_32bit(input, output, count);
1216 }
1217 return;
1218 }
1219 if sse::is_available() {
1220 unsafe {
1221 sse::unpack_32bit(input, output, count);
1222 }
1223 return;
1224 }
1225 }
1226
1227 scalar::unpack_32bit(input, output, count);
1228}
1229
1230#[inline]
1236pub fn delta_decode(output: &mut [u32], deltas: &[u32], first_value: u32, count: usize) {
1237 #[cfg(target_arch = "aarch64")]
1238 {
1239 if neon::is_available() {
1240 unsafe {
1241 neon::delta_decode(output, deltas, first_value, count);
1242 }
1243 return;
1244 }
1245 }
1246
1247 #[cfg(target_arch = "x86_64")]
1248 {
1249 if sse::is_available() {
1250 unsafe {
1251 sse::delta_decode(output, deltas, first_value, count);
1252 }
1253 return;
1254 }
1255 }
1256
1257 scalar::delta_decode(output, deltas, first_value, count);
1258}
1259
1260#[inline]
1264pub fn add_one(values: &mut [u32], count: usize) {
1265 #[cfg(target_arch = "aarch64")]
1266 {
1267 if neon::is_available() {
1268 unsafe {
1269 neon::add_one(values, count);
1270 }
1271 return;
1272 }
1273 }
1274
1275 #[cfg(target_arch = "x86_64")]
1276 {
1277 if avx2::is_available() {
1279 unsafe {
1280 avx2::add_one(values, count);
1281 }
1282 return;
1283 }
1284 if sse::is_available() {
1285 unsafe {
1286 sse::add_one(values, count);
1287 }
1288 return;
1289 }
1290 }
1291
1292 scalar::add_one(values, count);
1293}
1294
1295#[inline]
1297pub fn bits_needed(val: u32) -> u8 {
1298 if val == 0 {
1299 0
1300 } else {
1301 32 - val.leading_zeros() as u8
1302 }
1303}
1304
1305#[derive(Debug, Clone, Copy, PartialEq, Eq)]
1322#[repr(u8)]
1323pub enum RoundedBitWidth {
1324 Zero = 0,
1325 Bits8 = 8,
1326 Bits16 = 16,
1327 Bits32 = 32,
1328}
1329
1330impl RoundedBitWidth {
1331 #[inline]
1333 pub fn from_exact(bits: u8) -> Self {
1334 match bits {
1335 0 => RoundedBitWidth::Zero,
1336 1..=8 => RoundedBitWidth::Bits8,
1337 9..=16 => RoundedBitWidth::Bits16,
1338 _ => RoundedBitWidth::Bits32,
1339 }
1340 }
1341
1342 #[inline]
1344 pub fn from_u8(bits: u8) -> Self {
1345 match bits {
1346 0 => RoundedBitWidth::Zero,
1347 8 => RoundedBitWidth::Bits8,
1348 16 => RoundedBitWidth::Bits16,
1349 32 => RoundedBitWidth::Bits32,
1350 _ => RoundedBitWidth::Bits32, }
1352 }
1353
1354 #[inline]
1356 pub fn bytes_per_value(self) -> usize {
1357 match self {
1358 RoundedBitWidth::Zero => 0,
1359 RoundedBitWidth::Bits8 => 1,
1360 RoundedBitWidth::Bits16 => 2,
1361 RoundedBitWidth::Bits32 => 4,
1362 }
1363 }
1364
1365 #[inline]
1367 pub fn as_u8(self) -> u8 {
1368 self as u8
1369 }
1370}
1371
1372#[inline]
1374pub fn round_bit_width(bits: u8) -> u8 {
1375 RoundedBitWidth::from_exact(bits).as_u8()
1376}
1377
1378#[inline]
1383pub fn pack_rounded(values: &[u32], bit_width: RoundedBitWidth, output: &mut [u8]) -> usize {
1384 let count = values.len();
1385 match bit_width {
1386 RoundedBitWidth::Zero => 0,
1387 RoundedBitWidth::Bits8 => {
1388 for (i, &v) in values.iter().enumerate() {
1389 output[i] = v as u8;
1390 }
1391 count
1392 }
1393 RoundedBitWidth::Bits16 => {
1394 for (i, &v) in values.iter().enumerate() {
1395 let bytes = (v as u16).to_le_bytes();
1396 output[i * 2] = bytes[0];
1397 output[i * 2 + 1] = bytes[1];
1398 }
1399 count * 2
1400 }
1401 RoundedBitWidth::Bits32 => {
1402 for (i, &v) in values.iter().enumerate() {
1403 let bytes = v.to_le_bytes();
1404 output[i * 4] = bytes[0];
1405 output[i * 4 + 1] = bytes[1];
1406 output[i * 4 + 2] = bytes[2];
1407 output[i * 4 + 3] = bytes[3];
1408 }
1409 count * 4
1410 }
1411 }
1412}
1413
1414#[inline]
1418pub fn unpack_rounded(input: &[u8], bit_width: RoundedBitWidth, output: &mut [u32], count: usize) {
1419 match bit_width {
1420 RoundedBitWidth::Zero => {
1421 for out in output.iter_mut().take(count) {
1422 *out = 0;
1423 }
1424 }
1425 RoundedBitWidth::Bits8 => unpack_8bit(input, output, count),
1426 RoundedBitWidth::Bits16 => unpack_16bit(input, output, count),
1427 RoundedBitWidth::Bits32 => unpack_32bit(input, output, count),
1428 }
1429}
1430
1431#[inline]
1435pub fn unpack_rounded_delta_decode(
1436 input: &[u8],
1437 bit_width: RoundedBitWidth,
1438 output: &mut [u32],
1439 first_value: u32,
1440 count: usize,
1441) {
1442 match bit_width {
1443 RoundedBitWidth::Zero => {
1444 let mut val = first_value;
1446 for out in output.iter_mut().take(count) {
1447 *out = val;
1448 val = val.wrapping_add(1);
1449 }
1450 }
1451 RoundedBitWidth::Bits8 => unpack_8bit_delta_decode(input, output, first_value, count),
1452 RoundedBitWidth::Bits16 => unpack_16bit_delta_decode(input, output, first_value, count),
1453 RoundedBitWidth::Bits32 => {
1454 if count > 0 {
1456 output[0] = first_value;
1457 let mut carry = first_value;
1458 for i in 0..count - 1 {
1459 let idx = i * 4;
1460 let delta = u32::from_le_bytes([
1461 input[idx],
1462 input[idx + 1],
1463 input[idx + 2],
1464 input[idx + 3],
1465 ]);
1466 carry = carry.wrapping_add(delta).wrapping_add(1);
1467 output[i + 1] = carry;
1468 }
1469 }
1470 }
1471 }
1472}
1473
1474#[inline]
1483pub fn unpack_8bit_delta_decode(input: &[u8], output: &mut [u32], first_value: u32, count: usize) {
1484 if count == 0 {
1485 return;
1486 }
1487
1488 output[0] = first_value;
1489 if count == 1 {
1490 return;
1491 }
1492
1493 #[cfg(target_arch = "aarch64")]
1494 {
1495 if neon::is_available() {
1496 unsafe {
1497 neon::unpack_8bit_delta_decode(input, output, first_value, count);
1498 }
1499 return;
1500 }
1501 }
1502
1503 #[cfg(target_arch = "x86_64")]
1504 {
1505 if avx2::is_available() {
1506 unsafe {
1507 avx2::unpack_8bit_delta_decode(input, output, first_value, count);
1508 }
1509 return;
1510 }
1511 if sse::is_available() {
1512 unsafe {
1513 sse::unpack_8bit_delta_decode(input, output, first_value, count);
1514 }
1515 return;
1516 }
1517 }
1518
1519 let mut carry = first_value;
1521 for i in 0..count - 1 {
1522 carry = carry.wrapping_add(input[i] as u32).wrapping_add(1);
1523 output[i + 1] = carry;
1524 }
1525}
1526
1527#[inline]
1529pub fn unpack_16bit_delta_decode(input: &[u8], output: &mut [u32], first_value: u32, count: usize) {
1530 if count == 0 {
1531 return;
1532 }
1533
1534 output[0] = first_value;
1535 if count == 1 {
1536 return;
1537 }
1538
1539 #[cfg(target_arch = "aarch64")]
1540 {
1541 if neon::is_available() {
1542 unsafe {
1543 neon::unpack_16bit_delta_decode(input, output, first_value, count);
1544 }
1545 return;
1546 }
1547 }
1548
1549 #[cfg(target_arch = "x86_64")]
1550 {
1551 if avx2::is_available() {
1552 unsafe {
1553 avx2::unpack_16bit_delta_decode(input, output, first_value, count);
1554 }
1555 return;
1556 }
1557 if sse::is_available() {
1558 unsafe {
1559 sse::unpack_16bit_delta_decode(input, output, first_value, count);
1560 }
1561 return;
1562 }
1563 }
1564
1565 let mut carry = first_value;
1567 for i in 0..count - 1 {
1568 let idx = i * 2;
1569 let delta = u16::from_le_bytes([input[idx], input[idx + 1]]) as u32;
1570 carry = carry.wrapping_add(delta).wrapping_add(1);
1571 output[i + 1] = carry;
1572 }
1573}
1574
1575#[inline]
1580pub fn unpack_delta_decode(
1581 input: &[u8],
1582 bit_width: u8,
1583 output: &mut [u32],
1584 first_value: u32,
1585 count: usize,
1586) {
1587 if count == 0 {
1588 return;
1589 }
1590
1591 output[0] = first_value;
1592 if count == 1 {
1593 return;
1594 }
1595
1596 match bit_width {
1598 0 => {
1599 let mut val = first_value;
1601 for item in output.iter_mut().take(count).skip(1) {
1602 val = val.wrapping_add(1);
1603 *item = val;
1604 }
1605 }
1606 8 => unpack_8bit_delta_decode(input, output, first_value, count),
1607 16 => unpack_16bit_delta_decode(input, output, first_value, count),
1608 32 => {
1609 let mut carry = first_value;
1611 for i in 0..count - 1 {
1612 let idx = i * 4;
1613 let delta = u32::from_le_bytes([
1614 input[idx],
1615 input[idx + 1],
1616 input[idx + 2],
1617 input[idx + 3],
1618 ]);
1619 carry = carry.wrapping_add(delta).wrapping_add(1);
1620 output[i + 1] = carry;
1621 }
1622 }
1623 _ => {
1624 let mask = (1u64 << bit_width) - 1;
1626 let bit_width_usize = bit_width as usize;
1627 let mut bit_pos = 0usize;
1628 let input_ptr = input.as_ptr();
1629 let mut carry = first_value;
1630
1631 for i in 0..count - 1 {
1632 let byte_idx = bit_pos >> 3;
1633 let bit_offset = bit_pos & 7;
1634
1635 let word = unsafe { (input_ptr.add(byte_idx) as *const u64).read_unaligned() };
1637 let delta = ((word >> bit_offset) & mask) as u32;
1638
1639 carry = carry.wrapping_add(delta).wrapping_add(1);
1640 output[i + 1] = carry;
1641 bit_pos += bit_width_usize;
1642 }
1643 }
1644 }
1645}
1646
1647#[inline]
1655pub fn dequantize_uint8(input: &[u8], output: &mut [f32], scale: f32, min_val: f32, count: usize) {
1656 #[cfg(target_arch = "aarch64")]
1657 {
1658 if neon::is_available() {
1659 unsafe {
1660 dequantize_uint8_neon(input, output, scale, min_val, count);
1661 }
1662 return;
1663 }
1664 }
1665
1666 #[cfg(target_arch = "x86_64")]
1667 {
1668 if sse::is_available() {
1669 unsafe {
1670 dequantize_uint8_sse(input, output, scale, min_val, count);
1671 }
1672 return;
1673 }
1674 }
1675
1676 for i in 0..count {
1678 output[i] = input[i] as f32 * scale + min_val;
1679 }
1680}
1681
1682#[cfg(target_arch = "aarch64")]
1683#[target_feature(enable = "neon")]
1684#[allow(unsafe_op_in_unsafe_fn)]
1685unsafe fn dequantize_uint8_neon(
1686 input: &[u8],
1687 output: &mut [f32],
1688 scale: f32,
1689 min_val: f32,
1690 count: usize,
1691) {
1692 use std::arch::aarch64::*;
1693
1694 let scale_v = vdupq_n_f32(scale);
1695 let min_v = vdupq_n_f32(min_val);
1696
1697 let chunks = count / 16;
1698 let remainder = count % 16;
1699
1700 for chunk in 0..chunks {
1701 let base = chunk * 16;
1702 let in_ptr = input.as_ptr().add(base);
1703
1704 let bytes = vld1q_u8(in_ptr);
1706
1707 let low8 = vget_low_u8(bytes);
1709 let high8 = vget_high_u8(bytes);
1710
1711 let low16 = vmovl_u8(low8);
1712 let high16 = vmovl_u8(high8);
1713
1714 let u32_0 = vmovl_u16(vget_low_u16(low16));
1716 let u32_1 = vmovl_u16(vget_high_u16(low16));
1717 let u32_2 = vmovl_u16(vget_low_u16(high16));
1718 let u32_3 = vmovl_u16(vget_high_u16(high16));
1719
1720 let f32_0 = vfmaq_f32(min_v, vcvtq_f32_u32(u32_0), scale_v);
1722 let f32_1 = vfmaq_f32(min_v, vcvtq_f32_u32(u32_1), scale_v);
1723 let f32_2 = vfmaq_f32(min_v, vcvtq_f32_u32(u32_2), scale_v);
1724 let f32_3 = vfmaq_f32(min_v, vcvtq_f32_u32(u32_3), scale_v);
1725
1726 let out_ptr = output.as_mut_ptr().add(base);
1727 vst1q_f32(out_ptr, f32_0);
1728 vst1q_f32(out_ptr.add(4), f32_1);
1729 vst1q_f32(out_ptr.add(8), f32_2);
1730 vst1q_f32(out_ptr.add(12), f32_3);
1731 }
1732
1733 let base = chunks * 16;
1735 for i in 0..remainder {
1736 output[base + i] = input[base + i] as f32 * scale + min_val;
1737 }
1738}
1739
1740#[cfg(target_arch = "x86_64")]
1741#[target_feature(enable = "sse2", enable = "sse4.1")]
1742#[allow(unsafe_op_in_unsafe_fn)]
1743unsafe fn dequantize_uint8_sse(
1744 input: &[u8],
1745 output: &mut [f32],
1746 scale: f32,
1747 min_val: f32,
1748 count: usize,
1749) {
1750 use std::arch::x86_64::*;
1751
1752 let scale_v = _mm_set1_ps(scale);
1753 let min_v = _mm_set1_ps(min_val);
1754
1755 let chunks = count / 4;
1756 let remainder = count % 4;
1757
1758 for chunk in 0..chunks {
1759 let base = chunk * 4;
1760
1761 let bytes = _mm_cvtsi32_si128(std::ptr::read_unaligned(
1763 input.as_ptr().add(base) as *const i32
1764 ));
1765 let ints = _mm_cvtepu8_epi32(bytes);
1766 let floats = _mm_cvtepi32_ps(ints);
1767
1768 let scaled = _mm_add_ps(_mm_mul_ps(floats, scale_v), min_v);
1770
1771 _mm_storeu_ps(output.as_mut_ptr().add(base), scaled);
1772 }
1773
1774 let base = chunks * 4;
1776 for i in 0..remainder {
1777 output[base + i] = input[base + i] as f32 * scale + min_val;
1778 }
1779}
1780
1781#[inline]
1825fn dot_product_f32_scalar(a: &[f32], b: &[f32]) -> f32 {
1826 a.iter().zip(b).fold(0.0f32, |acc, (&x, &y)| {
1827 acc.algebraic_add(x.algebraic_mul(y))
1828 })
1829}
1830
1831#[inline]
1833fn fused_dot_norm_scalar(a: &[f32], b: &[f32]) -> (f32, f32) {
1834 a.iter()
1835 .zip(b)
1836 .fold((0.0f32, 0.0f32), |(dot, norm), (&x, &y)| {
1837 (
1838 dot.algebraic_add(x.algebraic_mul(y)),
1839 norm.algebraic_add(y.algebraic_mul(y)),
1840 )
1841 })
1842}
1843
1844#[inline]
1849pub fn squared_l2_f32(a: &[f32], b: &[f32]) -> f32 {
1850 a.iter().zip(b).fold(0.0f32, |acc, (&x, &y)| {
1851 let delta = x - y;
1852 acc.algebraic_add(delta.algebraic_mul(delta))
1853 })
1854}
1855
1856#[inline]
1858pub fn norm_squared_f32(v: &[f32]) -> f32 {
1859 v.iter()
1860 .fold(0.0f32, |acc, &x| acc.algebraic_add(x.algebraic_mul(x)))
1861}
1862
1863#[inline]
1865pub fn norm_f32(v: &[f32]) -> f32 {
1866 norm_squared_f32(v).sqrt()
1867}
1868
1869#[inline]
1871pub fn dot_product_f32(a: &[f32], b: &[f32], count: usize) -> f32 {
1872 assert!(
1873 count <= a.len() && count <= b.len(),
1874 "dot_product_f32 count {count} exceeds input lengths ({}, {})",
1875 a.len(),
1876 b.len()
1877 );
1878 #[cfg(target_arch = "aarch64")]
1879 {
1880 if neon::is_available() {
1881 return unsafe { dot_product_f32_neon(a, b, count) };
1882 }
1883 }
1884
1885 #[cfg(target_arch = "x86_64")]
1886 {
1887 if is_x86_feature_detected!("avx512f") {
1888 return unsafe { dot_product_f32_avx512(a, b, count) };
1889 }
1890 if is_x86_feature_detected!("avx2") && is_x86_feature_detected!("fma") {
1891 return unsafe { dot_product_f32_avx2(a, b, count) };
1892 }
1893 if sse::is_available() {
1894 return unsafe { dot_product_f32_sse(a, b, count) };
1895 }
1896 }
1897
1898 dot_product_f32_scalar(&a[..count], &b[..count])
1900}
1901
1902#[cfg(target_arch = "aarch64")]
1903#[target_feature(enable = "neon")]
1904#[allow(unsafe_op_in_unsafe_fn)]
1905unsafe fn dot_product_f32_neon(a: &[f32], b: &[f32], count: usize) -> f32 {
1906 use std::arch::aarch64::*;
1907
1908 let chunks16 = count / 16;
1909 let remainder = count % 16;
1910
1911 let mut acc0 = vdupq_n_f32(0.0);
1912 let mut acc1 = vdupq_n_f32(0.0);
1913 let mut acc2 = vdupq_n_f32(0.0);
1914 let mut acc3 = vdupq_n_f32(0.0);
1915
1916 for c in 0..chunks16 {
1917 let base = c * 16;
1918 acc0 = vfmaq_f32(
1919 acc0,
1920 vld1q_f32(a.as_ptr().add(base)),
1921 vld1q_f32(b.as_ptr().add(base)),
1922 );
1923 acc1 = vfmaq_f32(
1924 acc1,
1925 vld1q_f32(a.as_ptr().add(base + 4)),
1926 vld1q_f32(b.as_ptr().add(base + 4)),
1927 );
1928 acc2 = vfmaq_f32(
1929 acc2,
1930 vld1q_f32(a.as_ptr().add(base + 8)),
1931 vld1q_f32(b.as_ptr().add(base + 8)),
1932 );
1933 acc3 = vfmaq_f32(
1934 acc3,
1935 vld1q_f32(a.as_ptr().add(base + 12)),
1936 vld1q_f32(b.as_ptr().add(base + 12)),
1937 );
1938 }
1939
1940 let acc = vaddq_f32(vaddq_f32(acc0, acc1), vaddq_f32(acc2, acc3));
1941 let mut sum = vaddvq_f32(acc);
1942
1943 let base = chunks16 * 16;
1944 for i in 0..remainder {
1945 sum += a[base + i] * b[base + i];
1946 }
1947
1948 sum
1949}
1950
1951#[cfg(target_arch = "x86_64")]
1952#[target_feature(enable = "avx2", enable = "fma")]
1953#[allow(unsafe_op_in_unsafe_fn)]
1954unsafe fn dot_product_f32_avx2(a: &[f32], b: &[f32], count: usize) -> f32 {
1955 use std::arch::x86_64::*;
1956
1957 let chunks32 = count / 32;
1958 let remainder = count % 32;
1959
1960 let mut acc0 = _mm256_setzero_ps();
1961 let mut acc1 = _mm256_setzero_ps();
1962 let mut acc2 = _mm256_setzero_ps();
1963 let mut acc3 = _mm256_setzero_ps();
1964
1965 for c in 0..chunks32 {
1966 let base = c * 32;
1967 acc0 = _mm256_fmadd_ps(
1968 _mm256_loadu_ps(a.as_ptr().add(base)),
1969 _mm256_loadu_ps(b.as_ptr().add(base)),
1970 acc0,
1971 );
1972 acc1 = _mm256_fmadd_ps(
1973 _mm256_loadu_ps(a.as_ptr().add(base + 8)),
1974 _mm256_loadu_ps(b.as_ptr().add(base + 8)),
1975 acc1,
1976 );
1977 acc2 = _mm256_fmadd_ps(
1978 _mm256_loadu_ps(a.as_ptr().add(base + 16)),
1979 _mm256_loadu_ps(b.as_ptr().add(base + 16)),
1980 acc2,
1981 );
1982 acc3 = _mm256_fmadd_ps(
1983 _mm256_loadu_ps(a.as_ptr().add(base + 24)),
1984 _mm256_loadu_ps(b.as_ptr().add(base + 24)),
1985 acc3,
1986 );
1987 }
1988
1989 let acc = _mm256_add_ps(_mm256_add_ps(acc0, acc1), _mm256_add_ps(acc2, acc3));
1990
1991 let hi = _mm256_extractf128_ps(acc, 1);
1993 let lo = _mm256_castps256_ps128(acc);
1994 let sum128 = _mm_add_ps(lo, hi);
1995 let shuf = _mm_shuffle_ps(sum128, sum128, 0b10_11_00_01);
1996 let sums = _mm_add_ps(sum128, shuf);
1997 let shuf2 = _mm_movehl_ps(sums, sums);
1998 let final_sum = _mm_add_ss(sums, shuf2);
1999
2000 let mut sum = _mm_cvtss_f32(final_sum);
2001
2002 let base = chunks32 * 32;
2003 for i in 0..remainder {
2004 sum += a[base + i] * b[base + i];
2005 }
2006
2007 sum
2008}
2009
2010#[cfg(target_arch = "x86_64")]
2011#[target_feature(enable = "sse")]
2012#[allow(unsafe_op_in_unsafe_fn)]
2013unsafe fn dot_product_f32_sse(a: &[f32], b: &[f32], count: usize) -> f32 {
2014 use std::arch::x86_64::*;
2015
2016 let chunks = count / 4;
2017 let remainder = count % 4;
2018
2019 let mut acc = _mm_setzero_ps();
2020
2021 for chunk in 0..chunks {
2022 let base = chunk * 4;
2023 let va = _mm_loadu_ps(a.as_ptr().add(base));
2024 let vb = _mm_loadu_ps(b.as_ptr().add(base));
2025 acc = _mm_add_ps(acc, _mm_mul_ps(va, vb));
2026 }
2027
2028 let shuf = _mm_shuffle_ps(acc, acc, 0b10_11_00_01); let sums = _mm_add_ps(acc, shuf); let shuf2 = _mm_movehl_ps(sums, sums); let final_sum = _mm_add_ss(sums, shuf2); let mut sum = _mm_cvtss_f32(final_sum);
2035
2036 let base = chunks * 4;
2038 for i in 0..remainder {
2039 sum += a[base + i] * b[base + i];
2040 }
2041
2042 sum
2043}
2044
2045#[cfg(target_arch = "x86_64")]
2046#[target_feature(enable = "avx512f")]
2047#[allow(unsafe_op_in_unsafe_fn)]
2048unsafe fn dot_product_f32_avx512(a: &[f32], b: &[f32], count: usize) -> f32 {
2049 use std::arch::x86_64::*;
2050
2051 let chunks64 = count / 64;
2052 let remainder = count % 64;
2053
2054 let mut acc0 = _mm512_setzero_ps();
2055 let mut acc1 = _mm512_setzero_ps();
2056 let mut acc2 = _mm512_setzero_ps();
2057 let mut acc3 = _mm512_setzero_ps();
2058
2059 for c in 0..chunks64 {
2060 let base = c * 64;
2061 acc0 = _mm512_fmadd_ps(
2062 _mm512_loadu_ps(a.as_ptr().add(base)),
2063 _mm512_loadu_ps(b.as_ptr().add(base)),
2064 acc0,
2065 );
2066 acc1 = _mm512_fmadd_ps(
2067 _mm512_loadu_ps(a.as_ptr().add(base + 16)),
2068 _mm512_loadu_ps(b.as_ptr().add(base + 16)),
2069 acc1,
2070 );
2071 acc2 = _mm512_fmadd_ps(
2072 _mm512_loadu_ps(a.as_ptr().add(base + 32)),
2073 _mm512_loadu_ps(b.as_ptr().add(base + 32)),
2074 acc2,
2075 );
2076 acc3 = _mm512_fmadd_ps(
2077 _mm512_loadu_ps(a.as_ptr().add(base + 48)),
2078 _mm512_loadu_ps(b.as_ptr().add(base + 48)),
2079 acc3,
2080 );
2081 }
2082
2083 let acc = _mm512_add_ps(_mm512_add_ps(acc0, acc1), _mm512_add_ps(acc2, acc3));
2084 let mut sum = _mm512_reduce_add_ps(acc);
2085
2086 let base = chunks64 * 64;
2087 for i in 0..remainder {
2088 sum += a[base + i] * b[base + i];
2089 }
2090
2091 sum
2092}
2093
2094#[cfg(target_arch = "x86_64")]
2095#[target_feature(enable = "avx512f")]
2096#[allow(unsafe_op_in_unsafe_fn)]
2097unsafe fn fused_dot_norm_avx512(a: &[f32], b: &[f32], count: usize) -> (f32, f32) {
2098 use std::arch::x86_64::*;
2099
2100 let chunks64 = count / 64;
2101 let remainder = count % 64;
2102
2103 let mut d0 = _mm512_setzero_ps();
2104 let mut d1 = _mm512_setzero_ps();
2105 let mut d2 = _mm512_setzero_ps();
2106 let mut d3 = _mm512_setzero_ps();
2107 let mut n0 = _mm512_setzero_ps();
2108 let mut n1 = _mm512_setzero_ps();
2109 let mut n2 = _mm512_setzero_ps();
2110 let mut n3 = _mm512_setzero_ps();
2111
2112 for c in 0..chunks64 {
2113 let base = c * 64;
2114 let vb0 = _mm512_loadu_ps(b.as_ptr().add(base));
2115 d0 = _mm512_fmadd_ps(_mm512_loadu_ps(a.as_ptr().add(base)), vb0, d0);
2116 n0 = _mm512_fmadd_ps(vb0, vb0, n0);
2117 let vb1 = _mm512_loadu_ps(b.as_ptr().add(base + 16));
2118 d1 = _mm512_fmadd_ps(_mm512_loadu_ps(a.as_ptr().add(base + 16)), vb1, d1);
2119 n1 = _mm512_fmadd_ps(vb1, vb1, n1);
2120 let vb2 = _mm512_loadu_ps(b.as_ptr().add(base + 32));
2121 d2 = _mm512_fmadd_ps(_mm512_loadu_ps(a.as_ptr().add(base + 32)), vb2, d2);
2122 n2 = _mm512_fmadd_ps(vb2, vb2, n2);
2123 let vb3 = _mm512_loadu_ps(b.as_ptr().add(base + 48));
2124 d3 = _mm512_fmadd_ps(_mm512_loadu_ps(a.as_ptr().add(base + 48)), vb3, d3);
2125 n3 = _mm512_fmadd_ps(vb3, vb3, n3);
2126 }
2127
2128 let acc_dot = _mm512_add_ps(_mm512_add_ps(d0, d1), _mm512_add_ps(d2, d3));
2129 let acc_norm = _mm512_add_ps(_mm512_add_ps(n0, n1), _mm512_add_ps(n2, n3));
2130 let mut dot = _mm512_reduce_add_ps(acc_dot);
2131 let mut norm = _mm512_reduce_add_ps(acc_norm);
2132
2133 let base = chunks64 * 64;
2134 for i in 0..remainder {
2135 dot += a[base + i] * b[base + i];
2136 norm += b[base + i] * b[base + i];
2137 }
2138
2139 (dot, norm)
2140}
2141
2142#[inline]
2151fn fused_dot_norm(a: &[f32], b: &[f32], count: usize) -> (f32, f32) {
2152 #[cfg(target_arch = "aarch64")]
2153 {
2154 if neon::is_available() {
2155 return unsafe { fused_dot_norm_neon(a, b, count) };
2156 }
2157 }
2158
2159 #[cfg(target_arch = "x86_64")]
2160 {
2161 if is_x86_feature_detected!("avx512f") {
2162 return unsafe { fused_dot_norm_avx512(a, b, count) };
2163 }
2164 if is_x86_feature_detected!("avx2") && is_x86_feature_detected!("fma") {
2165 return unsafe { fused_dot_norm_avx2(a, b, count) };
2166 }
2167 if sse::is_available() {
2168 return unsafe { fused_dot_norm_sse(a, b, count) };
2169 }
2170 }
2171
2172 fused_dot_norm_scalar(&a[..count], &b[..count])
2174}
2175
2176#[cfg(target_arch = "aarch64")]
2177#[target_feature(enable = "neon")]
2178#[allow(unsafe_op_in_unsafe_fn)]
2179unsafe fn fused_dot_norm_neon(a: &[f32], b: &[f32], count: usize) -> (f32, f32) {
2180 use std::arch::aarch64::*;
2181
2182 let chunks16 = count / 16;
2183 let remainder = count % 16;
2184
2185 let mut d0 = vdupq_n_f32(0.0);
2186 let mut d1 = vdupq_n_f32(0.0);
2187 let mut d2 = vdupq_n_f32(0.0);
2188 let mut d3 = vdupq_n_f32(0.0);
2189 let mut n0 = vdupq_n_f32(0.0);
2190 let mut n1 = vdupq_n_f32(0.0);
2191 let mut n2 = vdupq_n_f32(0.0);
2192 let mut n3 = vdupq_n_f32(0.0);
2193
2194 for c in 0..chunks16 {
2195 let base = c * 16;
2196 let va0 = vld1q_f32(a.as_ptr().add(base));
2197 let vb0 = vld1q_f32(b.as_ptr().add(base));
2198 d0 = vfmaq_f32(d0, va0, vb0);
2199 n0 = vfmaq_f32(n0, vb0, vb0);
2200 let va1 = vld1q_f32(a.as_ptr().add(base + 4));
2201 let vb1 = vld1q_f32(b.as_ptr().add(base + 4));
2202 d1 = vfmaq_f32(d1, va1, vb1);
2203 n1 = vfmaq_f32(n1, vb1, vb1);
2204 let va2 = vld1q_f32(a.as_ptr().add(base + 8));
2205 let vb2 = vld1q_f32(b.as_ptr().add(base + 8));
2206 d2 = vfmaq_f32(d2, va2, vb2);
2207 n2 = vfmaq_f32(n2, vb2, vb2);
2208 let va3 = vld1q_f32(a.as_ptr().add(base + 12));
2209 let vb3 = vld1q_f32(b.as_ptr().add(base + 12));
2210 d3 = vfmaq_f32(d3, va3, vb3);
2211 n3 = vfmaq_f32(n3, vb3, vb3);
2212 }
2213
2214 let acc_dot = vaddq_f32(vaddq_f32(d0, d1), vaddq_f32(d2, d3));
2215 let acc_norm = vaddq_f32(vaddq_f32(n0, n1), vaddq_f32(n2, n3));
2216 let mut dot = vaddvq_f32(acc_dot);
2217 let mut norm = vaddvq_f32(acc_norm);
2218
2219 let base = chunks16 * 16;
2220 for i in 0..remainder {
2221 dot += a[base + i] * b[base + i];
2222 norm += b[base + i] * b[base + i];
2223 }
2224
2225 (dot, norm)
2226}
2227
2228#[cfg(target_arch = "x86_64")]
2229#[target_feature(enable = "avx2", enable = "fma")]
2230#[allow(unsafe_op_in_unsafe_fn)]
2231unsafe fn fused_dot_norm_avx2(a: &[f32], b: &[f32], count: usize) -> (f32, f32) {
2232 use std::arch::x86_64::*;
2233
2234 let chunks32 = count / 32;
2235 let remainder = count % 32;
2236
2237 let mut d0 = _mm256_setzero_ps();
2238 let mut d1 = _mm256_setzero_ps();
2239 let mut d2 = _mm256_setzero_ps();
2240 let mut d3 = _mm256_setzero_ps();
2241 let mut n0 = _mm256_setzero_ps();
2242 let mut n1 = _mm256_setzero_ps();
2243 let mut n2 = _mm256_setzero_ps();
2244 let mut n3 = _mm256_setzero_ps();
2245
2246 for c in 0..chunks32 {
2247 let base = c * 32;
2248 let vb0 = _mm256_loadu_ps(b.as_ptr().add(base));
2249 d0 = _mm256_fmadd_ps(_mm256_loadu_ps(a.as_ptr().add(base)), vb0, d0);
2250 n0 = _mm256_fmadd_ps(vb0, vb0, n0);
2251 let vb1 = _mm256_loadu_ps(b.as_ptr().add(base + 8));
2252 d1 = _mm256_fmadd_ps(_mm256_loadu_ps(a.as_ptr().add(base + 8)), vb1, d1);
2253 n1 = _mm256_fmadd_ps(vb1, vb1, n1);
2254 let vb2 = _mm256_loadu_ps(b.as_ptr().add(base + 16));
2255 d2 = _mm256_fmadd_ps(_mm256_loadu_ps(a.as_ptr().add(base + 16)), vb2, d2);
2256 n2 = _mm256_fmadd_ps(vb2, vb2, n2);
2257 let vb3 = _mm256_loadu_ps(b.as_ptr().add(base + 24));
2258 d3 = _mm256_fmadd_ps(_mm256_loadu_ps(a.as_ptr().add(base + 24)), vb3, d3);
2259 n3 = _mm256_fmadd_ps(vb3, vb3, n3);
2260 }
2261
2262 let acc_dot = _mm256_add_ps(_mm256_add_ps(d0, d1), _mm256_add_ps(d2, d3));
2263 let acc_norm = _mm256_add_ps(_mm256_add_ps(n0, n1), _mm256_add_ps(n2, n3));
2264
2265 let hi_d = _mm256_extractf128_ps(acc_dot, 1);
2267 let lo_d = _mm256_castps256_ps128(acc_dot);
2268 let sum_d = _mm_add_ps(lo_d, hi_d);
2269 let shuf_d = _mm_shuffle_ps(sum_d, sum_d, 0b10_11_00_01);
2270 let sums_d = _mm_add_ps(sum_d, shuf_d);
2271 let shuf2_d = _mm_movehl_ps(sums_d, sums_d);
2272 let mut dot = _mm_cvtss_f32(_mm_add_ss(sums_d, shuf2_d));
2273
2274 let hi_n = _mm256_extractf128_ps(acc_norm, 1);
2275 let lo_n = _mm256_castps256_ps128(acc_norm);
2276 let sum_n = _mm_add_ps(lo_n, hi_n);
2277 let shuf_n = _mm_shuffle_ps(sum_n, sum_n, 0b10_11_00_01);
2278 let sums_n = _mm_add_ps(sum_n, shuf_n);
2279 let shuf2_n = _mm_movehl_ps(sums_n, sums_n);
2280 let mut norm = _mm_cvtss_f32(_mm_add_ss(sums_n, shuf2_n));
2281
2282 let base = chunks32 * 32;
2283 for i in 0..remainder {
2284 dot += a[base + i] * b[base + i];
2285 norm += b[base + i] * b[base + i];
2286 }
2287
2288 (dot, norm)
2289}
2290
2291#[cfg(target_arch = "x86_64")]
2292#[target_feature(enable = "sse")]
2293#[allow(unsafe_op_in_unsafe_fn)]
2294unsafe fn fused_dot_norm_sse(a: &[f32], b: &[f32], count: usize) -> (f32, f32) {
2295 use std::arch::x86_64::*;
2296
2297 let chunks = count / 4;
2298 let remainder = count % 4;
2299
2300 let mut acc_dot = _mm_setzero_ps();
2301 let mut acc_norm = _mm_setzero_ps();
2302
2303 for chunk in 0..chunks {
2304 let base = chunk * 4;
2305 let va = _mm_loadu_ps(a.as_ptr().add(base));
2306 let vb = _mm_loadu_ps(b.as_ptr().add(base));
2307 acc_dot = _mm_add_ps(acc_dot, _mm_mul_ps(va, vb));
2308 acc_norm = _mm_add_ps(acc_norm, _mm_mul_ps(vb, vb));
2309 }
2310
2311 let shuf_d = _mm_shuffle_ps(acc_dot, acc_dot, 0b10_11_00_01);
2313 let sums_d = _mm_add_ps(acc_dot, shuf_d);
2314 let shuf2_d = _mm_movehl_ps(sums_d, sums_d);
2315 let final_d = _mm_add_ss(sums_d, shuf2_d);
2316 let mut dot = _mm_cvtss_f32(final_d);
2317
2318 let shuf_n = _mm_shuffle_ps(acc_norm, acc_norm, 0b10_11_00_01);
2319 let sums_n = _mm_add_ps(acc_norm, shuf_n);
2320 let shuf2_n = _mm_movehl_ps(sums_n, sums_n);
2321 let final_n = _mm_add_ss(sums_n, shuf2_n);
2322 let mut norm = _mm_cvtss_f32(final_n);
2323
2324 let base = chunks * 4;
2325 for i in 0..remainder {
2326 dot += a[base + i] * b[base + i];
2327 norm += b[base + i] * b[base + i];
2328 }
2329
2330 (dot, norm)
2331}
2332
2333#[inline]
2339pub fn fast_inv_sqrt(x: f32) -> f32 {
2340 let half = 0.5 * x;
2341 let i = 0x5F37_5A86_u32.wrapping_sub(x.to_bits() >> 1);
2342 let y = f32::from_bits(i);
2343 let y = y * (1.5 - half * y * y); y * (1.5 - half * y * y) }
2346
2347#[inline]
2358pub fn batch_cosine_scores(query: &[f32], vectors: &[f32], dim: usize, scores: &mut [f32]) {
2359 let n = scores.len();
2360 let required = n
2361 .checked_mul(dim)
2362 .expect("batch cosine vector length overflow");
2363 assert_eq!(query.len(), dim, "batch cosine query dimension mismatch");
2364 assert!(
2365 vectors.len() >= required,
2366 "batch cosine vectors are truncated: need {required}, got {}",
2367 vectors.len()
2368 );
2369
2370 if dim == 0 || n == 0 {
2371 return;
2372 }
2373
2374 let norm_q_sq = dot_product_f32(query, query, dim);
2376 if norm_q_sq < f32::EPSILON {
2377 for s in scores.iter_mut() {
2378 *s = 0.0;
2379 }
2380 return;
2381 }
2382 let inv_norm_q = fast_inv_sqrt(norm_q_sq);
2383
2384 for i in 0..n {
2385 let vec = &vectors[i * dim..(i + 1) * dim];
2386 let (dot, norm_v_sq) = fused_dot_norm(query, vec, dim);
2387 if norm_v_sq < f32::EPSILON {
2388 scores[i] = 0.0;
2389 } else {
2390 scores[i] = dot * inv_norm_q * fast_inv_sqrt(norm_v_sq);
2391 }
2392 }
2393}
2394
2395#[inline]
2401pub fn f32_to_f16(value: f32) -> u16 {
2402 let bits = value.to_bits();
2403 let sign = (bits >> 16) & 0x8000;
2404 let exp = ((bits >> 23) & 0xFF) as i32;
2405 let mantissa = bits & 0x7F_FFFF;
2406
2407 if exp == 255 {
2408 return (sign | 0x7C00 | ((mantissa >> 13) & 0x3FF)) as u16;
2410 }
2411
2412 let exp16 = exp - 127 + 15;
2413
2414 if exp16 >= 31 {
2415 return (sign | 0x7C00) as u16; }
2417
2418 if exp16 <= 0 {
2419 if exp16 < -10 {
2420 return sign as u16; }
2422 let shift = (1 - exp16) as u32;
2423 let m = (mantissa | 0x80_0000) >> shift;
2424 let round_bit = (m >> 12) & 1;
2426 let sticky = m & 0xFFF;
2427 let m13 = m >> 13;
2428 let rounded = m13 + (round_bit & (m13 | if sticky != 0 { 1 } else { 0 }));
2429 return (sign | rounded) as u16;
2430 }
2431
2432 let round_bit = (mantissa >> 12) & 1;
2434 let sticky = mantissa & 0xFFF;
2435 let m13 = mantissa >> 13;
2436 let rounded = m13 + (round_bit & (m13 | if sticky != 0 { 1 } else { 0 }));
2437 if rounded > 0x3FF {
2439 let exp16_inc = exp16 as u32 + 1;
2440 if exp16_inc >= 31 {
2441 return (sign | 0x7C00) as u16; }
2443 (sign | (exp16_inc << 10)) as u16
2444 } else {
2445 (sign | ((exp16 as u32) << 10) | rounded) as u16
2446 }
2447}
2448
2449#[inline]
2451pub fn f16_to_f32(half: u16) -> f32 {
2452 let sign = ((half & 0x8000) as u32) << 16;
2453 let exp = ((half >> 10) & 0x1F) as u32;
2454 let mantissa = (half & 0x3FF) as u32;
2455
2456 if exp == 0 {
2457 if mantissa == 0 {
2458 return f32::from_bits(sign);
2459 }
2460 let mut e = 0u32;
2462 let mut m = mantissa;
2463 while (m & 0x400) == 0 {
2464 m <<= 1;
2465 e += 1;
2466 }
2467 return f32::from_bits(sign | ((127 - 15 + 1 - e) << 23) | ((m & 0x3FF) << 13));
2468 }
2469
2470 if exp == 31 {
2471 return f32::from_bits(sign | 0x7F80_0000 | (mantissa << 13));
2472 }
2473
2474 f32::from_bits(sign | ((exp + 127 - 15) << 23) | (mantissa << 13))
2475}
2476
2477const U8_SCALE: f32 = 127.5;
2482const U8_INV_SCALE: f32 = 1.0 / 127.5;
2483
2484#[inline]
2486pub fn f32_to_u8_saturating(value: f32) -> u8 {
2487 ((value.clamp(-1.0, 1.0) + 1.0) * U8_SCALE) as u8
2488}
2489
2490#[inline]
2492pub fn u8_to_f32(byte: u8) -> f32 {
2493 byte as f32 * U8_INV_SCALE - 1.0
2494}
2495
2496pub fn batch_f32_to_f16(src: &[f32], dst: &mut [u16]) {
2502 debug_assert_eq!(src.len(), dst.len());
2503 for (s, d) in src.iter().zip(dst.iter_mut()) {
2504 *d = f32_to_f16(*s);
2505 }
2506}
2507
2508pub fn batch_f32_to_u8(src: &[f32], dst: &mut [u8]) {
2510 debug_assert_eq!(src.len(), dst.len());
2511 for (s, d) in src.iter().zip(dst.iter_mut()) {
2512 *d = f32_to_u8_saturating(*s);
2513 }
2514}
2515
2516#[cfg(target_arch = "aarch64")]
2521#[allow(unsafe_op_in_unsafe_fn)]
2522mod neon_quant {
2523 use std::arch::aarch64::*;
2524
2525 #[allow(clippy::incompatible_msrv)]
2531 #[target_feature(enable = "neon")]
2532 pub unsafe fn fused_dot_norm_f16(query_f16: &[u16], vec_f16: &[u16], dim: usize) -> (f32, f32) {
2533 let chunks16 = dim / 16;
2534 let remainder = dim % 16;
2535
2536 let mut acc_dot0 = vdupq_n_f32(0.0);
2538 let mut acc_dot1 = vdupq_n_f32(0.0);
2539 let mut acc_norm0 = vdupq_n_f32(0.0);
2540 let mut acc_norm1 = vdupq_n_f32(0.0);
2541
2542 for c in 0..chunks16 {
2543 let base = c * 16;
2544
2545 let v_raw0 = vld1q_u16(vec_f16.as_ptr().add(base));
2547 let v_lo0 = vcvt_f32_f16(vreinterpret_f16_u16(vget_low_u16(v_raw0)));
2548 let v_hi0 = vcvt_f32_f16(vreinterpret_f16_u16(vget_high_u16(v_raw0)));
2549 let q_raw0 = vld1q_u16(query_f16.as_ptr().add(base));
2550 let q_lo0 = vcvt_f32_f16(vreinterpret_f16_u16(vget_low_u16(q_raw0)));
2551 let q_hi0 = vcvt_f32_f16(vreinterpret_f16_u16(vget_high_u16(q_raw0)));
2552
2553 acc_dot0 = vfmaq_f32(acc_dot0, q_lo0, v_lo0);
2554 acc_dot0 = vfmaq_f32(acc_dot0, q_hi0, v_hi0);
2555 acc_norm0 = vfmaq_f32(acc_norm0, v_lo0, v_lo0);
2556 acc_norm0 = vfmaq_f32(acc_norm0, v_hi0, v_hi0);
2557
2558 let v_raw1 = vld1q_u16(vec_f16.as_ptr().add(base + 8));
2560 let v_lo1 = vcvt_f32_f16(vreinterpret_f16_u16(vget_low_u16(v_raw1)));
2561 let v_hi1 = vcvt_f32_f16(vreinterpret_f16_u16(vget_high_u16(v_raw1)));
2562 let q_raw1 = vld1q_u16(query_f16.as_ptr().add(base + 8));
2563 let q_lo1 = vcvt_f32_f16(vreinterpret_f16_u16(vget_low_u16(q_raw1)));
2564 let q_hi1 = vcvt_f32_f16(vreinterpret_f16_u16(vget_high_u16(q_raw1)));
2565
2566 acc_dot1 = vfmaq_f32(acc_dot1, q_lo1, v_lo1);
2567 acc_dot1 = vfmaq_f32(acc_dot1, q_hi1, v_hi1);
2568 acc_norm1 = vfmaq_f32(acc_norm1, v_lo1, v_lo1);
2569 acc_norm1 = vfmaq_f32(acc_norm1, v_hi1, v_hi1);
2570 }
2571
2572 let mut dot = vaddvq_f32(vaddq_f32(acc_dot0, acc_dot1));
2574 let mut norm = vaddvq_f32(vaddq_f32(acc_norm0, acc_norm1));
2575
2576 let base = chunks16 * 16;
2578 for i in 0..remainder {
2579 let v = super::f16_to_f32(*vec_f16.get_unchecked(base + i));
2580 let q = super::f16_to_f32(*query_f16.get_unchecked(base + i));
2581 dot += q * v;
2582 norm += v * v;
2583 }
2584
2585 (dot, norm)
2586 }
2587
2588 #[target_feature(enable = "neon")]
2591 pub unsafe fn fused_dot_norm_u8(query: &[f32], vec_u8: &[u8], dim: usize) -> (f32, f32) {
2592 let scale = vdupq_n_f32(super::U8_INV_SCALE);
2593 let offset = vdupq_n_f32(-1.0);
2594
2595 let chunks16 = dim / 16;
2596 let remainder = dim % 16;
2597
2598 let mut acc_dot = vdupq_n_f32(0.0);
2599 let mut acc_norm = vdupq_n_f32(0.0);
2600
2601 for c in 0..chunks16 {
2602 let base = c * 16;
2603
2604 let bytes = vld1q_u8(vec_u8.as_ptr().add(base));
2606
2607 let lo8 = vget_low_u8(bytes);
2609 let hi8 = vget_high_u8(bytes);
2610 let lo16 = vmovl_u8(lo8);
2611 let hi16 = vmovl_u8(hi8);
2612
2613 let f0 = vaddq_f32(
2614 vmulq_f32(vcvtq_f32_u32(vmovl_u16(vget_low_u16(lo16))), scale),
2615 offset,
2616 );
2617 let f1 = vaddq_f32(
2618 vmulq_f32(vcvtq_f32_u32(vmovl_u16(vget_high_u16(lo16))), scale),
2619 offset,
2620 );
2621 let f2 = vaddq_f32(
2622 vmulq_f32(vcvtq_f32_u32(vmovl_u16(vget_low_u16(hi16))), scale),
2623 offset,
2624 );
2625 let f3 = vaddq_f32(
2626 vmulq_f32(vcvtq_f32_u32(vmovl_u16(vget_high_u16(hi16))), scale),
2627 offset,
2628 );
2629
2630 let q0 = vld1q_f32(query.as_ptr().add(base));
2631 let q1 = vld1q_f32(query.as_ptr().add(base + 4));
2632 let q2 = vld1q_f32(query.as_ptr().add(base + 8));
2633 let q3 = vld1q_f32(query.as_ptr().add(base + 12));
2634
2635 acc_dot = vfmaq_f32(acc_dot, q0, f0);
2636 acc_dot = vfmaq_f32(acc_dot, q1, f1);
2637 acc_dot = vfmaq_f32(acc_dot, q2, f2);
2638 acc_dot = vfmaq_f32(acc_dot, q3, f3);
2639
2640 acc_norm = vfmaq_f32(acc_norm, f0, f0);
2641 acc_norm = vfmaq_f32(acc_norm, f1, f1);
2642 acc_norm = vfmaq_f32(acc_norm, f2, f2);
2643 acc_norm = vfmaq_f32(acc_norm, f3, f3);
2644 }
2645
2646 let mut dot = vaddvq_f32(acc_dot);
2647 let mut norm = vaddvq_f32(acc_norm);
2648
2649 let base = chunks16 * 16;
2650 for i in 0..remainder {
2651 let v = super::u8_to_f32(*vec_u8.get_unchecked(base + i));
2652 dot += *query.get_unchecked(base + i) * v;
2653 norm += v * v;
2654 }
2655
2656 (dot, norm)
2657 }
2658
2659 #[allow(clippy::incompatible_msrv)]
2661 #[target_feature(enable = "neon")]
2662 pub unsafe fn dot_product_f16(query_f16: &[u16], vec_f16: &[u16], dim: usize) -> f32 {
2663 let chunks8 = dim / 8;
2664 let remainder = dim % 8;
2665
2666 let mut acc = vdupq_n_f32(0.0);
2667
2668 for c in 0..chunks8 {
2669 let base = c * 8;
2670 let v_raw = vld1q_u16(vec_f16.as_ptr().add(base));
2671 let v_lo = vcvt_f32_f16(vreinterpret_f16_u16(vget_low_u16(v_raw)));
2672 let v_hi = vcvt_f32_f16(vreinterpret_f16_u16(vget_high_u16(v_raw)));
2673 let q_raw = vld1q_u16(query_f16.as_ptr().add(base));
2674 let q_lo = vcvt_f32_f16(vreinterpret_f16_u16(vget_low_u16(q_raw)));
2675 let q_hi = vcvt_f32_f16(vreinterpret_f16_u16(vget_high_u16(q_raw)));
2676 acc = vfmaq_f32(acc, q_lo, v_lo);
2677 acc = vfmaq_f32(acc, q_hi, v_hi);
2678 }
2679
2680 let mut dot = vaddvq_f32(acc);
2681 let base = chunks8 * 8;
2682 for i in 0..remainder {
2683 let v = super::f16_to_f32(*vec_f16.get_unchecked(base + i));
2684 let q = super::f16_to_f32(*query_f16.get_unchecked(base + i));
2685 dot += q * v;
2686 }
2687 dot
2688 }
2689
2690 #[target_feature(enable = "neon")]
2692 pub unsafe fn dot_product_u8(query: &[f32], vec_u8: &[u8], dim: usize) -> f32 {
2693 let scale = vdupq_n_f32(super::U8_INV_SCALE);
2694 let offset = vdupq_n_f32(-1.0);
2695 let chunks16 = dim / 16;
2696 let remainder = dim % 16;
2697
2698 let mut acc = vdupq_n_f32(0.0);
2699
2700 for c in 0..chunks16 {
2701 let base = c * 16;
2702 let bytes = vld1q_u8(vec_u8.as_ptr().add(base));
2703 let lo8 = vget_low_u8(bytes);
2704 let hi8 = vget_high_u8(bytes);
2705 let lo16 = vmovl_u8(lo8);
2706 let hi16 = vmovl_u8(hi8);
2707 let f0 = vaddq_f32(
2708 vmulq_f32(vcvtq_f32_u32(vmovl_u16(vget_low_u16(lo16))), scale),
2709 offset,
2710 );
2711 let f1 = vaddq_f32(
2712 vmulq_f32(vcvtq_f32_u32(vmovl_u16(vget_high_u16(lo16))), scale),
2713 offset,
2714 );
2715 let f2 = vaddq_f32(
2716 vmulq_f32(vcvtq_f32_u32(vmovl_u16(vget_low_u16(hi16))), scale),
2717 offset,
2718 );
2719 let f3 = vaddq_f32(
2720 vmulq_f32(vcvtq_f32_u32(vmovl_u16(vget_high_u16(hi16))), scale),
2721 offset,
2722 );
2723 let q0 = vld1q_f32(query.as_ptr().add(base));
2724 let q1 = vld1q_f32(query.as_ptr().add(base + 4));
2725 let q2 = vld1q_f32(query.as_ptr().add(base + 8));
2726 let q3 = vld1q_f32(query.as_ptr().add(base + 12));
2727 acc = vfmaq_f32(acc, q0, f0);
2728 acc = vfmaq_f32(acc, q1, f1);
2729 acc = vfmaq_f32(acc, q2, f2);
2730 acc = vfmaq_f32(acc, q3, f3);
2731 }
2732
2733 let mut dot = vaddvq_f32(acc);
2734 let base = chunks16 * 16;
2735 for i in 0..remainder {
2736 let v = super::u8_to_f32(*vec_u8.get_unchecked(base + i));
2737 dot += *query.get_unchecked(base + i) * v;
2738 }
2739 dot
2740 }
2741}
2742
2743#[allow(dead_code)]
2748fn fused_dot_norm_f16_scalar(query_f16: &[u16], vec_f16: &[u16], dim: usize) -> (f32, f32) {
2749 (0..dim).fold((0.0f32, 0.0f32), |(dot, norm), i| {
2750 let v = f16_to_f32(vec_f16[i]);
2751 let q = f16_to_f32(query_f16[i]);
2752 (
2753 dot.algebraic_add(q.algebraic_mul(v)),
2754 norm.algebraic_add(v.algebraic_mul(v)),
2755 )
2756 })
2757}
2758
2759#[allow(dead_code)]
2760fn fused_dot_norm_u8_scalar(query: &[f32], vec_u8: &[u8], dim: usize) -> (f32, f32) {
2761 (0..dim).fold((0.0f32, 0.0f32), |(dot, norm), i| {
2762 let v = u8_to_f32(vec_u8[i]);
2763 (
2764 dot.algebraic_add(query[i].algebraic_mul(v)),
2765 norm.algebraic_add(v.algebraic_mul(v)),
2766 )
2767 })
2768}
2769
2770#[allow(dead_code)]
2771fn dot_product_f16_scalar(query_f16: &[u16], vec_f16: &[u16], dim: usize) -> f32 {
2772 (0..dim).fold(0.0f32, |dot, i| {
2773 dot.algebraic_add(f16_to_f32(query_f16[i]).algebraic_mul(f16_to_f32(vec_f16[i])))
2774 })
2775}
2776
2777#[allow(dead_code)]
2778fn dot_product_u8_scalar(query: &[f32], vec_u8: &[u8], dim: usize) -> f32 {
2779 (0..dim).fold(0.0f32, |dot, i| {
2780 dot.algebraic_add(query[i].algebraic_mul(u8_to_f32(vec_u8[i])))
2781 })
2782}
2783
2784#[cfg(target_arch = "x86_64")]
2789#[target_feature(enable = "sse2", enable = "sse4.1")]
2790#[allow(unsafe_op_in_unsafe_fn)]
2791unsafe fn fused_dot_norm_f16_sse(query_f16: &[u16], vec_f16: &[u16], dim: usize) -> (f32, f32) {
2792 use std::arch::x86_64::*;
2793
2794 let chunks = dim / 4;
2795 let remainder = dim % 4;
2796
2797 let mut acc_dot = _mm_setzero_ps();
2798 let mut acc_norm = _mm_setzero_ps();
2799
2800 for chunk in 0..chunks {
2801 let base = chunk * 4;
2802 let v0 = f16_to_f32(*vec_f16.get_unchecked(base));
2804 let v1 = f16_to_f32(*vec_f16.get_unchecked(base + 1));
2805 let v2 = f16_to_f32(*vec_f16.get_unchecked(base + 2));
2806 let v3 = f16_to_f32(*vec_f16.get_unchecked(base + 3));
2807 let vb = _mm_set_ps(v3, v2, v1, v0);
2808
2809 let q0 = f16_to_f32(*query_f16.get_unchecked(base));
2810 let q1 = f16_to_f32(*query_f16.get_unchecked(base + 1));
2811 let q2 = f16_to_f32(*query_f16.get_unchecked(base + 2));
2812 let q3 = f16_to_f32(*query_f16.get_unchecked(base + 3));
2813 let va = _mm_set_ps(q3, q2, q1, q0);
2814
2815 acc_dot = _mm_add_ps(acc_dot, _mm_mul_ps(va, vb));
2816 acc_norm = _mm_add_ps(acc_norm, _mm_mul_ps(vb, vb));
2817 }
2818
2819 let shuf_d = _mm_shuffle_ps(acc_dot, acc_dot, 0b10_11_00_01);
2821 let sums_d = _mm_add_ps(acc_dot, shuf_d);
2822 let shuf2_d = _mm_movehl_ps(sums_d, sums_d);
2823 let mut dot = _mm_cvtss_f32(_mm_add_ss(sums_d, shuf2_d));
2824
2825 let shuf_n = _mm_shuffle_ps(acc_norm, acc_norm, 0b10_11_00_01);
2826 let sums_n = _mm_add_ps(acc_norm, shuf_n);
2827 let shuf2_n = _mm_movehl_ps(sums_n, sums_n);
2828 let mut norm = _mm_cvtss_f32(_mm_add_ss(sums_n, shuf2_n));
2829
2830 let base = chunks * 4;
2831 for i in 0..remainder {
2832 let v = f16_to_f32(*vec_f16.get_unchecked(base + i));
2833 let q = f16_to_f32(*query_f16.get_unchecked(base + i));
2834 dot += q * v;
2835 norm += v * v;
2836 }
2837
2838 (dot, norm)
2839}
2840
2841#[cfg(target_arch = "x86_64")]
2842#[target_feature(enable = "sse2", enable = "sse4.1")]
2843#[allow(unsafe_op_in_unsafe_fn)]
2844unsafe fn fused_dot_norm_u8_sse(query: &[f32], vec_u8: &[u8], dim: usize) -> (f32, f32) {
2845 use std::arch::x86_64::*;
2846
2847 let scale = _mm_set1_ps(U8_INV_SCALE);
2848 let offset = _mm_set1_ps(-1.0);
2849
2850 let chunks = dim / 4;
2851 let remainder = dim % 4;
2852
2853 let mut acc_dot = _mm_setzero_ps();
2854 let mut acc_norm = _mm_setzero_ps();
2855
2856 for chunk in 0..chunks {
2857 let base = chunk * 4;
2858
2859 let bytes = _mm_cvtsi32_si128(std::ptr::read_unaligned(
2861 vec_u8.as_ptr().add(base) as *const i32
2862 ));
2863 let ints = _mm_cvtepu8_epi32(bytes);
2864 let floats = _mm_cvtepi32_ps(ints);
2865 let vb = _mm_add_ps(_mm_mul_ps(floats, scale), offset);
2866
2867 let va = _mm_loadu_ps(query.as_ptr().add(base));
2868
2869 acc_dot = _mm_add_ps(acc_dot, _mm_mul_ps(va, vb));
2870 acc_norm = _mm_add_ps(acc_norm, _mm_mul_ps(vb, vb));
2871 }
2872
2873 let shuf_d = _mm_shuffle_ps(acc_dot, acc_dot, 0b10_11_00_01);
2875 let sums_d = _mm_add_ps(acc_dot, shuf_d);
2876 let shuf2_d = _mm_movehl_ps(sums_d, sums_d);
2877 let mut dot = _mm_cvtss_f32(_mm_add_ss(sums_d, shuf2_d));
2878
2879 let shuf_n = _mm_shuffle_ps(acc_norm, acc_norm, 0b10_11_00_01);
2880 let sums_n = _mm_add_ps(acc_norm, shuf_n);
2881 let shuf2_n = _mm_movehl_ps(sums_n, sums_n);
2882 let mut norm = _mm_cvtss_f32(_mm_add_ss(sums_n, shuf2_n));
2883
2884 let base = chunks * 4;
2885 for i in 0..remainder {
2886 let v = u8_to_f32(*vec_u8.get_unchecked(base + i));
2887 dot += *query.get_unchecked(base + i) * v;
2888 norm += v * v;
2889 }
2890
2891 (dot, norm)
2892}
2893
2894#[cfg(target_arch = "x86_64")]
2899#[target_feature(enable = "avx", enable = "f16c", enable = "fma")]
2900#[allow(unsafe_op_in_unsafe_fn)]
2901unsafe fn fused_dot_norm_f16_f16c(query_f16: &[u16], vec_f16: &[u16], dim: usize) -> (f32, f32) {
2902 use std::arch::x86_64::*;
2903
2904 let chunks16 = dim / 16;
2905 let remainder = dim % 16;
2906
2907 let mut acc_dot0 = _mm256_setzero_ps();
2909 let mut acc_dot1 = _mm256_setzero_ps();
2910 let mut acc_norm0 = _mm256_setzero_ps();
2911 let mut acc_norm1 = _mm256_setzero_ps();
2912
2913 for c in 0..chunks16 {
2914 let base = c * 16;
2915
2916 let v_raw0 = _mm_loadu_si128(vec_f16.as_ptr().add(base) as *const __m128i);
2918 let vb0 = _mm256_cvtph_ps(v_raw0);
2919 let q_raw0 = _mm_loadu_si128(query_f16.as_ptr().add(base) as *const __m128i);
2920 let qa0 = _mm256_cvtph_ps(q_raw0);
2921 acc_dot0 = _mm256_fmadd_ps(qa0, vb0, acc_dot0);
2922 acc_norm0 = _mm256_fmadd_ps(vb0, vb0, acc_norm0);
2923
2924 let v_raw1 = _mm_loadu_si128(vec_f16.as_ptr().add(base + 8) as *const __m128i);
2926 let vb1 = _mm256_cvtph_ps(v_raw1);
2927 let q_raw1 = _mm_loadu_si128(query_f16.as_ptr().add(base + 8) as *const __m128i);
2928 let qa1 = _mm256_cvtph_ps(q_raw1);
2929 acc_dot1 = _mm256_fmadd_ps(qa1, vb1, acc_dot1);
2930 acc_norm1 = _mm256_fmadd_ps(vb1, vb1, acc_norm1);
2931 }
2932
2933 let acc_dot = _mm256_add_ps(acc_dot0, acc_dot1);
2935 let acc_norm = _mm256_add_ps(acc_norm0, acc_norm1);
2936
2937 let hi_d = _mm256_extractf128_ps(acc_dot, 1);
2939 let lo_d = _mm256_castps256_ps128(acc_dot);
2940 let sum_d = _mm_add_ps(lo_d, hi_d);
2941 let shuf_d = _mm_shuffle_ps(sum_d, sum_d, 0b10_11_00_01);
2942 let sums_d = _mm_add_ps(sum_d, shuf_d);
2943 let shuf2_d = _mm_movehl_ps(sums_d, sums_d);
2944 let mut dot = _mm_cvtss_f32(_mm_add_ss(sums_d, shuf2_d));
2945
2946 let hi_n = _mm256_extractf128_ps(acc_norm, 1);
2947 let lo_n = _mm256_castps256_ps128(acc_norm);
2948 let sum_n = _mm_add_ps(lo_n, hi_n);
2949 let shuf_n = _mm_shuffle_ps(sum_n, sum_n, 0b10_11_00_01);
2950 let sums_n = _mm_add_ps(sum_n, shuf_n);
2951 let shuf2_n = _mm_movehl_ps(sums_n, sums_n);
2952 let mut norm = _mm_cvtss_f32(_mm_add_ss(sums_n, shuf2_n));
2953
2954 let base = chunks16 * 16;
2955 for i in 0..remainder {
2956 let v = f16_to_f32(*vec_f16.get_unchecked(base + i));
2957 let q = f16_to_f32(*query_f16.get_unchecked(base + i));
2958 dot += q * v;
2959 norm += v * v;
2960 }
2961
2962 (dot, norm)
2963}
2964
2965#[cfg(target_arch = "x86_64")]
2966#[target_feature(enable = "avx", enable = "f16c", enable = "fma")]
2967#[allow(unsafe_op_in_unsafe_fn)]
2968unsafe fn dot_product_f16_f16c(query_f16: &[u16], vec_f16: &[u16], dim: usize) -> f32 {
2969 use std::arch::x86_64::*;
2970
2971 let chunks = dim / 8;
2972 let remainder = dim % 8;
2973 let mut acc = _mm256_setzero_ps();
2974
2975 for chunk in 0..chunks {
2976 let base = chunk * 8;
2977 let v_raw = _mm_loadu_si128(vec_f16.as_ptr().add(base) as *const __m128i);
2978 let vb = _mm256_cvtph_ps(v_raw);
2979 let q_raw = _mm_loadu_si128(query_f16.as_ptr().add(base) as *const __m128i);
2980 let qa = _mm256_cvtph_ps(q_raw);
2981 acc = _mm256_fmadd_ps(qa, vb, acc);
2982 }
2983
2984 let hi = _mm256_extractf128_ps(acc, 1);
2985 let lo = _mm256_castps256_ps128(acc);
2986 let sum = _mm_add_ps(lo, hi);
2987 let shuf = _mm_shuffle_ps(sum, sum, 0b10_11_00_01);
2988 let sums = _mm_add_ps(sum, shuf);
2989 let shuf2 = _mm_movehl_ps(sums, sums);
2990 let mut dot = _mm_cvtss_f32(_mm_add_ss(sums, shuf2));
2991
2992 let base = chunks * 8;
2993 for i in 0..remainder {
2994 let v = f16_to_f32(*vec_f16.get_unchecked(base + i));
2995 let q = f16_to_f32(*query_f16.get_unchecked(base + i));
2996 dot += q * v;
2997 }
2998 dot
2999}
3000
3001#[cfg(target_arch = "x86_64")]
3002#[target_feature(enable = "sse2", enable = "sse4.1")]
3003#[allow(unsafe_op_in_unsafe_fn)]
3004unsafe fn dot_product_u8_sse(query: &[f32], vec_u8: &[u8], dim: usize) -> f32 {
3005 use std::arch::x86_64::*;
3006
3007 let scale = _mm_set1_ps(U8_INV_SCALE);
3008 let offset = _mm_set1_ps(-1.0);
3009 let chunks = dim / 4;
3010 let remainder = dim % 4;
3011 let mut acc = _mm_setzero_ps();
3012
3013 for chunk in 0..chunks {
3014 let base = chunk * 4;
3015 let bytes = _mm_cvtsi32_si128(std::ptr::read_unaligned(
3016 vec_u8.as_ptr().add(base) as *const i32
3017 ));
3018 let ints = _mm_cvtepu8_epi32(bytes);
3019 let floats = _mm_cvtepi32_ps(ints);
3020 let vb = _mm_add_ps(_mm_mul_ps(floats, scale), offset);
3021 let va = _mm_loadu_ps(query.as_ptr().add(base));
3022 acc = _mm_add_ps(acc, _mm_mul_ps(va, vb));
3023 }
3024
3025 let shuf = _mm_shuffle_ps(acc, acc, 0b10_11_00_01);
3026 let sums = _mm_add_ps(acc, shuf);
3027 let shuf2 = _mm_movehl_ps(sums, sums);
3028 let mut dot = _mm_cvtss_f32(_mm_add_ss(sums, shuf2));
3029
3030 let base = chunks * 4;
3031 for i in 0..remainder {
3032 dot += *query.get_unchecked(base + i) * u8_to_f32(*vec_u8.get_unchecked(base + i));
3033 }
3034 dot
3035}
3036
3037#[inline]
3042fn fused_dot_norm_f16(query_f16: &[u16], vec_f16: &[u16], dim: usize) -> (f32, f32) {
3043 #[cfg(target_arch = "aarch64")]
3044 {
3045 return unsafe { neon_quant::fused_dot_norm_f16(query_f16, vec_f16, dim) };
3046 }
3047
3048 #[cfg(target_arch = "x86_64")]
3049 {
3050 if is_x86_feature_detected!("f16c") && is_x86_feature_detected!("fma") {
3051 return unsafe { fused_dot_norm_f16_f16c(query_f16, vec_f16, dim) };
3052 }
3053 if sse::is_available() {
3054 return unsafe { fused_dot_norm_f16_sse(query_f16, vec_f16, dim) };
3055 }
3056 }
3057
3058 #[allow(unreachable_code)]
3059 fused_dot_norm_f16_scalar(query_f16, vec_f16, dim)
3060}
3061
3062#[inline]
3063fn fused_dot_norm_u8(query: &[f32], vec_u8: &[u8], dim: usize) -> (f32, f32) {
3064 #[cfg(target_arch = "aarch64")]
3065 {
3066 return unsafe { neon_quant::fused_dot_norm_u8(query, vec_u8, dim) };
3067 }
3068
3069 #[cfg(target_arch = "x86_64")]
3070 {
3071 if sse::is_available() {
3072 return unsafe { fused_dot_norm_u8_sse(query, vec_u8, dim) };
3073 }
3074 }
3075
3076 #[allow(unreachable_code)]
3077 fused_dot_norm_u8_scalar(query, vec_u8, dim)
3078}
3079
3080#[inline]
3083fn dot_product_f16_quant(query_f16: &[u16], vec_f16: &[u16], dim: usize) -> f32 {
3084 #[cfg(target_arch = "aarch64")]
3085 {
3086 return unsafe { neon_quant::dot_product_f16(query_f16, vec_f16, dim) };
3087 }
3088
3089 #[cfg(target_arch = "x86_64")]
3090 {
3091 if is_x86_feature_detected!("f16c") && is_x86_feature_detected!("fma") {
3092 return unsafe { dot_product_f16_f16c(query_f16, vec_f16, dim) };
3093 }
3094 }
3095
3096 #[allow(unreachable_code)]
3097 dot_product_f16_scalar(query_f16, vec_f16, dim)
3098}
3099
3100#[inline]
3101fn dot_product_u8_quant(query: &[f32], vec_u8: &[u8], dim: usize) -> f32 {
3102 #[cfg(target_arch = "aarch64")]
3103 {
3104 return unsafe { neon_quant::dot_product_u8(query, vec_u8, dim) };
3105 }
3106
3107 #[cfg(target_arch = "x86_64")]
3108 {
3109 if sse::is_available() {
3110 return unsafe { dot_product_u8_sse(query, vec_u8, dim) };
3111 }
3112 }
3113
3114 #[allow(unreachable_code)]
3115 dot_product_u8_scalar(query, vec_u8, dim)
3116}
3117
3118#[inline]
3129pub fn batch_cosine_scores_f16(query: &[f32], vectors_raw: &[u8], dim: usize, scores: &mut [f32]) {
3130 let n = scores.len();
3131 let vec_bytes = dim.checked_mul(2).expect("f16 vector size overflow");
3132 let required = n
3133 .checked_mul(vec_bytes)
3134 .expect("f16 batch byte length overflow");
3135 assert_eq!(
3136 query.len(),
3137 dim,
3138 "f16 batch cosine query dimension mismatch"
3139 );
3140 assert!(
3141 vectors_raw.len() >= required,
3142 "f16 batch cosine vectors are truncated: need {required} bytes, got {}",
3143 vectors_raw.len()
3144 );
3145 if required > 0 {
3146 assert!(
3147 (vectors_raw.as_ptr() as usize).is_multiple_of(std::mem::align_of::<u16>()),
3148 "f16 batch cosine vectors are not 2-byte aligned"
3149 );
3150 }
3151 if dim == 0 || n == 0 {
3152 return;
3153 }
3154
3155 let norm_q_sq = dot_product_f32(query, query, dim);
3157 if norm_q_sq < f32::EPSILON {
3158 for s in scores.iter_mut() {
3159 *s = 0.0;
3160 }
3161 return;
3162 }
3163 let inv_norm_q = fast_inv_sqrt(norm_q_sq);
3164
3165 let query_f16: Vec<u16> = query.iter().map(|&v| f32_to_f16(v)).collect();
3167
3168 for i in 0..n {
3169 let raw = &vectors_raw[i * vec_bytes..(i + 1) * vec_bytes];
3170 let f16_slice = unsafe { std::slice::from_raw_parts(raw.as_ptr() as *const u16, dim) };
3171
3172 let (dot, norm_v_sq) = fused_dot_norm_f16(&query_f16, f16_slice, dim);
3173 scores[i] = if norm_v_sq < f32::EPSILON {
3174 0.0
3175 } else {
3176 dot * inv_norm_q * fast_inv_sqrt(norm_v_sq)
3177 };
3178 }
3179}
3180
3181#[inline]
3188pub fn batch_cosine_scores_u8(query: &[f32], vectors_raw: &[u8], dim: usize, scores: &mut [f32]) {
3189 let n = scores.len();
3190 let required = n.checked_mul(dim).expect("u8 batch byte length overflow");
3191 assert_eq!(query.len(), dim, "u8 batch cosine query dimension mismatch");
3192 assert!(
3193 vectors_raw.len() >= required,
3194 "u8 batch cosine vectors are truncated: need {required} bytes, got {}",
3195 vectors_raw.len()
3196 );
3197 if dim == 0 || n == 0 {
3198 return;
3199 }
3200
3201 let norm_q_sq = dot_product_f32(query, query, dim);
3202 if norm_q_sq < f32::EPSILON {
3203 for s in scores.iter_mut() {
3204 *s = 0.0;
3205 }
3206 return;
3207 }
3208 let inv_norm_q = fast_inv_sqrt(norm_q_sq);
3209
3210 for i in 0..n {
3211 let u8_slice = &vectors_raw[i * dim..(i + 1) * dim];
3212
3213 let (dot, norm_v_sq) = fused_dot_norm_u8(query, u8_slice, dim);
3214 scores[i] = if norm_v_sq < f32::EPSILON {
3215 0.0
3216 } else {
3217 dot * inv_norm_q * fast_inv_sqrt(norm_v_sq)
3218 };
3219 }
3220}
3221
3222#[inline]
3231pub fn batch_dot_scores(query: &[f32], vectors: &[f32], dim: usize, scores: &mut [f32]) {
3232 let n = scores.len();
3233 let required = n
3234 .checked_mul(dim)
3235 .expect("batch dot vector length overflow");
3236 assert_eq!(query.len(), dim, "batch dot query dimension mismatch");
3237 assert!(
3238 vectors.len() >= required,
3239 "batch dot vectors are truncated: need {required}, got {}",
3240 vectors.len()
3241 );
3242
3243 if dim == 0 || n == 0 {
3244 return;
3245 }
3246
3247 let norm_q_sq = dot_product_f32(query, query, dim);
3248 if norm_q_sq < f32::EPSILON {
3249 for s in scores.iter_mut() {
3250 *s = 0.0;
3251 }
3252 return;
3253 }
3254 let inv_norm_q = fast_inv_sqrt(norm_q_sq);
3255
3256 for i in 0..n {
3257 let vec = &vectors[i * dim..(i + 1) * dim];
3258 let dot = dot_product_f32(query, vec, dim);
3259 scores[i] = dot * inv_norm_q;
3260 }
3261}
3262
3263#[inline]
3268pub fn batch_dot_scores_f16(query: &[f32], vectors_raw: &[u8], dim: usize, scores: &mut [f32]) {
3269 let n = scores.len();
3270 let vec_bytes = dim.checked_mul(2).expect("f16 vector size overflow");
3271 let required = n
3272 .checked_mul(vec_bytes)
3273 .expect("f16 batch byte length overflow");
3274 assert_eq!(query.len(), dim, "f16 batch dot query dimension mismatch");
3275 assert!(
3276 vectors_raw.len() >= required,
3277 "f16 batch dot vectors are truncated: need {required} bytes, got {}",
3278 vectors_raw.len()
3279 );
3280 if required > 0 {
3281 assert!(
3282 (vectors_raw.as_ptr() as usize).is_multiple_of(std::mem::align_of::<u16>()),
3283 "f16 batch dot vectors are not 2-byte aligned"
3284 );
3285 }
3286 if dim == 0 || n == 0 {
3287 return;
3288 }
3289
3290 let norm_q_sq = dot_product_f32(query, query, dim);
3291 if norm_q_sq < f32::EPSILON {
3292 for s in scores.iter_mut() {
3293 *s = 0.0;
3294 }
3295 return;
3296 }
3297 let inv_norm_q = fast_inv_sqrt(norm_q_sq);
3298
3299 let query_f16: Vec<u16> = query.iter().map(|&v| f32_to_f16(v)).collect();
3300 for i in 0..n {
3301 let raw = &vectors_raw[i * vec_bytes..(i + 1) * vec_bytes];
3302 let f16_slice = unsafe { std::slice::from_raw_parts(raw.as_ptr() as *const u16, dim) };
3303 let dot = dot_product_f16_quant(&query_f16, f16_slice, dim);
3304 scores[i] = dot * inv_norm_q;
3305 }
3306}
3307
3308#[inline]
3313pub fn batch_dot_scores_u8(query: &[f32], vectors_raw: &[u8], dim: usize, scores: &mut [f32]) {
3314 let n = scores.len();
3315 let required = n.checked_mul(dim).expect("u8 batch byte length overflow");
3316 assert_eq!(query.len(), dim, "u8 batch dot query dimension mismatch");
3317 assert!(
3318 vectors_raw.len() >= required,
3319 "u8 batch dot vectors are truncated: need {required} bytes, got {}",
3320 vectors_raw.len()
3321 );
3322 if dim == 0 || n == 0 {
3323 return;
3324 }
3325
3326 let norm_q_sq = dot_product_f32(query, query, dim);
3327 if norm_q_sq < f32::EPSILON {
3328 for s in scores.iter_mut() {
3329 *s = 0.0;
3330 }
3331 return;
3332 }
3333 let inv_norm_q = fast_inv_sqrt(norm_q_sq);
3334
3335 for i in 0..n {
3336 let u8_slice = &vectors_raw[i * dim..(i + 1) * dim];
3337 let dot = dot_product_u8_quant(query, u8_slice, dim);
3338 scores[i] = dot * inv_norm_q;
3339 }
3340}
3341
3342#[inline]
3348pub fn batch_cosine_scores_precomp(
3349 query: &[f32],
3350 vectors: &[f32],
3351 dim: usize,
3352 scores: &mut [f32],
3353 inv_norm_q: f32,
3354) {
3355 let n = scores.len();
3356 let required = n
3357 .checked_mul(dim)
3358 .expect("precomputed cosine vector length overflow");
3359 assert_eq!(
3360 query.len(),
3361 dim,
3362 "precomputed cosine query dimension mismatch"
3363 );
3364 assert!(
3365 vectors.len() >= required,
3366 "precomputed cosine vectors are truncated: need {required}, got {}",
3367 vectors.len()
3368 );
3369 for i in 0..n {
3370 let vec = &vectors[i * dim..(i + 1) * dim];
3371 let (dot, norm_v_sq) = fused_dot_norm(query, vec, dim);
3372 scores[i] = if norm_v_sq < f32::EPSILON {
3373 0.0
3374 } else {
3375 dot * inv_norm_q * fast_inv_sqrt(norm_v_sq)
3376 };
3377 }
3378}
3379
3380#[inline]
3382pub fn batch_cosine_scores_f16_precomp(
3383 query_f16: &[u16],
3384 vectors_raw: &[u8],
3385 dim: usize,
3386 scores: &mut [f32],
3387 inv_norm_q: f32,
3388) {
3389 let n = scores.len();
3390 let vec_bytes = dim.checked_mul(2).expect("f16 vector size overflow");
3391 let required = n
3392 .checked_mul(vec_bytes)
3393 .expect("precomputed f16 cosine batch byte length overflow");
3394 assert_eq!(
3395 query_f16.len(),
3396 dim,
3397 "precomputed f16 cosine query dimension mismatch"
3398 );
3399 assert!(
3400 vectors_raw.len() >= required,
3401 "precomputed f16 cosine vectors are truncated: need {required} bytes, got {}",
3402 vectors_raw.len()
3403 );
3404 if required > 0 {
3405 assert!(
3406 (vectors_raw.as_ptr() as usize).is_multiple_of(std::mem::align_of::<u16>()),
3407 "precomputed f16 cosine vectors are not 2-byte aligned"
3408 );
3409 }
3410 for i in 0..n {
3411 let raw = &vectors_raw[i * vec_bytes..(i + 1) * vec_bytes];
3412 let f16_slice = unsafe { std::slice::from_raw_parts(raw.as_ptr() as *const u16, dim) };
3413 let (dot, norm_v_sq) = fused_dot_norm_f16(query_f16, f16_slice, dim);
3414 scores[i] = if norm_v_sq < f32::EPSILON {
3415 0.0
3416 } else {
3417 dot * inv_norm_q * fast_inv_sqrt(norm_v_sq)
3418 };
3419 }
3420}
3421
3422#[inline]
3424pub fn batch_cosine_scores_u8_precomp(
3425 query: &[f32],
3426 vectors_raw: &[u8],
3427 dim: usize,
3428 scores: &mut [f32],
3429 inv_norm_q: f32,
3430) {
3431 let n = scores.len();
3432 let required = n
3433 .checked_mul(dim)
3434 .expect("precomputed u8 cosine batch byte length overflow");
3435 assert_eq!(
3436 query.len(),
3437 dim,
3438 "precomputed u8 cosine query dimension mismatch"
3439 );
3440 assert!(
3441 vectors_raw.len() >= required,
3442 "precomputed u8 cosine vectors are truncated: need {required} bytes, got {}",
3443 vectors_raw.len()
3444 );
3445 for i in 0..n {
3446 let u8_slice = &vectors_raw[i * dim..(i + 1) * dim];
3447 let (dot, norm_v_sq) = fused_dot_norm_u8(query, u8_slice, dim);
3448 scores[i] = if norm_v_sq < f32::EPSILON {
3449 0.0
3450 } else {
3451 dot * inv_norm_q * fast_inv_sqrt(norm_v_sq)
3452 };
3453 }
3454}
3455
3456#[inline]
3458pub fn batch_dot_scores_precomp(
3459 query: &[f32],
3460 vectors: &[f32],
3461 dim: usize,
3462 scores: &mut [f32],
3463 inv_norm_q: f32,
3464) {
3465 let n = scores.len();
3466 let required = n
3467 .checked_mul(dim)
3468 .expect("precomputed dot vector length overflow");
3469 assert_eq!(query.len(), dim, "precomputed dot query dimension mismatch");
3470 assert!(
3471 vectors.len() >= required,
3472 "precomputed dot vectors are truncated: need {required}, got {}",
3473 vectors.len()
3474 );
3475 for i in 0..n {
3476 let vec = &vectors[i * dim..(i + 1) * dim];
3477 scores[i] = dot_product_f32(query, vec, dim) * inv_norm_q;
3478 }
3479}
3480
3481#[inline]
3483pub fn batch_dot_scores_f16_precomp(
3484 query_f16: &[u16],
3485 vectors_raw: &[u8],
3486 dim: usize,
3487 scores: &mut [f32],
3488 inv_norm_q: f32,
3489) {
3490 let n = scores.len();
3491 let vec_bytes = dim.checked_mul(2).expect("f16 vector size overflow");
3492 let required = n
3493 .checked_mul(vec_bytes)
3494 .expect("precomputed f16 dot batch byte length overflow");
3495 assert_eq!(
3496 query_f16.len(),
3497 dim,
3498 "precomputed f16 dot query dimension mismatch"
3499 );
3500 assert!(
3501 vectors_raw.len() >= required,
3502 "precomputed f16 dot vectors are truncated: need {required} bytes, got {}",
3503 vectors_raw.len()
3504 );
3505 if required > 0 {
3506 assert!(
3507 (vectors_raw.as_ptr() as usize).is_multiple_of(std::mem::align_of::<u16>()),
3508 "precomputed f16 dot vectors are not 2-byte aligned"
3509 );
3510 }
3511 for i in 0..n {
3512 let raw = &vectors_raw[i * vec_bytes..(i + 1) * vec_bytes];
3513 let f16_slice = unsafe { std::slice::from_raw_parts(raw.as_ptr() as *const u16, dim) };
3514 scores[i] = dot_product_f16_quant(query_f16, f16_slice, dim) * inv_norm_q;
3515 }
3516}
3517
3518#[inline]
3520pub fn batch_dot_scores_u8_precomp(
3521 query: &[f32],
3522 vectors_raw: &[u8],
3523 dim: usize,
3524 scores: &mut [f32],
3525 inv_norm_q: f32,
3526) {
3527 let n = scores.len();
3528 let required = n
3529 .checked_mul(dim)
3530 .expect("precomputed u8 dot batch byte length overflow");
3531 assert_eq!(
3532 query.len(),
3533 dim,
3534 "precomputed u8 dot query dimension mismatch"
3535 );
3536 assert!(
3537 vectors_raw.len() >= required,
3538 "precomputed u8 dot vectors are truncated: need {required} bytes, got {}",
3539 vectors_raw.len()
3540 );
3541 for i in 0..n {
3542 let u8_slice = &vectors_raw[i * dim..(i + 1) * dim];
3543 scores[i] = dot_product_u8_quant(query, u8_slice, dim) * inv_norm_q;
3544 }
3545}
3546
3547#[inline]
3552pub fn cosine_similarity(a: &[f32], b: &[f32]) -> f32 {
3553 assert_eq!(a.len(), b.len(), "cosine vector dimension mismatch");
3554 let count = a.len();
3555
3556 if count == 0 {
3557 return 0.0;
3558 }
3559
3560 let dot = dot_product_f32(a, b, count);
3561 let norm_a = dot_product_f32(a, a, count);
3562 let norm_b = dot_product_f32(b, b, count);
3563
3564 let denom = (norm_a * norm_b).sqrt();
3565 if denom < f32::EPSILON {
3566 return 0.0;
3567 }
3568
3569 dot / denom
3570}
3571
3572#[cfg(target_arch = "x86_64")]
3581#[target_feature(enable = "avx512f,avx512vpopcntdq")]
3582#[allow(unsafe_op_in_unsafe_fn)]
3583unsafe fn hamming_distance_avx512(a: &[u8], b: &[u8]) -> u32 {
3584 use std::arch::x86_64::*;
3585
3586 let len = a.len();
3587 let chunks64 = len / 64;
3588 let mut acc = _mm512_setzero_si512();
3589
3590 for c in 0..chunks64 {
3591 let off = c * 64;
3592 let va = _mm512_loadu_si512(a.as_ptr().add(off) as *const __m512i);
3593 let vb = _mm512_loadu_si512(b.as_ptr().add(off) as *const __m512i);
3594 acc = _mm512_add_epi64(acc, _mm512_popcnt_epi64(_mm512_xor_si512(va, vb)));
3595 }
3596
3597 let base = chunks64 * 64;
3598 _mm512_reduce_add_epi64(acc) as u32 + hamming_distance_scalar(&a[base..], &b[base..])
3599}
3600
3601#[cfg(target_arch = "x86_64")]
3603#[target_feature(enable = "avx512f,avx512vpopcntdq")]
3604#[allow(unsafe_op_in_unsafe_fn)]
3605unsafe fn hamming_distance_x4_avx512(query: &[u8], rows: [&[u8]; 4]) -> [u32; 4] {
3606 use std::arch::x86_64::*;
3607
3608 let len = query.len();
3609 let chunks64 = len / 64;
3610 let mut acc = [_mm512_setzero_si512(); 4];
3611
3612 for c in 0..chunks64 {
3613 let off = c * 64;
3614 let vq = _mm512_loadu_si512(query.as_ptr().add(off) as *const __m512i);
3615 for r in 0..4 {
3616 let vr = _mm512_loadu_si512(rows[r].as_ptr().add(off) as *const __m512i);
3617 acc[r] = _mm512_add_epi64(acc[r], _mm512_popcnt_epi64(_mm512_xor_si512(vq, vr)));
3618 }
3619 }
3620
3621 let base = chunks64 * 64;
3622 let tail = &query[base..];
3623 [
3624 _mm512_reduce_add_epi64(acc[0]) as u32 + hamming_distance_scalar(tail, &rows[0][base..]),
3625 _mm512_reduce_add_epi64(acc[1]) as u32 + hamming_distance_scalar(tail, &rows[1][base..]),
3626 _mm512_reduce_add_epi64(acc[2]) as u32 + hamming_distance_scalar(tail, &rows[2][base..]),
3627 _mm512_reduce_add_epi64(acc[3]) as u32 + hamming_distance_scalar(tail, &rows[3][base..]),
3628 ]
3629}
3630
3631#[inline]
3633fn hamming_distance_x4_scalar(query: &[u8], rows: [&[u8]; 4]) -> [u32; 4] {
3634 let len = query.len();
3635 let chunks = len / 8;
3636 let mut total = [0u32; 4];
3637
3638 for i in 0..chunks {
3639 let off = i * 8;
3640 let vq = unsafe { std::ptr::read_unaligned(query.as_ptr().add(off) as *const u64) };
3641 for r in 0..4 {
3642 let vr = unsafe { std::ptr::read_unaligned(rows[r].as_ptr().add(off) as *const u64) };
3643 total[r] += (vq ^ vr).count_ones();
3644 }
3645 }
3646
3647 let base = chunks * 8;
3648 for k in base..len {
3649 let q = query[k];
3650 for r in 0..4 {
3651 total[r] += (q ^ rows[r][k]).count_ones();
3652 }
3653 }
3654
3655 total
3656}
3657
3658const HAMMING_ROWS_PER_KERNEL: usize = 4;
3662
3663#[derive(Clone, Copy, Debug, PartialEq, Eq)]
3670pub enum HammingKernel {
3671 #[cfg(target_arch = "x86_64")]
3672 Avx512,
3673 #[cfg(target_arch = "x86_64")]
3674 Avx2,
3675 #[cfg(target_arch = "aarch64")]
3676 Neon,
3677 Scalar,
3678}
3679
3680impl HammingKernel {
3681 #[inline]
3683 pub fn resolve() -> Self {
3684 #[cfg(target_arch = "x86_64")]
3685 {
3686 if is_x86_feature_detected!("avx512f") && is_x86_feature_detected!("avx512vpopcntdq") {
3687 return Self::Avx512;
3688 }
3689 if avx2::is_available() {
3690 return Self::Avx2;
3691 }
3692 Self::Scalar
3693 }
3694
3695 #[cfg(target_arch = "aarch64")]
3696 {
3697 Self::Neon
3698 }
3699
3700 #[cfg(not(any(target_arch = "aarch64", target_arch = "x86_64")))]
3701 {
3702 Self::Scalar
3703 }
3704 }
3705
3706 #[inline]
3708 pub fn distance(self, a: &[u8], b: &[u8]) -> u32 {
3709 debug_assert_eq!(a.len(), b.len(), "Hamming vector byte length mismatch");
3710 match self {
3711 #[cfg(target_arch = "x86_64")]
3712 Self::Avx512 => unsafe { hamming_distance_avx512(a, b) },
3713 #[cfg(target_arch = "x86_64")]
3714 Self::Avx2 => unsafe { avx2::hamming_distance(a, b) },
3715 #[cfg(target_arch = "aarch64")]
3716 Self::Neon => unsafe { neon::hamming_distance(a, b) },
3717 Self::Scalar => hamming_distance_scalar(a, b),
3718 }
3719 }
3720
3721 pub fn distances(self, query: &[u8], db: &[u8], byte_len: usize, out: &mut [u32]) {
3723 self.score_rows(query, db, byte_len, out, |index| index);
3724 }
3725
3726 pub fn gather_distances(
3731 self,
3732 query: &[u8],
3733 db: &[u8],
3734 byte_len: usize,
3735 ids: &[u32],
3736 out: &mut [u32],
3737 ) {
3738 assert_eq!(
3739 ids.len(),
3740 out.len(),
3741 "Hamming gather needs one output slot per row id"
3742 );
3743 self.score_rows(query, db, byte_len, out, |index| ids[index] as usize);
3744 }
3745
3746 #[inline]
3747 fn score_rows(
3748 self,
3749 query: &[u8],
3750 db: &[u8],
3751 byte_len: usize,
3752 out: &mut [u32],
3753 index_of: impl Fn(usize) -> usize,
3754 ) {
3755 assert_eq!(query.len(), byte_len, "Hamming query byte length mismatch");
3756 if byte_len == 0 || out.is_empty() {
3757 return;
3758 }
3759 let row = |index: usize| -> &[u8] {
3760 let start = index * byte_len;
3761 &db[start..start + byte_len]
3762 };
3763 macro_rules! score_with {
3764 ($one:expr, $four:expr) => {{
3765 let mut i = 0;
3766 while i + HAMMING_ROWS_PER_KERNEL <= out.len() {
3767 let quad = [
3768 row(index_of(i)),
3769 row(index_of(i + 1)),
3770 row(index_of(i + 2)),
3771 row(index_of(i + 3)),
3772 ];
3773 out[i..i + HAMMING_ROWS_PER_KERNEL].copy_from_slice(&$four(query, quad));
3774 i += HAMMING_ROWS_PER_KERNEL;
3775 }
3776 while i < out.len() {
3777 out[i] = $one(query, row(index_of(i)));
3778 i += 1;
3779 }
3780 }};
3781 }
3782 match self {
3783 #[cfg(target_arch = "x86_64")]
3784 Self::Avx512 => score_with!(
3785 |query, row| unsafe { hamming_distance_avx512(query, row) },
3786 |query, rows| unsafe { hamming_distance_x4_avx512(query, rows) }
3787 ),
3788 #[cfg(target_arch = "x86_64")]
3789 Self::Avx2 => score_with!(
3790 |query, row| unsafe { avx2::hamming_distance(query, row) },
3791 |query, rows| unsafe { avx2::hamming_distance_x4(query, rows) }
3792 ),
3793 #[cfg(target_arch = "aarch64")]
3794 Self::Neon => score_with!(
3795 |query, row| unsafe { neon::hamming_distance(query, row) },
3796 |query, rows| unsafe { neon::hamming_distance_x4(query, rows) }
3797 ),
3798 Self::Scalar => score_with!(hamming_distance_scalar, hamming_distance_x4_scalar),
3799 }
3800 }
3801}
3802
3803#[inline]
3810pub fn hamming_distance(a: &[u8], b: &[u8]) -> u32 {
3811 assert_eq!(a.len(), b.len(), "Hamming vector byte length mismatch");
3812 HammingKernel::resolve().distance(a, b)
3813}
3814
3815#[inline]
3818#[allow(dead_code)]
3819fn hamming_distance_scalar(a: &[u8], b: &[u8]) -> u32 {
3820 let len = a.len();
3821 let chunks = len / 8;
3822 let remainder = len % 8;
3823 let mut total = 0u32;
3824
3825 for i in 0..chunks {
3826 let off = i * 8;
3827 let va = unsafe { std::ptr::read_unaligned(a.as_ptr().add(off) as *const u64) };
3828 let vb = unsafe { std::ptr::read_unaligned(b.as_ptr().add(off) as *const u64) };
3829 total += (va ^ vb).count_ones();
3830 }
3831
3832 let base = chunks * 8;
3833 for i in 0..remainder {
3834 total += (a[base + i] ^ b[base + i]).count_ones();
3835 }
3836
3837 total
3838}
3839
3840pub fn batch_hamming_scores(
3846 query: &[u8],
3847 db: &[u8],
3848 byte_len: usize,
3849 dim_bits: usize,
3850 scores: &mut [f32],
3851) {
3852 let n = scores.len();
3853 let required = n
3854 .checked_mul(byte_len)
3855 .expect("Hamming batch byte length overflow");
3856 assert_eq!(query.len(), byte_len, "Hamming query byte length mismatch");
3857 assert!(
3858 db.len() >= required,
3859 "Hamming batch is truncated: need {required} bytes, got {}",
3860 db.len()
3861 );
3862
3863 if byte_len == 0 || n == 0 || dim_bits == 0 {
3864 return;
3865 }
3866
3867 scores_from_hamming(
3868 HammingKernel::resolve(),
3869 query,
3870 db,
3871 byte_len,
3872 dim_bits,
3873 scores,
3874 );
3875}
3876
3877pub fn scores_from_hamming(
3882 kernel: HammingKernel,
3883 query: &[u8],
3884 db: &[u8],
3885 byte_len: usize,
3886 dim_bits: usize,
3887 scores: &mut [f32],
3888) {
3889 if byte_len == 0 || scores.is_empty() || dim_bits == 0 {
3890 return;
3891 }
3892 let inv_dim = 1.0 / dim_bits as f32;
3893 let mut distances = [0u32; HAMMING_DISTANCE_BLOCK];
3896 for (block_index, block) in scores.chunks_mut(HAMMING_DISTANCE_BLOCK).enumerate() {
3897 let rows = &mut distances[..block.len()];
3898 kernel.distances(
3899 query,
3900 &db[block_index * HAMMING_DISTANCE_BLOCK * byte_len..],
3901 byte_len,
3902 rows,
3903 );
3904 for (score, &distance) in block.iter_mut().zip(rows.iter()) {
3905 *score = 1.0 - distance as f32 * inv_dim;
3906 }
3907 }
3908}
3909
3910const HAMMING_DISTANCE_BLOCK: usize = 64;
3912
3913pub fn batch_hamming_distances(query: &[u8], db: &[u8], byte_len: usize, out: &mut [u32]) {
3918 HammingKernel::resolve().distances(query, db, byte_len, out);
3919}
3920
3921#[cfg(test)]
3922mod tests {
3923 use super::*;
3924
3925 #[test]
3926 fn vector_simd_boundaries_reject_dimension_mismatches() {
3927 let vectors = vec![1.0f32; 6];
3928 let raw_f16 = vec![0u8; 12];
3929 let raw_u8 = vec![0u8; 6];
3930 let mut scores = vec![0.0f32; 2];
3931
3932 for invalid_query in [vec![1.0, 2.0], vec![1.0, 2.0, 3.0, 4.0]] {
3933 assert!(
3934 std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
3935 batch_cosine_scores(&invalid_query, &vectors, 3, &mut scores)
3936 }))
3937 .is_err()
3938 );
3939 assert!(
3940 std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
3941 batch_dot_scores_f16(&invalid_query, &raw_f16, 3, &mut scores)
3942 }))
3943 .is_err()
3944 );
3945 assert!(
3946 std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
3947 batch_cosine_scores_u8(&invalid_query, &raw_u8, 3, &mut scores)
3948 }))
3949 .is_err()
3950 );
3951 }
3952 }
3953
3954 #[test]
3955 fn vector_simd_boundaries_reject_truncated_storage() {
3956 let query = [1.0f32, 2.0, 3.0];
3957 let mut scores = [0.0f32; 2];
3958
3959 assert!(
3960 std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
3961 batch_dot_scores(&query, &[0.0; 5], 3, &mut scores)
3962 }))
3963 .is_err()
3964 );
3965 assert!(
3966 std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
3967 batch_cosine_scores_f16(&query, &[0u8; 11], 3, &mut scores)
3968 }))
3969 .is_err()
3970 );
3971 assert!(
3972 std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
3973 dot_product_f32(&query, &query, 4)
3974 }))
3975 .is_err()
3976 );
3977 }
3978
3979 #[test]
3980 fn test_unpack_8bit() {
3981 let input: Vec<u8> = (0..128).collect();
3982 let mut output = vec![0u32; 128];
3983 unpack_8bit(&input, &mut output, 128);
3984
3985 for (i, &v) in output.iter().enumerate() {
3986 assert_eq!(v, i as u32);
3987 }
3988 }
3989
3990 #[test]
3991 fn test_unpack_16bit() {
3992 let mut input = vec![0u8; 256];
3993 for i in 0..128 {
3994 let val = (i * 100) as u16;
3995 input[i * 2] = val as u8;
3996 input[i * 2 + 1] = (val >> 8) as u8;
3997 }
3998
3999 let mut output = vec![0u32; 128];
4000 unpack_16bit(&input, &mut output, 128);
4001
4002 for (i, &v) in output.iter().enumerate() {
4003 assert_eq!(v, (i * 100) as u32);
4004 }
4005 }
4006
4007 #[test]
4008 fn test_unpack_32bit() {
4009 let mut input = vec![0u8; 512];
4010 for i in 0..128 {
4011 let val = (i * 1000) as u32;
4012 let bytes = val.to_le_bytes();
4013 input[i * 4..i * 4 + 4].copy_from_slice(&bytes);
4014 }
4015
4016 let mut output = vec![0u32; 128];
4017 unpack_32bit(&input, &mut output, 128);
4018
4019 for (i, &v) in output.iter().enumerate() {
4020 assert_eq!(v, (i * 1000) as u32);
4021 }
4022 }
4023
4024 #[test]
4025 fn test_delta_decode() {
4026 let deltas = vec![4u32, 4, 9, 19];
4030 let mut output = vec![0u32; 5];
4031
4032 delta_decode(&mut output, &deltas, 10, 5);
4033
4034 assert_eq!(output, vec![10, 15, 20, 30, 50]);
4035 }
4036
4037 #[test]
4038 fn test_add_one() {
4039 let mut values = vec![0u32, 1, 2, 3, 4, 5, 6, 7];
4040 add_one(&mut values, 8);
4041
4042 assert_eq!(values, vec![1, 2, 3, 4, 5, 6, 7, 8]);
4043 }
4044
4045 #[test]
4046 fn test_bits_needed() {
4047 assert_eq!(bits_needed(0), 0);
4048 assert_eq!(bits_needed(1), 1);
4049 assert_eq!(bits_needed(2), 2);
4050 assert_eq!(bits_needed(3), 2);
4051 assert_eq!(bits_needed(4), 3);
4052 assert_eq!(bits_needed(255), 8);
4053 assert_eq!(bits_needed(256), 9);
4054 assert_eq!(bits_needed(u32::MAX), 32);
4055 }
4056
4057 #[test]
4058 fn test_unpack_8bit_delta_decode() {
4059 let input: Vec<u8> = vec![4, 4, 9, 19];
4063 let mut output = vec![0u32; 5];
4064
4065 unpack_8bit_delta_decode(&input, &mut output, 10, 5);
4066
4067 assert_eq!(output, vec![10, 15, 20, 30, 50]);
4068 }
4069
4070 #[test]
4071 fn test_unpack_16bit_delta_decode() {
4072 let mut input = vec![0u8; 8];
4076 for (i, &delta) in [499u16, 499, 999, 1999].iter().enumerate() {
4077 input[i * 2] = delta as u8;
4078 input[i * 2 + 1] = (delta >> 8) as u8;
4079 }
4080 let mut output = vec![0u32; 5];
4081
4082 unpack_16bit_delta_decode(&input, &mut output, 100, 5);
4083
4084 assert_eq!(output, vec![100, 600, 1100, 2100, 4100]);
4085 }
4086
4087 #[test]
4088 fn test_fused_vs_separate_8bit() {
4089 let input: Vec<u8> = (0..127).collect();
4091 let first_value = 1000u32;
4092 let count = 128;
4093
4094 let mut unpacked = vec![0u32; 128];
4096 unpack_8bit(&input, &mut unpacked, 127);
4097 let mut separate_output = vec![0u32; 128];
4098 delta_decode(&mut separate_output, &unpacked, first_value, count);
4099
4100 let mut fused_output = vec![0u32; 128];
4102 unpack_8bit_delta_decode(&input, &mut fused_output, first_value, count);
4103
4104 assert_eq!(separate_output, fused_output);
4105 }
4106
4107 #[test]
4108 fn test_round_bit_width() {
4109 assert_eq!(round_bit_width(0), 0);
4110 assert_eq!(round_bit_width(1), 8);
4111 assert_eq!(round_bit_width(5), 8);
4112 assert_eq!(round_bit_width(8), 8);
4113 assert_eq!(round_bit_width(9), 16);
4114 assert_eq!(round_bit_width(12), 16);
4115 assert_eq!(round_bit_width(16), 16);
4116 assert_eq!(round_bit_width(17), 32);
4117 assert_eq!(round_bit_width(24), 32);
4118 assert_eq!(round_bit_width(32), 32);
4119 }
4120
4121 #[test]
4122 fn test_rounded_bitwidth_from_exact() {
4123 assert_eq!(RoundedBitWidth::from_exact(0), RoundedBitWidth::Zero);
4124 assert_eq!(RoundedBitWidth::from_exact(1), RoundedBitWidth::Bits8);
4125 assert_eq!(RoundedBitWidth::from_exact(8), RoundedBitWidth::Bits8);
4126 assert_eq!(RoundedBitWidth::from_exact(9), RoundedBitWidth::Bits16);
4127 assert_eq!(RoundedBitWidth::from_exact(16), RoundedBitWidth::Bits16);
4128 assert_eq!(RoundedBitWidth::from_exact(17), RoundedBitWidth::Bits32);
4129 assert_eq!(RoundedBitWidth::from_exact(32), RoundedBitWidth::Bits32);
4130 }
4131
4132 #[test]
4133 fn test_pack_unpack_rounded_8bit() {
4134 let values: Vec<u32> = (0..128).map(|i| i % 256).collect();
4135 let mut packed = vec![0u8; 128];
4136
4137 let bytes_written = pack_rounded(&values, RoundedBitWidth::Bits8, &mut packed);
4138 assert_eq!(bytes_written, 128);
4139
4140 let mut unpacked = vec![0u32; 128];
4141 unpack_rounded(&packed, RoundedBitWidth::Bits8, &mut unpacked, 128);
4142
4143 assert_eq!(values, unpacked);
4144 }
4145
4146 #[test]
4147 fn test_pack_unpack_rounded_16bit() {
4148 let values: Vec<u32> = (0..128).map(|i| i * 100).collect();
4149 let mut packed = vec![0u8; 256];
4150
4151 let bytes_written = pack_rounded(&values, RoundedBitWidth::Bits16, &mut packed);
4152 assert_eq!(bytes_written, 256);
4153
4154 let mut unpacked = vec![0u32; 128];
4155 unpack_rounded(&packed, RoundedBitWidth::Bits16, &mut unpacked, 128);
4156
4157 assert_eq!(values, unpacked);
4158 }
4159
4160 #[test]
4161 fn test_pack_unpack_rounded_32bit() {
4162 let values: Vec<u32> = (0..128).map(|i| i * 100000).collect();
4163 let mut packed = vec![0u8; 512];
4164
4165 let bytes_written = pack_rounded(&values, RoundedBitWidth::Bits32, &mut packed);
4166 assert_eq!(bytes_written, 512);
4167
4168 let mut unpacked = vec![0u32; 128];
4169 unpack_rounded(&packed, RoundedBitWidth::Bits32, &mut unpacked, 128);
4170
4171 assert_eq!(values, unpacked);
4172 }
4173
4174 #[test]
4175 fn test_unpack_rounded_delta_decode() {
4176 let input: Vec<u8> = vec![4, 4, 9, 19];
4181 let mut output = vec![0u32; 5];
4182
4183 unpack_rounded_delta_decode(&input, RoundedBitWidth::Bits8, &mut output, 10, 5);
4184
4185 assert_eq!(output, vec![10, 15, 20, 30, 50]);
4186 }
4187
4188 #[test]
4189 fn test_unpack_rounded_delta_decode_zero() {
4190 let input: Vec<u8> = vec![];
4192 let mut output = vec![0u32; 5];
4193
4194 unpack_rounded_delta_decode(&input, RoundedBitWidth::Zero, &mut output, 100, 5);
4195
4196 assert_eq!(output, vec![100, 101, 102, 103, 104]);
4197 }
4198
4199 #[test]
4204 fn test_dequantize_uint8() {
4205 let input: Vec<u8> = vec![0, 128, 255, 64, 192];
4206 let mut output = vec![0.0f32; 5];
4207 let scale = 0.1;
4208 let min_val = 1.0;
4209
4210 dequantize_uint8(&input, &mut output, scale, min_val, 5);
4211
4212 assert!((output[0] - 1.0).abs() < 1e-6); assert!((output[1] - 13.8).abs() < 1e-6); assert!((output[2] - 26.5).abs() < 1e-6); assert!((output[3] - 7.4).abs() < 1e-6); assert!((output[4] - 20.2).abs() < 1e-6); }
4219
4220 #[test]
4221 fn test_dequantize_uint8_large() {
4222 let input: Vec<u8> = (0..128).collect();
4224 let mut output = vec![0.0f32; 128];
4225 let scale = 2.0;
4226 let min_val = -10.0;
4227
4228 dequantize_uint8(&input, &mut output, scale, min_val, 128);
4229
4230 for (i, &out) in output.iter().enumerate().take(128) {
4231 let expected = i as f32 * scale + min_val;
4232 assert!(
4233 (out - expected).abs() < 1e-5,
4234 "Mismatch at {}: expected {}, got {}",
4235 i,
4236 expected,
4237 out
4238 );
4239 }
4240 }
4241
4242 #[test]
4243 fn test_dot_product_f32() {
4244 let a = vec![1.0f32, 2.0, 3.0, 4.0, 5.0];
4245 let b = vec![2.0f32, 3.0, 4.0, 5.0, 6.0];
4246
4247 let result = dot_product_f32(&a, &b, 5);
4248
4249 assert!((result - 70.0).abs() < 1e-5);
4251 }
4252
4253 #[test]
4254 fn test_dot_product_f32_large() {
4255 let a: Vec<f32> = (0..128).map(|i| i as f32).collect();
4257 let b: Vec<f32> = (0..128).map(|i| (i + 1) as f32).collect();
4258
4259 let result = dot_product_f32(&a, &b, 128);
4260
4261 let expected: f32 = (0..128).map(|i| (i as f32) * ((i + 1) as f32)).sum();
4263 assert!(
4264 (result - expected).abs() < 1e-3,
4265 "Expected {}, got {}",
4266 expected,
4267 result
4268 );
4269 }
4270
4271 #[test]
4272 fn test_fused_dot_norm() {
4273 let a = vec![1.0f32, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0];
4274 let b = vec![2.0f32, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0];
4275 let (dot, norm_b) = fused_dot_norm(&a, &b, a.len());
4276
4277 let expected_dot: f32 = a.iter().zip(b.iter()).map(|(x, y)| x * y).sum();
4278 let expected_norm: f32 = b.iter().map(|x| x * x).sum();
4279 assert!(
4280 (dot - expected_dot).abs() < 1e-5,
4281 "dot: expected {}, got {}",
4282 expected_dot,
4283 dot
4284 );
4285 assert!(
4286 (norm_b - expected_norm).abs() < 1e-5,
4287 "norm: expected {}, got {}",
4288 expected_norm,
4289 norm_b
4290 );
4291 }
4292
4293 #[test]
4294 fn test_fused_dot_norm_large() {
4295 let a: Vec<f32> = (0..768).map(|i| (i as f32) * 0.01).collect();
4296 let b: Vec<f32> = (0..768).map(|i| (i as f32) * 0.02 + 0.5).collect();
4297 let (dot, norm_b) = fused_dot_norm(&a, &b, a.len());
4298
4299 let expected_dot: f32 = a.iter().zip(b.iter()).map(|(x, y)| x * y).sum();
4300 let expected_norm: f32 = b.iter().map(|x| x * x).sum();
4301 assert!(
4302 (dot - expected_dot).abs() < 1.0,
4303 "dot: expected {}, got {}",
4304 expected_dot,
4305 dot
4306 );
4307 assert!(
4308 (norm_b - expected_norm).abs() < 1.0,
4309 "norm: expected {}, got {}",
4310 expected_norm,
4311 norm_b
4312 );
4313 }
4314
4315 #[test]
4316 fn test_batch_cosine_scores() {
4317 let query = vec![1.0f32, 0.0, 0.0];
4319 let vectors = vec![
4320 1.0, 0.0, 0.0, 0.0, 1.0, 0.0, -1.0, 0.0, 0.0, 0.5, 0.5, 0.0, ];
4325 let mut scores = vec![0f32; 4];
4326 batch_cosine_scores(&query, &vectors, 3, &mut scores);
4327
4328 assert!((scores[0] - 1.0).abs() < 1e-5, "identical: {}", scores[0]);
4329 assert!(scores[1].abs() < 1e-5, "orthogonal: {}", scores[1]);
4330 assert!((scores[2] - (-1.0)).abs() < 1e-5, "opposite: {}", scores[2]);
4331 let expected_45 = 0.5f32 / (0.5f32.powi(2) + 0.5f32.powi(2)).sqrt();
4332 assert!(
4333 (scores[3] - expected_45).abs() < 1e-5,
4334 "45deg: expected {}, got {}",
4335 expected_45,
4336 scores[3]
4337 );
4338 }
4339
4340 #[test]
4341 fn test_batch_cosine_scores_matches_individual() {
4342 let query: Vec<f32> = (0..128).map(|i| (i as f32) * 0.1).collect();
4343 let n = 50;
4344 let dim = 128;
4345 let vectors: Vec<f32> = (0..n * dim).map(|i| ((i * 7 + 3) as f32) * 0.01).collect();
4346
4347 let mut batch_scores = vec![0f32; n];
4348 batch_cosine_scores(&query, &vectors, dim, &mut batch_scores);
4349
4350 for i in 0..n {
4351 let vec_i = &vectors[i * dim..(i + 1) * dim];
4352 let individual = cosine_similarity(&query, vec_i);
4353 assert!(
4354 (batch_scores[i] - individual).abs() < 1e-5,
4355 "vec {}: batch={}, individual={}",
4356 i,
4357 batch_scores[i],
4358 individual
4359 );
4360 }
4361 }
4362
4363 #[test]
4364 fn test_batch_cosine_scores_empty() {
4365 let query = vec![1.0f32, 2.0, 3.0];
4366 let vectors: Vec<f32> = vec![];
4367 let mut scores: Vec<f32> = vec![];
4368 batch_cosine_scores(&query, &vectors, 3, &mut scores);
4369 assert!(scores.is_empty());
4370 }
4371
4372 #[test]
4373 fn test_batch_cosine_scores_zero_query() {
4374 let query = vec![0.0f32, 0.0, 0.0];
4375 let vectors = vec![1.0f32, 2.0, 3.0, 4.0, 5.0, 6.0];
4376 let mut scores = vec![0f32; 2];
4377 batch_cosine_scores(&query, &vectors, 3, &mut scores);
4378 assert_eq!(scores[0], 0.0);
4379 assert_eq!(scores[1], 0.0);
4380 }
4381
4382 #[test]
4387 fn test_f16_roundtrip_normal() {
4388 for &v in &[0.0f32, 1.0, -1.0, 0.5, -0.5, 0.333, 65504.0] {
4389 let h = f32_to_f16(v);
4390 let back = f16_to_f32(h);
4391 let err = (back - v).abs() / v.abs().max(1e-6);
4392 assert!(
4393 err < 0.002,
4394 "f16 roundtrip {v} → {h:#06x} → {back}, rel err {err}"
4395 );
4396 }
4397 }
4398
4399 #[test]
4400 fn test_f16_special() {
4401 assert_eq!(f16_to_f32(f32_to_f16(0.0)), 0.0);
4403 assert_eq!(f32_to_f16(-0.0), 0x8000);
4405 assert!(f16_to_f32(f32_to_f16(f32::INFINITY)).is_infinite());
4407 assert!(f16_to_f32(f32_to_f16(f32::NAN)).is_nan());
4409 }
4410
4411 #[test]
4412 fn test_f16_embedding_range() {
4413 let values: Vec<f32> = (-100..=100).map(|i| i as f32 / 100.0).collect();
4415 for &v in &values {
4416 let back = f16_to_f32(f32_to_f16(v));
4417 assert!((back - v).abs() < 0.001, "f16 error for {v}: got {back}");
4418 }
4419 }
4420
4421 #[test]
4426 fn test_u8_roundtrip() {
4427 assert_eq!(f32_to_u8_saturating(-1.0), 0);
4429 assert_eq!(f32_to_u8_saturating(1.0), 255);
4430 assert_eq!(f32_to_u8_saturating(0.0), 127); assert_eq!(f32_to_u8_saturating(-2.0), 0);
4434 assert_eq!(f32_to_u8_saturating(2.0), 255);
4435 }
4436
4437 #[test]
4438 fn test_u8_dequantize() {
4439 assert!((u8_to_f32(0) - (-1.0)).abs() < 0.01);
4440 assert!((u8_to_f32(255) - 1.0).abs() < 0.01);
4441 assert!((u8_to_f32(127) - 0.0).abs() < 0.01);
4442 }
4443
4444 #[test]
4449 fn test_batch_cosine_scores_f16() {
4450 let query = vec![0.6f32, 0.8, 0.0, 0.0];
4451 let dim = 4;
4452 let vecs_f32 = vec![
4453 0.6f32, 0.8, 0.0, 0.0, 0.0, 0.0, 0.6, 0.8, ];
4456
4457 let mut f16_buf = vec![0u16; 8];
4459 batch_f32_to_f16(&vecs_f32, &mut f16_buf);
4460 let raw: &[u8] =
4461 unsafe { std::slice::from_raw_parts(f16_buf.as_ptr() as *const u8, f16_buf.len() * 2) };
4462
4463 let mut scores = vec![0f32; 2];
4464 batch_cosine_scores_f16(&query, raw, dim, &mut scores);
4465
4466 assert!(
4467 (scores[0] - 1.0).abs() < 0.01,
4468 "identical vectors: {}",
4469 scores[0]
4470 );
4471 assert!(scores[1].abs() < 0.01, "orthogonal vectors: {}", scores[1]);
4472 }
4473
4474 #[test]
4475 fn test_batch_cosine_scores_u8() {
4476 let query = vec![0.6f32, 0.8, 0.0, 0.0];
4477 let dim = 4;
4478 let vecs_f32 = vec![
4479 0.6f32, 0.8, 0.0, 0.0, -0.6, -0.8, 0.0, 0.0, ];
4482
4483 let mut u8_buf = vec![0u8; 8];
4485 batch_f32_to_u8(&vecs_f32, &mut u8_buf);
4486
4487 let mut scores = vec![0f32; 2];
4488 batch_cosine_scores_u8(&query, &u8_buf, dim, &mut scores);
4489
4490 assert!(scores[0] > 0.95, "similar vectors: {}", scores[0]);
4491 assert!(scores[1] < -0.95, "opposite vectors: {}", scores[1]);
4492 }
4493
4494 #[test]
4495 fn test_batch_cosine_scores_f16_large_dim() {
4496 let dim = 768;
4498 let query: Vec<f32> = (0..dim).map(|i| (i as f32 / dim as f32) - 0.5).collect();
4499 let vec2: Vec<f32> = query.iter().map(|x| x * 0.9 + 0.01).collect();
4500
4501 let mut all_vecs = query.clone();
4502 all_vecs.extend_from_slice(&vec2);
4503
4504 let mut f16_buf = vec![0u16; all_vecs.len()];
4505 batch_f32_to_f16(&all_vecs, &mut f16_buf);
4506 let raw: &[u8] =
4507 unsafe { std::slice::from_raw_parts(f16_buf.as_ptr() as *const u8, f16_buf.len() * 2) };
4508
4509 let mut scores = vec![0f32; 2];
4510 batch_cosine_scores_f16(&query, raw, dim, &mut scores);
4511
4512 assert!((scores[0] - 1.0).abs() < 0.01, "self-sim: {}", scores[0]);
4514 assert!(scores[1] > 0.99, "scaled-sim: {}", scores[1]);
4516 }
4517
4518 #[test]
4523 fn test_hamming_distance_identical() {
4524 let a = vec![0xAA; 64];
4525 assert_eq!(hamming_distance(&a, &a), 0);
4526 }
4527
4528 #[test]
4529 fn test_hamming_distance_opposite() {
4530 let a = vec![0xFF; 32];
4531 let b = vec![0x00; 32];
4532 assert_eq!(hamming_distance(&a, &b), 256);
4533 }
4534
4535 #[test]
4536 fn test_hamming_distance_known() {
4537 let a = vec![0xAA];
4539 let b = vec![0x55];
4540 assert_eq!(hamming_distance(&a, &b), 8);
4541
4542 let a = vec![0xFF, 0x00];
4544 let b = vec![0x00, 0x00];
4545 assert_eq!(hamming_distance(&a, &b), 8);
4546 }
4547
4548 #[test]
4549 fn test_hamming_distance_single_bit() {
4550 let a = vec![0x00; 16];
4551 let mut b = vec![0x00; 16];
4552 b[7] = 0x01; assert_eq!(hamming_distance(&a, &b), 1);
4554 }
4555
4556 #[test]
4557 fn test_hamming_distance_empty() {
4558 let a: Vec<u8> = vec![];
4559 assert_eq!(hamming_distance(&a, &a), 0);
4560 }
4561
4562 #[test]
4563 fn test_hamming_distance_remainder_path() {
4564 let a = vec![0xFF; 17];
4566 let b = vec![0x00; 17];
4567 assert_eq!(hamming_distance(&a, &b), 136); let a = vec![0xFF; 33];
4571 let b = vec![0x00; 33];
4572 assert_eq!(hamming_distance(&a, &b), 264); }
4574
4575 #[test]
4576 fn test_hamming_distance_large() {
4577 let a = vec![0xFF; 4096];
4579 let b = vec![0x00; 4096];
4580 assert_eq!(hamming_distance(&a, &b), 32768);
4581 }
4582
4583 #[test]
4584 fn test_hamming_distance_scalar_matches() {
4585 for size in [1, 7, 8, 15, 16, 31, 32, 63, 64, 100, 128, 255, 256] {
4587 let a: Vec<u8> = (0..size).map(|i| (i * 37 + 13) as u8).collect();
4588 let b: Vec<u8> = (0..size).map(|i| (i * 53 + 7) as u8).collect();
4589 let expected = hamming_distance_scalar(&a, &b);
4590 let got = hamming_distance(&a, &b);
4591 assert_eq!(got, expected, "mismatch at size {size}");
4592 }
4593 }
4594
4595 #[test]
4600 fn test_batch_hamming_scores_identical() {
4601 let query = vec![0xAA; 16];
4602 let db = vec![0xAA; 16]; let mut scores = vec![0f32; 1];
4604 batch_hamming_scores(&query, &db, 16, 128, &mut scores);
4605 assert!((scores[0] - 1.0).abs() < 1e-6, "identical: {}", scores[0]);
4606 }
4607
4608 #[test]
4609 fn test_batch_hamming_scores_opposite() {
4610 let query = vec![0xFF; 16];
4611 let db = vec![0x00; 16];
4612 let mut scores = vec![0f32; 1];
4613 batch_hamming_scores(&query, &db, 16, 128, &mut scores);
4614 assert!((scores[0] - 0.0).abs() < 1e-6, "opposite: {}", scores[0]);
4615 }
4616
4617 #[test]
4618 fn test_batch_hamming_scores_multiple() {
4619 let byte_len = 8;
4620 let dim_bits = 64;
4621 let query = vec![0xFF; byte_len];
4622 let mut db = Vec::new();
4623 db.extend_from_slice(&vec![0xFF; byte_len]); db.extend_from_slice(&vec![0x00; byte_len]); db.extend_from_slice(&vec![0x0F; byte_len]); let mut scores = vec![0f32; 3];
4628 batch_hamming_scores(&query, &db, byte_len, dim_bits, &mut scores);
4629
4630 assert!((scores[0] - 1.0).abs() < 1e-6, "identical: {}", scores[0]);
4631 assert!((scores[1] - 0.0).abs() < 1e-6, "opposite: {}", scores[1]);
4632 assert!((scores[2] - 0.5).abs() < 1e-6, "half: {}", scores[2]);
4633 }
4634
4635 #[test]
4636 fn test_batch_hamming_scores_empty() {
4637 let query = vec![0xFF; 8];
4638 let db: Vec<u8> = vec![];
4639 let mut scores: Vec<f32> = vec![];
4640 batch_hamming_scores(&query, &db, 8, 64, &mut scores);
4641 assert!(scores.is_empty());
4642 }
4643
4644 #[test]
4645 fn test_batch_hamming_scores_zero_byte_len() {
4646 let query: Vec<u8> = vec![];
4647 let db: Vec<u8> = vec![];
4648 let mut scores = vec![0f32; 1];
4649 batch_hamming_scores(&query, &db, 0, 0, &mut scores);
4650 assert_eq!(scores[0], 0.0);
4652 }
4653
4654 fn hamming_matrix(rows: usize, byte_len: usize) -> (Vec<u8>, Vec<u8>) {
4659 let query: Vec<u8> = (0..byte_len).map(|i| (i * 31 + 5) as u8).collect();
4660 let db: Vec<u8> = (0..rows * byte_len)
4661 .map(|i| (i * 97 + i / byte_len * 11 + 3) as u8)
4662 .collect();
4663 (query, db)
4664 }
4665
4666 #[test]
4669 fn batched_hamming_distances_match_scalar_for_every_row_count() {
4670 let kernel = HammingKernel::resolve();
4671 for byte_len in [1, 7, 8, 15, 16, 31, 32, 33, 63, 64, 65, 128, 320] {
4674 for rows in [1, 2, 3, 4, 5, 7, 8, 9, 64, 70] {
4675 let (query, db) = hamming_matrix(rows, byte_len);
4676 let mut got = vec![0u32; rows];
4677 kernel.distances(&query, &db, byte_len, &mut got);
4678 for (row, &distance) in got.iter().enumerate() {
4679 let expected =
4680 hamming_distance_scalar(&query, &db[row * byte_len..(row + 1) * byte_len]);
4681 assert_eq!(
4682 distance, expected,
4683 "row {row} of {rows} at byte_len {byte_len}"
4684 );
4685 }
4686 }
4687 }
4688 }
4689
4690 #[test]
4691 fn gathered_hamming_distances_follow_row_ids() {
4692 let kernel = HammingKernel::resolve();
4693 let byte_len = 320;
4694 let rows = 37;
4695 let (query, db) = hamming_matrix(rows, byte_len);
4696 let ids: Vec<u32> = [36, 0, 17, 17, 5, 31, 2, 9, 9, 36, 1].into_iter().collect();
4699 let mut got = vec![0u32; ids.len()];
4700 kernel.gather_distances(&query, &db, byte_len, &ids, &mut got);
4701 for (slot, &id) in ids.iter().enumerate() {
4702 let start = id as usize * byte_len;
4703 let expected = hamming_distance_scalar(&query, &db[start..start + byte_len]);
4704 assert_eq!(got[slot], expected, "slot {slot} for row {id}");
4705 }
4706 }
4707
4708 #[test]
4709 fn resolved_kernel_matches_scalar_pairwise() {
4710 let kernel = HammingKernel::resolve();
4711 for byte_len in [1, 8, 32, 64, 65, 320, 4096] {
4712 let (query, db) = hamming_matrix(1, byte_len);
4713 assert_eq!(
4714 kernel.distance(&query, &db),
4715 hamming_distance_scalar(&query, &db),
4716 "byte_len {byte_len}"
4717 );
4718 }
4719 }
4720
4721 #[test]
4722 fn scores_from_hamming_matches_batch_scores_across_blocks() {
4723 let kernel = HammingKernel::resolve();
4724 let byte_len = 320;
4725 let dim_bits = byte_len * 8;
4726 let rows = HAMMING_DISTANCE_BLOCK * 2 + 3;
4728 let (query, db) = hamming_matrix(rows, byte_len);
4729 let mut expected = vec![0f32; rows];
4730 for (row, score) in expected.iter_mut().enumerate() {
4731 let distance =
4732 hamming_distance_scalar(&query, &db[row * byte_len..(row + 1) * byte_len]);
4733 *score = 1.0 - distance as f32 / dim_bits as f32;
4734 }
4735 let mut got = vec![0f32; rows];
4736 scores_from_hamming(kernel, &query, &db, byte_len, dim_bits, &mut got);
4737 for (row, (&got, &want)) in got.iter().zip(expected.iter()).enumerate() {
4738 assert!((got - want).abs() < 1e-6, "row {row}: {got} vs {want}");
4739 }
4740 let mut public = vec![0f32; rows];
4741 batch_hamming_scores(&query, &db, byte_len, dim_bits, &mut public);
4742 assert_eq!(got, public);
4743 }
4744}
4745
4746#[inline]
4759pub fn find_first_ge_u32(slice: &[u32], target: u32) -> usize {
4760 #[cfg(target_arch = "aarch64")]
4761 {
4762 if neon::is_available() {
4763 return unsafe { find_first_ge_u32_neon(slice, target) };
4764 }
4765 }
4766
4767 #[cfg(target_arch = "x86_64")]
4768 {
4769 if sse::is_available() {
4770 return unsafe { find_first_ge_u32_sse(slice, target) };
4771 }
4772 }
4773
4774 slice.partition_point(|&d| d < target)
4776}
4777
4778#[cfg(target_arch = "aarch64")]
4779#[target_feature(enable = "neon")]
4780#[allow(unsafe_op_in_unsafe_fn)]
4781unsafe fn find_first_ge_u32_neon(slice: &[u32], target: u32) -> usize {
4782 use std::arch::aarch64::*;
4783
4784 let n = slice.len();
4785 let ptr = slice.as_ptr();
4786 let target_vec = vdupq_n_u32(target);
4787 let bit_mask: uint32x4_t = core::mem::transmute([1u32, 2u32, 4u32, 8u32]);
4789
4790 let chunks = n / 16;
4791 let mut base = 0usize;
4792
4793 for _ in 0..chunks {
4795 let v0 = vld1q_u32(ptr.add(base));
4796 let v1 = vld1q_u32(ptr.add(base + 4));
4797 let v2 = vld1q_u32(ptr.add(base + 8));
4798 let v3 = vld1q_u32(ptr.add(base + 12));
4799
4800 let c0 = vcgeq_u32(v0, target_vec);
4801 let c1 = vcgeq_u32(v1, target_vec);
4802 let c2 = vcgeq_u32(v2, target_vec);
4803 let c3 = vcgeq_u32(v3, target_vec);
4804
4805 let m0 = vaddvq_u32(vandq_u32(c0, bit_mask));
4806 if m0 != 0 {
4807 return base + m0.trailing_zeros() as usize;
4808 }
4809 let m1 = vaddvq_u32(vandq_u32(c1, bit_mask));
4810 if m1 != 0 {
4811 return base + 4 + m1.trailing_zeros() as usize;
4812 }
4813 let m2 = vaddvq_u32(vandq_u32(c2, bit_mask));
4814 if m2 != 0 {
4815 return base + 8 + m2.trailing_zeros() as usize;
4816 }
4817 let m3 = vaddvq_u32(vandq_u32(c3, bit_mask));
4818 if m3 != 0 {
4819 return base + 12 + m3.trailing_zeros() as usize;
4820 }
4821 base += 16;
4822 }
4823
4824 while base + 4 <= n {
4826 let vals = vld1q_u32(ptr.add(base));
4827 let cmp = vcgeq_u32(vals, target_vec);
4828 let mask = vaddvq_u32(vandq_u32(cmp, bit_mask));
4829 if mask != 0 {
4830 return base + mask.trailing_zeros() as usize;
4831 }
4832 base += 4;
4833 }
4834
4835 while base < n {
4837 if *slice.get_unchecked(base) >= target {
4838 return base;
4839 }
4840 base += 1;
4841 }
4842 n
4843}
4844
4845#[cfg(target_arch = "x86_64")]
4846#[target_feature(enable = "sse2")]
4847#[allow(unsafe_op_in_unsafe_fn)]
4848unsafe fn find_first_ge_u32_sse(slice: &[u32], target: u32) -> usize {
4849 use std::arch::x86_64::*;
4850
4851 let n = slice.len();
4852 let ptr = slice.as_ptr();
4853
4854 let sign_flip = _mm_set1_epi32(i32::MIN);
4856 let target_xor = _mm_xor_si128(_mm_set1_epi32(target as i32), sign_flip);
4857
4858 let chunks = n / 16;
4859 let mut base = 0usize;
4860
4861 for _ in 0..chunks {
4863 let v0 = _mm_xor_si128(_mm_loadu_si128(ptr.add(base) as *const __m128i), sign_flip);
4864 let v1 = _mm_xor_si128(
4865 _mm_loadu_si128(ptr.add(base + 4) as *const __m128i),
4866 sign_flip,
4867 );
4868 let v2 = _mm_xor_si128(
4869 _mm_loadu_si128(ptr.add(base + 8) as *const __m128i),
4870 sign_flip,
4871 );
4872 let v3 = _mm_xor_si128(
4873 _mm_loadu_si128(ptr.add(base + 12) as *const __m128i),
4874 sign_flip,
4875 );
4876
4877 let ge0 = _mm_or_si128(
4879 _mm_cmpeq_epi32(v0, target_xor),
4880 _mm_cmpgt_epi32(v0, target_xor),
4881 );
4882 let m0 = _mm_movemask_ps(_mm_castsi128_ps(ge0)) as u32;
4883 if m0 != 0 {
4884 return base + m0.trailing_zeros() as usize;
4885 }
4886
4887 let ge1 = _mm_or_si128(
4888 _mm_cmpeq_epi32(v1, target_xor),
4889 _mm_cmpgt_epi32(v1, target_xor),
4890 );
4891 let m1 = _mm_movemask_ps(_mm_castsi128_ps(ge1)) as u32;
4892 if m1 != 0 {
4893 return base + 4 + m1.trailing_zeros() as usize;
4894 }
4895
4896 let ge2 = _mm_or_si128(
4897 _mm_cmpeq_epi32(v2, target_xor),
4898 _mm_cmpgt_epi32(v2, target_xor),
4899 );
4900 let m2 = _mm_movemask_ps(_mm_castsi128_ps(ge2)) as u32;
4901 if m2 != 0 {
4902 return base + 8 + m2.trailing_zeros() as usize;
4903 }
4904
4905 let ge3 = _mm_or_si128(
4906 _mm_cmpeq_epi32(v3, target_xor),
4907 _mm_cmpgt_epi32(v3, target_xor),
4908 );
4909 let m3 = _mm_movemask_ps(_mm_castsi128_ps(ge3)) as u32;
4910 if m3 != 0 {
4911 return base + 12 + m3.trailing_zeros() as usize;
4912 }
4913 base += 16;
4914 }
4915
4916 while base + 4 <= n {
4918 let vals = _mm_xor_si128(_mm_loadu_si128(ptr.add(base) as *const __m128i), sign_flip);
4919 let ge = _mm_or_si128(
4920 _mm_cmpeq_epi32(vals, target_xor),
4921 _mm_cmpgt_epi32(vals, target_xor),
4922 );
4923 let mask = _mm_movemask_ps(_mm_castsi128_ps(ge)) as u32;
4924 if mask != 0 {
4925 return base + mask.trailing_zeros() as usize;
4926 }
4927 base += 4;
4928 }
4929
4930 while base < n {
4932 if *slice.get_unchecked(base) >= target {
4933 return base;
4934 }
4935 base += 1;
4936 }
4937 n
4938}
4939
4940#[cfg(test)]
4941mod find_first_ge_tests {
4942 use super::find_first_ge_u32;
4943
4944 #[test]
4945 fn test_find_first_ge_basic() {
4946 let data: Vec<u32> = (0..128).map(|i| i * 3).collect(); assert_eq!(find_first_ge_u32(&data, 0), 0);
4948 assert_eq!(find_first_ge_u32(&data, 1), 1); assert_eq!(find_first_ge_u32(&data, 3), 1);
4950 assert_eq!(find_first_ge_u32(&data, 4), 2); assert_eq!(find_first_ge_u32(&data, 381), 127);
4952 assert_eq!(find_first_ge_u32(&data, 382), 128); }
4954
4955 #[test]
4956 fn test_find_first_ge_matches_partition_point() {
4957 let data: Vec<u32> = vec![1, 5, 10, 15, 20, 25, 30, 35, 40, 45, 50, 55, 60, 65, 70, 75];
4958 for target in 0..80 {
4959 let expected = data.partition_point(|&d| d < target);
4960 let actual = find_first_ge_u32(&data, target);
4961 assert_eq!(actual, expected, "target={}", target);
4962 }
4963 }
4964
4965 #[test]
4966 fn test_find_first_ge_small_slices() {
4967 assert_eq!(find_first_ge_u32(&[], 5), 0);
4969 assert_eq!(find_first_ge_u32(&[10], 5), 0);
4971 assert_eq!(find_first_ge_u32(&[10], 10), 0);
4972 assert_eq!(find_first_ge_u32(&[10], 11), 1);
4973 assert_eq!(find_first_ge_u32(&[2, 4, 6], 5), 2);
4975 }
4976
4977 #[test]
4978 fn test_find_first_ge_full_block() {
4979 let data: Vec<u32> = (100..228).collect();
4981 assert_eq!(find_first_ge_u32(&data, 100), 0);
4982 assert_eq!(find_first_ge_u32(&data, 150), 50);
4983 assert_eq!(find_first_ge_u32(&data, 227), 127);
4984 assert_eq!(find_first_ge_u32(&data, 228), 128);
4985 assert_eq!(find_first_ge_u32(&data, 99), 0);
4986 }
4987
4988 #[test]
4989 fn test_find_first_ge_u32_max() {
4990 let data = vec![u32::MAX - 10, u32::MAX - 5, u32::MAX - 1, u32::MAX];
4992 assert_eq!(find_first_ge_u32(&data, u32::MAX - 10), 0);
4993 assert_eq!(find_first_ge_u32(&data, u32::MAX - 7), 1);
4994 assert_eq!(find_first_ge_u32(&data, u32::MAX), 3);
4995 }
4996}
4997
4998#[cfg(test)]
5002mod algebraic_reduction_tests {
5003 use super::*;
5004
5005 fn algebraic_test_vector(dim: usize, seed: u64) -> Vec<f32> {
5006 let mut state = seed | 1;
5007 (0..dim)
5008 .map(|_| {
5009 state = state
5010 .wrapping_mul(6364136223846793005)
5011 .wrapping_add(1442695040888963407);
5012 ((state >> 33) as f32 / (1u64 << 30) as f32) - 1.0
5013 })
5014 .collect()
5015 }
5016
5017 const ALGEBRAIC_TEST_DIMS: [usize; 12] = [0, 1, 3, 4, 7, 8, 15, 16, 17, 64, 384, 768];
5020
5021 #[test]
5022 fn test_algebraic_squared_l2_matches_f64_reference() {
5023 for dim in ALGEBRAIC_TEST_DIMS {
5024 let a = algebraic_test_vector(dim, 0x51ed_0001);
5025 let b = algebraic_test_vector(dim, 0x51ed_0002);
5026 let reference: f64 = a
5027 .iter()
5028 .zip(&b)
5029 .map(|(&x, &y)| {
5030 let delta = f64::from(x) - f64::from(y);
5031 delta * delta
5032 })
5033 .sum();
5034 let actual = squared_l2_f32(&a, &b);
5035 let tolerance = (reference * 1e-5).max(1e-6);
5036 assert!(
5037 (f64::from(actual) - reference).abs() <= tolerance,
5038 "dim {dim}: squared_l2_f32 {actual} drifted from f64 reference {reference}"
5039 );
5040 }
5041 }
5042
5043 #[test]
5044 fn test_algebraic_squared_l2_uses_shorter_length() {
5045 let a = [1.0f32, 2.0, 3.0, 4.0];
5046 let b = [1.0f32, 4.0];
5047 assert_eq!(squared_l2_f32(&a, &b), 4.0);
5048 assert_eq!(squared_l2_f32(&b, &a), 4.0);
5049 }
5050
5051 #[test]
5052 fn test_algebraic_norm_squared_matches_f64_reference() {
5053 for dim in ALGEBRAIC_TEST_DIMS {
5054 let v = algebraic_test_vector(dim, 0x51ed_0003);
5055 let reference: f64 = v.iter().map(|&x| f64::from(x) * f64::from(x)).sum();
5056 let actual = norm_squared_f32(&v);
5057 let tolerance = (reference * 1e-5).max(1e-6);
5058 assert!(
5059 (f64::from(actual) - reference).abs() <= tolerance,
5060 "dim {dim}: norm_squared_f32 {actual} drifted from f64 reference {reference}"
5061 );
5062 assert!(
5063 (f64::from(norm_f32(&v)) - reference.sqrt()).abs() <= tolerance.sqrt().max(1e-5)
5064 );
5065 }
5066 }
5067
5068 #[test]
5069 fn test_algebraic_dot_scalar_matches_simd_dispatch() {
5070 for dim in ALGEBRAIC_TEST_DIMS {
5071 let a = algebraic_test_vector(dim, 0x51ed_0004);
5072 let b = algebraic_test_vector(dim, 0x51ed_0005);
5073 let dispatched = dot_product_f32(&a, &b, dim);
5074 let scalar = dot_product_f32_scalar(&a, &b);
5075 let tolerance = (dispatched.abs() * 1e-5).max(1e-5);
5076 assert!(
5077 (dispatched - scalar).abs() <= tolerance,
5078 "dim {dim}: scalar dot {scalar} disagrees with dispatched {dispatched}"
5079 );
5080
5081 let (fused_dot, fused_norm) = fused_dot_norm(&a, &b, dim);
5082 let (scalar_dot, scalar_norm) = fused_dot_norm_scalar(&a, &b);
5083 assert!((fused_dot - scalar_dot).abs() <= tolerance);
5084 assert!((fused_norm - scalar_norm).abs() <= (fused_norm.abs() * 1e-5).max(1e-5));
5085 }
5086 }
5087
5088 #[test]
5093 fn test_algebraic_reductions_propagate_non_finite() {
5094 let finite = vec![1.0f32; 8];
5095
5096 let mut with_nan = finite.clone();
5097 with_nan[5] = f32::NAN;
5098 assert!(norm_squared_f32(&with_nan).is_nan());
5099 assert!(squared_l2_f32(&with_nan, &finite).is_nan());
5100 assert!(dot_product_f32_scalar(&with_nan, &finite).is_nan());
5101
5102 let mut with_inf = finite.clone();
5103 with_inf[2] = f32::INFINITY;
5104 assert!(norm_squared_f32(&with_inf).is_infinite());
5105 assert!(squared_l2_f32(&with_inf, &finite).is_infinite());
5106 assert!(dot_product_f32_scalar(&with_inf, &finite).is_infinite());
5107 }
5108}