1#[cfg(target_arch = "aarch64")]
18#[allow(unsafe_op_in_unsafe_fn)]
19mod neon {
20 use std::arch::aarch64::*;
21
22 #[target_feature(enable = "neon")]
24 pub unsafe fn unpack_8bit(input: &[u8], output: &mut [u32], count: usize) {
25 let chunks = count / 16;
26 let remainder = count % 16;
27
28 for chunk in 0..chunks {
29 let base = chunk * 16;
30 let in_ptr = input.as_ptr().add(base);
31
32 let bytes = vld1q_u8(in_ptr);
34
35 let low8 = vget_low_u8(bytes);
37 let high8 = vget_high_u8(bytes);
38
39 let low16 = vmovl_u8(low8);
40 let high16 = vmovl_u8(high8);
41
42 let v0 = vmovl_u16(vget_low_u16(low16));
43 let v1 = vmovl_u16(vget_high_u16(low16));
44 let v2 = vmovl_u16(vget_low_u16(high16));
45 let v3 = vmovl_u16(vget_high_u16(high16));
46
47 let out_ptr = output.as_mut_ptr().add(base);
48 vst1q_u32(out_ptr, v0);
49 vst1q_u32(out_ptr.add(4), v1);
50 vst1q_u32(out_ptr.add(8), v2);
51 vst1q_u32(out_ptr.add(12), v3);
52 }
53
54 let base = chunks * 16;
56 for i in 0..remainder {
57 output[base + i] = input[base + i] as u32;
58 }
59 }
60
61 #[target_feature(enable = "neon")]
63 pub unsafe fn unpack_16bit(input: &[u8], output: &mut [u32], count: usize) {
64 let chunks = count / 8;
65 let remainder = count % 8;
66
67 for chunk in 0..chunks {
68 let base = chunk * 8;
69 let in_ptr = input.as_ptr().add(base * 2) as *const u16;
70
71 let vals = vld1q_u16(in_ptr);
72 let low = vmovl_u16(vget_low_u16(vals));
73 let high = vmovl_u16(vget_high_u16(vals));
74
75 let out_ptr = output.as_mut_ptr().add(base);
76 vst1q_u32(out_ptr, low);
77 vst1q_u32(out_ptr.add(4), high);
78 }
79
80 let base = chunks * 8;
82 for i in 0..remainder {
83 let idx = (base + i) * 2;
84 output[base + i] = u16::from_le_bytes([input[idx], input[idx + 1]]) as u32;
85 }
86 }
87
88 #[target_feature(enable = "neon")]
90 pub unsafe fn unpack_32bit(input: &[u8], output: &mut [u32], count: usize) {
91 let chunks = count / 4;
92 let remainder = count % 4;
93
94 let in_ptr = input.as_ptr() as *const u32;
95 let out_ptr = output.as_mut_ptr();
96
97 for chunk in 0..chunks {
98 let vals = vld1q_u32(in_ptr.add(chunk * 4));
99 vst1q_u32(out_ptr.add(chunk * 4), vals);
100 }
101
102 let base = chunks * 4;
104 for i in 0..remainder {
105 let idx = (base + i) * 4;
106 output[base + i] =
107 u32::from_le_bytes([input[idx], input[idx + 1], input[idx + 2], input[idx + 3]]);
108 }
109 }
110
111 #[inline]
115 #[target_feature(enable = "neon")]
116 unsafe fn prefix_sum_4(v: uint32x4_t) -> uint32x4_t {
117 let shifted1 = vextq_u32(vdupq_n_u32(0), v, 3);
120 let sum1 = vaddq_u32(v, shifted1);
121
122 let shifted2 = vextq_u32(vdupq_n_u32(0), sum1, 2);
125 vaddq_u32(sum1, shifted2)
126 }
127
128 #[target_feature(enable = "neon")]
132 pub unsafe fn delta_decode(
133 output: &mut [u32],
134 deltas: &[u32],
135 first_doc_id: u32,
136 count: usize,
137 ) {
138 if count == 0 {
139 return;
140 }
141
142 output[0] = first_doc_id;
143 if count == 1 {
144 return;
145 }
146
147 let ones = vdupq_n_u32(1);
148 let mut carry = vdupq_n_u32(first_doc_id);
149
150 let full_groups = (count - 1) / 4;
151 let remainder = (count - 1) % 4;
152
153 for group in 0..full_groups {
154 let base = group * 4;
155
156 let d = vld1q_u32(deltas[base..].as_ptr());
158 let gaps = vaddq_u32(d, ones);
159
160 let prefix = prefix_sum_4(gaps);
162
163 let result = vaddq_u32(prefix, carry);
165
166 vst1q_u32(output[base + 1..].as_mut_ptr(), result);
168
169 carry = vdupq_n_u32(vgetq_lane_u32(result, 3));
171 }
172
173 let base = full_groups * 4;
175 let mut scalar_carry = vgetq_lane_u32(carry, 0);
176 for j in 0..remainder {
177 scalar_carry = scalar_carry.wrapping_add(deltas[base + j]).wrapping_add(1);
178 output[base + j + 1] = scalar_carry;
179 }
180 }
181
182 #[target_feature(enable = "neon")]
184 pub unsafe fn add_one(values: &mut [u32], count: usize) {
185 let ones = vdupq_n_u32(1);
186 let chunks = count / 4;
187 let remainder = count % 4;
188
189 for chunk in 0..chunks {
190 let base = chunk * 4;
191 let ptr = values.as_mut_ptr().add(base);
192 let v = vld1q_u32(ptr);
193 let result = vaddq_u32(v, ones);
194 vst1q_u32(ptr, result);
195 }
196
197 let base = chunks * 4;
198 for i in 0..remainder {
199 values[base + i] += 1;
200 }
201 }
202
203 #[target_feature(enable = "neon")]
206 pub unsafe fn unpack_8bit_delta_decode(
207 input: &[u8],
208 output: &mut [u32],
209 first_value: u32,
210 count: usize,
211 ) {
212 output[0] = first_value;
213 if count <= 1 {
214 return;
215 }
216
217 let ones = vdupq_n_u32(1);
218 let mut carry = vdupq_n_u32(first_value);
219
220 let full_groups = (count - 1) / 4;
221 let remainder = (count - 1) % 4;
222
223 for group in 0..full_groups {
224 let base = group * 4;
225
226 let raw = std::ptr::read_unaligned(input.as_ptr().add(base) as *const u32);
228 let bytes = vreinterpret_u8_u32(vdup_n_u32(raw));
229 let u16s = vmovl_u8(bytes); let d = vmovl_u16(vget_low_u16(u16s)); let gaps = vaddq_u32(d, ones);
234
235 let prefix = prefix_sum_4(gaps);
237
238 let result = vaddq_u32(prefix, carry);
240
241 vst1q_u32(output[base + 1..].as_mut_ptr(), result);
243
244 carry = vdupq_n_u32(vgetq_lane_u32(result, 3));
246 }
247
248 let base = full_groups * 4;
250 let mut scalar_carry = vgetq_lane_u32(carry, 0);
251 for j in 0..remainder {
252 scalar_carry = scalar_carry
253 .wrapping_add(input[base + j] as u32)
254 .wrapping_add(1);
255 output[base + j + 1] = scalar_carry;
256 }
257 }
258
259 #[target_feature(enable = "neon")]
261 pub unsafe fn unpack_16bit_delta_decode(
262 input: &[u8],
263 output: &mut [u32],
264 first_value: u32,
265 count: usize,
266 ) {
267 output[0] = first_value;
268 if count <= 1 {
269 return;
270 }
271
272 let ones = vdupq_n_u32(1);
273 let mut carry = vdupq_n_u32(first_value);
274
275 let full_groups = (count - 1) / 4;
276 let remainder = (count - 1) % 4;
277
278 for group in 0..full_groups {
279 let base = group * 4;
280 let in_ptr = input.as_ptr().add(base * 2) as *const u16;
281
282 let vals = vld1_u16(in_ptr);
284 let d = vmovl_u16(vals);
285
286 let gaps = vaddq_u32(d, ones);
288
289 let prefix = prefix_sum_4(gaps);
291
292 let result = vaddq_u32(prefix, carry);
294
295 vst1q_u32(output[base + 1..].as_mut_ptr(), result);
297
298 carry = vdupq_n_u32(vgetq_lane_u32(result, 3));
300 }
301
302 let base = full_groups * 4;
304 let mut scalar_carry = vgetq_lane_u32(carry, 0);
305 for j in 0..remainder {
306 let idx = (base + j) * 2;
307 let delta = u16::from_le_bytes([input[idx], input[idx + 1]]) as u32;
308 scalar_carry = scalar_carry.wrapping_add(delta).wrapping_add(1);
309 output[base + j + 1] = scalar_carry;
310 }
311 }
312
313 #[target_feature(enable = "neon")]
316 pub unsafe fn hamming_distance(a: &[u8], b: &[u8]) -> u32 {
317 let len = a.len();
318 let chunks16 = len / 16;
319 let mut total = 0u32;
320
321 let mut i = 0;
324 while i < chunks16 {
325 let batch_end = (i + 31).min(chunks16);
326 let mut acc = vdupq_n_u8(0);
327 for j in i..batch_end {
328 let off = j * 16;
329 let va = vld1q_u8(a.as_ptr().add(off));
330 let vb = vld1q_u8(b.as_ptr().add(off));
331 let popcnt = vcntq_u8(veorq_u8(va, vb));
332 acc = vaddq_u8(acc, popcnt);
333 }
334 let sum64 = vpaddlq_u32(vpaddlq_u16(vpaddlq_u8(acc)));
336 total += vgetq_lane_u64(sum64, 0) as u32 + vgetq_lane_u64(sum64, 1) as u32;
337 i = batch_end;
338 }
339
340 let base = chunks16 * 16;
342 for k in base..len {
343 total += (a[k] ^ b[k]).count_ones();
344 }
345
346 total
347 }
348
349 #[target_feature(enable = "neon")]
355 pub unsafe fn hamming_distance_x4(query: &[u8], rows: [&[u8]; 4]) -> [u32; 4] {
356 let len = query.len();
357 let chunks16 = len / 16;
358 let mut total = [0u32; 4];
359
360 let mut i = 0;
361 while i < chunks16 {
362 let batch_end = (i + 31).min(chunks16);
363 let mut acc = [vdupq_n_u8(0); 4];
364 for j in i..batch_end {
365 let off = j * 16;
366 let vq = vld1q_u8(query.as_ptr().add(off));
367 for r in 0..4 {
368 let vr = vld1q_u8(rows[r].as_ptr().add(off));
369 acc[r] = vaddq_u8(acc[r], vcntq_u8(veorq_u8(vq, vr)));
370 }
371 }
372 for r in 0..4 {
373 let sum64 = vpaddlq_u32(vpaddlq_u16(vpaddlq_u8(acc[r])));
374 total[r] += vgetq_lane_u64(sum64, 0) as u32 + vgetq_lane_u64(sum64, 1) as u32;
375 }
376 i = batch_end;
377 }
378
379 let base = chunks16 * 16;
382 if base < len {
383 let tail = &query[base..];
384 for r in 0..4 {
385 total[r] += super::hamming_distance_scalar(tail, &rows[r][base..]);
386 }
387 }
388
389 total
390 }
391
392 #[inline]
394 pub fn is_available() -> bool {
395 true
396 }
397}
398
399#[cfg(target_arch = "x86_64")]
404#[allow(unsafe_op_in_unsafe_fn)]
405mod sse {
406 use std::arch::x86_64::*;
407
408 #[target_feature(enable = "sse2", enable = "sse4.1")]
410 pub unsafe fn unpack_8bit(input: &[u8], output: &mut [u32], count: usize) {
411 let chunks = count / 16;
412 let remainder = count % 16;
413
414 for chunk in 0..chunks {
415 let base = chunk * 16;
416 let in_ptr = input.as_ptr().add(base);
417
418 let bytes = _mm_loadu_si128(in_ptr as *const __m128i);
419
420 let v0 = _mm_cvtepu8_epi32(bytes);
422 let v1 = _mm_cvtepu8_epi32(_mm_srli_si128(bytes, 4));
423 let v2 = _mm_cvtepu8_epi32(_mm_srli_si128(bytes, 8));
424 let v3 = _mm_cvtepu8_epi32(_mm_srli_si128(bytes, 12));
425
426 let out_ptr = output.as_mut_ptr().add(base);
427 _mm_storeu_si128(out_ptr as *mut __m128i, v0);
428 _mm_storeu_si128(out_ptr.add(4) as *mut __m128i, v1);
429 _mm_storeu_si128(out_ptr.add(8) as *mut __m128i, v2);
430 _mm_storeu_si128(out_ptr.add(12) as *mut __m128i, v3);
431 }
432
433 let base = chunks * 16;
434 for i in 0..remainder {
435 output[base + i] = input[base + i] as u32;
436 }
437 }
438
439 #[target_feature(enable = "sse2", enable = "sse4.1")]
441 pub unsafe fn unpack_16bit(input: &[u8], output: &mut [u32], count: usize) {
442 let chunks = count / 8;
443 let remainder = count % 8;
444
445 for chunk in 0..chunks {
446 let base = chunk * 8;
447 let in_ptr = input.as_ptr().add(base * 2);
448
449 let vals = _mm_loadu_si128(in_ptr as *const __m128i);
450 let low = _mm_cvtepu16_epi32(vals);
451 let high = _mm_cvtepu16_epi32(_mm_srli_si128(vals, 8));
452
453 let out_ptr = output.as_mut_ptr().add(base);
454 _mm_storeu_si128(out_ptr as *mut __m128i, low);
455 _mm_storeu_si128(out_ptr.add(4) as *mut __m128i, high);
456 }
457
458 let base = chunks * 8;
459 for i in 0..remainder {
460 let idx = (base + i) * 2;
461 output[base + i] = u16::from_le_bytes([input[idx], input[idx + 1]]) as u32;
462 }
463 }
464
465 #[target_feature(enable = "sse2")]
467 pub unsafe fn unpack_32bit(input: &[u8], output: &mut [u32], count: usize) {
468 let chunks = count / 4;
469 let remainder = count % 4;
470
471 let in_ptr = input.as_ptr() as *const __m128i;
472 let out_ptr = output.as_mut_ptr() as *mut __m128i;
473
474 for chunk in 0..chunks {
475 let vals = _mm_loadu_si128(in_ptr.add(chunk));
476 _mm_storeu_si128(out_ptr.add(chunk), vals);
477 }
478
479 let base = chunks * 4;
481 for i in 0..remainder {
482 let idx = (base + i) * 4;
483 output[base + i] =
484 u32::from_le_bytes([input[idx], input[idx + 1], input[idx + 2], input[idx + 3]]);
485 }
486 }
487
488 #[inline]
492 #[target_feature(enable = "sse2")]
493 unsafe fn prefix_sum_4(v: __m128i) -> __m128i {
494 let shifted1 = _mm_slli_si128(v, 4);
497 let sum1 = _mm_add_epi32(v, shifted1);
498
499 let shifted2 = _mm_slli_si128(sum1, 8);
502 _mm_add_epi32(sum1, shifted2)
503 }
504
505 #[target_feature(enable = "sse2", enable = "sse4.1")]
507 pub unsafe fn delta_decode(
508 output: &mut [u32],
509 deltas: &[u32],
510 first_doc_id: u32,
511 count: usize,
512 ) {
513 if count == 0 {
514 return;
515 }
516
517 output[0] = first_doc_id;
518 if count == 1 {
519 return;
520 }
521
522 let ones = _mm_set1_epi32(1);
523 let mut carry = _mm_set1_epi32(first_doc_id as i32);
524
525 let full_groups = (count - 1) / 4;
526 let remainder = (count - 1) % 4;
527
528 for group in 0..full_groups {
529 let base = group * 4;
530
531 let d = _mm_loadu_si128(deltas[base..].as_ptr() as *const __m128i);
533 let gaps = _mm_add_epi32(d, ones);
534
535 let prefix = prefix_sum_4(gaps);
537
538 let result = _mm_add_epi32(prefix, carry);
540
541 _mm_storeu_si128(output[base + 1..].as_mut_ptr() as *mut __m128i, result);
543
544 carry = _mm_shuffle_epi32(result, 0xFF); }
547
548 let base = full_groups * 4;
550 let mut scalar_carry = _mm_extract_epi32(carry, 0) as u32;
551 for j in 0..remainder {
552 scalar_carry = scalar_carry.wrapping_add(deltas[base + j]).wrapping_add(1);
553 output[base + j + 1] = scalar_carry;
554 }
555 }
556
557 #[target_feature(enable = "sse2")]
559 pub unsafe fn add_one(values: &mut [u32], count: usize) {
560 let ones = _mm_set1_epi32(1);
561 let chunks = count / 4;
562 let remainder = count % 4;
563
564 for chunk in 0..chunks {
565 let base = chunk * 4;
566 let ptr = values.as_mut_ptr().add(base) as *mut __m128i;
567 let v = _mm_loadu_si128(ptr);
568 let result = _mm_add_epi32(v, ones);
569 _mm_storeu_si128(ptr, result);
570 }
571
572 let base = chunks * 4;
573 for i in 0..remainder {
574 values[base + i] += 1;
575 }
576 }
577
578 #[target_feature(enable = "sse2", enable = "sse4.1")]
580 pub unsafe fn unpack_8bit_delta_decode(
581 input: &[u8],
582 output: &mut [u32],
583 first_value: u32,
584 count: usize,
585 ) {
586 output[0] = first_value;
587 if count <= 1 {
588 return;
589 }
590
591 let ones = _mm_set1_epi32(1);
592 let mut carry = _mm_set1_epi32(first_value as i32);
593
594 let full_groups = (count - 1) / 4;
595 let remainder = (count - 1) % 4;
596
597 for group in 0..full_groups {
598 let base = group * 4;
599
600 let bytes = _mm_cvtsi32_si128(std::ptr::read_unaligned(
602 input.as_ptr().add(base) as *const i32
603 ));
604 let d = _mm_cvtepu8_epi32(bytes);
605
606 let gaps = _mm_add_epi32(d, ones);
608
609 let prefix = prefix_sum_4(gaps);
611
612 let result = _mm_add_epi32(prefix, carry);
614
615 _mm_storeu_si128(output[base + 1..].as_mut_ptr() as *mut __m128i, result);
617
618 carry = _mm_shuffle_epi32(result, 0xFF);
620 }
621
622 let base = full_groups * 4;
624 let mut scalar_carry = _mm_extract_epi32(carry, 0) as u32;
625 for j in 0..remainder {
626 scalar_carry = scalar_carry
627 .wrapping_add(input[base + j] as u32)
628 .wrapping_add(1);
629 output[base + j + 1] = scalar_carry;
630 }
631 }
632
633 #[target_feature(enable = "sse2", enable = "sse4.1")]
635 pub unsafe fn unpack_16bit_delta_decode(
636 input: &[u8],
637 output: &mut [u32],
638 first_value: u32,
639 count: usize,
640 ) {
641 output[0] = first_value;
642 if count <= 1 {
643 return;
644 }
645
646 let ones = _mm_set1_epi32(1);
647 let mut carry = _mm_set1_epi32(first_value as i32);
648
649 let full_groups = (count - 1) / 4;
650 let remainder = (count - 1) % 4;
651
652 for group in 0..full_groups {
653 let base = group * 4;
654 let in_ptr = input.as_ptr().add(base * 2);
655
656 let vals = _mm_loadl_epi64(in_ptr as *const __m128i); let d = _mm_cvtepu16_epi32(vals);
659
660 let gaps = _mm_add_epi32(d, ones);
662
663 let prefix = prefix_sum_4(gaps);
665
666 let result = _mm_add_epi32(prefix, carry);
668
669 _mm_storeu_si128(output[base + 1..].as_mut_ptr() as *mut __m128i, result);
671
672 carry = _mm_shuffle_epi32(result, 0xFF);
674 }
675
676 let base = full_groups * 4;
678 let mut scalar_carry = _mm_extract_epi32(carry, 0) as u32;
679 for j in 0..remainder {
680 let idx = (base + j) * 2;
681 let delta = u16::from_le_bytes([input[idx], input[idx + 1]]) as u32;
682 scalar_carry = scalar_carry.wrapping_add(delta).wrapping_add(1);
683 output[base + j + 1] = scalar_carry;
684 }
685 }
686
687 #[inline]
689 pub fn is_available() -> bool {
690 is_x86_feature_detected!("sse4.1")
691 }
692}
693
694#[cfg(target_arch = "x86_64")]
699#[allow(unsafe_op_in_unsafe_fn)]
700mod avx2 {
701 use std::arch::x86_64::*;
702
703 #[target_feature(enable = "avx2")]
705 pub unsafe fn unpack_8bit(input: &[u8], output: &mut [u32], count: usize) {
706 let chunks = count / 32;
707 let remainder = count % 32;
708
709 for chunk in 0..chunks {
710 let base = chunk * 32;
711 let in_ptr = input.as_ptr().add(base);
712
713 let bytes_lo = _mm_loadu_si128(in_ptr as *const __m128i);
715 let bytes_hi = _mm_loadu_si128(in_ptr.add(16) as *const __m128i);
716
717 let v0 = _mm256_cvtepu8_epi32(bytes_lo);
719 let v1 = _mm256_cvtepu8_epi32(_mm_srli_si128(bytes_lo, 8));
720 let v2 = _mm256_cvtepu8_epi32(bytes_hi);
721 let v3 = _mm256_cvtepu8_epi32(_mm_srli_si128(bytes_hi, 8));
722
723 let out_ptr = output.as_mut_ptr().add(base);
724 _mm256_storeu_si256(out_ptr as *mut __m256i, v0);
725 _mm256_storeu_si256(out_ptr.add(8) as *mut __m256i, v1);
726 _mm256_storeu_si256(out_ptr.add(16) as *mut __m256i, v2);
727 _mm256_storeu_si256(out_ptr.add(24) as *mut __m256i, v3);
728 }
729
730 let base = chunks * 32;
732 for i in 0..remainder {
733 output[base + i] = input[base + i] as u32;
734 }
735 }
736
737 #[target_feature(enable = "avx2")]
739 pub unsafe fn unpack_16bit(input: &[u8], output: &mut [u32], count: usize) {
740 let chunks = count / 16;
741 let remainder = count % 16;
742
743 for chunk in 0..chunks {
744 let base = chunk * 16;
745 let in_ptr = input.as_ptr().add(base * 2);
746
747 let vals_lo = _mm_loadu_si128(in_ptr as *const __m128i);
749 let vals_hi = _mm_loadu_si128(in_ptr.add(16) as *const __m128i);
750
751 let v0 = _mm256_cvtepu16_epi32(vals_lo);
753 let v1 = _mm256_cvtepu16_epi32(vals_hi);
754
755 let out_ptr = output.as_mut_ptr().add(base);
756 _mm256_storeu_si256(out_ptr as *mut __m256i, v0);
757 _mm256_storeu_si256(out_ptr.add(8) as *mut __m256i, v1);
758 }
759
760 let base = chunks * 16;
762 for i in 0..remainder {
763 let idx = (base + i) * 2;
764 output[base + i] = u16::from_le_bytes([input[idx], input[idx + 1]]) as u32;
765 }
766 }
767
768 #[target_feature(enable = "avx2")]
770 pub unsafe fn unpack_32bit(input: &[u8], output: &mut [u32], count: usize) {
771 let chunks = count / 8;
772 let remainder = count % 8;
773
774 let in_ptr = input.as_ptr() as *const __m256i;
775 let out_ptr = output.as_mut_ptr() as *mut __m256i;
776
777 for chunk in 0..chunks {
778 let vals = _mm256_loadu_si256(in_ptr.add(chunk));
779 _mm256_storeu_si256(out_ptr.add(chunk), vals);
780 }
781
782 let base = chunks * 8;
784 for i in 0..remainder {
785 let idx = (base + i) * 4;
786 output[base + i] =
787 u32::from_le_bytes([input[idx], input[idx + 1], input[idx + 2], input[idx + 3]]);
788 }
789 }
790
791 #[target_feature(enable = "avx2")]
793 pub unsafe fn add_one(values: &mut [u32], count: usize) {
794 let ones = _mm256_set1_epi32(1);
795 let chunks = count / 8;
796 let remainder = count % 8;
797
798 for chunk in 0..chunks {
799 let base = chunk * 8;
800 let ptr = values.as_mut_ptr().add(base) as *mut __m256i;
801 let v = _mm256_loadu_si256(ptr);
802 let result = _mm256_add_epi32(v, ones);
803 _mm256_storeu_si256(ptr, result);
804 }
805
806 let base = chunks * 8;
807 for i in 0..remainder {
808 values[base + i] += 1;
809 }
810 }
811
812 #[inline]
816 #[target_feature(enable = "avx2")]
817 unsafe fn prefix_sum_8(v: __m256i) -> __m256i {
818 let s1 = _mm256_slli_si256(v, 4);
820 let r1 = _mm256_add_epi32(v, s1);
821
822 let s2 = _mm256_slli_si256(r1, 8);
824 let r2 = _mm256_add_epi32(r1, s2);
825
826 let lo_sum = _mm256_shuffle_epi32(r2, 0xFF);
829 let carry = _mm256_permute2x128_si256(lo_sum, lo_sum, 0x00);
831 let carry_hi = _mm256_blend_epi32::<0xF0>(_mm256_setzero_si256(), carry);
833 _mm256_add_epi32(r2, carry_hi)
834 }
835
836 #[target_feature(enable = "avx2")]
838 pub unsafe fn unpack_8bit_delta_decode(
839 input: &[u8],
840 output: &mut [u32],
841 first_value: u32,
842 count: usize,
843 ) {
844 output[0] = first_value;
845 if count <= 1 {
846 return;
847 }
848
849 let ones = _mm256_set1_epi32(1);
850 let mut carry = _mm256_set1_epi32(first_value as i32);
851 let broadcast_idx = _mm256_set1_epi32(7);
852
853 let full_groups = (count - 1) / 8;
854 let remainder = (count - 1) % 8;
855
856 for group in 0..full_groups {
857 let base = group * 8;
858
859 let bytes = _mm_loadl_epi64(input.as_ptr().add(base) as *const __m128i);
861 let d = _mm256_cvtepu8_epi32(bytes);
862
863 let gaps = _mm256_add_epi32(d, ones);
865
866 let prefix = prefix_sum_8(gaps);
868
869 let result = _mm256_add_epi32(prefix, carry);
871
872 _mm256_storeu_si256(output[base + 1..].as_mut_ptr() as *mut __m256i, result);
874
875 carry = _mm256_permutevar8x32_epi32(result, broadcast_idx);
877 }
878
879 let base = full_groups * 8;
881 let mut scalar_carry = _mm256_extract_epi32::<0>(carry) as u32;
882 for j in 0..remainder {
883 scalar_carry = scalar_carry
884 .wrapping_add(input[base + j] as u32)
885 .wrapping_add(1);
886 output[base + j + 1] = scalar_carry;
887 }
888 }
889
890 #[target_feature(enable = "avx2")]
892 pub unsafe fn unpack_16bit_delta_decode(
893 input: &[u8],
894 output: &mut [u32],
895 first_value: u32,
896 count: usize,
897 ) {
898 output[0] = first_value;
899 if count <= 1 {
900 return;
901 }
902
903 let ones = _mm256_set1_epi32(1);
904 let mut carry = _mm256_set1_epi32(first_value as i32);
905 let broadcast_idx = _mm256_set1_epi32(7);
906
907 let full_groups = (count - 1) / 8;
908 let remainder = (count - 1) % 8;
909
910 for group in 0..full_groups {
911 let base = group * 8;
912 let in_ptr = input.as_ptr().add(base * 2);
913
914 let vals = _mm_loadu_si128(in_ptr as *const __m128i);
916 let d = _mm256_cvtepu16_epi32(vals);
917
918 let gaps = _mm256_add_epi32(d, ones);
920
921 let prefix = prefix_sum_8(gaps);
923
924 let result = _mm256_add_epi32(prefix, carry);
926
927 _mm256_storeu_si256(output[base + 1..].as_mut_ptr() as *mut __m256i, result);
929
930 carry = _mm256_permutevar8x32_epi32(result, broadcast_idx);
932 }
933
934 let base = full_groups * 8;
936 let mut scalar_carry = _mm256_extract_epi32::<0>(carry) as u32;
937 for j in 0..remainder {
938 let idx = (base + j) * 2;
939 let delta = u16::from_le_bytes([input[idx], input[idx + 1]]) as u32;
940 scalar_carry = scalar_carry.wrapping_add(delta).wrapping_add(1);
941 output[base + j + 1] = scalar_carry;
942 }
943 }
944
945 #[target_feature(enable = "avx2")]
948 pub unsafe fn hamming_distance(a: &[u8], b: &[u8]) -> u32 {
949 let len = a.len();
950 let chunks32 = len / 32;
951 let low_mask = _mm256_set1_epi8(0x0f);
952 let lookup = _mm256_setr_epi8(
954 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,
955 3, 3, 4,
956 );
957 let mut total = 0u64;
958
959 let mut i = 0;
960 while i < chunks32 {
961 let batch_end = (i + 31).min(chunks32);
964 let mut acc = _mm256_setzero_si256();
965 for j in i..batch_end {
966 let off = j * 32;
967 let va = _mm256_loadu_si256(a.as_ptr().add(off) as *const __m256i);
968 let vb = _mm256_loadu_si256(b.as_ptr().add(off) as *const __m256i);
969 let xored = _mm256_xor_si256(va, vb);
970 let lo = _mm256_and_si256(xored, low_mask);
972 let hi = _mm256_and_si256(_mm256_srli_epi16(xored, 4), low_mask);
973 let popcnt = _mm256_add_epi8(
974 _mm256_shuffle_epi8(lookup, lo),
975 _mm256_shuffle_epi8(lookup, hi),
976 );
977 acc = _mm256_add_epi8(acc, popcnt);
978 }
979 let sad = _mm256_sad_epu8(acc, _mm256_setzero_si256());
981 total += _mm256_extract_epi64(sad, 0) as u64
982 + _mm256_extract_epi64(sad, 1) as u64
983 + _mm256_extract_epi64(sad, 2) as u64
984 + _mm256_extract_epi64(sad, 3) as u64;
985 i = batch_end;
986 }
987
988 let base = chunks32 * 32;
990 for k in base..len {
991 total += (a[k] ^ b[k]).count_ones() as u64;
992 }
993
994 total as u32
995 }
996
997 #[target_feature(enable = "avx2")]
1003 pub unsafe fn hamming_distance_x4(query: &[u8], rows: [&[u8]; 4]) -> [u32; 4] {
1004 let len = query.len();
1005 let chunks32 = len / 32;
1006 let low_mask = _mm256_set1_epi8(0x0f);
1007 let lookup = _mm256_setr_epi8(
1008 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,
1009 3, 3, 4,
1010 );
1011 let mut total = [0u64; 4];
1012
1013 let mut i = 0;
1014 while i < chunks32 {
1015 let batch_end = (i + 31).min(chunks32);
1016 let mut acc = [_mm256_setzero_si256(); 4];
1017 for j in i..batch_end {
1018 let off = j * 32;
1019 let vq = _mm256_loadu_si256(query.as_ptr().add(off) as *const __m256i);
1020 for r in 0..4 {
1021 let vr = _mm256_loadu_si256(rows[r].as_ptr().add(off) as *const __m256i);
1022 let xored = _mm256_xor_si256(vq, vr);
1023 let lo = _mm256_and_si256(xored, low_mask);
1024 let hi = _mm256_and_si256(_mm256_srli_epi16(xored, 4), low_mask);
1025 acc[r] = _mm256_add_epi8(
1026 acc[r],
1027 _mm256_add_epi8(
1028 _mm256_shuffle_epi8(lookup, lo),
1029 _mm256_shuffle_epi8(lookup, hi),
1030 ),
1031 );
1032 }
1033 }
1034 for r in 0..4 {
1035 let sad = _mm256_sad_epu8(acc[r], _mm256_setzero_si256());
1036 total[r] += _mm256_extract_epi64(sad, 0) as u64
1037 + _mm256_extract_epi64(sad, 1) as u64
1038 + _mm256_extract_epi64(sad, 2) as u64
1039 + _mm256_extract_epi64(sad, 3) as u64;
1040 }
1041 i = batch_end;
1042 }
1043
1044 let base = chunks32 * 32;
1047 if base < len {
1048 let tail = &query[base..];
1049 for r in 0..4 {
1050 total[r] += u64::from(super::hamming_distance_scalar(tail, &rows[r][base..]));
1051 }
1052 }
1053
1054 [
1055 total[0] as u32,
1056 total[1] as u32,
1057 total[2] as u32,
1058 total[3] as u32,
1059 ]
1060 }
1061
1062 #[inline]
1064 pub fn is_available() -> bool {
1065 is_x86_feature_detected!("avx2")
1066 }
1067}
1068
1069#[allow(dead_code)]
1074mod scalar {
1075 #[inline]
1077 pub fn unpack_8bit(input: &[u8], output: &mut [u32], count: usize) {
1078 for i in 0..count {
1079 output[i] = input[i] as u32;
1080 }
1081 }
1082
1083 #[inline]
1085 pub fn unpack_16bit(input: &[u8], output: &mut [u32], count: usize) {
1086 for (i, out) in output.iter_mut().enumerate().take(count) {
1087 let idx = i * 2;
1088 *out = u16::from_le_bytes([input[idx], input[idx + 1]]) as u32;
1089 }
1090 }
1091
1092 #[inline]
1094 pub fn unpack_32bit(input: &[u8], output: &mut [u32], count: usize) {
1095 for (i, out) in output.iter_mut().enumerate().take(count) {
1096 let idx = i * 4;
1097 *out = u32::from_le_bytes([input[idx], input[idx + 1], input[idx + 2], input[idx + 3]]);
1098 }
1099 }
1100
1101 #[inline]
1103 pub fn delta_decode(output: &mut [u32], deltas: &[u32], first_doc_id: u32, count: usize) {
1104 if count == 0 {
1105 return;
1106 }
1107
1108 output[0] = first_doc_id;
1109 let mut carry = first_doc_id;
1110
1111 for i in 0..count - 1 {
1112 carry = carry.wrapping_add(deltas[i]).wrapping_add(1);
1113 output[i + 1] = carry;
1114 }
1115 }
1116
1117 #[inline]
1119 pub fn add_one(values: &mut [u32], count: usize) {
1120 for val in values.iter_mut().take(count) {
1121 *val += 1;
1122 }
1123 }
1124}
1125
1126#[inline]
1132pub fn unpack_8bit(input: &[u8], output: &mut [u32], count: usize) {
1133 #[cfg(target_arch = "aarch64")]
1134 {
1135 if neon::is_available() {
1136 unsafe {
1137 neon::unpack_8bit(input, output, count);
1138 }
1139 return;
1140 }
1141 }
1142
1143 #[cfg(target_arch = "x86_64")]
1144 {
1145 if avx2::is_available() {
1147 unsafe {
1148 avx2::unpack_8bit(input, output, count);
1149 }
1150 return;
1151 }
1152 if sse::is_available() {
1153 unsafe {
1154 sse::unpack_8bit(input, output, count);
1155 }
1156 return;
1157 }
1158 }
1159
1160 scalar::unpack_8bit(input, output, count);
1161}
1162
1163#[inline]
1165pub fn unpack_16bit(input: &[u8], output: &mut [u32], count: usize) {
1166 #[cfg(target_arch = "aarch64")]
1167 {
1168 if neon::is_available() {
1169 unsafe {
1170 neon::unpack_16bit(input, output, count);
1171 }
1172 return;
1173 }
1174 }
1175
1176 #[cfg(target_arch = "x86_64")]
1177 {
1178 if avx2::is_available() {
1180 unsafe {
1181 avx2::unpack_16bit(input, output, count);
1182 }
1183 return;
1184 }
1185 if sse::is_available() {
1186 unsafe {
1187 sse::unpack_16bit(input, output, count);
1188 }
1189 return;
1190 }
1191 }
1192
1193 scalar::unpack_16bit(input, output, count);
1194}
1195
1196#[inline]
1198pub fn unpack_32bit(input: &[u8], output: &mut [u32], count: usize) {
1199 #[cfg(target_arch = "aarch64")]
1200 {
1201 if neon::is_available() {
1202 unsafe {
1203 neon::unpack_32bit(input, output, count);
1204 }
1205 return;
1206 }
1207 }
1208
1209 #[cfg(target_arch = "x86_64")]
1210 {
1211 if avx2::is_available() {
1213 unsafe {
1214 avx2::unpack_32bit(input, output, count);
1215 }
1216 return;
1217 }
1218 if sse::is_available() {
1219 unsafe {
1220 sse::unpack_32bit(input, output, count);
1221 }
1222 return;
1223 }
1224 }
1225
1226 scalar::unpack_32bit(input, output, count);
1227}
1228
1229#[inline]
1235pub fn delta_decode(output: &mut [u32], deltas: &[u32], first_value: u32, count: usize) {
1236 #[cfg(target_arch = "aarch64")]
1237 {
1238 if neon::is_available() {
1239 unsafe {
1240 neon::delta_decode(output, deltas, first_value, count);
1241 }
1242 return;
1243 }
1244 }
1245
1246 #[cfg(target_arch = "x86_64")]
1247 {
1248 if sse::is_available() {
1249 unsafe {
1250 sse::delta_decode(output, deltas, first_value, count);
1251 }
1252 return;
1253 }
1254 }
1255
1256 scalar::delta_decode(output, deltas, first_value, count);
1257}
1258
1259#[inline]
1263pub fn add_one(values: &mut [u32], count: usize) {
1264 #[cfg(target_arch = "aarch64")]
1265 {
1266 if neon::is_available() {
1267 unsafe {
1268 neon::add_one(values, count);
1269 }
1270 return;
1271 }
1272 }
1273
1274 #[cfg(target_arch = "x86_64")]
1275 {
1276 if avx2::is_available() {
1278 unsafe {
1279 avx2::add_one(values, count);
1280 }
1281 return;
1282 }
1283 if sse::is_available() {
1284 unsafe {
1285 sse::add_one(values, count);
1286 }
1287 return;
1288 }
1289 }
1290
1291 scalar::add_one(values, count);
1292}
1293
1294#[inline]
1296pub fn bits_needed(val: u32) -> u8 {
1297 if val == 0 {
1298 0
1299 } else {
1300 32 - val.leading_zeros() as u8
1301 }
1302}
1303
1304#[derive(Debug, Clone, Copy, PartialEq, Eq)]
1321#[repr(u8)]
1322pub enum RoundedBitWidth {
1323 Zero = 0,
1324 Bits8 = 8,
1325 Bits16 = 16,
1326 Bits32 = 32,
1327}
1328
1329impl RoundedBitWidth {
1330 #[inline]
1332 pub fn from_exact(bits: u8) -> Self {
1333 match bits {
1334 0 => RoundedBitWidth::Zero,
1335 1..=8 => RoundedBitWidth::Bits8,
1336 9..=16 => RoundedBitWidth::Bits16,
1337 _ => RoundedBitWidth::Bits32,
1338 }
1339 }
1340
1341 #[inline]
1343 pub fn from_u8(bits: u8) -> Self {
1344 match bits {
1345 0 => RoundedBitWidth::Zero,
1346 8 => RoundedBitWidth::Bits8,
1347 16 => RoundedBitWidth::Bits16,
1348 32 => RoundedBitWidth::Bits32,
1349 _ => RoundedBitWidth::Bits32, }
1351 }
1352
1353 #[inline]
1355 pub fn bytes_per_value(self) -> usize {
1356 match self {
1357 RoundedBitWidth::Zero => 0,
1358 RoundedBitWidth::Bits8 => 1,
1359 RoundedBitWidth::Bits16 => 2,
1360 RoundedBitWidth::Bits32 => 4,
1361 }
1362 }
1363
1364 #[inline]
1366 pub fn as_u8(self) -> u8 {
1367 self as u8
1368 }
1369}
1370
1371#[inline]
1373pub fn round_bit_width(bits: u8) -> u8 {
1374 RoundedBitWidth::from_exact(bits).as_u8()
1375}
1376
1377#[inline]
1382pub fn pack_rounded(values: &[u32], bit_width: RoundedBitWidth, output: &mut [u8]) -> usize {
1383 let count = values.len();
1384 match bit_width {
1385 RoundedBitWidth::Zero => 0,
1386 RoundedBitWidth::Bits8 => {
1387 for (i, &v) in values.iter().enumerate() {
1388 output[i] = v as u8;
1389 }
1390 count
1391 }
1392 RoundedBitWidth::Bits16 => {
1393 for (i, &v) in values.iter().enumerate() {
1394 let bytes = (v as u16).to_le_bytes();
1395 output[i * 2] = bytes[0];
1396 output[i * 2 + 1] = bytes[1];
1397 }
1398 count * 2
1399 }
1400 RoundedBitWidth::Bits32 => {
1401 for (i, &v) in values.iter().enumerate() {
1402 let bytes = v.to_le_bytes();
1403 output[i * 4] = bytes[0];
1404 output[i * 4 + 1] = bytes[1];
1405 output[i * 4 + 2] = bytes[2];
1406 output[i * 4 + 3] = bytes[3];
1407 }
1408 count * 4
1409 }
1410 }
1411}
1412
1413#[inline]
1417pub fn unpack_rounded(input: &[u8], bit_width: RoundedBitWidth, output: &mut [u32], count: usize) {
1418 match bit_width {
1419 RoundedBitWidth::Zero => {
1420 for out in output.iter_mut().take(count) {
1421 *out = 0;
1422 }
1423 }
1424 RoundedBitWidth::Bits8 => unpack_8bit(input, output, count),
1425 RoundedBitWidth::Bits16 => unpack_16bit(input, output, count),
1426 RoundedBitWidth::Bits32 => unpack_32bit(input, output, count),
1427 }
1428}
1429
1430#[inline]
1434pub fn unpack_rounded_delta_decode(
1435 input: &[u8],
1436 bit_width: RoundedBitWidth,
1437 output: &mut [u32],
1438 first_value: u32,
1439 count: usize,
1440) {
1441 match bit_width {
1442 RoundedBitWidth::Zero => {
1443 let mut val = first_value;
1445 for out in output.iter_mut().take(count) {
1446 *out = val;
1447 val = val.wrapping_add(1);
1448 }
1449 }
1450 RoundedBitWidth::Bits8 => unpack_8bit_delta_decode(input, output, first_value, count),
1451 RoundedBitWidth::Bits16 => unpack_16bit_delta_decode(input, output, first_value, count),
1452 RoundedBitWidth::Bits32 => {
1453 if count > 0 {
1455 output[0] = first_value;
1456 let mut carry = first_value;
1457 for i in 0..count - 1 {
1458 let idx = i * 4;
1459 let delta = u32::from_le_bytes([
1460 input[idx],
1461 input[idx + 1],
1462 input[idx + 2],
1463 input[idx + 3],
1464 ]);
1465 carry = carry.wrapping_add(delta).wrapping_add(1);
1466 output[i + 1] = carry;
1467 }
1468 }
1469 }
1470 }
1471}
1472
1473#[inline]
1482pub fn unpack_8bit_delta_decode(input: &[u8], output: &mut [u32], first_value: u32, count: usize) {
1483 if count == 0 {
1484 return;
1485 }
1486
1487 output[0] = first_value;
1488 if count == 1 {
1489 return;
1490 }
1491
1492 #[cfg(target_arch = "aarch64")]
1493 {
1494 if neon::is_available() {
1495 unsafe {
1496 neon::unpack_8bit_delta_decode(input, output, first_value, count);
1497 }
1498 return;
1499 }
1500 }
1501
1502 #[cfg(target_arch = "x86_64")]
1503 {
1504 if avx2::is_available() {
1505 unsafe {
1506 avx2::unpack_8bit_delta_decode(input, output, first_value, count);
1507 }
1508 return;
1509 }
1510 if sse::is_available() {
1511 unsafe {
1512 sse::unpack_8bit_delta_decode(input, output, first_value, count);
1513 }
1514 return;
1515 }
1516 }
1517
1518 let mut carry = first_value;
1520 for i in 0..count - 1 {
1521 carry = carry.wrapping_add(input[i] as u32).wrapping_add(1);
1522 output[i + 1] = carry;
1523 }
1524}
1525
1526#[inline]
1528pub fn unpack_16bit_delta_decode(input: &[u8], output: &mut [u32], first_value: u32, count: usize) {
1529 if count == 0 {
1530 return;
1531 }
1532
1533 output[0] = first_value;
1534 if count == 1 {
1535 return;
1536 }
1537
1538 #[cfg(target_arch = "aarch64")]
1539 {
1540 if neon::is_available() {
1541 unsafe {
1542 neon::unpack_16bit_delta_decode(input, output, first_value, count);
1543 }
1544 return;
1545 }
1546 }
1547
1548 #[cfg(target_arch = "x86_64")]
1549 {
1550 if avx2::is_available() {
1551 unsafe {
1552 avx2::unpack_16bit_delta_decode(input, output, first_value, count);
1553 }
1554 return;
1555 }
1556 if sse::is_available() {
1557 unsafe {
1558 sse::unpack_16bit_delta_decode(input, output, first_value, count);
1559 }
1560 return;
1561 }
1562 }
1563
1564 let mut carry = first_value;
1566 for i in 0..count - 1 {
1567 let idx = i * 2;
1568 let delta = u16::from_le_bytes([input[idx], input[idx + 1]]) as u32;
1569 carry = carry.wrapping_add(delta).wrapping_add(1);
1570 output[i + 1] = carry;
1571 }
1572}
1573
1574#[inline]
1579pub fn unpack_delta_decode(
1580 input: &[u8],
1581 bit_width: u8,
1582 output: &mut [u32],
1583 first_value: u32,
1584 count: usize,
1585) {
1586 if count == 0 {
1587 return;
1588 }
1589
1590 output[0] = first_value;
1591 if count == 1 {
1592 return;
1593 }
1594
1595 match bit_width {
1597 0 => {
1598 let mut val = first_value;
1600 for item in output.iter_mut().take(count).skip(1) {
1601 val = val.wrapping_add(1);
1602 *item = val;
1603 }
1604 }
1605 8 => unpack_8bit_delta_decode(input, output, first_value, count),
1606 16 => unpack_16bit_delta_decode(input, output, first_value, count),
1607 32 => {
1608 let mut carry = first_value;
1610 for i in 0..count - 1 {
1611 let idx = i * 4;
1612 let delta = u32::from_le_bytes([
1613 input[idx],
1614 input[idx + 1],
1615 input[idx + 2],
1616 input[idx + 3],
1617 ]);
1618 carry = carry.wrapping_add(delta).wrapping_add(1);
1619 output[i + 1] = carry;
1620 }
1621 }
1622 _ => {
1623 let mask = (1u64 << bit_width) - 1;
1625 let bit_width_usize = bit_width as usize;
1626 let mut bit_pos = 0usize;
1627 let input_ptr = input.as_ptr();
1628 let mut carry = first_value;
1629
1630 for i in 0..count - 1 {
1631 let byte_idx = bit_pos >> 3;
1632 let bit_offset = bit_pos & 7;
1633
1634 let word = unsafe { (input_ptr.add(byte_idx) as *const u64).read_unaligned() };
1636 let delta = ((word >> bit_offset) & mask) as u32;
1637
1638 carry = carry.wrapping_add(delta).wrapping_add(1);
1639 output[i + 1] = carry;
1640 bit_pos += bit_width_usize;
1641 }
1642 }
1643 }
1644}
1645
1646#[inline]
1654pub fn dequantize_uint8(input: &[u8], output: &mut [f32], scale: f32, min_val: f32, count: usize) {
1655 #[cfg(target_arch = "aarch64")]
1656 {
1657 if neon::is_available() {
1658 unsafe {
1659 dequantize_uint8_neon(input, output, scale, min_val, count);
1660 }
1661 return;
1662 }
1663 }
1664
1665 #[cfg(target_arch = "x86_64")]
1666 {
1667 if sse::is_available() {
1668 unsafe {
1669 dequantize_uint8_sse(input, output, scale, min_val, count);
1670 }
1671 return;
1672 }
1673 }
1674
1675 for i in 0..count {
1677 output[i] = input[i] as f32 * scale + min_val;
1678 }
1679}
1680
1681#[cfg(target_arch = "aarch64")]
1682#[target_feature(enable = "neon")]
1683#[allow(unsafe_op_in_unsafe_fn)]
1684unsafe fn dequantize_uint8_neon(
1685 input: &[u8],
1686 output: &mut [f32],
1687 scale: f32,
1688 min_val: f32,
1689 count: usize,
1690) {
1691 use std::arch::aarch64::*;
1692
1693 let scale_v = vdupq_n_f32(scale);
1694 let min_v = vdupq_n_f32(min_val);
1695
1696 let chunks = count / 16;
1697 let remainder = count % 16;
1698
1699 for chunk in 0..chunks {
1700 let base = chunk * 16;
1701 let in_ptr = input.as_ptr().add(base);
1702
1703 let bytes = vld1q_u8(in_ptr);
1705
1706 let low8 = vget_low_u8(bytes);
1708 let high8 = vget_high_u8(bytes);
1709
1710 let low16 = vmovl_u8(low8);
1711 let high16 = vmovl_u8(high8);
1712
1713 let u32_0 = vmovl_u16(vget_low_u16(low16));
1715 let u32_1 = vmovl_u16(vget_high_u16(low16));
1716 let u32_2 = vmovl_u16(vget_low_u16(high16));
1717 let u32_3 = vmovl_u16(vget_high_u16(high16));
1718
1719 let f32_0 = vfmaq_f32(min_v, vcvtq_f32_u32(u32_0), scale_v);
1721 let f32_1 = vfmaq_f32(min_v, vcvtq_f32_u32(u32_1), scale_v);
1722 let f32_2 = vfmaq_f32(min_v, vcvtq_f32_u32(u32_2), scale_v);
1723 let f32_3 = vfmaq_f32(min_v, vcvtq_f32_u32(u32_3), scale_v);
1724
1725 let out_ptr = output.as_mut_ptr().add(base);
1726 vst1q_f32(out_ptr, f32_0);
1727 vst1q_f32(out_ptr.add(4), f32_1);
1728 vst1q_f32(out_ptr.add(8), f32_2);
1729 vst1q_f32(out_ptr.add(12), f32_3);
1730 }
1731
1732 let base = chunks * 16;
1734 for i in 0..remainder {
1735 output[base + i] = input[base + i] as f32 * scale + min_val;
1736 }
1737}
1738
1739#[cfg(target_arch = "x86_64")]
1740#[target_feature(enable = "sse2", enable = "sse4.1")]
1741#[allow(unsafe_op_in_unsafe_fn)]
1742unsafe fn dequantize_uint8_sse(
1743 input: &[u8],
1744 output: &mut [f32],
1745 scale: f32,
1746 min_val: f32,
1747 count: usize,
1748) {
1749 use std::arch::x86_64::*;
1750
1751 let scale_v = _mm_set1_ps(scale);
1752 let min_v = _mm_set1_ps(min_val);
1753
1754 let chunks = count / 4;
1755 let remainder = count % 4;
1756
1757 for chunk in 0..chunks {
1758 let base = chunk * 4;
1759
1760 let bytes = _mm_cvtsi32_si128(std::ptr::read_unaligned(
1762 input.as_ptr().add(base) as *const i32
1763 ));
1764 let ints = _mm_cvtepu8_epi32(bytes);
1765 let floats = _mm_cvtepi32_ps(ints);
1766
1767 let scaled = _mm_add_ps(_mm_mul_ps(floats, scale_v), min_v);
1769
1770 _mm_storeu_ps(output.as_mut_ptr().add(base), scaled);
1771 }
1772
1773 let base = chunks * 4;
1775 for i in 0..remainder {
1776 output[base + i] = input[base + i] as f32 * scale + min_val;
1777 }
1778}
1779
1780#[inline]
1782pub fn dot_product_f32(a: &[f32], b: &[f32], count: usize) -> f32 {
1783 assert!(
1784 count <= a.len() && count <= b.len(),
1785 "dot_product_f32 count {count} exceeds input lengths ({}, {})",
1786 a.len(),
1787 b.len()
1788 );
1789 #[cfg(target_arch = "aarch64")]
1790 {
1791 if neon::is_available() {
1792 return unsafe { dot_product_f32_neon(a, b, count) };
1793 }
1794 }
1795
1796 #[cfg(target_arch = "x86_64")]
1797 {
1798 if is_x86_feature_detected!("avx512f") {
1799 return unsafe { dot_product_f32_avx512(a, b, count) };
1800 }
1801 if is_x86_feature_detected!("avx2") && is_x86_feature_detected!("fma") {
1802 return unsafe { dot_product_f32_avx2(a, b, count) };
1803 }
1804 if sse::is_available() {
1805 return unsafe { dot_product_f32_sse(a, b, count) };
1806 }
1807 }
1808
1809 let mut sum = 0.0f32;
1811 for i in 0..count {
1812 sum += a[i] * b[i];
1813 }
1814 sum
1815}
1816
1817#[cfg(target_arch = "aarch64")]
1818#[target_feature(enable = "neon")]
1819#[allow(unsafe_op_in_unsafe_fn)]
1820unsafe fn dot_product_f32_neon(a: &[f32], b: &[f32], count: usize) -> f32 {
1821 use std::arch::aarch64::*;
1822
1823 let chunks16 = count / 16;
1824 let remainder = count % 16;
1825
1826 let mut acc0 = vdupq_n_f32(0.0);
1827 let mut acc1 = vdupq_n_f32(0.0);
1828 let mut acc2 = vdupq_n_f32(0.0);
1829 let mut acc3 = vdupq_n_f32(0.0);
1830
1831 for c in 0..chunks16 {
1832 let base = c * 16;
1833 acc0 = vfmaq_f32(
1834 acc0,
1835 vld1q_f32(a.as_ptr().add(base)),
1836 vld1q_f32(b.as_ptr().add(base)),
1837 );
1838 acc1 = vfmaq_f32(
1839 acc1,
1840 vld1q_f32(a.as_ptr().add(base + 4)),
1841 vld1q_f32(b.as_ptr().add(base + 4)),
1842 );
1843 acc2 = vfmaq_f32(
1844 acc2,
1845 vld1q_f32(a.as_ptr().add(base + 8)),
1846 vld1q_f32(b.as_ptr().add(base + 8)),
1847 );
1848 acc3 = vfmaq_f32(
1849 acc3,
1850 vld1q_f32(a.as_ptr().add(base + 12)),
1851 vld1q_f32(b.as_ptr().add(base + 12)),
1852 );
1853 }
1854
1855 let acc = vaddq_f32(vaddq_f32(acc0, acc1), vaddq_f32(acc2, acc3));
1856 let mut sum = vaddvq_f32(acc);
1857
1858 let base = chunks16 * 16;
1859 for i in 0..remainder {
1860 sum += a[base + i] * b[base + i];
1861 }
1862
1863 sum
1864}
1865
1866#[cfg(target_arch = "x86_64")]
1867#[target_feature(enable = "avx2", enable = "fma")]
1868#[allow(unsafe_op_in_unsafe_fn)]
1869unsafe fn dot_product_f32_avx2(a: &[f32], b: &[f32], count: usize) -> f32 {
1870 use std::arch::x86_64::*;
1871
1872 let chunks32 = count / 32;
1873 let remainder = count % 32;
1874
1875 let mut acc0 = _mm256_setzero_ps();
1876 let mut acc1 = _mm256_setzero_ps();
1877 let mut acc2 = _mm256_setzero_ps();
1878 let mut acc3 = _mm256_setzero_ps();
1879
1880 for c in 0..chunks32 {
1881 let base = c * 32;
1882 acc0 = _mm256_fmadd_ps(
1883 _mm256_loadu_ps(a.as_ptr().add(base)),
1884 _mm256_loadu_ps(b.as_ptr().add(base)),
1885 acc0,
1886 );
1887 acc1 = _mm256_fmadd_ps(
1888 _mm256_loadu_ps(a.as_ptr().add(base + 8)),
1889 _mm256_loadu_ps(b.as_ptr().add(base + 8)),
1890 acc1,
1891 );
1892 acc2 = _mm256_fmadd_ps(
1893 _mm256_loadu_ps(a.as_ptr().add(base + 16)),
1894 _mm256_loadu_ps(b.as_ptr().add(base + 16)),
1895 acc2,
1896 );
1897 acc3 = _mm256_fmadd_ps(
1898 _mm256_loadu_ps(a.as_ptr().add(base + 24)),
1899 _mm256_loadu_ps(b.as_ptr().add(base + 24)),
1900 acc3,
1901 );
1902 }
1903
1904 let acc = _mm256_add_ps(_mm256_add_ps(acc0, acc1), _mm256_add_ps(acc2, acc3));
1905
1906 let hi = _mm256_extractf128_ps(acc, 1);
1908 let lo = _mm256_castps256_ps128(acc);
1909 let sum128 = _mm_add_ps(lo, hi);
1910 let shuf = _mm_shuffle_ps(sum128, sum128, 0b10_11_00_01);
1911 let sums = _mm_add_ps(sum128, shuf);
1912 let shuf2 = _mm_movehl_ps(sums, sums);
1913 let final_sum = _mm_add_ss(sums, shuf2);
1914
1915 let mut sum = _mm_cvtss_f32(final_sum);
1916
1917 let base = chunks32 * 32;
1918 for i in 0..remainder {
1919 sum += a[base + i] * b[base + i];
1920 }
1921
1922 sum
1923}
1924
1925#[cfg(target_arch = "x86_64")]
1926#[target_feature(enable = "sse")]
1927#[allow(unsafe_op_in_unsafe_fn)]
1928unsafe fn dot_product_f32_sse(a: &[f32], b: &[f32], count: usize) -> f32 {
1929 use std::arch::x86_64::*;
1930
1931 let chunks = count / 4;
1932 let remainder = count % 4;
1933
1934 let mut acc = _mm_setzero_ps();
1935
1936 for chunk in 0..chunks {
1937 let base = chunk * 4;
1938 let va = _mm_loadu_ps(a.as_ptr().add(base));
1939 let vb = _mm_loadu_ps(b.as_ptr().add(base));
1940 acc = _mm_add_ps(acc, _mm_mul_ps(va, vb));
1941 }
1942
1943 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);
1950
1951 let base = chunks * 4;
1953 for i in 0..remainder {
1954 sum += a[base + i] * b[base + i];
1955 }
1956
1957 sum
1958}
1959
1960#[cfg(target_arch = "x86_64")]
1961#[target_feature(enable = "avx512f")]
1962#[allow(unsafe_op_in_unsafe_fn)]
1963unsafe fn dot_product_f32_avx512(a: &[f32], b: &[f32], count: usize) -> f32 {
1964 use std::arch::x86_64::*;
1965
1966 let chunks64 = count / 64;
1967 let remainder = count % 64;
1968
1969 let mut acc0 = _mm512_setzero_ps();
1970 let mut acc1 = _mm512_setzero_ps();
1971 let mut acc2 = _mm512_setzero_ps();
1972 let mut acc3 = _mm512_setzero_ps();
1973
1974 for c in 0..chunks64 {
1975 let base = c * 64;
1976 acc0 = _mm512_fmadd_ps(
1977 _mm512_loadu_ps(a.as_ptr().add(base)),
1978 _mm512_loadu_ps(b.as_ptr().add(base)),
1979 acc0,
1980 );
1981 acc1 = _mm512_fmadd_ps(
1982 _mm512_loadu_ps(a.as_ptr().add(base + 16)),
1983 _mm512_loadu_ps(b.as_ptr().add(base + 16)),
1984 acc1,
1985 );
1986 acc2 = _mm512_fmadd_ps(
1987 _mm512_loadu_ps(a.as_ptr().add(base + 32)),
1988 _mm512_loadu_ps(b.as_ptr().add(base + 32)),
1989 acc2,
1990 );
1991 acc3 = _mm512_fmadd_ps(
1992 _mm512_loadu_ps(a.as_ptr().add(base + 48)),
1993 _mm512_loadu_ps(b.as_ptr().add(base + 48)),
1994 acc3,
1995 );
1996 }
1997
1998 let acc = _mm512_add_ps(_mm512_add_ps(acc0, acc1), _mm512_add_ps(acc2, acc3));
1999 let mut sum = _mm512_reduce_add_ps(acc);
2000
2001 let base = chunks64 * 64;
2002 for i in 0..remainder {
2003 sum += a[base + i] * b[base + i];
2004 }
2005
2006 sum
2007}
2008
2009#[cfg(target_arch = "x86_64")]
2010#[target_feature(enable = "avx512f")]
2011#[allow(unsafe_op_in_unsafe_fn)]
2012unsafe fn fused_dot_norm_avx512(a: &[f32], b: &[f32], count: usize) -> (f32, f32) {
2013 use std::arch::x86_64::*;
2014
2015 let chunks64 = count / 64;
2016 let remainder = count % 64;
2017
2018 let mut d0 = _mm512_setzero_ps();
2019 let mut d1 = _mm512_setzero_ps();
2020 let mut d2 = _mm512_setzero_ps();
2021 let mut d3 = _mm512_setzero_ps();
2022 let mut n0 = _mm512_setzero_ps();
2023 let mut n1 = _mm512_setzero_ps();
2024 let mut n2 = _mm512_setzero_ps();
2025 let mut n3 = _mm512_setzero_ps();
2026
2027 for c in 0..chunks64 {
2028 let base = c * 64;
2029 let vb0 = _mm512_loadu_ps(b.as_ptr().add(base));
2030 d0 = _mm512_fmadd_ps(_mm512_loadu_ps(a.as_ptr().add(base)), vb0, d0);
2031 n0 = _mm512_fmadd_ps(vb0, vb0, n0);
2032 let vb1 = _mm512_loadu_ps(b.as_ptr().add(base + 16));
2033 d1 = _mm512_fmadd_ps(_mm512_loadu_ps(a.as_ptr().add(base + 16)), vb1, d1);
2034 n1 = _mm512_fmadd_ps(vb1, vb1, n1);
2035 let vb2 = _mm512_loadu_ps(b.as_ptr().add(base + 32));
2036 d2 = _mm512_fmadd_ps(_mm512_loadu_ps(a.as_ptr().add(base + 32)), vb2, d2);
2037 n2 = _mm512_fmadd_ps(vb2, vb2, n2);
2038 let vb3 = _mm512_loadu_ps(b.as_ptr().add(base + 48));
2039 d3 = _mm512_fmadd_ps(_mm512_loadu_ps(a.as_ptr().add(base + 48)), vb3, d3);
2040 n3 = _mm512_fmadd_ps(vb3, vb3, n3);
2041 }
2042
2043 let acc_dot = _mm512_add_ps(_mm512_add_ps(d0, d1), _mm512_add_ps(d2, d3));
2044 let acc_norm = _mm512_add_ps(_mm512_add_ps(n0, n1), _mm512_add_ps(n2, n3));
2045 let mut dot = _mm512_reduce_add_ps(acc_dot);
2046 let mut norm = _mm512_reduce_add_ps(acc_norm);
2047
2048 let base = chunks64 * 64;
2049 for i in 0..remainder {
2050 dot += a[base + i] * b[base + i];
2051 norm += b[base + i] * b[base + i];
2052 }
2053
2054 (dot, norm)
2055}
2056
2057#[inline]
2066fn fused_dot_norm(a: &[f32], b: &[f32], count: usize) -> (f32, f32) {
2067 #[cfg(target_arch = "aarch64")]
2068 {
2069 if neon::is_available() {
2070 return unsafe { fused_dot_norm_neon(a, b, count) };
2071 }
2072 }
2073
2074 #[cfg(target_arch = "x86_64")]
2075 {
2076 if is_x86_feature_detected!("avx512f") {
2077 return unsafe { fused_dot_norm_avx512(a, b, count) };
2078 }
2079 if is_x86_feature_detected!("avx2") && is_x86_feature_detected!("fma") {
2080 return unsafe { fused_dot_norm_avx2(a, b, count) };
2081 }
2082 if sse::is_available() {
2083 return unsafe { fused_dot_norm_sse(a, b, count) };
2084 }
2085 }
2086
2087 let mut dot = 0.0f32;
2089 let mut norm_b = 0.0f32;
2090 for i in 0..count {
2091 dot += a[i] * b[i];
2092 norm_b += b[i] * b[i];
2093 }
2094 (dot, norm_b)
2095}
2096
2097#[cfg(target_arch = "aarch64")]
2098#[target_feature(enable = "neon")]
2099#[allow(unsafe_op_in_unsafe_fn)]
2100unsafe fn fused_dot_norm_neon(a: &[f32], b: &[f32], count: usize) -> (f32, f32) {
2101 use std::arch::aarch64::*;
2102
2103 let chunks16 = count / 16;
2104 let remainder = count % 16;
2105
2106 let mut d0 = vdupq_n_f32(0.0);
2107 let mut d1 = vdupq_n_f32(0.0);
2108 let mut d2 = vdupq_n_f32(0.0);
2109 let mut d3 = vdupq_n_f32(0.0);
2110 let mut n0 = vdupq_n_f32(0.0);
2111 let mut n1 = vdupq_n_f32(0.0);
2112 let mut n2 = vdupq_n_f32(0.0);
2113 let mut n3 = vdupq_n_f32(0.0);
2114
2115 for c in 0..chunks16 {
2116 let base = c * 16;
2117 let va0 = vld1q_f32(a.as_ptr().add(base));
2118 let vb0 = vld1q_f32(b.as_ptr().add(base));
2119 d0 = vfmaq_f32(d0, va0, vb0);
2120 n0 = vfmaq_f32(n0, vb0, vb0);
2121 let va1 = vld1q_f32(a.as_ptr().add(base + 4));
2122 let vb1 = vld1q_f32(b.as_ptr().add(base + 4));
2123 d1 = vfmaq_f32(d1, va1, vb1);
2124 n1 = vfmaq_f32(n1, vb1, vb1);
2125 let va2 = vld1q_f32(a.as_ptr().add(base + 8));
2126 let vb2 = vld1q_f32(b.as_ptr().add(base + 8));
2127 d2 = vfmaq_f32(d2, va2, vb2);
2128 n2 = vfmaq_f32(n2, vb2, vb2);
2129 let va3 = vld1q_f32(a.as_ptr().add(base + 12));
2130 let vb3 = vld1q_f32(b.as_ptr().add(base + 12));
2131 d3 = vfmaq_f32(d3, va3, vb3);
2132 n3 = vfmaq_f32(n3, vb3, vb3);
2133 }
2134
2135 let acc_dot = vaddq_f32(vaddq_f32(d0, d1), vaddq_f32(d2, d3));
2136 let acc_norm = vaddq_f32(vaddq_f32(n0, n1), vaddq_f32(n2, n3));
2137 let mut dot = vaddvq_f32(acc_dot);
2138 let mut norm = vaddvq_f32(acc_norm);
2139
2140 let base = chunks16 * 16;
2141 for i in 0..remainder {
2142 dot += a[base + i] * b[base + i];
2143 norm += b[base + i] * b[base + i];
2144 }
2145
2146 (dot, norm)
2147}
2148
2149#[cfg(target_arch = "x86_64")]
2150#[target_feature(enable = "avx2", enable = "fma")]
2151#[allow(unsafe_op_in_unsafe_fn)]
2152unsafe fn fused_dot_norm_avx2(a: &[f32], b: &[f32], count: usize) -> (f32, f32) {
2153 use std::arch::x86_64::*;
2154
2155 let chunks32 = count / 32;
2156 let remainder = count % 32;
2157
2158 let mut d0 = _mm256_setzero_ps();
2159 let mut d1 = _mm256_setzero_ps();
2160 let mut d2 = _mm256_setzero_ps();
2161 let mut d3 = _mm256_setzero_ps();
2162 let mut n0 = _mm256_setzero_ps();
2163 let mut n1 = _mm256_setzero_ps();
2164 let mut n2 = _mm256_setzero_ps();
2165 let mut n3 = _mm256_setzero_ps();
2166
2167 for c in 0..chunks32 {
2168 let base = c * 32;
2169 let vb0 = _mm256_loadu_ps(b.as_ptr().add(base));
2170 d0 = _mm256_fmadd_ps(_mm256_loadu_ps(a.as_ptr().add(base)), vb0, d0);
2171 n0 = _mm256_fmadd_ps(vb0, vb0, n0);
2172 let vb1 = _mm256_loadu_ps(b.as_ptr().add(base + 8));
2173 d1 = _mm256_fmadd_ps(_mm256_loadu_ps(a.as_ptr().add(base + 8)), vb1, d1);
2174 n1 = _mm256_fmadd_ps(vb1, vb1, n1);
2175 let vb2 = _mm256_loadu_ps(b.as_ptr().add(base + 16));
2176 d2 = _mm256_fmadd_ps(_mm256_loadu_ps(a.as_ptr().add(base + 16)), vb2, d2);
2177 n2 = _mm256_fmadd_ps(vb2, vb2, n2);
2178 let vb3 = _mm256_loadu_ps(b.as_ptr().add(base + 24));
2179 d3 = _mm256_fmadd_ps(_mm256_loadu_ps(a.as_ptr().add(base + 24)), vb3, d3);
2180 n3 = _mm256_fmadd_ps(vb3, vb3, n3);
2181 }
2182
2183 let acc_dot = _mm256_add_ps(_mm256_add_ps(d0, d1), _mm256_add_ps(d2, d3));
2184 let acc_norm = _mm256_add_ps(_mm256_add_ps(n0, n1), _mm256_add_ps(n2, n3));
2185
2186 let hi_d = _mm256_extractf128_ps(acc_dot, 1);
2188 let lo_d = _mm256_castps256_ps128(acc_dot);
2189 let sum_d = _mm_add_ps(lo_d, hi_d);
2190 let shuf_d = _mm_shuffle_ps(sum_d, sum_d, 0b10_11_00_01);
2191 let sums_d = _mm_add_ps(sum_d, shuf_d);
2192 let shuf2_d = _mm_movehl_ps(sums_d, sums_d);
2193 let mut dot = _mm_cvtss_f32(_mm_add_ss(sums_d, shuf2_d));
2194
2195 let hi_n = _mm256_extractf128_ps(acc_norm, 1);
2196 let lo_n = _mm256_castps256_ps128(acc_norm);
2197 let sum_n = _mm_add_ps(lo_n, hi_n);
2198 let shuf_n = _mm_shuffle_ps(sum_n, sum_n, 0b10_11_00_01);
2199 let sums_n = _mm_add_ps(sum_n, shuf_n);
2200 let shuf2_n = _mm_movehl_ps(sums_n, sums_n);
2201 let mut norm = _mm_cvtss_f32(_mm_add_ss(sums_n, shuf2_n));
2202
2203 let base = chunks32 * 32;
2204 for i in 0..remainder {
2205 dot += a[base + i] * b[base + i];
2206 norm += b[base + i] * b[base + i];
2207 }
2208
2209 (dot, norm)
2210}
2211
2212#[cfg(target_arch = "x86_64")]
2213#[target_feature(enable = "sse")]
2214#[allow(unsafe_op_in_unsafe_fn)]
2215unsafe fn fused_dot_norm_sse(a: &[f32], b: &[f32], count: usize) -> (f32, f32) {
2216 use std::arch::x86_64::*;
2217
2218 let chunks = count / 4;
2219 let remainder = count % 4;
2220
2221 let mut acc_dot = _mm_setzero_ps();
2222 let mut acc_norm = _mm_setzero_ps();
2223
2224 for chunk in 0..chunks {
2225 let base = chunk * 4;
2226 let va = _mm_loadu_ps(a.as_ptr().add(base));
2227 let vb = _mm_loadu_ps(b.as_ptr().add(base));
2228 acc_dot = _mm_add_ps(acc_dot, _mm_mul_ps(va, vb));
2229 acc_norm = _mm_add_ps(acc_norm, _mm_mul_ps(vb, vb));
2230 }
2231
2232 let shuf_d = _mm_shuffle_ps(acc_dot, acc_dot, 0b10_11_00_01);
2234 let sums_d = _mm_add_ps(acc_dot, shuf_d);
2235 let shuf2_d = _mm_movehl_ps(sums_d, sums_d);
2236 let final_d = _mm_add_ss(sums_d, shuf2_d);
2237 let mut dot = _mm_cvtss_f32(final_d);
2238
2239 let shuf_n = _mm_shuffle_ps(acc_norm, acc_norm, 0b10_11_00_01);
2240 let sums_n = _mm_add_ps(acc_norm, shuf_n);
2241 let shuf2_n = _mm_movehl_ps(sums_n, sums_n);
2242 let final_n = _mm_add_ss(sums_n, shuf2_n);
2243 let mut norm = _mm_cvtss_f32(final_n);
2244
2245 let base = chunks * 4;
2246 for i in 0..remainder {
2247 dot += a[base + i] * b[base + i];
2248 norm += b[base + i] * b[base + i];
2249 }
2250
2251 (dot, norm)
2252}
2253
2254#[inline]
2260pub fn fast_inv_sqrt(x: f32) -> f32 {
2261 let half = 0.5 * x;
2262 let i = 0x5F37_5A86_u32.wrapping_sub(x.to_bits() >> 1);
2263 let y = f32::from_bits(i);
2264 let y = y * (1.5 - half * y * y); y * (1.5 - half * y * y) }
2267
2268#[inline]
2279pub fn batch_cosine_scores(query: &[f32], vectors: &[f32], dim: usize, scores: &mut [f32]) {
2280 let n = scores.len();
2281 let required = n
2282 .checked_mul(dim)
2283 .expect("batch cosine vector length overflow");
2284 assert_eq!(query.len(), dim, "batch cosine query dimension mismatch");
2285 assert!(
2286 vectors.len() >= required,
2287 "batch cosine vectors are truncated: need {required}, got {}",
2288 vectors.len()
2289 );
2290
2291 if dim == 0 || n == 0 {
2292 return;
2293 }
2294
2295 let norm_q_sq = dot_product_f32(query, query, dim);
2297 if norm_q_sq < f32::EPSILON {
2298 for s in scores.iter_mut() {
2299 *s = 0.0;
2300 }
2301 return;
2302 }
2303 let inv_norm_q = fast_inv_sqrt(norm_q_sq);
2304
2305 for i in 0..n {
2306 let vec = &vectors[i * dim..(i + 1) * dim];
2307 let (dot, norm_v_sq) = fused_dot_norm(query, vec, dim);
2308 if norm_v_sq < f32::EPSILON {
2309 scores[i] = 0.0;
2310 } else {
2311 scores[i] = dot * inv_norm_q * fast_inv_sqrt(norm_v_sq);
2312 }
2313 }
2314}
2315
2316#[inline]
2322pub fn f32_to_f16(value: f32) -> u16 {
2323 let bits = value.to_bits();
2324 let sign = (bits >> 16) & 0x8000;
2325 let exp = ((bits >> 23) & 0xFF) as i32;
2326 let mantissa = bits & 0x7F_FFFF;
2327
2328 if exp == 255 {
2329 return (sign | 0x7C00 | ((mantissa >> 13) & 0x3FF)) as u16;
2331 }
2332
2333 let exp16 = exp - 127 + 15;
2334
2335 if exp16 >= 31 {
2336 return (sign | 0x7C00) as u16; }
2338
2339 if exp16 <= 0 {
2340 if exp16 < -10 {
2341 return sign as u16; }
2343 let shift = (1 - exp16) as u32;
2344 let m = (mantissa | 0x80_0000) >> shift;
2345 let round_bit = (m >> 12) & 1;
2347 let sticky = m & 0xFFF;
2348 let m13 = m >> 13;
2349 let rounded = m13 + (round_bit & (m13 | if sticky != 0 { 1 } else { 0 }));
2350 return (sign | rounded) as u16;
2351 }
2352
2353 let round_bit = (mantissa >> 12) & 1;
2355 let sticky = mantissa & 0xFFF;
2356 let m13 = mantissa >> 13;
2357 let rounded = m13 + (round_bit & (m13 | if sticky != 0 { 1 } else { 0 }));
2358 if rounded > 0x3FF {
2360 let exp16_inc = exp16 as u32 + 1;
2361 if exp16_inc >= 31 {
2362 return (sign | 0x7C00) as u16; }
2364 (sign | (exp16_inc << 10)) as u16
2365 } else {
2366 (sign | ((exp16 as u32) << 10) | rounded) as u16
2367 }
2368}
2369
2370#[inline]
2372pub fn f16_to_f32(half: u16) -> f32 {
2373 let sign = ((half & 0x8000) as u32) << 16;
2374 let exp = ((half >> 10) & 0x1F) as u32;
2375 let mantissa = (half & 0x3FF) as u32;
2376
2377 if exp == 0 {
2378 if mantissa == 0 {
2379 return f32::from_bits(sign);
2380 }
2381 let mut e = 0u32;
2383 let mut m = mantissa;
2384 while (m & 0x400) == 0 {
2385 m <<= 1;
2386 e += 1;
2387 }
2388 return f32::from_bits(sign | ((127 - 15 + 1 - e) << 23) | ((m & 0x3FF) << 13));
2389 }
2390
2391 if exp == 31 {
2392 return f32::from_bits(sign | 0x7F80_0000 | (mantissa << 13));
2393 }
2394
2395 f32::from_bits(sign | ((exp + 127 - 15) << 23) | (mantissa << 13))
2396}
2397
2398const U8_SCALE: f32 = 127.5;
2403const U8_INV_SCALE: f32 = 1.0 / 127.5;
2404
2405#[inline]
2407pub fn f32_to_u8_saturating(value: f32) -> u8 {
2408 ((value.clamp(-1.0, 1.0) + 1.0) * U8_SCALE) as u8
2409}
2410
2411#[inline]
2413pub fn u8_to_f32(byte: u8) -> f32 {
2414 byte as f32 * U8_INV_SCALE - 1.0
2415}
2416
2417pub fn batch_f32_to_f16(src: &[f32], dst: &mut [u16]) {
2423 debug_assert_eq!(src.len(), dst.len());
2424 for (s, d) in src.iter().zip(dst.iter_mut()) {
2425 *d = f32_to_f16(*s);
2426 }
2427}
2428
2429pub fn batch_f32_to_u8(src: &[f32], dst: &mut [u8]) {
2431 debug_assert_eq!(src.len(), dst.len());
2432 for (s, d) in src.iter().zip(dst.iter_mut()) {
2433 *d = f32_to_u8_saturating(*s);
2434 }
2435}
2436
2437#[cfg(target_arch = "aarch64")]
2442#[allow(unsafe_op_in_unsafe_fn)]
2443mod neon_quant {
2444 use std::arch::aarch64::*;
2445
2446 #[allow(clippy::incompatible_msrv)]
2452 #[target_feature(enable = "neon")]
2453 pub unsafe fn fused_dot_norm_f16(query_f16: &[u16], vec_f16: &[u16], dim: usize) -> (f32, f32) {
2454 let chunks16 = dim / 16;
2455 let remainder = dim % 16;
2456
2457 let mut acc_dot0 = vdupq_n_f32(0.0);
2459 let mut acc_dot1 = vdupq_n_f32(0.0);
2460 let mut acc_norm0 = vdupq_n_f32(0.0);
2461 let mut acc_norm1 = vdupq_n_f32(0.0);
2462
2463 for c in 0..chunks16 {
2464 let base = c * 16;
2465
2466 let v_raw0 = vld1q_u16(vec_f16.as_ptr().add(base));
2468 let v_lo0 = vcvt_f32_f16(vreinterpret_f16_u16(vget_low_u16(v_raw0)));
2469 let v_hi0 = vcvt_f32_f16(vreinterpret_f16_u16(vget_high_u16(v_raw0)));
2470 let q_raw0 = vld1q_u16(query_f16.as_ptr().add(base));
2471 let q_lo0 = vcvt_f32_f16(vreinterpret_f16_u16(vget_low_u16(q_raw0)));
2472 let q_hi0 = vcvt_f32_f16(vreinterpret_f16_u16(vget_high_u16(q_raw0)));
2473
2474 acc_dot0 = vfmaq_f32(acc_dot0, q_lo0, v_lo0);
2475 acc_dot0 = vfmaq_f32(acc_dot0, q_hi0, v_hi0);
2476 acc_norm0 = vfmaq_f32(acc_norm0, v_lo0, v_lo0);
2477 acc_norm0 = vfmaq_f32(acc_norm0, v_hi0, v_hi0);
2478
2479 let v_raw1 = vld1q_u16(vec_f16.as_ptr().add(base + 8));
2481 let v_lo1 = vcvt_f32_f16(vreinterpret_f16_u16(vget_low_u16(v_raw1)));
2482 let v_hi1 = vcvt_f32_f16(vreinterpret_f16_u16(vget_high_u16(v_raw1)));
2483 let q_raw1 = vld1q_u16(query_f16.as_ptr().add(base + 8));
2484 let q_lo1 = vcvt_f32_f16(vreinterpret_f16_u16(vget_low_u16(q_raw1)));
2485 let q_hi1 = vcvt_f32_f16(vreinterpret_f16_u16(vget_high_u16(q_raw1)));
2486
2487 acc_dot1 = vfmaq_f32(acc_dot1, q_lo1, v_lo1);
2488 acc_dot1 = vfmaq_f32(acc_dot1, q_hi1, v_hi1);
2489 acc_norm1 = vfmaq_f32(acc_norm1, v_lo1, v_lo1);
2490 acc_norm1 = vfmaq_f32(acc_norm1, v_hi1, v_hi1);
2491 }
2492
2493 let mut dot = vaddvq_f32(vaddq_f32(acc_dot0, acc_dot1));
2495 let mut norm = vaddvq_f32(vaddq_f32(acc_norm0, acc_norm1));
2496
2497 let base = chunks16 * 16;
2499 for i in 0..remainder {
2500 let v = super::f16_to_f32(*vec_f16.get_unchecked(base + i));
2501 let q = super::f16_to_f32(*query_f16.get_unchecked(base + i));
2502 dot += q * v;
2503 norm += v * v;
2504 }
2505
2506 (dot, norm)
2507 }
2508
2509 #[target_feature(enable = "neon")]
2512 pub unsafe fn fused_dot_norm_u8(query: &[f32], vec_u8: &[u8], dim: usize) -> (f32, f32) {
2513 let scale = vdupq_n_f32(super::U8_INV_SCALE);
2514 let offset = vdupq_n_f32(-1.0);
2515
2516 let chunks16 = dim / 16;
2517 let remainder = dim % 16;
2518
2519 let mut acc_dot = vdupq_n_f32(0.0);
2520 let mut acc_norm = vdupq_n_f32(0.0);
2521
2522 for c in 0..chunks16 {
2523 let base = c * 16;
2524
2525 let bytes = vld1q_u8(vec_u8.as_ptr().add(base));
2527
2528 let lo8 = vget_low_u8(bytes);
2530 let hi8 = vget_high_u8(bytes);
2531 let lo16 = vmovl_u8(lo8);
2532 let hi16 = vmovl_u8(hi8);
2533
2534 let f0 = vaddq_f32(
2535 vmulq_f32(vcvtq_f32_u32(vmovl_u16(vget_low_u16(lo16))), scale),
2536 offset,
2537 );
2538 let f1 = vaddq_f32(
2539 vmulq_f32(vcvtq_f32_u32(vmovl_u16(vget_high_u16(lo16))), scale),
2540 offset,
2541 );
2542 let f2 = vaddq_f32(
2543 vmulq_f32(vcvtq_f32_u32(vmovl_u16(vget_low_u16(hi16))), scale),
2544 offset,
2545 );
2546 let f3 = vaddq_f32(
2547 vmulq_f32(vcvtq_f32_u32(vmovl_u16(vget_high_u16(hi16))), scale),
2548 offset,
2549 );
2550
2551 let q0 = vld1q_f32(query.as_ptr().add(base));
2552 let q1 = vld1q_f32(query.as_ptr().add(base + 4));
2553 let q2 = vld1q_f32(query.as_ptr().add(base + 8));
2554 let q3 = vld1q_f32(query.as_ptr().add(base + 12));
2555
2556 acc_dot = vfmaq_f32(acc_dot, q0, f0);
2557 acc_dot = vfmaq_f32(acc_dot, q1, f1);
2558 acc_dot = vfmaq_f32(acc_dot, q2, f2);
2559 acc_dot = vfmaq_f32(acc_dot, q3, f3);
2560
2561 acc_norm = vfmaq_f32(acc_norm, f0, f0);
2562 acc_norm = vfmaq_f32(acc_norm, f1, f1);
2563 acc_norm = vfmaq_f32(acc_norm, f2, f2);
2564 acc_norm = vfmaq_f32(acc_norm, f3, f3);
2565 }
2566
2567 let mut dot = vaddvq_f32(acc_dot);
2568 let mut norm = vaddvq_f32(acc_norm);
2569
2570 let base = chunks16 * 16;
2571 for i in 0..remainder {
2572 let v = super::u8_to_f32(*vec_u8.get_unchecked(base + i));
2573 dot += *query.get_unchecked(base + i) * v;
2574 norm += v * v;
2575 }
2576
2577 (dot, norm)
2578 }
2579
2580 #[allow(clippy::incompatible_msrv)]
2582 #[target_feature(enable = "neon")]
2583 pub unsafe fn dot_product_f16(query_f16: &[u16], vec_f16: &[u16], dim: usize) -> f32 {
2584 let chunks8 = dim / 8;
2585 let remainder = dim % 8;
2586
2587 let mut acc = vdupq_n_f32(0.0);
2588
2589 for c in 0..chunks8 {
2590 let base = c * 8;
2591 let v_raw = vld1q_u16(vec_f16.as_ptr().add(base));
2592 let v_lo = vcvt_f32_f16(vreinterpret_f16_u16(vget_low_u16(v_raw)));
2593 let v_hi = vcvt_f32_f16(vreinterpret_f16_u16(vget_high_u16(v_raw)));
2594 let q_raw = vld1q_u16(query_f16.as_ptr().add(base));
2595 let q_lo = vcvt_f32_f16(vreinterpret_f16_u16(vget_low_u16(q_raw)));
2596 let q_hi = vcvt_f32_f16(vreinterpret_f16_u16(vget_high_u16(q_raw)));
2597 acc = vfmaq_f32(acc, q_lo, v_lo);
2598 acc = vfmaq_f32(acc, q_hi, v_hi);
2599 }
2600
2601 let mut dot = vaddvq_f32(acc);
2602 let base = chunks8 * 8;
2603 for i in 0..remainder {
2604 let v = super::f16_to_f32(*vec_f16.get_unchecked(base + i));
2605 let q = super::f16_to_f32(*query_f16.get_unchecked(base + i));
2606 dot += q * v;
2607 }
2608 dot
2609 }
2610
2611 #[target_feature(enable = "neon")]
2613 pub unsafe fn dot_product_u8(query: &[f32], vec_u8: &[u8], dim: usize) -> f32 {
2614 let scale = vdupq_n_f32(super::U8_INV_SCALE);
2615 let offset = vdupq_n_f32(-1.0);
2616 let chunks16 = dim / 16;
2617 let remainder = dim % 16;
2618
2619 let mut acc = vdupq_n_f32(0.0);
2620
2621 for c in 0..chunks16 {
2622 let base = c * 16;
2623 let bytes = vld1q_u8(vec_u8.as_ptr().add(base));
2624 let lo8 = vget_low_u8(bytes);
2625 let hi8 = vget_high_u8(bytes);
2626 let lo16 = vmovl_u8(lo8);
2627 let hi16 = vmovl_u8(hi8);
2628 let f0 = vaddq_f32(
2629 vmulq_f32(vcvtq_f32_u32(vmovl_u16(vget_low_u16(lo16))), scale),
2630 offset,
2631 );
2632 let f1 = vaddq_f32(
2633 vmulq_f32(vcvtq_f32_u32(vmovl_u16(vget_high_u16(lo16))), scale),
2634 offset,
2635 );
2636 let f2 = vaddq_f32(
2637 vmulq_f32(vcvtq_f32_u32(vmovl_u16(vget_low_u16(hi16))), scale),
2638 offset,
2639 );
2640 let f3 = vaddq_f32(
2641 vmulq_f32(vcvtq_f32_u32(vmovl_u16(vget_high_u16(hi16))), scale),
2642 offset,
2643 );
2644 let q0 = vld1q_f32(query.as_ptr().add(base));
2645 let q1 = vld1q_f32(query.as_ptr().add(base + 4));
2646 let q2 = vld1q_f32(query.as_ptr().add(base + 8));
2647 let q3 = vld1q_f32(query.as_ptr().add(base + 12));
2648 acc = vfmaq_f32(acc, q0, f0);
2649 acc = vfmaq_f32(acc, q1, f1);
2650 acc = vfmaq_f32(acc, q2, f2);
2651 acc = vfmaq_f32(acc, q3, f3);
2652 }
2653
2654 let mut dot = vaddvq_f32(acc);
2655 let base = chunks16 * 16;
2656 for i in 0..remainder {
2657 let v = super::u8_to_f32(*vec_u8.get_unchecked(base + i));
2658 dot += *query.get_unchecked(base + i) * v;
2659 }
2660 dot
2661 }
2662}
2663
2664#[allow(dead_code)]
2669fn fused_dot_norm_f16_scalar(query_f16: &[u16], vec_f16: &[u16], dim: usize) -> (f32, f32) {
2670 let mut dot = 0.0f32;
2671 let mut norm = 0.0f32;
2672 for i in 0..dim {
2673 let v = f16_to_f32(vec_f16[i]);
2674 let q = f16_to_f32(query_f16[i]);
2675 dot += q * v;
2676 norm += v * v;
2677 }
2678 (dot, norm)
2679}
2680
2681#[allow(dead_code)]
2682fn fused_dot_norm_u8_scalar(query: &[f32], vec_u8: &[u8], dim: usize) -> (f32, f32) {
2683 let mut dot = 0.0f32;
2684 let mut norm = 0.0f32;
2685 for i in 0..dim {
2686 let v = u8_to_f32(vec_u8[i]);
2687 dot += query[i] * v;
2688 norm += v * v;
2689 }
2690 (dot, norm)
2691}
2692
2693#[allow(dead_code)]
2694fn dot_product_f16_scalar(query_f16: &[u16], vec_f16: &[u16], dim: usize) -> f32 {
2695 let mut dot = 0.0f32;
2696 for i in 0..dim {
2697 dot += f16_to_f32(query_f16[i]) * f16_to_f32(vec_f16[i]);
2698 }
2699 dot
2700}
2701
2702#[allow(dead_code)]
2703fn dot_product_u8_scalar(query: &[f32], vec_u8: &[u8], dim: usize) -> f32 {
2704 let mut dot = 0.0f32;
2705 for i in 0..dim {
2706 dot += query[i] * u8_to_f32(vec_u8[i]);
2707 }
2708 dot
2709}
2710
2711#[cfg(target_arch = "x86_64")]
2716#[target_feature(enable = "sse2", enable = "sse4.1")]
2717#[allow(unsafe_op_in_unsafe_fn)]
2718unsafe fn fused_dot_norm_f16_sse(query_f16: &[u16], vec_f16: &[u16], dim: usize) -> (f32, f32) {
2719 use std::arch::x86_64::*;
2720
2721 let chunks = dim / 4;
2722 let remainder = dim % 4;
2723
2724 let mut acc_dot = _mm_setzero_ps();
2725 let mut acc_norm = _mm_setzero_ps();
2726
2727 for chunk in 0..chunks {
2728 let base = chunk * 4;
2729 let v0 = f16_to_f32(*vec_f16.get_unchecked(base));
2731 let v1 = f16_to_f32(*vec_f16.get_unchecked(base + 1));
2732 let v2 = f16_to_f32(*vec_f16.get_unchecked(base + 2));
2733 let v3 = f16_to_f32(*vec_f16.get_unchecked(base + 3));
2734 let vb = _mm_set_ps(v3, v2, v1, v0);
2735
2736 let q0 = f16_to_f32(*query_f16.get_unchecked(base));
2737 let q1 = f16_to_f32(*query_f16.get_unchecked(base + 1));
2738 let q2 = f16_to_f32(*query_f16.get_unchecked(base + 2));
2739 let q3 = f16_to_f32(*query_f16.get_unchecked(base + 3));
2740 let va = _mm_set_ps(q3, q2, q1, q0);
2741
2742 acc_dot = _mm_add_ps(acc_dot, _mm_mul_ps(va, vb));
2743 acc_norm = _mm_add_ps(acc_norm, _mm_mul_ps(vb, vb));
2744 }
2745
2746 let shuf_d = _mm_shuffle_ps(acc_dot, acc_dot, 0b10_11_00_01);
2748 let sums_d = _mm_add_ps(acc_dot, shuf_d);
2749 let shuf2_d = _mm_movehl_ps(sums_d, sums_d);
2750 let mut dot = _mm_cvtss_f32(_mm_add_ss(sums_d, shuf2_d));
2751
2752 let shuf_n = _mm_shuffle_ps(acc_norm, acc_norm, 0b10_11_00_01);
2753 let sums_n = _mm_add_ps(acc_norm, shuf_n);
2754 let shuf2_n = _mm_movehl_ps(sums_n, sums_n);
2755 let mut norm = _mm_cvtss_f32(_mm_add_ss(sums_n, shuf2_n));
2756
2757 let base = chunks * 4;
2758 for i in 0..remainder {
2759 let v = f16_to_f32(*vec_f16.get_unchecked(base + i));
2760 let q = f16_to_f32(*query_f16.get_unchecked(base + i));
2761 dot += q * v;
2762 norm += v * v;
2763 }
2764
2765 (dot, norm)
2766}
2767
2768#[cfg(target_arch = "x86_64")]
2769#[target_feature(enable = "sse2", enable = "sse4.1")]
2770#[allow(unsafe_op_in_unsafe_fn)]
2771unsafe fn fused_dot_norm_u8_sse(query: &[f32], vec_u8: &[u8], dim: usize) -> (f32, f32) {
2772 use std::arch::x86_64::*;
2773
2774 let scale = _mm_set1_ps(U8_INV_SCALE);
2775 let offset = _mm_set1_ps(-1.0);
2776
2777 let chunks = dim / 4;
2778 let remainder = dim % 4;
2779
2780 let mut acc_dot = _mm_setzero_ps();
2781 let mut acc_norm = _mm_setzero_ps();
2782
2783 for chunk in 0..chunks {
2784 let base = chunk * 4;
2785
2786 let bytes = _mm_cvtsi32_si128(std::ptr::read_unaligned(
2788 vec_u8.as_ptr().add(base) as *const i32
2789 ));
2790 let ints = _mm_cvtepu8_epi32(bytes);
2791 let floats = _mm_cvtepi32_ps(ints);
2792 let vb = _mm_add_ps(_mm_mul_ps(floats, scale), offset);
2793
2794 let va = _mm_loadu_ps(query.as_ptr().add(base));
2795
2796 acc_dot = _mm_add_ps(acc_dot, _mm_mul_ps(va, vb));
2797 acc_norm = _mm_add_ps(acc_norm, _mm_mul_ps(vb, vb));
2798 }
2799
2800 let shuf_d = _mm_shuffle_ps(acc_dot, acc_dot, 0b10_11_00_01);
2802 let sums_d = _mm_add_ps(acc_dot, shuf_d);
2803 let shuf2_d = _mm_movehl_ps(sums_d, sums_d);
2804 let mut dot = _mm_cvtss_f32(_mm_add_ss(sums_d, shuf2_d));
2805
2806 let shuf_n = _mm_shuffle_ps(acc_norm, acc_norm, 0b10_11_00_01);
2807 let sums_n = _mm_add_ps(acc_norm, shuf_n);
2808 let shuf2_n = _mm_movehl_ps(sums_n, sums_n);
2809 let mut norm = _mm_cvtss_f32(_mm_add_ss(sums_n, shuf2_n));
2810
2811 let base = chunks * 4;
2812 for i in 0..remainder {
2813 let v = u8_to_f32(*vec_u8.get_unchecked(base + i));
2814 dot += *query.get_unchecked(base + i) * v;
2815 norm += v * v;
2816 }
2817
2818 (dot, norm)
2819}
2820
2821#[cfg(target_arch = "x86_64")]
2826#[target_feature(enable = "avx", enable = "f16c", enable = "fma")]
2827#[allow(unsafe_op_in_unsafe_fn)]
2828unsafe fn fused_dot_norm_f16_f16c(query_f16: &[u16], vec_f16: &[u16], dim: usize) -> (f32, f32) {
2829 use std::arch::x86_64::*;
2830
2831 let chunks16 = dim / 16;
2832 let remainder = dim % 16;
2833
2834 let mut acc_dot0 = _mm256_setzero_ps();
2836 let mut acc_dot1 = _mm256_setzero_ps();
2837 let mut acc_norm0 = _mm256_setzero_ps();
2838 let mut acc_norm1 = _mm256_setzero_ps();
2839
2840 for c in 0..chunks16 {
2841 let base = c * 16;
2842
2843 let v_raw0 = _mm_loadu_si128(vec_f16.as_ptr().add(base) as *const __m128i);
2845 let vb0 = _mm256_cvtph_ps(v_raw0);
2846 let q_raw0 = _mm_loadu_si128(query_f16.as_ptr().add(base) as *const __m128i);
2847 let qa0 = _mm256_cvtph_ps(q_raw0);
2848 acc_dot0 = _mm256_fmadd_ps(qa0, vb0, acc_dot0);
2849 acc_norm0 = _mm256_fmadd_ps(vb0, vb0, acc_norm0);
2850
2851 let v_raw1 = _mm_loadu_si128(vec_f16.as_ptr().add(base + 8) as *const __m128i);
2853 let vb1 = _mm256_cvtph_ps(v_raw1);
2854 let q_raw1 = _mm_loadu_si128(query_f16.as_ptr().add(base + 8) as *const __m128i);
2855 let qa1 = _mm256_cvtph_ps(q_raw1);
2856 acc_dot1 = _mm256_fmadd_ps(qa1, vb1, acc_dot1);
2857 acc_norm1 = _mm256_fmadd_ps(vb1, vb1, acc_norm1);
2858 }
2859
2860 let acc_dot = _mm256_add_ps(acc_dot0, acc_dot1);
2862 let acc_norm = _mm256_add_ps(acc_norm0, acc_norm1);
2863
2864 let hi_d = _mm256_extractf128_ps(acc_dot, 1);
2866 let lo_d = _mm256_castps256_ps128(acc_dot);
2867 let sum_d = _mm_add_ps(lo_d, hi_d);
2868 let shuf_d = _mm_shuffle_ps(sum_d, sum_d, 0b10_11_00_01);
2869 let sums_d = _mm_add_ps(sum_d, shuf_d);
2870 let shuf2_d = _mm_movehl_ps(sums_d, sums_d);
2871 let mut dot = _mm_cvtss_f32(_mm_add_ss(sums_d, shuf2_d));
2872
2873 let hi_n = _mm256_extractf128_ps(acc_norm, 1);
2874 let lo_n = _mm256_castps256_ps128(acc_norm);
2875 let sum_n = _mm_add_ps(lo_n, hi_n);
2876 let shuf_n = _mm_shuffle_ps(sum_n, sum_n, 0b10_11_00_01);
2877 let sums_n = _mm_add_ps(sum_n, shuf_n);
2878 let shuf2_n = _mm_movehl_ps(sums_n, sums_n);
2879 let mut norm = _mm_cvtss_f32(_mm_add_ss(sums_n, shuf2_n));
2880
2881 let base = chunks16 * 16;
2882 for i in 0..remainder {
2883 let v = f16_to_f32(*vec_f16.get_unchecked(base + i));
2884 let q = f16_to_f32(*query_f16.get_unchecked(base + i));
2885 dot += q * v;
2886 norm += v * v;
2887 }
2888
2889 (dot, norm)
2890}
2891
2892#[cfg(target_arch = "x86_64")]
2893#[target_feature(enable = "avx", enable = "f16c", enable = "fma")]
2894#[allow(unsafe_op_in_unsafe_fn)]
2895unsafe fn dot_product_f16_f16c(query_f16: &[u16], vec_f16: &[u16], dim: usize) -> f32 {
2896 use std::arch::x86_64::*;
2897
2898 let chunks = dim / 8;
2899 let remainder = dim % 8;
2900 let mut acc = _mm256_setzero_ps();
2901
2902 for chunk in 0..chunks {
2903 let base = chunk * 8;
2904 let v_raw = _mm_loadu_si128(vec_f16.as_ptr().add(base) as *const __m128i);
2905 let vb = _mm256_cvtph_ps(v_raw);
2906 let q_raw = _mm_loadu_si128(query_f16.as_ptr().add(base) as *const __m128i);
2907 let qa = _mm256_cvtph_ps(q_raw);
2908 acc = _mm256_fmadd_ps(qa, vb, acc);
2909 }
2910
2911 let hi = _mm256_extractf128_ps(acc, 1);
2912 let lo = _mm256_castps256_ps128(acc);
2913 let sum = _mm_add_ps(lo, hi);
2914 let shuf = _mm_shuffle_ps(sum, sum, 0b10_11_00_01);
2915 let sums = _mm_add_ps(sum, shuf);
2916 let shuf2 = _mm_movehl_ps(sums, sums);
2917 let mut dot = _mm_cvtss_f32(_mm_add_ss(sums, shuf2));
2918
2919 let base = chunks * 8;
2920 for i in 0..remainder {
2921 let v = f16_to_f32(*vec_f16.get_unchecked(base + i));
2922 let q = f16_to_f32(*query_f16.get_unchecked(base + i));
2923 dot += q * v;
2924 }
2925 dot
2926}
2927
2928#[cfg(target_arch = "x86_64")]
2929#[target_feature(enable = "sse2", enable = "sse4.1")]
2930#[allow(unsafe_op_in_unsafe_fn)]
2931unsafe fn dot_product_u8_sse(query: &[f32], vec_u8: &[u8], dim: usize) -> f32 {
2932 use std::arch::x86_64::*;
2933
2934 let scale = _mm_set1_ps(U8_INV_SCALE);
2935 let offset = _mm_set1_ps(-1.0);
2936 let chunks = dim / 4;
2937 let remainder = dim % 4;
2938 let mut acc = _mm_setzero_ps();
2939
2940 for chunk in 0..chunks {
2941 let base = chunk * 4;
2942 let bytes = _mm_cvtsi32_si128(std::ptr::read_unaligned(
2943 vec_u8.as_ptr().add(base) as *const i32
2944 ));
2945 let ints = _mm_cvtepu8_epi32(bytes);
2946 let floats = _mm_cvtepi32_ps(ints);
2947 let vb = _mm_add_ps(_mm_mul_ps(floats, scale), offset);
2948 let va = _mm_loadu_ps(query.as_ptr().add(base));
2949 acc = _mm_add_ps(acc, _mm_mul_ps(va, vb));
2950 }
2951
2952 let shuf = _mm_shuffle_ps(acc, acc, 0b10_11_00_01);
2953 let sums = _mm_add_ps(acc, shuf);
2954 let shuf2 = _mm_movehl_ps(sums, sums);
2955 let mut dot = _mm_cvtss_f32(_mm_add_ss(sums, shuf2));
2956
2957 let base = chunks * 4;
2958 for i in 0..remainder {
2959 dot += *query.get_unchecked(base + i) * u8_to_f32(*vec_u8.get_unchecked(base + i));
2960 }
2961 dot
2962}
2963
2964#[inline]
2969fn fused_dot_norm_f16(query_f16: &[u16], vec_f16: &[u16], dim: usize) -> (f32, f32) {
2970 #[cfg(target_arch = "aarch64")]
2971 {
2972 return unsafe { neon_quant::fused_dot_norm_f16(query_f16, vec_f16, dim) };
2973 }
2974
2975 #[cfg(target_arch = "x86_64")]
2976 {
2977 if is_x86_feature_detected!("f16c") && is_x86_feature_detected!("fma") {
2978 return unsafe { fused_dot_norm_f16_f16c(query_f16, vec_f16, dim) };
2979 }
2980 if sse::is_available() {
2981 return unsafe { fused_dot_norm_f16_sse(query_f16, vec_f16, dim) };
2982 }
2983 }
2984
2985 #[allow(unreachable_code)]
2986 fused_dot_norm_f16_scalar(query_f16, vec_f16, dim)
2987}
2988
2989#[inline]
2990fn fused_dot_norm_u8(query: &[f32], vec_u8: &[u8], dim: usize) -> (f32, f32) {
2991 #[cfg(target_arch = "aarch64")]
2992 {
2993 return unsafe { neon_quant::fused_dot_norm_u8(query, vec_u8, dim) };
2994 }
2995
2996 #[cfg(target_arch = "x86_64")]
2997 {
2998 if sse::is_available() {
2999 return unsafe { fused_dot_norm_u8_sse(query, vec_u8, dim) };
3000 }
3001 }
3002
3003 #[allow(unreachable_code)]
3004 fused_dot_norm_u8_scalar(query, vec_u8, dim)
3005}
3006
3007#[inline]
3010fn dot_product_f16_quant(query_f16: &[u16], vec_f16: &[u16], dim: usize) -> f32 {
3011 #[cfg(target_arch = "aarch64")]
3012 {
3013 return unsafe { neon_quant::dot_product_f16(query_f16, vec_f16, dim) };
3014 }
3015
3016 #[cfg(target_arch = "x86_64")]
3017 {
3018 if is_x86_feature_detected!("f16c") && is_x86_feature_detected!("fma") {
3019 return unsafe { dot_product_f16_f16c(query_f16, vec_f16, dim) };
3020 }
3021 }
3022
3023 #[allow(unreachable_code)]
3024 dot_product_f16_scalar(query_f16, vec_f16, dim)
3025}
3026
3027#[inline]
3028fn dot_product_u8_quant(query: &[f32], vec_u8: &[u8], dim: usize) -> f32 {
3029 #[cfg(target_arch = "aarch64")]
3030 {
3031 return unsafe { neon_quant::dot_product_u8(query, vec_u8, dim) };
3032 }
3033
3034 #[cfg(target_arch = "x86_64")]
3035 {
3036 if sse::is_available() {
3037 return unsafe { dot_product_u8_sse(query, vec_u8, dim) };
3038 }
3039 }
3040
3041 #[allow(unreachable_code)]
3042 dot_product_u8_scalar(query, vec_u8, dim)
3043}
3044
3045#[inline]
3056pub fn batch_cosine_scores_f16(query: &[f32], vectors_raw: &[u8], dim: usize, scores: &mut [f32]) {
3057 let n = scores.len();
3058 let vec_bytes = dim.checked_mul(2).expect("f16 vector size overflow");
3059 let required = n
3060 .checked_mul(vec_bytes)
3061 .expect("f16 batch byte length overflow");
3062 assert_eq!(
3063 query.len(),
3064 dim,
3065 "f16 batch cosine query dimension mismatch"
3066 );
3067 assert!(
3068 vectors_raw.len() >= required,
3069 "f16 batch cosine vectors are truncated: need {required} bytes, got {}",
3070 vectors_raw.len()
3071 );
3072 if required > 0 {
3073 assert!(
3074 (vectors_raw.as_ptr() as usize).is_multiple_of(std::mem::align_of::<u16>()),
3075 "f16 batch cosine vectors are not 2-byte aligned"
3076 );
3077 }
3078 if dim == 0 || n == 0 {
3079 return;
3080 }
3081
3082 let norm_q_sq = dot_product_f32(query, query, dim);
3084 if norm_q_sq < f32::EPSILON {
3085 for s in scores.iter_mut() {
3086 *s = 0.0;
3087 }
3088 return;
3089 }
3090 let inv_norm_q = fast_inv_sqrt(norm_q_sq);
3091
3092 let query_f16: Vec<u16> = query.iter().map(|&v| f32_to_f16(v)).collect();
3094
3095 for i in 0..n {
3096 let raw = &vectors_raw[i * vec_bytes..(i + 1) * vec_bytes];
3097 let f16_slice = unsafe { std::slice::from_raw_parts(raw.as_ptr() as *const u16, dim) };
3098
3099 let (dot, norm_v_sq) = fused_dot_norm_f16(&query_f16, f16_slice, dim);
3100 scores[i] = if norm_v_sq < f32::EPSILON {
3101 0.0
3102 } else {
3103 dot * inv_norm_q * fast_inv_sqrt(norm_v_sq)
3104 };
3105 }
3106}
3107
3108#[inline]
3115pub fn batch_cosine_scores_u8(query: &[f32], vectors_raw: &[u8], dim: usize, scores: &mut [f32]) {
3116 let n = scores.len();
3117 let required = n.checked_mul(dim).expect("u8 batch byte length overflow");
3118 assert_eq!(query.len(), dim, "u8 batch cosine query dimension mismatch");
3119 assert!(
3120 vectors_raw.len() >= required,
3121 "u8 batch cosine vectors are truncated: need {required} bytes, got {}",
3122 vectors_raw.len()
3123 );
3124 if dim == 0 || n == 0 {
3125 return;
3126 }
3127
3128 let norm_q_sq = dot_product_f32(query, query, dim);
3129 if norm_q_sq < f32::EPSILON {
3130 for s in scores.iter_mut() {
3131 *s = 0.0;
3132 }
3133 return;
3134 }
3135 let inv_norm_q = fast_inv_sqrt(norm_q_sq);
3136
3137 for i in 0..n {
3138 let u8_slice = &vectors_raw[i * dim..(i + 1) * dim];
3139
3140 let (dot, norm_v_sq) = fused_dot_norm_u8(query, u8_slice, dim);
3141 scores[i] = if norm_v_sq < f32::EPSILON {
3142 0.0
3143 } else {
3144 dot * inv_norm_q * fast_inv_sqrt(norm_v_sq)
3145 };
3146 }
3147}
3148
3149#[inline]
3158pub fn batch_dot_scores(query: &[f32], vectors: &[f32], dim: usize, scores: &mut [f32]) {
3159 let n = scores.len();
3160 let required = n
3161 .checked_mul(dim)
3162 .expect("batch dot vector length overflow");
3163 assert_eq!(query.len(), dim, "batch dot query dimension mismatch");
3164 assert!(
3165 vectors.len() >= required,
3166 "batch dot vectors are truncated: need {required}, got {}",
3167 vectors.len()
3168 );
3169
3170 if dim == 0 || n == 0 {
3171 return;
3172 }
3173
3174 let norm_q_sq = dot_product_f32(query, query, dim);
3175 if norm_q_sq < f32::EPSILON {
3176 for s in scores.iter_mut() {
3177 *s = 0.0;
3178 }
3179 return;
3180 }
3181 let inv_norm_q = fast_inv_sqrt(norm_q_sq);
3182
3183 for i in 0..n {
3184 let vec = &vectors[i * dim..(i + 1) * dim];
3185 let dot = dot_product_f32(query, vec, dim);
3186 scores[i] = dot * inv_norm_q;
3187 }
3188}
3189
3190#[inline]
3195pub fn batch_dot_scores_f16(query: &[f32], vectors_raw: &[u8], dim: usize, scores: &mut [f32]) {
3196 let n = scores.len();
3197 let vec_bytes = dim.checked_mul(2).expect("f16 vector size overflow");
3198 let required = n
3199 .checked_mul(vec_bytes)
3200 .expect("f16 batch byte length overflow");
3201 assert_eq!(query.len(), dim, "f16 batch dot query dimension mismatch");
3202 assert!(
3203 vectors_raw.len() >= required,
3204 "f16 batch dot vectors are truncated: need {required} bytes, got {}",
3205 vectors_raw.len()
3206 );
3207 if required > 0 {
3208 assert!(
3209 (vectors_raw.as_ptr() as usize).is_multiple_of(std::mem::align_of::<u16>()),
3210 "f16 batch dot vectors are not 2-byte aligned"
3211 );
3212 }
3213 if dim == 0 || n == 0 {
3214 return;
3215 }
3216
3217 let norm_q_sq = dot_product_f32(query, query, dim);
3218 if norm_q_sq < f32::EPSILON {
3219 for s in scores.iter_mut() {
3220 *s = 0.0;
3221 }
3222 return;
3223 }
3224 let inv_norm_q = fast_inv_sqrt(norm_q_sq);
3225
3226 let query_f16: Vec<u16> = query.iter().map(|&v| f32_to_f16(v)).collect();
3227 for i in 0..n {
3228 let raw = &vectors_raw[i * vec_bytes..(i + 1) * vec_bytes];
3229 let f16_slice = unsafe { std::slice::from_raw_parts(raw.as_ptr() as *const u16, dim) };
3230 let dot = dot_product_f16_quant(&query_f16, f16_slice, dim);
3231 scores[i] = dot * inv_norm_q;
3232 }
3233}
3234
3235#[inline]
3240pub fn batch_dot_scores_u8(query: &[f32], vectors_raw: &[u8], dim: usize, scores: &mut [f32]) {
3241 let n = scores.len();
3242 let required = n.checked_mul(dim).expect("u8 batch byte length overflow");
3243 assert_eq!(query.len(), dim, "u8 batch dot query dimension mismatch");
3244 assert!(
3245 vectors_raw.len() >= required,
3246 "u8 batch dot vectors are truncated: need {required} bytes, got {}",
3247 vectors_raw.len()
3248 );
3249 if dim == 0 || n == 0 {
3250 return;
3251 }
3252
3253 let norm_q_sq = dot_product_f32(query, query, dim);
3254 if norm_q_sq < f32::EPSILON {
3255 for s in scores.iter_mut() {
3256 *s = 0.0;
3257 }
3258 return;
3259 }
3260 let inv_norm_q = fast_inv_sqrt(norm_q_sq);
3261
3262 for i in 0..n {
3263 let u8_slice = &vectors_raw[i * dim..(i + 1) * dim];
3264 let dot = dot_product_u8_quant(query, u8_slice, dim);
3265 scores[i] = dot * inv_norm_q;
3266 }
3267}
3268
3269#[inline]
3275pub fn batch_cosine_scores_precomp(
3276 query: &[f32],
3277 vectors: &[f32],
3278 dim: usize,
3279 scores: &mut [f32],
3280 inv_norm_q: f32,
3281) {
3282 let n = scores.len();
3283 let required = n
3284 .checked_mul(dim)
3285 .expect("precomputed cosine vector length overflow");
3286 assert_eq!(
3287 query.len(),
3288 dim,
3289 "precomputed cosine query dimension mismatch"
3290 );
3291 assert!(
3292 vectors.len() >= required,
3293 "precomputed cosine vectors are truncated: need {required}, got {}",
3294 vectors.len()
3295 );
3296 for i in 0..n {
3297 let vec = &vectors[i * dim..(i + 1) * dim];
3298 let (dot, norm_v_sq) = fused_dot_norm(query, vec, dim);
3299 scores[i] = if norm_v_sq < f32::EPSILON {
3300 0.0
3301 } else {
3302 dot * inv_norm_q * fast_inv_sqrt(norm_v_sq)
3303 };
3304 }
3305}
3306
3307#[inline]
3309pub fn batch_cosine_scores_f16_precomp(
3310 query_f16: &[u16],
3311 vectors_raw: &[u8],
3312 dim: usize,
3313 scores: &mut [f32],
3314 inv_norm_q: f32,
3315) {
3316 let n = scores.len();
3317 let vec_bytes = dim.checked_mul(2).expect("f16 vector size overflow");
3318 let required = n
3319 .checked_mul(vec_bytes)
3320 .expect("precomputed f16 cosine batch byte length overflow");
3321 assert_eq!(
3322 query_f16.len(),
3323 dim,
3324 "precomputed f16 cosine query dimension mismatch"
3325 );
3326 assert!(
3327 vectors_raw.len() >= required,
3328 "precomputed f16 cosine vectors are truncated: need {required} bytes, got {}",
3329 vectors_raw.len()
3330 );
3331 if required > 0 {
3332 assert!(
3333 (vectors_raw.as_ptr() as usize).is_multiple_of(std::mem::align_of::<u16>()),
3334 "precomputed f16 cosine vectors are not 2-byte aligned"
3335 );
3336 }
3337 for i in 0..n {
3338 let raw = &vectors_raw[i * vec_bytes..(i + 1) * vec_bytes];
3339 let f16_slice = unsafe { std::slice::from_raw_parts(raw.as_ptr() as *const u16, dim) };
3340 let (dot, norm_v_sq) = fused_dot_norm_f16(query_f16, f16_slice, dim);
3341 scores[i] = if norm_v_sq < f32::EPSILON {
3342 0.0
3343 } else {
3344 dot * inv_norm_q * fast_inv_sqrt(norm_v_sq)
3345 };
3346 }
3347}
3348
3349#[inline]
3351pub fn batch_cosine_scores_u8_precomp(
3352 query: &[f32],
3353 vectors_raw: &[u8],
3354 dim: usize,
3355 scores: &mut [f32],
3356 inv_norm_q: f32,
3357) {
3358 let n = scores.len();
3359 let required = n
3360 .checked_mul(dim)
3361 .expect("precomputed u8 cosine batch byte length overflow");
3362 assert_eq!(
3363 query.len(),
3364 dim,
3365 "precomputed u8 cosine query dimension mismatch"
3366 );
3367 assert!(
3368 vectors_raw.len() >= required,
3369 "precomputed u8 cosine vectors are truncated: need {required} bytes, got {}",
3370 vectors_raw.len()
3371 );
3372 for i in 0..n {
3373 let u8_slice = &vectors_raw[i * dim..(i + 1) * dim];
3374 let (dot, norm_v_sq) = fused_dot_norm_u8(query, u8_slice, dim);
3375 scores[i] = if norm_v_sq < f32::EPSILON {
3376 0.0
3377 } else {
3378 dot * inv_norm_q * fast_inv_sqrt(norm_v_sq)
3379 };
3380 }
3381}
3382
3383#[inline]
3385pub fn batch_dot_scores_precomp(
3386 query: &[f32],
3387 vectors: &[f32],
3388 dim: usize,
3389 scores: &mut [f32],
3390 inv_norm_q: f32,
3391) {
3392 let n = scores.len();
3393 let required = n
3394 .checked_mul(dim)
3395 .expect("precomputed dot vector length overflow");
3396 assert_eq!(query.len(), dim, "precomputed dot query dimension mismatch");
3397 assert!(
3398 vectors.len() >= required,
3399 "precomputed dot vectors are truncated: need {required}, got {}",
3400 vectors.len()
3401 );
3402 for i in 0..n {
3403 let vec = &vectors[i * dim..(i + 1) * dim];
3404 scores[i] = dot_product_f32(query, vec, dim) * inv_norm_q;
3405 }
3406}
3407
3408#[inline]
3410pub fn batch_dot_scores_f16_precomp(
3411 query_f16: &[u16],
3412 vectors_raw: &[u8],
3413 dim: usize,
3414 scores: &mut [f32],
3415 inv_norm_q: f32,
3416) {
3417 let n = scores.len();
3418 let vec_bytes = dim.checked_mul(2).expect("f16 vector size overflow");
3419 let required = n
3420 .checked_mul(vec_bytes)
3421 .expect("precomputed f16 dot batch byte length overflow");
3422 assert_eq!(
3423 query_f16.len(),
3424 dim,
3425 "precomputed f16 dot query dimension mismatch"
3426 );
3427 assert!(
3428 vectors_raw.len() >= required,
3429 "precomputed f16 dot vectors are truncated: need {required} bytes, got {}",
3430 vectors_raw.len()
3431 );
3432 if required > 0 {
3433 assert!(
3434 (vectors_raw.as_ptr() as usize).is_multiple_of(std::mem::align_of::<u16>()),
3435 "precomputed f16 dot vectors are not 2-byte aligned"
3436 );
3437 }
3438 for i in 0..n {
3439 let raw = &vectors_raw[i * vec_bytes..(i + 1) * vec_bytes];
3440 let f16_slice = unsafe { std::slice::from_raw_parts(raw.as_ptr() as *const u16, dim) };
3441 scores[i] = dot_product_f16_quant(query_f16, f16_slice, dim) * inv_norm_q;
3442 }
3443}
3444
3445#[inline]
3447pub fn batch_dot_scores_u8_precomp(
3448 query: &[f32],
3449 vectors_raw: &[u8],
3450 dim: usize,
3451 scores: &mut [f32],
3452 inv_norm_q: f32,
3453) {
3454 let n = scores.len();
3455 let required = n
3456 .checked_mul(dim)
3457 .expect("precomputed u8 dot batch byte length overflow");
3458 assert_eq!(
3459 query.len(),
3460 dim,
3461 "precomputed u8 dot query dimension mismatch"
3462 );
3463 assert!(
3464 vectors_raw.len() >= required,
3465 "precomputed u8 dot vectors are truncated: need {required} bytes, got {}",
3466 vectors_raw.len()
3467 );
3468 for i in 0..n {
3469 let u8_slice = &vectors_raw[i * dim..(i + 1) * dim];
3470 scores[i] = dot_product_u8_quant(query, u8_slice, dim) * inv_norm_q;
3471 }
3472}
3473
3474#[inline]
3479pub fn cosine_similarity(a: &[f32], b: &[f32]) -> f32 {
3480 assert_eq!(a.len(), b.len(), "cosine vector dimension mismatch");
3481 let count = a.len();
3482
3483 if count == 0 {
3484 return 0.0;
3485 }
3486
3487 let dot = dot_product_f32(a, b, count);
3488 let norm_a = dot_product_f32(a, a, count);
3489 let norm_b = dot_product_f32(b, b, count);
3490
3491 let denom = (norm_a * norm_b).sqrt();
3492 if denom < f32::EPSILON {
3493 return 0.0;
3494 }
3495
3496 dot / denom
3497}
3498
3499#[cfg(target_arch = "x86_64")]
3508#[target_feature(enable = "avx512f,avx512vpopcntdq")]
3509#[allow(unsafe_op_in_unsafe_fn)]
3510unsafe fn hamming_distance_avx512(a: &[u8], b: &[u8]) -> u32 {
3511 use std::arch::x86_64::*;
3512
3513 let len = a.len();
3514 let chunks64 = len / 64;
3515 let mut acc = _mm512_setzero_si512();
3516
3517 for c in 0..chunks64 {
3518 let off = c * 64;
3519 let va = _mm512_loadu_si512(a.as_ptr().add(off) as *const __m512i);
3520 let vb = _mm512_loadu_si512(b.as_ptr().add(off) as *const __m512i);
3521 acc = _mm512_add_epi64(acc, _mm512_popcnt_epi64(_mm512_xor_si512(va, vb)));
3522 }
3523
3524 let base = chunks64 * 64;
3525 _mm512_reduce_add_epi64(acc) as u32 + hamming_distance_scalar(&a[base..], &b[base..])
3526}
3527
3528#[cfg(target_arch = "x86_64")]
3530#[target_feature(enable = "avx512f,avx512vpopcntdq")]
3531#[allow(unsafe_op_in_unsafe_fn)]
3532unsafe fn hamming_distance_x4_avx512(query: &[u8], rows: [&[u8]; 4]) -> [u32; 4] {
3533 use std::arch::x86_64::*;
3534
3535 let len = query.len();
3536 let chunks64 = len / 64;
3537 let mut acc = [_mm512_setzero_si512(); 4];
3538
3539 for c in 0..chunks64 {
3540 let off = c * 64;
3541 let vq = _mm512_loadu_si512(query.as_ptr().add(off) as *const __m512i);
3542 for r in 0..4 {
3543 let vr = _mm512_loadu_si512(rows[r].as_ptr().add(off) as *const __m512i);
3544 acc[r] = _mm512_add_epi64(acc[r], _mm512_popcnt_epi64(_mm512_xor_si512(vq, vr)));
3545 }
3546 }
3547
3548 let base = chunks64 * 64;
3549 let tail = &query[base..];
3550 [
3551 _mm512_reduce_add_epi64(acc[0]) as u32 + hamming_distance_scalar(tail, &rows[0][base..]),
3552 _mm512_reduce_add_epi64(acc[1]) as u32 + hamming_distance_scalar(tail, &rows[1][base..]),
3553 _mm512_reduce_add_epi64(acc[2]) as u32 + hamming_distance_scalar(tail, &rows[2][base..]),
3554 _mm512_reduce_add_epi64(acc[3]) as u32 + hamming_distance_scalar(tail, &rows[3][base..]),
3555 ]
3556}
3557
3558#[inline]
3560fn hamming_distance_x4_scalar(query: &[u8], rows: [&[u8]; 4]) -> [u32; 4] {
3561 let len = query.len();
3562 let chunks = len / 8;
3563 let mut total = [0u32; 4];
3564
3565 for i in 0..chunks {
3566 let off = i * 8;
3567 let vq = unsafe { std::ptr::read_unaligned(query.as_ptr().add(off) as *const u64) };
3568 for r in 0..4 {
3569 let vr = unsafe { std::ptr::read_unaligned(rows[r].as_ptr().add(off) as *const u64) };
3570 total[r] += (vq ^ vr).count_ones();
3571 }
3572 }
3573
3574 let base = chunks * 8;
3575 for k in base..len {
3576 let q = query[k];
3577 for r in 0..4 {
3578 total[r] += (q ^ rows[r][k]).count_ones();
3579 }
3580 }
3581
3582 total
3583}
3584
3585const HAMMING_ROWS_PER_KERNEL: usize = 4;
3589
3590#[derive(Clone, Copy, Debug, PartialEq, Eq)]
3597pub enum HammingKernel {
3598 #[cfg(target_arch = "x86_64")]
3599 Avx512,
3600 #[cfg(target_arch = "x86_64")]
3601 Avx2,
3602 #[cfg(target_arch = "aarch64")]
3603 Neon,
3604 Scalar,
3605}
3606
3607impl HammingKernel {
3608 #[inline]
3610 pub fn resolve() -> Self {
3611 #[cfg(target_arch = "x86_64")]
3612 {
3613 if is_x86_feature_detected!("avx512f") && is_x86_feature_detected!("avx512vpopcntdq") {
3614 return Self::Avx512;
3615 }
3616 if avx2::is_available() {
3617 return Self::Avx2;
3618 }
3619 Self::Scalar
3620 }
3621
3622 #[cfg(target_arch = "aarch64")]
3623 {
3624 Self::Neon
3625 }
3626
3627 #[cfg(not(any(target_arch = "aarch64", target_arch = "x86_64")))]
3628 {
3629 Self::Scalar
3630 }
3631 }
3632
3633 #[inline]
3635 pub fn distance(self, a: &[u8], b: &[u8]) -> u32 {
3636 debug_assert_eq!(a.len(), b.len(), "Hamming vector byte length mismatch");
3637 match self {
3638 #[cfg(target_arch = "x86_64")]
3639 Self::Avx512 => unsafe { hamming_distance_avx512(a, b) },
3640 #[cfg(target_arch = "x86_64")]
3641 Self::Avx2 => unsafe { avx2::hamming_distance(a, b) },
3642 #[cfg(target_arch = "aarch64")]
3643 Self::Neon => unsafe { neon::hamming_distance(a, b) },
3644 Self::Scalar => hamming_distance_scalar(a, b),
3645 }
3646 }
3647
3648 pub fn distances(self, query: &[u8], db: &[u8], byte_len: usize, out: &mut [u32]) {
3650 self.score_rows(query, db, byte_len, out, |index| index);
3651 }
3652
3653 pub fn gather_distances(
3658 self,
3659 query: &[u8],
3660 db: &[u8],
3661 byte_len: usize,
3662 ids: &[u32],
3663 out: &mut [u32],
3664 ) {
3665 assert_eq!(
3666 ids.len(),
3667 out.len(),
3668 "Hamming gather needs one output slot per row id"
3669 );
3670 self.score_rows(query, db, byte_len, out, |index| ids[index] as usize);
3671 }
3672
3673 #[inline]
3674 fn score_rows(
3675 self,
3676 query: &[u8],
3677 db: &[u8],
3678 byte_len: usize,
3679 out: &mut [u32],
3680 index_of: impl Fn(usize) -> usize,
3681 ) {
3682 assert_eq!(query.len(), byte_len, "Hamming query byte length mismatch");
3683 if byte_len == 0 || out.is_empty() {
3684 return;
3685 }
3686 let row = |index: usize| -> &[u8] {
3687 let start = index * byte_len;
3688 &db[start..start + byte_len]
3689 };
3690 macro_rules! score_with {
3691 ($one:expr, $four:expr) => {{
3692 let mut i = 0;
3693 while i + HAMMING_ROWS_PER_KERNEL <= out.len() {
3694 let quad = [
3695 row(index_of(i)),
3696 row(index_of(i + 1)),
3697 row(index_of(i + 2)),
3698 row(index_of(i + 3)),
3699 ];
3700 out[i..i + HAMMING_ROWS_PER_KERNEL].copy_from_slice(&$four(query, quad));
3701 i += HAMMING_ROWS_PER_KERNEL;
3702 }
3703 while i < out.len() {
3704 out[i] = $one(query, row(index_of(i)));
3705 i += 1;
3706 }
3707 }};
3708 }
3709 match self {
3710 #[cfg(target_arch = "x86_64")]
3711 Self::Avx512 => score_with!(
3712 |query, row| unsafe { hamming_distance_avx512(query, row) },
3713 |query, rows| unsafe { hamming_distance_x4_avx512(query, rows) }
3714 ),
3715 #[cfg(target_arch = "x86_64")]
3716 Self::Avx2 => score_with!(
3717 |query, row| unsafe { avx2::hamming_distance(query, row) },
3718 |query, rows| unsafe { avx2::hamming_distance_x4(query, rows) }
3719 ),
3720 #[cfg(target_arch = "aarch64")]
3721 Self::Neon => score_with!(
3722 |query, row| unsafe { neon::hamming_distance(query, row) },
3723 |query, rows| unsafe { neon::hamming_distance_x4(query, rows) }
3724 ),
3725 Self::Scalar => score_with!(hamming_distance_scalar, hamming_distance_x4_scalar),
3726 }
3727 }
3728}
3729
3730#[inline]
3737pub fn hamming_distance(a: &[u8], b: &[u8]) -> u32 {
3738 assert_eq!(a.len(), b.len(), "Hamming vector byte length mismatch");
3739 HammingKernel::resolve().distance(a, b)
3740}
3741
3742#[inline]
3745#[allow(dead_code)]
3746fn hamming_distance_scalar(a: &[u8], b: &[u8]) -> u32 {
3747 let len = a.len();
3748 let chunks = len / 8;
3749 let remainder = len % 8;
3750 let mut total = 0u32;
3751
3752 for i in 0..chunks {
3753 let off = i * 8;
3754 let va = unsafe { std::ptr::read_unaligned(a.as_ptr().add(off) as *const u64) };
3755 let vb = unsafe { std::ptr::read_unaligned(b.as_ptr().add(off) as *const u64) };
3756 total += (va ^ vb).count_ones();
3757 }
3758
3759 let base = chunks * 8;
3760 for i in 0..remainder {
3761 total += (a[base + i] ^ b[base + i]).count_ones();
3762 }
3763
3764 total
3765}
3766
3767pub fn batch_hamming_scores(
3773 query: &[u8],
3774 db: &[u8],
3775 byte_len: usize,
3776 dim_bits: usize,
3777 scores: &mut [f32],
3778) {
3779 let n = scores.len();
3780 let required = n
3781 .checked_mul(byte_len)
3782 .expect("Hamming batch byte length overflow");
3783 assert_eq!(query.len(), byte_len, "Hamming query byte length mismatch");
3784 assert!(
3785 db.len() >= required,
3786 "Hamming batch is truncated: need {required} bytes, got {}",
3787 db.len()
3788 );
3789
3790 if byte_len == 0 || n == 0 || dim_bits == 0 {
3791 return;
3792 }
3793
3794 scores_from_hamming(
3795 HammingKernel::resolve(),
3796 query,
3797 db,
3798 byte_len,
3799 dim_bits,
3800 scores,
3801 );
3802}
3803
3804pub fn scores_from_hamming(
3809 kernel: HammingKernel,
3810 query: &[u8],
3811 db: &[u8],
3812 byte_len: usize,
3813 dim_bits: usize,
3814 scores: &mut [f32],
3815) {
3816 if byte_len == 0 || scores.is_empty() || dim_bits == 0 {
3817 return;
3818 }
3819 let inv_dim = 1.0 / dim_bits as f32;
3820 let mut distances = [0u32; HAMMING_DISTANCE_BLOCK];
3823 for (block_index, block) in scores.chunks_mut(HAMMING_DISTANCE_BLOCK).enumerate() {
3824 let rows = &mut distances[..block.len()];
3825 kernel.distances(
3826 query,
3827 &db[block_index * HAMMING_DISTANCE_BLOCK * byte_len..],
3828 byte_len,
3829 rows,
3830 );
3831 for (score, &distance) in block.iter_mut().zip(rows.iter()) {
3832 *score = 1.0 - distance as f32 * inv_dim;
3833 }
3834 }
3835}
3836
3837const HAMMING_DISTANCE_BLOCK: usize = 64;
3839
3840pub fn batch_hamming_distances(query: &[u8], db: &[u8], byte_len: usize, out: &mut [u32]) {
3845 HammingKernel::resolve().distances(query, db, byte_len, out);
3846}
3847
3848#[cfg(test)]
3849mod tests {
3850 use super::*;
3851
3852 #[test]
3853 fn vector_simd_boundaries_reject_dimension_mismatches() {
3854 let vectors = vec![1.0f32; 6];
3855 let raw_f16 = vec![0u8; 12];
3856 let raw_u8 = vec![0u8; 6];
3857 let mut scores = vec![0.0f32; 2];
3858
3859 for invalid_query in [vec![1.0, 2.0], vec![1.0, 2.0, 3.0, 4.0]] {
3860 assert!(
3861 std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
3862 batch_cosine_scores(&invalid_query, &vectors, 3, &mut scores)
3863 }))
3864 .is_err()
3865 );
3866 assert!(
3867 std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
3868 batch_dot_scores_f16(&invalid_query, &raw_f16, 3, &mut scores)
3869 }))
3870 .is_err()
3871 );
3872 assert!(
3873 std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
3874 batch_cosine_scores_u8(&invalid_query, &raw_u8, 3, &mut scores)
3875 }))
3876 .is_err()
3877 );
3878 }
3879 }
3880
3881 #[test]
3882 fn vector_simd_boundaries_reject_truncated_storage() {
3883 let query = [1.0f32, 2.0, 3.0];
3884 let mut scores = [0.0f32; 2];
3885
3886 assert!(
3887 std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
3888 batch_dot_scores(&query, &[0.0; 5], 3, &mut scores)
3889 }))
3890 .is_err()
3891 );
3892 assert!(
3893 std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
3894 batch_cosine_scores_f16(&query, &[0u8; 11], 3, &mut scores)
3895 }))
3896 .is_err()
3897 );
3898 assert!(
3899 std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
3900 dot_product_f32(&query, &query, 4)
3901 }))
3902 .is_err()
3903 );
3904 }
3905
3906 #[test]
3907 fn test_unpack_8bit() {
3908 let input: Vec<u8> = (0..128).collect();
3909 let mut output = vec![0u32; 128];
3910 unpack_8bit(&input, &mut output, 128);
3911
3912 for (i, &v) in output.iter().enumerate() {
3913 assert_eq!(v, i as u32);
3914 }
3915 }
3916
3917 #[test]
3918 fn test_unpack_16bit() {
3919 let mut input = vec![0u8; 256];
3920 for i in 0..128 {
3921 let val = (i * 100) as u16;
3922 input[i * 2] = val as u8;
3923 input[i * 2 + 1] = (val >> 8) as u8;
3924 }
3925
3926 let mut output = vec![0u32; 128];
3927 unpack_16bit(&input, &mut output, 128);
3928
3929 for (i, &v) in output.iter().enumerate() {
3930 assert_eq!(v, (i * 100) as u32);
3931 }
3932 }
3933
3934 #[test]
3935 fn test_unpack_32bit() {
3936 let mut input = vec![0u8; 512];
3937 for i in 0..128 {
3938 let val = (i * 1000) as u32;
3939 let bytes = val.to_le_bytes();
3940 input[i * 4..i * 4 + 4].copy_from_slice(&bytes);
3941 }
3942
3943 let mut output = vec![0u32; 128];
3944 unpack_32bit(&input, &mut output, 128);
3945
3946 for (i, &v) in output.iter().enumerate() {
3947 assert_eq!(v, (i * 1000) as u32);
3948 }
3949 }
3950
3951 #[test]
3952 fn test_delta_decode() {
3953 let deltas = vec![4u32, 4, 9, 19];
3957 let mut output = vec![0u32; 5];
3958
3959 delta_decode(&mut output, &deltas, 10, 5);
3960
3961 assert_eq!(output, vec![10, 15, 20, 30, 50]);
3962 }
3963
3964 #[test]
3965 fn test_add_one() {
3966 let mut values = vec![0u32, 1, 2, 3, 4, 5, 6, 7];
3967 add_one(&mut values, 8);
3968
3969 assert_eq!(values, vec![1, 2, 3, 4, 5, 6, 7, 8]);
3970 }
3971
3972 #[test]
3973 fn test_bits_needed() {
3974 assert_eq!(bits_needed(0), 0);
3975 assert_eq!(bits_needed(1), 1);
3976 assert_eq!(bits_needed(2), 2);
3977 assert_eq!(bits_needed(3), 2);
3978 assert_eq!(bits_needed(4), 3);
3979 assert_eq!(bits_needed(255), 8);
3980 assert_eq!(bits_needed(256), 9);
3981 assert_eq!(bits_needed(u32::MAX), 32);
3982 }
3983
3984 #[test]
3985 fn test_unpack_8bit_delta_decode() {
3986 let input: Vec<u8> = vec![4, 4, 9, 19];
3990 let mut output = vec![0u32; 5];
3991
3992 unpack_8bit_delta_decode(&input, &mut output, 10, 5);
3993
3994 assert_eq!(output, vec![10, 15, 20, 30, 50]);
3995 }
3996
3997 #[test]
3998 fn test_unpack_16bit_delta_decode() {
3999 let mut input = vec![0u8; 8];
4003 for (i, &delta) in [499u16, 499, 999, 1999].iter().enumerate() {
4004 input[i * 2] = delta as u8;
4005 input[i * 2 + 1] = (delta >> 8) as u8;
4006 }
4007 let mut output = vec![0u32; 5];
4008
4009 unpack_16bit_delta_decode(&input, &mut output, 100, 5);
4010
4011 assert_eq!(output, vec![100, 600, 1100, 2100, 4100]);
4012 }
4013
4014 #[test]
4015 fn test_fused_vs_separate_8bit() {
4016 let input: Vec<u8> = (0..127).collect();
4018 let first_value = 1000u32;
4019 let count = 128;
4020
4021 let mut unpacked = vec![0u32; 128];
4023 unpack_8bit(&input, &mut unpacked, 127);
4024 let mut separate_output = vec![0u32; 128];
4025 delta_decode(&mut separate_output, &unpacked, first_value, count);
4026
4027 let mut fused_output = vec![0u32; 128];
4029 unpack_8bit_delta_decode(&input, &mut fused_output, first_value, count);
4030
4031 assert_eq!(separate_output, fused_output);
4032 }
4033
4034 #[test]
4035 fn test_round_bit_width() {
4036 assert_eq!(round_bit_width(0), 0);
4037 assert_eq!(round_bit_width(1), 8);
4038 assert_eq!(round_bit_width(5), 8);
4039 assert_eq!(round_bit_width(8), 8);
4040 assert_eq!(round_bit_width(9), 16);
4041 assert_eq!(round_bit_width(12), 16);
4042 assert_eq!(round_bit_width(16), 16);
4043 assert_eq!(round_bit_width(17), 32);
4044 assert_eq!(round_bit_width(24), 32);
4045 assert_eq!(round_bit_width(32), 32);
4046 }
4047
4048 #[test]
4049 fn test_rounded_bitwidth_from_exact() {
4050 assert_eq!(RoundedBitWidth::from_exact(0), RoundedBitWidth::Zero);
4051 assert_eq!(RoundedBitWidth::from_exact(1), RoundedBitWidth::Bits8);
4052 assert_eq!(RoundedBitWidth::from_exact(8), RoundedBitWidth::Bits8);
4053 assert_eq!(RoundedBitWidth::from_exact(9), RoundedBitWidth::Bits16);
4054 assert_eq!(RoundedBitWidth::from_exact(16), RoundedBitWidth::Bits16);
4055 assert_eq!(RoundedBitWidth::from_exact(17), RoundedBitWidth::Bits32);
4056 assert_eq!(RoundedBitWidth::from_exact(32), RoundedBitWidth::Bits32);
4057 }
4058
4059 #[test]
4060 fn test_pack_unpack_rounded_8bit() {
4061 let values: Vec<u32> = (0..128).map(|i| i % 256).collect();
4062 let mut packed = vec![0u8; 128];
4063
4064 let bytes_written = pack_rounded(&values, RoundedBitWidth::Bits8, &mut packed);
4065 assert_eq!(bytes_written, 128);
4066
4067 let mut unpacked = vec![0u32; 128];
4068 unpack_rounded(&packed, RoundedBitWidth::Bits8, &mut unpacked, 128);
4069
4070 assert_eq!(values, unpacked);
4071 }
4072
4073 #[test]
4074 fn test_pack_unpack_rounded_16bit() {
4075 let values: Vec<u32> = (0..128).map(|i| i * 100).collect();
4076 let mut packed = vec![0u8; 256];
4077
4078 let bytes_written = pack_rounded(&values, RoundedBitWidth::Bits16, &mut packed);
4079 assert_eq!(bytes_written, 256);
4080
4081 let mut unpacked = vec![0u32; 128];
4082 unpack_rounded(&packed, RoundedBitWidth::Bits16, &mut unpacked, 128);
4083
4084 assert_eq!(values, unpacked);
4085 }
4086
4087 #[test]
4088 fn test_pack_unpack_rounded_32bit() {
4089 let values: Vec<u32> = (0..128).map(|i| i * 100000).collect();
4090 let mut packed = vec![0u8; 512];
4091
4092 let bytes_written = pack_rounded(&values, RoundedBitWidth::Bits32, &mut packed);
4093 assert_eq!(bytes_written, 512);
4094
4095 let mut unpacked = vec![0u32; 128];
4096 unpack_rounded(&packed, RoundedBitWidth::Bits32, &mut unpacked, 128);
4097
4098 assert_eq!(values, unpacked);
4099 }
4100
4101 #[test]
4102 fn test_unpack_rounded_delta_decode() {
4103 let input: Vec<u8> = vec![4, 4, 9, 19];
4108 let mut output = vec![0u32; 5];
4109
4110 unpack_rounded_delta_decode(&input, RoundedBitWidth::Bits8, &mut output, 10, 5);
4111
4112 assert_eq!(output, vec![10, 15, 20, 30, 50]);
4113 }
4114
4115 #[test]
4116 fn test_unpack_rounded_delta_decode_zero() {
4117 let input: Vec<u8> = vec![];
4119 let mut output = vec![0u32; 5];
4120
4121 unpack_rounded_delta_decode(&input, RoundedBitWidth::Zero, &mut output, 100, 5);
4122
4123 assert_eq!(output, vec![100, 101, 102, 103, 104]);
4124 }
4125
4126 #[test]
4131 fn test_dequantize_uint8() {
4132 let input: Vec<u8> = vec![0, 128, 255, 64, 192];
4133 let mut output = vec![0.0f32; 5];
4134 let scale = 0.1;
4135 let min_val = 1.0;
4136
4137 dequantize_uint8(&input, &mut output, scale, min_val, 5);
4138
4139 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); }
4146
4147 #[test]
4148 fn test_dequantize_uint8_large() {
4149 let input: Vec<u8> = (0..128).collect();
4151 let mut output = vec![0.0f32; 128];
4152 let scale = 2.0;
4153 let min_val = -10.0;
4154
4155 dequantize_uint8(&input, &mut output, scale, min_val, 128);
4156
4157 for (i, &out) in output.iter().enumerate().take(128) {
4158 let expected = i as f32 * scale + min_val;
4159 assert!(
4160 (out - expected).abs() < 1e-5,
4161 "Mismatch at {}: expected {}, got {}",
4162 i,
4163 expected,
4164 out
4165 );
4166 }
4167 }
4168
4169 #[test]
4170 fn test_dot_product_f32() {
4171 let a = vec![1.0f32, 2.0, 3.0, 4.0, 5.0];
4172 let b = vec![2.0f32, 3.0, 4.0, 5.0, 6.0];
4173
4174 let result = dot_product_f32(&a, &b, 5);
4175
4176 assert!((result - 70.0).abs() < 1e-5);
4178 }
4179
4180 #[test]
4181 fn test_dot_product_f32_large() {
4182 let a: Vec<f32> = (0..128).map(|i| i as f32).collect();
4184 let b: Vec<f32> = (0..128).map(|i| (i + 1) as f32).collect();
4185
4186 let result = dot_product_f32(&a, &b, 128);
4187
4188 let expected: f32 = (0..128).map(|i| (i as f32) * ((i + 1) as f32)).sum();
4190 assert!(
4191 (result - expected).abs() < 1e-3,
4192 "Expected {}, got {}",
4193 expected,
4194 result
4195 );
4196 }
4197
4198 #[test]
4199 fn test_fused_dot_norm() {
4200 let a = vec![1.0f32, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0];
4201 let b = vec![2.0f32, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0];
4202 let (dot, norm_b) = fused_dot_norm(&a, &b, a.len());
4203
4204 let expected_dot: f32 = a.iter().zip(b.iter()).map(|(x, y)| x * y).sum();
4205 let expected_norm: f32 = b.iter().map(|x| x * x).sum();
4206 assert!(
4207 (dot - expected_dot).abs() < 1e-5,
4208 "dot: expected {}, got {}",
4209 expected_dot,
4210 dot
4211 );
4212 assert!(
4213 (norm_b - expected_norm).abs() < 1e-5,
4214 "norm: expected {}, got {}",
4215 expected_norm,
4216 norm_b
4217 );
4218 }
4219
4220 #[test]
4221 fn test_fused_dot_norm_large() {
4222 let a: Vec<f32> = (0..768).map(|i| (i as f32) * 0.01).collect();
4223 let b: Vec<f32> = (0..768).map(|i| (i as f32) * 0.02 + 0.5).collect();
4224 let (dot, norm_b) = fused_dot_norm(&a, &b, a.len());
4225
4226 let expected_dot: f32 = a.iter().zip(b.iter()).map(|(x, y)| x * y).sum();
4227 let expected_norm: f32 = b.iter().map(|x| x * x).sum();
4228 assert!(
4229 (dot - expected_dot).abs() < 1.0,
4230 "dot: expected {}, got {}",
4231 expected_dot,
4232 dot
4233 );
4234 assert!(
4235 (norm_b - expected_norm).abs() < 1.0,
4236 "norm: expected {}, got {}",
4237 expected_norm,
4238 norm_b
4239 );
4240 }
4241
4242 #[test]
4243 fn test_batch_cosine_scores() {
4244 let query = vec![1.0f32, 0.0, 0.0];
4246 let vectors = vec![
4247 1.0, 0.0, 0.0, 0.0, 1.0, 0.0, -1.0, 0.0, 0.0, 0.5, 0.5, 0.0, ];
4252 let mut scores = vec![0f32; 4];
4253 batch_cosine_scores(&query, &vectors, 3, &mut scores);
4254
4255 assert!((scores[0] - 1.0).abs() < 1e-5, "identical: {}", scores[0]);
4256 assert!(scores[1].abs() < 1e-5, "orthogonal: {}", scores[1]);
4257 assert!((scores[2] - (-1.0)).abs() < 1e-5, "opposite: {}", scores[2]);
4258 let expected_45 = 0.5f32 / (0.5f32.powi(2) + 0.5f32.powi(2)).sqrt();
4259 assert!(
4260 (scores[3] - expected_45).abs() < 1e-5,
4261 "45deg: expected {}, got {}",
4262 expected_45,
4263 scores[3]
4264 );
4265 }
4266
4267 #[test]
4268 fn test_batch_cosine_scores_matches_individual() {
4269 let query: Vec<f32> = (0..128).map(|i| (i as f32) * 0.1).collect();
4270 let n = 50;
4271 let dim = 128;
4272 let vectors: Vec<f32> = (0..n * dim).map(|i| ((i * 7 + 3) as f32) * 0.01).collect();
4273
4274 let mut batch_scores = vec![0f32; n];
4275 batch_cosine_scores(&query, &vectors, dim, &mut batch_scores);
4276
4277 for i in 0..n {
4278 let vec_i = &vectors[i * dim..(i + 1) * dim];
4279 let individual = cosine_similarity(&query, vec_i);
4280 assert!(
4281 (batch_scores[i] - individual).abs() < 1e-5,
4282 "vec {}: batch={}, individual={}",
4283 i,
4284 batch_scores[i],
4285 individual
4286 );
4287 }
4288 }
4289
4290 #[test]
4291 fn test_batch_cosine_scores_empty() {
4292 let query = vec![1.0f32, 2.0, 3.0];
4293 let vectors: Vec<f32> = vec![];
4294 let mut scores: Vec<f32> = vec![];
4295 batch_cosine_scores(&query, &vectors, 3, &mut scores);
4296 assert!(scores.is_empty());
4297 }
4298
4299 #[test]
4300 fn test_batch_cosine_scores_zero_query() {
4301 let query = vec![0.0f32, 0.0, 0.0];
4302 let vectors = vec![1.0f32, 2.0, 3.0, 4.0, 5.0, 6.0];
4303 let mut scores = vec![0f32; 2];
4304 batch_cosine_scores(&query, &vectors, 3, &mut scores);
4305 assert_eq!(scores[0], 0.0);
4306 assert_eq!(scores[1], 0.0);
4307 }
4308
4309 #[test]
4314 fn test_f16_roundtrip_normal() {
4315 for &v in &[0.0f32, 1.0, -1.0, 0.5, -0.5, 0.333, 65504.0] {
4316 let h = f32_to_f16(v);
4317 let back = f16_to_f32(h);
4318 let err = (back - v).abs() / v.abs().max(1e-6);
4319 assert!(
4320 err < 0.002,
4321 "f16 roundtrip {v} → {h:#06x} → {back}, rel err {err}"
4322 );
4323 }
4324 }
4325
4326 #[test]
4327 fn test_f16_special() {
4328 assert_eq!(f16_to_f32(f32_to_f16(0.0)), 0.0);
4330 assert_eq!(f32_to_f16(-0.0), 0x8000);
4332 assert!(f16_to_f32(f32_to_f16(f32::INFINITY)).is_infinite());
4334 assert!(f16_to_f32(f32_to_f16(f32::NAN)).is_nan());
4336 }
4337
4338 #[test]
4339 fn test_f16_embedding_range() {
4340 let values: Vec<f32> = (-100..=100).map(|i| i as f32 / 100.0).collect();
4342 for &v in &values {
4343 let back = f16_to_f32(f32_to_f16(v));
4344 assert!((back - v).abs() < 0.001, "f16 error for {v}: got {back}");
4345 }
4346 }
4347
4348 #[test]
4353 fn test_u8_roundtrip() {
4354 assert_eq!(f32_to_u8_saturating(-1.0), 0);
4356 assert_eq!(f32_to_u8_saturating(1.0), 255);
4357 assert_eq!(f32_to_u8_saturating(0.0), 127); assert_eq!(f32_to_u8_saturating(-2.0), 0);
4361 assert_eq!(f32_to_u8_saturating(2.0), 255);
4362 }
4363
4364 #[test]
4365 fn test_u8_dequantize() {
4366 assert!((u8_to_f32(0) - (-1.0)).abs() < 0.01);
4367 assert!((u8_to_f32(255) - 1.0).abs() < 0.01);
4368 assert!((u8_to_f32(127) - 0.0).abs() < 0.01);
4369 }
4370
4371 #[test]
4376 fn test_batch_cosine_scores_f16() {
4377 let query = vec![0.6f32, 0.8, 0.0, 0.0];
4378 let dim = 4;
4379 let vecs_f32 = vec![
4380 0.6f32, 0.8, 0.0, 0.0, 0.0, 0.0, 0.6, 0.8, ];
4383
4384 let mut f16_buf = vec![0u16; 8];
4386 batch_f32_to_f16(&vecs_f32, &mut f16_buf);
4387 let raw: &[u8] =
4388 unsafe { std::slice::from_raw_parts(f16_buf.as_ptr() as *const u8, f16_buf.len() * 2) };
4389
4390 let mut scores = vec![0f32; 2];
4391 batch_cosine_scores_f16(&query, raw, dim, &mut scores);
4392
4393 assert!(
4394 (scores[0] - 1.0).abs() < 0.01,
4395 "identical vectors: {}",
4396 scores[0]
4397 );
4398 assert!(scores[1].abs() < 0.01, "orthogonal vectors: {}", scores[1]);
4399 }
4400
4401 #[test]
4402 fn test_batch_cosine_scores_u8() {
4403 let query = vec![0.6f32, 0.8, 0.0, 0.0];
4404 let dim = 4;
4405 let vecs_f32 = vec![
4406 0.6f32, 0.8, 0.0, 0.0, -0.6, -0.8, 0.0, 0.0, ];
4409
4410 let mut u8_buf = vec![0u8; 8];
4412 batch_f32_to_u8(&vecs_f32, &mut u8_buf);
4413
4414 let mut scores = vec![0f32; 2];
4415 batch_cosine_scores_u8(&query, &u8_buf, dim, &mut scores);
4416
4417 assert!(scores[0] > 0.95, "similar vectors: {}", scores[0]);
4418 assert!(scores[1] < -0.95, "opposite vectors: {}", scores[1]);
4419 }
4420
4421 #[test]
4422 fn test_batch_cosine_scores_f16_large_dim() {
4423 let dim = 768;
4425 let query: Vec<f32> = (0..dim).map(|i| (i as f32 / dim as f32) - 0.5).collect();
4426 let vec2: Vec<f32> = query.iter().map(|x| x * 0.9 + 0.01).collect();
4427
4428 let mut all_vecs = query.clone();
4429 all_vecs.extend_from_slice(&vec2);
4430
4431 let mut f16_buf = vec![0u16; all_vecs.len()];
4432 batch_f32_to_f16(&all_vecs, &mut f16_buf);
4433 let raw: &[u8] =
4434 unsafe { std::slice::from_raw_parts(f16_buf.as_ptr() as *const u8, f16_buf.len() * 2) };
4435
4436 let mut scores = vec![0f32; 2];
4437 batch_cosine_scores_f16(&query, raw, dim, &mut scores);
4438
4439 assert!((scores[0] - 1.0).abs() < 0.01, "self-sim: {}", scores[0]);
4441 assert!(scores[1] > 0.99, "scaled-sim: {}", scores[1]);
4443 }
4444
4445 #[test]
4450 fn test_hamming_distance_identical() {
4451 let a = vec![0xAA; 64];
4452 assert_eq!(hamming_distance(&a, &a), 0);
4453 }
4454
4455 #[test]
4456 fn test_hamming_distance_opposite() {
4457 let a = vec![0xFF; 32];
4458 let b = vec![0x00; 32];
4459 assert_eq!(hamming_distance(&a, &b), 256);
4460 }
4461
4462 #[test]
4463 fn test_hamming_distance_known() {
4464 let a = vec![0xAA];
4466 let b = vec![0x55];
4467 assert_eq!(hamming_distance(&a, &b), 8);
4468
4469 let a = vec![0xFF, 0x00];
4471 let b = vec![0x00, 0x00];
4472 assert_eq!(hamming_distance(&a, &b), 8);
4473 }
4474
4475 #[test]
4476 fn test_hamming_distance_single_bit() {
4477 let a = vec![0x00; 16];
4478 let mut b = vec![0x00; 16];
4479 b[7] = 0x01; assert_eq!(hamming_distance(&a, &b), 1);
4481 }
4482
4483 #[test]
4484 fn test_hamming_distance_empty() {
4485 let a: Vec<u8> = vec![];
4486 assert_eq!(hamming_distance(&a, &a), 0);
4487 }
4488
4489 #[test]
4490 fn test_hamming_distance_remainder_path() {
4491 let a = vec![0xFF; 17];
4493 let b = vec![0x00; 17];
4494 assert_eq!(hamming_distance(&a, &b), 136); let a = vec![0xFF; 33];
4498 let b = vec![0x00; 33];
4499 assert_eq!(hamming_distance(&a, &b), 264); }
4501
4502 #[test]
4503 fn test_hamming_distance_large() {
4504 let a = vec![0xFF; 4096];
4506 let b = vec![0x00; 4096];
4507 assert_eq!(hamming_distance(&a, &b), 32768);
4508 }
4509
4510 #[test]
4511 fn test_hamming_distance_scalar_matches() {
4512 for size in [1, 7, 8, 15, 16, 31, 32, 63, 64, 100, 128, 255, 256] {
4514 let a: Vec<u8> = (0..size).map(|i| (i * 37 + 13) as u8).collect();
4515 let b: Vec<u8> = (0..size).map(|i| (i * 53 + 7) as u8).collect();
4516 let expected = hamming_distance_scalar(&a, &b);
4517 let got = hamming_distance(&a, &b);
4518 assert_eq!(got, expected, "mismatch at size {size}");
4519 }
4520 }
4521
4522 #[test]
4527 fn test_batch_hamming_scores_identical() {
4528 let query = vec![0xAA; 16];
4529 let db = vec![0xAA; 16]; let mut scores = vec![0f32; 1];
4531 batch_hamming_scores(&query, &db, 16, 128, &mut scores);
4532 assert!((scores[0] - 1.0).abs() < 1e-6, "identical: {}", scores[0]);
4533 }
4534
4535 #[test]
4536 fn test_batch_hamming_scores_opposite() {
4537 let query = vec![0xFF; 16];
4538 let db = vec![0x00; 16];
4539 let mut scores = vec![0f32; 1];
4540 batch_hamming_scores(&query, &db, 16, 128, &mut scores);
4541 assert!((scores[0] - 0.0).abs() < 1e-6, "opposite: {}", scores[0]);
4542 }
4543
4544 #[test]
4545 fn test_batch_hamming_scores_multiple() {
4546 let byte_len = 8;
4547 let dim_bits = 64;
4548 let query = vec![0xFF; byte_len];
4549 let mut db = Vec::new();
4550 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];
4555 batch_hamming_scores(&query, &db, byte_len, dim_bits, &mut scores);
4556
4557 assert!((scores[0] - 1.0).abs() < 1e-6, "identical: {}", scores[0]);
4558 assert!((scores[1] - 0.0).abs() < 1e-6, "opposite: {}", scores[1]);
4559 assert!((scores[2] - 0.5).abs() < 1e-6, "half: {}", scores[2]);
4560 }
4561
4562 #[test]
4563 fn test_batch_hamming_scores_empty() {
4564 let query = vec![0xFF; 8];
4565 let db: Vec<u8> = vec![];
4566 let mut scores: Vec<f32> = vec![];
4567 batch_hamming_scores(&query, &db, 8, 64, &mut scores);
4568 assert!(scores.is_empty());
4569 }
4570
4571 #[test]
4572 fn test_batch_hamming_scores_zero_byte_len() {
4573 let query: Vec<u8> = vec![];
4574 let db: Vec<u8> = vec![];
4575 let mut scores = vec![0f32; 1];
4576 batch_hamming_scores(&query, &db, 0, 0, &mut scores);
4577 assert_eq!(scores[0], 0.0);
4579 }
4580
4581 fn hamming_matrix(rows: usize, byte_len: usize) -> (Vec<u8>, Vec<u8>) {
4586 let query: Vec<u8> = (0..byte_len).map(|i| (i * 31 + 5) as u8).collect();
4587 let db: Vec<u8> = (0..rows * byte_len)
4588 .map(|i| (i * 97 + i / byte_len * 11 + 3) as u8)
4589 .collect();
4590 (query, db)
4591 }
4592
4593 #[test]
4596 fn batched_hamming_distances_match_scalar_for_every_row_count() {
4597 let kernel = HammingKernel::resolve();
4598 for byte_len in [1, 7, 8, 15, 16, 31, 32, 33, 63, 64, 65, 128, 320] {
4601 for rows in [1, 2, 3, 4, 5, 7, 8, 9, 64, 70] {
4602 let (query, db) = hamming_matrix(rows, byte_len);
4603 let mut got = vec![0u32; rows];
4604 kernel.distances(&query, &db, byte_len, &mut got);
4605 for (row, &distance) in got.iter().enumerate() {
4606 let expected =
4607 hamming_distance_scalar(&query, &db[row * byte_len..(row + 1) * byte_len]);
4608 assert_eq!(
4609 distance, expected,
4610 "row {row} of {rows} at byte_len {byte_len}"
4611 );
4612 }
4613 }
4614 }
4615 }
4616
4617 #[test]
4618 fn gathered_hamming_distances_follow_row_ids() {
4619 let kernel = HammingKernel::resolve();
4620 let byte_len = 320;
4621 let rows = 37;
4622 let (query, db) = hamming_matrix(rows, byte_len);
4623 let ids: Vec<u32> = [36, 0, 17, 17, 5, 31, 2, 9, 9, 36, 1].into_iter().collect();
4626 let mut got = vec![0u32; ids.len()];
4627 kernel.gather_distances(&query, &db, byte_len, &ids, &mut got);
4628 for (slot, &id) in ids.iter().enumerate() {
4629 let start = id as usize * byte_len;
4630 let expected = hamming_distance_scalar(&query, &db[start..start + byte_len]);
4631 assert_eq!(got[slot], expected, "slot {slot} for row {id}");
4632 }
4633 }
4634
4635 #[test]
4636 fn resolved_kernel_matches_scalar_pairwise() {
4637 let kernel = HammingKernel::resolve();
4638 for byte_len in [1, 8, 32, 64, 65, 320, 4096] {
4639 let (query, db) = hamming_matrix(1, byte_len);
4640 assert_eq!(
4641 kernel.distance(&query, &db),
4642 hamming_distance_scalar(&query, &db),
4643 "byte_len {byte_len}"
4644 );
4645 }
4646 }
4647
4648 #[test]
4649 fn scores_from_hamming_matches_batch_scores_across_blocks() {
4650 let kernel = HammingKernel::resolve();
4651 let byte_len = 320;
4652 let dim_bits = byte_len * 8;
4653 let rows = HAMMING_DISTANCE_BLOCK * 2 + 3;
4655 let (query, db) = hamming_matrix(rows, byte_len);
4656 let mut expected = vec![0f32; rows];
4657 for (row, score) in expected.iter_mut().enumerate() {
4658 let distance =
4659 hamming_distance_scalar(&query, &db[row * byte_len..(row + 1) * byte_len]);
4660 *score = 1.0 - distance as f32 / dim_bits as f32;
4661 }
4662 let mut got = vec![0f32; rows];
4663 scores_from_hamming(kernel, &query, &db, byte_len, dim_bits, &mut got);
4664 for (row, (&got, &want)) in got.iter().zip(expected.iter()).enumerate() {
4665 assert!((got - want).abs() < 1e-6, "row {row}: {got} vs {want}");
4666 }
4667 let mut public = vec![0f32; rows];
4668 batch_hamming_scores(&query, &db, byte_len, dim_bits, &mut public);
4669 assert_eq!(got, public);
4670 }
4671}
4672
4673#[inline]
4686pub fn find_first_ge_u32(slice: &[u32], target: u32) -> usize {
4687 #[cfg(target_arch = "aarch64")]
4688 {
4689 if neon::is_available() {
4690 return unsafe { find_first_ge_u32_neon(slice, target) };
4691 }
4692 }
4693
4694 #[cfg(target_arch = "x86_64")]
4695 {
4696 if sse::is_available() {
4697 return unsafe { find_first_ge_u32_sse(slice, target) };
4698 }
4699 }
4700
4701 slice.partition_point(|&d| d < target)
4703}
4704
4705#[cfg(target_arch = "aarch64")]
4706#[target_feature(enable = "neon")]
4707#[allow(unsafe_op_in_unsafe_fn)]
4708unsafe fn find_first_ge_u32_neon(slice: &[u32], target: u32) -> usize {
4709 use std::arch::aarch64::*;
4710
4711 let n = slice.len();
4712 let ptr = slice.as_ptr();
4713 let target_vec = vdupq_n_u32(target);
4714 let bit_mask: uint32x4_t = core::mem::transmute([1u32, 2u32, 4u32, 8u32]);
4716
4717 let chunks = n / 16;
4718 let mut base = 0usize;
4719
4720 for _ in 0..chunks {
4722 let v0 = vld1q_u32(ptr.add(base));
4723 let v1 = vld1q_u32(ptr.add(base + 4));
4724 let v2 = vld1q_u32(ptr.add(base + 8));
4725 let v3 = vld1q_u32(ptr.add(base + 12));
4726
4727 let c0 = vcgeq_u32(v0, target_vec);
4728 let c1 = vcgeq_u32(v1, target_vec);
4729 let c2 = vcgeq_u32(v2, target_vec);
4730 let c3 = vcgeq_u32(v3, target_vec);
4731
4732 let m0 = vaddvq_u32(vandq_u32(c0, bit_mask));
4733 if m0 != 0 {
4734 return base + m0.trailing_zeros() as usize;
4735 }
4736 let m1 = vaddvq_u32(vandq_u32(c1, bit_mask));
4737 if m1 != 0 {
4738 return base + 4 + m1.trailing_zeros() as usize;
4739 }
4740 let m2 = vaddvq_u32(vandq_u32(c2, bit_mask));
4741 if m2 != 0 {
4742 return base + 8 + m2.trailing_zeros() as usize;
4743 }
4744 let m3 = vaddvq_u32(vandq_u32(c3, bit_mask));
4745 if m3 != 0 {
4746 return base + 12 + m3.trailing_zeros() as usize;
4747 }
4748 base += 16;
4749 }
4750
4751 while base + 4 <= n {
4753 let vals = vld1q_u32(ptr.add(base));
4754 let cmp = vcgeq_u32(vals, target_vec);
4755 let mask = vaddvq_u32(vandq_u32(cmp, bit_mask));
4756 if mask != 0 {
4757 return base + mask.trailing_zeros() as usize;
4758 }
4759 base += 4;
4760 }
4761
4762 while base < n {
4764 if *slice.get_unchecked(base) >= target {
4765 return base;
4766 }
4767 base += 1;
4768 }
4769 n
4770}
4771
4772#[cfg(target_arch = "x86_64")]
4773#[target_feature(enable = "sse2")]
4774#[allow(unsafe_op_in_unsafe_fn)]
4775unsafe fn find_first_ge_u32_sse(slice: &[u32], target: u32) -> usize {
4776 use std::arch::x86_64::*;
4777
4778 let n = slice.len();
4779 let ptr = slice.as_ptr();
4780
4781 let sign_flip = _mm_set1_epi32(i32::MIN);
4783 let target_xor = _mm_xor_si128(_mm_set1_epi32(target as i32), sign_flip);
4784
4785 let chunks = n / 16;
4786 let mut base = 0usize;
4787
4788 for _ in 0..chunks {
4790 let v0 = _mm_xor_si128(_mm_loadu_si128(ptr.add(base) as *const __m128i), sign_flip);
4791 let v1 = _mm_xor_si128(
4792 _mm_loadu_si128(ptr.add(base + 4) as *const __m128i),
4793 sign_flip,
4794 );
4795 let v2 = _mm_xor_si128(
4796 _mm_loadu_si128(ptr.add(base + 8) as *const __m128i),
4797 sign_flip,
4798 );
4799 let v3 = _mm_xor_si128(
4800 _mm_loadu_si128(ptr.add(base + 12) as *const __m128i),
4801 sign_flip,
4802 );
4803
4804 let ge0 = _mm_or_si128(
4806 _mm_cmpeq_epi32(v0, target_xor),
4807 _mm_cmpgt_epi32(v0, target_xor),
4808 );
4809 let m0 = _mm_movemask_ps(_mm_castsi128_ps(ge0)) as u32;
4810 if m0 != 0 {
4811 return base + m0.trailing_zeros() as usize;
4812 }
4813
4814 let ge1 = _mm_or_si128(
4815 _mm_cmpeq_epi32(v1, target_xor),
4816 _mm_cmpgt_epi32(v1, target_xor),
4817 );
4818 let m1 = _mm_movemask_ps(_mm_castsi128_ps(ge1)) as u32;
4819 if m1 != 0 {
4820 return base + 4 + m1.trailing_zeros() as usize;
4821 }
4822
4823 let ge2 = _mm_or_si128(
4824 _mm_cmpeq_epi32(v2, target_xor),
4825 _mm_cmpgt_epi32(v2, target_xor),
4826 );
4827 let m2 = _mm_movemask_ps(_mm_castsi128_ps(ge2)) as u32;
4828 if m2 != 0 {
4829 return base + 8 + m2.trailing_zeros() as usize;
4830 }
4831
4832 let ge3 = _mm_or_si128(
4833 _mm_cmpeq_epi32(v3, target_xor),
4834 _mm_cmpgt_epi32(v3, target_xor),
4835 );
4836 let m3 = _mm_movemask_ps(_mm_castsi128_ps(ge3)) as u32;
4837 if m3 != 0 {
4838 return base + 12 + m3.trailing_zeros() as usize;
4839 }
4840 base += 16;
4841 }
4842
4843 while base + 4 <= n {
4845 let vals = _mm_xor_si128(_mm_loadu_si128(ptr.add(base) as *const __m128i), sign_flip);
4846 let ge = _mm_or_si128(
4847 _mm_cmpeq_epi32(vals, target_xor),
4848 _mm_cmpgt_epi32(vals, target_xor),
4849 );
4850 let mask = _mm_movemask_ps(_mm_castsi128_ps(ge)) as u32;
4851 if mask != 0 {
4852 return base + mask.trailing_zeros() as usize;
4853 }
4854 base += 4;
4855 }
4856
4857 while base < n {
4859 if *slice.get_unchecked(base) >= target {
4860 return base;
4861 }
4862 base += 1;
4863 }
4864 n
4865}
4866
4867#[cfg(test)]
4868mod find_first_ge_tests {
4869 use super::find_first_ge_u32;
4870
4871 #[test]
4872 fn test_find_first_ge_basic() {
4873 let data: Vec<u32> = (0..128).map(|i| i * 3).collect(); assert_eq!(find_first_ge_u32(&data, 0), 0);
4875 assert_eq!(find_first_ge_u32(&data, 1), 1); assert_eq!(find_first_ge_u32(&data, 3), 1);
4877 assert_eq!(find_first_ge_u32(&data, 4), 2); assert_eq!(find_first_ge_u32(&data, 381), 127);
4879 assert_eq!(find_first_ge_u32(&data, 382), 128); }
4881
4882 #[test]
4883 fn test_find_first_ge_matches_partition_point() {
4884 let data: Vec<u32> = vec![1, 5, 10, 15, 20, 25, 30, 35, 40, 45, 50, 55, 60, 65, 70, 75];
4885 for target in 0..80 {
4886 let expected = data.partition_point(|&d| d < target);
4887 let actual = find_first_ge_u32(&data, target);
4888 assert_eq!(actual, expected, "target={}", target);
4889 }
4890 }
4891
4892 #[test]
4893 fn test_find_first_ge_small_slices() {
4894 assert_eq!(find_first_ge_u32(&[], 5), 0);
4896 assert_eq!(find_first_ge_u32(&[10], 5), 0);
4898 assert_eq!(find_first_ge_u32(&[10], 10), 0);
4899 assert_eq!(find_first_ge_u32(&[10], 11), 1);
4900 assert_eq!(find_first_ge_u32(&[2, 4, 6], 5), 2);
4902 }
4903
4904 #[test]
4905 fn test_find_first_ge_full_block() {
4906 let data: Vec<u32> = (100..228).collect();
4908 assert_eq!(find_first_ge_u32(&data, 100), 0);
4909 assert_eq!(find_first_ge_u32(&data, 150), 50);
4910 assert_eq!(find_first_ge_u32(&data, 227), 127);
4911 assert_eq!(find_first_ge_u32(&data, 228), 128);
4912 assert_eq!(find_first_ge_u32(&data, 99), 0);
4913 }
4914
4915 #[test]
4916 fn test_find_first_ge_u32_max() {
4917 let data = vec![u32::MAX - 10, u32::MAX - 5, u32::MAX - 1, u32::MAX];
4919 assert_eq!(find_first_ge_u32(&data, u32::MAX - 10), 0);
4920 assert_eq!(find_first_ge_u32(&data, u32::MAX - 7), 1);
4921 assert_eq!(find_first_ge_u32(&data, u32::MAX), 3);
4922 }
4923}