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]
3114pub fn batch_cosine_scores_u8(query: &[f32], vectors_raw: &[u8], dim: usize, scores: &mut [f32]) {
3115 let n = scores.len();
3116 let required = n.checked_mul(dim).expect("u8 batch byte length overflow");
3117 assert_eq!(query.len(), dim, "u8 batch cosine query dimension mismatch");
3118 assert!(
3119 vectors_raw.len() >= required,
3120 "u8 batch cosine vectors are truncated: need {required} bytes, got {}",
3121 vectors_raw.len()
3122 );
3123 if dim == 0 || n == 0 {
3124 return;
3125 }
3126
3127 let norm_q_sq = dot_product_f32(query, query, dim);
3128 if norm_q_sq < f32::EPSILON {
3129 for s in scores.iter_mut() {
3130 *s = 0.0;
3131 }
3132 return;
3133 }
3134 let inv_norm_q = fast_inv_sqrt(norm_q_sq);
3135
3136 for i in 0..n {
3137 let u8_slice = &vectors_raw[i * dim..(i + 1) * dim];
3138
3139 let (dot, norm_v_sq) = fused_dot_norm_u8(query, u8_slice, dim);
3140 scores[i] = if norm_v_sq < f32::EPSILON {
3141 0.0
3142 } else {
3143 dot * inv_norm_q * fast_inv_sqrt(norm_v_sq)
3144 };
3145 }
3146}
3147
3148#[inline]
3157pub fn batch_dot_scores(query: &[f32], vectors: &[f32], dim: usize, scores: &mut [f32]) {
3158 let n = scores.len();
3159 let required = n
3160 .checked_mul(dim)
3161 .expect("batch dot vector length overflow");
3162 assert_eq!(query.len(), dim, "batch dot query dimension mismatch");
3163 assert!(
3164 vectors.len() >= required,
3165 "batch dot vectors are truncated: need {required}, got {}",
3166 vectors.len()
3167 );
3168
3169 if dim == 0 || n == 0 {
3170 return;
3171 }
3172
3173 let norm_q_sq = dot_product_f32(query, query, dim);
3174 if norm_q_sq < f32::EPSILON {
3175 for s in scores.iter_mut() {
3176 *s = 0.0;
3177 }
3178 return;
3179 }
3180 let inv_norm_q = fast_inv_sqrt(norm_q_sq);
3181
3182 for i in 0..n {
3183 let vec = &vectors[i * dim..(i + 1) * dim];
3184 let dot = dot_product_f32(query, vec, dim);
3185 scores[i] = dot * inv_norm_q;
3186 }
3187}
3188
3189#[inline]
3194pub fn batch_dot_scores_f16(query: &[f32], vectors_raw: &[u8], dim: usize, scores: &mut [f32]) {
3195 let n = scores.len();
3196 let vec_bytes = dim.checked_mul(2).expect("f16 vector size overflow");
3197 let required = n
3198 .checked_mul(vec_bytes)
3199 .expect("f16 batch byte length overflow");
3200 assert_eq!(query.len(), dim, "f16 batch dot query dimension mismatch");
3201 assert!(
3202 vectors_raw.len() >= required,
3203 "f16 batch dot vectors are truncated: need {required} bytes, got {}",
3204 vectors_raw.len()
3205 );
3206 if required > 0 {
3207 assert!(
3208 (vectors_raw.as_ptr() as usize).is_multiple_of(std::mem::align_of::<u16>()),
3209 "f16 batch dot vectors are not 2-byte aligned"
3210 );
3211 }
3212 if dim == 0 || n == 0 {
3213 return;
3214 }
3215
3216 let norm_q_sq = dot_product_f32(query, query, dim);
3217 if norm_q_sq < f32::EPSILON {
3218 for s in scores.iter_mut() {
3219 *s = 0.0;
3220 }
3221 return;
3222 }
3223 let inv_norm_q = fast_inv_sqrt(norm_q_sq);
3224
3225 let query_f16: Vec<u16> = query.iter().map(|&v| f32_to_f16(v)).collect();
3226 for i in 0..n {
3227 let raw = &vectors_raw[i * vec_bytes..(i + 1) * vec_bytes];
3228 let f16_slice = unsafe { std::slice::from_raw_parts(raw.as_ptr() as *const u16, dim) };
3229 let dot = dot_product_f16_quant(&query_f16, f16_slice, dim);
3230 scores[i] = dot * inv_norm_q;
3231 }
3232}
3233
3234#[inline]
3239pub fn batch_dot_scores_u8(query: &[f32], vectors_raw: &[u8], dim: usize, scores: &mut [f32]) {
3240 let n = scores.len();
3241 let required = n.checked_mul(dim).expect("u8 batch byte length overflow");
3242 assert_eq!(query.len(), dim, "u8 batch dot query dimension mismatch");
3243 assert!(
3244 vectors_raw.len() >= required,
3245 "u8 batch dot vectors are truncated: need {required} bytes, got {}",
3246 vectors_raw.len()
3247 );
3248 if dim == 0 || n == 0 {
3249 return;
3250 }
3251
3252 let norm_q_sq = dot_product_f32(query, query, dim);
3253 if norm_q_sq < f32::EPSILON {
3254 for s in scores.iter_mut() {
3255 *s = 0.0;
3256 }
3257 return;
3258 }
3259 let inv_norm_q = fast_inv_sqrt(norm_q_sq);
3260
3261 for i in 0..n {
3262 let u8_slice = &vectors_raw[i * dim..(i + 1) * dim];
3263 let dot = dot_product_u8_quant(query, u8_slice, dim);
3264 scores[i] = dot * inv_norm_q;
3265 }
3266}
3267
3268#[inline]
3274pub fn batch_cosine_scores_precomp(
3275 query: &[f32],
3276 vectors: &[f32],
3277 dim: usize,
3278 scores: &mut [f32],
3279 inv_norm_q: f32,
3280) {
3281 let n = scores.len();
3282 let required = n
3283 .checked_mul(dim)
3284 .expect("precomputed cosine vector length overflow");
3285 assert_eq!(
3286 query.len(),
3287 dim,
3288 "precomputed cosine query dimension mismatch"
3289 );
3290 assert!(
3291 vectors.len() >= required,
3292 "precomputed cosine vectors are truncated: need {required}, got {}",
3293 vectors.len()
3294 );
3295 for i in 0..n {
3296 let vec = &vectors[i * dim..(i + 1) * dim];
3297 let (dot, norm_v_sq) = fused_dot_norm(query, vec, dim);
3298 scores[i] = if norm_v_sq < f32::EPSILON {
3299 0.0
3300 } else {
3301 dot * inv_norm_q * fast_inv_sqrt(norm_v_sq)
3302 };
3303 }
3304}
3305
3306#[inline]
3308pub fn batch_cosine_scores_f16_precomp(
3309 query_f16: &[u16],
3310 vectors_raw: &[u8],
3311 dim: usize,
3312 scores: &mut [f32],
3313 inv_norm_q: f32,
3314) {
3315 let n = scores.len();
3316 let vec_bytes = dim.checked_mul(2).expect("f16 vector size overflow");
3317 let required = n
3318 .checked_mul(vec_bytes)
3319 .expect("precomputed f16 cosine batch byte length overflow");
3320 assert_eq!(
3321 query_f16.len(),
3322 dim,
3323 "precomputed f16 cosine query dimension mismatch"
3324 );
3325 assert!(
3326 vectors_raw.len() >= required,
3327 "precomputed f16 cosine vectors are truncated: need {required} bytes, got {}",
3328 vectors_raw.len()
3329 );
3330 if required > 0 {
3331 assert!(
3332 (vectors_raw.as_ptr() as usize).is_multiple_of(std::mem::align_of::<u16>()),
3333 "precomputed f16 cosine vectors are not 2-byte aligned"
3334 );
3335 }
3336 for i in 0..n {
3337 let raw = &vectors_raw[i * vec_bytes..(i + 1) * vec_bytes];
3338 let f16_slice = unsafe { std::slice::from_raw_parts(raw.as_ptr() as *const u16, dim) };
3339 let (dot, norm_v_sq) = fused_dot_norm_f16(query_f16, f16_slice, dim);
3340 scores[i] = if norm_v_sq < f32::EPSILON {
3341 0.0
3342 } else {
3343 dot * inv_norm_q * fast_inv_sqrt(norm_v_sq)
3344 };
3345 }
3346}
3347
3348#[inline]
3350pub fn batch_cosine_scores_u8_precomp(
3351 query: &[f32],
3352 vectors_raw: &[u8],
3353 dim: usize,
3354 scores: &mut [f32],
3355 inv_norm_q: f32,
3356) {
3357 let n = scores.len();
3358 let required = n
3359 .checked_mul(dim)
3360 .expect("precomputed u8 cosine batch byte length overflow");
3361 assert_eq!(
3362 query.len(),
3363 dim,
3364 "precomputed u8 cosine query dimension mismatch"
3365 );
3366 assert!(
3367 vectors_raw.len() >= required,
3368 "precomputed u8 cosine vectors are truncated: need {required} bytes, got {}",
3369 vectors_raw.len()
3370 );
3371 for i in 0..n {
3372 let u8_slice = &vectors_raw[i * dim..(i + 1) * dim];
3373 let (dot, norm_v_sq) = fused_dot_norm_u8(query, u8_slice, dim);
3374 scores[i] = if norm_v_sq < f32::EPSILON {
3375 0.0
3376 } else {
3377 dot * inv_norm_q * fast_inv_sqrt(norm_v_sq)
3378 };
3379 }
3380}
3381
3382#[inline]
3384pub fn batch_dot_scores_precomp(
3385 query: &[f32],
3386 vectors: &[f32],
3387 dim: usize,
3388 scores: &mut [f32],
3389 inv_norm_q: f32,
3390) {
3391 let n = scores.len();
3392 let required = n
3393 .checked_mul(dim)
3394 .expect("precomputed dot vector length overflow");
3395 assert_eq!(query.len(), dim, "precomputed dot query dimension mismatch");
3396 assert!(
3397 vectors.len() >= required,
3398 "precomputed dot vectors are truncated: need {required}, got {}",
3399 vectors.len()
3400 );
3401 for i in 0..n {
3402 let vec = &vectors[i * dim..(i + 1) * dim];
3403 scores[i] = dot_product_f32(query, vec, dim) * inv_norm_q;
3404 }
3405}
3406
3407#[inline]
3409pub fn batch_dot_scores_f16_precomp(
3410 query_f16: &[u16],
3411 vectors_raw: &[u8],
3412 dim: usize,
3413 scores: &mut [f32],
3414 inv_norm_q: f32,
3415) {
3416 let n = scores.len();
3417 let vec_bytes = dim.checked_mul(2).expect("f16 vector size overflow");
3418 let required = n
3419 .checked_mul(vec_bytes)
3420 .expect("precomputed f16 dot batch byte length overflow");
3421 assert_eq!(
3422 query_f16.len(),
3423 dim,
3424 "precomputed f16 dot query dimension mismatch"
3425 );
3426 assert!(
3427 vectors_raw.len() >= required,
3428 "precomputed f16 dot vectors are truncated: need {required} bytes, got {}",
3429 vectors_raw.len()
3430 );
3431 if required > 0 {
3432 assert!(
3433 (vectors_raw.as_ptr() as usize).is_multiple_of(std::mem::align_of::<u16>()),
3434 "precomputed f16 dot vectors are not 2-byte aligned"
3435 );
3436 }
3437 for i in 0..n {
3438 let raw = &vectors_raw[i * vec_bytes..(i + 1) * vec_bytes];
3439 let f16_slice = unsafe { std::slice::from_raw_parts(raw.as_ptr() as *const u16, dim) };
3440 scores[i] = dot_product_f16_quant(query_f16, f16_slice, dim) * inv_norm_q;
3441 }
3442}
3443
3444#[inline]
3446pub fn batch_dot_scores_u8_precomp(
3447 query: &[f32],
3448 vectors_raw: &[u8],
3449 dim: usize,
3450 scores: &mut [f32],
3451 inv_norm_q: f32,
3452) {
3453 let n = scores.len();
3454 let required = n
3455 .checked_mul(dim)
3456 .expect("precomputed u8 dot batch byte length overflow");
3457 assert_eq!(
3458 query.len(),
3459 dim,
3460 "precomputed u8 dot query dimension mismatch"
3461 );
3462 assert!(
3463 vectors_raw.len() >= required,
3464 "precomputed u8 dot vectors are truncated: need {required} bytes, got {}",
3465 vectors_raw.len()
3466 );
3467 for i in 0..n {
3468 let u8_slice = &vectors_raw[i * dim..(i + 1) * dim];
3469 scores[i] = dot_product_u8_quant(query, u8_slice, dim) * inv_norm_q;
3470 }
3471}
3472
3473#[inline]
3478pub fn cosine_similarity(a: &[f32], b: &[f32]) -> f32 {
3479 assert_eq!(a.len(), b.len(), "cosine vector dimension mismatch");
3480 let count = a.len();
3481
3482 if count == 0 {
3483 return 0.0;
3484 }
3485
3486 let dot = dot_product_f32(a, b, count);
3487 let norm_a = dot_product_f32(a, a, count);
3488 let norm_b = dot_product_f32(b, b, count);
3489
3490 let denom = (norm_a * norm_b).sqrt();
3491 if denom < f32::EPSILON {
3492 return 0.0;
3493 }
3494
3495 dot / denom
3496}
3497
3498#[cfg(target_arch = "x86_64")]
3507#[target_feature(enable = "avx512f,avx512vpopcntdq")]
3508#[allow(unsafe_op_in_unsafe_fn)]
3509unsafe fn hamming_distance_avx512(a: &[u8], b: &[u8]) -> u32 {
3510 use std::arch::x86_64::*;
3511
3512 let len = a.len();
3513 let chunks64 = len / 64;
3514 let mut acc = _mm512_setzero_si512();
3515
3516 for c in 0..chunks64 {
3517 let off = c * 64;
3518 let va = _mm512_loadu_si512(a.as_ptr().add(off) as *const __m512i);
3519 let vb = _mm512_loadu_si512(b.as_ptr().add(off) as *const __m512i);
3520 acc = _mm512_add_epi64(acc, _mm512_popcnt_epi64(_mm512_xor_si512(va, vb)));
3521 }
3522
3523 let base = chunks64 * 64;
3524 _mm512_reduce_add_epi64(acc) as u32 + hamming_distance_scalar(&a[base..], &b[base..])
3525}
3526
3527#[cfg(target_arch = "x86_64")]
3529#[target_feature(enable = "avx512f,avx512vpopcntdq")]
3530#[allow(unsafe_op_in_unsafe_fn)]
3531unsafe fn hamming_distance_x4_avx512(query: &[u8], rows: [&[u8]; 4]) -> [u32; 4] {
3532 use std::arch::x86_64::*;
3533
3534 let len = query.len();
3535 let chunks64 = len / 64;
3536 let mut acc = [_mm512_setzero_si512(); 4];
3537
3538 for c in 0..chunks64 {
3539 let off = c * 64;
3540 let vq = _mm512_loadu_si512(query.as_ptr().add(off) as *const __m512i);
3541 for r in 0..4 {
3542 let vr = _mm512_loadu_si512(rows[r].as_ptr().add(off) as *const __m512i);
3543 acc[r] = _mm512_add_epi64(acc[r], _mm512_popcnt_epi64(_mm512_xor_si512(vq, vr)));
3544 }
3545 }
3546
3547 let base = chunks64 * 64;
3548 let tail = &query[base..];
3549 [
3550 _mm512_reduce_add_epi64(acc[0]) as u32 + hamming_distance_scalar(tail, &rows[0][base..]),
3551 _mm512_reduce_add_epi64(acc[1]) as u32 + hamming_distance_scalar(tail, &rows[1][base..]),
3552 _mm512_reduce_add_epi64(acc[2]) as u32 + hamming_distance_scalar(tail, &rows[2][base..]),
3553 _mm512_reduce_add_epi64(acc[3]) as u32 + hamming_distance_scalar(tail, &rows[3][base..]),
3554 ]
3555}
3556
3557#[inline]
3559fn hamming_distance_x4_scalar(query: &[u8], rows: [&[u8]; 4]) -> [u32; 4] {
3560 let len = query.len();
3561 let chunks = len / 8;
3562 let mut total = [0u32; 4];
3563
3564 for i in 0..chunks {
3565 let off = i * 8;
3566 let vq = unsafe { std::ptr::read_unaligned(query.as_ptr().add(off) as *const u64) };
3567 for r in 0..4 {
3568 let vr = unsafe { std::ptr::read_unaligned(rows[r].as_ptr().add(off) as *const u64) };
3569 total[r] += (vq ^ vr).count_ones();
3570 }
3571 }
3572
3573 let base = chunks * 8;
3574 for k in base..len {
3575 let q = query[k];
3576 for r in 0..4 {
3577 total[r] += (q ^ rows[r][k]).count_ones();
3578 }
3579 }
3580
3581 total
3582}
3583
3584const HAMMING_ROWS_PER_KERNEL: usize = 4;
3588
3589#[derive(Clone, Copy, Debug, PartialEq, Eq)]
3596pub enum HammingKernel {
3597 #[cfg(target_arch = "x86_64")]
3598 Avx512,
3599 #[cfg(target_arch = "x86_64")]
3600 Avx2,
3601 #[cfg(target_arch = "aarch64")]
3602 Neon,
3603 Scalar,
3604}
3605
3606impl HammingKernel {
3607 #[inline]
3609 pub fn resolve() -> Self {
3610 #[cfg(target_arch = "x86_64")]
3611 {
3612 if is_x86_feature_detected!("avx512f") && is_x86_feature_detected!("avx512vpopcntdq") {
3613 return Self::Avx512;
3614 }
3615 if avx2::is_available() {
3616 return Self::Avx2;
3617 }
3618 Self::Scalar
3619 }
3620
3621 #[cfg(target_arch = "aarch64")]
3622 {
3623 Self::Neon
3624 }
3625
3626 #[cfg(not(any(target_arch = "aarch64", target_arch = "x86_64")))]
3627 {
3628 Self::Scalar
3629 }
3630 }
3631
3632 #[inline]
3634 pub fn distance(self, a: &[u8], b: &[u8]) -> u32 {
3635 debug_assert_eq!(a.len(), b.len(), "Hamming vector byte length mismatch");
3636 match self {
3637 #[cfg(target_arch = "x86_64")]
3638 Self::Avx512 => unsafe { hamming_distance_avx512(a, b) },
3639 #[cfg(target_arch = "x86_64")]
3640 Self::Avx2 => unsafe { avx2::hamming_distance(a, b) },
3641 #[cfg(target_arch = "aarch64")]
3642 Self::Neon => unsafe { neon::hamming_distance(a, b) },
3643 Self::Scalar => hamming_distance_scalar(a, b),
3644 }
3645 }
3646
3647 pub fn distances(self, query: &[u8], db: &[u8], byte_len: usize, out: &mut [u32]) {
3649 self.score_rows(query, db, byte_len, out, |index| index);
3650 }
3651
3652 pub fn gather_distances(
3657 self,
3658 query: &[u8],
3659 db: &[u8],
3660 byte_len: usize,
3661 ids: &[u32],
3662 out: &mut [u32],
3663 ) {
3664 assert_eq!(
3665 ids.len(),
3666 out.len(),
3667 "Hamming gather needs one output slot per row id"
3668 );
3669 self.score_rows(query, db, byte_len, out, |index| ids[index] as usize);
3670 }
3671
3672 #[inline]
3673 fn score_rows(
3674 self,
3675 query: &[u8],
3676 db: &[u8],
3677 byte_len: usize,
3678 out: &mut [u32],
3679 index_of: impl Fn(usize) -> usize,
3680 ) {
3681 assert_eq!(query.len(), byte_len, "Hamming query byte length mismatch");
3682 if byte_len == 0 || out.is_empty() {
3683 return;
3684 }
3685 let row = |index: usize| -> &[u8] {
3686 let start = index * byte_len;
3687 &db[start..start + byte_len]
3688 };
3689 macro_rules! score_with {
3690 ($one:expr, $four:expr) => {{
3691 let mut i = 0;
3692 while i + HAMMING_ROWS_PER_KERNEL <= out.len() {
3693 let quad = [
3694 row(index_of(i)),
3695 row(index_of(i + 1)),
3696 row(index_of(i + 2)),
3697 row(index_of(i + 3)),
3698 ];
3699 out[i..i + HAMMING_ROWS_PER_KERNEL].copy_from_slice(&$four(query, quad));
3700 i += HAMMING_ROWS_PER_KERNEL;
3701 }
3702 while i < out.len() {
3703 out[i] = $one(query, row(index_of(i)));
3704 i += 1;
3705 }
3706 }};
3707 }
3708 match self {
3709 #[cfg(target_arch = "x86_64")]
3710 Self::Avx512 => score_with!(
3711 |query, row| unsafe { hamming_distance_avx512(query, row) },
3712 |query, rows| unsafe { hamming_distance_x4_avx512(query, rows) }
3713 ),
3714 #[cfg(target_arch = "x86_64")]
3715 Self::Avx2 => score_with!(
3716 |query, row| unsafe { avx2::hamming_distance(query, row) },
3717 |query, rows| unsafe { avx2::hamming_distance_x4(query, rows) }
3718 ),
3719 #[cfg(target_arch = "aarch64")]
3720 Self::Neon => score_with!(
3721 |query, row| unsafe { neon::hamming_distance(query, row) },
3722 |query, rows| unsafe { neon::hamming_distance_x4(query, rows) }
3723 ),
3724 Self::Scalar => score_with!(hamming_distance_scalar, hamming_distance_x4_scalar),
3725 }
3726 }
3727}
3728
3729#[inline]
3736pub fn hamming_distance(a: &[u8], b: &[u8]) -> u32 {
3737 assert_eq!(a.len(), b.len(), "Hamming vector byte length mismatch");
3738 HammingKernel::resolve().distance(a, b)
3739}
3740
3741#[inline]
3744#[allow(dead_code)]
3745fn hamming_distance_scalar(a: &[u8], b: &[u8]) -> u32 {
3746 let len = a.len();
3747 let chunks = len / 8;
3748 let remainder = len % 8;
3749 let mut total = 0u32;
3750
3751 for i in 0..chunks {
3752 let off = i * 8;
3753 let va = unsafe { std::ptr::read_unaligned(a.as_ptr().add(off) as *const u64) };
3754 let vb = unsafe { std::ptr::read_unaligned(b.as_ptr().add(off) as *const u64) };
3755 total += (va ^ vb).count_ones();
3756 }
3757
3758 let base = chunks * 8;
3759 for i in 0..remainder {
3760 total += (a[base + i] ^ b[base + i]).count_ones();
3761 }
3762
3763 total
3764}
3765
3766pub fn batch_hamming_scores(
3772 query: &[u8],
3773 db: &[u8],
3774 byte_len: usize,
3775 dim_bits: usize,
3776 scores: &mut [f32],
3777) {
3778 let n = scores.len();
3779 let required = n
3780 .checked_mul(byte_len)
3781 .expect("Hamming batch byte length overflow");
3782 assert_eq!(query.len(), byte_len, "Hamming query byte length mismatch");
3783 assert!(
3784 db.len() >= required,
3785 "Hamming batch is truncated: need {required} bytes, got {}",
3786 db.len()
3787 );
3788
3789 if byte_len == 0 || n == 0 || dim_bits == 0 {
3790 return;
3791 }
3792
3793 scores_from_hamming(
3794 HammingKernel::resolve(),
3795 query,
3796 db,
3797 byte_len,
3798 dim_bits,
3799 scores,
3800 );
3801}
3802
3803pub fn scores_from_hamming(
3808 kernel: HammingKernel,
3809 query: &[u8],
3810 db: &[u8],
3811 byte_len: usize,
3812 dim_bits: usize,
3813 scores: &mut [f32],
3814) {
3815 if byte_len == 0 || scores.is_empty() || dim_bits == 0 {
3816 return;
3817 }
3818 let inv_dim = 1.0 / dim_bits as f32;
3819 let mut distances = [0u32; HAMMING_DISTANCE_BLOCK];
3822 for (block_index, block) in scores.chunks_mut(HAMMING_DISTANCE_BLOCK).enumerate() {
3823 let rows = &mut distances[..block.len()];
3824 kernel.distances(
3825 query,
3826 &db[block_index * HAMMING_DISTANCE_BLOCK * byte_len..],
3827 byte_len,
3828 rows,
3829 );
3830 for (score, &distance) in block.iter_mut().zip(rows.iter()) {
3831 *score = 1.0 - distance as f32 * inv_dim;
3832 }
3833 }
3834}
3835
3836const HAMMING_DISTANCE_BLOCK: usize = 64;
3838
3839pub fn batch_hamming_distances(query: &[u8], db: &[u8], byte_len: usize, out: &mut [u32]) {
3844 HammingKernel::resolve().distances(query, db, byte_len, out);
3845}
3846
3847#[cfg(test)]
3848mod tests {
3849 use super::*;
3850
3851 #[test]
3852 fn vector_simd_boundaries_reject_dimension_mismatches() {
3853 let vectors = vec![1.0f32; 6];
3854 let raw_f16 = vec![0u8; 12];
3855 let raw_u8 = vec![0u8; 6];
3856 let mut scores = vec![0.0f32; 2];
3857
3858 for invalid_query in [vec![1.0, 2.0], vec![1.0, 2.0, 3.0, 4.0]] {
3859 assert!(
3860 std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
3861 batch_cosine_scores(&invalid_query, &vectors, 3, &mut scores)
3862 }))
3863 .is_err()
3864 );
3865 assert!(
3866 std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
3867 batch_dot_scores_f16(&invalid_query, &raw_f16, 3, &mut scores)
3868 }))
3869 .is_err()
3870 );
3871 assert!(
3872 std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
3873 batch_cosine_scores_u8(&invalid_query, &raw_u8, 3, &mut scores)
3874 }))
3875 .is_err()
3876 );
3877 }
3878 }
3879
3880 #[test]
3881 fn vector_simd_boundaries_reject_truncated_storage() {
3882 let query = [1.0f32, 2.0, 3.0];
3883 let mut scores = [0.0f32; 2];
3884
3885 assert!(
3886 std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
3887 batch_dot_scores(&query, &[0.0; 5], 3, &mut scores)
3888 }))
3889 .is_err()
3890 );
3891 assert!(
3892 std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
3893 batch_cosine_scores_f16(&query, &[0u8; 11], 3, &mut scores)
3894 }))
3895 .is_err()
3896 );
3897 assert!(
3898 std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
3899 dot_product_f32(&query, &query, 4)
3900 }))
3901 .is_err()
3902 );
3903 }
3904
3905 #[test]
3906 fn test_unpack_8bit() {
3907 let input: Vec<u8> = (0..128).collect();
3908 let mut output = vec![0u32; 128];
3909 unpack_8bit(&input, &mut output, 128);
3910
3911 for (i, &v) in output.iter().enumerate() {
3912 assert_eq!(v, i as u32);
3913 }
3914 }
3915
3916 #[test]
3917 fn test_unpack_16bit() {
3918 let mut input = vec![0u8; 256];
3919 for i in 0..128 {
3920 let val = (i * 100) as u16;
3921 input[i * 2] = val as u8;
3922 input[i * 2 + 1] = (val >> 8) as u8;
3923 }
3924
3925 let mut output = vec![0u32; 128];
3926 unpack_16bit(&input, &mut output, 128);
3927
3928 for (i, &v) in output.iter().enumerate() {
3929 assert_eq!(v, (i * 100) as u32);
3930 }
3931 }
3932
3933 #[test]
3934 fn test_unpack_32bit() {
3935 let mut input = vec![0u8; 512];
3936 for i in 0..128 {
3937 let val = (i * 1000) as u32;
3938 let bytes = val.to_le_bytes();
3939 input[i * 4..i * 4 + 4].copy_from_slice(&bytes);
3940 }
3941
3942 let mut output = vec![0u32; 128];
3943 unpack_32bit(&input, &mut output, 128);
3944
3945 for (i, &v) in output.iter().enumerate() {
3946 assert_eq!(v, (i * 1000) as u32);
3947 }
3948 }
3949
3950 #[test]
3951 fn test_delta_decode() {
3952 let deltas = vec![4u32, 4, 9, 19];
3956 let mut output = vec![0u32; 5];
3957
3958 delta_decode(&mut output, &deltas, 10, 5);
3959
3960 assert_eq!(output, vec![10, 15, 20, 30, 50]);
3961 }
3962
3963 #[test]
3964 fn test_add_one() {
3965 let mut values = vec![0u32, 1, 2, 3, 4, 5, 6, 7];
3966 add_one(&mut values, 8);
3967
3968 assert_eq!(values, vec![1, 2, 3, 4, 5, 6, 7, 8]);
3969 }
3970
3971 #[test]
3972 fn test_bits_needed() {
3973 assert_eq!(bits_needed(0), 0);
3974 assert_eq!(bits_needed(1), 1);
3975 assert_eq!(bits_needed(2), 2);
3976 assert_eq!(bits_needed(3), 2);
3977 assert_eq!(bits_needed(4), 3);
3978 assert_eq!(bits_needed(255), 8);
3979 assert_eq!(bits_needed(256), 9);
3980 assert_eq!(bits_needed(u32::MAX), 32);
3981 }
3982
3983 #[test]
3984 fn test_unpack_8bit_delta_decode() {
3985 let input: Vec<u8> = vec![4, 4, 9, 19];
3989 let mut output = vec![0u32; 5];
3990
3991 unpack_8bit_delta_decode(&input, &mut output, 10, 5);
3992
3993 assert_eq!(output, vec![10, 15, 20, 30, 50]);
3994 }
3995
3996 #[test]
3997 fn test_unpack_16bit_delta_decode() {
3998 let mut input = vec![0u8; 8];
4002 for (i, &delta) in [499u16, 499, 999, 1999].iter().enumerate() {
4003 input[i * 2] = delta as u8;
4004 input[i * 2 + 1] = (delta >> 8) as u8;
4005 }
4006 let mut output = vec![0u32; 5];
4007
4008 unpack_16bit_delta_decode(&input, &mut output, 100, 5);
4009
4010 assert_eq!(output, vec![100, 600, 1100, 2100, 4100]);
4011 }
4012
4013 #[test]
4014 fn test_fused_vs_separate_8bit() {
4015 let input: Vec<u8> = (0..127).collect();
4017 let first_value = 1000u32;
4018 let count = 128;
4019
4020 let mut unpacked = vec![0u32; 128];
4022 unpack_8bit(&input, &mut unpacked, 127);
4023 let mut separate_output = vec![0u32; 128];
4024 delta_decode(&mut separate_output, &unpacked, first_value, count);
4025
4026 let mut fused_output = vec![0u32; 128];
4028 unpack_8bit_delta_decode(&input, &mut fused_output, first_value, count);
4029
4030 assert_eq!(separate_output, fused_output);
4031 }
4032
4033 #[test]
4034 fn test_round_bit_width() {
4035 assert_eq!(round_bit_width(0), 0);
4036 assert_eq!(round_bit_width(1), 8);
4037 assert_eq!(round_bit_width(5), 8);
4038 assert_eq!(round_bit_width(8), 8);
4039 assert_eq!(round_bit_width(9), 16);
4040 assert_eq!(round_bit_width(12), 16);
4041 assert_eq!(round_bit_width(16), 16);
4042 assert_eq!(round_bit_width(17), 32);
4043 assert_eq!(round_bit_width(24), 32);
4044 assert_eq!(round_bit_width(32), 32);
4045 }
4046
4047 #[test]
4048 fn test_rounded_bitwidth_from_exact() {
4049 assert_eq!(RoundedBitWidth::from_exact(0), RoundedBitWidth::Zero);
4050 assert_eq!(RoundedBitWidth::from_exact(1), RoundedBitWidth::Bits8);
4051 assert_eq!(RoundedBitWidth::from_exact(8), RoundedBitWidth::Bits8);
4052 assert_eq!(RoundedBitWidth::from_exact(9), RoundedBitWidth::Bits16);
4053 assert_eq!(RoundedBitWidth::from_exact(16), RoundedBitWidth::Bits16);
4054 assert_eq!(RoundedBitWidth::from_exact(17), RoundedBitWidth::Bits32);
4055 assert_eq!(RoundedBitWidth::from_exact(32), RoundedBitWidth::Bits32);
4056 }
4057
4058 #[test]
4059 fn test_pack_unpack_rounded_8bit() {
4060 let values: Vec<u32> = (0..128).map(|i| i % 256).collect();
4061 let mut packed = vec![0u8; 128];
4062
4063 let bytes_written = pack_rounded(&values, RoundedBitWidth::Bits8, &mut packed);
4064 assert_eq!(bytes_written, 128);
4065
4066 let mut unpacked = vec![0u32; 128];
4067 unpack_rounded(&packed, RoundedBitWidth::Bits8, &mut unpacked, 128);
4068
4069 assert_eq!(values, unpacked);
4070 }
4071
4072 #[test]
4073 fn test_pack_unpack_rounded_16bit() {
4074 let values: Vec<u32> = (0..128).map(|i| i * 100).collect();
4075 let mut packed = vec![0u8; 256];
4076
4077 let bytes_written = pack_rounded(&values, RoundedBitWidth::Bits16, &mut packed);
4078 assert_eq!(bytes_written, 256);
4079
4080 let mut unpacked = vec![0u32; 128];
4081 unpack_rounded(&packed, RoundedBitWidth::Bits16, &mut unpacked, 128);
4082
4083 assert_eq!(values, unpacked);
4084 }
4085
4086 #[test]
4087 fn test_pack_unpack_rounded_32bit() {
4088 let values: Vec<u32> = (0..128).map(|i| i * 100000).collect();
4089 let mut packed = vec![0u8; 512];
4090
4091 let bytes_written = pack_rounded(&values, RoundedBitWidth::Bits32, &mut packed);
4092 assert_eq!(bytes_written, 512);
4093
4094 let mut unpacked = vec![0u32; 128];
4095 unpack_rounded(&packed, RoundedBitWidth::Bits32, &mut unpacked, 128);
4096
4097 assert_eq!(values, unpacked);
4098 }
4099
4100 #[test]
4101 fn test_unpack_rounded_delta_decode() {
4102 let input: Vec<u8> = vec![4, 4, 9, 19];
4107 let mut output = vec![0u32; 5];
4108
4109 unpack_rounded_delta_decode(&input, RoundedBitWidth::Bits8, &mut output, 10, 5);
4110
4111 assert_eq!(output, vec![10, 15, 20, 30, 50]);
4112 }
4113
4114 #[test]
4115 fn test_unpack_rounded_delta_decode_zero() {
4116 let input: Vec<u8> = vec![];
4118 let mut output = vec![0u32; 5];
4119
4120 unpack_rounded_delta_decode(&input, RoundedBitWidth::Zero, &mut output, 100, 5);
4121
4122 assert_eq!(output, vec![100, 101, 102, 103, 104]);
4123 }
4124
4125 #[test]
4130 fn test_dequantize_uint8() {
4131 let input: Vec<u8> = vec![0, 128, 255, 64, 192];
4132 let mut output = vec![0.0f32; 5];
4133 let scale = 0.1;
4134 let min_val = 1.0;
4135
4136 dequantize_uint8(&input, &mut output, scale, min_val, 5);
4137
4138 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); }
4145
4146 #[test]
4147 fn test_dequantize_uint8_large() {
4148 let input: Vec<u8> = (0..128).collect();
4150 let mut output = vec![0.0f32; 128];
4151 let scale = 2.0;
4152 let min_val = -10.0;
4153
4154 dequantize_uint8(&input, &mut output, scale, min_val, 128);
4155
4156 for (i, &out) in output.iter().enumerate().take(128) {
4157 let expected = i as f32 * scale + min_val;
4158 assert!(
4159 (out - expected).abs() < 1e-5,
4160 "Mismatch at {}: expected {}, got {}",
4161 i,
4162 expected,
4163 out
4164 );
4165 }
4166 }
4167
4168 #[test]
4169 fn test_dot_product_f32() {
4170 let a = vec![1.0f32, 2.0, 3.0, 4.0, 5.0];
4171 let b = vec![2.0f32, 3.0, 4.0, 5.0, 6.0];
4172
4173 let result = dot_product_f32(&a, &b, 5);
4174
4175 assert!((result - 70.0).abs() < 1e-5);
4177 }
4178
4179 #[test]
4180 fn test_dot_product_f32_large() {
4181 let a: Vec<f32> = (0..128).map(|i| i as f32).collect();
4183 let b: Vec<f32> = (0..128).map(|i| (i + 1) as f32).collect();
4184
4185 let result = dot_product_f32(&a, &b, 128);
4186
4187 let expected: f32 = (0..128).map(|i| (i as f32) * ((i + 1) as f32)).sum();
4189 assert!(
4190 (result - expected).abs() < 1e-3,
4191 "Expected {}, got {}",
4192 expected,
4193 result
4194 );
4195 }
4196
4197 #[test]
4198 fn test_fused_dot_norm() {
4199 let a = vec![1.0f32, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0];
4200 let b = vec![2.0f32, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0];
4201 let (dot, norm_b) = fused_dot_norm(&a, &b, a.len());
4202
4203 let expected_dot: f32 = a.iter().zip(b.iter()).map(|(x, y)| x * y).sum();
4204 let expected_norm: f32 = b.iter().map(|x| x * x).sum();
4205 assert!(
4206 (dot - expected_dot).abs() < 1e-5,
4207 "dot: expected {}, got {}",
4208 expected_dot,
4209 dot
4210 );
4211 assert!(
4212 (norm_b - expected_norm).abs() < 1e-5,
4213 "norm: expected {}, got {}",
4214 expected_norm,
4215 norm_b
4216 );
4217 }
4218
4219 #[test]
4220 fn test_fused_dot_norm_large() {
4221 let a: Vec<f32> = (0..768).map(|i| (i as f32) * 0.01).collect();
4222 let b: Vec<f32> = (0..768).map(|i| (i as f32) * 0.02 + 0.5).collect();
4223 let (dot, norm_b) = fused_dot_norm(&a, &b, a.len());
4224
4225 let expected_dot: f32 = a.iter().zip(b.iter()).map(|(x, y)| x * y).sum();
4226 let expected_norm: f32 = b.iter().map(|x| x * x).sum();
4227 assert!(
4228 (dot - expected_dot).abs() < 1.0,
4229 "dot: expected {}, got {}",
4230 expected_dot,
4231 dot
4232 );
4233 assert!(
4234 (norm_b - expected_norm).abs() < 1.0,
4235 "norm: expected {}, got {}",
4236 expected_norm,
4237 norm_b
4238 );
4239 }
4240
4241 #[test]
4242 fn test_batch_cosine_scores() {
4243 let query = vec![1.0f32, 0.0, 0.0];
4245 let vectors = vec![
4246 1.0, 0.0, 0.0, 0.0, 1.0, 0.0, -1.0, 0.0, 0.0, 0.5, 0.5, 0.0, ];
4251 let mut scores = vec![0f32; 4];
4252 batch_cosine_scores(&query, &vectors, 3, &mut scores);
4253
4254 assert!((scores[0] - 1.0).abs() < 1e-5, "identical: {}", scores[0]);
4255 assert!(scores[1].abs() < 1e-5, "orthogonal: {}", scores[1]);
4256 assert!((scores[2] - (-1.0)).abs() < 1e-5, "opposite: {}", scores[2]);
4257 let expected_45 = 0.5f32 / (0.5f32.powi(2) + 0.5f32.powi(2)).sqrt();
4258 assert!(
4259 (scores[3] - expected_45).abs() < 1e-5,
4260 "45deg: expected {}, got {}",
4261 expected_45,
4262 scores[3]
4263 );
4264 }
4265
4266 #[test]
4267 fn test_batch_cosine_scores_matches_individual() {
4268 let query: Vec<f32> = (0..128).map(|i| (i as f32) * 0.1).collect();
4269 let n = 50;
4270 let dim = 128;
4271 let vectors: Vec<f32> = (0..n * dim).map(|i| ((i * 7 + 3) as f32) * 0.01).collect();
4272
4273 let mut batch_scores = vec![0f32; n];
4274 batch_cosine_scores(&query, &vectors, dim, &mut batch_scores);
4275
4276 for i in 0..n {
4277 let vec_i = &vectors[i * dim..(i + 1) * dim];
4278 let individual = cosine_similarity(&query, vec_i);
4279 assert!(
4280 (batch_scores[i] - individual).abs() < 1e-5,
4281 "vec {}: batch={}, individual={}",
4282 i,
4283 batch_scores[i],
4284 individual
4285 );
4286 }
4287 }
4288
4289 #[test]
4290 fn test_batch_cosine_scores_empty() {
4291 let query = vec![1.0f32, 2.0, 3.0];
4292 let vectors: Vec<f32> = vec![];
4293 let mut scores: Vec<f32> = vec![];
4294 batch_cosine_scores(&query, &vectors, 3, &mut scores);
4295 assert!(scores.is_empty());
4296 }
4297
4298 #[test]
4299 fn test_batch_cosine_scores_zero_query() {
4300 let query = vec![0.0f32, 0.0, 0.0];
4301 let vectors = vec![1.0f32, 2.0, 3.0, 4.0, 5.0, 6.0];
4302 let mut scores = vec![0f32; 2];
4303 batch_cosine_scores(&query, &vectors, 3, &mut scores);
4304 assert_eq!(scores[0], 0.0);
4305 assert_eq!(scores[1], 0.0);
4306 }
4307
4308 #[test]
4313 fn test_f16_roundtrip_normal() {
4314 for &v in &[0.0f32, 1.0, -1.0, 0.5, -0.5, 0.333, 65504.0] {
4315 let h = f32_to_f16(v);
4316 let back = f16_to_f32(h);
4317 let err = (back - v).abs() / v.abs().max(1e-6);
4318 assert!(
4319 err < 0.002,
4320 "f16 roundtrip {v} → {h:#06x} → {back}, rel err {err}"
4321 );
4322 }
4323 }
4324
4325 #[test]
4326 fn test_f16_special() {
4327 assert_eq!(f16_to_f32(f32_to_f16(0.0)), 0.0);
4329 assert_eq!(f32_to_f16(-0.0), 0x8000);
4331 assert!(f16_to_f32(f32_to_f16(f32::INFINITY)).is_infinite());
4333 assert!(f16_to_f32(f32_to_f16(f32::NAN)).is_nan());
4335 }
4336
4337 #[test]
4338 fn test_f16_embedding_range() {
4339 let values: Vec<f32> = (-100..=100).map(|i| i as f32 / 100.0).collect();
4341 for &v in &values {
4342 let back = f16_to_f32(f32_to_f16(v));
4343 assert!((back - v).abs() < 0.001, "f16 error for {v}: got {back}");
4344 }
4345 }
4346
4347 #[test]
4352 fn test_u8_roundtrip() {
4353 assert_eq!(f32_to_u8_saturating(-1.0), 0);
4355 assert_eq!(f32_to_u8_saturating(1.0), 255);
4356 assert_eq!(f32_to_u8_saturating(0.0), 127); assert_eq!(f32_to_u8_saturating(-2.0), 0);
4360 assert_eq!(f32_to_u8_saturating(2.0), 255);
4361 }
4362
4363 #[test]
4364 fn test_u8_dequantize() {
4365 assert!((u8_to_f32(0) - (-1.0)).abs() < 0.01);
4366 assert!((u8_to_f32(255) - 1.0).abs() < 0.01);
4367 assert!((u8_to_f32(127) - 0.0).abs() < 0.01);
4368 }
4369
4370 #[test]
4375 fn test_batch_cosine_scores_f16() {
4376 let query = vec![0.6f32, 0.8, 0.0, 0.0];
4377 let dim = 4;
4378 let vecs_f32 = vec![
4379 0.6f32, 0.8, 0.0, 0.0, 0.0, 0.0, 0.6, 0.8, ];
4382
4383 let mut f16_buf = vec![0u16; 8];
4385 batch_f32_to_f16(&vecs_f32, &mut f16_buf);
4386 let raw: &[u8] =
4387 unsafe { std::slice::from_raw_parts(f16_buf.as_ptr() as *const u8, f16_buf.len() * 2) };
4388
4389 let mut scores = vec![0f32; 2];
4390 batch_cosine_scores_f16(&query, raw, dim, &mut scores);
4391
4392 assert!(
4393 (scores[0] - 1.0).abs() < 0.01,
4394 "identical vectors: {}",
4395 scores[0]
4396 );
4397 assert!(scores[1].abs() < 0.01, "orthogonal vectors: {}", scores[1]);
4398 }
4399
4400 #[test]
4401 fn test_batch_cosine_scores_u8() {
4402 let query = vec![0.6f32, 0.8, 0.0, 0.0];
4403 let dim = 4;
4404 let vecs_f32 = vec![
4405 0.6f32, 0.8, 0.0, 0.0, -0.6, -0.8, 0.0, 0.0, ];
4408
4409 let mut u8_buf = vec![0u8; 8];
4411 batch_f32_to_u8(&vecs_f32, &mut u8_buf);
4412
4413 let mut scores = vec![0f32; 2];
4414 batch_cosine_scores_u8(&query, &u8_buf, dim, &mut scores);
4415
4416 assert!(scores[0] > 0.95, "similar vectors: {}", scores[0]);
4417 assert!(scores[1] < -0.95, "opposite vectors: {}", scores[1]);
4418 }
4419
4420 #[test]
4421 fn test_batch_cosine_scores_f16_large_dim() {
4422 let dim = 768;
4424 let query: Vec<f32> = (0..dim).map(|i| (i as f32 / dim as f32) - 0.5).collect();
4425 let vec2: Vec<f32> = query.iter().map(|x| x * 0.9 + 0.01).collect();
4426
4427 let mut all_vecs = query.clone();
4428 all_vecs.extend_from_slice(&vec2);
4429
4430 let mut f16_buf = vec![0u16; all_vecs.len()];
4431 batch_f32_to_f16(&all_vecs, &mut f16_buf);
4432 let raw: &[u8] =
4433 unsafe { std::slice::from_raw_parts(f16_buf.as_ptr() as *const u8, f16_buf.len() * 2) };
4434
4435 let mut scores = vec![0f32; 2];
4436 batch_cosine_scores_f16(&query, raw, dim, &mut scores);
4437
4438 assert!((scores[0] - 1.0).abs() < 0.01, "self-sim: {}", scores[0]);
4440 assert!(scores[1] > 0.99, "scaled-sim: {}", scores[1]);
4442 }
4443
4444 #[test]
4449 fn test_hamming_distance_identical() {
4450 let a = vec![0xAA; 64];
4451 assert_eq!(hamming_distance(&a, &a), 0);
4452 }
4453
4454 #[test]
4455 fn test_hamming_distance_opposite() {
4456 let a = vec![0xFF; 32];
4457 let b = vec![0x00; 32];
4458 assert_eq!(hamming_distance(&a, &b), 256);
4459 }
4460
4461 #[test]
4462 fn test_hamming_distance_known() {
4463 let a = vec![0xAA];
4465 let b = vec![0x55];
4466 assert_eq!(hamming_distance(&a, &b), 8);
4467
4468 let a = vec![0xFF, 0x00];
4470 let b = vec![0x00, 0x00];
4471 assert_eq!(hamming_distance(&a, &b), 8);
4472 }
4473
4474 #[test]
4475 fn test_hamming_distance_single_bit() {
4476 let a = vec![0x00; 16];
4477 let mut b = vec![0x00; 16];
4478 b[7] = 0x01; assert_eq!(hamming_distance(&a, &b), 1);
4480 }
4481
4482 #[test]
4483 fn test_hamming_distance_empty() {
4484 let a: Vec<u8> = vec![];
4485 assert_eq!(hamming_distance(&a, &a), 0);
4486 }
4487
4488 #[test]
4489 fn test_hamming_distance_remainder_path() {
4490 let a = vec![0xFF; 17];
4492 let b = vec![0x00; 17];
4493 assert_eq!(hamming_distance(&a, &b), 136); let a = vec![0xFF; 33];
4497 let b = vec![0x00; 33];
4498 assert_eq!(hamming_distance(&a, &b), 264); }
4500
4501 #[test]
4502 fn test_hamming_distance_large() {
4503 let a = vec![0xFF; 4096];
4505 let b = vec![0x00; 4096];
4506 assert_eq!(hamming_distance(&a, &b), 32768);
4507 }
4508
4509 #[test]
4510 fn test_hamming_distance_scalar_matches() {
4511 for size in [1, 7, 8, 15, 16, 31, 32, 63, 64, 100, 128, 255, 256] {
4513 let a: Vec<u8> = (0..size).map(|i| (i * 37 + 13) as u8).collect();
4514 let b: Vec<u8> = (0..size).map(|i| (i * 53 + 7) as u8).collect();
4515 let expected = hamming_distance_scalar(&a, &b);
4516 let got = hamming_distance(&a, &b);
4517 assert_eq!(got, expected, "mismatch at size {size}");
4518 }
4519 }
4520
4521 #[test]
4526 fn test_batch_hamming_scores_identical() {
4527 let query = vec![0xAA; 16];
4528 let db = vec![0xAA; 16]; let mut scores = vec![0f32; 1];
4530 batch_hamming_scores(&query, &db, 16, 128, &mut scores);
4531 assert!((scores[0] - 1.0).abs() < 1e-6, "identical: {}", scores[0]);
4532 }
4533
4534 #[test]
4535 fn test_batch_hamming_scores_opposite() {
4536 let query = vec![0xFF; 16];
4537 let db = vec![0x00; 16];
4538 let mut scores = vec![0f32; 1];
4539 batch_hamming_scores(&query, &db, 16, 128, &mut scores);
4540 assert!((scores[0] - 0.0).abs() < 1e-6, "opposite: {}", scores[0]);
4541 }
4542
4543 #[test]
4544 fn test_batch_hamming_scores_multiple() {
4545 let byte_len = 8;
4546 let dim_bits = 64;
4547 let query = vec![0xFF; byte_len];
4548 let mut db = Vec::new();
4549 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];
4554 batch_hamming_scores(&query, &db, byte_len, dim_bits, &mut scores);
4555
4556 assert!((scores[0] - 1.0).abs() < 1e-6, "identical: {}", scores[0]);
4557 assert!((scores[1] - 0.0).abs() < 1e-6, "opposite: {}", scores[1]);
4558 assert!((scores[2] - 0.5).abs() < 1e-6, "half: {}", scores[2]);
4559 }
4560
4561 #[test]
4562 fn test_batch_hamming_scores_empty() {
4563 let query = vec![0xFF; 8];
4564 let db: Vec<u8> = vec![];
4565 let mut scores: Vec<f32> = vec![];
4566 batch_hamming_scores(&query, &db, 8, 64, &mut scores);
4567 assert!(scores.is_empty());
4568 }
4569
4570 #[test]
4571 fn test_batch_hamming_scores_zero_byte_len() {
4572 let query: Vec<u8> = vec![];
4573 let db: Vec<u8> = vec![];
4574 let mut scores = vec![0f32; 1];
4575 batch_hamming_scores(&query, &db, 0, 0, &mut scores);
4576 assert_eq!(scores[0], 0.0);
4578 }
4579
4580 fn hamming_matrix(rows: usize, byte_len: usize) -> (Vec<u8>, Vec<u8>) {
4585 let query: Vec<u8> = (0..byte_len).map(|i| (i * 31 + 5) as u8).collect();
4586 let db: Vec<u8> = (0..rows * byte_len)
4587 .map(|i| (i * 97 + i / byte_len * 11 + 3) as u8)
4588 .collect();
4589 (query, db)
4590 }
4591
4592 #[test]
4595 fn batched_hamming_distances_match_scalar_for_every_row_count() {
4596 let kernel = HammingKernel::resolve();
4597 for byte_len in [1, 7, 8, 15, 16, 31, 32, 33, 63, 64, 65, 128, 320] {
4600 for rows in [1, 2, 3, 4, 5, 7, 8, 9, 64, 70] {
4601 let (query, db) = hamming_matrix(rows, byte_len);
4602 let mut got = vec![0u32; rows];
4603 kernel.distances(&query, &db, byte_len, &mut got);
4604 for (row, &distance) in got.iter().enumerate() {
4605 let expected =
4606 hamming_distance_scalar(&query, &db[row * byte_len..(row + 1) * byte_len]);
4607 assert_eq!(
4608 distance, expected,
4609 "row {row} of {rows} at byte_len {byte_len}"
4610 );
4611 }
4612 }
4613 }
4614 }
4615
4616 #[test]
4617 fn gathered_hamming_distances_follow_row_ids() {
4618 let kernel = HammingKernel::resolve();
4619 let byte_len = 320;
4620 let rows = 37;
4621 let (query, db) = hamming_matrix(rows, byte_len);
4622 let ids: Vec<u32> = [36, 0, 17, 17, 5, 31, 2, 9, 9, 36, 1].into_iter().collect();
4625 let mut got = vec![0u32; ids.len()];
4626 kernel.gather_distances(&query, &db, byte_len, &ids, &mut got);
4627 for (slot, &id) in ids.iter().enumerate() {
4628 let start = id as usize * byte_len;
4629 let expected = hamming_distance_scalar(&query, &db[start..start + byte_len]);
4630 assert_eq!(got[slot], expected, "slot {slot} for row {id}");
4631 }
4632 }
4633
4634 #[test]
4635 fn resolved_kernel_matches_scalar_pairwise() {
4636 let kernel = HammingKernel::resolve();
4637 for byte_len in [1, 8, 32, 64, 65, 320, 4096] {
4638 let (query, db) = hamming_matrix(1, byte_len);
4639 assert_eq!(
4640 kernel.distance(&query, &db),
4641 hamming_distance_scalar(&query, &db),
4642 "byte_len {byte_len}"
4643 );
4644 }
4645 }
4646
4647 #[test]
4648 fn scores_from_hamming_matches_batch_scores_across_blocks() {
4649 let kernel = HammingKernel::resolve();
4650 let byte_len = 320;
4651 let dim_bits = byte_len * 8;
4652 let rows = HAMMING_DISTANCE_BLOCK * 2 + 3;
4654 let (query, db) = hamming_matrix(rows, byte_len);
4655 let mut expected = vec![0f32; rows];
4656 for (row, score) in expected.iter_mut().enumerate() {
4657 let distance =
4658 hamming_distance_scalar(&query, &db[row * byte_len..(row + 1) * byte_len]);
4659 *score = 1.0 - distance as f32 / dim_bits as f32;
4660 }
4661 let mut got = vec![0f32; rows];
4662 scores_from_hamming(kernel, &query, &db, byte_len, dim_bits, &mut got);
4663 for (row, (&got, &want)) in got.iter().zip(expected.iter()).enumerate() {
4664 assert!((got - want).abs() < 1e-6, "row {row}: {got} vs {want}");
4665 }
4666 let mut public = vec![0f32; rows];
4667 batch_hamming_scores(&query, &db, byte_len, dim_bits, &mut public);
4668 assert_eq!(got, public);
4669 }
4670}
4671
4672#[inline]
4685pub fn find_first_ge_u32(slice: &[u32], target: u32) -> usize {
4686 #[cfg(target_arch = "aarch64")]
4687 {
4688 if neon::is_available() {
4689 return unsafe { find_first_ge_u32_neon(slice, target) };
4690 }
4691 }
4692
4693 #[cfg(target_arch = "x86_64")]
4694 {
4695 if sse::is_available() {
4696 return unsafe { find_first_ge_u32_sse(slice, target) };
4697 }
4698 }
4699
4700 slice.partition_point(|&d| d < target)
4702}
4703
4704#[cfg(target_arch = "aarch64")]
4705#[target_feature(enable = "neon")]
4706#[allow(unsafe_op_in_unsafe_fn)]
4707unsafe fn find_first_ge_u32_neon(slice: &[u32], target: u32) -> usize {
4708 use std::arch::aarch64::*;
4709
4710 let n = slice.len();
4711 let ptr = slice.as_ptr();
4712 let target_vec = vdupq_n_u32(target);
4713 let bit_mask: uint32x4_t = core::mem::transmute([1u32, 2u32, 4u32, 8u32]);
4715
4716 let chunks = n / 16;
4717 let mut base = 0usize;
4718
4719 for _ in 0..chunks {
4721 let v0 = vld1q_u32(ptr.add(base));
4722 let v1 = vld1q_u32(ptr.add(base + 4));
4723 let v2 = vld1q_u32(ptr.add(base + 8));
4724 let v3 = vld1q_u32(ptr.add(base + 12));
4725
4726 let c0 = vcgeq_u32(v0, target_vec);
4727 let c1 = vcgeq_u32(v1, target_vec);
4728 let c2 = vcgeq_u32(v2, target_vec);
4729 let c3 = vcgeq_u32(v3, target_vec);
4730
4731 let m0 = vaddvq_u32(vandq_u32(c0, bit_mask));
4732 if m0 != 0 {
4733 return base + m0.trailing_zeros() as usize;
4734 }
4735 let m1 = vaddvq_u32(vandq_u32(c1, bit_mask));
4736 if m1 != 0 {
4737 return base + 4 + m1.trailing_zeros() as usize;
4738 }
4739 let m2 = vaddvq_u32(vandq_u32(c2, bit_mask));
4740 if m2 != 0 {
4741 return base + 8 + m2.trailing_zeros() as usize;
4742 }
4743 let m3 = vaddvq_u32(vandq_u32(c3, bit_mask));
4744 if m3 != 0 {
4745 return base + 12 + m3.trailing_zeros() as usize;
4746 }
4747 base += 16;
4748 }
4749
4750 while base + 4 <= n {
4752 let vals = vld1q_u32(ptr.add(base));
4753 let cmp = vcgeq_u32(vals, target_vec);
4754 let mask = vaddvq_u32(vandq_u32(cmp, bit_mask));
4755 if mask != 0 {
4756 return base + mask.trailing_zeros() as usize;
4757 }
4758 base += 4;
4759 }
4760
4761 while base < n {
4763 if *slice.get_unchecked(base) >= target {
4764 return base;
4765 }
4766 base += 1;
4767 }
4768 n
4769}
4770
4771#[cfg(target_arch = "x86_64")]
4772#[target_feature(enable = "sse2")]
4773#[allow(unsafe_op_in_unsafe_fn)]
4774unsafe fn find_first_ge_u32_sse(slice: &[u32], target: u32) -> usize {
4775 use std::arch::x86_64::*;
4776
4777 let n = slice.len();
4778 let ptr = slice.as_ptr();
4779
4780 let sign_flip = _mm_set1_epi32(i32::MIN);
4782 let target_xor = _mm_xor_si128(_mm_set1_epi32(target as i32), sign_flip);
4783
4784 let chunks = n / 16;
4785 let mut base = 0usize;
4786
4787 for _ in 0..chunks {
4789 let v0 = _mm_xor_si128(_mm_loadu_si128(ptr.add(base) as *const __m128i), sign_flip);
4790 let v1 = _mm_xor_si128(
4791 _mm_loadu_si128(ptr.add(base + 4) as *const __m128i),
4792 sign_flip,
4793 );
4794 let v2 = _mm_xor_si128(
4795 _mm_loadu_si128(ptr.add(base + 8) as *const __m128i),
4796 sign_flip,
4797 );
4798 let v3 = _mm_xor_si128(
4799 _mm_loadu_si128(ptr.add(base + 12) as *const __m128i),
4800 sign_flip,
4801 );
4802
4803 let ge0 = _mm_or_si128(
4805 _mm_cmpeq_epi32(v0, target_xor),
4806 _mm_cmpgt_epi32(v0, target_xor),
4807 );
4808 let m0 = _mm_movemask_ps(_mm_castsi128_ps(ge0)) as u32;
4809 if m0 != 0 {
4810 return base + m0.trailing_zeros() as usize;
4811 }
4812
4813 let ge1 = _mm_or_si128(
4814 _mm_cmpeq_epi32(v1, target_xor),
4815 _mm_cmpgt_epi32(v1, target_xor),
4816 );
4817 let m1 = _mm_movemask_ps(_mm_castsi128_ps(ge1)) as u32;
4818 if m1 != 0 {
4819 return base + 4 + m1.trailing_zeros() as usize;
4820 }
4821
4822 let ge2 = _mm_or_si128(
4823 _mm_cmpeq_epi32(v2, target_xor),
4824 _mm_cmpgt_epi32(v2, target_xor),
4825 );
4826 let m2 = _mm_movemask_ps(_mm_castsi128_ps(ge2)) as u32;
4827 if m2 != 0 {
4828 return base + 8 + m2.trailing_zeros() as usize;
4829 }
4830
4831 let ge3 = _mm_or_si128(
4832 _mm_cmpeq_epi32(v3, target_xor),
4833 _mm_cmpgt_epi32(v3, target_xor),
4834 );
4835 let m3 = _mm_movemask_ps(_mm_castsi128_ps(ge3)) as u32;
4836 if m3 != 0 {
4837 return base + 12 + m3.trailing_zeros() as usize;
4838 }
4839 base += 16;
4840 }
4841
4842 while base + 4 <= n {
4844 let vals = _mm_xor_si128(_mm_loadu_si128(ptr.add(base) as *const __m128i), sign_flip);
4845 let ge = _mm_or_si128(
4846 _mm_cmpeq_epi32(vals, target_xor),
4847 _mm_cmpgt_epi32(vals, target_xor),
4848 );
4849 let mask = _mm_movemask_ps(_mm_castsi128_ps(ge)) as u32;
4850 if mask != 0 {
4851 return base + mask.trailing_zeros() as usize;
4852 }
4853 base += 4;
4854 }
4855
4856 while base < n {
4858 if *slice.get_unchecked(base) >= target {
4859 return base;
4860 }
4861 base += 1;
4862 }
4863 n
4864}
4865
4866#[cfg(test)]
4867mod find_first_ge_tests {
4868 use super::find_first_ge_u32;
4869
4870 #[test]
4871 fn test_find_first_ge_basic() {
4872 let data: Vec<u32> = (0..128).map(|i| i * 3).collect(); assert_eq!(find_first_ge_u32(&data, 0), 0);
4874 assert_eq!(find_first_ge_u32(&data, 1), 1); assert_eq!(find_first_ge_u32(&data, 3), 1);
4876 assert_eq!(find_first_ge_u32(&data, 4), 2); assert_eq!(find_first_ge_u32(&data, 381), 127);
4878 assert_eq!(find_first_ge_u32(&data, 382), 128); }
4880
4881 #[test]
4882 fn test_find_first_ge_matches_partition_point() {
4883 let data: Vec<u32> = vec![1, 5, 10, 15, 20, 25, 30, 35, 40, 45, 50, 55, 60, 65, 70, 75];
4884 for target in 0..80 {
4885 let expected = data.partition_point(|&d| d < target);
4886 let actual = find_first_ge_u32(&data, target);
4887 assert_eq!(actual, expected, "target={}", target);
4888 }
4889 }
4890
4891 #[test]
4892 fn test_find_first_ge_small_slices() {
4893 assert_eq!(find_first_ge_u32(&[], 5), 0);
4895 assert_eq!(find_first_ge_u32(&[10], 5), 0);
4897 assert_eq!(find_first_ge_u32(&[10], 10), 0);
4898 assert_eq!(find_first_ge_u32(&[10], 11), 1);
4899 assert_eq!(find_first_ge_u32(&[2, 4, 6], 5), 2);
4901 }
4902
4903 #[test]
4904 fn test_find_first_ge_full_block() {
4905 let data: Vec<u32> = (100..228).collect();
4907 assert_eq!(find_first_ge_u32(&data, 100), 0);
4908 assert_eq!(find_first_ge_u32(&data, 150), 50);
4909 assert_eq!(find_first_ge_u32(&data, 227), 127);
4910 assert_eq!(find_first_ge_u32(&data, 228), 128);
4911 assert_eq!(find_first_ge_u32(&data, 99), 0);
4912 }
4913
4914 #[test]
4915 fn test_find_first_ge_u32_max() {
4916 let data = vec![u32::MAX - 10, u32::MAX - 5, u32::MAX - 1, u32::MAX];
4918 assert_eq!(find_first_ge_u32(&data, u32::MAX - 10), 0);
4919 assert_eq!(find_first_ge_u32(&data, u32::MAX - 7), 1);
4920 assert_eq!(find_first_ge_u32(&data, u32::MAX), 3);
4921 }
4922}